Compare commits
396
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7efbb04709 | ||
|
|
d06fdd78cc | ||
|
|
7a9ea70580 | ||
|
|
64fe5e81f2 | ||
|
|
d29f4703e9 | ||
|
|
c41ad9bd75 | ||
|
|
af8cf142d7 | ||
|
|
dfac776dd5 | ||
|
|
31218127e0 | ||
|
|
57c18bc298 | ||
|
|
6a0b85f108 | ||
|
|
1ff09ccf9c | ||
|
|
ff4dc116c3 | ||
|
|
a6c875156b | ||
|
|
029e5d47d6 | ||
|
|
489d55fa62 | ||
|
|
47a7d43809 | ||
|
|
21928d7c95 | ||
|
|
552feb24b2 | ||
|
|
3cc3713829 | ||
|
|
03a1fd08c0 | ||
|
|
0db0d3ec93 | ||
|
|
fa69b5b73b | ||
|
|
dbf3b9d9bc | ||
|
|
1942cf2cd8 | ||
|
|
441cf2272e | ||
|
|
4dc8348348 | ||
|
|
1bb7e03429 | ||
|
|
7b161740d0 | ||
|
|
6428a83713 | ||
|
|
705757f362 | ||
|
|
dd6766c37d | ||
|
|
8ca1f20f93 | ||
|
|
a23187a456 | ||
|
|
18650712fa | ||
|
|
bd853be356 | ||
|
|
17d0a0c068 | ||
|
|
be6caf9dcf | ||
|
|
91ee2537e2 | ||
|
|
c1fae2ae1b | ||
|
|
434dbee64c | ||
|
|
c3dfaf3dd9 | ||
|
|
234f5a0621 | ||
|
|
d7c32f9633 | ||
|
|
dcef72480b | ||
|
|
3a65d5ed29 | ||
|
|
1265563fc8 | ||
|
|
036ba19b65 | ||
|
|
6c8a92c318 | ||
|
|
912f8e22d1 | ||
|
|
5126e0bb08 | ||
|
|
41d71a9818 | ||
|
|
9fe339dfb1 | ||
|
|
7d1aab8b8d | ||
|
|
2ea5df3db0 | ||
|
|
820019f3c6 | ||
|
|
0679d38e8e | ||
|
|
0d2dbe756d | ||
|
|
ea460c062d | ||
|
|
6ba98c040e | ||
|
|
b82fb603c2 | ||
|
|
c1fadc008a | ||
|
|
4754efc419 | ||
|
|
a404071d3b | ||
|
|
e07ca1bd0a | ||
|
|
d61e8ffff6 | ||
|
|
c20e08fecc | ||
|
|
2142d43f6f | ||
|
|
0c38d46ee9 | ||
|
|
375ec091f3 | ||
|
|
eb5b29f067 | ||
|
|
000fb9ddbe | ||
|
|
dfa02c8b61 | ||
|
|
d28b852732 | ||
|
|
6e03dd2d4c | ||
|
|
b11e053323 | ||
|
|
11154d48b7 | ||
|
|
a483682da6 | ||
|
|
8a7a70b828 | ||
|
|
dddd571a99 | ||
|
|
8bd9ce194b | ||
|
|
0e0fd14ed4 | ||
|
|
946a80bfb8 | ||
|
|
07af5b4a37 | ||
|
|
59b008a461 | ||
|
|
188f8fcb32 | ||
|
|
d7059d2840 | ||
|
|
5adaf60a9f | ||
|
|
3fb33b9873 | ||
|
|
ca97f947d9 | ||
|
|
347b9c23c4 | ||
|
|
ebe438b25b | ||
|
|
c05ca6529f | ||
|
|
6a6d4f2261 | ||
|
|
cf7aece729 | ||
|
|
36e94454e6 | ||
|
|
5bfef3caa4 | ||
|
|
c1908a0205 | ||
|
|
54354e4735 | ||
|
|
4df4b398f5 | ||
|
|
81aa4979b3 | ||
|
|
0108cd7019 | ||
|
|
aad0af9756 | ||
|
|
bd49f6a7c7 | ||
|
|
263eba1342 | ||
|
|
418524051d | ||
|
|
327aa05e80 | ||
|
|
9bf0023163 | ||
|
|
d64921221e | ||
|
|
7bb3d1925b | ||
|
|
eca0fd5a29 | ||
|
|
44398633bb | ||
|
|
52b893b93b | ||
|
|
52c216c1fb | ||
|
|
d9bc612a3c | ||
|
|
1fe43ec3fb | ||
|
|
ca87e29891 | ||
|
|
d2b14e51ee | ||
|
|
9676fc1699 | ||
|
|
73a5b576c9 | ||
|
|
685e50b9af | ||
|
|
c27fafb298 | ||
|
|
37fa0adf93 | ||
|
|
b95055ba28 | ||
|
|
4b78761d20 | ||
|
|
1549987015 | ||
|
|
b0e5d103d8 | ||
|
|
5ab6bdc9b6 | ||
|
|
ed1083803b | ||
|
|
ec3b415c42 | ||
|
|
3a188971b9 | ||
|
|
b3e32ce5b6 | ||
|
|
d0bbfa1015 | ||
|
|
043a68fa44 | ||
|
|
7542035e44 | ||
|
|
bee6f5d114 | ||
|
|
e20b1815fc | ||
|
|
395823f7b6 | ||
|
|
a910fd8a0b | ||
|
|
64ee029a18 | ||
|
|
25df91ca53 | ||
|
|
8987fb8267 | ||
|
|
86ff344c93 | ||
|
|
b6c7b2a2e9 | ||
|
|
9efce6a470 | ||
|
|
6cc7484b78 | ||
|
|
920afcce20 | ||
|
|
6d82c8a554 | ||
|
|
4b52b10e2e | ||
|
|
ac06df1ade | ||
|
|
5f1a867650 | ||
|
|
d284119246 | ||
|
|
84b9aa56ce | ||
|
|
f58ecee0b9 | ||
|
|
b9d16fe86f | ||
|
|
2ccde7a5cd | ||
|
|
d2ca26afaf | ||
|
|
b52feb305e | ||
|
|
558b2f8b67 | ||
|
|
383f0caa16 | ||
|
|
16e8c4e216 | ||
|
|
c6a1b507ba | ||
|
|
cb4fe70b63 | ||
|
|
408d7c34c8 | ||
|
|
7f2df54e01 | ||
|
|
1758d8510b | ||
|
|
7c3a5619f9 | ||
|
|
2b5b47f97d | ||
|
|
b2508bf2d6 | ||
|
|
b0fb1111ec | ||
|
|
4bf126bf52 | ||
|
|
e4449326e6 | ||
|
|
7bab4db28d | ||
|
|
e29ee58603 | ||
|
|
82e67315a4 | ||
|
|
2e82edee14 | ||
|
|
0338ab4d73 | ||
|
|
f4106fff55 | ||
|
|
d4f700ed22 | ||
|
|
041f0f13f4 | ||
|
|
9dee1c594b | ||
|
|
a1ebb2d9c6 | ||
|
|
06ddf041e5 | ||
|
|
d251fcb7d2 | ||
|
|
e839c65537 | ||
|
|
0cde79b831 | ||
|
|
8767a518db | ||
|
|
10ed288d80 | ||
|
|
7143684a81 | ||
|
|
11da432db2 | ||
|
|
73d3231bbd | ||
|
|
59d825dfca | ||
|
|
ae7099fd02 | ||
|
|
56db6d7cf6 | ||
|
|
0eb52762ae | ||
|
|
e93c560288 | ||
|
|
1f213d00f7 | ||
|
|
8f51f99dde | ||
|
|
5cc1482a72 | ||
|
|
639d84ad32 | ||
|
|
1c74f795a6 | ||
|
|
b4f9fbe1b5 | ||
|
|
b992ba996d | ||
|
|
f00d3c7f66 | ||
|
|
e97b615547 | ||
|
|
29cc1d7fdc | ||
|
|
016b5f0363 | ||
|
|
dd7e252452 | ||
|
|
fda1a77f70 | ||
|
|
0accef8e98 | ||
|
|
c77e2368de | ||
|
|
38ef0247c2 | ||
|
|
a158b819f3 | ||
|
|
381963c28a | ||
|
|
ba158c9cdb | ||
|
|
767a2c0061 | ||
|
|
36334f27a1 | ||
|
|
6a479dddb9 | ||
|
|
7058d1aad7 | ||
|
|
265192e509 | ||
|
|
8f2cee4568 | ||
|
|
92f433c904 | ||
|
|
f8ce15b9bf | ||
|
|
d68f618969 | ||
|
|
33e9db64a1 | ||
|
|
0c8699dc20 | ||
|
|
251c451fc3 | ||
|
|
8ed49e4387 | ||
|
|
82afa76182 | ||
|
|
4c307ce4e7 | ||
|
|
252b013243 | ||
|
|
a6b0f82124 | ||
|
|
a27754fb15 | ||
|
|
40fe7aac86 | ||
|
|
7393400f34 | ||
|
|
3b7d18d474 | ||
|
|
6e18858e32 | ||
|
|
84e67efbf4 | ||
|
|
ef2e8ab7ff | ||
|
|
cc45e16904 | ||
|
|
82b01ace5e | ||
|
|
962140eef6 | ||
|
|
5e73d5ff62 | ||
|
|
2c47b8b0d5 | ||
|
|
00968a1ce4 | ||
|
|
b7080a16cf | ||
|
|
c0aed313f4 | ||
|
|
2c53629420 | ||
|
|
760bfc7447 | ||
|
|
cce0a2cb39 | ||
|
|
1c1cf4ce56 | ||
|
|
a99a1ebf9b | ||
|
|
12a6739fc9 | ||
|
|
d8ee10a78d | ||
|
|
ab01144b26 | ||
|
|
072b3278ba | ||
|
|
74e82a3ea8 | ||
|
|
a5b752a983 | ||
|
|
7b878f89a7 | ||
|
|
1b92c8230f | ||
|
|
0178d91333 | ||
|
|
017b8d7271 | ||
|
|
21176f8ee8 | ||
|
|
74bdfc9475 | ||
|
|
b0038e9855 | ||
|
|
6eb85570af | ||
|
|
85599f3ef5 | ||
|
|
dd83bffeef | ||
|
|
911d27fc5f | ||
|
|
a0af096081 | ||
|
|
fcb2c958e7 | ||
|
|
758f346d30 | ||
|
|
78d32cd16c | ||
|
|
fb475cc5bc | ||
|
|
1621e5d261 | ||
|
|
91e095afa9 | ||
|
|
815d99f5ba | ||
|
|
2452f72e75 | ||
|
|
dae18b1faf | ||
|
|
6a10b6241d | ||
|
|
47992d843b | ||
|
|
93100ed314 | ||
|
|
01eda51880 | ||
|
|
e63d028a5a | ||
|
|
58b5677617 | ||
|
|
b6608076ff | ||
|
|
4fe477eaa3 | ||
|
|
a7d1f26f98 | ||
|
|
0673131a80 | ||
|
|
6e02a0829f | ||
|
|
bc813692c6 | ||
|
|
701de3293d | ||
|
|
9dafadc7eb | ||
|
|
d0b77f5bee | ||
|
|
7194f98b19 | ||
|
|
34ba3c676e | ||
|
|
0379b4c823 | ||
|
|
422e0fd809 | ||
|
|
91bf32842e | ||
|
|
e4afa5a61b | ||
|
|
680305aea4 | ||
|
|
9e06237e40 | ||
|
|
8e66c397a9 | ||
|
|
f7c7a62e5f | ||
|
|
c056edaa90 | ||
|
|
63a92bef5f | ||
|
|
1533c0915d | ||
|
|
8eb2937cdb | ||
|
|
d9a372a92e | ||
|
|
199ae146ab | ||
|
|
b4874672fa | ||
|
|
f8d277697d | ||
|
|
4db8a12362 | ||
|
|
cabcb3bb0b | ||
|
|
5f0b715517 | ||
|
|
0ba613c3ce | ||
|
|
04703d2153 | ||
|
|
20da6d7609 | ||
|
|
387c09e91e | ||
|
|
cbce937042 | ||
|
|
246803bcfe | ||
|
|
625c331e80 | ||
|
|
f21944d789 | ||
|
|
0672fba279 | ||
|
|
53a52afe8b | ||
|
|
a2166ee4ff | ||
|
|
26bfd2ece4 | ||
|
|
0d60f0c638 | ||
|
|
735172f806 | ||
|
|
2c2a20b290 | ||
|
|
1a09a9cccd | ||
|
|
f183b09b93 | ||
|
|
5543992d7d | ||
|
|
dcabd76911 | ||
|
|
1a51184a32 | ||
|
|
c7e5095a86 | ||
|
|
f187d32351 | ||
|
|
ee81c65e4b | ||
|
|
ccd3eb24c9 | ||
|
|
51cb32896f | ||
|
|
af42382983 | ||
|
|
955b0c523c | ||
|
|
80281e2543 | ||
|
|
bb4dd4f393 | ||
|
|
ba5ddd59af | ||
|
|
aab7032071 | ||
|
|
adb6dcd683 | ||
|
|
ae93c93182 | ||
|
|
1c0c53c062 | ||
|
|
731add1fdf | ||
|
|
a7f82453a3 | ||
|
|
2bd6be8d85 | ||
|
|
e1014cc790 | ||
|
|
da2125cf13 | ||
|
|
f4bac2d41d | ||
|
|
39abf0ad3f | ||
|
|
8426b0c359 | ||
|
|
f4a0a31f70 | ||
|
|
65862c4fef | ||
|
|
ef548833fd | ||
|
|
55f70e1d27 | ||
|
|
b8cfddd7b6 | ||
|
|
52cb9a2bae | ||
|
|
539101af38 | ||
|
|
efa37cb15f | ||
|
|
faaa97d4a0 | ||
|
|
70804fe2c6 | ||
|
|
ae2532b165 | ||
|
|
81865bf873 | ||
|
|
9e47759347 | ||
|
|
d8665d7ab0 | ||
|
|
9bde15331e | ||
|
|
2fb2de1aa8 | ||
|
|
d68bd07423 | ||
|
|
44972d3215 | ||
|
|
aa308ad201 | ||
|
|
cb73790c27 | ||
|
|
b1fe23fbe4 | ||
|
|
ce81217381 | ||
|
|
94619ce52b | ||
|
|
27aa6bbf46 | ||
|
|
acf4d5c860 | ||
|
|
551932991d | ||
|
|
5ee53c512f | ||
|
|
4efa204727 | ||
|
|
c3bb647640 | ||
|
|
12851bc7ee | ||
|
|
a5f4d30ea6 | ||
|
|
0c9bc765ce | ||
|
|
cd34efa596 | ||
|
|
2b521c3a09 | ||
|
|
9681d96195 | ||
|
|
a32ecfeb33 | ||
|
|
bf73a1dfbe | ||
|
|
ca2ce5c16d | ||
|
|
7b17da7a0c |
@@ -1,6 +1,7 @@
|
||||
{
|
||||
"$schema": "https://anthropic.com/claude-code/marketplace.schema.json",
|
||||
"name": "hindsight",
|
||||
"version": "0.7.5",
|
||||
"description": "Official Hindsight integrations for Claude Code",
|
||||
"owner": {
|
||||
"name": "vectorize-io"
|
||||
@@ -10,6 +11,11 @@
|
||||
"name": "hindsight-memory",
|
||||
"description": "Automatic long-term memory for Claude Code via Hindsight",
|
||||
"source": "./hindsight-integrations/claude-code"
|
||||
},
|
||||
{
|
||||
"name": "hindsight-zcode",
|
||||
"description": "No-MCP long-term memory for ZCode via Hindsight hooks",
|
||||
"source": "./hindsight-integrations/zcode"
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
@@ -78,6 +78,11 @@ results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
- **Authentication/tenancy is enforced inside each engine method, not assumed by the handler.** Every engine method that touches bank-scoped data must authenticate via `request_context` — typically `await self._authenticate_tenant(request_context)` (often indirectly through `get_bank_profile(...)`) — so the correct tenant schema is resolved before any query runs. Handlers must thread `request_context` through to the engine method; never query a tenant-scoped table assuming the schema is already set.
|
||||
- Engine methods return typed models (Pydantic/dataclass), not raw dicts (see Type Safety).
|
||||
|
||||
### Database Locking
|
||||
- **Never use PostgreSQL advisory locks** (`pg_advisory_lock`, `pg_try_advisory_lock`, `pg_advisory_xact_lock`, `pg_advisory_unlock`, …) in migrations, engine code, or anything else. Hindsight runs against connection poolers and managed/PG-compatible services where advisory locks are unreliable or unsupported: session-level locks silently leak or vanish when a pooler hands the session to another client, and callers can block forever on a lock the server never grants. Reject any new occurrence, including ones that look "safe" because they are transaction-scoped.
|
||||
- The pre-existing usage in `hindsight_api/migrations.py` is grandfathered, not a precedent — it is tracked for removal. Don't copy it.
|
||||
- Design the concurrency out instead of locking around it: give each process its own object to write (e.g. per-schema DDL rather than a shared `public.` object), make the operation idempotent, or use a real row/table constraint (`INSERT ... ON CONFLICT`, `SELECT ... FOR UPDATE` in a fixed order). See #2690 for a migration that reached for `pg_advisory_xact_lock` and had to be reverted.
|
||||
|
||||
### Branch Hygiene
|
||||
- **Always start new feature branches from `origin/main`** — rebase to ensure a clean base.
|
||||
- **Only include commits relevant to the PR/branch/feature** — no unrelated changes. If the branch contains commits that don't belong, they must be removed before merging.
|
||||
@@ -204,6 +209,14 @@ in `hindsight-api-slim/hindsight_api/config.py`):
|
||||
The `test_bundled_template_matches_repo_root` sync test fails on drift; if the
|
||||
root file changed without re-copying, flag it as a **must fix**.
|
||||
|
||||
### 11c. Check for advisory locks
|
||||
|
||||
Grep the diff for `advisory` (`git diff main...HEAD | grep -in advisory`). Any new
|
||||
`pg_advisory_lock` / `pg_try_advisory_lock` / `pg_advisory_xact_lock` /
|
||||
`pg_advisory_unlock` call is a **must fix** — see Database Locking above. Point the
|
||||
author at the alternatives (per-process objects, idempotent DDL, row-level
|
||||
constraints) rather than just asking them to drop the lock.
|
||||
|
||||
### 12. Review against other coding standards
|
||||
|
||||
Check the diff for violations of the standards listed above:
|
||||
|
||||
+99
-1
@@ -2,7 +2,7 @@
|
||||
# Copy this file to .env and fill in your values
|
||||
|
||||
# LLM Configuration (Required)
|
||||
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio, vertexai, minimax, deepseek, zai, volcano
|
||||
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio, vertexai, minimax, deepseek, zai, atlas, volcano
|
||||
HINDSIGHT_API_LLM_PROVIDER=openai
|
||||
HINDSIGHT_API_LLM_API_KEY=your-api-key-here
|
||||
HINDSIGHT_API_LLM_MODEL=gpt-4o-mini
|
||||
@@ -10,6 +10,32 @@ HINDSIGHT_API_LLM_BASE_URL=https://api.openai.com/v1
|
||||
# Reasoning effort for providers/models that support it. Examples: low, medium, high, xhigh.
|
||||
# HINDSIGHT_API_LLM_REASONING_EFFORT=low
|
||||
|
||||
# Sampling temperature for internal LLM calls. Set a number in [0.0, 2.0], or `none`
|
||||
# to omit the temperature parameter entirely (required for models that reject explicit
|
||||
# temperatures, e.g. Azure gpt-5.5). The global override below applies to every operation;
|
||||
# per-operation overrides (defaults: verification=0.0, retain=0.1, reflect=0.9,
|
||||
# consolidation=0.0) take precedence.
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE=none
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE_VERIFICATION=0.0
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE_RETAIN=0.1
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE_REFLECT=0.9
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE_CONSOLIDATION=0.0
|
||||
|
||||
# Grammar-enforce structured output (json_schema strict) instead of the soft
|
||||
# schema-in-prompt path. Helps weaker self-hosted models that emit prose preambles
|
||||
# or invalid JSON. The global override below applies to every operation;
|
||||
# per-operation overrides take precedence, in both directions -- set one to false
|
||||
# to opt that operation out while the global flag is on.
|
||||
# HINDSIGHT_API_LLM_STRICT_SCHEMA=false
|
||||
# HINDSIGHT_API_LLM_STRICT_SCHEMA_RETAIN=true
|
||||
# HINDSIGHT_API_LLM_STRICT_SCHEMA_REFLECT=true
|
||||
# HINDSIGHT_API_LLM_STRICT_SCHEMA_CONSOLIDATION=true
|
||||
|
||||
# Diagnostic: on any LLM 4xx, log the exact assembled request ([LLM_4XX_DUMP]) --
|
||||
# serialized request config (message bodies stripped) + capped per-message previews.
|
||||
# For debugging otherwise-unreproducible rejected calls. Off by default.
|
||||
# HINDSIGHT_API_LLM_DEBUG_DUMP_4XX=false
|
||||
|
||||
# Example: Anthropic Claude configuration
|
||||
# HINDSIGHT_API_LLM_PROVIDER=anthropic
|
||||
# HINDSIGHT_API_LLM_API_KEY=your-anthropic-api-key
|
||||
@@ -37,12 +63,39 @@ HINDSIGHT_API_LLM_BASE_URL=https://api.openai.com/v1
|
||||
# HINDSIGHT_API_LLM_API_KEY=your-zai-api-key
|
||||
# HINDSIGHT_API_LLM_MODEL=glm-4.5-flash # or glm-4.5-air for the paid tier
|
||||
|
||||
# Example: Atlas Cloud configuration (OpenAI-compatible, https://www.atlascloud.ai)
|
||||
# HINDSIGHT_API_LLM_PROVIDER=atlas
|
||||
# HINDSIGHT_API_LLM_API_KEY=your-atlascloud-api-key
|
||||
# HINDSIGHT_API_LLM_MODEL=deepseek-ai/deepseek-v4-pro # reasoning model; also Qwen / GLM / Kimi / MiniMax, etc.
|
||||
|
||||
# Example: LM Studio local configuration (Qwen 2.5 32B recommended)
|
||||
# HINDSIGHT_API_LLM_PROVIDER=lmstudio
|
||||
# HINDSIGHT_API_LLM_API_KEY=lmstudio
|
||||
# HINDSIGHT_API_LLM_BASE_URL=http://localhost:1234/v1
|
||||
# HINDSIGHT_API_LLM_MODEL=qwen2.5-32b-instruct
|
||||
|
||||
# Example: Ollama local configuration (native provider)
|
||||
# HINDSIGHT_API_LLM_PROVIDER=ollama
|
||||
# HINDSIGHT_API_LLM_BASE_URL=http://localhost:11434/v1
|
||||
# HINDSIGHT_API_LLM_MODEL=gemma3:12b
|
||||
# Native Ollama context-window override (num_ctx). Leave unset to let Ollama use
|
||||
# the model Modelfile / server default; set a positive integer only to force a
|
||||
# specific context size (e.g. 16384 to keep the previous request behavior).
|
||||
# HINDSIGHT_API_LLM_OLLAMA_NUM_CTX=16384
|
||||
|
||||
# Multi-LLM strategies: configure extra LLMs by index alongside the primary above,
|
||||
# then pick a routing strategy. Unset = single primary LLM (default). Members are
|
||||
# numbered from 1; indices must be contiguous. Each operation can override with a
|
||||
# RETAIN_/REFLECT_/CONSOLIDATION_ prefix (e.g. HINDSIGHT_API_RETAIN_LLM_1_PROVIDER).
|
||||
# HINDSIGHT_API_LLM_1_PROVIDER=groq
|
||||
# HINDSIGHT_API_LLM_1_API_KEY=your-groq-api-key
|
||||
# HINDSIGHT_API_LLM_1_MODEL=openai/gpt-oss-120b
|
||||
# HINDSIGHT_API_LLM_2_PROVIDER=anthropic
|
||||
# HINDSIGHT_API_LLM_2_API_KEY=your-anthropic-api-key
|
||||
# Strategy JSON: {"mode": "failover"} or {"mode": "round-robin"}.
|
||||
# Round-robin accepts optional positive-int "weights" (one per member, primary first).
|
||||
# HINDSIGHT_API_LLM_STRATEGY={"mode": "failover"}
|
||||
|
||||
# API Configuration (Optional)
|
||||
HINDSIGHT_API_HOST=0.0.0.0
|
||||
HINDSIGHT_API_PORT=8888
|
||||
@@ -51,6 +104,10 @@ HINDSIGHT_API_LOG_LEVEL=info
|
||||
# Unset uses HINDSIGHT_API_RETAIN_CHUNK_SIZE as the structured-chunk limit.
|
||||
# HINDSIGHT_API_RETAIN_STRUCTURED_CHUNK_SIZE=
|
||||
|
||||
# When true, a retain operation that hit any fact-extraction errors is marked
|
||||
# 'failed' (not 'completed'), surfacing silently-dropped facts. Default false.
|
||||
# HINDSIGHT_API_FAIL_ON_EXTRACTION_ERRORS=false
|
||||
|
||||
# Dry-run extraction preview endpoint (POST /memories/dry-run-extract). Enabled by default; it makes
|
||||
# a real LLM call but stores nothing. Set to false to remove the endpoint (returns 404).
|
||||
# HINDSIGHT_API_ENABLE_DRY_RUN_EXTRACT=true
|
||||
@@ -66,7 +123,10 @@ HINDSIGHT_API_LOG_LEVEL=info
|
||||
# HINDSIGHT_API_READ_DATABASE_URL= # Optional read-replica URL. When set, recall queries (semantic, BM25, graph, temporal) flow through a separate pool against this URL, offloading the primary. Typically points to a read-only endpoint (CNPG's <cluster>-ro service or Aurora reader endpoint).
|
||||
# HINDSIGHT_API_MIGRATION_DATABASE_URL= # Direct PostgreSQL URL for migrations (bypasses PgBouncer). Falls back to DATABASE_URL.
|
||||
# HINDSIGHT_API_DATABASE_SCHEMA=public # PostgreSQL schema name (default: public)
|
||||
# HINDSIGHT_API_DB_MAX_PARALLEL_WORKERS_PER_GATHER= # Optional cap on Postgres planner parallelism for this process's pool connections. Unset leaves the server default; 0 makes background/bulk queries run serially (useful on worker processes sharing a primary with latency-sensitive traffic).
|
||||
# HINDSIGHT_API_MIGRATION_CONCURRENCY=1 # Tenant schemas to migrate concurrently (PG only, each in its own process; per-schema work stays sequential). Each worker has ~1-2s startup cost + uses ~3 DB connections, so it only pays off with many schemas (tens+) or slow migrations; keep concurrency*3 <= spare max_connections. Default: 1 (sequential).
|
||||
# HINDSIGHT_API_OPERATION_RETENTION_DAYS=30 # Prune terminal operation rows, payloads, and metadata after this many days; 0 (the default) keeps them forever.
|
||||
# HINDSIGHT_API_OPERATION_CLEANUP_BATCH_SIZE=1000 # Maximum expired terminal rows deleted per tenant schema in each cleanup cycle; must be positive.
|
||||
|
||||
# Vector Extension (Optional - uses pgvector by default)
|
||||
# Options: "pgvector" (default), "vchord", "pgvectorscale" (DiskANN)
|
||||
@@ -86,12 +146,31 @@ HINDSIGHT_API_LOG_LEVEL=info
|
||||
# chinese_lindera/lindera(chinese), japanese_lindera/lindera(japanese),
|
||||
# korean_lindera/lindera(korean), ngram(min,max), edge_ngram(min,max)
|
||||
# HINDSIGHT_API_TEXT_SEARCH_EXTENSION_PG_SEARCH_TOKENIZER=
|
||||
# Optional cap on the number of terms in the native PostgreSQL BM25 tsquery.
|
||||
# Long queries OR-join every normalized token, which can match too much of a
|
||||
# large bank. 0 (default) keeps the historical uncapped behavior; a positive
|
||||
# value bounds only the native backend (other BM25 backends get the raw query).
|
||||
# HINDSIGHT_API_BM25_MAX_QUERY_TERMS=0
|
||||
|
||||
# File Parser (Optional - uses markitdown by default)
|
||||
# HINDSIGHT_API_FILE_PARSER=markitdown
|
||||
# Enable image OCR for MarkItDown using an OpenAI-compatible OCR/vision endpoint.
|
||||
# These OCR settings are independent from HINDSIGHT_API_LLM_* because MarkItDown
|
||||
# uses the OpenAI SDK directly and requires Chat Completions image input support.
|
||||
# When OCR is enabled, API_KEY, BASE_URL, and MODEL are required.
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_ENABLED=false
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_API_KEY=
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_BASE_URL=
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_MODEL=
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_PROMPT=
|
||||
|
||||
# Embeddings Configuration (Optional - uses local by default)
|
||||
# Provider: "local" (default), "onnx", "tei", "openai", "cohere", "google", "openrouter", "zeroentropy", "litellm", or "litellm-sdk"
|
||||
# HINDSIGHT_API_EMBEDDINGS_PROVIDER=local
|
||||
# For local provider:
|
||||
# HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL=BAAI/bge-small-en-v1.5
|
||||
# Force CPU if local embeddings hit MPS/XPC instability on macOS:
|
||||
# HINDSIGHT_API_EMBEDDINGS_LOCAL_FORCE_CPU=false
|
||||
# For ONNX provider (local CPU embeddings without an Ollama/TEI sidecar):
|
||||
# HINDSIGHT_API_EMBEDDINGS_PROVIDER=onnx
|
||||
# HINDSIGHT_API_EMBEDDINGS_ONNX_MODEL_ID=intfloat/multilingual-e5-small
|
||||
@@ -131,8 +210,13 @@ HINDSIGHT_API_LOG_LEVEL=info
|
||||
# Reranker Configuration (Optional - uses local by default)
|
||||
# Provider: "local" (default) or "tei" (HuggingFace Text Embeddings Inference)
|
||||
# HINDSIGHT_API_RERANKER_PROVIDER=local
|
||||
# Trusted gateway attribution (disabled by default). When enabled, remote
|
||||
# reranker requests include X-Hindsight-Bank-Id with the current bank ID.
|
||||
# HINDSIGHT_API_RERANKER_SEND_BANK_AS_HEADER=false
|
||||
# For local provider:
|
||||
# HINDSIGHT_API_RERANKER_LOCAL_MODEL=cross-encoder/ms-marco-MiniLM-L-6-v2
|
||||
# Force CPU if the local reranker hits MPS/XPC instability on macOS:
|
||||
# HINDSIGHT_API_RERANKER_LOCAL_FORCE_CPU=false
|
||||
# For TEI provider:
|
||||
# HINDSIGHT_API_RERANKER_TEI_URL=http://localhost:8081
|
||||
|
||||
@@ -150,6 +234,20 @@ HINDSIGHT_API_LOG_LEVEL=info
|
||||
# Custom service name and environment (optional, defaults: hindsight-api, development)
|
||||
# HINDSIGHT_API_OTEL_SERVICE_NAME=hindsight-production
|
||||
# HINDSIGHT_API_OTEL_DEPLOYMENT_ENVIRONMENT=production
|
||||
#
|
||||
# Expose async-operation queue + consolidation-backlog gauges on /metrics.
|
||||
# Runs periodic per-schema COUNT queries on a background task (disabled by default).
|
||||
# HINDSIGHT_API_METRICS_BACKLOG_ENABLED=true
|
||||
#
|
||||
# Runtime-stall observability (enabled by default). When a liveness probe fails,
|
||||
# these tell you WHY: a blocked event loop vs DB connection-pool exhaustion.
|
||||
# The loop watchdog logs the offending stack when the loop is unresponsive; the
|
||||
# DB-pool acquire timing logs (and exposes hindsight.db.pool.waiting) when
|
||||
# callers queue for a connection. Both are cheap; tune or disable if needed.
|
||||
# HINDSIGHT_API_LOOP_WATCHDOG_ENABLED=false
|
||||
# HINDSIGHT_API_LOOP_WATCHDOG_STALL_THRESHOLD_MS=1000
|
||||
# HINDSIGHT_API_LOOP_WATCHDOG_POLL_INTERVAL_MS=250
|
||||
# HINDSIGHT_API_DB_ACQUIRE_WARN_THRESHOLD_MS=1000
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Control Plane (Optional)
|
||||
|
||||
@@ -1,6 +0,0 @@
|
||||
version: 2
|
||||
updates:
|
||||
- package-ecosystem: "github-actions"
|
||||
directory: "/"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
@@ -75,7 +75,7 @@ jobs:
|
||||
if: steps.type.outputs.type == 'plugin'
|
||||
run: |
|
||||
echo "Plugin integration ${{ steps.info.outputs.integration }} v${{ steps.info.outputs.version }} — no package to publish."
|
||||
echo "Users install via: claude plugin marketplace add vectorize-io/hindsight --sparse hindsight-integrations"
|
||||
echo "Users install via: claude plugin marketplace add vectorize-io/hindsight"
|
||||
|
||||
# ── TypeScript integrations (ai-sdk, chat, openclaw) ────────────────────
|
||||
|
||||
|
||||
@@ -266,7 +266,7 @@ jobs:
|
||||
strategy:
|
||||
matrix:
|
||||
include:
|
||||
- os: ubuntu-latest
|
||||
- os: ubuntu-22.04
|
||||
target: x86_64-unknown-linux-gnu
|
||||
artifact_name: hindsight
|
||||
asset_name: hindsight-linux-amd64
|
||||
@@ -278,7 +278,7 @@ jobs:
|
||||
target: aarch64-apple-darwin
|
||||
artifact_name: hindsight
|
||||
asset_name: hindsight-darwin-arm64
|
||||
- os: ubuntu-24.04-arm
|
||||
- os: ubuntu-22.04-arm
|
||||
target: aarch64-unknown-linux-gnu
|
||||
artifact_name: hindsight
|
||||
asset_name: hindsight-linux-arm64
|
||||
|
||||
+295
-5
@@ -38,24 +38,31 @@ jobs:
|
||||
integrations-claude-code: ${{ steps.filter.outputs.integrations-claude-code }}
|
||||
integrations-cline: ${{ steps.filter.outputs.integrations-cline }}
|
||||
integrations-codex: ${{ steps.filter.outputs.integrations-codex }}
|
||||
integrations-github-copilot: ${{ steps.filter.outputs.integrations-github-copilot }}
|
||||
integrations-continue: ${{ steps.filter.outputs.integrations-continue }}
|
||||
integrations-cursor-cli: ${{ steps.filter.outputs.integrations-cursor-cli }}
|
||||
integrations-zcode: ${{ steps.filter.outputs.integrations-zcode }}
|
||||
integrations-crewai: ${{ steps.filter.outputs.integrations-crewai }}
|
||||
integrations-litellm: ${{ steps.filter.outputs.integrations-litellm }}
|
||||
integrations-pydantic-ai: ${{ steps.filter.outputs.integrations-pydantic-ai }}
|
||||
integrations-ag2: ${{ steps.filter.outputs.integrations-ag2 }}
|
||||
integrations-autogen: ${{ steps.filter.outputs.integrations-autogen }}
|
||||
integrations-aider: ${{ steps.filter.outputs.integrations-aider }}
|
||||
integrations-langgraph: ${{ steps.filter.outputs.integrations-langgraph }}
|
||||
integrations-llamaindex: ${{ steps.filter.outputs.integrations-llamaindex }}
|
||||
integrations-paperclip: ${{ steps.filter.outputs.integrations-paperclip }}
|
||||
integrations-opencode: ${{ steps.filter.outputs.integrations-opencode }}
|
||||
integrations-eve: ${{ steps.filter.outputs.integrations-eve }}
|
||||
integrations-cursor: ${{ steps.filter.outputs.integrations-cursor }}
|
||||
integrations-zed: ${{ steps.filter.outputs.integrations-zed }}
|
||||
integrations-n8n: ${{ steps.filter.outputs.integrations-n8n }}
|
||||
integrations-zapier: ${{ steps.filter.outputs.integrations-zapier }}
|
||||
integrations-cloudflare-oauth-proxy: ${{ steps.filter.outputs.integrations-cloudflare-oauth-proxy }}
|
||||
integrations-superagent: ${{ steps.filter.outputs.integrations-superagent }}
|
||||
integrations-lockfiles: ${{ steps.filter.outputs.integrations-lockfiles }}
|
||||
integrations-openai-agents: ${{ steps.filter.outputs.integrations-openai-agents }}
|
||||
integrations-openhands: ${{ steps.filter.outputs.integrations-openhands }}
|
||||
integrations-devin-desktop: ${{ steps.filter.outputs.integrations-devin-desktop }}
|
||||
integrations-pipecat: ${{ steps.filter.outputs.integrations-pipecat }}
|
||||
integrations-agentcore: ${{ steps.filter.outputs.integrations-agentcore }}
|
||||
integrations-smolagents: ${{ steps.filter.outputs.integrations-smolagents }}
|
||||
@@ -140,6 +147,8 @@ jobs:
|
||||
- 'hindsight-integrations/cline/**'
|
||||
integrations-codex:
|
||||
- 'hindsight-integrations/codex/**'
|
||||
integrations-github-copilot:
|
||||
- 'hindsight-integrations/github-copilot/**'
|
||||
integrations-continue:
|
||||
- 'hindsight-integrations/continue/**'
|
||||
integrations-cursor-cli:
|
||||
@@ -154,6 +163,8 @@ jobs:
|
||||
- 'hindsight-integrations/ag2/**'
|
||||
integrations-autogen:
|
||||
- 'hindsight-integrations/autogen/**'
|
||||
integrations-aider:
|
||||
- 'hindsight-integrations/aider/**'
|
||||
integrations-langgraph:
|
||||
- 'hindsight-integrations/langgraph/**'
|
||||
integrations-llamaindex:
|
||||
@@ -164,8 +175,14 @@ jobs:
|
||||
- 'hindsight-integrations/paperclip/**'
|
||||
integrations-opencode:
|
||||
- 'hindsight-integrations/opencode/**'
|
||||
integrations-eve:
|
||||
- 'hindsight-integrations/eve/**'
|
||||
integrations-cursor:
|
||||
- 'hindsight-integrations/cursor/**'
|
||||
integrations-zed:
|
||||
- 'hindsight-integrations/zed/**'
|
||||
integrations-zcode:
|
||||
- 'hindsight-integrations/zcode/**'
|
||||
integrations-n8n:
|
||||
- 'hindsight-integrations/n8n/**'
|
||||
integrations-zapier:
|
||||
@@ -180,6 +197,10 @@ jobs:
|
||||
- 'scripts/check-integration-lockfiles.sh'
|
||||
integrations-openai-agents:
|
||||
- 'hindsight-integrations/openai-agents/**'
|
||||
integrations-openhands:
|
||||
- 'hindsight-integrations/openhands/**'
|
||||
integrations-devin-desktop:
|
||||
- 'hindsight-integrations/devin-desktop/**'
|
||||
integrations-pipecat:
|
||||
- 'hindsight-integrations/pipecat/**'
|
||||
integrations-agentcore:
|
||||
@@ -265,6 +286,18 @@ jobs:
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv build
|
||||
|
||||
# `uv build` only packages the source; it does not prove the dependency set
|
||||
# resolves or that the code imports on this interpreter. Install into a fresh
|
||||
# env and run a byte-compile + import smoke test so the matrix actually
|
||||
# exercises each Python version (notably 3.14).
|
||||
- name: Install and smoke-test on Python ${{ matrix.python-version }}
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: |
|
||||
uv venv --python ${{ matrix.python-version }} .venv-smoke
|
||||
VIRTUAL_ENV=.venv-smoke uv pip install .
|
||||
.venv-smoke/bin/python -m compileall -q hindsight_api
|
||||
.venv-smoke/bin/python -c "import hindsight_api, hindsight_api.main, hindsight_api.config; from hindsight_api.engine import memory_engine, llm_wrapper; print('import OK')"
|
||||
|
||||
build-typescript-client:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -488,6 +521,32 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/cursor
|
||||
run: python -m pytest tests/ -v
|
||||
|
||||
test-zed-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
github.event_name != 'pull_request_review' &&
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-zed == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '22'
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/zed
|
||||
# Config-only integration with no dependencies — it uses Node's built-in
|
||||
# test runner. The runtime MCP bridge is `npx mcp-remote` (Node), so this
|
||||
# integration requires only Node.js (no Python).
|
||||
run: npm test
|
||||
|
||||
test-omo-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -551,6 +610,45 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/cline
|
||||
run: uv run pytest tests -v
|
||||
|
||||
test-github-copilot-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-github-copilot == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build github-copilot integration
|
||||
working-directory: ./hindsight-integrations/github-copilot
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/github-copilot
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/github-copilot
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-codex-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -614,6 +712,43 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/cursor-cli
|
||||
run: uv run pytest tests -v
|
||||
|
||||
test-zcode-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-zcode == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build zcode integration
|
||||
working-directory: ./hindsight-integrations/zcode
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/zcode
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/zcode
|
||||
run: uv run pytest tests -v
|
||||
|
||||
build-ai-sdk-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -708,6 +843,37 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/opencode
|
||||
run: npm run build
|
||||
|
||||
test-eve-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-eve == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '24'
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/eve
|
||||
run: npm ci
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/eve
|
||||
run: npm test
|
||||
|
||||
- name: Build
|
||||
working-directory: ./hindsight-integrations/eve
|
||||
run: npm run build
|
||||
|
||||
test-n8n-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -1110,10 +1276,10 @@ jobs:
|
||||
|
||||
build-docs:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.docs == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
# Keep the production docs build as an unconditional PR check. OpenAPI
|
||||
# generation used to build the site again inside verify-generated-files;
|
||||
# running the existing job for every PR preserves that coverage without
|
||||
# serializing two full Docusaurus builds in the generated-files check.
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
@@ -3021,6 +3187,45 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/ag2
|
||||
run: uv run pytest tests -v
|
||||
|
||||
test-aider-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-aider == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build aider integration
|
||||
working-directory: ./hindsight-integrations/aider
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/aider
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/aider
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-autogen-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -3669,6 +3874,84 @@ jobs:
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-openhands-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-openhands == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build openhands integration
|
||||
working-directory: ./hindsight-integrations/openhands
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/openhands
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/openhands
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-devin-desktop-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-devin-desktop == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build devin-desktop integration
|
||||
working-directory: ./hindsight-integrations/devin-desktop
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/devin-desktop
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/devin-desktop
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-claude-agent-sdk-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -4535,7 +4818,8 @@ jobs:
|
||||
cd ../hindsight-embed && uv sync --frozen --index-strategy unsafe-best-match
|
||||
|
||||
- name: Run generate-openapi
|
||||
run: ./scripts/generate-openapi.sh
|
||||
working-directory: hindsight-dev
|
||||
run: uv run generate-openapi
|
||||
|
||||
- name: Run generate-bank-template-schema
|
||||
run: ./scripts/generate-bank-template-schema.sh
|
||||
@@ -4670,11 +4954,14 @@ jobs:
|
||||
- test-claude-code-integration
|
||||
- test-cursor-integration
|
||||
- test-cline-integration
|
||||
- test-github-copilot-integration
|
||||
- test-codex-integration
|
||||
- test-cursor-cli-integration
|
||||
- test-zcode-integration
|
||||
- build-ai-sdk-integration
|
||||
- test-ai-sdk-integration-deno
|
||||
- test-opencode-integration
|
||||
- test-eve-integration
|
||||
- test-omo-integration
|
||||
- test-cloudflare-oauth-proxy-integration
|
||||
- build-chat-integration
|
||||
@@ -4703,6 +4990,7 @@ jobs:
|
||||
- test-openclaw-integration
|
||||
- test-integration
|
||||
- test-ag2-integration
|
||||
- test-aider-integration
|
||||
- test-autogen-integration
|
||||
- test-continue-integration
|
||||
- test-smolagents-integration
|
||||
@@ -4717,6 +5005,8 @@ jobs:
|
||||
- test-pydantic-ai-integration
|
||||
- test-llamaindex-integration
|
||||
- test-openai-agents-integration
|
||||
- test-openhands-integration
|
||||
- test-devin-desktop-integration
|
||||
- test-agentcore-integration
|
||||
- test-haystack-integration
|
||||
- test-pip-slim
|
||||
|
||||
@@ -6,6 +6,7 @@ dist/
|
||||
wheels/
|
||||
*.egg-info
|
||||
.mcp.json
|
||||
.playwright-mcp/
|
||||
.osgrep
|
||||
# Virtual environments
|
||||
.venv
|
||||
@@ -40,6 +41,7 @@ nltk_data/
|
||||
logs/
|
||||
|
||||
.DS_Store
|
||||
.sesskey
|
||||
|
||||
# Generated docs files
|
||||
hindsight-docs/static/llms-full.txt
|
||||
|
||||
@@ -70,7 +70,7 @@ docker run -it --pull always --name hindsight --restart unless-stopped -p 8888:8
|
||||
>API: http://localhost:8888
|
||||
>UI: http://localhost:9999
|
||||
|
||||
You can modify the LLM provider by setting `HINDSIGHT_API_LLM_PROVIDER`. Valid options are `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio`, and `minimax`. The documentation provides more details on [supported models](https://hindsight.vectorize.io/developer/models).
|
||||
You can modify the LLM provider by setting `HINDSIGHT_API_LLM_PROVIDER`. Valid options are `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio`, `minimax`, and `atlas` ([Atlas Cloud](https://www.atlascloud.ai/?utm_source=github&utm_medium=link&utm_campaign=hindsight)). The documentation provides more details on [supported models](https://hindsight.vectorize.io/developer/models).
|
||||
|
||||
|
||||
|
||||
@@ -250,7 +250,7 @@ Recall performs 4 retrieval strategies in parallel:
|
||||
- Graph: Entity/temporal/causal links
|
||||
- Temporal: Time range filtering
|
||||
|
||||

|
||||

|
||||
|
||||
The individual results from the retrievals are merged, then ordered by relevance using reciprocal rank fusion and a cross-encoder reranking model.
|
||||
|
||||
@@ -276,7 +276,7 @@ client = Hindsight(base_url="http://localhost:8888")
|
||||
client.reflect(bank_id="my-bank", query="What should I know about Alice?")
|
||||
```
|
||||
|
||||

|
||||

|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -41,28 +41,43 @@ RUN apt-get update && apt-get install -y \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& pip install --no-cache-dir uv
|
||||
|
||||
# Copy dependency files and README (required by pyproject.toml)
|
||||
# Copy the workspace lock and member metadata before source code so dependency
|
||||
# installation stays cacheable while matching the versions tested in CI.
|
||||
COPY pyproject.toml uv.lock ./
|
||||
COPY hindsight-all/pyproject.toml ./hindsight-all/
|
||||
COPY hindsight-api/pyproject.toml ./hindsight-api/
|
||||
COPY hindsight-api-slim/pyproject.toml ./api/
|
||||
COPY hindsight-api-slim/README.md ./api/
|
||||
|
||||
WORKDIR /app/api
|
||||
COPY hindsight-all-slim/pyproject.toml ./hindsight-all-slim/
|
||||
COPY hindsight-dev/pyproject.toml ./hindsight-dev/
|
||||
COPY hindsight-clients/python/pyproject.toml ./hindsight-clients/python/
|
||||
COPY hindsight-embed/pyproject.toml ./hindsight-embed/
|
||||
RUN ln -s api hindsight-api-slim
|
||||
|
||||
# Sync dependencies using appropriate extras based on INCLUDE_LOCAL_MODELS
|
||||
# local-ml: torch, sentence-transformers, transformers, einops, flashrank, mlx (optional)
|
||||
# embedded-db: pg0-embedded (always included for embedded PostgreSQL support)
|
||||
# ONNX Runtime embeddings are intentionally not bundled into the official
|
||||
# standalone image; install the local-onnx extra in custom images when needed.
|
||||
ENV UV_PROJECT_ENVIRONMENT=/app/api/.venv
|
||||
RUN if [ "$INCLUDE_LOCAL_MODELS" = "true" ]; then \
|
||||
uv sync --extra local-ml --extra embedded-db; \
|
||||
uv sync --locked --package hindsight-api-slim --no-install-package hindsight-api-slim --extra local-ml --extra embedded-db; \
|
||||
else \
|
||||
uv sync --extra embedded-db; \
|
||||
uv sync --locked --package hindsight-api-slim --no-install-package hindsight-api-slim --extra embedded-db; \
|
||||
fi
|
||||
|
||||
# Copy source code (alembic migrations are inside hindsight_api/)
|
||||
WORKDIR /app/api
|
||||
COPY hindsight-api-slim/hindsight_api ./hindsight_api
|
||||
|
||||
# Install the local package (uv sync only installed dependencies, not the package itself)
|
||||
RUN uv pip install -e .
|
||||
# Install the local package from the same validated lock after source is present.
|
||||
WORKDIR /app
|
||||
RUN if [ "$INCLUDE_LOCAL_MODELS" = "true" ]; then \
|
||||
uv sync --locked --package hindsight-api-slim --extra local-ml --extra embedded-db; \
|
||||
else \
|
||||
uv sync --locked --package hindsight-api-slim --extra embedded-db; \
|
||||
fi \
|
||||
&& uv pip check --python /app/api/.venv/bin/python
|
||||
|
||||
# =============================================================================
|
||||
# Stage: SDK Builder (needed for Control Plane)
|
||||
@@ -145,6 +160,8 @@ FROM python:3.11-slim AS api-only
|
||||
WORKDIR /app
|
||||
|
||||
# Note: libicu version varies by Debian version - try common versions in order
|
||||
# Runtime images use uv directly; remove pip build tooling after installation so
|
||||
# vulnerable setuptools-vendored packages and wheel are not shipped in production.
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
procps \
|
||||
@@ -154,7 +171,8 @@ RUN apt-get update && apt-get install -y \
|
||||
libossp-uuid16 \
|
||||
&& (apt-get install -y libicu72 2>/dev/null || apt-get install -y libicu74 2>/dev/null || apt-get install -y libicu76 2>/dev/null || true) \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& pip install --no-cache-dir uv
|
||||
&& pip install --no-cache-dir uv \
|
||||
&& pip uninstall --yes setuptools wheel
|
||||
|
||||
RUN useradd -m -s /bin/bash hindsight
|
||||
|
||||
@@ -292,6 +310,8 @@ WORKDIR /app
|
||||
|
||||
# Install Node.js, curl, uv, and system dependencies
|
||||
# Note: libicu version varies by Debian version - try common versions in order
|
||||
# Runtime images use uv directly; remove pip build tooling after installation so
|
||||
# vulnerable setuptools-vendored packages and wheel are not shipped in production.
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
procps \
|
||||
@@ -303,7 +323,8 @@ RUN apt-get update && apt-get install -y \
|
||||
&& curl -fsSL https://deb.nodesource.com/setup_20.x | bash - \
|
||||
&& apt-get install -y nodejs \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& pip install --no-cache-dir uv
|
||||
&& pip install --no-cache-dir uv \
|
||||
&& pip uninstall --yes setuptools wheel
|
||||
|
||||
RUN useradd -m -s /bin/bash hindsight
|
||||
|
||||
|
||||
@@ -2,8 +2,8 @@ apiVersion: v2
|
||||
name: hindsight
|
||||
description: Hindsight helm chart
|
||||
type: application
|
||||
version: 0.8.2
|
||||
appVersion: "0.8.2"
|
||||
version: 0.8.5
|
||||
appVersion: "0.8.5"
|
||||
keywords:
|
||||
- ai
|
||||
- memory
|
||||
|
||||
@@ -60,13 +60,13 @@ spec:
|
||||
valueFrom:
|
||||
fieldRef:
|
||||
fieldPath: metadata.name
|
||||
{{- /* Inherit LLM config from api.env */}}
|
||||
{{- range $key, $value := .Values.api.env }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
{{- /* Worker-specific env vars */}}
|
||||
{{- range $key, $value := .Values.worker.env }}
|
||||
{{- /* Explicitly set port to override K8s service discovery env var (HINDSIGHT_API_PORT) */}}
|
||||
- name: HINDSIGHT_API_PORT
|
||||
value: {{ .Values.worker.service.targetPort | quote }}
|
||||
{{- /* Inherit LLM config from api.env, then apply worker-specific env.
|
||||
Merge (worker.env wins) so a key set in both does not emit a
|
||||
duplicate env entry, which server-side apply rejects. */}}
|
||||
{{- range $key, $value := merge (deepCopy (.Values.worker.env | default dict)) (.Values.api.env | default dict) }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@vectorize-io/hindsight-all",
|
||||
"version": "0.8.2",
|
||||
"version": "0.8.5",
|
||||
"description": "Node.js programmatic lifecycle manager for Hindsight — embeds a local hindsight daemon in a Node application. Pair with @vectorize-io/hindsight-client for memory operations.",
|
||||
"main": "dist/index.js",
|
||||
"types": "dist/index.d.ts",
|
||||
|
||||
@@ -4,12 +4,12 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "hindsight-all-slim"
|
||||
version = "0.8.2"
|
||||
version = "0.8.5"
|
||||
description = "Hindsight: Agent Memory That Works Like Human Memory - Slim All-in-One Bundle"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
"hindsight-api-slim==0.8.2",
|
||||
"hindsight-api-slim==0.8.5",
|
||||
"hindsight-client>=0.0.7",
|
||||
"hindsight-embed>=0.1.0",
|
||||
]
|
||||
|
||||
@@ -4,12 +4,12 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "hindsight-all"
|
||||
version = "0.8.2"
|
||||
version = "0.8.5"
|
||||
description = "Hindsight: Agent Memory That Works Like Human Memory - All-in-One Bundle"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
"hindsight-api-slim[all]==0.8.2",
|
||||
"hindsight-api-slim[all]==0.8.5",
|
||||
"hindsight-client>=0.0.7",
|
||||
"hindsight-embed>=0.1.0",
|
||||
]
|
||||
@@ -21,7 +21,7 @@ hindsight-embed = { workspace = true }
|
||||
|
||||
[project.optional-dependencies]
|
||||
local-llm = [
|
||||
"hindsight-api-slim[local-llm]==0.8.2",
|
||||
"hindsight-api-slim[local-llm]==0.8.5",
|
||||
]
|
||||
test = [
|
||||
"pytest>=7.0.0",
|
||||
|
||||
@@ -121,7 +121,7 @@ This runs a stdio-based MCP server that can be used directly with MCP-compatible
|
||||
- **Entity Graph** — Automatic entity extraction and relationship tracking
|
||||
- **Temporal Reasoning** — Native support for time-based queries
|
||||
- **Disposition Traits** — Configurable skepticism, literalism, and empathy influence opinion formation
|
||||
- **Three Memory Types** — World facts, bank actions, and formed opinions with confidence scores
|
||||
- **Three Memory Types** — World facts, experience facts (the bank's own actions), and observations
|
||||
|
||||
## Documentation
|
||||
|
||||
|
||||
@@ -53,4 +53,4 @@ __all__ = [
|
||||
"RemoteTEICrossEncoder",
|
||||
"LLMConfig",
|
||||
]
|
||||
__version__ = "0.8.2"
|
||||
__version__ = "0.8.5"
|
||||
|
||||
@@ -9,6 +9,7 @@ import io
|
||||
import json
|
||||
import logging
|
||||
import zipfile
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
@@ -18,8 +19,10 @@ import typer
|
||||
|
||||
from ..config import DEFAULT_DATABASE_SCHEMA, HindsightConfig
|
||||
from ..engine.memory_engine import _current_schema
|
||||
from ..engine.retain.bank_utils import _vector_index_clause
|
||||
from ..engine.schema import fq_table_explicit as _fq_table
|
||||
from ..engine.transfer import export_bank
|
||||
from ..engine.vector_index_health import SchemaVectorIndexResult, repair_vector_indexes
|
||||
from ..extensions import TenantExtension, load_extension
|
||||
from ..pg0 import parse_pg0_url, resolve_database_url
|
||||
|
||||
@@ -65,7 +68,88 @@ BACKUP_TABLES = [
|
||||
"graph_maintenance_queue",
|
||||
]
|
||||
|
||||
MANIFEST_VERSION = "1"
|
||||
MANIFEST_VERSION = "2"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BackupColumn:
|
||||
"""A PostgreSQL column shape required to decode a binary COPY stream."""
|
||||
|
||||
name: str
|
||||
type_name: str
|
||||
|
||||
|
||||
async def _table_columns(conn: asyncpg.Connection, schema: str, table: str) -> list[BackupColumn]:
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT a.attname AS name, pg_catalog.format_type(a.atttypid, a.atttypmod) AS type_name
|
||||
FROM pg_catalog.pg_attribute AS a
|
||||
JOIN pg_catalog.pg_class AS c ON c.oid = a.attrelid
|
||||
JOIN pg_catalog.pg_namespace AS n ON n.oid = c.relnamespace
|
||||
WHERE n.nspname = $1 AND c.relname = $2 AND a.attnum > 0 AND NOT a.attisdropped
|
||||
AND a.attgenerated = ''
|
||||
ORDER BY a.attnum
|
||||
""",
|
||||
schema,
|
||||
table,
|
||||
)
|
||||
return [BackupColumn(name=row["name"], type_name=row["type_name"]) for row in rows]
|
||||
|
||||
|
||||
async def _validate_restore_schema(
|
||||
conn: asyncpg.Connection, manifest: dict[str, Any], schema: str
|
||||
) -> dict[str, list[str]]:
|
||||
"""Validate every COPY stream against the target before destructive work starts.
|
||||
|
||||
Type equality is an exact ``format_type`` string match. This is deliberately
|
||||
stricter than binary-COPY wire compatibility (e.g. ``varchar`` and ``text``
|
||||
share a binary format yet compare unequal here): we would rather fail a
|
||||
genuinely-restorable backup with a clear, actionable error than silently risk
|
||||
a subtle binary mismatch. Restores blocked this way can be recovered by
|
||||
aligning the target schema.
|
||||
"""
|
||||
restore_columns: dict[str, list[str]] = {}
|
||||
errors: list[str] = []
|
||||
for table, table_manifest in manifest["tables"].items():
|
||||
source_columns = [BackupColumn(**column) for column in table_manifest["columns"]]
|
||||
target_by_name = {column.name: column for column in await _table_columns(conn, schema, table)}
|
||||
missing = [column.name for column in source_columns if column.name not in target_by_name]
|
||||
mismatched = [
|
||||
f"{column.name} ({column.type_name} in backup, {target_by_name[column.name].type_name} in target)"
|
||||
for column in source_columns
|
||||
if column.name in target_by_name and target_by_name[column.name].type_name != column.type_name
|
||||
]
|
||||
if missing:
|
||||
errors.append(f"{table}: target is missing backup columns {', '.join(missing)}")
|
||||
if mismatched:
|
||||
errors.append(f"{table}: incompatible column types: {', '.join(mismatched)}")
|
||||
restore_columns[table] = [column.name for column in source_columns]
|
||||
|
||||
if errors:
|
||||
details = "; ".join(errors)
|
||||
raise ValueError(f"Backup schema is incompatible with target schema '{schema}': {details}")
|
||||
return restore_columns
|
||||
|
||||
|
||||
def _effective_backup_tables() -> list[str]:
|
||||
"""Core backup tables plus any bank-scoped tables a loaded extension declares.
|
||||
|
||||
``BACKUP_TABLES`` covers only the tables core owns. An extension that
|
||||
provisions its own bank-scoped tables (via ``TenantExtension``) declares
|
||||
them through ``extra_bank_tables()`` so they aren't dropped on restore.
|
||||
Extension tables are appended *after* the core set so restore's forward
|
||||
COPY inserts them after their FK parents (e.g. ``banks``) and the reversed
|
||||
TRUNCATE clears them before those parents.
|
||||
"""
|
||||
tables = list(BACKUP_TABLES)
|
||||
tenant_extension = load_extension("TENANT", TenantExtension)
|
||||
if tenant_extension is not None:
|
||||
seen = set(tables)
|
||||
for spec in tenant_extension.extra_bank_tables():
|
||||
if spec.include_in_backup and spec.name not in seen:
|
||||
tables.append(spec.name)
|
||||
seen.add(spec.name)
|
||||
return tables
|
||||
|
||||
|
||||
async def _admin_connect(db_url: str) -> asyncpg.Connection:
|
||||
@@ -76,7 +160,8 @@ async def _admin_connect(db_url: str) -> asyncpg.Connection:
|
||||
is the only step needed to connect. JSON codecs are registered so ``jsonb``
|
||||
columns decode to Python objects (used by the export row dumps).
|
||||
"""
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
_pg0 = parse_pg0_url(db_url)
|
||||
is_pg0, instance_name = _pg0.is_pg0, _pg0.instance_name
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
conn = await asyncpg.connect(await resolve_database_url(db_url))
|
||||
@@ -85,8 +170,18 @@ async def _admin_connect(db_url: str) -> asyncpg.Connection:
|
||||
return conn
|
||||
|
||||
|
||||
async def _backup(database_url: str, output_path: Path, schema: str = "public") -> dict[str, Any]:
|
||||
"""Backup all tables to a zip file using binary COPY protocol."""
|
||||
async def _backup(
|
||||
database_url: str,
|
||||
output_path: Path,
|
||||
schema: str = "public",
|
||||
backup_tables: list[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Backup all tables to a zip file using binary COPY protocol.
|
||||
|
||||
``backup_tables`` defaults to the core ``BACKUP_TABLES``; callers pass the
|
||||
extension-augmented list from ``_effective_backup_tables()``.
|
||||
"""
|
||||
backup_tables = backup_tables if backup_tables is not None else BACKUP_TABLES
|
||||
conn = await asyncpg.connect(database_url)
|
||||
try:
|
||||
tables: dict[str, Any] = {}
|
||||
@@ -103,14 +198,24 @@ async def _backup(database_url: str, output_path: Path, schema: str = "public")
|
||||
# entities table was backed up.
|
||||
async with conn.transaction(isolation="repeatable_read"):
|
||||
with zipfile.ZipFile(output_path, "w", zipfile.ZIP_DEFLATED) as zf:
|
||||
for i, table in enumerate(BACKUP_TABLES, 1):
|
||||
typer.echo(f" [{i}/{len(BACKUP_TABLES)}] Backing up {table}...", nl=False)
|
||||
for i, table in enumerate(backup_tables, 1):
|
||||
typer.echo(f" [{i}/{len(backup_tables)}] Backing up {table}...", nl=False)
|
||||
|
||||
buffer = io.BytesIO()
|
||||
|
||||
# Use binary COPY for exact type preservation
|
||||
columns = await _table_columns(conn, schema, table)
|
||||
|
||||
# Pin the ordered columns into both the stream and manifest.
|
||||
# PostgreSQL binary COPY does not encode column identities, so
|
||||
# restore must validate this shape before truncating any data.
|
||||
# asyncpg requires schema_name as separate parameter
|
||||
await conn.copy_from_table(table, schema_name=schema, output=buffer, format="binary")
|
||||
await conn.copy_from_table(
|
||||
table,
|
||||
schema_name=schema,
|
||||
columns=[column.name for column in columns],
|
||||
output=buffer,
|
||||
format="binary",
|
||||
)
|
||||
|
||||
data = buffer.getvalue()
|
||||
zf.writestr(f"{table}.bin", data)
|
||||
@@ -121,6 +226,7 @@ async def _backup(database_url: str, output_path: Path, schema: str = "public")
|
||||
tables[table] = {
|
||||
"rows": row_count,
|
||||
"size_bytes": len(data),
|
||||
"columns": [{"name": column.name, "type_name": column.type_name} for column in columns],
|
||||
}
|
||||
|
||||
typer.echo(f" {row_count} rows")
|
||||
@@ -132,8 +238,20 @@ async def _backup(database_url: str, output_path: Path, schema: str = "public")
|
||||
await conn.close()
|
||||
|
||||
|
||||
async def _restore(database_url: str, input_path: Path, schema: str = "public") -> dict[str, Any]:
|
||||
"""Restore all tables from a zip file using binary COPY protocol."""
|
||||
async def _restore(
|
||||
database_url: str,
|
||||
input_path: Path,
|
||||
schema: str = "public",
|
||||
backup_tables: list[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Restore all tables from a zip file using binary COPY protocol.
|
||||
|
||||
``backup_tables`` defaults to the core ``BACKUP_TABLES``; callers pass the
|
||||
extension-augmented list from ``_effective_backup_tables()``. Tables named
|
||||
here but absent from the archive are truncated then skipped for restore, so
|
||||
a stale extension registration never leaves pre-restore rows behind.
|
||||
"""
|
||||
backup_tables = backup_tables if backup_tables is not None else BACKUP_TABLES
|
||||
conn = await asyncpg.connect(database_url)
|
||||
try:
|
||||
with zipfile.ZipFile(input_path, "r") as zf:
|
||||
@@ -142,29 +260,40 @@ async def _restore(database_url: str, input_path: Path, schema: str = "public")
|
||||
if manifest.get("version") != MANIFEST_VERSION:
|
||||
raise ValueError(f"Unsupported backup version: {manifest.get('version')}")
|
||||
|
||||
# Complete the compatibility check before entering the transaction
|
||||
# that truncates tables. This turns historical schema drift into an
|
||||
# actionable error without risking the target's existing data.
|
||||
restore_columns = await _validate_restore_schema(conn, manifest, schema)
|
||||
|
||||
# Use a transaction for atomic restore - either all tables are
|
||||
# restored or none are, preventing partial/inconsistent state.
|
||||
async with conn.transaction():
|
||||
typer.echo(" Clearing existing data...")
|
||||
# Truncate tables in reverse order (respects FK constraints)
|
||||
for table in reversed(BACKUP_TABLES):
|
||||
for table in reversed(backup_tables):
|
||||
qualified_table = _fq_table(table, schema)
|
||||
await conn.execute(f"TRUNCATE TABLE {qualified_table} CASCADE")
|
||||
|
||||
# Restore tables in forward order
|
||||
for i, table in enumerate(BACKUP_TABLES, 1):
|
||||
for i, table in enumerate(backup_tables, 1):
|
||||
filename = f"{table}.bin"
|
||||
if filename not in zf.namelist():
|
||||
typer.echo(f" [{i}/{len(BACKUP_TABLES)}] {table}: skipped (not in backup)")
|
||||
typer.echo(f" [{i}/{len(backup_tables)}] {table}: skipped (not in backup)")
|
||||
continue
|
||||
|
||||
expected_rows = manifest["tables"].get(table, {}).get("rows", "?")
|
||||
typer.echo(f" [{i}/{len(BACKUP_TABLES)}] Restoring {table}... {expected_rows} rows")
|
||||
typer.echo(f" [{i}/{len(backup_tables)}] Restoring {table}... {expected_rows} rows")
|
||||
|
||||
data = zf.read(filename)
|
||||
buffer = io.BytesIO(data)
|
||||
# asyncpg requires schema_name as separate parameter
|
||||
await conn.copy_to_table(table, schema_name=schema, source=buffer, format="binary")
|
||||
await conn.copy_to_table(
|
||||
table,
|
||||
schema_name=schema,
|
||||
columns=restore_columns[table],
|
||||
source=buffer,
|
||||
format="binary",
|
||||
)
|
||||
|
||||
# Refresh materialized view
|
||||
typer.echo(" Refreshing materialized views...")
|
||||
@@ -177,20 +306,22 @@ async def _restore(database_url: str, input_path: Path, schema: str = "public")
|
||||
|
||||
async def _run_backup(db_url: str, output: Path, schema: str = "public") -> dict[str, Any]:
|
||||
"""Resolve database URL and run backup."""
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
_pg0 = parse_pg0_url(db_url)
|
||||
is_pg0, instance_name = _pg0.is_pg0, _pg0.instance_name
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
return await _backup(resolved_url, output, schema)
|
||||
return await _backup(resolved_url, output, schema, backup_tables=_effective_backup_tables())
|
||||
|
||||
|
||||
async def _run_restore(db_url: str, input_file: Path, schema: str = "public") -> dict[str, Any]:
|
||||
"""Resolve database URL and run restore."""
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
_pg0 = parse_pg0_url(db_url)
|
||||
is_pg0, instance_name = _pg0.is_pg0, _pg0.instance_name
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
return await _restore(resolved_url, input_file, schema)
|
||||
return await _restore(resolved_url, input_file, schema, backup_tables=_effective_backup_tables())
|
||||
|
||||
|
||||
@app.command()
|
||||
@@ -214,7 +345,7 @@ def backup(
|
||||
manifest = asyncio.run(_run_backup(config.database_url, output, schema))
|
||||
|
||||
total_rows = sum(t["rows"] for t in manifest["tables"].values())
|
||||
typer.echo(f"Backed up {total_rows} rows across {len(BACKUP_TABLES)} tables")
|
||||
typer.echo(f"Backed up {total_rows} rows across {len(manifest['tables'])} tables")
|
||||
typer.echo(f"Backup saved to {output}")
|
||||
|
||||
|
||||
@@ -247,7 +378,7 @@ def restore(
|
||||
manifest = asyncio.run(_run_restore(config.database_url, input_file, schema))
|
||||
|
||||
total_rows = sum(t["rows"] for t in manifest["tables"].values())
|
||||
typer.echo(f"Restored {total_rows} rows across {len(BACKUP_TABLES)} tables")
|
||||
typer.echo(f"Restored {total_rows} rows across {len(manifest['tables'])} tables")
|
||||
typer.echo("Restore complete")
|
||||
|
||||
|
||||
@@ -256,21 +387,22 @@ async def _run_migration(
|
||||
schema: str | None = None,
|
||||
base_schema: str = DEFAULT_DATABASE_SCHEMA,
|
||||
embedding_dimension: int | None = None,
|
||||
ensure_extensions: bool = True,
|
||||
) -> list[str]:
|
||||
"""Resolve database URL and run migrations for one schema or all discovered schemas."""
|
||||
from ..migrations import run_migrations_for_schemas
|
||||
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
_pg0 = parse_pg0_url(db_url)
|
||||
is_pg0, instance_name = _pg0.is_pg0, _pg0.instance_name
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
|
||||
config = HindsightConfig.from_env()
|
||||
tenant_extension = load_extension("TENANT", TenantExtension)
|
||||
if schema:
|
||||
schemas = [schema]
|
||||
else:
|
||||
tenant_extension = load_extension("TENANT", TenantExtension)
|
||||
|
||||
schemas = [base_schema or DEFAULT_DATABASE_SCHEMA]
|
||||
if tenant_extension:
|
||||
tenants = await tenant_extension.list_tenants()
|
||||
@@ -292,12 +424,39 @@ async def _run_migration(
|
||||
vector_extension=config.vector_extension,
|
||||
text_search_extension=config.text_search_extension,
|
||||
pg_search_tokenizer=config.text_search_extension_pg_search_tokenizer,
|
||||
ensure_extensions=True,
|
||||
ensure_extensions=ensure_extensions,
|
||||
)
|
||||
|
||||
# After core migrations, provision any extension-owned bank-scoped tables
|
||||
# per schema so extension schema evolves on the same lifecycle as core
|
||||
# schema (rather than via a lazy first-request path).
|
||||
if tenant_extension is not None:
|
||||
await _provision_extra_bank_tables(resolved_url, schemas, tenant_extension)
|
||||
|
||||
return schemas
|
||||
|
||||
|
||||
async def _provision_extra_bank_tables(
|
||||
resolved_url: str, schemas: list[str], tenant_extension: TenantExtension
|
||||
) -> None:
|
||||
"""Run the tenant extension's table provisioner for each migrated schema.
|
||||
|
||||
Fires after core migrations complete so extension-owned bank tables are
|
||||
created/evolved on the same lifecycle as core schema. A failure aborts the
|
||||
migration command (and names the offending schema) rather than being
|
||||
swallowed — provisioning is idempotent, so the operator can fix and re-run.
|
||||
"""
|
||||
for schema in schemas:
|
||||
conn = await asyncpg.connect(resolved_url)
|
||||
try:
|
||||
await tenant_extension.provision_bank_tables(conn, schema)
|
||||
except Exception as e:
|
||||
typer.echo(f" Failed to provision extension tables for schema '{schema}': {e}", err=True)
|
||||
raise
|
||||
finally:
|
||||
await conn.close()
|
||||
|
||||
|
||||
@app.command(name="run-db-migration")
|
||||
def run_db_migration(
|
||||
schema: str | None = typer.Option(
|
||||
@@ -311,6 +470,18 @@ def run_db_migration(
|
||||
"--embedding-dimension",
|
||||
help="Expected embedding dimension to enforce after migrations. Omit to skip dimension sync.",
|
||||
),
|
||||
skip_extension_reconcile: bool = typer.Option(
|
||||
False,
|
||||
"--skip-extension-reconcile",
|
||||
help=(
|
||||
"Skip the post-migration vector / text-search index reconcile. This step only does "
|
||||
"work when the configured backend (HINDSIGHT_API_VECTOR_EXTENSION / "
|
||||
"HINDSIGHT_API_TEXT_SEARCH_EXTENSION) differs from a schema's existing indexes — a "
|
||||
"rare, operator-driven change. Skipping it makes a no-change re-migration over many "
|
||||
"tenant schemas much faster. Only use when you have NOT changed the backend; a "
|
||||
"backend change still needs a normal run to reshape the indexes."
|
||||
),
|
||||
),
|
||||
):
|
||||
"""Run database migrations to the latest version."""
|
||||
config = HindsightConfig.from_env()
|
||||
@@ -324,6 +495,8 @@ def run_db_migration(
|
||||
typer.echo(f"Running database migrations for schema: {schema}...")
|
||||
else:
|
||||
typer.echo("Running database migrations for base schema and all discovered tenant schemas...")
|
||||
if skip_extension_reconcile:
|
||||
typer.echo("Skipping post-migration extension reconcile (--skip-extension-reconcile).")
|
||||
|
||||
schemas = asyncio.run(
|
||||
_run_migration(
|
||||
@@ -331,12 +504,141 @@ def run_db_migration(
|
||||
schema=schema,
|
||||
base_schema=config.database_schema,
|
||||
embedding_dimension=embedding_dimension,
|
||||
ensure_extensions=not skip_extension_reconcile,
|
||||
)
|
||||
)
|
||||
|
||||
typer.echo(f"Database migrations completed successfully for {len(schemas)} schema(s)")
|
||||
|
||||
|
||||
async def _resolve_schemas(base_schema: str | None) -> list[str]:
|
||||
"""Base schema plus every discovered tenant schema, de-duplicated in order."""
|
||||
schemas = [base_schema or DEFAULT_DATABASE_SCHEMA]
|
||||
tenant_extension = load_extension("TENANT", TenantExtension)
|
||||
if tenant_extension:
|
||||
tenants = await tenant_extension.list_tenants()
|
||||
schemas.extend(tenant.schema for tenant in tenants if tenant.schema)
|
||||
return list(dict.fromkeys(schemas))
|
||||
|
||||
|
||||
async def _run_repair_bank(
|
||||
db_url: str,
|
||||
*,
|
||||
base_schema: str,
|
||||
schema: str | None,
|
||||
bank_id: str | None,
|
||||
dry_run: bool,
|
||||
) -> list[SchemaVectorIndexResult]:
|
||||
"""Reconcile per-(bank, fact_type) vector index coverage over a raw connection.
|
||||
|
||||
A single autocommit connection is used because ``CREATE INDEX CONCURRENTLY``
|
||||
(used by ``repair_vector_indexes``) cannot run inside a transaction block.
|
||||
"""
|
||||
schemas = [schema] if schema else await _resolve_schemas(base_schema)
|
||||
index_clause = _vector_index_clause()
|
||||
# Guarded by the command, but assert so this helper is never called for a
|
||||
# backend without per-bank indexes.
|
||||
assert index_clause is not None
|
||||
|
||||
conn = await _admin_connect(db_url)
|
||||
try:
|
||||
results = await repair_vector_indexes(conn, schemas, index_clause, dry_run=dry_run, bank_id=bank_id)
|
||||
for result in results:
|
||||
typer.echo(
|
||||
f" schema '{result.schema}': {result.banks_scanned} bank(s) scanned, "
|
||||
f"{result.already_present} present, {result.created} created, "
|
||||
f"{result.skipped} to-create (dry-run), {result.failed} failed"
|
||||
)
|
||||
return results
|
||||
finally:
|
||||
await conn.close()
|
||||
|
||||
|
||||
@app.command(name="repair-bank")
|
||||
def repair_bank(
|
||||
bank_id: str | None = typer.Option(
|
||||
None,
|
||||
"--bank",
|
||||
"-b",
|
||||
help="Bank id to repair. Mutually exclusive with --all.",
|
||||
),
|
||||
all_banks: bool = typer.Option(
|
||||
False,
|
||||
"--all",
|
||||
help="Repair every bank in the base schema and all discovered tenant schemas.",
|
||||
),
|
||||
schema: str | None = typer.Option(
|
||||
None,
|
||||
"--schema",
|
||||
"-s",
|
||||
help="Limit to a single schema. Defaults to the base schema plus discovered tenant schemas.",
|
||||
),
|
||||
dry_run: bool = typer.Option(
|
||||
False,
|
||||
"--dry-run",
|
||||
help="Report what would be repaired without creating or dropping any index.",
|
||||
),
|
||||
):
|
||||
"""Verify and repair a bank's per-(bank, fact_type) vector index coverage.
|
||||
|
||||
Per-bank partial vector indexes are created when a bank is first created
|
||||
(instant on an empty bank). Banks that arrive populated — via logical
|
||||
restore, a cross-version upgrade, or a vector-extension switch — never hit
|
||||
that path, so their recall silently falls back to a global index +
|
||||
post-filter (slower, under-returning). This command detects missing OR
|
||||
invalid coverage (an INVALID leftover or an index whose access method
|
||||
drifted counts as missing) and rebuilds it with CREATE INDEX CONCURRENTLY,
|
||||
so it never blocks the live fleet. Idempotent and safe to re-run — the
|
||||
escape hatch after a restore, upgrade, or backend switch.
|
||||
"""
|
||||
if bool(bank_id) == all_banks:
|
||||
typer.echo("Error: pass exactly one of --bank <id> or --all.", err=True)
|
||||
raise typer.Exit(2)
|
||||
|
||||
config = HindsightConfig.from_env()
|
||||
if not config.database_url:
|
||||
typer.echo("Error: Database URL not configured.", err=True)
|
||||
typer.echo("Set HINDSIGHT_API_DATABASE_URL environment variable.", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
# Backend guard: backends with a single global vector index (AlloyDB ScaNN,
|
||||
# Oracle) have no per-bank indexes to repair.
|
||||
if _vector_index_clause() is None:
|
||||
typer.echo("Configured vector backend does not use per-bank vector indexes — nothing to repair.")
|
||||
return
|
||||
|
||||
target = f"bank '{bank_id}'" if bank_id else "all banks"
|
||||
scope = f"schema '{schema}'" if schema else "base schema and all discovered tenant schemas"
|
||||
typer.echo(f"Repairing per-bank vector indexes for {target} across {scope}...")
|
||||
if dry_run:
|
||||
typer.echo("Dry run: no indexes will be created or dropped.")
|
||||
|
||||
results = asyncio.run(
|
||||
_run_repair_bank(
|
||||
config.database_url,
|
||||
base_schema=config.database_schema,
|
||||
schema=schema,
|
||||
bank_id=bank_id,
|
||||
dry_run=dry_run,
|
||||
)
|
||||
)
|
||||
|
||||
total_banks = sum(r.banks_scanned for r in results)
|
||||
total_present = sum(r.already_present for r in results)
|
||||
total_created = sum(r.created for r in results)
|
||||
total_skipped = sum(r.skipped for r in results)
|
||||
total_failed = sum(r.failed for r in results)
|
||||
typer.echo(
|
||||
f"Done: {len(results)} schema(s), {total_banks} bank(s) scanned, "
|
||||
f"{total_present} already present, {total_created} created, "
|
||||
f"{total_skipped} to-create (dry-run), {total_failed} failed"
|
||||
)
|
||||
if total_failed:
|
||||
failed_names = [name for r in results for name in r.failed_indexes]
|
||||
typer.echo(f"Failed indexes (dropped, retry with a re-run): {', '.join(failed_names)}", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
|
||||
async def _run_export_bank(db_url: str, bank_id: str, output: Path, schema: str, include_history: bool) -> int:
|
||||
"""Export a whole bank to a ZIP archive."""
|
||||
conn = await _admin_connect(db_url)
|
||||
@@ -344,7 +646,14 @@ async def _run_export_bank(db_url: str, bank_id: str, output: Path, schema: str,
|
||||
# export_bank resolves table names via fq_table (the _current_schema
|
||||
# contextvar); set it so the raw connection targets the right schema.
|
||||
_current_schema.set(schema)
|
||||
data = await export_bank(conn, bank_id, include_history=include_history)
|
||||
# _admin_connect registers JSON codecs, so row dumps already contain
|
||||
# decoded Python values (including JSON scalar strings).
|
||||
data = await export_bank(
|
||||
conn,
|
||||
bank_id,
|
||||
include_history=include_history,
|
||||
bank_rows_json_encoding="decoded",
|
||||
)
|
||||
finally:
|
||||
await conn.close()
|
||||
|
||||
@@ -456,7 +765,8 @@ def import_bank_command(
|
||||
|
||||
async def _decommission_worker(db_url: str, worker_id: str, schema: str = "public") -> int:
|
||||
"""Release all tasks owned by a worker, setting them back to pending status."""
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
_pg0 = parse_pg0_url(db_url)
|
||||
is_pg0, instance_name = _pg0.is_pg0, _pg0.instance_name
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
@@ -515,7 +825,8 @@ def decommission_worker(
|
||||
|
||||
async def _decommission_all_workers(db_url: str, schema: str = "public") -> list[dict[str, Any]]:
|
||||
"""Release all processing tasks from all workers, setting them back to pending status."""
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
_pg0 = parse_pg0_url(db_url)
|
||||
is_pg0, instance_name = _pg0.is_pg0, _pg0.instance_name
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
@@ -580,7 +891,8 @@ def decommission_workers(
|
||||
|
||||
async def _worker_status(db_url: str, schema: str = "public") -> list[dict[str, Any]]:
|
||||
"""Get all processing tasks grouped by worker with their last update time."""
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
_pg0 = parse_pg0_url(db_url)
|
||||
is_pg0, instance_name = _pg0.is_pg0, _pg0.instance_name
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
|
||||
@@ -96,7 +96,8 @@ def get_database_url() -> str:
|
||||
# for the sync engine used during migrations.
|
||||
database_url = to_libpq_url(database_url)
|
||||
|
||||
config.set_main_option("sqlalchemy.url", database_url)
|
||||
# Alembic stores options through ConfigParser, where '%' is interpolation.
|
||||
config.set_main_option("sqlalchemy.url", database_url.replace("%", "%%"))
|
||||
return database_url
|
||||
|
||||
|
||||
|
||||
+82
@@ -0,0 +1,82 @@
|
||||
"""Add indexes for terminal cleanup and newest-first operation listing.
|
||||
|
||||
Revision ID: a8c1e4f7b0d3
|
||||
Revises: e7c3a9f1b2d5
|
||||
Create Date: 2026-07-14
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "a8c1e4f7b0d3"
|
||||
down_revision: str | Sequence[str] | None = "e7c3a9f1b2d5"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _pg_schema_prefix() -> str:
|
||||
"""Schema-qualifier for PostgreSQL multi-tenant migration runs."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
# These can be large tables in long-running installations. Concurrent DDL
|
||||
# keeps operation submission, polling, and status reads available.
|
||||
with op.get_context().autocommit_block():
|
||||
op.execute(
|
||||
"CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_async_operations_terminal_cleanup "
|
||||
f"ON {schema}async_operations (updated_at, operation_id) "
|
||||
"WHERE status IN ('completed', 'failed', 'cancelled')"
|
||||
)
|
||||
op.execute(
|
||||
"CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_async_operations_bank_created_desc "
|
||||
f"ON {schema}async_operations (bank_id, created_at DESC)"
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
with op.get_context().autocommit_block():
|
||||
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}idx_async_operations_bank_created_desc")
|
||||
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}idx_async_operations_terminal_cleanup")
|
||||
|
||||
|
||||
def _oracle_create_index(sql: str) -> None:
|
||||
"""Create an index idempotently for rerun-safe Oracle migrations."""
|
||||
block = (
|
||||
"BEGIN "
|
||||
"EXECUTE IMMEDIATE :stmt; "
|
||||
"EXCEPTION WHEN OTHERS THEN "
|
||||
"IF SQLCODE = -955 THEN NULL; ELSE RAISE; END IF; "
|
||||
"END;"
|
||||
)
|
||||
op.get_bind().exec_driver_sql(block, {"stmt": sql})
|
||||
|
||||
|
||||
def _oracle_upgrade() -> None:
|
||||
# Oracle migrations run with CURRENT_SCHEMA set to each tenant, so table
|
||||
# and index names intentionally remain unqualified here.
|
||||
_oracle_create_index(
|
||||
"CREATE INDEX idx_async_operations_terminal_cleanup ON async_operations (updated_at, operation_id, status)"
|
||||
)
|
||||
_oracle_create_index(
|
||||
"CREATE INDEX idx_async_operations_bank_created_desc ON async_operations (bank_id, created_at DESC)"
|
||||
)
|
||||
|
||||
|
||||
def _oracle_downgrade() -> None:
|
||||
op.execute("DROP INDEX idx_async_operations_bank_created_desc")
|
||||
op.execute("DROP INDEX idx_async_operations_terminal_cleanup")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade, oracle=_oracle_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade, oracle=_oracle_downgrade)
|
||||
+61
@@ -0,0 +1,61 @@
|
||||
"""Add bank_stats_cache table for distributed get_bank_stats caching
|
||||
|
||||
Revision ID: b57a7c9e0d13
|
||||
Revises: c3f7a1b9d2e4
|
||||
Create Date: 2026-07-01
|
||||
|
||||
get_bank_stats aggregates over memory_links / unit_entities — a multi-second scan
|
||||
on banks with millions of rows. The result was cached per-process (in-memory), so
|
||||
every API worker recomputed it once per TTL and the first caller after expiry
|
||||
stalled. This table backs a shared, cross-process TTL cache: one worker's compute
|
||||
is written here and served to all the others.
|
||||
|
||||
PostgreSQL only. Oracle keeps the in-process cache (the runtime picks the backing
|
||||
store by dialect), so the Oracle upgrade slot is intentionally absent.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "b57a7c9e0d13"
|
||||
down_revision: str | Sequence[str] | None = "c3f7a1b9d2e4"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Schema-qualifier for raw SQL on PG (multi-tenant search_path)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
# One row per bank: payload is the full get_bank_stats result, computed_at
|
||||
# drives logical TTL expiry. Rows are overwritten in place (ON CONFLICT), so
|
||||
# the table never grows beyond the number of banks and needs no purge job.
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {schema}bank_stats_cache (
|
||||
bank_id TEXT PRIMARY KEY,
|
||||
payload JSONB NOT NULL,
|
||||
computed_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}bank_stats_cache")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade) # oracle slot intentionally absent → no-op
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
+259
@@ -0,0 +1,259 @@
|
||||
"""Install the maintenance discovery routines into the configured schema.
|
||||
|
||||
The three discovery routines driving the background maintenance loop —
|
||||
``banks_needing_consolidation()``, ``schemas_with_expired_rows(...)`` and
|
||||
``mental_models_with_cron()`` — were installed into ``public`` and gated on the
|
||||
run being the base run (no ``target_schema``) or an explicit
|
||||
``target_schema='public'`` run (``e5f6a7b8c9d0`` → ``b2d4f6a8c1e3`` →
|
||||
``c7e9f1a3b5d2``, ``f4d1c2b3a5e6``).
|
||||
|
||||
That leaves a **single-tenant deployment migrated into a dedicated, non-**
|
||||
``public`` **schema** (``HINDSIGHT_API_DATABASE_SCHEMA=<non-public>``) with no
|
||||
routines at all: the runtime migrates only that one schema, so ``target_schema``
|
||||
is never falsy or ``public``, the gate never opens, and the maintenance loop
|
||||
logs, forever::
|
||||
|
||||
function public.banks_needing_consolidation() does not exist
|
||||
function public.schemas_with_expired_rows(...) does not exist
|
||||
|
||||
The revision is stamped applied, so redeploying the same version does not help
|
||||
(issue #2638; #2056 only fixed the ``public``/base-run case).
|
||||
|
||||
**The bug was the hardcoded literal, not the gating.** These routines are
|
||||
database-global — each enumerates ``pg_class`` across every schema and dispatches
|
||||
per schema — so exactly one copy should exist, and the maintenance loop calls the
|
||||
one in ``get_config().database_schema`` (see ``fq_routine``). The old gate
|
||||
installed into whichever schema was named ``public`` instead of whichever schema
|
||||
the deployment is actually configured to use. Comparing ``target_schema`` against
|
||||
the configured schema instead of the literal fixes #2638 at the source.
|
||||
|
||||
That also keeps the property the gate existed for: exactly one migration run
|
||||
satisfies the predicate, so concurrent per-schema runs never issue competing
|
||||
``CREATE OR REPLACE`` against the same ``pg_proc`` row and cannot hit
|
||||
``tuple concurrently updated``. No cross-process coordination is required — in
|
||||
particular no advisory lock, which is unusable here because Hindsight runs behind
|
||||
connection poolers and managed PG services (see #2817).
|
||||
|
||||
Runs targeting any *other* schema drop the routines from that schema rather than
|
||||
merely skipping. An earlier revision of this migration installed a copy into
|
||||
every schema it touched, which left one dead duplicate per tenant on any database
|
||||
that ran it; the drop makes the next migration pass clean those up instead of
|
||||
leaving them behind forever.
|
||||
|
||||
PostgreSQL only: the maintenance loop and worker poller are PG-only, so the
|
||||
Oracle slot is intentionally absent (mirrors ``e5f6a7b8c9d0``).
|
||||
|
||||
Revision ID: b6d2f8a4c1e7
|
||||
Revises: a8c1e4f7b0d3
|
||||
Create Date: 2026-07-20
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
from hindsight_api.config import get_config
|
||||
|
||||
revision: str = "b6d2f8a4c1e7"
|
||||
down_revision: str | Sequence[str] | None = "a8c1e4f7b0d3"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _configured_schema() -> str:
|
||||
"""The one schema this deployment's routines live in and are called from."""
|
||||
return get_config().database_schema or "public"
|
||||
|
||||
|
||||
def _target_schema() -> str | None:
|
||||
return context.config.get_main_option("target_schema")
|
||||
|
||||
|
||||
def _is_install_run() -> bool:
|
||||
"""True for the single run that owns the routines.
|
||||
|
||||
The base run (no ``target_schema``) and the run targeting the configured
|
||||
schema are the same deployment-level run; every other target is a tenant
|
||||
schema that must not carry its own copy.
|
||||
"""
|
||||
target = _target_schema()
|
||||
return not target or target == _configured_schema()
|
||||
|
||||
|
||||
def _prefix(schema: str | None) -> str:
|
||||
"""Qualifier for ``schema``, or ``""`` to fall back to ``search_path``."""
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
if not _is_install_run():
|
||||
_drop_stray_copies()
|
||||
return
|
||||
schema = _prefix(_target_schema())
|
||||
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE OR REPLACE FUNCTION {schema}banks_needing_consolidation()
|
||||
RETURNS TABLE(schema_name text, bank_id text)
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
BEGIN
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = 'memory_units' AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
RETURN QUERY EXECUTE format($q$
|
||||
SELECT %1$L::text, m.bank_id
|
||||
FROM %1$I.memory_units m
|
||||
JOIN %1$I.banks b ON b.bank_id = m.bank_id
|
||||
WHERE m.consolidated_at IS NULL
|
||||
AND m.consolidation_failed_at IS NULL
|
||||
AND m.fact_type IN ('experience', 'world')
|
||||
AND COALESCE(b.config -> 'enable_auto_consolidation', 'true'::jsonb) <> 'false'::jsonb
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM %1$I.async_operations o
|
||||
WHERE o.bank_id = m.bank_id
|
||||
AND o.operation_type = 'consolidation'
|
||||
AND o.status IN ('pending', 'processing')
|
||||
)
|
||||
GROUP BY m.bank_id
|
||||
$q$, sch);
|
||||
EXCEPTION
|
||||
-- Schema or its tables vanished between the pg_class
|
||||
-- snapshot and this query (tenant dropped or migrating).
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE OR REPLACE FUNCTION {schema}schemas_with_expired_rows(
|
||||
p_table text, p_ts_col text, p_days int
|
||||
)
|
||||
RETURNS SETOF text
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
has_expired boolean;
|
||||
BEGIN
|
||||
IF p_days IS NULL OR p_days <= 0 THEN
|
||||
RETURN;
|
||||
END IF;
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = p_table AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
EXECUTE format(
|
||||
'SELECT EXISTS (SELECT 1 FROM %I.%I WHERE %I < NOW() - make_interval(days => $1))',
|
||||
sch, p_table, p_ts_col
|
||||
) INTO has_expired USING p_days;
|
||||
EXCEPTION
|
||||
-- Schema or its table vanished mid-scan; skip it.
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
IF has_expired THEN
|
||||
RETURN NEXT sch;
|
||||
END IF;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE OR REPLACE FUNCTION {schema}mental_models_with_cron()
|
||||
RETURNS TABLE(schema_name text, bank_id text, mental_model_id text,
|
||||
refresh_cron text, last_refreshed_at timestamptz)
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
BEGIN
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = 'mental_models' AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
RETURN QUERY EXECUTE format($q$
|
||||
SELECT %1$L::text, mm.bank_id::text, mm.id::text,
|
||||
mm.trigger->>'refresh_cron', mm.last_refreshed_at
|
||||
FROM %1$I.mental_models mm
|
||||
WHERE COALESCE(mm.trigger->>'refresh_cron', '') <> ''
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM %1$I.async_operations o
|
||||
WHERE o.bank_id = mm.bank_id
|
||||
AND o.operation_type = 'refresh_mental_model'
|
||||
AND o.status IN ('pending', 'processing')
|
||||
AND o.task_payload->>'mental_model_id' = mm.id::text
|
||||
)
|
||||
$q$, sch);
|
||||
EXCEPTION
|
||||
-- Schema or its tables vanished between the pg_class
|
||||
-- snapshot and this query (tenant dropped or migrating).
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _drop_routines(schema: str | None) -> None:
|
||||
prefix = _prefix(schema)
|
||||
op.execute(f"DROP FUNCTION IF EXISTS {prefix}mental_models_with_cron()")
|
||||
op.execute(f"DROP FUNCTION IF EXISTS {prefix}schemas_with_expired_rows(text, text, int)")
|
||||
op.execute(f"DROP FUNCTION IF EXISTS {prefix}banks_needing_consolidation()")
|
||||
|
||||
|
||||
def _drop_stray_copies() -> None:
|
||||
"""Remove per-tenant duplicates left by the first cut of this migration.
|
||||
|
||||
That version installed a copy into every schema it touched, so a database
|
||||
that ran it carries one dead duplicate per tenant — only the copy in the
|
||||
configured schema is ever called. Dropping here means the next migration pass
|
||||
cleans them up; without it they would persist for the life of the database.
|
||||
|
||||
Safe on a database that never had them: ``DROP FUNCTION IF EXISTS`` is a
|
||||
no-op, and this branch never runs for the configured schema.
|
||||
"""
|
||||
_drop_routines(_target_schema())
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
# Only drop what this migration uniquely owns. When the configured schema is
|
||||
# ``public`` the copies there belong to e5f6a7b8c9d0 / f4d1c2b3a5e6, which are
|
||||
# still applied at this point and drop them on their own downgrade — removing
|
||||
# them here would strand those migrations without the functions they claim to
|
||||
# have installed.
|
||||
if not _is_install_run() or _configured_schema() == "public":
|
||||
return
|
||||
_drop_routines(_target_schema())
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
+144
@@ -0,0 +1,144 @@
|
||||
"""Backfill search_vector for native-backend observations.
|
||||
|
||||
Observations created or updated by the consolidator landed with a NULL
|
||||
``search_vector`` under the ``native`` text-search backend: the
|
||||
single-row INSERT/UPDATE paths in ``consolidator.py`` never populated the
|
||||
tsvector (only the batch raw-fact path in ``ops_postgresql.insert_facts_batch``
|
||||
did). Those observations were therefore invisible to the BM25 retrieval arm
|
||||
until they were re-written by a later consolidation pass. The writer is fixed
|
||||
in the same change set (all four consolidator sites now call
|
||||
``to_tsvector($lang, COALESCE(text, ''))``); this migration repairs the
|
||||
historical residue so existing observations become BM25-searchable without a
|
||||
re-ingest.
|
||||
|
||||
Scope mirrors the writer fix exactly:
|
||||
* Only the ``native`` backend is touched. The gate is the column *type*:
|
||||
under ``native`` ``search_vector`` is a regular (non-generated) tsvector
|
||||
column; under ``vchord`` it is a ``bm25vector`` and under
|
||||
``pg_textsearch`` / ``pgroonga`` / ``pg_search`` it is a dummy ``text``
|
||||
column. ``_is_regular_tsvector`` is true only for ``native``, so every
|
||||
other backend is a no-op.
|
||||
* The tsvector is built from the observation's own ``text`` only — matching
|
||||
the consolidator INSERT/UPDATE paths (entity / source / temporal signals
|
||||
are intentionally excluded; the other retrieval arms cover those).
|
||||
* Only ``fact_type = 'observation'`` rows with a NULL ``search_vector`` are
|
||||
rewritten. Raw facts already carry a populated tsvector, and the
|
||||
``IS NULL`` predicate makes the migration idempotent and re-runnable.
|
||||
|
||||
The configured ``HINDSIGHT_API_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE`` is used
|
||||
so backfilled rows are lexically identical to newly-created observations. The
|
||||
value is validated as a PG identifier (mirroring
|
||||
``HindsightConfig.validate``) before being embedded as a SQL literal.
|
||||
|
||||
This is a single UPDATE per schema: it locks the targeted observation rows for
|
||||
its duration. It is one-time and only touches unpopulated rows, so subsequent
|
||||
online writes (which now carry the tsvector via the writer fix) are unaffected.
|
||||
|
||||
Oracle slot is intentionally absent: the consolidator INSERT/UPDATE paths that
|
||||
this repairs are PostgreSQL-specific (``ops_postgresql``), and the native
|
||||
tsvector ``search_vector`` column only exists on PostgreSQL. There is no Oracle
|
||||
residue to repair.
|
||||
|
||||
Revision ID: c3f7a1b9d2e4
|
||||
Revises: f4d1c2b3a5e6
|
||||
Create Date: 2026-06-29
|
||||
"""
|
||||
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
from sqlalchemy import Connection, text
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
from hindsight_api.config import (
|
||||
DEFAULT_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE,
|
||||
ENV_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE,
|
||||
)
|
||||
|
||||
revision: str = "c3f7a1b9d2e4"
|
||||
down_revision: str | Sequence[str] | None = "f4d1c2b3a5e6"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
# Matches HindsightConfig.validate(): a tsvector regconfig name embedded as a
|
||||
# SQL literal must be a bare PG identifier.
|
||||
_PG_IDENTIFIER = re.compile(r"[a-zA-Z_][a-zA-Z0-9_]*")
|
||||
|
||||
|
||||
def _schema_prefix() -> str:
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _schema_name() -> str:
|
||||
return (context.config.get_main_option("target_schema") or "public").strip('"')
|
||||
|
||||
|
||||
def _native_language() -> str:
|
||||
"""Configured native tsvector language, validated as a PG identifier."""
|
||||
lang = os.getenv(
|
||||
ENV_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE,
|
||||
DEFAULT_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE,
|
||||
)
|
||||
if not _PG_IDENTIFIER.fullmatch(lang):
|
||||
return DEFAULT_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE
|
||||
return lang
|
||||
|
||||
|
||||
def _is_regular_tsvector(conn: Connection, schema: str, table: str) -> bool:
|
||||
"""True iff ``schema.table.search_vector`` is a non-generated tsvector column.
|
||||
|
||||
This is the ``native`` backend signature. ``vchord`` (bm25vector) and
|
||||
``pg_textsearch`` / ``pgroonga`` / ``pg_search`` (dummy text column) all
|
||||
fail this check, so the backfill is a no-op for them.
|
||||
"""
|
||||
row = conn.execute(
|
||||
text(
|
||||
"""
|
||||
SELECT is_generated, udt_name
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = :schema
|
||||
AND table_name = :table
|
||||
AND column_name = 'search_vector'
|
||||
"""
|
||||
),
|
||||
{"schema": schema, "table": table},
|
||||
).fetchone()
|
||||
if not row:
|
||||
return False
|
||||
is_generated, udt_name = row[0], row[1]
|
||||
return udt_name == "tsvector" and is_generated != "ALWAYS"
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
schema_name = _schema_name()
|
||||
if not _is_regular_tsvector(conn, schema_name, "memory_units"):
|
||||
# Non-native backend (or column absent) — nothing to backfill.
|
||||
return
|
||||
schema_prefix = _schema_prefix()
|
||||
lang = _native_language()
|
||||
op.execute(
|
||||
f"""
|
||||
UPDATE {schema_prefix}memory_units
|
||||
SET search_vector = to_tsvector('{lang}'::regconfig, COALESCE(text, ''))
|
||||
WHERE fact_type = 'observation' AND search_vector IS NULL
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
# No-op: backfilled rows are indistinguishable from observations that were
|
||||
# populated by the post-fix writer, and reverting either to NULL would
|
||||
# re-break BM25 retrieval. The column simply stays populated.
|
||||
pass
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
+158
@@ -0,0 +1,158 @@
|
||||
"""Make maintenance routines resilient to schemas that vanish mid-scan.
|
||||
|
||||
``public.banks_needing_consolidation()`` and
|
||||
``public.schemas_with_expired_rows(...)`` snapshot the set of schemas owning a
|
||||
target table from ``pg_class`` and then run a dynamic query against each schema
|
||||
in turn. That is a time-of-check/time-of-use race: a schema (or its tables) can
|
||||
be dropped — a tenant being deleted, or a tenant migration that recreates
|
||||
tables — between the snapshot and the per-schema query, which then aborts the
|
||||
whole routine with::
|
||||
|
||||
relation "<schema>.memory_units" does not exist
|
||||
relation "<schema>.audit_log" does not exist
|
||||
|
||||
In the test suite this surfaces as cross-worker contamination: the multi-tenant
|
||||
maintenance test creates and drops ~100 ``mt<hash>_NNN`` schemas while
|
||||
``test_maintenance_routines`` (on another xdist worker, same DB) calls the
|
||||
routines. In production the background maintenance loop hits the same race when
|
||||
a tenant is removed or mid-migration.
|
||||
|
||||
Wrap each per-schema query in its own ``BEGIN ... EXCEPTION`` block so a schema
|
||||
that disappears (``undefined_table`` / ``invalid_schema_name`` /
|
||||
``undefined_column``) is skipped instead of aborting the scan. The routines stay
|
||||
``CREATE OR REPLACE`` and PostgreSQL-only, and are (re)installed only on the run
|
||||
that targets the shared ``public`` schema — same gating as the original
|
||||
install (``e5f6a7b8c9d0``) and its repair (``b2d4f6a8c1e3``).
|
||||
|
||||
Revision ID: c7e9f1a3b5d2
|
||||
Revises: e1f2a3b4c5d6
|
||||
Create Date: 2026-06-19
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "c7e9f1a3b5d2"
|
||||
down_revision: str | Sequence[str] | None = "e1f2a3b4c5d6"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _should_install_public_routines(target_schema: str | None) -> bool:
|
||||
"""True for the run that must (re)create the shared ``public.*`` routines.
|
||||
|
||||
The routines physically live in ``public``, so they are installed exactly
|
||||
once — on the base run (no ``target_schema``) or the run that explicitly
|
||||
targets ``public``. Mirrors ``b2d4f6a8c1e3``.
|
||||
"""
|
||||
return not target_schema or target_schema == "public"
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
if not _should_install_public_routines(context.config.get_main_option("target_schema")):
|
||||
return
|
||||
|
||||
# Same body as b2d4f6a8c1e3, but each per-schema query runs in its own
|
||||
# subtransaction so a schema dropped mid-scan is skipped, not fatal.
|
||||
op.execute(
|
||||
"""
|
||||
CREATE OR REPLACE FUNCTION public.banks_needing_consolidation()
|
||||
RETURNS TABLE(schema_name text, bank_id text)
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
BEGIN
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = 'memory_units' AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
RETURN QUERY EXECUTE format($q$
|
||||
SELECT %1$L::text, m.bank_id
|
||||
FROM %1$I.memory_units m
|
||||
JOIN %1$I.banks b ON b.bank_id = m.bank_id
|
||||
WHERE m.consolidated_at IS NULL
|
||||
AND m.consolidation_failed_at IS NULL
|
||||
AND m.fact_type IN ('experience', 'world')
|
||||
AND COALESCE(b.config -> 'enable_auto_consolidation', 'true'::jsonb) <> 'false'::jsonb
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM %1$I.async_operations o
|
||||
WHERE o.bank_id = m.bank_id
|
||||
AND o.operation_type = 'consolidation'
|
||||
AND o.status IN ('pending', 'processing')
|
||||
)
|
||||
GROUP BY m.bank_id
|
||||
$q$, sch);
|
||||
EXCEPTION
|
||||
-- Schema or its tables vanished between the pg_class
|
||||
-- snapshot and this query (tenant dropped or migrating).
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
op.execute(
|
||||
"""
|
||||
CREATE OR REPLACE FUNCTION public.schemas_with_expired_rows(
|
||||
p_table text, p_ts_col text, p_days int
|
||||
)
|
||||
RETURNS SETOF text
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
has_expired boolean;
|
||||
BEGIN
|
||||
IF p_days IS NULL OR p_days <= 0 THEN
|
||||
RETURN;
|
||||
END IF;
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = p_table AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
EXECUTE format(
|
||||
'SELECT EXISTS (SELECT 1 FROM %I.%I WHERE %I < NOW() - make_interval(days => $1))',
|
||||
sch, p_table, p_ts_col
|
||||
) INTO has_expired USING p_days;
|
||||
EXCEPTION
|
||||
-- Schema or its table vanished mid-scan; skip it.
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
IF has_expired THEN
|
||||
RETURN NEXT sch;
|
||||
END IF;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
# No-op: e5f6a7b8c9d0 owns these functions' lifecycle and drops them on its
|
||||
# own downgrade. This migration only re-installs them (the resilient body is
|
||||
# a strict superset of the previous behaviour), so there is nothing to undo
|
||||
# without racing that migration's DROP.
|
||||
pass
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
+150
@@ -0,0 +1,150 @@
|
||||
"""Add the ``schemas_with_expired_operations`` cross-tenant discovery routine.
|
||||
|
||||
The worker's terminal-operation cleanup (``a8c1e4f7b0d3``) opens a connection
|
||||
and a prune transaction against *every* tenant schema on every cleanup cycle,
|
||||
whether or not that tenant has anything to prune. At thousands of tenants that
|
||||
is a per-cycle query storm whose cost is paid entirely by idle schemas.
|
||||
|
||||
This is the same problem ``public.schemas_with_expired_rows`` already solves for
|
||||
the ``audit_log`` / ``llm_requests`` retention sweeps (``e5f6a7b8c9d0``): one
|
||||
round-trip returns just the schemas that actually hold expired rows, and the
|
||||
caller then does real work only there. ``async_operations`` needs its own
|
||||
routine rather than reusing that one because eligibility is not "row older than
|
||||
N days" — pending and processing rows are never prunable, so the status filter
|
||||
has to be part of the predicate.
|
||||
|
||||
Install policy mirrors ``b6d2f8a4c1e7`` (#2638/#2824), the current behaviour for
|
||||
the sibling routines: the routine is database-global — it enumerates ``pg_class``
|
||||
across every schema and dispatches per schema — so exactly one copy should exist,
|
||||
installed into the schema this deployment is *configured* to use and called from
|
||||
there via ``fq_routine``. Gating on the literal ``"public"`` instead of the
|
||||
configured schema is what left single-tenant deployments in a dedicated
|
||||
non-``public`` schema without the routine (#2638).
|
||||
|
||||
Exactly one migration run satisfies that predicate, so concurrent per-schema runs
|
||||
never issue competing ``CREATE OR REPLACE`` against the same ``pg_proc`` row and
|
||||
cannot hit ``tuple concurrently updated``. No cross-process coordination is
|
||||
required — in particular no advisory lock, which is unusable here because
|
||||
Hindsight runs behind connection poolers and managed PG services (see #2817).
|
||||
|
||||
Each per-schema probe runs in its own ``BEGIN ... EXCEPTION`` block so a tenant
|
||||
dropped mid-scan is skipped instead of aborting the sweep (see ``c7e9f1a3b5d2``).
|
||||
|
||||
Revision ID: d7b2f8a1c934
|
||||
Revises: b6d2f8a4c1e7
|
||||
Create Date: 2026-07-20
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
from hindsight_api.config import get_config
|
||||
|
||||
revision: str = "d7b2f8a1c934"
|
||||
down_revision: str | Sequence[str] | None = "b6d2f8a4c1e7"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _configured_schema() -> str:
|
||||
"""The one schema this deployment's routines live in and are called from."""
|
||||
return get_config().database_schema or "public"
|
||||
|
||||
|
||||
def _target_schema() -> str | None:
|
||||
return context.config.get_main_option("target_schema")
|
||||
|
||||
|
||||
def _is_install_run() -> bool:
|
||||
"""True for the single run that owns the routine (mirrors b6d2f8a4c1e7)."""
|
||||
target = _target_schema()
|
||||
return not target or target == _configured_schema()
|
||||
|
||||
|
||||
def _prefix(schema: str | None) -> str:
|
||||
"""Qualifier for ``schema``, or ``""`` to fall back to ``search_path``."""
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _drop_routine(schema: str | None) -> None:
|
||||
op.execute(f"DROP FUNCTION IF EXISTS {_prefix(schema)}schemas_with_expired_operations(int)")
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
if not _is_install_run():
|
||||
# Tenant schemas must not carry their own copy: the routine is
|
||||
# database-global and only the configured schema's copy is ever called.
|
||||
# Dropping (rather than skipping) also cleans up after any interim build
|
||||
# of this branch that installed per-schema copies.
|
||||
_drop_routine(_target_schema())
|
||||
return
|
||||
schema = _prefix(_target_schema())
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE OR REPLACE FUNCTION {schema}schemas_with_expired_operations(p_days int)
|
||||
RETURNS SETOF text
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
has_expired boolean;
|
||||
BEGIN
|
||||
-- Zero (or negative) retention means "keep forever": report nothing
|
||||
-- so the caller skips the sweep entirely.
|
||||
IF p_days IS NULL OR p_days <= 0 THEN
|
||||
RETURN;
|
||||
END IF;
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = 'async_operations' AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
-- Matches the worker's prune predicate: only terminal rows
|
||||
-- are eligible, so a schema holding nothing but pending or
|
||||
-- processing work is correctly reported as having nothing
|
||||
-- to prune. Uses idx_async_operations_terminal_cleanup.
|
||||
EXECUTE format(
|
||||
'SELECT EXISTS ('
|
||||
' SELECT 1 FROM %I.async_operations'
|
||||
' WHERE status IN (''completed'', ''failed'', ''cancelled'')'
|
||||
' AND updated_at < NOW() - make_interval(days => $1)'
|
||||
')',
|
||||
sch
|
||||
) INTO has_expired USING p_days;
|
||||
EXCEPTION
|
||||
-- Schema or its table vanished between the pg_class
|
||||
-- snapshot and this probe (tenant dropped or migrating).
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
IF has_expired THEN
|
||||
RETURN NEXT sch;
|
||||
END IF;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
# This migration is the sole creator of this routine — no older migration
|
||||
# owns a copy the way e5f6a7b8c9d0 owns the public sibling routines — so the
|
||||
# install run's own copy is always ours to drop.
|
||||
if not _is_install_run():
|
||||
return
|
||||
_drop_routine(_target_schema())
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Oracle slot intentionally absent: this mirrors the PostgreSQL-only
|
||||
# maintenance routines, and the Oracle worker keeps its per-schema sweep.
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
+96
@@ -0,0 +1,96 @@
|
||||
"""Drop the search_vector column from the curation archive (invalidated_memory_units).
|
||||
|
||||
The archive is cold storage, never a recall surface, and carries no text-search
|
||||
index. Like ``embedding`` (dropped in d4f6a8c2e1b3), ``search_vector`` is a
|
||||
recall-surface column whose type follows the configured text-search backend, so
|
||||
it has no business living on the archive. Earlier curation code copied the live
|
||||
row's ``search_vector`` into ``invalidated_memory_units`` on invalidate; the
|
||||
engine now leaves it out on invalidate and recomputes it on revert, so the
|
||||
column is dead weight.
|
||||
|
||||
Dropping it removes a latent failure mode (#2503): under a non-native backend
|
||||
(pgroonga / pg_textsearch / pg_search / vchord) ``ensure_text_search_extension``
|
||||
reconciles ``memory_units.search_vector`` to ``text`` / ``bm25vector`` but never
|
||||
touched the archive, which the ``LIKE memory_units`` clone (c9a1b2d3e4f5) created
|
||||
as ``tsvector``. The type mismatch then broke the curation INSERT … SELECT
|
||||
round-trip:
|
||||
|
||||
column "search_vector" is of type tsvector but expression is of type text
|
||||
|
||||
With no column at all, there is nothing to mismatch. Unlike ``embedding`` (whose
|
||||
creation sites already omit it), the ``LIKE`` clone still adds ``search_vector``,
|
||||
so this migration does real work on both fresh and existing PostgreSQL databases.
|
||||
|
||||
DROP COLUMN is a metadata-only operation on both PostgreSQL and Oracle 23ai (no
|
||||
table rewrite), so it is cheap even across many tenant schemas. The downgrade
|
||||
re-adds an empty ``tsvector`` column (its original creation type).
|
||||
|
||||
Revision ID: e7c3a9f1b2d5
|
||||
Revises: b57a7c9e0d13
|
||||
Create Date: 2026-07-02
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "e7c3a9f1b2d5"
|
||||
down_revision: str | Sequence[str] | None = "b57a7c9e0d13"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _pg_schema_prefix() -> str:
|
||||
"""Schema-qualifier for raw SQL on PG (multi-tenant search_path)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
op.execute(f"ALTER TABLE {schema}invalidated_memory_units DROP COLUMN IF EXISTS search_vector")
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
# Re-add as the original tsvector creation type; comes back empty regardless.
|
||||
op.execute(f"ALTER TABLE {schema}invalidated_memory_units ADD COLUMN IF NOT EXISTS search_vector tsvector")
|
||||
|
||||
|
||||
def _oracle_upgrade() -> None:
|
||||
# Oracle has no `DROP COLUMN IF EXISTS`; swallow ORA-00904 (column does not
|
||||
# exist) so the migration is idempotent and safe on a schema whose baseline
|
||||
# may already omit the column.
|
||||
op.execute(
|
||||
"""
|
||||
BEGIN
|
||||
EXECUTE IMMEDIATE 'ALTER TABLE invalidated_memory_units DROP COLUMN search_vector';
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
IF SQLCODE != -904 THEN RAISE; END IF;
|
||||
END;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _oracle_downgrade() -> None:
|
||||
# Swallow ORA-01430 (column already exists) for idempotency. Oracle stores
|
||||
# search_vector as CLOB (see the Oracle baseline), so re-add it as CLOB.
|
||||
op.execute(
|
||||
"""
|
||||
BEGIN
|
||||
EXECUTE IMMEDIATE 'ALTER TABLE invalidated_memory_units ADD (search_vector CLOB)';
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
IF SQLCODE != -1430 THEN RAISE; END IF;
|
||||
END;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade, oracle=_oracle_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade, oracle=_oracle_downgrade)
|
||||
+110
@@ -0,0 +1,110 @@
|
||||
"""Add server-side routine for cron-scheduled mental model refresh.
|
||||
|
||||
Installs ``public.mental_models_with_cron()`` — a discovery routine that returns
|
||||
every mental model carrying a non-empty ``trigger->>'refresh_cron'`` across all
|
||||
tenant schemas in one round-trip (the same per-schema scan as the other
|
||||
maintenance routines from ``e5f6a7b8c9d0``). The maintenance loop evaluates each
|
||||
candidate's cron expression in Python (``croniter``) against ``last_refreshed_at``
|
||||
to decide whether a scheduled refresh is due — cron arithmetic isn't expressible
|
||||
in plain SQL — and only the cron *candidate set* is discovered here.
|
||||
|
||||
Models that already have a ``refresh_mental_model`` operation pending/processing
|
||||
are excluded so a slow refresh isn't double-queued (mirrors the in-flight guard
|
||||
in ``banks_needing_consolidation``). Each per-schema query runs in its own
|
||||
``BEGIN ... EXCEPTION`` subtransaction so a schema dropped mid-scan (tenant
|
||||
deletion / migration) is skipped, not fatal — same resilience as
|
||||
``c7e9f1a3b5d2``.
|
||||
|
||||
Read-only (STABLE) discovery routine — the caller performs the refresh enqueue —
|
||||
so installing it never mutates data. PostgreSQL only: the worker poller and the
|
||||
maintenance loop are PG-only (Oracle slot intentionally absent, mirroring
|
||||
``e5f6a7b8c9d0``). The routine lives in ``public`` and is CREATE OR REPLACE, so
|
||||
it is installed exactly once (base / ``public`` run) to avoid the
|
||||
``tuple concurrently updated`` race on concurrent per-tenant runs.
|
||||
|
||||
Revision ID: f4d1c2b3a5e6
|
||||
Revises: c7e9f1a3b5d2
|
||||
Create Date: 2026-06-23
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "f4d1c2b3a5e6"
|
||||
down_revision: str | Sequence[str] | None = "c7e9f1a3b5d2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _should_install_public_routines(target_schema: str | None) -> bool:
|
||||
"""True for the run that must (re)create the shared ``public.*`` routine.
|
||||
|
||||
The routine physically lives in ``public``, so it is installed exactly once —
|
||||
on the base run (no ``target_schema``) or the run that explicitly targets
|
||||
``public``. Mirrors ``c7e9f1a3b5d2``.
|
||||
"""
|
||||
return not target_schema or target_schema == "public"
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
if not _should_install_public_routines(context.config.get_main_option("target_schema")):
|
||||
return
|
||||
|
||||
op.execute(
|
||||
"""
|
||||
CREATE OR REPLACE FUNCTION public.mental_models_with_cron()
|
||||
RETURNS TABLE(schema_name text, bank_id text, mental_model_id text,
|
||||
refresh_cron text, last_refreshed_at timestamptz)
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
BEGIN
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = 'mental_models' AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
RETURN QUERY EXECUTE format($q$
|
||||
SELECT %1$L::text, mm.bank_id::text, mm.id::text,
|
||||
mm.trigger->>'refresh_cron', mm.last_refreshed_at
|
||||
FROM %1$I.mental_models mm
|
||||
WHERE COALESCE(mm.trigger->>'refresh_cron', '') <> ''
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM %1$I.async_operations o
|
||||
WHERE o.bank_id = mm.bank_id
|
||||
AND o.operation_type = 'refresh_mental_model'
|
||||
AND o.status IN ('pending', 'processing')
|
||||
AND o.task_payload->>'mental_model_id' = mm.id::text
|
||||
)
|
||||
$q$, sch);
|
||||
EXCEPTION
|
||||
-- Schema or its tables vanished between the pg_class
|
||||
-- snapshot and this query (tenant dropped or migrating).
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
if not _should_install_public_routines(context.config.get_main_option("target_schema")):
|
||||
return
|
||||
op.execute("DROP FUNCTION IF EXISTS public.mental_models_with_cron()")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
@@ -27,7 +27,7 @@ from hindsight_api.engine.audit import (
|
||||
AuditLogStatsResponse,
|
||||
)
|
||||
from hindsight_api.engine.llm_trace import LLMRequestListResponse, LLMRequestStatsResponse
|
||||
from hindsight_api.extensions import AuthenticationError
|
||||
from hindsight_api.extensions import AuthenticationError, PrecheckOperation
|
||||
|
||||
|
||||
def _parse_metadata(metadata: Any) -> dict[str, Any]:
|
||||
@@ -52,6 +52,7 @@ from fastapi.routing import APIRoute
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
from hindsight_api.config import RETAIN_EXTRACTION_MODES
|
||||
|
||||
|
||||
def _annotation_is_nullable(annotation: Any) -> bool:
|
||||
@@ -148,17 +149,29 @@ def FieldWithDefault(default_factory: Callable, **kwargs) -> Any:
|
||||
|
||||
|
||||
from hindsight_api.config import get_config
|
||||
from hindsight_api.engine.memory_engine import Budget, _current_schema, _get_tiktoken_encoding
|
||||
from hindsight_api.engine.memory_engine import (
|
||||
Budget,
|
||||
RetainOperationConflictError,
|
||||
_current_schema,
|
||||
_get_tiktoken_encoding,
|
||||
)
|
||||
from hindsight_api.engine.providers.none_llm import LLMNotAvailableError
|
||||
from hindsight_api.engine.response_models import (
|
||||
VALID_RECALL_FACT_TYPES,
|
||||
DryRunExtractionResult,
|
||||
MemoryFact,
|
||||
MinScores,
|
||||
RecallScores,
|
||||
TokenUsage,
|
||||
)
|
||||
from hindsight_api.engine.search.tags import TagGroup, TagsMatch
|
||||
from hindsight_api.extensions import HttpExtension, OperationValidationError, load_extension
|
||||
from hindsight_api.metrics import create_metrics_collector, get_metrics_collector, initialize_metrics
|
||||
from hindsight_api.metrics import (
|
||||
create_metrics_collector,
|
||||
get_metrics_collector,
|
||||
initialize_metrics,
|
||||
normalize_http_endpoint,
|
||||
)
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -265,6 +278,16 @@ class RecallRequest(BaseModel):
|
||||
default=None,
|
||||
description="List of fact types to recall: 'world', 'experience', 'observation'. Defaults to world and experience if not specified.",
|
||||
)
|
||||
prefer_observations: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"When recalling raw facts ('world'/'experience') together with 'observation', drop any raw "
|
||||
"fact that an observation in the results was consolidated from, so the observation supersedes "
|
||||
"it and you don't get duplicate content. The freed slots are backfilled with the next results, "
|
||||
"keeping the result count at the requested budget. Disabled by default; set to true to enable. "
|
||||
"No effect unless 'observation' and at least one raw type are both requested."
|
||||
),
|
||||
)
|
||||
budget: Budget = Budget.MID
|
||||
max_tokens: int = 4096
|
||||
trace: bool = False
|
||||
@@ -281,18 +304,31 @@ class RecallRequest(BaseModel):
|
||||
)
|
||||
tags: list[str] | None = Field(
|
||||
default=None,
|
||||
description="Filter memories by tags. If not specified, all memories are returned.",
|
||||
description="Filter memories by tags. If not specified, all memories are returned. "
|
||||
"Omitting tags (or passing []) together with tags_match='exact' filters to "
|
||||
"untagged/global observations only (the scope written by observation_scopes='shared').",
|
||||
)
|
||||
tags_match: TagsMatch = Field(
|
||||
default="any",
|
||||
description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), "
|
||||
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
|
||||
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged), "
|
||||
"'exact' (set-equality on the full scope, excludes untagged). With 'exact' and no tags "
|
||||
"(or []), the empty global scope is selected and only untagged memories match.",
|
||||
)
|
||||
tag_groups: list[TagGroup] | None = Field(
|
||||
default=None,
|
||||
description="Compound tag filter using boolean groups. Groups in the list are AND-ed. "
|
||||
"Each group is a leaf {tags, match} or compound {and: [...]}, {or: [...]}, {not: ...}.",
|
||||
)
|
||||
min_scores: MinScores | None = Field(
|
||||
default=None,
|
||||
description="Optional per-stage score floors (all inclusive, AND-ed). `semantic` and `keyword` are "
|
||||
"retrieval-level cutoffs pushed into the SQL arms (overriding the global similarity/BM25 minimums for "
|
||||
"this request); `reranker` and `final` are post-ranking filters on the scored results. Any field left "
|
||||
"unset imposes no floor; omitting `min_scores` entirely (the default) applies no score filtering. Use "
|
||||
"with care — the reranker's absolute scores are not calibrated across queries (a clearly-relevant match "
|
||||
"may score ~0.001 even though it is ranked first).",
|
||||
)
|
||||
|
||||
@field_validator("query")
|
||||
@classmethod
|
||||
@@ -348,6 +384,7 @@ class RecallResult(BaseModel):
|
||||
source_fact_ids: list[str] | None = (
|
||||
None # IDs of source facts (observation type only, when source_facts is enabled)
|
||||
)
|
||||
scores: RecallScores | None = None # Per-stage recall scores (final/reranker/semantic/text)
|
||||
|
||||
|
||||
class EntityObservationResponse(BaseModel):
|
||||
@@ -690,6 +727,25 @@ class RetainRequest(BaseModel):
|
||||
description="Deprecated. Use item-level tags instead.",
|
||||
deprecated=True,
|
||||
)
|
||||
operation_id: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Optional client-supplied UUID used as the identity of an async retain operation. "
|
||||
"Re-submitting with the same operation_id returns the original operation and creates no new "
|
||||
"work, so retrying after a lost or timed-out acknowledgement will not enqueue a duplicate. "
|
||||
"Reusing an id that belongs to a different operation returns HTTP 409. Ignored for synchronous retain."
|
||||
),
|
||||
)
|
||||
|
||||
@field_validator("operation_id")
|
||||
@classmethod
|
||||
def validate_operation_id(cls, value: str | None) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return str(uuid.UUID(value))
|
||||
except (ValueError, AttributeError, TypeError) as exc:
|
||||
raise ValueError("operation_id must be a valid UUID") from exc
|
||||
|
||||
|
||||
class FileRetainMetadata(BaseModel):
|
||||
@@ -1214,7 +1270,7 @@ class CreateBankRequest(BaseModel):
|
||||
)
|
||||
retain_extraction_mode: str | None = Field(
|
||||
default=None,
|
||||
description="Fact extraction mode: 'concise' (default), 'verbose', or 'custom'.",
|
||||
description="Fact extraction mode: 'concise' (default), 'verbose', 'custom', 'verbatim', or 'chunks'.",
|
||||
)
|
||||
retain_custom_instructions: str | None = Field(
|
||||
default=None,
|
||||
@@ -1402,6 +1458,7 @@ class ListMemoryUnitsResponse(BaseModel):
|
||||
"date": "2024-01-15T10:30:00Z",
|
||||
"type": "world",
|
||||
"entities": "Alice (PERSON), Google (ORGANIZATION)",
|
||||
"metadata": {"source": "slack", "channel": "engineering"},
|
||||
}
|
||||
],
|
||||
"total": 150,
|
||||
@@ -1442,6 +1499,13 @@ class DryRunExtractRequest(BaseModel):
|
||||
entities_allow_free_form: bool | None = None
|
||||
llm_output_language: str | None = None
|
||||
|
||||
@field_validator("content")
|
||||
@classmethod
|
||||
def validate_content(cls, v: str) -> str:
|
||||
if not v.strip():
|
||||
raise ValueError("content cannot be empty")
|
||||
return v
|
||||
|
||||
|
||||
class ListDocumentsResponse(BaseModel):
|
||||
"""Response model for list documents endpoint."""
|
||||
@@ -1628,8 +1692,8 @@ class UpdateMemoryRequest(BaseModel):
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _require_an_edit(self) -> "UpdateMemoryRequest":
|
||||
if all(
|
||||
v is None
|
||||
has_value_edit = any(
|
||||
v is not None
|
||||
for v in (
|
||||
self.text,
|
||||
self.context,
|
||||
@@ -1639,7 +1703,9 @@ class UpdateMemoryRequest(BaseModel):
|
||||
self.entities,
|
||||
self.state,
|
||||
)
|
||||
):
|
||||
)
|
||||
has_date_clear = bool({"occurred_start", "occurred_end"} & self.model_fields_set)
|
||||
if not has_value_edit and not has_date_clear:
|
||||
raise ValueError("Provide at least one field to update.")
|
||||
if self.state is not None and self.state not in ("valid", "invalidated"):
|
||||
raise ValueError("state must be 'valid' or 'invalidated'.")
|
||||
@@ -1941,6 +2007,17 @@ class MentalModelTrigger(BaseModel):
|
||||
default=False,
|
||||
description="If true, refresh this mental model after observations consolidation (real-time mode)",
|
||||
)
|
||||
refresh_cron: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Cron expression (UTC, standard 5-field syntax, e.g. '0 3 * * *' for daily at 03:00 UTC) "
|
||||
"for refreshing this mental model on a fixed schedule. Mutually exclusive with "
|
||||
"refresh_after_consolidation — a model refreshes either after consolidation or on a cron "
|
||||
"schedule, not both. A scheduled refresh only runs when the model is stale (new memories in "
|
||||
"its scope since the last refresh); if nothing changed, the tick is skipped to avoid a "
|
||||
"wasted LLM call. null = no schedule."
|
||||
),
|
||||
)
|
||||
fact_types: list[Literal["world", "experience", "observation"]] | None = Field(
|
||||
default=None,
|
||||
description="Filter which fact types are retrieved during reflect. None means all types (world, experience, observation).",
|
||||
@@ -1999,6 +2076,31 @@ class MentalModelTrigger(BaseModel):
|
||||
raise ValueError("fact_types must not be empty. Use null to include all fact types.")
|
||||
return v
|
||||
|
||||
@field_validator("refresh_cron")
|
||||
@classmethod
|
||||
def validate_refresh_cron(cls, v: str | None) -> str | None:
|
||||
if v is None:
|
||||
return v
|
||||
v = v.strip()
|
||||
if not v:
|
||||
return None
|
||||
from croniter import croniter
|
||||
|
||||
if not croniter.is_valid(v):
|
||||
raise ValueError(f"refresh_cron is not a valid cron expression: {v!r}")
|
||||
return v
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_refresh_exclusivity(self) -> "MentalModelTrigger":
|
||||
# A mental model refreshes either after consolidation (real-time) or on a
|
||||
# cron schedule, never both — the two triggers would race and double-refresh.
|
||||
if self.refresh_after_consolidation and self.refresh_cron:
|
||||
raise ValueError(
|
||||
"refresh_after_consolidation and refresh_cron are mutually exclusive: "
|
||||
"a mental model refreshes either after consolidation or on a cron schedule, not both."
|
||||
)
|
||||
return self
|
||||
|
||||
|
||||
class MentalModelResponse(BaseModel):
|
||||
"""Response model for a mental model (stored reflect response)."""
|
||||
@@ -2129,7 +2231,8 @@ class BankTemplateConfig(BaseModel):
|
||||
reflect_mission: str | None = Field(default=None, description="Mission/context for Reflect operations")
|
||||
retain_mission: str | None = Field(default=None, description="Steers what gets extracted during retain")
|
||||
retain_extraction_mode: str | None = Field(
|
||||
default=None, description="Fact extraction mode: 'concise' (default), 'verbose', or 'custom'"
|
||||
default=None,
|
||||
description="Fact extraction mode: 'concise' (default), 'verbose', 'custom', 'verbatim', or 'chunks'",
|
||||
)
|
||||
retain_custom_instructions: str | None = Field(
|
||||
default=None, description="Custom extraction prompt (when mode='custom')"
|
||||
@@ -2218,6 +2321,14 @@ class BankTemplateConfig(BaseModel):
|
||||
recall_budget_max: int | None = Field(
|
||||
default=None, description="Ceiling for the adaptive function (after clamping)"
|
||||
)
|
||||
audit_log_enabled: bool | None = Field(
|
||||
default=None, description="Enable audit logging for this bank (overrides the server default)"
|
||||
)
|
||||
store_document_text: bool | None = Field(
|
||||
default=None,
|
||||
description="Persist raw source text (documents.original_text / chunks.chunk_text). "
|
||||
"Set false to keep only derived facts.",
|
||||
)
|
||||
|
||||
def get_config_updates(self) -> dict[str, Any]:
|
||||
"""Return only the fields that were explicitly set (non-None)."""
|
||||
@@ -2355,10 +2466,10 @@ def validate_bank_template(manifest: "BankTemplateManifest") -> list[str]:
|
||||
if manifest.bank:
|
||||
bank = manifest.bank
|
||||
if bank.retain_extraction_mode is not None:
|
||||
valid_modes = ("concise", "verbose", "custom", "chunks")
|
||||
if bank.retain_extraction_mode not in valid_modes:
|
||||
if bank.retain_extraction_mode not in RETAIN_EXTRACTION_MODES:
|
||||
errors.append(
|
||||
f"bank.retain_extraction_mode: must be one of {valid_modes}, got '{bank.retain_extraction_mode}'"
|
||||
"bank.retain_extraction_mode: "
|
||||
f"must be one of {RETAIN_EXTRACTION_MODES}, got '{bank.retain_extraction_mode}'"
|
||||
)
|
||||
if bank.retain_custom_instructions and bank.retain_extraction_mode != "custom":
|
||||
errors.append("bank.retain_custom_instructions: requires retain_extraction_mode='custom'")
|
||||
@@ -2534,6 +2645,10 @@ class OperationResponse(BaseModel):
|
||||
task_type: str
|
||||
items_count: int
|
||||
document_id: str | None = None
|
||||
filename: str | None = Field(
|
||||
default=None,
|
||||
description="Original filename for file-conversion operations (file_convert_retain); null for other task types.",
|
||||
)
|
||||
created_at: str
|
||||
updated_at: str | None = Field(
|
||||
default=None,
|
||||
@@ -2647,6 +2762,24 @@ class RetryOperationResponse(BaseModel):
|
||||
operation_id: str
|
||||
|
||||
|
||||
class DeleteOperationResponse(BaseModel):
|
||||
"""Response model for delete operation endpoint."""
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"success": True,
|
||||
"message": "Operation 550e8400-e29b-41d4-a716-446655440000 deleted",
|
||||
"operation_id": "550e8400-e29b-41d4-a716-446655440000",
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
success: bool
|
||||
message: str
|
||||
operation_id: str
|
||||
|
||||
|
||||
class ChildOperationStatus(BaseModel):
|
||||
"""Status of a child operation (for batch operations)."""
|
||||
|
||||
@@ -2738,7 +2871,7 @@ class FeaturesInfo(BaseModel):
|
||||
file_upload_api: bool = Field(description="Whether file upload/conversion API is enabled")
|
||||
document_export_api: bool = Field(description="Whether the document export endpoint is enabled")
|
||||
document_import_api: bool = Field(description="Whether the document import endpoint is enabled")
|
||||
audit_log: bool = Field(description="Whether audit logging is enabled")
|
||||
audit_log: bool = Field(description="Whether audit logging is enabled by default (overridable per bank)")
|
||||
llm_trace: bool = Field(description="Whether per-bank LLM request tracing is enabled")
|
||||
store_document_text: bool = Field(
|
||||
description="Whether raw source text is persisted. When false, document/chunk source text is not stored."
|
||||
@@ -2908,10 +3041,15 @@ def _make_audited_http(audit_logger_getter: Callable[[], AuditLogger | None]):
|
||||
@wraps(func)
|
||||
async def wrapper(*args, **kwargs):
|
||||
al = audit_logger_getter()
|
||||
if al is None or not al.is_enabled(action):
|
||||
# Cheap bank-independent pre-filter first, then the per-bank
|
||||
# decision (audit_log_enabled is overridable per bank).
|
||||
if al is None or not al.action_allowed(action):
|
||||
return await func(*args, **kwargs)
|
||||
|
||||
bank_id = kwargs.get("bank_id")
|
||||
if not await al.should_log(action, bank_id, kwargs.get("request_context")):
|
||||
return await func(*args, **kwargs)
|
||||
|
||||
started_at = _dt.now(_tz.utc)
|
||||
|
||||
req_data = None
|
||||
@@ -2996,6 +3134,7 @@ def create_app(
|
||||
config = get_config()
|
||||
poller = None
|
||||
poller_task = None
|
||||
loop_watchdog = None
|
||||
|
||||
# Initialize OpenTelemetry metrics
|
||||
try:
|
||||
@@ -3039,12 +3178,22 @@ def create_app(
|
||||
metrics_collector.set_db_pool(memory._pool)
|
||||
logging.info("DB pool metrics configured")
|
||||
|
||||
# Start the event-loop stall watchdog (logs the culprit stack if a task
|
||||
# blocks the loop, so a failing /health can be told apart from pool exhaustion).
|
||||
from ..loop_watchdog import start_loop_watchdog
|
||||
|
||||
loop_watchdog = start_loop_watchdog(asyncio.get_running_loop())
|
||||
|
||||
# Start worker poller if the backend supports it.
|
||||
# All current backends (PostgreSQL, Oracle) support async worker/poller.
|
||||
if config.worker_enabled and memory._backend.supports_worker_poller:
|
||||
from ..config import DEFAULT_DATABASE_SCHEMA
|
||||
from ..utils import warn_if_container_default_worker_id
|
||||
|
||||
warn_if_container_default_worker_id(config.worker_id)
|
||||
worker_id = config.worker_id or socket.gethostname()
|
||||
worker_id_source = "HINDSIGHT_API_WORKER_ID" if config.worker_id else "hostname (default)"
|
||||
logging.info(f"Worker id: {worker_id} (source: {worker_id_source})")
|
||||
# Convert default schema to None for SQL compatibility (no schema prefix)
|
||||
schema = None if config.database_schema == DEFAULT_DATABASE_SCHEMA else config.database_schema
|
||||
poller = WorkerPoller(
|
||||
@@ -3079,6 +3228,10 @@ def create_app(
|
||||
|
||||
yield
|
||||
|
||||
# Stop the loop watchdog first so it doesn't fire during teardown.
|
||||
if loop_watchdog is not None:
|
||||
loop_watchdog.stop()
|
||||
|
||||
# Shutdown worker poller if running
|
||||
if poller is not None:
|
||||
await poller.shutdown_graceful(timeout=30.0)
|
||||
@@ -3237,15 +3390,9 @@ def create_app(
|
||||
@app.middleware("http")
|
||||
async def http_metrics_middleware(request, call_next):
|
||||
"""Record HTTP request metrics."""
|
||||
# Normalize endpoint path to reduce cardinality
|
||||
# Replace UUIDs and numeric IDs with placeholders
|
||||
import re
|
||||
|
||||
path = request.url.path
|
||||
# Replace UUIDs
|
||||
path = re.sub(r"/[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}", "/{id}", path)
|
||||
# Replace numeric IDs
|
||||
path = re.sub(r"/\d+(?=/|$)", "/{id}", path)
|
||||
# Template id segments (bank ids, UUIDs, numeric ids) so the endpoint
|
||||
# metric label stays bounded-cardinality.
|
||||
path = normalize_http_endpoint(request.url.path)
|
||||
|
||||
status_code = [500] # Default to 500, will be updated
|
||||
metrics_collector = get_metrics_collector()
|
||||
@@ -3304,7 +3451,7 @@ def _register_routes(app: FastAPI):
|
||||
api_key = authorization.strip()
|
||||
return RequestContext(api_key=api_key)
|
||||
|
||||
def precheck_for(operation: str):
|
||||
def precheck_for(operation: PrecheckOperation):
|
||||
"""
|
||||
Build a FastAPI dependency that runs ``OperationValidator.precheck``.
|
||||
|
||||
@@ -3333,6 +3480,7 @@ def _register_routes(app: FastAPI):
|
||||
|
||||
async def _precheck_dep(
|
||||
bank_id: str,
|
||||
request: Request,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
) -> None:
|
||||
validator = getattr(app.state.memory, "_operation_validator", None)
|
||||
@@ -3341,10 +3489,20 @@ def _register_routes(app: FastAPI):
|
||||
from hindsight_api.extensions import PrecheckContext
|
||||
|
||||
await app.state.memory._authenticate_tenant(request_context)
|
||||
cl_header = request.headers.get("content-length")
|
||||
content_length: int | None = None
|
||||
if cl_header is not None:
|
||||
try:
|
||||
parsed = int(cl_header)
|
||||
except ValueError:
|
||||
parsed = -1
|
||||
if parsed >= 0:
|
||||
content_length = parsed
|
||||
ctx = PrecheckContext(
|
||||
operation=operation,
|
||||
bank_id=bank_id,
|
||||
request_context=request_context,
|
||||
content_length=content_length,
|
||||
)
|
||||
result = await validator.precheck(ctx)
|
||||
if not result.allowed:
|
||||
@@ -3399,8 +3557,8 @@ def _register_routes(app: FastAPI):
|
||||
Returns version info and feature flags that can be used by clients
|
||||
to determine which capabilities are available.
|
||||
|
||||
Note: observations flag shows the global default. Individual banks
|
||||
may override this setting via bank-specific configuration.
|
||||
Note: the observations and audit_log flags show the global default.
|
||||
Individual banks may override these via bank-specific configuration.
|
||||
"""
|
||||
from hindsight_api import __version__
|
||||
from hindsight_api.config import _get_raw_config
|
||||
@@ -3448,7 +3606,7 @@ def _register_routes(app: FastAPI):
|
||||
async def api_graph(
|
||||
bank_id: str,
|
||||
type: str | None = None,
|
||||
limit: int = 1000,
|
||||
limit: int = Query(default=1000, ge=0),
|
||||
q: str | None = None,
|
||||
tags: list[str] | None = Query(None),
|
||||
tags_match: str = "all_strict",
|
||||
@@ -3485,7 +3643,7 @@ def _register_routes(app: FastAPI):
|
||||
"/v1/default/banks/{bank_id}/memories/list",
|
||||
response_model=ListMemoryUnitsResponse,
|
||||
summary="List memory units",
|
||||
description="List memory units with pagination and optional full-text search. Supports filtering by type. Results are sorted by most recent first (mentioned_at DESC, then created_at DESC).",
|
||||
description="List memory units with pagination and optional full-text search. Supports filtering by type, source document, and linked entity ID. Results are sorted by most recent first (mentioned_at DESC, then created_at DESC).",
|
||||
operation_id="list_memories",
|
||||
tags=["Memory"],
|
||||
)
|
||||
@@ -3496,8 +3654,11 @@ def _register_routes(app: FastAPI):
|
||||
consolidation_state: str | None = None,
|
||||
state: str | None = None,
|
||||
document_id: str | None = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
entity_id: str | None = None,
|
||||
tags: list[str] | None = Query(default=None),
|
||||
tags_match: TagsMatch = Query(default="any"),
|
||||
limit: int = Query(default=100, ge=0),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""
|
||||
@@ -3512,6 +3673,14 @@ def _register_routes(app: FastAPI):
|
||||
q: Search query for full-text search (searches text and context)
|
||||
consolidation_state: Filter by consolidation state for source memories
|
||||
(world/experience). One of 'failed', 'pending', or 'done'.
|
||||
document_id: Filter to a single source document.
|
||||
entity_id: Filter to memory units linked to this entity ID (via stored
|
||||
entity links, not text/semantic match). Combining with
|
||||
state='invalidated' returns no results (the archive has no links).
|
||||
tags: Optional list of tag names to filter by.
|
||||
tags_match: How to combine tags: 'any' (OR, default) or 'all' (AND) both
|
||||
also include untagged memories; 'any_strict'/'all_strict' exclude
|
||||
untagged; 'exact' matches the tag set exactly.
|
||||
limit: Maximum number of results (default: 100)
|
||||
offset: Offset for pagination (default: 0)
|
||||
"""
|
||||
@@ -3523,6 +3692,9 @@ def _register_routes(app: FastAPI):
|
||||
consolidation_state=consolidation_state,
|
||||
state=state,
|
||||
document_id=document_id,
|
||||
entity_id=entity_id,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
request_context=request_context,
|
||||
@@ -3544,7 +3716,7 @@ def _register_routes(app: FastAPI):
|
||||
async def _require_dry_run_enabled() -> None:
|
||||
"""Feature-flag gate for dry-run extraction.
|
||||
|
||||
Declared as a dependency BEFORE ``precheck_for("dry_run_extract")`` so a
|
||||
Declared as a dependency BEFORE ``precheck_for(PrecheckOperation.DRY_RUN_EXTRACT)`` so a
|
||||
disabled route returns 404 regardless of tenant/billing state — FastAPI
|
||||
resolves path-operation dependencies in signature order, so this runs
|
||||
first and preserves the original "disabled → 404" contract.
|
||||
@@ -3574,7 +3746,7 @@ def _register_routes(app: FastAPI):
|
||||
body: DryRunExtractRequest,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
_enabled: None = Depends(_require_dry_run_enabled),
|
||||
_precheck: None = Depends(precheck_for("dry_run_extract")),
|
||||
_precheck: None = Depends(precheck_for(PrecheckOperation.DRY_RUN_EXTRACT)),
|
||||
):
|
||||
try:
|
||||
override_fields = (
|
||||
@@ -3655,6 +3827,7 @@ def _register_routes(app: FastAPI):
|
||||
operation_id="update_memory",
|
||||
tags=["Memory"],
|
||||
)
|
||||
@audited("update_memory")
|
||||
async def api_update_memory(
|
||||
bank_id: str,
|
||||
memory_id: str,
|
||||
@@ -3663,13 +3836,23 @@ def _register_routes(app: FastAPI):
|
||||
):
|
||||
"""Curate a single memory unit (edit text / invalidate / revert)."""
|
||||
try:
|
||||
occurred_start = (
|
||||
""
|
||||
if "occurred_start" in request.model_fields_set and request.occurred_start is None
|
||||
else request.occurred_start
|
||||
)
|
||||
occurred_end = (
|
||||
""
|
||||
if "occurred_end" in request.model_fields_set and request.occurred_end is None
|
||||
else request.occurred_end
|
||||
)
|
||||
data = await app.state.memory.update_memory_unit(
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
text=request.text,
|
||||
context=request.context,
|
||||
occurred_start=request.occurred_start,
|
||||
occurred_end=request.occurred_end,
|
||||
occurred_start=occurred_start,
|
||||
occurred_end=occurred_end,
|
||||
new_fact_type=request.fact_type,
|
||||
entities=request.entities,
|
||||
state=request.state,
|
||||
@@ -3742,7 +3925,7 @@ def _register_routes(app: FastAPI):
|
||||
request: RecallRequest,
|
||||
http_request: Request,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
_precheck: None = Depends(precheck_for("recall")),
|
||||
_precheck: None = Depends(precheck_for(PrecheckOperation.RECALL)),
|
||||
):
|
||||
"""Run a recall and return results with trace."""
|
||||
import time
|
||||
@@ -3809,6 +3992,7 @@ def _register_routes(app: FastAPI):
|
||||
max_tokens=request.max_tokens,
|
||||
enable_trace=request.trace,
|
||||
fact_type=fact_types,
|
||||
prefer_observations=request.prefer_observations,
|
||||
question_date=question_date,
|
||||
include_entities=include_entities,
|
||||
max_entity_tokens=max_entity_tokens,
|
||||
@@ -3821,6 +4005,7 @@ def _register_routes(app: FastAPI):
|
||||
tags=request.tags,
|
||||
tags_match=request.tags_match,
|
||||
tag_groups=request.tag_groups,
|
||||
min_scores=request.min_scores,
|
||||
),
|
||||
operation="recall",
|
||||
bank_id=bank_id,
|
||||
@@ -3842,6 +4027,7 @@ def _register_routes(app: FastAPI):
|
||||
chunk_id=fact.chunk_id,
|
||||
tags=fact.tags,
|
||||
source_fact_ids=fact.source_fact_ids,
|
||||
scores=fact.scores,
|
||||
)
|
||||
|
||||
recall_results = [_fact_to_result(fact) for fact in core_result.results]
|
||||
@@ -3943,7 +4129,7 @@ def _register_routes(app: FastAPI):
|
||||
request: ReflectRequest,
|
||||
http_request: Request,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
_precheck: None = Depends(precheck_for("reflect")),
|
||||
_precheck: None = Depends(precheck_for(PrecheckOperation.REFLECT)),
|
||||
):
|
||||
metrics = get_metrics_collector()
|
||||
|
||||
@@ -4101,11 +4287,17 @@ def _register_routes(app: FastAPI):
|
||||
)
|
||||
async def api_stats(
|
||||
bank_id: str,
|
||||
refresh: bool = Query(
|
||||
default=False,
|
||||
description="Force a fresh recompute, bypassing the cached value (and refreshing the cache).",
|
||||
),
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""Get statistics about memory nodes and links for a memory bank."""
|
||||
try:
|
||||
stats = await app.state.memory.get_bank_stats(bank_id, request_context=request_context)
|
||||
stats = await app.state.memory.get_bank_stats(
|
||||
bank_id, request_context=request_context, force_refresh=refresh
|
||||
)
|
||||
nodes_by_type = stats["node_counts"]
|
||||
links_by_type = stats["link_counts"]
|
||||
links_by_fact_type = stats["link_counts_by_fact_type"]
|
||||
@@ -4231,8 +4423,8 @@ def _register_routes(app: FastAPI):
|
||||
)
|
||||
async def api_list_entities(
|
||||
bank_id: str,
|
||||
limit: int = Query(default=100, description="Maximum number of entities to return"),
|
||||
offset: int = Query(default=0, description="Offset for pagination"),
|
||||
limit: int = Query(default=100, ge=0, description="Maximum number of entities to return"),
|
||||
offset: int = Query(default=0, ge=0, description="Offset for pagination"),
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""List entities for a memory bank with pagination."""
|
||||
@@ -4267,7 +4459,7 @@ def _register_routes(app: FastAPI):
|
||||
)
|
||||
async def api_entity_graph(
|
||||
bank_id: str,
|
||||
limit: int = Query(default=1000, description="Maximum number of co-occurrence edges to return"),
|
||||
limit: int = Query(default=1000, ge=0, description="Maximum number of co-occurrence edges to return"),
|
||||
min_count: int = Query(default=1, description="Minimum cooccurrence_count to include an edge"),
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
@@ -4490,7 +4682,7 @@ def _register_routes(app: FastAPI):
|
||||
bank_id: str,
|
||||
body: CreateMentalModelRequest,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
_precheck: None = Depends(precheck_for("mental_model_create")),
|
||||
_precheck: None = Depends(precheck_for(PrecheckOperation.MENTAL_MODEL_CREATE)),
|
||||
):
|
||||
"""Create a mental model (async - returns operation_id)."""
|
||||
try:
|
||||
@@ -4539,7 +4731,7 @@ def _register_routes(app: FastAPI):
|
||||
bank_id: str,
|
||||
mental_model_id: str,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
_precheck: None = Depends(precheck_for("mental_model_refresh")),
|
||||
_precheck: None = Depends(precheck_for(PrecheckOperation.MENTAL_MODEL_REFRESH)),
|
||||
):
|
||||
"""Refresh a mental model by re-running its source query (async)."""
|
||||
try:
|
||||
@@ -4894,8 +5086,8 @@ def _register_routes(app: FastAPI):
|
||||
tags_match: str = Query(
|
||||
"any_strict", description="How to match tags: 'any', 'all', 'any_strict', 'all_strict'"
|
||||
),
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
limit: int = Query(default=100, ge=0),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""
|
||||
@@ -5079,8 +5271,8 @@ def _register_routes(app: FastAPI):
|
||||
default="memories",
|
||||
description="Where to read tags from: 'memories' (memory_units, default) or 'mental_models'.",
|
||||
),
|
||||
limit: int = Query(default=100, description="Maximum number of tags to return"),
|
||||
offset: int = Query(default=0, description="Offset for pagination"),
|
||||
limit: int = Query(default=100, ge=0, description="Maximum number of tags to return"),
|
||||
offset: int = Query(default=0, ge=0, description="Offset for pagination"),
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""
|
||||
@@ -5311,8 +5503,9 @@ def _register_routes(app: FastAPI):
|
||||
"/v1/default/banks/{bank_id}/operations/{operation_id}",
|
||||
response_model=OperationStatusResponse,
|
||||
summary="Get operation status",
|
||||
description="Get the status of a specific async operation. Returns 'pending', 'completed', or 'failed'. "
|
||||
"Completed operations are removed from storage, so 'completed' means the operation finished successfully.",
|
||||
description="Get the status of a specific async operation. Returns 'pending', 'processing', 'completed', "
|
||||
"'failed', or 'cancelled'. Completed operations remain queryable with their payload for the configured "
|
||||
"retention window and are pruned afterward.",
|
||||
operation_id="get_operation_status",
|
||||
tags=["Operations"],
|
||||
)
|
||||
@@ -5417,6 +5610,42 @@ def _register_routes(app: FastAPI):
|
||||
logger.error(f"Error in POST /v1/default/banks/{bank_id}/operations/{operation_id}/retry: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.delete(
|
||||
"/v1/default/banks/{bank_id}/operations/{operation_id}/delete",
|
||||
response_model=DeleteOperationResponse,
|
||||
summary="Delete a terminal async operation",
|
||||
description="Permanently remove a failed, cancelled, or completed async operation record",
|
||||
operation_id="delete_operation",
|
||||
tags=["Operations"],
|
||||
)
|
||||
@audited("delete_operation", request_param=None)
|
||||
async def api_delete_operation(
|
||||
bank_id: str, operation_id: str, request_context: RequestContext = Depends(get_request_context)
|
||||
):
|
||||
"""Delete a terminal async operation record."""
|
||||
try:
|
||||
try:
|
||||
uuid.UUID(operation_id)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=400, detail=f"Invalid operation_id format: {operation_id}")
|
||||
|
||||
result = await app.state.memory.delete_operation(bank_id, operation_id, request_context=request_context)
|
||||
return DeleteOperationResponse(**result)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
except OperationValidationError as e:
|
||||
raise HTTPException(status_code=e.status_code, detail=e.reason)
|
||||
except (AuthenticationError, HTTPException):
|
||||
raise
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
|
||||
logger.error(
|
||||
f"Error in DELETE /v1/default/banks/{bank_id}/operations/{operation_id}/delete: {error_detail}"
|
||||
)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.get(
|
||||
"/v1/default/banks/{bank_id}/profile",
|
||||
response_model=BankProfileResponse,
|
||||
@@ -5552,8 +5781,12 @@ def _register_routes(app: FastAPI):
|
||||
):
|
||||
"""Create or update an agent with disposition and mission."""
|
||||
try:
|
||||
# Ensure bank exists by getting profile (auto-creates with defaults)
|
||||
await app.state.memory.get_bank_profile(bank_id, request_context=request_context)
|
||||
# Ensure bank exists, validating create_bank only when this call
|
||||
# actually creates a missing bank.
|
||||
await app.state.memory._ensure_bank_exists(
|
||||
bank_id,
|
||||
request_context,
|
||||
)
|
||||
|
||||
# Update name if provided (stored in DB for display only, deprecated)
|
||||
if request.name is not None:
|
||||
@@ -5608,8 +5841,13 @@ def _register_routes(app: FastAPI):
|
||||
):
|
||||
"""Partially update an agent's profile (name, mission, disposition)."""
|
||||
try:
|
||||
# Ensure bank exists
|
||||
await app.state.memory.get_bank_profile(bank_id, request_context=request_context)
|
||||
# PATCH is update-only; missing banks must not be created as a
|
||||
# side effect of reading the profile.
|
||||
existing_profile = await app.state.memory.get_bank_profile(
|
||||
bank_id, request_context=request_context, create_if_missing=False
|
||||
)
|
||||
if existing_profile is None:
|
||||
raise HTTPException(status_code=404, detail=f"Bank '{bank_id}' not found")
|
||||
|
||||
# Update name if provided (stored in DB for display only, deprecated)
|
||||
if request.name is not None:
|
||||
@@ -5625,7 +5863,11 @@ def _register_routes(app: FastAPI):
|
||||
await app.state.memory._config_resolver.update_bank_config(bank_id, config_updates, request_context)
|
||||
|
||||
# Get final profile
|
||||
final_profile = await app.state.memory.get_bank_profile(bank_id, request_context=request_context)
|
||||
final_profile = await app.state.memory.get_bank_profile(
|
||||
bank_id, request_context=request_context, create_if_missing=False
|
||||
)
|
||||
if final_profile is None:
|
||||
raise HTTPException(status_code=404, detail=f"Bank '{bank_id}' not found")
|
||||
disposition_dict = (
|
||||
final_profile["disposition"].model_dump()
|
||||
if hasattr(final_profile["disposition"], "model_dump")
|
||||
@@ -5736,8 +5978,12 @@ def _register_routes(app: FastAPI):
|
||||
dry_run=True,
|
||||
)
|
||||
|
||||
# Ensure bank exists (auto-creates with defaults if needed)
|
||||
await app.state.memory.get_bank_profile(bank_id, request_context=request_context)
|
||||
# Ensure bank exists, validating create_bank only when this import
|
||||
# actually creates a missing target bank.
|
||||
await app.state.memory._ensure_bank_exists(
|
||||
bank_id,
|
||||
request_context,
|
||||
)
|
||||
|
||||
return await apply_bank_template_manifest(
|
||||
memory=app.state.memory,
|
||||
@@ -6116,9 +6362,11 @@ def _register_routes(app: FastAPI):
|
||||
# Authenticate and set schema context for multi-tenant DB queries
|
||||
await app.state.memory._authenticate_tenant(request_context)
|
||||
if app.state.memory._operation_validator:
|
||||
from hindsight_api.extensions import BankReadContext
|
||||
from hindsight_api.extensions import BankReadContext, BankReadOperation
|
||||
|
||||
ctx = BankReadContext(bank_id=bank_id, operation="get_bank_config", request_context=request_context)
|
||||
ctx = BankReadContext(
|
||||
bank_id=bank_id, operation=BankReadOperation.GET_BANK_CONFIG, request_context=request_context
|
||||
)
|
||||
await app.state.memory._validate_operation(
|
||||
app.state.memory._operation_validator.validate_bank_read(ctx)
|
||||
)
|
||||
@@ -6164,9 +6412,11 @@ def _register_routes(app: FastAPI):
|
||||
# Authenticate and set schema context for multi-tenant DB queries
|
||||
await app.state.memory._authenticate_tenant(request_context)
|
||||
if app.state.memory._operation_validator:
|
||||
from hindsight_api.extensions import BankWriteContext
|
||||
from hindsight_api.extensions import BankWriteContext, BankWriteOperation
|
||||
|
||||
ctx = BankWriteContext(bank_id=bank_id, operation="update_bank_config", request_context=request_context)
|
||||
ctx = BankWriteContext(
|
||||
bank_id=bank_id, operation=BankWriteOperation.UPDATE_BANK_CONFIG, request_context=request_context
|
||||
)
|
||||
await app.state.memory._validate_operation(
|
||||
app.state.memory._operation_validator.validate_bank_write(ctx)
|
||||
)
|
||||
@@ -6223,9 +6473,11 @@ def _register_routes(app: FastAPI):
|
||||
# Authenticate and set schema context for multi-tenant DB queries
|
||||
await app.state.memory._authenticate_tenant(request_context)
|
||||
if app.state.memory._operation_validator:
|
||||
from hindsight_api.extensions import BankWriteContext
|
||||
from hindsight_api.extensions import BankWriteContext, BankWriteOperation
|
||||
|
||||
ctx = BankWriteContext(bank_id=bank_id, operation="reset_bank_config", request_context=request_context)
|
||||
ctx = BankWriteContext(
|
||||
bank_id=bank_id, operation=BankWriteOperation.RESET_BANK_CONFIG, request_context=request_context
|
||||
)
|
||||
await app.state.memory._validate_operation(
|
||||
app.state.memory._operation_validator.validate_bank_write(ctx)
|
||||
)
|
||||
@@ -6596,7 +6848,7 @@ def _register_routes(app: FastAPI):
|
||||
bank_id: str,
|
||||
request: RetainRequest,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
_precheck: None = Depends(precheck_for("retain")),
|
||||
_precheck: None = Depends(precheck_for(PrecheckOperation.RETAIN)),
|
||||
):
|
||||
"""Retain memories with optional async processing."""
|
||||
metrics = get_metrics_collector()
|
||||
@@ -6630,6 +6882,11 @@ def _register_routes(app: FastAPI):
|
||||
strategy_groups[effective].append(content_dict)
|
||||
|
||||
if request.async_:
|
||||
if request.operation_id is not None and len(strategy_groups) != 1:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="operation_id requires all retain items to resolve to a single strategy",
|
||||
)
|
||||
# Async processing: one submit per strategy group
|
||||
all_operation_ids = []
|
||||
total_items_count = 0
|
||||
@@ -6640,6 +6897,7 @@ def _register_routes(app: FastAPI):
|
||||
document_tags=request.document_tags,
|
||||
strategy=group_strategy,
|
||||
request_context=request_context,
|
||||
operation_id=request.operation_id,
|
||||
)
|
||||
all_operation_ids.append(result["operation_id"])
|
||||
total_items_count += result["items_count"]
|
||||
@@ -6706,6 +6964,10 @@ def _register_routes(app: FastAPI):
|
||||
)
|
||||
except OperationValidationError as e:
|
||||
raise HTTPException(status_code=e.status_code, detail=e.reason)
|
||||
except RetainOperationConflictError as e:
|
||||
# Caller reused an async retain operation_id that already belongs to
|
||||
# a different operation.
|
||||
raise HTTPException(status_code=409, detail=str(e))
|
||||
except (AuthenticationError, HTTPException):
|
||||
raise
|
||||
except ValueError as e:
|
||||
@@ -6749,7 +7011,7 @@ def _register_routes(app: FastAPI):
|
||||
description="Upload files (PDF, DOCX, etc.), convert them to markdown, and retain as memories.\n\n"
|
||||
"This endpoint handles file upload, conversion, and memory creation in a single operation.\n\n"
|
||||
"**Features:**\n"
|
||||
"- Supports PDF, DOCX, PPTX, XLSX, images (with OCR), audio (with transcription)\n"
|
||||
"- Supports PDF, DOCX, PPTX, XLSX, images (parser-dependent OCR), audio (with transcription)\n"
|
||||
"- Automatic file-to-markdown conversion using pluggable parsers\n"
|
||||
"- Files stored in object storage (PostgreSQL by default, S3 for production)\n"
|
||||
"- Each file becomes a separate document with optional metadata/tags\n"
|
||||
@@ -6779,7 +7041,7 @@ def _register_routes(app: FastAPI):
|
||||
files: list[UploadFile] = File(..., description="Files to upload and convert"),
|
||||
request: str = Form(..., description="JSON string with FileRetainRequest model"),
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
_precheck: None = Depends(precheck_for("files_retain")),
|
||||
_precheck: None = Depends(precheck_for(PrecheckOperation.FILES_RETAIN)),
|
||||
):
|
||||
"""Upload and convert files to memories."""
|
||||
from hindsight_api.config import get_config
|
||||
|
||||
@@ -9,7 +9,7 @@ from fastmcp import FastMCP
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
from hindsight_api import __version__ as HINDSIGHT_VERSION
|
||||
from hindsight_api.config import _get_raw_config
|
||||
from hindsight_api.config import DEFAULT_MCP_RECALL_DESCRIPTION, DEFAULT_MCP_RETAIN_DESCRIPTION, _get_raw_config
|
||||
from hindsight_api.engine.memory_engine import _current_schema
|
||||
from hindsight_api.extensions import MCPExtension, load_extension
|
||||
from hindsight_api.extensions.tenant import AuthenticationError
|
||||
@@ -78,6 +78,19 @@ def get_current_mcp_authenticated() -> bool:
|
||||
return _current_mcp_authenticated.get()
|
||||
|
||||
|
||||
def _build_mcp_tool_descriptions(extra_instructions: str | None) -> tuple[str | None, str | None]:
|
||||
"""Return custom retain/recall descriptions when server-level MCP instructions are set."""
|
||||
if not isinstance(extra_instructions, str):
|
||||
return None, None
|
||||
|
||||
extra_instructions = extra_instructions.strip()
|
||||
if not extra_instructions:
|
||||
return None, None
|
||||
|
||||
suffix = f"\n\nAdditional instructions: {extra_instructions}"
|
||||
return DEFAULT_MCP_RETAIN_DESCRIPTION + suffix, DEFAULT_MCP_RECALL_DESCRIPTION + suffix
|
||||
|
||||
|
||||
def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||
"""
|
||||
Create and configure the Hindsight MCP server.
|
||||
@@ -135,6 +148,10 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||
allowed = frozenset(global_config.mcp_enabled_tools)
|
||||
base_tools = (base_tools if base_tools is not None else _ALL_TOOLS) & allowed
|
||||
|
||||
retain_description, recall_description = _build_mcp_tool_descriptions(
|
||||
getattr(global_config, "mcp_instructions", None)
|
||||
)
|
||||
|
||||
# Configure and register tools using shared module
|
||||
config = MCPToolsConfig(
|
||||
bank_id_resolver=get_current_bank_id,
|
||||
@@ -144,6 +161,8 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||
mcp_authenticated_resolver=get_current_mcp_authenticated, # Propagate MCP pre-auth flag
|
||||
include_bank_id_param=multi_bank,
|
||||
tools=base_tools,
|
||||
retain_description=retain_description,
|
||||
recall_description=recall_description,
|
||||
)
|
||||
|
||||
register_mcp_tools(mcp, memory, config)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -8,6 +8,7 @@ Config values are resolved on every request to ensure consistency across
|
||||
multiple API servers.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import asdict, replace
|
||||
@@ -127,6 +128,21 @@ class ConfigResolver:
|
||||
# Return full config object (dataclass doesn't have __init__ that accepts kwargs, so we update the object)
|
||||
# Create a new config instance by copying the global config and updating fields
|
||||
resolved_config = HindsightConfig(**config_dict)
|
||||
# Multi-LLM chains are static credential fields (never tenant/bank-overridable),
|
||||
# but asdict() above flattened their member dataclasses into plain dicts. Restore
|
||||
# the original typed objects from the global config so the resolved object stays
|
||||
# well-typed for any consumer that reads them.
|
||||
resolved_config = replace(
|
||||
resolved_config,
|
||||
llm_members=self._global_config.llm_members,
|
||||
llm_strategy=self._global_config.llm_strategy,
|
||||
retain_llm_members=self._global_config.retain_llm_members,
|
||||
retain_llm_strategy=self._global_config.retain_llm_strategy,
|
||||
reflect_llm_members=self._global_config.reflect_llm_members,
|
||||
reflect_llm_strategy=self._global_config.reflect_llm_strategy,
|
||||
consolidation_llm_members=self._global_config.consolidation_llm_members,
|
||||
consolidation_llm_strategy=self._global_config.consolidation_llm_strategy,
|
||||
)
|
||||
validate_retain_chunking_config(
|
||||
resolved_config.retain_chunk_size,
|
||||
resolved_config.retain_structured_chunk_size,
|
||||
@@ -161,26 +177,83 @@ class ConfigResolver:
|
||||
resolved_config = await self.resolve_full_config(bank_id, context)
|
||||
config_dict = asdict(resolved_config)
|
||||
|
||||
# SECURITY: Filter to only configurable fields (exclude static/infrastructure)
|
||||
filtered = {k: v for k, v in config_dict.items() if k in self._configurable_fields}
|
||||
# SECURITY: drop static/infrastructure + credential fields, then permission-filter.
|
||||
filtered = self._strip_static_and_credential_fields(config_dict)
|
||||
return await self._apply_permission_filter(filtered, bank_id, context)
|
||||
|
||||
# SECURITY: Remove ALL credential fields (API keys, base URLs, etc.)
|
||||
filtered = {k: v for k, v in filtered.items() if k not in self._credential_fields}
|
||||
def _strip_static_and_credential_fields(self, config_dict: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Keep only configurable, non-credential fields.
|
||||
|
||||
# PERMISSIONS: Further filter based on tenant/bank permissions
|
||||
SECURITY: excludes static/infrastructure fields and ALL credential fields
|
||||
(API keys, base URLs, etc.) so a resolved config is safe to return over the API.
|
||||
"""
|
||||
return {
|
||||
k: v for k, v in config_dict.items() if k in self._configurable_fields and k not in self._credential_fields
|
||||
}
|
||||
|
||||
async def _apply_permission_filter(
|
||||
self, filtered: dict[str, Any], bank_id: str, context: RequestContext | None
|
||||
) -> dict[str, Any]:
|
||||
"""Further restrict already-stripped config to the tenant/bank permission allow-list.
|
||||
|
||||
On extension error, leaves ``filtered`` unchanged (parity with the historical
|
||||
single-bank path: a permissions lookup failure must not leak or drop fields).
|
||||
"""
|
||||
if not (self.tenant_extension and context):
|
||||
return filtered
|
||||
try:
|
||||
allowed_fields = await self.tenant_extension.get_allowed_config_fields(context, bank_id)
|
||||
if allowed_fields is not None: # None means "allow all"
|
||||
filtered = {k: v for k, v in filtered.items() if k in allowed_fields}
|
||||
logger.debug(
|
||||
f"Applied permission filter for bank {bank_id}: allowed={len(allowed_fields)} fields, "
|
||||
f"returned={len(filtered)} fields"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load permissions for bank {bank_id}: {e}")
|
||||
return filtered
|
||||
|
||||
async def get_bank_configs(
|
||||
self, bank_ids: list[str], context: RequestContext | None = None
|
||||
) -> dict[str, dict[str, Any]]:
|
||||
"""Batch variant of :meth:`get_bank_config` for many banks.
|
||||
|
||||
Equivalent to calling ``get_bank_config`` per bank, but resolves the
|
||||
global + tenant base once and loads every bank's ``banks.config`` JSONB
|
||||
in a single query, instead of one config round-trip per bank. Used by
|
||||
``list_banks`` to overlay disposition + mission without an N+1.
|
||||
|
||||
Returns a mapping of bank_id -> filtered configurable-field dict. A bank
|
||||
with no config row still appears, mapped to the global+tenant base.
|
||||
"""
|
||||
if not bank_ids:
|
||||
return {}
|
||||
|
||||
# Global + tenant base, resolved once (tenant override is per-request, not per-bank).
|
||||
base_dict = asdict(self._global_config)
|
||||
if self.tenant_extension and context:
|
||||
try:
|
||||
allowed_fields = await self.tenant_extension.get_allowed_config_fields(context, bank_id)
|
||||
if allowed_fields is not None: # None means "allow all"
|
||||
filtered = {k: v for k, v in filtered.items() if k in allowed_fields}
|
||||
logger.debug(
|
||||
f"Applied permission filter for bank {bank_id}: allowed={len(allowed_fields)} fields, "
|
||||
f"returned={len(filtered)} fields"
|
||||
)
|
||||
tenant_overrides = await self.tenant_extension.get_tenant_config(context)
|
||||
if tenant_overrides:
|
||||
normalized_tenant = normalize_config_dict(tenant_overrides)
|
||||
base_dict.update({k: v for k, v in normalized_tenant.items() if k in self._configurable_fields})
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load permissions for bank {bank_id}: {e}")
|
||||
logger.warning(f"Failed to load tenant config for bulk resolve: {e}")
|
||||
|
||||
return filtered
|
||||
# All bank overrides in one query, then merge + strip per bank.
|
||||
bank_overrides = await self._load_bank_configs(bank_ids)
|
||||
stripped = {
|
||||
bank_id: self._strip_static_and_credential_fields({**base_dict, **bank_overrides.get(bank_id, {})})
|
||||
for bank_id in bank_ids
|
||||
}
|
||||
|
||||
# Permission filter is per-bank; resolve concurrently when an extension is present.
|
||||
if not (self.tenant_extension and context):
|
||||
return stripped
|
||||
permission_filtered = await asyncio.gather(
|
||||
*(self._apply_permission_filter(stripped[bank_id], bank_id, context) for bank_id in bank_ids)
|
||||
)
|
||||
return dict(zip(bank_ids, permission_filtered, strict=True))
|
||||
|
||||
async def _load_bank_config(self, bank_id: str) -> dict[str, Any]:
|
||||
"""
|
||||
@@ -219,6 +292,45 @@ class ConfigResolver:
|
||||
|
||||
return {}
|
||||
|
||||
async def _load_bank_configs(self, bank_ids: list[str]) -> dict[str, dict[str, Any]]:
|
||||
"""Bulk variant of :meth:`_load_bank_config`: load many banks' overrides in one query.
|
||||
|
||||
Returns a mapping of bank_id -> normalized active overrides. Banks with no row
|
||||
(or an empty/all-tombstone config) are simply absent from the mapping.
|
||||
"""
|
||||
result: dict[str, dict[str, Any]] = {}
|
||||
if not bank_ids:
|
||||
return result
|
||||
try:
|
||||
async with self._backend.acquire() as conn:
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT bank_id, config FROM {fq_table("banks")} WHERE bank_id = ANY($1)
|
||||
""",
|
||||
bank_ids,
|
||||
)
|
||||
for row in rows:
|
||||
config_data = row["config"]
|
||||
if not config_data:
|
||||
continue
|
||||
# Handle case where JSONB is returned as JSON string
|
||||
if isinstance(config_data, str):
|
||||
config_data = json.loads(config_data)
|
||||
|
||||
# Normalize keys (handle both env var format and Python field format)
|
||||
normalized = normalize_config_dict(config_data)
|
||||
|
||||
# Only active overrides for configurable fields. JSON null is a tombstone
|
||||
# for "Server Default" in the bank-config UI and must not override defaults.
|
||||
overrides = {
|
||||
k: v for k, v in normalized.items() if k in self._configurable_fields and v is not None
|
||||
}
|
||||
if overrides:
|
||||
result[row["bank_id"]] = overrides
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to bulk-load bank configs: {e}")
|
||||
return result
|
||||
|
||||
async def update_bank_config(
|
||||
self, bank_id: str, updates: dict[str, Any], context: RequestContext | None = None
|
||||
) -> None:
|
||||
@@ -305,6 +417,9 @@ class ConfigResolver:
|
||||
# Validate recall budget fields
|
||||
_validate_recall_budget_updates(normalized_updates)
|
||||
|
||||
# Validate disposition trait fields (1-5 integer scale)
|
||||
_validate_disposition_updates(normalized_updates)
|
||||
|
||||
chunking_fields_updated = (
|
||||
"retain_chunk_size" in normalized_updates
|
||||
or "retain_structured_chunk_size" in normalized_updates
|
||||
@@ -419,6 +534,31 @@ def _validate_recall_budget_updates(updates: dict[str, Any]) -> None:
|
||||
)
|
||||
|
||||
|
||||
_DISPOSITION_KEYS = (
|
||||
"disposition_skepticism",
|
||||
"disposition_literalism",
|
||||
"disposition_empathy",
|
||||
)
|
||||
|
||||
|
||||
def _validate_disposition_updates(updates: dict[str, Any]) -> None:
|
||||
"""Validate disposition trait config updates. Raises ValueError on invalid input.
|
||||
|
||||
Each trait is an integer on a 1-5 scale (or None to clear the per-bank
|
||||
override). The read overlay injects the stored value verbatim into a strict
|
||||
``DispositionTraits(int, ge=1, le=5)``; an out-of-contract value (a float, a
|
||||
0-1 scale, or an int outside 1-5) accepted here would later 500 the whole
|
||||
bank list when any bank profile is serialized (issue #2348).
|
||||
"""
|
||||
for key in _DISPOSITION_KEYS:
|
||||
if key in updates:
|
||||
value = updates[key]
|
||||
if value is None:
|
||||
continue
|
||||
if not isinstance(value, int) or isinstance(value, bool) or not (1 <= value <= 5):
|
||||
raise ValueError(f"{key} must be an integer between 1 and 5, got {value!r}")
|
||||
|
||||
|
||||
def apply_strategy(config: HindsightConfig, strategy_name: str) -> HindsightConfig:
|
||||
"""
|
||||
Apply a named retain strategy's overrides on top of a resolved config.
|
||||
|
||||
@@ -10,7 +10,7 @@ import asyncio
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Awaitable, Callable
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
@@ -19,6 +19,7 @@ from typing import Any
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..engine.db_utils import acquire_with_retry
|
||||
from ..models import RequestContext
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -119,23 +120,60 @@ class AuditLogger:
|
||||
schema_getter: Callable[[], str],
|
||||
enabled: bool,
|
||||
allowed_actions: list[str],
|
||||
bank_enabled_resolver: Callable[[str, RequestContext | None], Awaitable[bool]] | None = None,
|
||||
) -> None:
|
||||
self._pool_getter = pool_getter
|
||||
self._schema_getter = schema_getter
|
||||
self._enabled = enabled
|
||||
self._allowed_actions: frozenset[str] | None = frozenset(allowed_actions) if allowed_actions else None
|
||||
# Resolves the hierarchical ``audit_log_enabled`` for one bank
|
||||
# (env -> tenant -> bank). None means "no per-bank resolution wired",
|
||||
# in which case the global value alone decides.
|
||||
self._bank_enabled_resolver = bank_enabled_resolver
|
||||
|
||||
def is_enabled(self, action: str) -> bool:
|
||||
"""Check if audit logging is enabled for this action."""
|
||||
if not self._enabled:
|
||||
def action_allowed(self, action: str) -> bool:
|
||||
"""Global action-allowlist check. Cheap, synchronous, bank-independent.
|
||||
|
||||
The allowlist is deployment-wide, so this is a valid pre-filter to skip
|
||||
work for actions that can never be audited. It deliberately does NOT
|
||||
consult the enabled flag: that is per-bank overridable, so a bank may
|
||||
turn auditing ON even when the deployment default is off.
|
||||
"""
|
||||
if self._allowed_actions is None:
|
||||
return True
|
||||
return action in self._allowed_actions
|
||||
|
||||
async def should_log(self, action: str, bank_id: str | None, context: RequestContext | None = None) -> bool:
|
||||
"""Full audit decision: action allowlist AND the bank's resolved switch.
|
||||
|
||||
``audit_log_enabled`` is hierarchical (env -> tenant -> bank), so the
|
||||
effective value depends on which bank the action targets. Falls back to
|
||||
the global value when there is no bank in scope or no resolver wired.
|
||||
"""
|
||||
if not self.action_allowed(action):
|
||||
return False
|
||||
if self._allowed_actions is not None:
|
||||
return action in self._allowed_actions
|
||||
return True
|
||||
if bank_id is None or self._bank_enabled_resolver is None:
|
||||
return self._enabled
|
||||
try:
|
||||
return await self._bank_enabled_resolver(bank_id, context)
|
||||
except Exception as e:
|
||||
# Never let a config-resolution failure break the request. Fall back
|
||||
# to the deployment default: a transient DB blip must not silently
|
||||
# create an audit gap for a bank meant to be audited. The tradeoff is
|
||||
# the opt-out direction — a bank that overrode to false under a
|
||||
# default-on deployment will be audited during the outage. We accept
|
||||
# that: a few extra audit rows during a DB blip is the safer failure
|
||||
# than dropping records that compliance may require.
|
||||
logger.warning(f"Audit config resolution failed for bank={bank_id}: {e}; using global default")
|
||||
return self._enabled
|
||||
|
||||
def log_fire_and_forget(self, entry: AuditEntry) -> None:
|
||||
"""Schedule an audit write as a background task."""
|
||||
if not self.is_enabled(entry.action):
|
||||
"""Schedule an audit write as a background task.
|
||||
|
||||
Assumes the caller already made the audit decision via ``should_log``;
|
||||
only the bank-independent allowlist is re-checked here.
|
||||
"""
|
||||
if not self.action_allowed(entry.action):
|
||||
return
|
||||
try:
|
||||
asyncio.create_task(self._safe_log(entry))
|
||||
@@ -182,6 +220,7 @@ async def audit_context(
|
||||
bank_id: str | None = None,
|
||||
request: dict[str, Any] | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
context: RequestContext | None = None,
|
||||
):
|
||||
"""Async context manager that times the operation and writes audit on exit.
|
||||
|
||||
@@ -190,7 +229,7 @@ async def audit_context(
|
||||
result = await do_work()
|
||||
entry.response = result_dict
|
||||
"""
|
||||
if audit_logger is None or not audit_logger.is_enabled(action):
|
||||
if audit_logger is None or not await audit_logger.should_log(action, bank_id, context):
|
||||
entry = AuditEntry(action=action, transport=transport, bank_id=bank_id)
|
||||
yield entry
|
||||
return
|
||||
|
||||
@@ -13,6 +13,8 @@ but operators should opt in with that in mind.
|
||||
|
||||
from typing import Any
|
||||
|
||||
RERANKER_BANK_ID_HEADER = "X-Hindsight-Bank-Id"
|
||||
|
||||
|
||||
def apply_bank_attribution(request: dict[str, Any]) -> None:
|
||||
"""Tag ``request`` with ``user=<bank_id>`` for per-bank cost attribution.
|
||||
@@ -32,3 +34,14 @@ def apply_bank_attribution(request: dict[str, Any]) -> None:
|
||||
bank_id = get_current_bank_id()
|
||||
if bank_id:
|
||||
request["user"] = bank_id
|
||||
|
||||
|
||||
def reranker_bank_attribution_headers() -> dict[str, str]:
|
||||
"""Return the fixed per-bank header for trusted remote reranker endpoints."""
|
||||
from ..config import get_config
|
||||
from .memory_engine import get_current_bank_id
|
||||
|
||||
if not get_config().reranker_send_bank_as_header:
|
||||
return {}
|
||||
bank_id = get_current_bank_id()
|
||||
return {RERANKER_BANK_ID_HEADER: bank_id} if bank_id else {}
|
||||
|
||||
@@ -13,9 +13,18 @@ in-flight task so that N concurrent callers produce one query rather than N.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from typing import Any, Awaitable, Callable
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable
|
||||
|
||||
from .db_utils import acquire_with_retry
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .db.base import DatabaseBackend
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BankStatsCache:
|
||||
@@ -66,17 +75,28 @@ class BankStatsCache:
|
||||
schema: str,
|
||||
bank_id: str,
|
||||
loader: Callable[[], Awaitable[dict[str, Any]]],
|
||||
*,
|
||||
force_refresh: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""Return cached stats for `(schema, bank_id)` or call `loader()`.
|
||||
|
||||
Concurrent misses on the same key are coalesced onto a single
|
||||
in-flight loader.
|
||||
in-flight loader. When ``force_refresh`` is set the cached value is
|
||||
ignored: the loader runs and its result replaces the cached entry.
|
||||
"""
|
||||
if not self.enabled:
|
||||
return await loader()
|
||||
|
||||
key = (schema, bank_id)
|
||||
|
||||
if force_refresh:
|
||||
value = await loader()
|
||||
async with self._lock:
|
||||
self._store_unlocked(key, value)
|
||||
# Supersede any loader that was in flight for this key.
|
||||
self._in_flight.pop(key, None)
|
||||
return value
|
||||
|
||||
async with self._lock:
|
||||
cached = self._get_fresh_unlocked(key)
|
||||
if cached is not None:
|
||||
@@ -96,7 +116,10 @@ class BankStatsCache:
|
||||
value = await loader()
|
||||
except BaseException as exc:
|
||||
async with self._lock:
|
||||
self._in_flight.pop(key, None)
|
||||
# Invalidation may have detached this loader and allowed a new
|
||||
# one to claim the key. Never remove that newer loader's slot.
|
||||
if self._in_flight.get(key) is in_flight:
|
||||
self._in_flight.pop(key, None)
|
||||
if not in_flight.done():
|
||||
in_flight.set_exception(exc)
|
||||
# Suppress "Future exception was never retrieved" when no other
|
||||
@@ -106,8 +129,12 @@ class BankStatsCache:
|
||||
raise
|
||||
|
||||
async with self._lock:
|
||||
self._store_unlocked(key, value)
|
||||
self._in_flight.pop(key, None)
|
||||
# Only the loader that still owns the key may populate the cache.
|
||||
# An invalidated loader can finish for its original callers, but its
|
||||
# pre-invalidation result must not overwrite a newer load.
|
||||
if self._in_flight.get(key) is in_flight:
|
||||
self._store_unlocked(key, value)
|
||||
self._in_flight.pop(key, None)
|
||||
if not in_flight.done():
|
||||
in_flight.set_result(value)
|
||||
return value
|
||||
@@ -115,8 +142,113 @@ class BankStatsCache:
|
||||
async def invalidate(self, schema: str, bank_id: str) -> None:
|
||||
"""Drop any cached stats for `(schema, bank_id)`."""
|
||||
async with self._lock:
|
||||
self._entries.pop((schema, bank_id), None)
|
||||
key = (schema, bank_id)
|
||||
self._entries.pop(key, None)
|
||||
# Detach rather than cancel: existing callers may finish with the
|
||||
# snapshot they requested, while post-invalidation callers reload.
|
||||
self._in_flight.pop(key, None)
|
||||
|
||||
async def clear(self) -> None:
|
||||
async with self._lock:
|
||||
self._entries.clear()
|
||||
self._in_flight.clear()
|
||||
|
||||
|
||||
class DistributedBankStatsCache:
|
||||
"""Table-backed (cross-process) TTL cache for `get_bank_stats`.
|
||||
|
||||
Same ``get_or_load`` / ``invalidate`` / ``clear`` contract as
|
||||
:class:`BankStatsCache`, but the store is the per-schema ``bank_stats_cache``
|
||||
table instead of a per-process dict — so one worker's computation is shared
|
||||
with every other worker, and no caller recomputes while a fresh row exists.
|
||||
|
||||
On a hit, a call is a single primary-key ``SELECT`` (sub-millisecond); only a
|
||||
miss runs the (expensive) ``loader`` and writes the row back. Concurrent
|
||||
misses are *not* coalesced across processes (that would need a lock): they
|
||||
each compute and ``UPSERT``, last write wins — all results are correct, at the
|
||||
cost of a brief redundant compute at expiry.
|
||||
|
||||
Every DB touch is best-effort: if the cache table is unreachable or missing
|
||||
(e.g. a schema mid-migration), the call degrades to computing without caching
|
||||
rather than failing ``get_bank_stats``. PostgreSQL only — the engine keeps the
|
||||
in-process :class:`BankStatsCache` for Oracle.
|
||||
"""
|
||||
|
||||
def __init__(self, *, backend: "DatabaseBackend", ttl_seconds: float) -> None:
|
||||
self._backend = backend
|
||||
self._ttl = float(ttl_seconds)
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return self._ttl > 0
|
||||
|
||||
@staticmethod
|
||||
def _qualified(schema: str) -> str:
|
||||
return f'"{schema}".bank_stats_cache' if schema else "bank_stats_cache"
|
||||
|
||||
async def get_or_load(
|
||||
self,
|
||||
schema: str,
|
||||
bank_id: str,
|
||||
loader: Callable[[], Awaitable[dict[str, Any]]],
|
||||
*,
|
||||
force_refresh: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
if not self.enabled:
|
||||
return await loader()
|
||||
|
||||
table = self._qualified(schema)
|
||||
|
||||
# 1. Fresh row? Single PK lookup; ``payload::text`` sidesteps any
|
||||
# jsonb->object codec so we always decode the same way. Skipped when
|
||||
# the caller forces a refresh — then we recompute and overwrite below.
|
||||
if not force_refresh:
|
||||
try:
|
||||
async with acquire_with_retry(self._backend) as conn:
|
||||
row = await conn.fetchrow(
|
||||
f"SELECT payload::text AS payload FROM {table} "
|
||||
f"WHERE bank_id = $1 AND computed_at > now() - make_interval(secs => $2::double precision)",
|
||||
bank_id,
|
||||
self._ttl,
|
||||
)
|
||||
if row is not None:
|
||||
return json.loads(row["payload"])
|
||||
except Exception as exc: # noqa: BLE001 — cache read must never break the endpoint
|
||||
logger.debug("bank_stats_cache read failed for %s.%s (%s); computing uncached", schema, bank_id, exc)
|
||||
return await loader()
|
||||
|
||||
# 2. Miss — compute, then write the row back (best-effort).
|
||||
value = await loader()
|
||||
try:
|
||||
async with acquire_with_retry(self._backend) as conn:
|
||||
await conn.execute(
|
||||
f"INSERT INTO {table} (bank_id, payload, computed_at) VALUES ($1, $2::jsonb, now()) "
|
||||
f"ON CONFLICT (bank_id) DO UPDATE SET payload = EXCLUDED.payload, computed_at = now()",
|
||||
bank_id,
|
||||
json.dumps(value),
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 — a failed write just means no caching this round
|
||||
logger.warning("bank_stats_cache write failed for %s.%s (%s)", schema, bank_id, exc)
|
||||
return value
|
||||
|
||||
async def invalidate(self, schema: str, bank_id: str) -> None:
|
||||
"""Drop the cached row so the next read recomputes."""
|
||||
if not self.enabled:
|
||||
return
|
||||
try:
|
||||
async with acquire_with_retry(self._backend) as conn:
|
||||
await conn.execute(f"DELETE FROM {self._qualified(schema)} WHERE bank_id = $1", bank_id)
|
||||
except Exception as exc: # noqa: BLE001 — invalidation must never break the write path
|
||||
logger.debug("bank_stats_cache invalidate failed for %s.%s (%s)", schema, bank_id, exc)
|
||||
|
||||
async def clear(self) -> None:
|
||||
"""Drop all cached rows in the current schema (best-effort)."""
|
||||
if not self.enabled:
|
||||
return
|
||||
from .memory_engine import get_current_schema
|
||||
|
||||
try:
|
||||
async with acquire_with_retry(self._backend) as conn:
|
||||
await conn.execute(f"DELETE FROM {self._qualified(get_current_schema())}")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("bank_stats_cache clear failed (%s)", exc)
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
"""Shared causal-link taxonomy.
|
||||
|
||||
Retain writes only the canonical relationship. Transfer import/export also
|
||||
preserves historical relationship types so existing banks keep their graph
|
||||
semantics without allowing new retain output to create those types.
|
||||
"""
|
||||
|
||||
CANONICAL_CAUSAL_LINK_TYPE = "caused_by"
|
||||
LEGACY_CAUSAL_LINK_TYPE_NAMES = ("causes", "enables", "prevents")
|
||||
|
||||
CANONICAL_CAUSAL_LINK_TYPES = frozenset({CANONICAL_CAUSAL_LINK_TYPE})
|
||||
LEGACY_CAUSAL_LINK_TYPES = frozenset(LEGACY_CAUSAL_LINK_TYPE_NAMES)
|
||||
CAUSAL_LINK_TYPES = (CANONICAL_CAUSAL_LINK_TYPE, *LEGACY_CAUSAL_LINK_TYPE_NAMES)
|
||||
@@ -109,6 +109,11 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
end.replace(hour=23, minute=59, second=59, microsecond=999999),
|
||||
)
|
||||
|
||||
def safe_constraint(start: datetime | None, end: datetime | None) -> DateRange | NoTemporalConstraintSentinel:
|
||||
if start is None or end is None:
|
||||
return NO_TEMPORAL_CONSTRAINT
|
||||
return constraint(start, end)
|
||||
|
||||
def subtract_months(months: int) -> datetime:
|
||||
month_index = reference_date.month - months - 1
|
||||
year = reference_date.year + month_index // 12
|
||||
@@ -126,11 +131,21 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
day = min(base_date.day, calendar.monthrange(year, month)[1])
|
||||
return base_date.replace(year=year, month=month, day=day)
|
||||
|
||||
def add_years(base_date: datetime, years: int) -> datetime:
|
||||
def add_years(base_date: datetime, years: int) -> datetime | None:
|
||||
year = base_date.year + years
|
||||
if year < datetime.min.year or year > datetime.max.year:
|
||||
return None
|
||||
day = min(base_date.day, calendar.monthrange(year, base_date.month)[1])
|
||||
return base_date.replace(year=year, day=day)
|
||||
|
||||
def add_days(base_date: datetime | None, days: int) -> datetime | None:
|
||||
if base_date is None:
|
||||
return None
|
||||
try:
|
||||
return base_date + timedelta(days=days)
|
||||
except OverflowError:
|
||||
return None
|
||||
|
||||
def has_chinese_temporal_context(match: re.Match[str]) -> bool:
|
||||
if match.end() >= len(query):
|
||||
return True
|
||||
@@ -438,6 +453,11 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
return NO_TEMPORAL_CONSTRAINT
|
||||
return constraint(start, reference_date)
|
||||
|
||||
def safe_since_constraint(start: datetime | None) -> DateRange | NoTemporalConstraintSentinel:
|
||||
if start is None:
|
||||
return NO_TEMPORAL_CONSTRAINT
|
||||
return since_constraint(start)
|
||||
|
||||
def since_from_period(
|
||||
period: DateRange | None,
|
||||
) -> DateRange | NoTemporalConstraintSentinel | None:
|
||||
@@ -450,7 +470,7 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
return None
|
||||
return since_constraint(day)
|
||||
|
||||
def relative_offset_datetime(amount: int, unit: str, direction: int) -> datetime:
|
||||
def relative_offset_datetime(amount: int, unit: str, direction: int) -> datetime | None:
|
||||
if unit in ("天", "日"):
|
||||
return reference_date + timedelta(days=direction * amount)
|
||||
if unit in ("周", "星期", "礼拜"):
|
||||
@@ -459,15 +479,15 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
return add_months(reference_date, direction * amount)
|
||||
return add_years(reference_date, direction * amount)
|
||||
|
||||
def point_constraint_at_offset(amount: int, unit: str, direction: int) -> DateRange:
|
||||
def point_constraint_at_offset(amount: int, unit: str, direction: int) -> DateRange | NoTemporalConstraintSentinel:
|
||||
d = relative_offset_datetime(amount, unit, direction)
|
||||
return constraint(d, d)
|
||||
return safe_constraint(d, d)
|
||||
|
||||
def window_to_reference(amount: int, unit: str) -> DateRange:
|
||||
return constraint(relative_offset_datetime(amount, unit, -1), reference_date)
|
||||
def window_to_reference(amount: int, unit: str) -> DateRange | NoTemporalConstraintSentinel:
|
||||
return safe_constraint(relative_offset_datetime(amount, unit, -1), reference_date)
|
||||
|
||||
def window_from_reference(amount: int, unit: str) -> DateRange:
|
||||
return constraint(reference_date, relative_offset_datetime(amount, unit, 1))
|
||||
def window_from_reference(amount: int, unit: str) -> DateRange | NoTemporalConstraintSentinel:
|
||||
return safe_constraint(reference_date, relative_offset_datetime(amount, unit, 1))
|
||||
|
||||
# Chinese rule guide
|
||||
#
|
||||
@@ -781,8 +801,8 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
if relative_year_fixed_day_since_match:
|
||||
year = relative_year_number(relative_year_fixed_day_since_match.group(1))
|
||||
base = add_years(reference_date, year - reference_date.year)
|
||||
d = base + timedelta(days=fixed_day_offset(relative_year_fixed_day_since_match.group(2)))
|
||||
return since_constraint(d)
|
||||
d = add_days(base, fixed_day_offset(relative_year_fixed_day_since_match.group(2)))
|
||||
return safe_since_constraint(d)
|
||||
|
||||
fixed_day_since_match = chinese_search(
|
||||
rf"(大大后天|大后天|后天|明天|明日|今天|今日|本日|当日|当天|昨天|昨日|大大前天|大前天|前天)"
|
||||
@@ -799,7 +819,7 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
amount = parse_chinese_number(exact_relative_since_match.group(1))
|
||||
unit = exact_relative_since_match.group(2)
|
||||
if amount is not None:
|
||||
return since_constraint(relative_offset_datetime(amount, unit, -1))
|
||||
return safe_since_constraint(relative_offset_datetime(amount, unit, -1))
|
||||
|
||||
weekend_since_match = chinese_search(
|
||||
rf"(?<![上下大小每个各隔])"
|
||||
@@ -899,8 +919,8 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
if relative_year_daypart_since_match:
|
||||
year = relative_year_number(relative_year_daypart_since_match.group(1))
|
||||
base = add_years(reference_date, year - reference_date.year)
|
||||
d = base + timedelta(days=daypart_day_offset(relative_year_daypart_since_match.group(2)))
|
||||
return since_constraint(d)
|
||||
d = add_days(base, daypart_day_offset(relative_year_daypart_since_match.group(2)))
|
||||
return safe_since_constraint(d)
|
||||
|
||||
daypart_since_match = chinese_search(
|
||||
rf"(昨晚|昨夜|前晚|前夜|今晚|今早|今晨|明早|明晚|明夜){chinese_since_suffix_pattern}"
|
||||
@@ -915,17 +935,17 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
if relative_year_daypart_match:
|
||||
year = relative_year_number(relative_year_daypart_match.group(1))
|
||||
base = add_years(reference_date, year - reference_date.year)
|
||||
d = base + timedelta(days=daypart_day_offset(relative_year_daypart_match.group(2)))
|
||||
return constraint(d, d)
|
||||
d = add_days(base, daypart_day_offset(relative_year_daypart_match.group(2)))
|
||||
return safe_constraint(d, d)
|
||||
|
||||
# Day-part abbreviations still resolve only to date granularity.
|
||||
if chinese_search(r"昨晚|昨夜"):
|
||||
d = reference_date + timedelta(days=daypart_day_offset("昨晚"))
|
||||
return constraint(d, d)
|
||||
d = add_days(reference_date, daypart_day_offset("昨晚"))
|
||||
return safe_constraint(d, d)
|
||||
|
||||
if chinese_search(r"前晚|前夜"):
|
||||
d = reference_date + timedelta(days=daypart_day_offset("前晚"))
|
||||
return constraint(d, d)
|
||||
d = add_days(reference_date, daypart_day_offset("前晚"))
|
||||
return safe_constraint(d, d)
|
||||
|
||||
if chinese_search(r"今晚|今早|今晨"):
|
||||
return constraint(reference_date, reference_date)
|
||||
@@ -941,8 +961,8 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
if relative_year_fixed_day_match:
|
||||
year = relative_year_number(relative_year_fixed_day_match.group(1))
|
||||
base = add_years(reference_date, year - reference_date.year)
|
||||
d = base + timedelta(days=fixed_day_offset(relative_year_fixed_day_match.group(2)))
|
||||
return constraint(d, d)
|
||||
d = add_days(base, fixed_day_offset(relative_year_fixed_day_match.group(2)))
|
||||
return safe_constraint(d, d)
|
||||
|
||||
if chinese_search(r"昨天|昨日"):
|
||||
d = reference_date - timedelta(days=1)
|
||||
@@ -1085,7 +1105,7 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
end_amount = parse_chinese_number(amount_text[-1])
|
||||
unit = adjacent_fuzzy_future_match.group(2)
|
||||
if start_amount is not None and end_amount is not None:
|
||||
return constraint(
|
||||
return safe_constraint(
|
||||
relative_offset_datetime(start_amount, unit, 1),
|
||||
relative_offset_datetime(end_amount, unit, 1),
|
||||
)
|
||||
@@ -1093,7 +1113,7 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
few_future_match = chinese_search(rf"[几数]个?(天|日|周|星期|礼拜|月|年){chinese_relative_future_suffix_pattern}")
|
||||
if few_future_match:
|
||||
unit = few_future_match.group(1)
|
||||
return constraint(relative_offset_datetime(2, unit, 1), relative_offset_datetime(5, unit, 1))
|
||||
return safe_constraint(relative_offset_datetime(2, unit, 1), relative_offset_datetime(5, unit, 1))
|
||||
|
||||
exact_future_match = chinese_search(
|
||||
rf"(?<![{_CHINESE_NUMERAL_PREFIX_CHARS}])([0-9]+|[{_CHINESE_NUMERAL_CHARS}]+)个?(天|日|周|星期|礼拜|月|年){chinese_relative_future_suffix_pattern}"
|
||||
@@ -1113,7 +1133,7 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
second_amount = parse_chinese_number(adjacent_fuzzy_past_match.group(2))
|
||||
unit = adjacent_fuzzy_past_match.group(3)
|
||||
if first_amount is not None and second_amount is not None and second_amount == first_amount + 1:
|
||||
return constraint(
|
||||
return safe_constraint(
|
||||
relative_offset_datetime(second_amount, unit, -1),
|
||||
relative_offset_datetime(first_amount, unit, -1),
|
||||
)
|
||||
@@ -1144,7 +1164,7 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
return constraint(reference_date - timedelta(days=150), reference_date - timedelta(days=60))
|
||||
|
||||
if chinese_search(r"一两年前|[两二]三年前|三两年前"):
|
||||
return constraint(add_years(reference_date, -3), add_years(reference_date, -1))
|
||||
return safe_constraint(add_years(reference_date, -3), add_years(reference_date, -1))
|
||||
|
||||
rolling_this_adjacent_match = chinese_search(
|
||||
r"这(一两|[两二]三|三两|三四|四五|五六|六七|七八|八九|九十)个?(天|日|周|星期|礼拜|月|年)"
|
||||
@@ -1154,7 +1174,7 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
end_amount = 3 if amount_text in ("一两", "三两") else parse_chinese_number(amount_text[-1])
|
||||
unit = rolling_this_adjacent_match.group(2)
|
||||
if end_amount is not None:
|
||||
return constraint(relative_offset_datetime(end_amount, unit, -1), reference_date)
|
||||
return safe_constraint(relative_offset_datetime(end_amount, unit, -1), reference_date)
|
||||
|
||||
rolling_this_count_match = chinese_search(rf"这([0-9]+|[{_CHINESE_NUMERAL_CHARS}]+)个?(天|日|周|星期|礼拜|月|年)")
|
||||
if rolling_this_count_match:
|
||||
@@ -1191,7 +1211,7 @@ def extract_chinese_period(query: str, reference_date: datetime) -> DateRange |
|
||||
end_amount = 3 if amount_text in ("一两", "三两") else parse_chinese_number(amount_text[-1])
|
||||
unit = rolling_past_adjacent_match.group(3)
|
||||
if end_amount is not None:
|
||||
return constraint(relative_offset_datetime(end_amount, unit, -1), reference_date)
|
||||
return safe_constraint(relative_offset_datetime(end_amount, unit, -1), reference_date)
|
||||
|
||||
rolling_past_few_match = chinese_search(r"(过去|近|最近)几个?(天|日|周|星期|礼拜|月|年)")
|
||||
if rolling_past_few_match:
|
||||
|
||||
@@ -28,6 +28,7 @@ from fnmatch import fnmatchcase
|
||||
from itertools import combinations
|
||||
from typing import TYPE_CHECKING, Any, Literal
|
||||
|
||||
import asyncpg
|
||||
from pydantic import BaseModel, field_validator
|
||||
|
||||
from ...config import get_config
|
||||
@@ -98,10 +99,33 @@ _DEDUP_TOP_K = 5
|
||||
class _DedupDecision(BaseModel):
|
||||
"""Focused 1-by-1 verdict for whether a new observation duplicates an existing one."""
|
||||
|
||||
action: Literal["merge", "keep"]
|
||||
action: Literal["merge", "keep"] = "keep"
|
||||
text: str = "" # the synthesized merged observation (when action == "merge")
|
||||
reason: str = ""
|
||||
|
||||
@field_validator("action", mode="before")
|
||||
@classmethod
|
||||
def _normalize_action(cls, value: object) -> str:
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in {"merge", "keep"}:
|
||||
return normalized
|
||||
|
||||
logger.warning("Invalid consolidation dedup action %r; defaulting to keep", value)
|
||||
return "keep"
|
||||
|
||||
|
||||
def _dedup_decision_from_response(raw: Any) -> _DedupDecision:
|
||||
try:
|
||||
if isinstance(raw, _DedupDecision):
|
||||
return raw
|
||||
if isinstance(raw, str):
|
||||
return _DedupDecision.model_validate_json(raw)
|
||||
return _DedupDecision.model_validate(raw)
|
||||
except ValueError as exc:
|
||||
logger.warning("Invalid consolidation dedup response %r; defaulting to keep: %s", raw, exc)
|
||||
return _DedupDecision(action="keep", reason="invalid structured response")
|
||||
|
||||
|
||||
_DEDUP_PROMPT = """You reconcile long-term memory observations. A NEW observation is about to be \
|
||||
stored, and it is highly similar to an EXISTING one:
|
||||
@@ -109,9 +133,20 @@ stored, and it is highly similar to an EXISTING one:
|
||||
[NEW] {new}
|
||||
[EXISTING] {existing}
|
||||
|
||||
If they assert the SAME fact (wording aside), respond action="merge" and provide `text`: a single \
|
||||
observation that preserves EVERY detail from both. If they differ in ANY important detail — a \
|
||||
number/quantity, a named entity or language, a negation, or a condition — respond action="keep"."""
|
||||
Respond with ONLY one valid JSON object matching one of these shapes:
|
||||
|
||||
For duplicate facts:
|
||||
{{"action": "merge", "text": "...", "reason": "..."}}
|
||||
|
||||
For distinct facts:
|
||||
{{"action": "keep", "text": "", "reason": "..."}}
|
||||
|
||||
Do NOT use key=value lines, markdown fences, or any text outside the JSON object.
|
||||
|
||||
If they assert the SAME fact (wording aside), set "action" to "merge" and provide "text": a \
|
||||
single observation that preserves EVERY detail from both. If they differ in ANY important detail \
|
||||
— a number/quantity, a named entity or language, a negation, or a condition — set "action" to \
|
||||
"keep" and "text" to an empty string."""
|
||||
|
||||
|
||||
def _dedup_active(config: Any) -> bool:
|
||||
@@ -189,10 +224,13 @@ async def _dedup_adjudicate(
|
||||
if best_id is None:
|
||||
return _DedupOutcome(best_id=None, merged_text="", should_merge=False)
|
||||
|
||||
decision: _DedupDecision = await dedup_llm_config.call(
|
||||
messages=[{"role": "user", "content": _DEDUP_PROMPT.format(new=anchor_text, existing=best_text)}],
|
||||
response_format=_DedupDecision,
|
||||
scope="consolidation_dedup",
|
||||
decision = _dedup_decision_from_response(
|
||||
await dedup_llm_config.call(
|
||||
messages=[{"role": "user", "content": _DEDUP_PROMPT.format(new=anchor_text, existing=best_text)}],
|
||||
response_format=_DedupDecision,
|
||||
scope="consolidation_dedup",
|
||||
strict_schema=get_config().llm_strict_schema_consolidation,
|
||||
)
|
||||
)
|
||||
if decision.action != "merge":
|
||||
return _DedupOutcome(best_id=best_id, merged_text="", should_merge=False)
|
||||
@@ -224,13 +262,18 @@ async def _dedup_reconcile_create(
|
||||
# Fold the new source facts into the twin and persist the merged text. We keep the twin's
|
||||
# existing embedding: the merged text is >= threshold similar, so the stored vector stays
|
||||
# representative and we avoid a re-embed + a dialect-specific vector UPDATE.
|
||||
search_vector_clause = (
|
||||
f",\n search_vector = to_tsvector('{config.text_search_extension_native_language}'::regconfig, COALESCE($1, ''))"
|
||||
if config.text_search_extension == "native"
|
||||
else ""
|
||||
)
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("memory_units")}
|
||||
SET text = $1,
|
||||
source_memory_ids = (SELECT array_agg(DISTINCT e) FROM unnest(source_memory_ids || $2::uuid[]) e),
|
||||
proof_count = (SELECT count(DISTINCT e) FROM unnest(source_memory_ids || $2::uuid[]) e),
|
||||
updated_at = now()
|
||||
updated_at = now(){search_vector_clause}
|
||||
WHERE id = $3::uuid
|
||||
""",
|
||||
outcome.merged_text,
|
||||
@@ -279,6 +322,11 @@ async def _dedup_reconcile_update(
|
||||
# the create path) then delete the now-redundant updated row. The all_strict/any tag match
|
||||
# guarantees twin and updated share scope, so dropping the updated row's tags loses no
|
||||
# visibility. Temporal fields follow the surviving twin (minimal scope; matches create).
|
||||
search_vector_clause = (
|
||||
f",\n search_vector = to_tsvector('{config.text_search_extension_native_language}'::regconfig, COALESCE($1, ''))"
|
||||
if config.text_search_extension == "native"
|
||||
else ""
|
||||
)
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("memory_units")} t
|
||||
@@ -289,7 +337,7 @@ async def _dedup_reconcile_update(
|
||||
proof_count = (
|
||||
SELECT count(DISTINCT e) FROM unnest(t.source_memory_ids || u.source_memory_ids) e
|
||||
),
|
||||
updated_at = now()
|
||||
updated_at = now(){search_vector_clause}
|
||||
FROM {fq_table("memory_units")} u
|
||||
WHERE t.id = $2::uuid AND u.id = $3::uuid
|
||||
""",
|
||||
@@ -449,6 +497,13 @@ class _CreateAction(BaseModel):
|
||||
def sanitize_text(cls, v: str) -> str:
|
||||
return sanitize_llm_output(v) or ""
|
||||
|
||||
@field_validator("source_fact_ids", mode="before")
|
||||
@classmethod
|
||||
def ensure_list(cls, v: str | list[str]) -> list[str]:
|
||||
if isinstance(v, str):
|
||||
return [v]
|
||||
return v
|
||||
|
||||
|
||||
class _UpdateAction(BaseModel):
|
||||
text: str
|
||||
@@ -461,6 +516,13 @@ class _UpdateAction(BaseModel):
|
||||
def sanitize_text(cls, v: str) -> str:
|
||||
return sanitize_llm_output(v) or ""
|
||||
|
||||
@field_validator("source_fact_ids", mode="before")
|
||||
@classmethod
|
||||
def ensure_list(cls, v: str | list[str]) -> list[str]:
|
||||
if isinstance(v, str):
|
||||
return [v]
|
||||
return v
|
||||
|
||||
|
||||
class _DeleteAction(BaseModel):
|
||||
observation_id: str # UUID of the observation to remove
|
||||
@@ -640,6 +702,7 @@ class ConsolidationPerfLog:
|
||||
self.start_time = time.time()
|
||||
self.lines: list[str] = []
|
||||
self.timings: dict[str, float] = {}
|
||||
self.timing_counts: dict[str, int] = {}
|
||||
self.llm_calls: int = 0
|
||||
self.total_obs_in_context: int = 0
|
||||
self.total_prompt_chars: int = 0
|
||||
@@ -649,11 +712,13 @@ class ConsolidationPerfLog:
|
||||
self.lines.append(message)
|
||||
|
||||
def record_timing(self, key: str, duration: float) -> None:
|
||||
"""Record a timing measurement."""
|
||||
if key in self.timings:
|
||||
self.timings[key] += duration
|
||||
else:
|
||||
self.timings[key] = duration
|
||||
"""Record a timing measurement.
|
||||
|
||||
Tracks both total seconds and call count so the summary can
|
||||
distinguish one slow call from many fast calls in aggregate.
|
||||
"""
|
||||
self.timings[key] = self.timings.get(key, 0.0) + duration
|
||||
self.timing_counts[key] = self.timing_counts.get(key, 0) + 1
|
||||
|
||||
def record_llm_call(self, obs_count: int, prompt_chars: int) -> None:
|
||||
"""Record stats for a single LLM call."""
|
||||
@@ -676,6 +741,8 @@ class ConsolidationPerfLog:
|
||||
"""
|
||||
for key, value in other.timings.items():
|
||||
self.timings[key] = self.timings.get(key, 0.0) + value
|
||||
for key, count in other.timing_counts.items():
|
||||
self.timing_counts[key] = self.timing_counts.get(key, 0) + count
|
||||
self.llm_calls += other.llm_calls
|
||||
self.total_obs_in_context += other.total_obs_in_context
|
||||
self.total_prompt_chars += other.total_prompt_chars
|
||||
@@ -1276,16 +1343,22 @@ async def _run_consolidation_job(
|
||||
f"{stats['skipped']} skipped)"
|
||||
)
|
||||
|
||||
# Add timing breakdown
|
||||
# Add timing breakdown. Each phase is recorded once per call, so the count
|
||||
# disambiguates a single slow call from many fast calls — important for
|
||||
# operators triaging "the recall phase took 15s" log lines, where the
|
||||
# total is the sum of many serial sub-calls rather than one slow query.
|
||||
def _fmt(key: str) -> str:
|
||||
total = perf.timings[key]
|
||||
count = perf.timing_counts.get(key, 0)
|
||||
if count > 1:
|
||||
avg_ms = total * 1000.0 / count
|
||||
return f"{key}={total:.3f}s ({count} calls, avg={avg_ms:.0f}ms)"
|
||||
return f"{key}={total:.3f}s"
|
||||
|
||||
timing_parts = []
|
||||
if "recall" in perf.timings:
|
||||
timing_parts.append(f"recall={perf.timings['recall']:.3f}s")
|
||||
if "llm" in perf.timings:
|
||||
timing_parts.append(f"llm={perf.timings['llm']:.3f}s")
|
||||
if "embedding" in perf.timings:
|
||||
timing_parts.append(f"embedding={perf.timings['embedding']:.3f}s")
|
||||
if "db_write" in perf.timings:
|
||||
timing_parts.append(f"db_write={perf.timings['db_write']:.3f}s")
|
||||
for key in ("recall", "llm", "embedding", "db_write"):
|
||||
if key in perf.timings:
|
||||
timing_parts.append(_fmt(key))
|
||||
|
||||
if perf.llm_calls > 0:
|
||||
timing_parts.append(f"avg_obs={perf.total_obs_in_context / perf.llm_calls:.1f}")
|
||||
@@ -1731,15 +1804,22 @@ async def _append_observation_history(
|
||||
history from growing without bound.
|
||||
"""
|
||||
obs_uuid = uuid.UUID(observation_id)
|
||||
await conn.execute(
|
||||
f"""
|
||||
try:
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("observation_history")} (observation_id, bank_id, content, changed_at)
|
||||
VALUES ($1, $2, $3::jsonb, now())
|
||||
""",
|
||||
obs_uuid,
|
||||
bank_id,
|
||||
json.dumps(asdict(snapshot)),
|
||||
)
|
||||
obs_uuid,
|
||||
bank_id,
|
||||
json.dumps(asdict(snapshot)),
|
||||
)
|
||||
except asyncpg.exceptions.ForeignKeyViolationError:
|
||||
logger.warning(
|
||||
f"FK violation writing observation_history for {observation_id}: "
|
||||
"observation was removed before history could be written (race with parallel consolidation). Skipping."
|
||||
)
|
||||
return
|
||||
if max_entries and max_entries > 0:
|
||||
await conn.execute(
|
||||
f"""
|
||||
@@ -1820,6 +1900,12 @@ async def _execute_update_action(
|
||||
|
||||
config = get_config()
|
||||
|
||||
search_vector_clause = (
|
||||
f",\n search_vector = to_tsvector('{config.text_search_extension_native_language}'::regconfig, COALESCE($1, ''))"
|
||||
if config.text_search_extension == "native"
|
||||
else ""
|
||||
)
|
||||
|
||||
t0 = time.time()
|
||||
await conn.execute(
|
||||
f"""
|
||||
@@ -1832,7 +1918,7 @@ async def _execute_update_action(
|
||||
updated_at = now(),
|
||||
occurred_start = LEAST(occurred_start, COALESCE($6, occurred_start)),
|
||||
occurred_end = GREATEST(occurred_end, COALESCE($7, occurred_end)),
|
||||
mentioned_at = GREATEST(mentioned_at, COALESCE($8, mentioned_at))
|
||||
mentioned_at = GREATEST(mentioned_at, COALESCE($8, mentioned_at)){search_vector_clause}
|
||||
WHERE id = $5
|
||||
""",
|
||||
new_text,
|
||||
@@ -2213,6 +2299,11 @@ async def _consolidate_batch_with_llm(
|
||||
],
|
||||
"response_format": response_model,
|
||||
"scope": "consolidation",
|
||||
# Resolved per operation (HINDSIGHT_API_LLM_STRICT_SCHEMA_CONSOLIDATION, falling
|
||||
# back to the global flag) so an operator can grammar-enforce consolidation's
|
||||
# structured output -- which narrows the raw-JSON failure mode behind #2668 --
|
||||
# without forcing strict schema on operations whose model can't satisfy it.
|
||||
"strict_schema": config.llm_strict_schema_consolidation,
|
||||
}
|
||||
# Only request an explicit output budget when configured. Left unset by default the key is
|
||||
# omitted, so each provider keeps its implicit default (backwards compatible). Operators on
|
||||
@@ -2308,16 +2399,20 @@ async def _create_observation_directly(
|
||||
tokenize($3, 'llmlingua2')::bm25_catalog.bm25vector)
|
||||
RETURNING id
|
||||
"""
|
||||
else: # native, pg_textsearch, pgroonga, or pg_search
|
||||
# pg_textsearch / pgroonga / pg_search: indexes operate on base text
|
||||
# columns directly, so the dummy search_vector column is left NULL.
|
||||
# Native: the migration p4q5r6s7t8u9 dropped the GENERATED expression on
|
||||
# search_vector to allow per-deployment language configuration; the
|
||||
# batch insert path in ops_postgresql.insert_facts_batch now populates
|
||||
# it via to_tsvector($lang, ...). This single-observation INSERT does
|
||||
# not, so observations under the native backend currently land with
|
||||
# NULL search_vector and are not BM25-searchable until reflected/
|
||||
# re-ingested. Tracking a separate fix for that gap.
|
||||
elif config.text_search_extension == "native":
|
||||
# Native: search_vector is populated with to_tsvector() using the
|
||||
# configured native language dictionary, matching the batch insert
|
||||
# path in ops_postgresql.insert_facts_batch.
|
||||
query = f"""
|
||||
INSERT INTO {fq_table("memory_units")} (
|
||||
id, bank_id, text, fact_type, embedding, proof_count, source_memory_ids,
|
||||
tags, event_date, occurred_start, occurred_end, mentioned_at, search_vector
|
||||
)
|
||||
VALUES ($1, $2, $3, 'observation', $4::vector, 1, $5, $6, $7, $8, $9, $10,
|
||||
to_tsvector('{config.text_search_extension_native_language}'::regconfig, COALESCE($3, '')))
|
||||
RETURNING id
|
||||
"""
|
||||
else: # pg_textsearch, pgroonga, pg_search: indexes operate on base text columns directly
|
||||
query = f"""
|
||||
INSERT INTO {fq_table("memory_units")} (
|
||||
id, bank_id, text, fact_type, embedding, proof_count, source_memory_ids,
|
||||
|
||||
@@ -7,11 +7,13 @@ Configuration via environment variables - see hindsight_api.config for all env v
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import gc
|
||||
import logging
|
||||
import os
|
||||
import warnings
|
||||
from abc import ABC, abstractmethod
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -45,6 +47,7 @@ from ..config import (
|
||||
ENV_RERANKER_TEI_URL,
|
||||
ENV_RERANKER_ZEROENTROPY_API_KEY,
|
||||
)
|
||||
from .bank_attribution import reranker_bank_attribution_headers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -86,6 +89,12 @@ def _resolve_malloc_trim():
|
||||
_malloc_trim = _resolve_malloc_trim()
|
||||
|
||||
|
||||
def _release_rerank_heap() -> None:
|
||||
"""Release transient Python and native heap memory after local reranking."""
|
||||
gc.collect()
|
||||
_malloc_trim()
|
||||
|
||||
|
||||
class CrossEncoderModel(ABC):
|
||||
"""
|
||||
Abstract base class for cross-encoder reranking.
|
||||
@@ -212,7 +221,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
device = "cpu"
|
||||
logger.info("Reranker: forcing CPU mode (HINDSIGHT_API_RERANKER_LOCAL_FORCE_CPU=1)")
|
||||
else:
|
||||
# Check for GPU (CUDA) or Apple Silicon (MPS)
|
||||
# Check for GPU (CUDA), Apple Silicon (MPS), or Intel XPU
|
||||
# Wrap in try-except to gracefully handle any device detection issues
|
||||
# (e.g., in CI environments or when PyTorch is built without GPU support)
|
||||
device = "cpu" # Default to CPU
|
||||
@@ -220,10 +229,13 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
has_gpu = torch.cuda.is_available() or (
|
||||
hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
|
||||
)
|
||||
# Intel Arc XPU support — torch.xpu is available when the XPU build is loaded
|
||||
if not has_gpu and hasattr(torch, "xpu"):
|
||||
has_gpu = torch.xpu.is_available()
|
||||
if has_gpu:
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS/XPU
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to detect GPU/MPS, falling back to CPU: {e}")
|
||||
logger.warning(f"Failed to detect GPU/MPS/XPU, falling back to CPU: {e}")
|
||||
|
||||
# Patch transformers 5.x compatibility for models using XLM-RoBERTa
|
||||
# (e.g., jina-reranker-v2-base-multilingual). transformers 5.x removed
|
||||
@@ -312,7 +324,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
scores = self._model.predict(pairs, batch_size=self.batch_size, show_progress_bar=False)
|
||||
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
|
||||
finally:
|
||||
_malloc_trim()
|
||||
_release_rerank_heap()
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
@@ -481,6 +493,7 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
semaphore,
|
||||
"POST",
|
||||
f"{self.base_url}/rerank",
|
||||
headers=reranker_bank_attribution_headers(),
|
||||
json={
|
||||
"query": query,
|
||||
"texts": texts,
|
||||
@@ -621,7 +634,11 @@ class _CohereCompatibleRerankClient:
|
||||
if self.include_top_n:
|
||||
body["top_n"] = len(texts)
|
||||
|
||||
response = await self._async_client.post(self.rerank_url, json=body)
|
||||
response = await self._async_client.post(
|
||||
self.rerank_url,
|
||||
headers=reranker_bank_attribution_headers(),
|
||||
json=body,
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
@@ -987,11 +1004,11 @@ class FlashRankCrossEncoder(CrossEncoderModel):
|
||||
|
||||
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""Synchronous predict - processes each query group."""
|
||||
from flashrank import RerankRequest
|
||||
|
||||
if not pairs:
|
||||
return []
|
||||
|
||||
from flashrank import RerankRequest
|
||||
|
||||
try:
|
||||
# Group pairs by query
|
||||
query_groups: dict[str, list[tuple[int, str]]] = {}
|
||||
@@ -1020,7 +1037,7 @@ class FlashRankCrossEncoder(CrossEncoderModel):
|
||||
|
||||
return all_scores
|
||||
finally:
|
||||
_malloc_trim()
|
||||
_release_rerank_heap()
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
@@ -1148,6 +1165,7 @@ class LiteLLMCrossEncoder(CrossEncoderModel):
|
||||
# LiteLLM /rerank follows Cohere API format
|
||||
response = await self._async_client.post(
|
||||
f"{self.api_base}/rerank",
|
||||
headers=reranker_bank_attribution_headers(),
|
||||
json={
|
||||
"model": self.model,
|
||||
"query": query,
|
||||
@@ -1266,10 +1284,11 @@ class LiteLLMSDKCrossEncoder(CrossEncoderModel):
|
||||
indices = [idx for idx, _ in indexed_texts]
|
||||
|
||||
# Build kwargs for rerank call
|
||||
rerank_kwargs = {
|
||||
rerank_kwargs: dict[str, Any] = {
|
||||
"model": self.model,
|
||||
"query": query,
|
||||
"documents": texts,
|
||||
"headers": reranker_bank_attribution_headers(),
|
||||
}
|
||||
if self.api_key:
|
||||
rerank_kwargs["api_key"] = self.api_key
|
||||
@@ -1278,21 +1297,9 @@ class LiteLLMSDKCrossEncoder(CrossEncoderModel):
|
||||
|
||||
response = await self._litellm.arerank(**rerank_kwargs)
|
||||
|
||||
# Map scores back to original positions
|
||||
# Response format: RerankResponse with results list
|
||||
# Each result is a TypedDict with "index" and "relevance_score"
|
||||
if hasattr(response, "results") and response.results:
|
||||
for result in response.results:
|
||||
# Results are TypedDicts, use dict-style access
|
||||
original_idx = result["index"]
|
||||
score = result.get("relevance_score", result.get("score", 0.0))
|
||||
all_scores[indices[original_idx]] = score
|
||||
elif isinstance(response, list):
|
||||
# Direct list of scores (unlikely but defensive)
|
||||
for i, score in enumerate(response):
|
||||
all_scores[indices[i]] = score
|
||||
else:
|
||||
logger.warning(f"Unexpected response format from LiteLLM rerank: {type(response)}")
|
||||
for result in response.results:
|
||||
original_idx = result["index"]
|
||||
all_scores[indices[original_idx]] = result["relevance_score"]
|
||||
|
||||
return all_scores
|
||||
|
||||
|
||||
@@ -307,6 +307,17 @@ class DatabaseBackend(ABC):
|
||||
"""Close the connection pool and release all resources."""
|
||||
...
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def is_ready(self) -> bool:
|
||||
"""Whether the pool exists and can serve connections.
|
||||
|
||||
False before :meth:`initialize` and after :meth:`shutdown`. Best-effort
|
||||
callers (tracing, auditing) check this to skip work during those windows
|
||||
instead of acquiring and interpreting the resulting error.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
@asynccontextmanager
|
||||
async def acquire(self) -> AsyncIterator[DatabaseConnection]:
|
||||
|
||||
@@ -18,6 +18,7 @@ and mirrors Django's ``DatabaseOperations`` architecture.
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from .base import DatabaseConnection
|
||||
@@ -172,6 +173,25 @@ class DataAccessOps(ABC):
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def bulk_reassert_entities(
|
||||
self,
|
||||
conn: DatabaseConnection,
|
||||
table: str,
|
||||
bank_id: str,
|
||||
entity_ids: list[str],
|
||||
canonical_names: list[str],
|
||||
) -> None:
|
||||
"""Lock resolved parents and re-create any pruned since Phase-1 resolution.
|
||||
|
||||
Closes the retain Phase-1/prune race (#2662): existing rows are locked
|
||||
(PG ``FOR KEY SHARE`` / Oracle ``FOR UPDATE``) so a concurrent
|
||||
``prune_orphan_entities`` blocks until the caller's transaction commits,
|
||||
while rows already deleted are re-inserted idempotently. ``entity_ids``
|
||||
must be sorted by the caller for a stable lock order.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def bulk_insert_unit_entities(
|
||||
self,
|
||||
@@ -484,6 +504,23 @@ class DataAccessOps(ABC):
|
||||
|
||||
# -- Task claiming operations ------------------------------------------
|
||||
|
||||
@abstractmethod
|
||||
async def prune_terminal_operations(
|
||||
self,
|
||||
conn: DatabaseConnection,
|
||||
table: str,
|
||||
cutoff: datetime,
|
||||
*,
|
||||
batch_size: int,
|
||||
) -> int:
|
||||
"""Delete one deterministic batch of terminal operations older than ``cutoff``.
|
||||
|
||||
Implementations must lock candidates without waiting on rows another
|
||||
worker is pruning, never select pending/processing rows, and return the
|
||||
number deleted. The caller provides a transaction around this method.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def claim_tasks(
|
||||
self,
|
||||
|
||||
@@ -13,6 +13,8 @@ from .base import DatabaseConnection
|
||||
from .ops import DataAccessOps, TagListingParts
|
||||
from .result import DictResultRow as ResultRow
|
||||
|
||||
ORACLE_IN_LIST_LIMIT = 1000
|
||||
|
||||
|
||||
class OracleOps(DataAccessOps):
|
||||
"""Oracle-specific data access operations."""
|
||||
@@ -216,7 +218,7 @@ class OracleOps(DataAccessOps):
|
||||
for orig_name in missing_names:
|
||||
row = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT id, LOWER(canonical_name) AS name_lower
|
||||
SELECT id, canonical_name, LOWER(canonical_name) AS name_lower
|
||||
FROM {table}
|
||||
WHERE bank_id = $1 AND LOWER(canonical_name) = LOWER($2)
|
||||
""",
|
||||
@@ -224,10 +226,37 @@ class OracleOps(DataAccessOps):
|
||||
orig_name,
|
||||
)
|
||||
if row:
|
||||
# Wrap in a dict-like to include input_name for downstream compat
|
||||
results.append(row)
|
||||
return results
|
||||
|
||||
async def bulk_reassert_entities(
|
||||
self,
|
||||
conn: DatabaseConnection,
|
||||
table: str,
|
||||
bank_id: str,
|
||||
entity_ids: list[str],
|
||||
canonical_names: list[str],
|
||||
) -> None:
|
||||
# Oracle has no FOR KEY SHARE; FOR UPDATE is the row-lock equivalent that
|
||||
# blocks a concurrent prune DELETE until this transaction commits. Lock
|
||||
# each surviving parent in the caller's stable id order (pruned ids are
|
||||
# simply absent here), then re-insert any that vanished. The translation
|
||||
# layer rewrites ON CONFLICT DO NOTHING to strip-and-catch ORA-00001, so
|
||||
# a name recreated under a new id is suppressed rather than raising.
|
||||
for entity_id in entity_ids:
|
||||
await conn.fetchrow(
|
||||
f"SELECT id FROM {table} WHERE id = $1 FOR UPDATE",
|
||||
entity_id,
|
||||
)
|
||||
await conn.executemany(
|
||||
f"""
|
||||
INSERT INTO {table} (id, bank_id, canonical_name)
|
||||
VALUES ($1, $2, $3)
|
||||
ON CONFLICT DO NOTHING
|
||||
""",
|
||||
[(entity_id, bank_id, canonical_name) for entity_id, canonical_name in zip(entity_ids, canonical_names)],
|
||||
)
|
||||
|
||||
async def bulk_insert_unit_entities(
|
||||
self,
|
||||
conn: DatabaseConnection,
|
||||
@@ -256,13 +285,21 @@ class OracleOps(DataAccessOps):
|
||||
# Oracle doesn't support ON CONFLICT; rely on the PK and the
|
||||
# IGNORE_ROW_ON_DUPKEY_INDEX hint to skip duplicates server-side.
|
||||
# The hint name must match the PK constraint exactly.
|
||||
#
|
||||
# Sort to enforce a global lock-acquisition order on the
|
||||
# (bank_id, unit_id) PK. Without this, two concurrent
|
||||
# transactions inserting overlapping unit_id sets in different
|
||||
# orders can deadlock on the unique-check row locks. Sorting
|
||||
# gives every concurrent caller the same lock order, so
|
||||
# conflicting inserts queue cleanly instead of cycling.
|
||||
sorted_unit_ids = sorted(unit_ids)
|
||||
await conn.executemany(
|
||||
f"""
|
||||
INSERT /*+ IGNORE_ROW_ON_DUPKEY_INDEX({table}, pk_graph_maintenance_queue) */
|
||||
INTO {table} (bank_id, unit_id)
|
||||
VALUES ($1, $2)
|
||||
""",
|
||||
[(bank_id, uid) for uid in unit_ids],
|
||||
[(bank_id, uid) for uid in sorted_unit_ids],
|
||||
)
|
||||
|
||||
async def claim_graph_maintenance_batch(
|
||||
@@ -321,6 +358,12 @@ class OracleOps(DataAccessOps):
|
||||
entities_table: str,
|
||||
bank_id: str,
|
||||
) -> int:
|
||||
# NB: the Postgres path additionally selects victims FOR UPDATE in sorted
|
||||
# (entity_id_1, entity_id_2) order to prevent the #2529 deadlock against
|
||||
# retain's sorted cooccurrence upsert. Oracle's DELETE can't carry that
|
||||
# ordered-lock CTE the same way, so here we rely on the Pass 2/3 retry
|
||||
# wrap in run_graph_maintenance_job (retry_with_backoff is ORA-00060
|
||||
# deadlock-aware) to recover instead. Deliberate dialect asymmetry.
|
||||
deleted = await conn.execute(
|
||||
f"""
|
||||
DELETE FROM {ec_table}
|
||||
@@ -439,6 +482,14 @@ class OracleOps(DataAccessOps):
|
||||
FROM {ue_table} ue_target
|
||||
WHERE ue_target.entity_id = se.entity_id
|
||||
AND ue_target.unit_id != ALL($1::uuid[])
|
||||
-- Filter before applying the cap: candidates from other fact
|
||||
-- types must not consume this entity's bounded fan-out.
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM {mu_table} mu_target
|
||||
WHERE mu_target.id = ue_target.unit_id
|
||||
AND mu_target.fact_type = $2
|
||||
)
|
||||
ORDER BY ue_target.unit_id DESC
|
||||
FETCH FIRST {per_entity_limit} ROWS ONLY
|
||||
) t
|
||||
@@ -451,7 +502,6 @@ class OracleOps(DataAccessOps):
|
||||
es.score, 'entity' AS source
|
||||
FROM entity_scores es
|
||||
JOIN {mu_table} mu ON mu.id = es.unit_id
|
||||
WHERE mu.fact_type = $2
|
||||
ORDER BY es.score DESC
|
||||
FETCH FIRST $3 ROWS ONLY
|
||||
)"""
|
||||
@@ -816,6 +866,157 @@ class OracleOps(DataAccessOps):
|
||||
|
||||
# -- Task claiming operations ------------------------------------------
|
||||
|
||||
async def prune_terminal_operations(
|
||||
self,
|
||||
conn: DatabaseConnection,
|
||||
table: str,
|
||||
cutoff: datetime,
|
||||
*,
|
||||
batch_size: int,
|
||||
) -> int:
|
||||
# Oracle rejects a row-limited SELECT ... FOR UPDATE (ORA-02014). Pick
|
||||
# the deterministic bounded IDs first, then lock only that candidate
|
||||
# set and re-check eligibility before deleting in the same transaction.
|
||||
# Clamp to Oracle's 1000-expression IN-list limit because the adapter
|
||||
# expands the candidate UUID list into individual bind variables.
|
||||
# Cancelled children cannot complete parent aggregation, so retain the
|
||||
# parent guard only for completed/failed children. Before removing a
|
||||
# cancelled child, preserve its signal by cancelling a pending parent
|
||||
# in this transaction and refreshing the parent's retention window.
|
||||
# Validate metadata before HEXTORAW: CASE makes malformed UUIDs yield
|
||||
# NULL while keeping the indexed RAW parent.operation_id key unwrapped.
|
||||
effective_batch_size = min(batch_size, ORACLE_IN_LIST_LIMIT)
|
||||
candidates = await conn.fetch(
|
||||
f"""
|
||||
SELECT candidate_operation.operation_id
|
||||
FROM {table} candidate_operation
|
||||
WHERE candidate_operation.status IN ('completed', 'failed', 'cancelled')
|
||||
AND candidate_operation.updated_at < $1
|
||||
AND (
|
||||
candidate_operation.status = 'cancelled'
|
||||
OR NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM {table} parent
|
||||
WHERE parent.operation_id = CASE
|
||||
WHEN REGEXP_LIKE(
|
||||
JSON_VALUE(
|
||||
candidate_operation.result_metadata,
|
||||
'$.parent_operation_id' RETURNING VARCHAR2(36) NULL ON ERROR
|
||||
),
|
||||
'^[0-9A-Fa-f]{{8}}-[0-9A-Fa-f]{{4}}-[0-9A-Fa-f]{{4}}-[0-9A-Fa-f]{{4}}-[0-9A-Fa-f]{{12}}$'
|
||||
)
|
||||
THEN HEXTORAW(REPLACE(
|
||||
JSON_VALUE(
|
||||
candidate_operation.result_metadata,
|
||||
'$.parent_operation_id' RETURNING VARCHAR2(36) NULL ON ERROR
|
||||
),
|
||||
'-',
|
||||
''
|
||||
))
|
||||
ELSE NULL
|
||||
END
|
||||
AND parent.bank_id = candidate_operation.bank_id
|
||||
)
|
||||
)
|
||||
ORDER BY candidate_operation.updated_at, candidate_operation.operation_id
|
||||
LIMIT $2
|
||||
""",
|
||||
cutoff,
|
||||
effective_batch_size,
|
||||
)
|
||||
if not candidates:
|
||||
return 0
|
||||
|
||||
candidate_ids = [row["operation_id"] for row in candidates]
|
||||
locked = await conn.fetch(
|
||||
f"""
|
||||
SELECT candidate_operation.operation_id
|
||||
FROM {table} candidate_operation
|
||||
WHERE candidate_operation.operation_id = ANY($1)
|
||||
AND candidate_operation.status IN ('completed', 'failed', 'cancelled')
|
||||
AND candidate_operation.updated_at < $2
|
||||
AND (
|
||||
candidate_operation.status = 'cancelled'
|
||||
OR NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM {table} parent
|
||||
WHERE parent.operation_id = CASE
|
||||
WHEN REGEXP_LIKE(
|
||||
JSON_VALUE(
|
||||
candidate_operation.result_metadata,
|
||||
'$.parent_operation_id' RETURNING VARCHAR2(36) NULL ON ERROR
|
||||
),
|
||||
'^[0-9A-Fa-f]{{8}}-[0-9A-Fa-f]{{4}}-[0-9A-Fa-f]{{4}}-[0-9A-Fa-f]{{4}}-[0-9A-Fa-f]{{12}}$'
|
||||
)
|
||||
THEN HEXTORAW(REPLACE(
|
||||
JSON_VALUE(
|
||||
candidate_operation.result_metadata,
|
||||
'$.parent_operation_id' RETURNING VARCHAR2(36) NULL ON ERROR
|
||||
),
|
||||
'-',
|
||||
''
|
||||
))
|
||||
ELSE NULL
|
||||
END
|
||||
AND parent.bank_id = candidate_operation.bank_id
|
||||
)
|
||||
)
|
||||
ORDER BY candidate_operation.updated_at, candidate_operation.operation_id
|
||||
FOR UPDATE OF candidate_operation.operation_id SKIP LOCKED
|
||||
""",
|
||||
candidate_ids,
|
||||
cutoff,
|
||||
)
|
||||
if not locked:
|
||||
return 0
|
||||
operation_ids = [row["operation_id"] for row in locked]
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {table} parent
|
||||
SET status = 'cancelled',
|
||||
updated_at = now(),
|
||||
completed_at = COALESCE(parent.completed_at, now()),
|
||||
error_message = COALESCE(
|
||||
parent.error_message,
|
||||
'Cancelled because a child operation was cancelled'
|
||||
)
|
||||
WHERE parent.status = 'pending'
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM {table} candidate_operation
|
||||
WHERE candidate_operation.operation_id = ANY($1)
|
||||
AND candidate_operation.status = 'cancelled'
|
||||
AND candidate_operation.updated_at < $2
|
||||
AND candidate_operation.bank_id = parent.bank_id
|
||||
AND parent.operation_id = CASE
|
||||
WHEN REGEXP_LIKE(
|
||||
JSON_VALUE(
|
||||
candidate_operation.result_metadata,
|
||||
'$.parent_operation_id' RETURNING VARCHAR2(36) NULL ON ERROR
|
||||
),
|
||||
'^[0-9A-Fa-f]{{8}}-[0-9A-Fa-f]{{4}}-[0-9A-Fa-f]{{4}}-[0-9A-Fa-f]{{4}}-[0-9A-Fa-f]{{12}}$'
|
||||
)
|
||||
THEN HEXTORAW(REPLACE(
|
||||
JSON_VALUE(
|
||||
candidate_operation.result_metadata,
|
||||
'$.parent_operation_id' RETURNING VARCHAR2(36) NULL ON ERROR
|
||||
),
|
||||
'-',
|
||||
''
|
||||
))
|
||||
ELSE NULL
|
||||
END
|
||||
)
|
||||
""",
|
||||
operation_ids,
|
||||
cutoff,
|
||||
)
|
||||
await conn.execute(
|
||||
f"DELETE FROM {table} WHERE operation_id = ANY($1)",
|
||||
operation_ids,
|
||||
)
|
||||
return len(operation_ids)
|
||||
|
||||
async def _claim_consolidation_tasks(
|
||||
self,
|
||||
conn,
|
||||
|
||||
@@ -4,11 +4,40 @@ Uses unnest(), LATERAL, DISTINCT ON, and native array operations for
|
||||
efficient batch operations.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from .base import DatabaseConnection
|
||||
from .ops import DataAccessOps, TagListingParts
|
||||
from .result import ResultRow
|
||||
|
||||
|
||||
def pg_search_vector_expr(
|
||||
config,
|
||||
*,
|
||||
text_col: str = "text",
|
||||
context_col: str = "context",
|
||||
signals_col: str = "text_signals",
|
||||
) -> str | None:
|
||||
"""SQL expression that builds ``search_vector`` for the configured PG text-search backend.
|
||||
|
||||
Single source of truth shared by the batch insert (over the ``input_data``
|
||||
CTE columns) and the curation revert recompute (over a ``memory_units`` row),
|
||||
so the two can never drift. Returns ``None`` for backends that leave
|
||||
``search_vector`` unpopulated — pgroonga / pg_textsearch / pg_search index the
|
||||
base text columns directly and keep only a dummy column, so there is nothing
|
||||
to build.
|
||||
|
||||
``text_search_extension_native_language`` is validated as a PG identifier in
|
||||
``HindsightConfig.validate()``, so embedding it as a SQL literal is safe.
|
||||
"""
|
||||
combined = f"COALESCE({text_col}, '') || ' ' || COALESCE({context_col}, '') || ' ' || COALESCE({signals_col}, '')"
|
||||
if config.text_search_extension == "vchord":
|
||||
return f"tokenize({combined}, 'llmlingua2')::bm25_catalog.bm25vector"
|
||||
if config.text_search_extension == "native":
|
||||
return f"to_tsvector('{config.text_search_extension_native_language}'::regconfig, {combined})"
|
||||
return None
|
||||
|
||||
|
||||
class PostgreSQLOps(DataAccessOps):
|
||||
"""PostgreSQL-specific data access operations using unnest and LATERAL."""
|
||||
|
||||
@@ -93,101 +122,39 @@ class PostgreSQLOps(DataAccessOps):
|
||||
config = get_config()
|
||||
table = self._get_mu_table()
|
||||
|
||||
if config.text_search_extension == "vchord":
|
||||
query = f"""
|
||||
WITH input_data AS (
|
||||
SELECT * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::jsonb[], $11::text[], $12::text[], $13::jsonb[], $14::jsonb[], $15::text[]
|
||||
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id, tags_json,
|
||||
observation_scopes_json, text_signals)
|
||||
)
|
||||
INSERT INTO {table} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id, tags,
|
||||
observation_scopes, text_signals, search_vector)
|
||||
SELECT
|
||||
$1,
|
||||
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id,
|
||||
COALESCE(
|
||||
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
|
||||
'{{}}'::varchar[]
|
||||
),
|
||||
observation_scopes_json,
|
||||
text_signals,
|
||||
tokenize(
|
||||
COALESCE(text, '') || ' ' || COALESCE(context, '') || ' ' || COALESCE(text_signals, ''),
|
||||
'llmlingua2'
|
||||
)::bm25_catalog.bm25vector
|
||||
FROM input_data
|
||||
RETURNING id
|
||||
"""
|
||||
elif config.text_search_extension == "native":
|
||||
# search_vector is a regular tsvector column populated here using the
|
||||
# configured native dictionary. It used to be GENERATED ALWAYS with
|
||||
# a hardcoded 'english', which prevented per-deployment language
|
||||
# configuration. text_search_extension_native_language is validated
|
||||
# in HindsightConfig.validate() as a PG identifier, so embedding it
|
||||
# as a SQL literal is safe.
|
||||
query = f"""
|
||||
WITH input_data AS (
|
||||
SELECT * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::jsonb[], $11::text[], $12::text[], $13::jsonb[], $14::jsonb[], $15::text[]
|
||||
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id, tags_json,
|
||||
observation_scopes_json, text_signals)
|
||||
)
|
||||
INSERT INTO {table} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id, tags,
|
||||
observation_scopes, text_signals, search_vector)
|
||||
SELECT
|
||||
$1,
|
||||
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id,
|
||||
COALESCE(
|
||||
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
|
||||
'{{}}'::varchar[]
|
||||
),
|
||||
observation_scopes_json,
|
||||
text_signals,
|
||||
to_tsvector(
|
||||
'{config.text_search_extension_native_language}'::regconfig,
|
||||
COALESCE(text, '') || ' ' || COALESCE(context, '') || ' ' || COALESCE(text_signals, '')
|
||||
)
|
||||
FROM input_data
|
||||
RETURNING id
|
||||
"""
|
||||
else:
|
||||
# pg_textsearch, pgroonga, and pg_search: search_vector is a dummy
|
||||
# TEXT column; the actual full-text index operates on the base text
|
||||
# columns directly, so we don't populate search_vector at insert time.
|
||||
query = f"""
|
||||
WITH input_data AS (
|
||||
SELECT * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::jsonb[], $11::text[], $12::text[], $13::jsonb[], $14::jsonb[], $15::text[]
|
||||
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id, tags_json,
|
||||
observation_scopes_json, text_signals)
|
||||
)
|
||||
INSERT INTO {table} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id, tags,
|
||||
observation_scopes, text_signals)
|
||||
SELECT
|
||||
$1,
|
||||
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id,
|
||||
COALESCE(
|
||||
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
|
||||
'{{}}'::varchar[]
|
||||
),
|
||||
observation_scopes_json,
|
||||
text_signals
|
||||
FROM input_data
|
||||
RETURNING id
|
||||
"""
|
||||
# search_vector is populated inline for backends that store a real vector
|
||||
# (native tsvector, vchord bm25vector). pgroonga / pg_textsearch / pg_search
|
||||
# index the base text columns directly and keep only a dummy column, so the
|
||||
# expression is None and the column is left out of the insert entirely.
|
||||
# Same expression is reused by curation revert (see pg_search_vector_expr).
|
||||
sv_expr = pg_search_vector_expr(config)
|
||||
sv_insert_col = ", search_vector" if sv_expr else ""
|
||||
sv_select_val = f",\n {sv_expr}" if sv_expr else ""
|
||||
query = f"""
|
||||
WITH input_data AS (
|
||||
SELECT * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::jsonb[], $11::text[], $12::text[], $13::jsonb[], $14::jsonb[], $15::text[]
|
||||
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id, tags_json,
|
||||
observation_scopes_json, text_signals)
|
||||
)
|
||||
INSERT INTO {table} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id, tags,
|
||||
observation_scopes, text_signals{sv_insert_col})
|
||||
SELECT
|
||||
$1,
|
||||
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, metadata, chunk_id, document_id,
|
||||
COALESCE(
|
||||
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
|
||||
'{{}}'::varchar[]
|
||||
),
|
||||
observation_scopes_json,
|
||||
text_signals{sv_select_val}
|
||||
FROM input_data
|
||||
RETURNING id
|
||||
"""
|
||||
|
||||
results = await conn.fetch(
|
||||
query,
|
||||
@@ -310,7 +277,7 @@ class PostgreSQLOps(DataAccessOps):
|
||||
) -> list[ResultRow]:
|
||||
return await conn.fetch(
|
||||
f"""
|
||||
SELECT e.id, LOWER(e.canonical_name) AS name_lower, inputs.input_name
|
||||
SELECT e.id, e.canonical_name, LOWER(e.canonical_name) AS name_lower, inputs.input_name
|
||||
FROM {table} e
|
||||
JOIN (
|
||||
SELECT LOWER(n) AS input_name_lower, n AS input_name
|
||||
@@ -322,6 +289,42 @@ class PostgreSQLOps(DataAccessOps):
|
||||
missing_names,
|
||||
)
|
||||
|
||||
async def bulk_reassert_entities(
|
||||
self,
|
||||
conn: DatabaseConnection,
|
||||
table: str,
|
||||
bank_id: str,
|
||||
entity_ids: list[str],
|
||||
canonical_names: list[str],
|
||||
) -> None:
|
||||
# One statement, one round-trip (same shape as bulk_insert_links):
|
||||
# * the CTE takes FOR KEY SHARE on every parent that still exists,
|
||||
# held to COMMIT, so a concurrent prune_orphan_entities DELETE blocks
|
||||
# until the caller's unit_entities insert has committed;
|
||||
# * the INSERT re-creates only the parents that were already pruned
|
||||
# (NOT IN locked), carrying the canonical_name resolved in Phase 1.
|
||||
# ON CONFLICT DO NOTHING (no target) keeps the rare case where another
|
||||
# worker recreated the name under a new id from raising — that row stays
|
||||
# absent and its unit link is the sole casualty, never the whole batch.
|
||||
await conn.execute(
|
||||
f"""
|
||||
WITH locked AS (
|
||||
SELECT id FROM {table}
|
||||
WHERE id = ANY($2::uuid[])
|
||||
ORDER BY id
|
||||
FOR KEY SHARE
|
||||
)
|
||||
INSERT INTO {table} (id, bank_id, canonical_name)
|
||||
SELECT t.entity_id, $1, t.canonical_name
|
||||
FROM unnest($2::uuid[], $3::text[]) AS t(entity_id, canonical_name)
|
||||
WHERE t.entity_id NOT IN (SELECT id FROM locked)
|
||||
ON CONFLICT DO NOTHING
|
||||
""",
|
||||
bank_id,
|
||||
entity_ids,
|
||||
canonical_names,
|
||||
)
|
||||
|
||||
async def bulk_insert_unit_entities(
|
||||
self,
|
||||
conn: DatabaseConnection,
|
||||
@@ -348,6 +351,15 @@ class PostgreSQLOps(DataAccessOps):
|
||||
) -> None:
|
||||
if not unit_ids:
|
||||
return
|
||||
# Sort to enforce a global lock-acquisition order on the
|
||||
# (bank_id, unit_id) unique-key. Without this, two concurrent
|
||||
# transactions inserting overlapping unit_id sets in different
|
||||
# orders can deadlock on the ON CONFLICT row locks — Postgres
|
||||
# acquires a short-lived lock per row being checked, and cycle
|
||||
# detection then aborts one transaction. Sorting gives every
|
||||
# concurrent caller the same lock order, so conflicting inserts
|
||||
# queue cleanly instead of cycling.
|
||||
sorted_unit_ids = sorted(unit_ids)
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {table} (bank_id, unit_id)
|
||||
@@ -355,7 +367,7 @@ class PostgreSQLOps(DataAccessOps):
|
||||
ON CONFLICT (bank_id, unit_id) DO NOTHING
|
||||
""",
|
||||
bank_id,
|
||||
unit_ids,
|
||||
sorted_unit_ids,
|
||||
)
|
||||
|
||||
async def claim_graph_maintenance_batch(
|
||||
@@ -416,19 +428,48 @@ class PostgreSQLOps(DataAccessOps):
|
||||
# Scope by joining through entities.bank_id (entity_cooccurrences itself
|
||||
# has no bank_id column — entities don't span banks, so scoping via
|
||||
# entity_id_1 is sufficient).
|
||||
#
|
||||
# Ordered locking (deadlock avoidance, #2529): retain's concurrent
|
||||
# cooccurrence upsert (entity_resolver._flush_pending) locks rows in
|
||||
# sorted (entity_id_1, entity_id_2) order — sorted specifically to give
|
||||
# every writer one consistent lock-acquisition order. A plain
|
||||
# `DELETE ... USING` scans/locks in whatever order the join plan picks,
|
||||
# so it could lock the same rows in the opposite order and cycle. We
|
||||
# instead select the victims in that same sorted order `FOR UPDATE`
|
||||
# first — the locking clause materialises the CTE and places LockRows
|
||||
# above the Sort, so locks are acquired ascending, matching the upsert —
|
||||
# then delete the already-locked rows. Same order on both sides ⇒ no
|
||||
# cycle (the deadlock is prevented, not merely retried). The Pass 2/3
|
||||
# retry wrap in run_graph_maintenance_job stays as a backstop for the
|
||||
# residual paths (FK cascade from prune_orphan_entities, Oracle).
|
||||
#
|
||||
# The staleness predicate is an INTERSECT of the two entities' unit sets
|
||||
# rather than the equivalent `unit_entities u1 JOIN u2 ON u1.unit_id =
|
||||
# u2.unit_id` self-join (#2473): both INTERSECT branches resolve as Index
|
||||
# Only Scans on idx_unit_entities_entity_unit (entity_id, unit_id), so the
|
||||
# per-pair cost is bounded by the two entities' degrees. The self-join let
|
||||
# the planner pick an anti-join that rescanned a high-degree hub entity's
|
||||
# membership set for every pair — 28-30min on a bank with a ~100K-membership
|
||||
# hub, even when zero rows were stale. Don't "simplify" it back.
|
||||
result = await conn.execute(
|
||||
f"""
|
||||
WITH victims AS (
|
||||
SELECT c.entity_id_1, c.entity_id_2
|
||||
FROM {ec_table} c
|
||||
JOIN {entities_table} e ON e.id = c.entity_id_1
|
||||
WHERE e.bank_id = $1
|
||||
AND NOT EXISTS (
|
||||
SELECT unit_id FROM {ue_table} WHERE entity_id = c.entity_id_1
|
||||
INTERSECT
|
||||
SELECT unit_id FROM {ue_table} WHERE entity_id = c.entity_id_2
|
||||
)
|
||||
ORDER BY c.entity_id_1, c.entity_id_2
|
||||
FOR UPDATE OF c
|
||||
)
|
||||
DELETE FROM {ec_table} c
|
||||
USING {entities_table} e
|
||||
WHERE e.id = c.entity_id_1
|
||||
AND e.bank_id = $1
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM {ue_table} u1
|
||||
JOIN {ue_table} u2 ON u1.unit_id = u2.unit_id
|
||||
WHERE u1.entity_id = c.entity_id_1
|
||||
AND u2.entity_id = c.entity_id_2
|
||||
)
|
||||
USING victims v
|
||||
WHERE c.entity_id_1 = v.entity_id_1
|
||||
AND c.entity_id_2 = v.entity_id_2
|
||||
""",
|
||||
bank_id,
|
||||
)
|
||||
@@ -532,11 +573,18 @@ class PostgreSQLOps(DataAccessOps):
|
||||
FROM {ue_table} ue_target
|
||||
WHERE ue_target.entity_id = se.entity_id
|
||||
AND ue_target.unit_id != ALL($1::uuid[])
|
||||
-- Filter before applying the cap: candidates from other fact
|
||||
-- types must not consume this entity's bounded fan-out.
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM {mu_table} mu_target
|
||||
WHERE mu_target.id = ue_target.unit_id
|
||||
AND mu_target.fact_type = $2
|
||||
)
|
||||
ORDER BY ue_target.unit_id DESC
|
||||
LIMIT {per_entity_limit}
|
||||
) t
|
||||
JOIN {mu_table} mu ON mu.id = t.unit_id
|
||||
WHERE mu.fact_type = $2
|
||||
GROUP BY mu.id
|
||||
ORDER BY score DESC
|
||||
LIMIT $3
|
||||
@@ -752,10 +800,16 @@ class PostgreSQLOps(DataAccessOps):
|
||||
internal_id: str,
|
||||
fact_types: dict[str, str],
|
||||
) -> None:
|
||||
# CONCURRENTLY so the drop takes ShareUpdateExclusive, not ACCESS
|
||||
# EXCLUSIVE, on the shared memory_units table. A plain DROP INDEX blocks
|
||||
# (and deadlocks with) every other bank's concurrent reads/writes on the
|
||||
# table; CONCURRENTLY does not conflict with DML. The caller
|
||||
# (delete_bank) runs this on an autocommit connection after its delete
|
||||
# transaction has committed — CONCURRENTLY cannot run inside a tx.
|
||||
for ft, suffix in fact_types.items():
|
||||
uid = str(internal_id).replace("-", "")[:16]
|
||||
idx = f"idx_mu_emb_{suffix}_{uid}"
|
||||
await conn.execute(f"DROP INDEX IF EXISTS {schema}.{idx}")
|
||||
await conn.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}.{idx}")
|
||||
|
||||
def get_entity_resolution_strategy(self) -> str:
|
||||
return "trigram"
|
||||
@@ -885,6 +939,93 @@ class PostgreSQLOps(DataAccessOps):
|
||||
|
||||
# -- Task claiming operations ------------------------------------------
|
||||
|
||||
async def prune_terminal_operations(
|
||||
self,
|
||||
conn: DatabaseConnection,
|
||||
table: str,
|
||||
cutoff: datetime,
|
||||
*,
|
||||
batch_size: int,
|
||||
) -> int:
|
||||
# Lock only the bounded candidate set. SKIP LOCKED lets multiple
|
||||
# workers prune disjoint batches without waiting or double-deleting.
|
||||
# Cancelled children cannot complete parent aggregation, so retain the
|
||||
# parent guard only for completed/failed children. Before removing a
|
||||
# cancelled child, preserve its signal by cancelling a pending parent
|
||||
# in this transaction and refreshing the parent's retention window.
|
||||
candidates = await conn.fetch(
|
||||
f"""
|
||||
SELECT candidate_operation.operation_id
|
||||
FROM {table} candidate_operation
|
||||
WHERE candidate_operation.status IN ('completed', 'failed', 'cancelled')
|
||||
AND candidate_operation.updated_at < $1
|
||||
AND (
|
||||
candidate_operation.status = 'cancelled'
|
||||
OR NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM {table} parent
|
||||
WHERE parent.operation_id = CASE
|
||||
WHEN candidate_operation.result_metadata->>'parent_operation_id'
|
||||
~* '^[0-9a-f]{{8}}-[0-9a-f]{{4}}-[0-9a-f]{{4}}-[0-9a-f]{{4}}-[0-9a-f]{{12}}$'
|
||||
THEN (candidate_operation.result_metadata->>'parent_operation_id')::uuid
|
||||
ELSE NULL
|
||||
END
|
||||
AND parent.bank_id = candidate_operation.bank_id
|
||||
)
|
||||
)
|
||||
ORDER BY candidate_operation.updated_at, candidate_operation.operation_id
|
||||
LIMIT $2
|
||||
FOR UPDATE OF candidate_operation SKIP LOCKED
|
||||
""",
|
||||
cutoff,
|
||||
batch_size,
|
||||
)
|
||||
if not candidates:
|
||||
return 0
|
||||
|
||||
candidate_ids = [row["operation_id"] for row in candidates]
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {table} parent
|
||||
SET status = 'cancelled',
|
||||
updated_at = now(),
|
||||
completed_at = COALESCE(parent.completed_at, now()),
|
||||
error_message = COALESCE(
|
||||
parent.error_message,
|
||||
'Cancelled because a child operation was cancelled'
|
||||
)
|
||||
WHERE parent.status = 'pending'
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM {table} candidate_operation
|
||||
WHERE candidate_operation.operation_id = ANY($1)
|
||||
AND candidate_operation.status = 'cancelled'
|
||||
AND candidate_operation.updated_at < $2
|
||||
AND candidate_operation.bank_id = parent.bank_id
|
||||
AND parent.operation_id = CASE
|
||||
WHEN candidate_operation.result_metadata->>'parent_operation_id'
|
||||
~* '^[0-9a-f]{{8}}-[0-9a-f]{{4}}-[0-9a-f]{{4}}-[0-9a-f]{{4}}-[0-9a-f]{{12}}$'
|
||||
THEN (candidate_operation.result_metadata->>'parent_operation_id')::uuid
|
||||
ELSE NULL
|
||||
END
|
||||
)
|
||||
""",
|
||||
candidate_ids,
|
||||
cutoff,
|
||||
)
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
DELETE FROM {table}
|
||||
WHERE operation_id = ANY($1)
|
||||
AND status IN ('completed', 'failed', 'cancelled')
|
||||
AND updated_at < $2
|
||||
RETURNING operation_id
|
||||
""",
|
||||
candidate_ids,
|
||||
cutoff,
|
||||
)
|
||||
return len(rows)
|
||||
|
||||
async def _claim_consolidation_tasks(
|
||||
self,
|
||||
conn,
|
||||
|
||||
@@ -23,6 +23,8 @@ from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any, NamedTuple
|
||||
|
||||
from .pool_instrumentation import PoolStats, acquire_conn
|
||||
|
||||
|
||||
class _OracleJSONEncoder(json.JSONEncoder):
|
||||
"""JSON encoder that handles datetime and UUID objects."""
|
||||
@@ -1242,6 +1244,11 @@ class OracleBackend(DatabaseBackend):
|
||||
def __init__(self) -> None:
|
||||
self._pool: Any = None
|
||||
self._oracledb: Any = None
|
||||
# Oracle pooled sessions retain CURRENT_SCHEMA across checkouts. Cache
|
||||
# SESSION_USER so default-schema acquisitions can explicitly reset a
|
||||
# connection that was previously used for a tenant schema.
|
||||
self._default_schema: str | None = None
|
||||
self._acquire_warn_threshold_s: float = 1.0
|
||||
|
||||
async def initialize(
|
||||
self,
|
||||
@@ -1257,6 +1264,10 @@ class OracleBackend(DatabaseBackend):
|
||||
oracledb = _import_oracledb()
|
||||
self._oracledb = oracledb
|
||||
|
||||
from ...config import get_config
|
||||
|
||||
self._acquire_warn_threshold_s = get_config().db_acquire_warn_threshold_ms / 1000.0
|
||||
|
||||
# Parse URL-format DSN (oracle://user:pass@host:port/service)
|
||||
from urllib.parse import urlparse
|
||||
|
||||
@@ -1277,11 +1288,17 @@ class OracleBackend(DatabaseBackend):
|
||||
logger.info(f"Oracle pool created (min={min_size}, max={max_size})")
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
if self._pool is not None:
|
||||
await self._pool.close(force=True)
|
||||
self._pool = None
|
||||
# Drop the reference before awaiting close() so is_ready flips False for
|
||||
# the whole teardown, not just after it completes (see PostgreSQLBackend).
|
||||
pool, self._pool = self._pool, None
|
||||
if pool is not None:
|
||||
await pool.close(force=True)
|
||||
logger.info("Oracle pool closed")
|
||||
|
||||
@property
|
||||
def is_ready(self) -> bool:
|
||||
return self._pool is not None
|
||||
|
||||
async def _set_session_schema(self, conn: Any) -> None:
|
||||
"""Set the session schema on an Oracle connection.
|
||||
|
||||
@@ -1294,15 +1311,41 @@ class OracleBackend(DatabaseBackend):
|
||||
from ..memory_engine import get_current_schema
|
||||
|
||||
schema = get_current_schema()
|
||||
if schema and schema != "public":
|
||||
cursor = conn.cursor()
|
||||
await cursor.execute(f'ALTER SESSION SET CURRENT_SCHEMA = "{schema}"')
|
||||
await cursor.close()
|
||||
cursor = conn.cursor()
|
||||
try:
|
||||
if self._default_schema is None:
|
||||
await cursor.execute("SELECT SYS_CONTEXT('USERENV', 'SESSION_USER') FROM DUAL")
|
||||
row = await cursor.fetchone()
|
||||
if not row or not row[0]:
|
||||
raise RuntimeError("Oracle did not return SESSION_USER while resetting CURRENT_SCHEMA")
|
||||
self._default_schema = str(row[0])
|
||||
|
||||
target_schema = self._default_schema if not schema or schema == "public" else schema
|
||||
safe_schema = target_schema.replace('"', '""')
|
||||
await cursor.execute(f'ALTER SESSION SET CURRENT_SCHEMA = "{safe_schema}"')
|
||||
finally:
|
||||
# oracledb's AsyncCursor.close() is synchronous (not a coroutine);
|
||||
# awaiting it raises "object NoneType can't be used in 'await'
|
||||
# expression" and aborts every acquire().
|
||||
cursor.close()
|
||||
|
||||
def _pool_stats(self) -> PoolStats | None:
|
||||
"""Snapshot for slow-acquire logs, from oracledb pool attributes."""
|
||||
pool = self._pool
|
||||
if pool is None:
|
||||
return None
|
||||
try:
|
||||
busy = pool.busy
|
||||
return PoolStats(in_use=busy, max=pool.max, idle=pool.opened - busy)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@asynccontextmanager
|
||||
async def acquire(self) -> AsyncIterator[OracleConnection]:
|
||||
pool = self._ensure_pool()
|
||||
conn = await pool.acquire()
|
||||
conn = await acquire_conn(
|
||||
pool.acquire(), pool_stats=self._pool_stats, warn_threshold_s=self._acquire_warn_threshold_s
|
||||
)
|
||||
try:
|
||||
await self._set_session_schema(conn)
|
||||
yield OracleConnection(conn)
|
||||
@@ -1318,7 +1361,9 @@ class OracleBackend(DatabaseBackend):
|
||||
@asynccontextmanager
|
||||
async def transaction(self) -> AsyncIterator[OracleConnection]:
|
||||
pool = self._ensure_pool()
|
||||
conn = await pool.acquire()
|
||||
conn = await acquire_conn(
|
||||
pool.acquire(), pool_stats=self._pool_stats, warn_threshold_s=self._acquire_warn_threshold_s
|
||||
)
|
||||
try:
|
||||
await self._set_session_schema(conn)
|
||||
yield OracleConnection(conn)
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
"""Instrumentation for database connection-pool acquisition.
|
||||
|
||||
asyncpg exposes pool *size* and *idle* counts, but not how many callers are
|
||||
currently **queued waiting** for a connection — and that queue depth is the
|
||||
signal that actually distinguishes a saturated pool from a healthy one. When the
|
||||
pool is exhausted, ``/health`` (which itself acquires a connection to run
|
||||
``SELECT 1``) blocks in ``pool.acquire()`` until a connection frees or the acquire
|
||||
times out, so a liveness probe can fail **with the event loop completely idle**.
|
||||
|
||||
This module tracks the process-wide count of in-flight acquisitions that have not
|
||||
yet obtained a connection, and times each acquire so a slow one logs with full
|
||||
pool stats. It is the DB-side counterpart to ``loop_watchdog`` (which covers loop
|
||||
stalls); together, a stuck ``/health`` can be attributed to either a blocked loop
|
||||
or pool exhaustion from the logs alone.
|
||||
|
||||
The counter is a plain int mutated only from the event-loop thread (asyncpg
|
||||
acquisitions are awaited on the loop), so no lock is needed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger("hindsight.db.pool")
|
||||
|
||||
_waiting = 0 # callers currently blocked in pool.acquire(), process-wide
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PoolStats:
|
||||
"""Point-in-time connection-pool utilization snapshot."""
|
||||
|
||||
in_use: int
|
||||
max: int
|
||||
idle: int
|
||||
|
||||
|
||||
def waiting_count() -> int:
|
||||
"""Number of callers currently blocked waiting to acquire a pooled connection."""
|
||||
return _waiting
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def instrument_acquire(
|
||||
acquire_cm: Any,
|
||||
*,
|
||||
pool_stats: Callable[[], PoolStats | None] | None = None,
|
||||
warn_threshold_s: float,
|
||||
) -> AsyncIterator[Any]:
|
||||
"""Wrap a pool's ``acquire()`` context manager with wait tracking + slow-acquire logging.
|
||||
|
||||
Args:
|
||||
acquire_cm: an async context manager yielding a connection (e.g. the object
|
||||
returned by ``asyncpg.Pool.acquire()``).
|
||||
pool_stats: optional zero-arg callable returning a ``PoolStats`` snapshot for
|
||||
the slow-acquire log line.
|
||||
warn_threshold_s: log a warning when the acquire itself takes at least this long.
|
||||
|
||||
Yields:
|
||||
The acquired connection.
|
||||
"""
|
||||
global _waiting
|
||||
_waiting += 1
|
||||
start = time.monotonic()
|
||||
acquired = False
|
||||
try:
|
||||
async with acquire_cm as conn:
|
||||
acquired = True
|
||||
_waiting -= 1
|
||||
_record_acquire_wait(time.monotonic() - start, pool_stats, warn_threshold_s)
|
||||
yield conn
|
||||
finally:
|
||||
# If __aenter__ raised (acquire timeout / cancellation), we never
|
||||
# decremented above — do it here so the waiter count can't leak.
|
||||
if not acquired:
|
||||
_waiting -= 1
|
||||
|
||||
|
||||
async def acquire_conn(
|
||||
acquire_awaitable: Any,
|
||||
*,
|
||||
pool_stats: Callable[[], PoolStats | None] | None = None,
|
||||
warn_threshold_s: float,
|
||||
) -> Any:
|
||||
"""Await a pool acquire that returns a connection, with wait tracking + slow log.
|
||||
|
||||
For pools whose acquire is ``conn = await pool.acquire()`` (oracledb) rather than
|
||||
an async context manager (asyncpg — use ``instrument_acquire`` for those). The
|
||||
caller is responsible for releasing the returned connection.
|
||||
"""
|
||||
global _waiting
|
||||
_waiting += 1
|
||||
start = time.monotonic()
|
||||
try:
|
||||
conn = await acquire_awaitable
|
||||
finally:
|
||||
_waiting -= 1
|
||||
_record_acquire_wait(time.monotonic() - start, pool_stats, warn_threshold_s)
|
||||
return conn
|
||||
|
||||
|
||||
def _record_acquire_wait(
|
||||
wait_s: float,
|
||||
pool_stats: Callable[[], PoolStats | None] | None,
|
||||
warn_threshold_s: float,
|
||||
) -> None:
|
||||
try:
|
||||
from ...metrics import get_metrics_collector
|
||||
|
||||
get_metrics_collector().record_db_acquire_wait(wait_s)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if wait_s < warn_threshold_s:
|
||||
return
|
||||
|
||||
stats: PoolStats | None = None
|
||||
if pool_stats is not None:
|
||||
try:
|
||||
stats = pool_stats()
|
||||
except Exception:
|
||||
stats = None
|
||||
logger.warning(
|
||||
"slow DB pool acquire: waited %.3fs for a connection "
|
||||
"(in_use=%s max=%s idle=%s waiting=%s). The pool is likely saturated; "
|
||||
"/health can stall on connection acquisition while the event loop is free.",
|
||||
wait_s,
|
||||
stats.in_use if stats else None,
|
||||
stats.max if stats else None,
|
||||
stats.idle if stats else None,
|
||||
_waiting,
|
||||
)
|
||||
@@ -15,6 +15,7 @@ from typing import Any
|
||||
import asyncpg # noqa: F401
|
||||
|
||||
from .base import DatabaseBackend, DatabaseConnection
|
||||
from .pool_instrumentation import PoolStats, instrument_acquire
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -76,6 +77,7 @@ class PostgreSQLBackend(DatabaseBackend):
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._pool: asyncpg.Pool | None = None
|
||||
self._acquire_warn_threshold_s: float = 1.0
|
||||
|
||||
async def initialize(
|
||||
self,
|
||||
@@ -88,6 +90,9 @@ class PostgreSQLBackend(DatabaseBackend):
|
||||
statement_cache_size: int = 0,
|
||||
init_callback: Any | None = None,
|
||||
) -> None:
|
||||
from ...config import get_config
|
||||
|
||||
self._acquire_warn_threshold_s = get_config().db_acquire_warn_threshold_ms / 1000.0
|
||||
self._pool = await asyncpg.create_pool(
|
||||
dsn,
|
||||
min_size=min_size,
|
||||
@@ -95,7 +100,12 @@ class PostgreSQLBackend(DatabaseBackend):
|
||||
command_timeout=command_timeout,
|
||||
statement_cache_size=statement_cache_size,
|
||||
timeout=acquire_timeout,
|
||||
# init runs once per new connection; setup runs on every acquire,
|
||||
# after asyncpg's release-time RESET ALL. Passing init_callback as
|
||||
# both keeps the per-connection session GUCs (hnsw.ef_search, etc.)
|
||||
# applied after a connection is reused, not just on first creation.
|
||||
init=init_callback,
|
||||
setup=init_callback,
|
||||
)
|
||||
logger.info(
|
||||
f"PostgreSQL pool created (min={min_size}, max={max_size}, "
|
||||
@@ -103,21 +113,41 @@ class PostgreSQLBackend(DatabaseBackend):
|
||||
)
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
if self._pool is not None:
|
||||
await self._pool.close()
|
||||
self._pool = None
|
||||
# Drop the reference *before* awaiting close(): closing is not
|
||||
# instantaneous, and anything acquiring during that window would
|
||||
# otherwise get an asyncpg "pool is closing" error rather than seeing
|
||||
# is_ready False.
|
||||
pool, self._pool = self._pool, None
|
||||
if pool is not None:
|
||||
await pool.close()
|
||||
logger.info("PostgreSQL pool closed")
|
||||
|
||||
@property
|
||||
def is_ready(self) -> bool:
|
||||
return self._pool is not None
|
||||
|
||||
def _pool_stats(self) -> PoolStats | None:
|
||||
"""Snapshot for slow-acquire logs. in_use = live connections minus idle ones."""
|
||||
pool = self._pool
|
||||
if pool is None:
|
||||
return None
|
||||
idle = pool.get_idle_size()
|
||||
return PoolStats(in_use=pool.get_size() - idle, max=pool.get_max_size(), idle=idle)
|
||||
|
||||
@asynccontextmanager
|
||||
async def acquire(self) -> AsyncIterator[PostgresConnection]:
|
||||
pool = self._ensure_pool()
|
||||
async with pool.acquire() as conn:
|
||||
async with instrument_acquire(
|
||||
pool.acquire(), pool_stats=self._pool_stats, warn_threshold_s=self._acquire_warn_threshold_s
|
||||
) as conn:
|
||||
yield PostgresConnection(conn)
|
||||
|
||||
@asynccontextmanager
|
||||
async def transaction(self) -> AsyncIterator[PostgresConnection]:
|
||||
pool = self._ensure_pool()
|
||||
async with pool.acquire() as conn:
|
||||
async with instrument_acquire(
|
||||
pool.acquire(), pool_stats=self._pool_stats, warn_threshold_s=self._acquire_warn_threshold_s
|
||||
) as conn:
|
||||
async with conn.transaction():
|
||||
yield PostgresConnection(conn)
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ Database utility functions for connection management with retry logic.
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import random
|
||||
import time
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import AsyncExitStack, asynccontextmanager
|
||||
@@ -16,6 +17,20 @@ DEFAULT_MAX_RETRIES = 3
|
||||
DEFAULT_BASE_DELAY = 0.5 # seconds
|
||||
DEFAULT_MAX_DELAY = 5.0 # seconds
|
||||
|
||||
|
||||
def _backoff_delay(attempt: int, base_delay: float, max_delay: float) -> float:
|
||||
"""Exponential backoff with equal jitter.
|
||||
|
||||
Deterministic backoff makes concurrent retriers wake in lock-step and
|
||||
re-collide on the very same rows, re-triggering the deadlock they just
|
||||
backed off from. "Equal jitter" — half the window fixed, half random —
|
||||
keeps a floor (so we don't hot-spin) while decorrelating the wake-ups, so
|
||||
two contenders that deadlocked together are very unlikely to retry in sync.
|
||||
"""
|
||||
ceil = min(base_delay * (2**attempt), max_delay)
|
||||
return ceil / 2 + random.uniform(0, ceil / 2)
|
||||
|
||||
|
||||
# Retryable exception types (checked by class name to avoid hard imports)
|
||||
_RETRYABLE_EXCEPTION_NAMES = frozenset(
|
||||
{
|
||||
@@ -78,7 +93,7 @@ async def retry_with_backoff(
|
||||
raise
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
delay = min(base_delay * (2**attempt), max_delay)
|
||||
delay = _backoff_delay(attempt, base_delay, max_delay)
|
||||
if type(e).__name__ == "DeadlockDetectedError" or _is_oracle_deadlock(e):
|
||||
logger.warning(
|
||||
"Deadlock detected during parallel document processing — "
|
||||
@@ -136,7 +151,7 @@ async def acquire_with_retry(backend_or_pool: Any, max_retries: int = DEFAULT_MA
|
||||
if not _is_retryable(e):
|
||||
raise
|
||||
if attempt < max_retries:
|
||||
delay = min(DEFAULT_BASE_DELAY * (2**attempt), DEFAULT_MAX_DELAY)
|
||||
delay = _backoff_delay(attempt, DEFAULT_BASE_DELAY, DEFAULT_MAX_DELAY)
|
||||
logger.warning(
|
||||
f"Database acquire failed (attempt {attempt + 1}/{max_retries + 1}): {e}. "
|
||||
f"Retrying in {delay:.1f}s..."
|
||||
|
||||
@@ -76,6 +76,25 @@ class _ZeroEntropyEmbedResponse(BaseModel):
|
||||
results: list[_ZeroEntropyEmbedResult]
|
||||
|
||||
|
||||
def _truncate_to_tokens(text: str, max_tokens: int) -> tuple[str, int]:
|
||||
"""Truncate ``text`` to at most ``max_tokens`` cl100k_base tokens.
|
||||
|
||||
tiktoken is an approximation of any given provider's tokenizer, so set
|
||||
``max_tokens`` with a little headroom below the model's real limit.
|
||||
|
||||
Returns the (possibly truncated) text and the original token count (so the
|
||||
caller can report how much was dropped); the count equals ``len(tokens)``
|
||||
whether or not truncation occurred.
|
||||
"""
|
||||
from .token_encoding import get_token_encoding
|
||||
|
||||
enc = get_token_encoding()
|
||||
tokens = enc.encode(text)
|
||||
if len(tokens) <= max_tokens:
|
||||
return text, len(tokens)
|
||||
return enc.decode(tokens[:max_tokens]), len(tokens)
|
||||
|
||||
|
||||
class Embeddings(ABC):
|
||||
"""
|
||||
Abstract base class for embedding generation.
|
||||
@@ -190,7 +209,7 @@ class LocalSTEmbeddings(Embeddings):
|
||||
device = "cpu"
|
||||
logger.info("Embeddings: forcing CPU mode")
|
||||
else:
|
||||
# Check for GPU (CUDA) or Apple Silicon (MPS)
|
||||
# Check for GPU (CUDA), Apple Silicon (MPS), or Intel XPU
|
||||
# Wrap in try-except to gracefully handle any device detection issues
|
||||
# (e.g., in CI environments or when PyTorch is built without GPU support)
|
||||
device = "cpu" # Default to CPU
|
||||
@@ -198,10 +217,13 @@ class LocalSTEmbeddings(Embeddings):
|
||||
has_gpu = torch.cuda.is_available() or (
|
||||
hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
|
||||
)
|
||||
# Intel Arc XPU support — torch.xpu is available when the XPU build is loaded
|
||||
if not has_gpu and hasattr(torch, "xpu"):
|
||||
has_gpu = torch.xpu.is_available()
|
||||
if has_gpu:
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS/XPU
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to detect GPU/MPS, falling back to CPU: {e}")
|
||||
logger.warning(f"Failed to detect GPU/MPS/XPU, falling back to CPU: {e}")
|
||||
|
||||
# Suppress verbose transformers warnings during model loading
|
||||
# This suppresses the "UNEXPECTED" warnings from BertModel which are harmless
|
||||
@@ -709,7 +731,8 @@ class OpenAIEmbeddings(Embeddings):
|
||||
|
||||
class CodexOAuthEmbeddings(OpenAIEmbeddings):
|
||||
"""
|
||||
OpenAI embeddings using the Codex/ChatGPT OAuth token from ``~/.codex/auth.json``.
|
||||
OpenAI embeddings using the Codex/ChatGPT OAuth token from the Codex
|
||||
``auth.json`` (``$CODEX_HOME/auth.json``, or ``~/.codex/auth.json`` when unset).
|
||||
|
||||
Codex OAuth is an LLM-provider auth path in Hindsight, but the same bearer token
|
||||
can also authenticate against the standard OpenAI embeddings endpoint. This keeps
|
||||
@@ -1198,6 +1221,7 @@ class LiteLLMSDKEmbeddings(Embeddings):
|
||||
batch_size: int = 100,
|
||||
timeout: float = 60.0,
|
||||
encoding_format: str | None = "float",
|
||||
max_input_tokens: int | None = None,
|
||||
):
|
||||
"""
|
||||
Initialize LiteLLM SDK embeddings client.
|
||||
@@ -1212,6 +1236,10 @@ class LiteLLMSDKEmbeddings(Embeddings):
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
encoding_format: Encoding format for embeddings (default: "float").
|
||||
Set to None or empty string to omit (needed for Voyage AI, Gemini).
|
||||
max_input_tokens: If set, truncate each input text to this many tokens
|
||||
(tiktoken cl100k_base) before embedding. Needed for models with a
|
||||
fixed input-token limit (e.g. Bedrock Titan V2's hard 8192 cap),
|
||||
where an oversized text would otherwise fail permanently (#2501).
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
@@ -1220,6 +1248,7 @@ class LiteLLMSDKEmbeddings(Embeddings):
|
||||
self.batch_size = batch_size
|
||||
self.timeout = timeout
|
||||
self.encoding_format = encoding_format or None
|
||||
self.max_input_tokens = max_input_tokens
|
||||
self._litellm = None # Will be set during initialization
|
||||
self._dimension: int | None = None
|
||||
|
||||
@@ -1296,6 +1325,33 @@ class LiteLLMSDKEmbeddings(Embeddings):
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
# Truncate oversized inputs before hitting the provider. Models with a
|
||||
# fixed input-token limit (e.g. Bedrock Titan V2, 8192) reject an
|
||||
# oversized text with a permanent error rather than truncating it
|
||||
# server-side, which strands the caller (e.g. a delta mental model whose
|
||||
# content grew past the cap) with no recovery path. See #2501.
|
||||
if self.max_input_tokens is not None:
|
||||
truncated_texts = []
|
||||
original_token_counts = []
|
||||
for t in texts:
|
||||
new_text, original_tokens = _truncate_to_tokens(t, self.max_input_tokens)
|
||||
truncated_texts.append(new_text)
|
||||
if original_tokens > self.max_input_tokens:
|
||||
original_token_counts.append(original_tokens)
|
||||
texts = truncated_texts
|
||||
if original_token_counts:
|
||||
logger.warning(
|
||||
"Embeddings: truncated %d of %d input(s) to %d tokens for model %s "
|
||||
"(largest was ~%d tokens); embedded content is incomplete. "
|
||||
"This usually means a mental model's content has grown past the model's "
|
||||
"input limit — see issue #2501.",
|
||||
len(original_token_counts),
|
||||
len(texts),
|
||||
self.max_input_tokens,
|
||||
self.model,
|
||||
max(original_token_counts),
|
||||
)
|
||||
|
||||
all_embeddings = []
|
||||
|
||||
# Process in batches
|
||||
@@ -1634,6 +1690,20 @@ def create_embeddings_from_env() -> Embeddings:
|
||||
batch_size=config.embeddings_openai_batch_size,
|
||||
dimensions=config.embeddings_openai_dimensions,
|
||||
)
|
||||
elif provider == "requesty":
|
||||
api_key = config.embeddings_requesty_api_key
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"HINDSIGHT_API_EMBEDDINGS_REQUESTY_API_KEY, HINDSIGHT_API_REQUESTY_API_KEY, "
|
||||
f"or {ENV_LLM_API_KEY} is required when {ENV_EMBEDDINGS_PROVIDER} is 'requesty'"
|
||||
)
|
||||
return OpenAIEmbeddings(
|
||||
api_key=api_key,
|
||||
model=config.embeddings_requesty_model,
|
||||
base_url="https://router.requesty.ai/v1",
|
||||
batch_size=config.embeddings_openai_batch_size,
|
||||
dimensions=config.embeddings_openai_dimensions,
|
||||
)
|
||||
elif provider == "zeroentropy":
|
||||
api_key = config.embeddings_zeroentropy_api_key
|
||||
if not api_key:
|
||||
@@ -1673,6 +1743,7 @@ def create_embeddings_from_env() -> Embeddings:
|
||||
api_base=config.embeddings_litellm_sdk_api_base,
|
||||
output_dimensions=config.embeddings_litellm_sdk_output_dimensions,
|
||||
encoding_format=config.embeddings_litellm_sdk_encoding_format,
|
||||
max_input_tokens=config.embeddings_litellm_sdk_max_input_tokens,
|
||||
)
|
||||
elif provider == "google":
|
||||
vertexai_project_id = config.embeddings_vertexai_project_id
|
||||
@@ -1697,6 +1768,6 @@ def create_embeddings_from_env() -> Embeddings:
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown embeddings provider: {provider}. "
|
||||
f"Supported: 'local', 'onnx', 'tei', 'openai', 'openai-codex', 'openrouter', 'cohere', 'google', "
|
||||
f"Supported: 'local', 'onnx', 'tei', 'openai', 'openai-codex', 'openrouter', 'requesty', 'cohere', 'google', "
|
||||
f"'zeroentropy', 'litellm', 'litellm-sdk'"
|
||||
)
|
||||
|
||||
@@ -9,10 +9,11 @@ import asyncio
|
||||
import json
|
||||
import logging
|
||||
from collections import defaultdict
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from difflib import SequenceMatcher
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from .db_utils import acquire_with_retry
|
||||
from .memory_engine import fq_table
|
||||
@@ -25,6 +26,7 @@ from .retain.entity_labels import (
|
||||
from .retain.entity_labels import (
|
||||
parse_entity_labels as _parse_entity_labels,
|
||||
)
|
||||
from .retain.types import ResolvedEntity
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -75,6 +77,22 @@ def _later_date(a: datetime | None, b: datetime | None) -> datetime | None:
|
||||
return a if a > b else b
|
||||
|
||||
|
||||
def _canonical_cooccurrence_pairs(entity_list: list[str]) -> Iterator[tuple[str, str]]:
|
||||
"""Yield each distinct pair of ``entity_list`` as ``(a, b)`` with ``a < b``.
|
||||
|
||||
Canonical ordering matches the entity_cooccurrences PK and check constraint.
|
||||
The pair is ordered into fresh locals rather than by swapping the loop
|
||||
variables: ``entity_id_1`` is the outer iterate, so swapping it would leak
|
||||
into the remaining inner iterations and build later pairs off the wrong
|
||||
element.
|
||||
"""
|
||||
for i, entity_id_1 in enumerate(entity_list):
|
||||
for entity_id_2 in entity_list[i + 1 :]:
|
||||
if entity_id_1 == entity_id_2:
|
||||
continue
|
||||
yield (entity_id_1, entity_id_2) if entity_id_1 < entity_id_2 else (entity_id_2, entity_id_1)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _CooccurrencePair:
|
||||
"""A (entity_id_1, entity_id_2) pair observed in a retain batch (for post-txn flush)."""
|
||||
@@ -230,7 +248,7 @@ class EntityResolver:
|
||||
unit_event_date,
|
||||
conn=None,
|
||||
entity_labels: list | None = None,
|
||||
) -> list[str]:
|
||||
) -> list[ResolvedEntity]:
|
||||
"""
|
||||
Resolve multiple entities in batch (MUCH faster than sequential).
|
||||
|
||||
@@ -245,7 +263,8 @@ class EntityResolver:
|
||||
conn: Optional connection to use (if None, acquires from pool)
|
||||
|
||||
Returns:
|
||||
List of entity IDs in same order as input
|
||||
Resolved entity identities (id + stored canonical name) in the same
|
||||
order as input.
|
||||
"""
|
||||
if not entities_data:
|
||||
return []
|
||||
@@ -271,7 +290,7 @@ class EntityResolver:
|
||||
unit_event_date,
|
||||
taxonomy_lookup: set[str] | None = None,
|
||||
labels_cfg=None,
|
||||
) -> list[str]:
|
||||
) -> list[ResolvedEntity]:
|
||||
if self.entity_lookup == "trigram":
|
||||
# Route to backend-specific fuzzy strategy.
|
||||
# Non-PG backends (Oracle) use UTL_MATCH instead of pg_trgm.
|
||||
@@ -311,7 +330,7 @@ class EntityResolver:
|
||||
unit_event_date,
|
||||
taxonomy_lookup: set[str] | None = None,
|
||||
labels_cfg=None,
|
||||
) -> list[str]:
|
||||
) -> list[ResolvedEntity]:
|
||||
"""Original strategy: load all bank entities then match in Python."""
|
||||
# Query ALL candidates for this bank
|
||||
all_entities = await conn.fetch(
|
||||
@@ -395,7 +414,7 @@ class EntityResolver:
|
||||
unit_event_date,
|
||||
taxonomy_lookup: set[str] | None = None,
|
||||
labels_cfg=None,
|
||||
) -> list[str]:
|
||||
) -> list[ResolvedEntity]:
|
||||
"""
|
||||
Trigram strategy: fetch only similar candidates per entity name using pg_trgm.
|
||||
|
||||
@@ -499,7 +518,7 @@ class EntityResolver:
|
||||
unit_event_date: datetime | None,
|
||||
taxonomy_lookup: set[str] | None = None,
|
||||
labels_cfg=None,
|
||||
) -> list[str]:
|
||||
) -> list[ResolvedEntity]:
|
||||
"""
|
||||
Oracle strategy: fetch similar candidates using UTL_MATCH.JARO_WINKLER_SIMILARITY.
|
||||
|
||||
@@ -607,11 +626,14 @@ class EntityResolver:
|
||||
cooccurrence_map: dict[str, set[str]],
|
||||
taxonomy_lookup: set[str] | None = None,
|
||||
labels_cfg=None,
|
||||
) -> list[str]:
|
||||
) -> list[ResolvedEntity]:
|
||||
"""Shared scoring + upsert logic used by both lookup strategies."""
|
||||
|
||||
# Resolve each entity using pre-fetched candidates
|
||||
entity_ids = [None] * len(entities_data)
|
||||
# Resolve each entity using pre-fetched candidates. A slot stays None
|
||||
# only if find-or-create fails to produce a row for a mention (a DB
|
||||
# inconsistency); it surfaces as a clear error at the reassert boundary
|
||||
# rather than a silent NOT NULL violation deeper in Phase 2.
|
||||
resolved: list[ResolvedEntity | None] = [None] * len(entities_data)
|
||||
entities_to_update: list[_EntityStat] = []
|
||||
entities_to_create: list[_EntityToCreate] = []
|
||||
|
||||
@@ -638,21 +660,23 @@ class EntityResolver:
|
||||
|
||||
if is_label:
|
||||
# Exact case-insensitive match only for label entities
|
||||
exact_match = None
|
||||
exact_match: ResolvedEntity | None = None
|
||||
entity_text_lower = entity_text.lower()
|
||||
for candidate_id, canonical_name, metadata, last_seen, mention_count in candidates:
|
||||
if canonical_name.lower() == entity_text_lower:
|
||||
exact_match = candidate_id
|
||||
exact_match = ResolvedEntity(entity_id=candidate_id, canonical_name=canonical_name)
|
||||
break
|
||||
if exact_match:
|
||||
entity_ids[idx] = exact_match
|
||||
entities_to_update.append(_EntityStat(entity_id=exact_match, event_date=entity_event_date))
|
||||
resolved[idx] = exact_match
|
||||
entities_to_update.append(
|
||||
_EntityStat(entity_id=exact_match.entity_id, event_date=entity_event_date)
|
||||
)
|
||||
else:
|
||||
entities_to_create.append(_EntityToCreate(idx=idx, name=entity_text, event_date=entity_event_date))
|
||||
continue
|
||||
|
||||
# Score candidates
|
||||
best_candidate = None
|
||||
best_candidate: ResolvedEntity | None = None
|
||||
best_score = 0.0
|
||||
|
||||
nearby_entity_set = {e["text"].lower() for e in nearby_entities if e["text"] != entity_text}
|
||||
@@ -685,14 +709,14 @@ class EntityResolver:
|
||||
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_candidate = candidate_id
|
||||
best_candidate = ResolvedEntity(entity_id=candidate_id, canonical_name=canonical_name)
|
||||
|
||||
# Apply unified threshold
|
||||
threshold = 0.6
|
||||
|
||||
if best_score > threshold:
|
||||
entity_ids[idx] = best_candidate
|
||||
entities_to_update.append(_EntityStat(entity_id=best_candidate, event_date=entity_event_date))
|
||||
if best_score > threshold and best_candidate is not None:
|
||||
resolved[idx] = best_candidate
|
||||
entities_to_update.append(_EntityStat(entity_id=best_candidate.entity_id, event_date=entity_event_date))
|
||||
else:
|
||||
entities_to_create.append(
|
||||
_EntityToCreate(idx=idx, name=entity_data["text"], event_date=entity_event_date)
|
||||
@@ -725,6 +749,9 @@ class EntityResolver:
|
||||
sorted_groups = sorted(groups.items())
|
||||
entity_names = [g.name for _, g in sorted_groups]
|
||||
entity_dates = [g.event_date for _, g in sorted_groups]
|
||||
# Stored canonical name per lowercase key, so a resurrected parent
|
||||
# keeps the name it was created/matched with rather than a fallback.
|
||||
canonical_by_name = {name_lower: g.name for name_lower, g in sorted_groups}
|
||||
|
||||
# INSERT ... ON CONFLICT DO NOTHING — no row lock on already-existing entities.
|
||||
# mention_count starts at 0 here; flush_pending_stats() is the sole source of
|
||||
@@ -759,11 +786,14 @@ class EntityResolver:
|
||||
)
|
||||
for row in existing_rows:
|
||||
id_by_name[row["name_lower"]] = row["id"]
|
||||
canonical_by_name[row["name_lower"]] = row["canonical_name"]
|
||||
# Also index by Python's lower() of the original input name so the
|
||||
# assignment loop (which uses Python-lowercased keys) finds it even
|
||||
# when Python and the database produce different lowercase strings.
|
||||
if "input_name" in row:
|
||||
id_by_name[row["input_name"].lower()] = row["id"]
|
||||
input_name_lower = row["input_name"].lower()
|
||||
id_by_name[input_name_lower] = row["id"]
|
||||
canonical_by_name[input_name_lower] = row["canonical_name"]
|
||||
|
||||
# Assign entity IDs back and queue one stat per original mention so that
|
||||
# flush_pending_stats() increments mention_count by the true mention count,
|
||||
@@ -771,245 +801,63 @@ class EntityResolver:
|
||||
for name_lower, g in sorted_groups:
|
||||
entity_id = id_by_name.get(name_lower)
|
||||
if entity_id:
|
||||
canonical_name = canonical_by_name.get(name_lower, g.name)
|
||||
for original_idx in g.indices:
|
||||
entity_ids[original_idx] = entity_id
|
||||
pending.append(_EntityStat(entity_id=entity_id, event_date=g.event_date))
|
||||
resolved[original_idx] = ResolvedEntity(entity_id=entity_id, canonical_name=canonical_name)
|
||||
pending.append(_EntityStat(entity_id=str(entity_id), event_date=g.event_date))
|
||||
|
||||
# Accumulate into the resolver's pending list; the orchestrator flushes
|
||||
# these with await entity_resolver.flush_pending_stats() after the txn.
|
||||
key = self._task_key()
|
||||
self._pending_stats.setdefault(key, []).extend(pending)
|
||||
|
||||
return entity_ids
|
||||
missing = [i for i, entity in enumerate(resolved) if entity is None]
|
||||
if missing:
|
||||
raise RuntimeError(
|
||||
f"Entity resolution produced no row for {len(missing)} mention(s) "
|
||||
f"(indices {missing[:5]}); refusing to link units to a missing parent."
|
||||
)
|
||||
return cast(list[ResolvedEntity], resolved)
|
||||
|
||||
async def resolve_entity(
|
||||
async def reassert_entities_batch(
|
||||
self,
|
||||
bank_id: str,
|
||||
entity_text: str,
|
||||
context: str,
|
||||
nearby_entities: list[dict],
|
||||
unit_event_date,
|
||||
) -> str:
|
||||
"""
|
||||
Resolve an entity to a canonical entity ID.
|
||||
|
||||
Args:
|
||||
bank_id: bank ID (entities are scoped to agents)
|
||||
entity_text: Entity text ("Alice", "Google", etc.)
|
||||
context: Context where entity appears
|
||||
nearby_entities: Other entities in the same unit
|
||||
unit_event_date: When this unit was created
|
||||
|
||||
Returns:
|
||||
Entity ID (creates new entity if needed)
|
||||
"""
|
||||
async with acquire_with_retry(self.pool) as conn:
|
||||
# Find candidate entities with similar name
|
||||
candidates = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, canonical_name, metadata, last_seen
|
||||
FROM {fq_table("entities")}
|
||||
WHERE bank_id = $1
|
||||
AND (
|
||||
canonical_name ILIKE $2
|
||||
OR canonical_name ILIKE $3
|
||||
OR $2 ILIKE canonical_name || '%%'
|
||||
)
|
||||
ORDER BY mention_count DESC
|
||||
""",
|
||||
bank_id,
|
||||
entity_text,
|
||||
f"%{entity_text}%",
|
||||
)
|
||||
|
||||
if not candidates:
|
||||
# New entity - create it
|
||||
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
|
||||
|
||||
# Score candidates based on:
|
||||
# 1. Name similarity
|
||||
# 2. Context overlap (TODO: could use embeddings)
|
||||
# 3. Co-occurring entities
|
||||
# 4. Temporal proximity
|
||||
|
||||
best_candidate = None
|
||||
best_score = 0.0
|
||||
|
||||
nearby_entity_set = {e["text"].lower() for e in nearby_entities if e["text"] != entity_text}
|
||||
|
||||
for row in candidates:
|
||||
candidate_id = row["id"]
|
||||
canonical_name = row["canonical_name"]
|
||||
last_seen = row["last_seen"]
|
||||
score = 0.0
|
||||
|
||||
# 1. Name similarity (0-1)
|
||||
name_similarity = SequenceMatcher(None, entity_text.lower(), canonical_name.lower()).ratio()
|
||||
score += name_similarity * 0.5
|
||||
|
||||
# 2. Co-occurring entities (0-0.5)
|
||||
# Get entities that co-occurred with this candidate before
|
||||
# Use the materialized co-occurrence cache for fast lookup
|
||||
co_entity_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT e.canonical_name, ec.cooccurrence_count
|
||||
FROM {fq_table("entity_cooccurrences")} ec
|
||||
JOIN {fq_table("entities")} e ON (
|
||||
CASE
|
||||
WHEN ec.entity_id_1 = $1 THEN ec.entity_id_2
|
||||
WHEN ec.entity_id_2 = $1 THEN ec.entity_id_1
|
||||
END = e.id
|
||||
)
|
||||
WHERE ec.entity_id_1 = $1 OR ec.entity_id_2 = $1
|
||||
""",
|
||||
candidate_id,
|
||||
)
|
||||
co_entities = {r["canonical_name"].lower() for r in co_entity_rows}
|
||||
|
||||
# Check overlap with nearby entities
|
||||
overlap = len(nearby_entity_set & co_entities)
|
||||
if nearby_entity_set:
|
||||
co_entity_score = overlap / len(nearby_entity_set)
|
||||
score += co_entity_score * 0.3
|
||||
|
||||
# 3. Temporal proximity (0-0.2)
|
||||
if last_seen:
|
||||
# Normalize both to UTC-aware to avoid naive/aware mismatch
|
||||
# (Oracle returns naive datetimes from fromisoformat)
|
||||
_evt = unit_event_date if unit_event_date.tzinfo else unit_event_date.replace(tzinfo=UTC)
|
||||
_seen = last_seen if last_seen.tzinfo else last_seen.replace(tzinfo=UTC)
|
||||
days_diff = abs((_evt - _seen).total_seconds() / 86400)
|
||||
if days_diff < 7: # Within a week
|
||||
temporal_score = max(0, 1.0 - (days_diff / 7))
|
||||
score += temporal_score * 0.2
|
||||
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_candidate = candidate_id
|
||||
|
||||
# Threshold for considering it the same entity
|
||||
threshold = 0.6
|
||||
|
||||
if best_score > threshold:
|
||||
# Update entity
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("entities")}
|
||||
SET mention_count = mention_count + 1,
|
||||
last_seen = $1
|
||||
WHERE id = $2
|
||||
""",
|
||||
unit_event_date,
|
||||
best_candidate,
|
||||
)
|
||||
return best_candidate
|
||||
else:
|
||||
# Not confident - create new entity
|
||||
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
|
||||
|
||||
async def _create_entity(
|
||||
self,
|
||||
resolved_entities: list[ResolvedEntity],
|
||||
conn,
|
||||
bank_id: str,
|
||||
entity_text: str,
|
||||
event_date,
|
||||
) -> str:
|
||||
) -> None:
|
||||
"""Lock (and, if pruned, re-create) resolved parents before linking units.
|
||||
|
||||
Phase-1 resolution and the Phase-2 ``unit_entities`` insert run on
|
||||
different transactions. In the gap, ``prune_orphan_entities`` can delete
|
||||
a just-resolved parent — it legitimately has no ``unit_entities`` row
|
||||
yet — and the Phase-2 FK insert then fails, dropping the whole batch as
|
||||
non-retryable (silent memory loss, #2662).
|
||||
|
||||
Called on the Phase-2 connection immediately before
|
||||
``link_units_to_entities_batch``, this locks the parents that still
|
||||
exist (so the pruner blocks until we commit) and re-inserts any that
|
||||
already vanished, in one round-trip. An entity referenced by a live unit
|
||||
is by definition not an orphan, so resurrecting it is correct.
|
||||
"""
|
||||
Create a new entity or get existing one if it already exists.
|
||||
# Deduplicate by id and lock in a stable order so concurrent reasserts
|
||||
# acquire row locks consistently (same convention as bulk_insert_links).
|
||||
seen: set[str] = set()
|
||||
unique: list[ResolvedEntity] = []
|
||||
for entity in sorted(resolved_entities, key=lambda e: e.entity_id):
|
||||
if entity.entity_id in seen:
|
||||
continue
|
||||
seen.add(entity.entity_id)
|
||||
unique.append(entity)
|
||||
|
||||
Uses INSERT ... ON CONFLICT to handle race conditions where
|
||||
two concurrent transactions try to create the same entity.
|
||||
if not unique:
|
||||
return
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
bank_id: bank ID
|
||||
entity_text: Entity text
|
||||
event_date: When first seen
|
||||
|
||||
Returns:
|
||||
Entity ID
|
||||
"""
|
||||
entity_id = await conn.fetchval(
|
||||
f"""
|
||||
INSERT INTO {fq_table("entities")} (bank_id, canonical_name, first_seen, last_seen, mention_count)
|
||||
VALUES ($1, $2, COALESCE($3, now()), COALESCE($4, now()), 1)
|
||||
ON CONFLICT (bank_id, LOWER(canonical_name))
|
||||
DO UPDATE SET
|
||||
mention_count = {fq_table("entities")}.mention_count + 1,
|
||||
last_seen = EXCLUDED.last_seen
|
||||
RETURNING id
|
||||
""",
|
||||
await self._ops.bulk_reassert_entities(
|
||||
conn,
|
||||
fq_table("entities"),
|
||||
bank_id,
|
||||
entity_text,
|
||||
event_date,
|
||||
event_date,
|
||||
)
|
||||
return entity_id
|
||||
|
||||
async def link_unit_to_entity(self, unit_id: str, entity_id: str):
|
||||
"""
|
||||
Link a memory unit to an entity.
|
||||
Also updates co-occurrence cache with other entities in the same unit.
|
||||
|
||||
Args:
|
||||
unit_id: Memory unit ID
|
||||
entity_id: Entity ID
|
||||
"""
|
||||
async with acquire_with_retry(self.pool) as conn:
|
||||
# Insert unit-entity link
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("unit_entities")} (unit_id, entity_id)
|
||||
VALUES ($1, $2)
|
||||
ON CONFLICT DO NOTHING
|
||||
""",
|
||||
unit_id,
|
||||
entity_id,
|
||||
)
|
||||
|
||||
# Update co-occurrence cache: find other entities in this unit
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT entity_id
|
||||
FROM {fq_table("unit_entities")}
|
||||
WHERE unit_id = $1 AND entity_id != $2
|
||||
""",
|
||||
unit_id,
|
||||
entity_id,
|
||||
)
|
||||
|
||||
other_entities = [row["entity_id"] for row in rows]
|
||||
|
||||
# Update co-occurrences for each pair
|
||||
for other_entity_id in other_entities:
|
||||
await self._update_cooccurrence(conn, entity_id, other_entity_id)
|
||||
|
||||
async def _update_cooccurrence(self, conn, entity_id_1: str, entity_id_2: str):
|
||||
"""
|
||||
Update the co-occurrence cache for two entities.
|
||||
|
||||
Uses CHECK constraint ordering (entity_id_1 < entity_id_2) to avoid duplicates.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
entity_id_1: First entity ID
|
||||
entity_id_2: Second entity ID
|
||||
"""
|
||||
# Ensure consistent ordering (smaller UUID first)
|
||||
if entity_id_1 > entity_id_2:
|
||||
entity_id_1, entity_id_2 = entity_id_2, entity_id_1
|
||||
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("entity_cooccurrences")} (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
|
||||
VALUES ($1, $2, 1, NOW())
|
||||
ON CONFLICT (entity_id_1, entity_id_2)
|
||||
DO UPDATE SET
|
||||
cooccurrence_count = {fq_table("entity_cooccurrences")}.cooccurrence_count + 1,
|
||||
last_cooccurred = NOW()
|
||||
""",
|
||||
entity_id_1,
|
||||
entity_id_2,
|
||||
[entity.entity_id for entity in unique],
|
||||
[entity.canonical_name for entity in unique],
|
||||
)
|
||||
|
||||
async def link_units_to_entities_batch(
|
||||
@@ -1083,20 +931,12 @@ class EntityResolver:
|
||||
for unit_id, entity_ids in unit_to_entities.items():
|
||||
entity_list = list(entity_ids)
|
||||
event_date = unit_event_date.get(unit_id)
|
||||
for i, entity_id_1 in enumerate(entity_list):
|
||||
for entity_id_2 in entity_list[i + 1 :]:
|
||||
if entity_id_1 == entity_id_2:
|
||||
continue
|
||||
# Canonical ordering (entity_id_1 < entity_id_2) matches the
|
||||
# entity_cooccurrences PK and check constraint.
|
||||
if entity_id_1 > entity_id_2:
|
||||
entity_id_1, entity_id_2 = entity_id_2, entity_id_1
|
||||
key = (entity_id_1, entity_id_2)
|
||||
prev = cooccurrence_pairs.get(key, _SENTINEL_MISSING)
|
||||
if prev is _SENTINEL_MISSING:
|
||||
cooccurrence_pairs[key] = event_date
|
||||
else:
|
||||
cooccurrence_pairs[key] = _later_date(prev, event_date)
|
||||
for key in _canonical_cooccurrence_pairs(entity_list):
|
||||
prev = cooccurrence_pairs.get(key, _SENTINEL_MISSING)
|
||||
if prev is _SENTINEL_MISSING:
|
||||
cooccurrence_pairs[key] = event_date
|
||||
else:
|
||||
cooccurrence_pairs[key] = _later_date(prev, event_date)
|
||||
|
||||
# Accumulate co-occurrence pairs for post-transaction flush.
|
||||
# The actual INSERT/UPDATE is deferred to flush_pending_stats() to avoid
|
||||
|
||||
@@ -66,6 +66,20 @@ MAX_SEMANTIC_LINKS_PER_UNIT = 50
|
||||
# under 1s.
|
||||
_DRAIN_BATCH_SIZE = 50
|
||||
|
||||
# Retry budget for the idempotent Pass 2/3 entity/cooccurrence sweep. Higher
|
||||
# than db_utils' default (3) because the sweep has no client waiting on it and
|
||||
# is safe to rerun, so we'd rather spend a longer jittered-backoff tail than
|
||||
# drop a maintenance pass and leak stale graph rows (see run_graph_maintenance_job).
|
||||
_SWEEP_MAX_RETRIES = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class _SweepCounts:
|
||||
"""Prune counts returned by the Pass 2/3 sweep (avoids a bare tuple return)."""
|
||||
|
||||
orphan_entities_pruned: int
|
||||
stale_cooccurrences_pruned: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class JobResult:
|
||||
@@ -203,27 +217,51 @@ async def run_graph_maintenance_job(
|
||||
|
||||
# --- Pass 2 & 3: entity / cooccurrence sweeps ---
|
||||
# Bank-wide single-statement deletes. Cheap when there's nothing to do.
|
||||
#
|
||||
# Unlike Pass 1's queue claim, these DELETEs aren't protected by any
|
||||
# consistent lock-ordering guarantee: prune_stale_cooccurrences scans
|
||||
# entity_cooccurrences via a join/NOT EXISTS plan, while retain's
|
||||
# concurrent cooccurrence upserts (entity_resolver._flush_pending) lock
|
||||
# the same rows in sorted (entity_id_1, entity_id_2) order. When a sweep
|
||||
# and a concurrent upsert touch overlapping rows in opposite orders,
|
||||
# Postgres detects a genuine circular wait and aborts one side with
|
||||
# DeadlockDetectedError. Both prunes are idempotent bank-wide sweeps —
|
||||
# rerunning only deletes what's still stale — so retrying the whole
|
||||
# transaction on deadlock is safe.
|
||||
from .db_utils import retry_with_backoff
|
||||
from .memory_engine import acquire_with_retry
|
||||
|
||||
async with acquire_with_retry(backend) as conn:
|
||||
async with conn.transaction():
|
||||
result.orphan_entities_pruned = await ops.prune_orphan_entities(
|
||||
conn,
|
||||
fq_table("entities"),
|
||||
fq_table("unit_entities"),
|
||||
bank_id,
|
||||
)
|
||||
# The orphan prune above cascades cooccurrences via FK. The
|
||||
# explicit cooccurrence pass below catches the *stale-count*
|
||||
# case: both entities still exist but no current unit witnesses
|
||||
# them together.
|
||||
result.stale_cooccurrences_pruned = await ops.prune_stale_cooccurrences(
|
||||
conn,
|
||||
fq_table("entity_cooccurrences"),
|
||||
fq_table("unit_entities"),
|
||||
fq_table("entities"),
|
||||
bank_id,
|
||||
)
|
||||
async def _run_sweep() -> _SweepCounts:
|
||||
async with acquire_with_retry(backend) as conn:
|
||||
async with conn.transaction():
|
||||
orphan_pruned = await ops.prune_orphan_entities(
|
||||
conn,
|
||||
fq_table("entities"),
|
||||
fq_table("unit_entities"),
|
||||
bank_id,
|
||||
)
|
||||
# The orphan prune above cascades cooccurrences via FK. The
|
||||
# explicit cooccurrence pass below catches the *stale-count*
|
||||
# case: both entities still exist but no current unit
|
||||
# witnesses them together.
|
||||
stale_pruned = await ops.prune_stale_cooccurrences(
|
||||
conn,
|
||||
fq_table("entity_cooccurrences"),
|
||||
fq_table("unit_entities"),
|
||||
fq_table("entities"),
|
||||
bank_id,
|
||||
)
|
||||
return _SweepCounts(orphan_entities_pruned=orphan_pruned, stale_cooccurrences_pruned=stale_pruned)
|
||||
|
||||
# A larger retry budget than the default (3): this is idempotent background
|
||||
# maintenance with no client waiting on it, so a longer retry tail costs
|
||||
# nothing, whereas a dropped sweep silently leaks orphan entities / stale
|
||||
# cooccurrences until the next run. With jittered backoff a single sweep
|
||||
# contending against continuous retain upserts effectively never exhausts
|
||||
# this budget (each retry independently clears with high probability).
|
||||
sweep = await retry_with_backoff(_run_sweep, max_retries=_SWEEP_MAX_RETRIES)
|
||||
result.orphan_entities_pruned = sweep.orphan_entities_pruned
|
||||
result.stale_cooccurrences_pruned = sweep.stale_cooccurrences_pruned
|
||||
|
||||
elapsed = time.time() - job_start
|
||||
logger.info(
|
||||
|
||||
@@ -275,6 +275,8 @@ class MemoryEngineInterface(ABC):
|
||||
*,
|
||||
fact_type: str | None = None,
|
||||
search_query: str | None = None,
|
||||
entity_id: str | None = None,
|
||||
created_before: datetime | None = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
request_context: "RequestContext",
|
||||
@@ -286,6 +288,8 @@ class MemoryEngineInterface(ABC):
|
||||
bank_id: The memory bank ID.
|
||||
fact_type: Filter by fact type.
|
||||
search_query: Full-text search query.
|
||||
entity_id: Filter to memory units linked to this entity ID.
|
||||
created_before: Keep units with ``created_at`` before this instant.
|
||||
limit: Maximum results.
|
||||
offset: Pagination offset.
|
||||
request_context: Request context for authentication.
|
||||
@@ -449,6 +453,7 @@ class MemoryEngineInterface(ABC):
|
||||
bank_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
force_refresh: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Get statistics about memory nodes and links for a bank.
|
||||
@@ -456,6 +461,8 @@ class MemoryEngineInterface(ABC):
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
request_context: Request context for authentication.
|
||||
force_refresh: Bypass the cached value and recompute (also refreshes
|
||||
the cache for subsequent callers).
|
||||
|
||||
Returns:
|
||||
Dict with node_counts, link_counts, link_counts_by_fact_type
|
||||
@@ -562,6 +569,30 @@ class MemoryEngineInterface(ABC):
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete_operation(
|
||||
self,
|
||||
bank_id: str,
|
||||
operation_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Delete a terminal async operation record.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
operation_id: The operation ID to delete.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with success status and message.
|
||||
|
||||
Raises:
|
||||
ValueError: If operation not found.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def update_bank(
|
||||
self,
|
||||
|
||||
@@ -6,11 +6,53 @@ enabling support for multiple LLM backends (OpenAI, Anthropic, Gemini, Codex, et
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from enum import StrEnum
|
||||
from typing import Any, Self
|
||||
|
||||
from .response_models import LLMToolCallResult
|
||||
|
||||
|
||||
class LLMToolChoiceMode(StrEnum):
|
||||
"""Canonical tool-selection modes shared by every LLM provider."""
|
||||
|
||||
AUTO = "auto"
|
||||
NONE = "none"
|
||||
REQUIRED = "required"
|
||||
NAMED = "named"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LLMToolChoice:
|
||||
"""Typed internal tool selection serialized only at provider boundaries."""
|
||||
|
||||
mode: LLMToolChoiceMode
|
||||
function_name: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.mode is LLMToolChoiceMode.NAMED:
|
||||
if self.function_name is None or not self.function_name or self.function_name != self.function_name.strip():
|
||||
raise ValueError("Named tool choice requires a non-empty canonical function name")
|
||||
elif self.function_name is not None:
|
||||
raise ValueError(f"Tool choice mode {self.mode.value!r} cannot include a function name")
|
||||
|
||||
@classmethod
|
||||
def named(cls, function_name: str) -> Self:
|
||||
return cls(mode=LLMToolChoiceMode.NAMED, function_name=function_name)
|
||||
|
||||
@property
|
||||
def selected_function_name(self) -> str:
|
||||
if self.function_name is None:
|
||||
raise ValueError("Tool choice does not select a named function")
|
||||
return self.function_name
|
||||
|
||||
|
||||
LLM_TOOL_CHOICE_AUTO = LLMToolChoice(mode=LLMToolChoiceMode.AUTO)
|
||||
LLM_TOOL_CHOICE_NONE = LLMToolChoice(mode=LLMToolChoiceMode.NONE)
|
||||
LLM_TOOL_CHOICE_REQUIRED = LLMToolChoice(mode=LLMToolChoiceMode.REQUIRED)
|
||||
|
||||
|
||||
class LLMInterface(ABC):
|
||||
"""
|
||||
Abstract interface for LLM providers.
|
||||
@@ -113,8 +155,9 @@ class LLMInterface(ABC):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
cached_prefix: str | None = None,
|
||||
cached_prefix_message_count: int = 0,
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make an LLM API call with tool/function calling support.
|
||||
@@ -128,7 +171,7 @@ class LLMInterface(ABC):
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or specific function.
|
||||
tool_choice: Canonical tool-selection policy.
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
@@ -184,6 +227,45 @@ class LLMInterface(ABC):
|
||||
"""
|
||||
return None
|
||||
|
||||
# ── Step-by-step incremental prompt caching (optional) ─────────────────────
|
||||
#
|
||||
# For agentic loops (reflect) the dominant cost is the conversation prefix
|
||||
# re-sent every turn, not the static system prefix. Providers that can cache
|
||||
# a *growing* prefix implement these: the caller rolls one cache per step
|
||||
# (each covering the previous step's full input), passes its handle plus the
|
||||
# message count it covers to ``call_with_tools`` so only the new turns are
|
||||
# sent fresh, and tears the caches down when the loop ends. Default no-ops so
|
||||
# non-supporting providers transparently run uncached.
|
||||
|
||||
def supports_incremental_prompt_cache(self) -> bool:
|
||||
"""Whether this provider can cache a growing multi-turn conversation prefix."""
|
||||
return False
|
||||
|
||||
async def create_incremental_cache(
|
||||
self,
|
||||
*,
|
||||
session_id: str,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
) -> str | None:
|
||||
"""Cache ``system + tools + messages`` and return an opaque handle, or None.
|
||||
|
||||
The handle is passed back to ``call_with_tools(cached_prefix=...,
|
||||
cached_prefix_message_count=len(messages))``. Caches are grouped under
|
||||
``session_id`` for teardown via ``delete_cache_session``. Returns None
|
||||
when caching is unavailable or the prefix is too small — caller falls
|
||||
back to an uncached call.
|
||||
"""
|
||||
return None
|
||||
|
||||
async def delete_cached_prefix(self, name: str) -> None:
|
||||
"""Best-effort delete of a single cache handle (a superseded step)."""
|
||||
return None
|
||||
|
||||
async def delete_cache_session(self, session_id: str) -> None:
|
||||
"""Best-effort teardown of every cache created under ``session_id``."""
|
||||
return None
|
||||
|
||||
async def submit_batch(
|
||||
self,
|
||||
requests: list[dict[str, Any]],
|
||||
@@ -252,3 +334,11 @@ class OutputTooLongError(Exception):
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class ProviderRateLimitResetError(Exception):
|
||||
"""Raised when an upstream provider says quota will reopen at a known time."""
|
||||
|
||||
def __init__(self, retry_at: datetime, message: str = "") -> None:
|
||||
self.retry_at = retry_at
|
||||
super().__init__(message)
|
||||
|
||||
@@ -76,6 +76,51 @@ _request_ctx: ContextVar[dict[str, Any] | None] = ContextVar("hindsight_llm_requ
|
||||
_call_metadata_ctx: ContextVar[dict[str, Any] | None] = ContextVar("hindsight_llm_call_metadata_ctx", default=None)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMResponseUsage:
|
||||
"""Provider-reported token usage for the in-flight LLM call.
|
||||
|
||||
Stashed by provider implementations as soon as a response is received —
|
||||
*before* local JSON parsing / schema validation, which may still fail. The
|
||||
wrapper reads it to attach real token counts to an error trace when the
|
||||
provider call itself succeeded but the structured output couldn't be parsed
|
||||
or validated (providers charge for those tokens regardless). See #2387.
|
||||
"""
|
||||
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cached_tokens: int = 0
|
||||
|
||||
|
||||
# Per-call provider usage, set by providers right after a response is received.
|
||||
_response_usage_ctx: ContextVar[LLMResponseUsage | None] = ContextVar("hindsight_llm_response_usage_ctx", default=None)
|
||||
|
||||
|
||||
def set_response_usage(usage: LLMResponseUsage | None) -> Token:
|
||||
"""Bind provider-reported usage for the current call. Returns a reset token."""
|
||||
return _response_usage_ctx.set(usage)
|
||||
|
||||
|
||||
def stash_response_usage(usage: LLMResponseUsage | None) -> None:
|
||||
"""Record provider-reported usage so an error trace can attach it later.
|
||||
|
||||
Called by provider implementations once a response (with usage) is in hand,
|
||||
before parsing/validation that may raise. Overwrites any prior value from an
|
||||
earlier retry attempt so the last attempt's usage wins.
|
||||
"""
|
||||
_response_usage_ctx.set(usage)
|
||||
|
||||
|
||||
def reset_response_usage(token: Token) -> None:
|
||||
"""Unwind a binding made by :func:`set_response_usage`."""
|
||||
_response_usage_ctx.reset(token)
|
||||
|
||||
|
||||
def current_response_usage() -> LLMResponseUsage | None:
|
||||
"""Return the active call's provider-reported usage, or None."""
|
||||
return _response_usage_ctx.get()
|
||||
|
||||
|
||||
def set_trace_context(ctx: LLMTraceContext | None) -> Token:
|
||||
"""Bind trace attribution to the current context. Returns a reset token."""
|
||||
return _trace_ctx.set(ctx)
|
||||
@@ -331,6 +376,26 @@ class LLMTraceRecorder:
|
||||
# INSERTs it patches — but it must not block on unrelated operations).
|
||||
self._pending: dict[str | None, set[asyncio.Task]] = {}
|
||||
|
||||
def _writable(self) -> Any | None:
|
||||
"""Return the pool to write through, or None if writing isn't possible.
|
||||
|
||||
Covers the two lifecycle windows in which best-effort trace writes must
|
||||
be skipped rather than attempted: before the backend pool is created
|
||||
(``initialize()`` verifies the LLM before the DB is up) and during/after
|
||||
shutdown. Writes already in flight need no handling — the pools close
|
||||
gracefully, waiting for their connections to be released.
|
||||
"""
|
||||
pool = self._pool_getter()
|
||||
if pool is None:
|
||||
return None
|
||||
# Backends declare readiness explicitly; a raw pool (some callers pass
|
||||
# one directly) has no lifecycle flag and is assumed usable.
|
||||
from .db.base import DatabaseBackend
|
||||
|
||||
if isinstance(pool, DatabaseBackend) and not pool.is_ready:
|
||||
return None
|
||||
return pool
|
||||
|
||||
def is_enabled(self, scope: str) -> bool:
|
||||
"""Whether tracing is active for the given call scope."""
|
||||
if not self._enabled:
|
||||
@@ -428,7 +493,7 @@ class LLMTraceRecorder:
|
||||
|
||||
async def _safe_write(self, record: LLMRequestRecord) -> None:
|
||||
"""Write a trace row. Errors are logged, never raised."""
|
||||
pool = self._pool_getter()
|
||||
pool = self._writable()
|
||||
if pool is None:
|
||||
logger.debug("LLM trace skipped: pool not available")
|
||||
return
|
||||
@@ -523,8 +588,9 @@ class LLMTraceRecorder:
|
||||
# so the UPDATE patches rows that already exist rather than racing ahead
|
||||
# of them (without blocking on unrelated operations' pending writes).
|
||||
await self._flush_pending(trace_id)
|
||||
pool = self._pool_getter()
|
||||
pool = self._writable()
|
||||
if pool is None:
|
||||
logger.debug("LLM trace memory_id attach skipped: pool not available")
|
||||
return
|
||||
try:
|
||||
schema = self._schema_getter()
|
||||
|
||||
@@ -10,9 +10,10 @@ import re
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import AsyncExitStack
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from json_repair import repair_json
|
||||
|
||||
# Vertex AI imports (conditional - for LLMProvider to pass credentials to GeminiLLM)
|
||||
try:
|
||||
from google.oauth2 import service_account
|
||||
@@ -28,13 +29,11 @@ from ..config import (
|
||||
ENV_REFLECT_LLM_MAX_CONCURRENT,
|
||||
ENV_RETAIN_LLM_MAX_CONCURRENT,
|
||||
)
|
||||
from .llm_interface import LLM_TOOL_CHOICE_AUTO, LLMToolChoice, LLMToolChoiceMode
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .response_models import LLMToolCallResult
|
||||
|
||||
# Seed applied to every Groq request for deterministic behavior.
|
||||
DEFAULT_LLM_SEED = 4242
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Disable httpx logging
|
||||
@@ -114,7 +113,7 @@ def _request_params(
|
||||
temperature: float | None = None,
|
||||
scope: str | None = None,
|
||||
response_format: Any | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
tool_choice: LLMToolChoice | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Build the requested-params bag for tracing — only values the caller set.
|
||||
|
||||
@@ -129,8 +128,8 @@ def _request_params(
|
||||
params["temperature"] = temperature
|
||||
if response_format is not None:
|
||||
params["response_schema"] = getattr(response_format, "__name__", None) or "structured"
|
||||
if tool_choice is not None and tool_choice != "auto":
|
||||
params["tool_choice"] = tool_choice if isinstance(tool_choice, str) else "named"
|
||||
if tool_choice is not None and tool_choice.mode is not LLMToolChoiceMode.AUTO:
|
||||
params["tool_choice"] = tool_choice.function_name or tool_choice.mode.value
|
||||
return params or None
|
||||
|
||||
|
||||
@@ -185,6 +184,14 @@ def parse_llm_json(raw: str) -> Any:
|
||||
1. Markdown code fences (```json ... ```) — strip them before parsing.
|
||||
2. Embedded control characters (\\x00-\\x1f, \\x7f) — replace with space
|
||||
and retry if the initial parse fails.
|
||||
3. Structural malformation (trailing commas, unterminated strings, single
|
||||
quotes, invalid ``\\escape`` sequences) — repaired as a last resort via
|
||||
``json_repair`` (#2547/#2544).
|
||||
|
||||
The repair pass is purely *structural*: it fixes JSON that ``json.loads``
|
||||
cannot parse at all. It deliberately does NOT touch content semantics —
|
||||
degenerate-but-valid JSON (repetition loops or leaked scaffolding inside
|
||||
string values) parses fine here and is out of scope for this helper.
|
||||
|
||||
Args:
|
||||
raw: Raw text returned by the LLM.
|
||||
@@ -193,7 +200,8 @@ def parse_llm_json(raw: str) -> Any:
|
||||
Parsed Python object (dict, list, etc.).
|
||||
|
||||
Raises:
|
||||
json.JSONDecodeError: If the text cannot be parsed even after cleanup.
|
||||
json.JSONDecodeError: If the text cannot be parsed even after cleanup
|
||||
and structural repair (e.g. repair yields an empty result).
|
||||
"""
|
||||
text = raw.strip()
|
||||
|
||||
@@ -210,7 +218,19 @@ def parse_llm_json(raw: str) -> Any:
|
||||
# Some models (e.g. Gemini) embed raw control characters inside JSON
|
||||
# string values. Replacing them with a space usually produces valid JSON.
|
||||
cleaned = re.sub(r"[\x00-\x1f\x7f]", " ", text)
|
||||
|
||||
try:
|
||||
return json.loads(cleaned)
|
||||
except json.JSONDecodeError:
|
||||
# Last resort: structural repair of malformed JSON. ``repair_json`` never
|
||||
# raises — unrecoverable input yields an empty result ("" / {} / []). Keep
|
||||
# failing loudly in that case rather than let an empty object masquerade
|
||||
# as a successful parse: callers (retry ladders, the #1833 fail-loud path)
|
||||
# rely on JSONDecodeError to retry or surface the failure.
|
||||
repaired = repair_json(cleaned, return_objects=True)
|
||||
if not repaired:
|
||||
raise
|
||||
return repaired
|
||||
|
||||
|
||||
_PROVIDERS_WITHOUT_API_KEY = frozenset(
|
||||
@@ -236,6 +256,17 @@ def requires_api_key(provider: str) -> bool:
|
||||
return provider.lower() not in _PROVIDERS_WITHOUT_API_KEY
|
||||
|
||||
|
||||
def _validate_ollama_num_ctx(value: Any) -> int | None:
|
||||
"""Validate a native Ollama context-window override."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError(f"ollama_num_ctx must be a positive integer, got {value!r}")
|
||||
if value < 1:
|
||||
raise ValueError(f"ollama_num_ctx must be >= 1, got {value}")
|
||||
return value
|
||||
|
||||
|
||||
def create_llm_provider(
|
||||
provider: str,
|
||||
api_key: str,
|
||||
@@ -253,6 +284,9 @@ def create_llm_provider(
|
||||
gemini_safety_settings: list | None = None,
|
||||
prompt_cache_enabled: bool = False,
|
||||
litellmrouter_config: dict[str, Any] | None = None,
|
||||
gemini_service_tier: str | None = None,
|
||||
timeout: float | None = None,
|
||||
ollama_num_ctx: int | None = None,
|
||||
) -> Any: # Returns LLMInterface
|
||||
"""
|
||||
Factory function to create the appropriate LLM provider implementation.
|
||||
@@ -266,21 +300,34 @@ def create_llm_provider(
|
||||
groq_service_tier: Groq service tier (for Groq provider) - "on_demand", "flex", or "auto".
|
||||
openai_service_tier: OpenAI service tier (for OpenAI provider) - None (default) or "flex" (50% cheaper).
|
||||
bedrock_service_tier: Bedrock service tier (for Bedrock provider) - None (default), "flex", "priority", or "reserved".
|
||||
gemini_service_tier: Gemini service tier (for Gemini provider) - None (default) or "flex" (50% cheaper).
|
||||
ollama_num_ctx: Native Ollama context window override. None lets Ollama use the
|
||||
model/server default.
|
||||
extra_body: Extra request-body params merged into the provider's native
|
||||
call. Threaded into OpenAI-compatible, Fireworks, Anthropic, Gemini/
|
||||
VertexAI and LiteLLM providers (each merges them in its own parameter
|
||||
space). Keys must use each provider's native names (e.g. ``max_tokens``
|
||||
for OpenAI/Anthropic vs ``max_output_tokens`` for Gemini).
|
||||
default_headers: Custom headers passed as ``default_headers`` to provider SDK clients
|
||||
(used by operators routing through proxies / request-tracing middleware). Currently
|
||||
wired into the Anthropic provider; other providers may opt in as needed.
|
||||
default_headers: Custom headers passed to provider SDK clients (used by operators
|
||||
routing through proxies / request-tracing middleware). Wired into the Anthropic
|
||||
provider (SDK ``default_headers``) and the LiteLLM-backed providers — ``litellm``,
|
||||
``litellmrouter`` and ``bedrock`` — as the LiteLLM ``extra_headers`` completion
|
||||
kwarg; other providers may opt in as needed.
|
||||
vertexai_project_id: Vertex AI project ID (for VertexAI provider).
|
||||
vertexai_region: Vertex AI region (for VertexAI provider).
|
||||
vertexai_credentials: Vertex AI credentials object (for VertexAI provider).
|
||||
timeout: Per-request LLM timeout in seconds (resolved by the caller from the
|
||||
per-operation/global config). Threaded into the providers that honour a
|
||||
configurable request timeout (LiteLLM, LiteLLM Router, OpenAI-compatible,
|
||||
Nous). ``None`` lets each provider fall back to its own default
|
||||
(``HINDSIGHT_API_LLM_TIMEOUT`` / ``DEFAULT_LLM_TIMEOUT`` for those four;
|
||||
Anthropic and Gemini keep their provider-specific defaults).
|
||||
|
||||
Returns:
|
||||
LLMInterface implementation for the specified provider.
|
||||
"""
|
||||
ollama_num_ctx = _validate_ollama_num_ctx(ollama_num_ctx)
|
||||
|
||||
from .providers import (
|
||||
AnthropicLLM,
|
||||
ClaudeCodeLLM,
|
||||
@@ -296,6 +343,12 @@ def create_llm_provider(
|
||||
)
|
||||
|
||||
provider_lower = provider.lower()
|
||||
if provider_lower == "gemini":
|
||||
from ..config import parse_gemini_service_tier
|
||||
|
||||
gemini_service_tier = parse_gemini_service_tier(gemini_service_tier)
|
||||
else:
|
||||
gemini_service_tier = None
|
||||
|
||||
if provider_lower == "openai-codex":
|
||||
return CodexLLM(
|
||||
@@ -344,6 +397,7 @@ def create_llm_provider(
|
||||
vertexai_region=vertexai_region,
|
||||
vertexai_credentials=vertexai_credentials,
|
||||
gemini_safety_settings=gemini_safety_settings,
|
||||
gemini_service_tier=gemini_service_tier,
|
||||
prompt_cache_enabled=prompt_cache_enabled,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
@@ -367,6 +421,8 @@ def create_llm_provider(
|
||||
model=model,
|
||||
reasoning_effort=reasoning_effort,
|
||||
extra_body=extra_body,
|
||||
default_headers=default_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
elif provider_lower == "litellmrouter":
|
||||
@@ -385,6 +441,8 @@ def create_llm_provider(
|
||||
config=litellmrouter_config,
|
||||
reasoning_effort=reasoning_effort,
|
||||
extra_body=extra_body,
|
||||
default_headers=default_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
elif provider_lower == "bedrock":
|
||||
@@ -397,7 +455,9 @@ def create_llm_provider(
|
||||
model=bedrock_model,
|
||||
reasoning_effort=reasoning_effort,
|
||||
extra_body=extra_body,
|
||||
default_headers=default_headers,
|
||||
bedrock_service_tier=bedrock_service_tier,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
elif provider_lower == "llamacpp":
|
||||
@@ -444,6 +504,7 @@ def create_llm_provider(
|
||||
model=model,
|
||||
reasoning_effort=reasoning_effort,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
elif provider_lower in (
|
||||
@@ -456,8 +517,10 @@ def create_llm_provider(
|
||||
"deepseek",
|
||||
"volcano",
|
||||
"openrouter",
|
||||
"requesty",
|
||||
"zai",
|
||||
"opencode-go",
|
||||
"atlas",
|
||||
):
|
||||
return OpenAICompatibleLLM(
|
||||
provider=provider,
|
||||
@@ -468,6 +531,8 @@ def create_llm_provider(
|
||||
groq_service_tier=groq_service_tier,
|
||||
openai_service_tier=openai_service_tier,
|
||||
extra_body=extra_body,
|
||||
ollama_num_ctx=ollama_num_ctx,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
else:
|
||||
@@ -496,6 +561,15 @@ class LLMProvider:
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
default_headers: dict[str, str] | None = None,
|
||||
litellmrouter_config: dict[str, Any] | None = None,
|
||||
gemini_service_tier: str | None = None,
|
||||
vertexai_project_id: str | None = None,
|
||||
vertexai_region: str | None = None,
|
||||
vertexai_service_account_key: str | None = None,
|
||||
timeout: float | None = None,
|
||||
max_retries: int | None = None,
|
||||
initial_backoff: float | None = None,
|
||||
max_backoff: float | None = None,
|
||||
ollama_num_ctx: int | None = None,
|
||||
):
|
||||
"""
|
||||
Initialize LLM provider.
|
||||
@@ -509,29 +583,63 @@ class LLMProvider:
|
||||
groq_service_tier: Groq service tier ("on_demand", "flex", "auto") - from config.
|
||||
openai_service_tier: OpenAI service tier (None or "flex") - from config.
|
||||
bedrock_service_tier: Bedrock service tier (None, "flex", "priority", "reserved") - from config.
|
||||
gemini_service_tier: Gemini service tier (None or "flex") - from config.
|
||||
ollama_num_ctx: Native Ollama context window override. ``None`` lets Ollama
|
||||
use the model/server default.
|
||||
gemini_safety_settings: Safety settings for Gemini/VertexAI providers.
|
||||
extra_body: Extra request-body params merged into the provider's native call
|
||||
(OpenAI-compatible, Fireworks, Anthropic, Gemini/VertexAI, LiteLLM).
|
||||
default_headers: Custom headers passed as ``default_headers`` to provider SDK clients.
|
||||
Used by operators routing through proxies / request-tracing middleware. Falls
|
||||
back to ``HindsightConfig.llm_default_headers`` (env: ``HINDSIGHT_API_LLM_DEFAULT_HEADERS``)
|
||||
when ``None``.
|
||||
Used by operators routing through proxies / request-tracing middleware.
|
||||
litellmrouter_config: Provider-specific config for ``provider="litellmrouter"``.
|
||||
JSON object passed verbatim to ``litellm.Router(**config)`` — see
|
||||
https://docs.litellm.ai/docs/routing. Ignored unless ``provider == "litellmrouter"``.
|
||||
When None and the provider is ``litellmrouter``, falls back to
|
||||
``HindsightConfig.llm_litellmrouter_config``.
|
||||
vertexai_project_id: Vertex AI project ID for ``provider="vertexai"`` (required for
|
||||
that provider).
|
||||
vertexai_region: Vertex AI region for ``provider="vertexai"`` (defaults to
|
||||
``"us-central1"`` when ``None``).
|
||||
vertexai_service_account_key: Path to a Vertex AI service-account key file for
|
||||
``provider="vertexai"`` (uses ADC when ``None``).
|
||||
timeout: Per-request LLM timeout in seconds. Resolved by the caller from the
|
||||
per-operation/global config (``retain_llm_timeout`` falling back to
|
||||
``llm_timeout``, etc.). ``None`` lets each provider apply its own default.
|
||||
max_retries: Default retry-attempt budget for ``call`` / ``call_with_tools``
|
||||
when the per-call argument is omitted. Resolved by the caller from the
|
||||
per-operation/global config (``reflect_llm_max_retries`` falling back to
|
||||
``llm_max_retries``, etc.). ``None`` keeps each method's own fallback.
|
||||
initial_backoff: Default initial retry backoff (seconds), same resolution as
|
||||
``max_retries``. ``None`` keeps each method's own fallback.
|
||||
max_backoff: Default maximum retry backoff (seconds), same resolution as
|
||||
``max_retries``. ``None`` keeps each method's own fallback.
|
||||
|
||||
This constructor uses every argument as passed and does not read global
|
||||
``HindsightConfig``: resolving the server-level default for a ``None`` argument is the
|
||||
caller's responsibility (see ``MemoryEngine``'s per-op builds, ``_member_to_llm``, and
|
||||
``LLMProvider.from_env``). Keeping it config-free makes a provider's effective settings a
|
||||
pure function of its arguments — which is what lets each member of a multi-LLM chain be
|
||||
configured independently.
|
||||
"""
|
||||
self.provider = provider.lower()
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.reasoning_effort = reasoning_effort
|
||||
# Per-request timeout (seconds). Used verbatim — the caller resolves the
|
||||
# per-operation/global fallback. ``None`` defers to the provider default.
|
||||
self.timeout = timeout
|
||||
# Default retry policy for call()/call_with_tools(). The caller resolves the
|
||||
# per-operation/global fallback; ``None`` keeps each method's own fallback so
|
||||
# providers built without a resolved config (from_env, tests) are unchanged.
|
||||
self.max_retries = max_retries
|
||||
self.initial_backoff = initial_backoff
|
||||
self.max_backoff = max_backoff
|
||||
self.litellmrouter_config = litellmrouter_config
|
||||
# Service tiers from hierarchical config (not env vars)
|
||||
self.groq_service_tier = groq_service_tier
|
||||
self.openai_service_tier = openai_service_tier
|
||||
self.bedrock_service_tier = bedrock_service_tier
|
||||
self.gemini_service_tier = gemini_service_tier
|
||||
self.ollama_num_ctx = _validate_ollama_num_ctx(ollama_num_ctx)
|
||||
# Gemini safety settings (instance default; can be overridden per-request via context var)
|
||||
self.gemini_safety_settings = gemini_safety_settings
|
||||
# Gemini prompt caching: when True, retain extraction (and any future
|
||||
@@ -542,16 +650,9 @@ class LLMProvider:
|
||||
# Extra body params for OpenAI-compatible providers (e.g. chat_template_kwargs)
|
||||
self.extra_body = extra_body
|
||||
# Default headers passed to provider SDK clients (e.g. proxy auth, request tracing).
|
||||
# Same pattern as ``gemini_safety_settings``: explicit override wins; otherwise read
|
||||
# the static server-level default from ``HindsightConfig`` via ``_get_raw_config()``.
|
||||
# Used verbatim — callers resolve the global fallback (see _member_to_llm /
|
||||
# the per-op builds in MemoryEngine, and LLMProvider.from_env).
|
||||
self.default_headers = default_headers
|
||||
if self.default_headers is None:
|
||||
from ..config import _get_raw_config
|
||||
|
||||
try:
|
||||
self.default_headers = _get_raw_config().llm_default_headers
|
||||
except Exception:
|
||||
pass # Config may not be initialized in test environments
|
||||
|
||||
# Validate provider
|
||||
valid_providers = [
|
||||
@@ -575,8 +676,10 @@ class LLMProvider:
|
||||
"bedrock",
|
||||
"volcano",
|
||||
"openrouter",
|
||||
"requesty",
|
||||
"zai",
|
||||
"opencode-go",
|
||||
"atlas",
|
||||
"fireworks",
|
||||
"nous",
|
||||
]
|
||||
@@ -599,32 +702,31 @@ class LLMProvider:
|
||||
self.base_url = "https://api.deepseek.com"
|
||||
elif self.provider == "openrouter":
|
||||
self.base_url = "https://openrouter.ai/api/v1"
|
||||
elif self.provider == "requesty":
|
||||
self.base_url = "https://router.requesty.ai/v1"
|
||||
elif self.provider == "zai":
|
||||
self.base_url = "https://api.z.ai/api/coding/paas/v4"
|
||||
elif self.provider == "opencode-go":
|
||||
self.base_url = "https://opencode.ai/zen/go/v1"
|
||||
elif self.provider == "atlas":
|
||||
self.base_url = "https://api.atlascloud.ai/v1"
|
||||
elif self.provider == "nous":
|
||||
self.base_url = "https://inference-api.nousresearch.com/v1"
|
||||
|
||||
# Prepare Vertex AI config (if applicable)
|
||||
vertexai_project_id = None
|
||||
vertexai_region = None
|
||||
# Prepare Vertex AI config (if applicable). Values are used as passed; the
|
||||
# caller resolves the global-config fallback (MemoryEngine builds /
|
||||
# _member_to_llm / from_env). The region keeps a constant default here.
|
||||
vertexai_credentials = None
|
||||
|
||||
if self.provider == "vertexai":
|
||||
from ..config import get_config
|
||||
|
||||
config = get_config()
|
||||
|
||||
vertexai_project_id = config.llm_vertexai_project_id
|
||||
if not vertexai_project_id:
|
||||
raise ValueError(
|
||||
"HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID is required for Vertex AI provider. "
|
||||
"Set it to your GCP project ID."
|
||||
)
|
||||
|
||||
vertexai_region = config.llm_vertexai_region or "us-central1"
|
||||
service_account_key = config.llm_vertexai_service_account_key
|
||||
vertexai_region = vertexai_region or "us-central1"
|
||||
service_account_key = vertexai_service_account_key
|
||||
|
||||
# Load explicit service account credentials if provided
|
||||
if service_account_key:
|
||||
@@ -648,45 +750,20 @@ class LLMProvider:
|
||||
f"model={self.model}, auth={'service_account' if service_account_key else 'ADC'}"
|
||||
)
|
||||
|
||||
# For Gemini/VertexAI providers: read safety settings from global config if not explicitly provided
|
||||
# Use _get_raw_config() to bypass StaticConfigProxy (which blocks configurable fields),
|
||||
# since LLMProvider initialization legitimately needs the server-level default.
|
||||
if self.provider in ("gemini", "vertexai") and self.gemini_safety_settings is None:
|
||||
from ..config import _get_raw_config
|
||||
# Normalize the Gemini service tier (pure: maps/validates the passed value,
|
||||
# no global config read). Non-Gemini providers never carry a tier. The
|
||||
# server-level default is resolved by the caller, like the other fields.
|
||||
if self.provider == "gemini":
|
||||
from ..config import parse_gemini_service_tier
|
||||
|
||||
try:
|
||||
raw_config = _get_raw_config()
|
||||
self.gemini_safety_settings = raw_config.llm_gemini_safety_settings
|
||||
except Exception:
|
||||
pass # Config may not be initialized in test environments
|
||||
self.gemini_service_tier = parse_gemini_service_tier(self.gemini_service_tier)
|
||||
else:
|
||||
self.gemini_service_tier = None
|
||||
|
||||
# Prompt-prefix caching is a provider-agnostic toggle (default on): resolve
|
||||
# it from the static server config for every provider when the caller didn't
|
||||
# pass an explicit override. Providers that don't support caching ignore the
|
||||
# value; only those that implement get_or_create_cached_prefix act on it.
|
||||
if not self.prompt_cache_enabled:
|
||||
from ..config import DEFAULT_LLM_PROMPT_CACHE_ENABLED, _get_raw_config
|
||||
|
||||
try:
|
||||
raw_config = _get_raw_config()
|
||||
self.prompt_cache_enabled = bool(
|
||||
getattr(raw_config, "llm_prompt_cache_enabled", DEFAULT_LLM_PROMPT_CACHE_ENABLED)
|
||||
)
|
||||
except Exception:
|
||||
pass # Config may not be initialized in test environments
|
||||
|
||||
# For litellmrouter: prefer an explicit chain from the caller (per-op
|
||||
# construction in MemoryEngine threads the right chain through). If the caller
|
||||
# didn't supply one, fall back to the global ``llm_litellmrouter_config`` so
|
||||
# ad-hoc constructions (e.g. ``LLMProvider.from_env()``) keep working.
|
||||
# gemini_safety_settings / prompt_cache_enabled / litellmrouter_config are
|
||||
# used as passed — the caller resolves the global-config fallback. Providers
|
||||
# that don't support prompt caching ignore the flag.
|
||||
router_config: dict[str, Any] | None = self.litellmrouter_config
|
||||
if self.provider == "litellmrouter" and router_config is None:
|
||||
from ..config import _get_raw_config
|
||||
|
||||
try:
|
||||
router_config = _get_raw_config().llm_litellmrouter_config
|
||||
except Exception:
|
||||
router_config = None
|
||||
|
||||
# Create provider implementation using factory
|
||||
self._provider_impl = create_llm_provider(
|
||||
@@ -698,6 +775,7 @@ class LLMProvider:
|
||||
groq_service_tier=self.groq_service_tier,
|
||||
openai_service_tier=self.openai_service_tier,
|
||||
bedrock_service_tier=self.bedrock_service_tier,
|
||||
gemini_service_tier=self.gemini_service_tier,
|
||||
extra_body=self.extra_body,
|
||||
default_headers=self.default_headers,
|
||||
vertexai_project_id=vertexai_project_id,
|
||||
@@ -706,6 +784,8 @@ class LLMProvider:
|
||||
gemini_safety_settings=self.gemini_safety_settings,
|
||||
prompt_cache_enabled=self.prompt_cache_enabled,
|
||||
litellmrouter_config=router_config,
|
||||
ollama_num_ctx=self.ollama_num_ctx,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
|
||||
# Backward compatibility: Keep mock provider properties
|
||||
@@ -762,11 +842,11 @@ class LLMProvider:
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 60.0,
|
||||
max_retries: int | None = None,
|
||||
initial_backoff: float | None = None,
|
||||
max_backoff: float | None = None,
|
||||
skip_validation: bool = False,
|
||||
strict_schema: bool = False,
|
||||
strict_schema: bool | None = None,
|
||||
return_usage: bool = False,
|
||||
cached_prefix: str | None = None,
|
||||
) -> Any:
|
||||
@@ -779,14 +859,18 @@ class LLMProvider:
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
max_retries: Maximum retry attempts. ``None`` uses the provider's configured
|
||||
default (per-operation/global ``llm_max_retries``), else 10.
|
||||
initial_backoff: Initial backoff time in seconds. ``None`` uses the provider's
|
||||
configured default (``llm_initial_backoff``), else 1.0.
|
||||
max_backoff: Maximum backoff time in seconds. ``None`` uses the provider's
|
||||
configured default (``llm_max_backoff``), else 60.0.
|
||||
skip_validation: Return raw JSON without Pydantic validation.
|
||||
strict_schema: Per-call override requesting grammar-enforced (json_schema strict)
|
||||
structured output instead of the soft json_object path. The server-level
|
||||
HINDSIGHT_API_LLM_STRICT_SCHEMA flag is OR-ed in here so it applies to every call;
|
||||
providers without a strict mode ignore it.
|
||||
structured output instead of the soft json_object path. None (the default)
|
||||
inherits the server-level HINDSIGHT_API_LLM_STRICT_SCHEMA flag; an explicit
|
||||
True or False wins over it, so a caller can force strict output on -- or off --
|
||||
for its own scope. Providers without a strict mode ignore it.
|
||||
return_usage: If True, return tuple (result, TokenUsage) instead of just result.
|
||||
|
||||
Returns:
|
||||
@@ -806,15 +890,33 @@ class LLMProvider:
|
||||
structured = "+structured" if response_format is not None else ""
|
||||
set_stage(f"llm.{self.provider}.{scope}{structured}")
|
||||
|
||||
# Resolve the retry policy: explicit per-call arg wins, else the provider's
|
||||
# configured per-operation/global default, else this method's own fallback.
|
||||
max_retries = (
|
||||
max_retries if max_retries is not None else (self.max_retries if self.max_retries is not None else 10)
|
||||
)
|
||||
initial_backoff = (
|
||||
initial_backoff
|
||||
if initial_backoff is not None
|
||||
else (self.initial_backoff if self.initial_backoff is not None else 1.0)
|
||||
)
|
||||
max_backoff = (
|
||||
max_backoff if max_backoff is not None else (self.max_backoff if self.max_backoff is not None else 60.0)
|
||||
)
|
||||
|
||||
# Resolve strict-schema once, here, rather than in each provider: the
|
||||
# per-call argument OR the server-level HINDSIGHT_API_LLM_STRICT_SCHEMA
|
||||
# flag. Providers with a json_schema response_format (OpenAI-compatible,
|
||||
# per-call argument, falling back to the server-level
|
||||
# HINDSIGHT_API_LLM_STRICT_SCHEMA flag when the caller expressed no
|
||||
# preference. Providers with a json_schema response_format (OpenAI-compatible,
|
||||
# LiteLLM) then grammar-enforce structured output instead of the fragile
|
||||
# soft json_object path; Gemini already enforces its native response_schema,
|
||||
# and providers without a strict mode simply ignore the flag.
|
||||
from ..config import get_config
|
||||
|
||||
strict_schema = strict_schema or get_config().llm_strict_schema
|
||||
# An explicit per-call value wins in BOTH directions -- `or` would have made a
|
||||
# per-call False indistinguishable from "unset", silently ignoring any caller
|
||||
# that opts out while the global flag is on.
|
||||
strict_schema = strict_schema if strict_schema is not None else get_config().llm_strict_schema
|
||||
|
||||
# LLM call observability flows through the OTel GenAI recorder
|
||||
# (tracing.get_span_recorder().record_llm_call). Provider implementations
|
||||
@@ -822,7 +924,13 @@ class LLMProvider:
|
||||
# The requested params are stashed in a contextvar (only what the caller
|
||||
# actually set) so the recorder can attach them to either path.
|
||||
from ..tracing import get_span_recorder
|
||||
from .llm_trace import reset_request_context, set_request_context
|
||||
from .llm_trace import (
|
||||
current_response_usage,
|
||||
reset_request_context,
|
||||
reset_response_usage,
|
||||
set_request_context,
|
||||
set_response_usage,
|
||||
)
|
||||
|
||||
call_start = time.monotonic()
|
||||
request_token = set_request_context(
|
||||
@@ -833,6 +941,9 @@ class LLMProvider:
|
||||
response_format=response_format,
|
||||
)
|
||||
)
|
||||
# Cleared per call; the provider stashes real usage once a response is in
|
||||
# hand so the error path below can attach it if parsing/validation fails.
|
||||
usage_token = set_response_usage(None)
|
||||
try:
|
||||
async with AsyncExitStack() as stack:
|
||||
for sem in _semaphores_for_scope(scope):
|
||||
@@ -860,14 +971,19 @@ class LLMProvider:
|
||||
**cache_kwarg,
|
||||
)
|
||||
except Exception as e:
|
||||
# The provider call may have succeeded (and incurred token
|
||||
# cost) before local parsing/validation raised; attach the
|
||||
# provider-reported usage to the error trace when available.
|
||||
usage = current_response_usage()
|
||||
get_span_recorder().record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
messages=messages,
|
||||
response_content=None,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
input_tokens=usage.input_tokens if usage else 0,
|
||||
output_tokens=usage.output_tokens if usage else 0,
|
||||
cached_tokens=usage.cached_tokens if usage else 0,
|
||||
duration=time.monotonic() - call_start,
|
||||
error=e,
|
||||
)
|
||||
@@ -883,6 +999,7 @@ class LLMProvider:
|
||||
self._mock_calls = self._provider_impl.get_mock_calls()
|
||||
finally:
|
||||
reset_request_context(request_token)
|
||||
reset_response_usage(usage_token)
|
||||
|
||||
return result
|
||||
|
||||
@@ -893,11 +1010,12 @@ class LLMProvider:
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "tools",
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
max_retries: int | None = None,
|
||||
initial_backoff: float | None = None,
|
||||
max_backoff: float | None = None,
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
cached_prefix: str | None = None,
|
||||
cached_prefix_message_count: int = 0,
|
||||
) -> "LLMToolCallResult":
|
||||
"""
|
||||
Make an LLM API call with tool/function calling support.
|
||||
@@ -908,10 +1026,13 @@ class LLMProvider:
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or {"type": "function", "function": {"name": "..."}}
|
||||
max_retries: Maximum retry attempts. ``None`` uses the provider's configured
|
||||
default (per-operation/global ``llm_max_retries``), else 5.
|
||||
initial_backoff: Initial backoff time in seconds. ``None`` uses the provider's
|
||||
configured default (``llm_initial_backoff``), else 1.0.
|
||||
max_backoff: Maximum backoff time in seconds. ``None`` uses the provider's
|
||||
configured default (``llm_max_backoff``), else 30.0.
|
||||
tool_choice: Canonical tool-selection policy.
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
@@ -920,9 +1041,29 @@ class LLMProvider:
|
||||
|
||||
set_stage(f"llm.{self.provider}.{scope}+tools")
|
||||
|
||||
# Resolve the retry policy: explicit per-call arg wins, else the provider's
|
||||
# configured per-operation/global default, else this method's own fallback.
|
||||
max_retries = (
|
||||
max_retries if max_retries is not None else (self.max_retries if self.max_retries is not None else 5)
|
||||
)
|
||||
initial_backoff = (
|
||||
initial_backoff
|
||||
if initial_backoff is not None
|
||||
else (self.initial_backoff if self.initial_backoff is not None else 1.0)
|
||||
)
|
||||
max_backoff = (
|
||||
max_backoff if max_backoff is not None else (self.max_backoff if self.max_backoff is not None else 30.0)
|
||||
)
|
||||
|
||||
# Failures forwarded to the GenAI recorder; successes recorded by providers.
|
||||
from ..tracing import get_span_recorder
|
||||
from .llm_trace import reset_request_context, set_request_context
|
||||
from .llm_trace import (
|
||||
current_response_usage,
|
||||
reset_request_context,
|
||||
reset_response_usage,
|
||||
set_request_context,
|
||||
set_response_usage,
|
||||
)
|
||||
|
||||
call_start = time.monotonic()
|
||||
request_token = set_request_context(
|
||||
@@ -933,15 +1074,23 @@ class LLMProvider:
|
||||
tool_choice=tool_choice,
|
||||
)
|
||||
)
|
||||
# Cleared per call; the provider stashes real usage once a response is in
|
||||
# hand so the error path below can attach it if parsing/validation fails.
|
||||
usage_token = set_response_usage(None)
|
||||
try:
|
||||
async with AsyncExitStack() as stack:
|
||||
for sem in _semaphores_for_scope(scope):
|
||||
await stack.enter_async_context(sem)
|
||||
|
||||
# cached_prefix is only set for providers that returned a handle
|
||||
# from get_or_create_cached_prefix(); forward it only when present
|
||||
# so non-caching providers keep their signature (same as call()).
|
||||
cache_kwarg = {"cached_prefix": cached_prefix} if cached_prefix is not None else {}
|
||||
# from get_or_create_cached_prefix() / create_incremental_cache();
|
||||
# forward it (plus how many leading messages it covers) only when
|
||||
# present so non-caching providers keep their signature.
|
||||
cache_kwarg = (
|
||||
{"cached_prefix": cached_prefix, "cached_prefix_message_count": cached_prefix_message_count}
|
||||
if cached_prefix is not None
|
||||
else {}
|
||||
)
|
||||
try:
|
||||
# Delegate to provider implementation
|
||||
result = await self._provider_impl.call_with_tools(
|
||||
@@ -957,14 +1106,19 @@ class LLMProvider:
|
||||
**cache_kwarg,
|
||||
)
|
||||
except Exception as e:
|
||||
# The provider call may have succeeded (and incurred token
|
||||
# cost) before local parsing/validation raised; attach the
|
||||
# provider-reported usage to the error trace when available.
|
||||
usage = current_response_usage()
|
||||
get_span_recorder().record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
messages=messages,
|
||||
response_content=None,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
input_tokens=usage.input_tokens if usage else 0,
|
||||
output_tokens=usage.output_tokens if usage else 0,
|
||||
cached_tokens=usage.cached_tokens if usage else 0,
|
||||
duration=time.monotonic() - call_start,
|
||||
error=e,
|
||||
)
|
||||
@@ -980,6 +1134,7 @@ class LLMProvider:
|
||||
self._mock_calls = self._provider_impl.get_mock_calls()
|
||||
finally:
|
||||
reset_request_context(request_token)
|
||||
reset_response_usage(usage_token)
|
||||
|
||||
return result
|
||||
|
||||
@@ -1023,7 +1178,9 @@ class LLMProvider:
|
||||
|
||||
def _load_codex_auth(self) -> tuple[str, str]:
|
||||
"""
|
||||
Load OAuth credentials from ~/.codex/auth.json.
|
||||
Load OAuth credentials from the Codex ``auth.json``.
|
||||
|
||||
Honors ``CODEX_HOME`` (falling back to ``~/.codex``).
|
||||
|
||||
Returns:
|
||||
Tuple of (access_token, account_id).
|
||||
@@ -1032,7 +1189,9 @@ class LLMProvider:
|
||||
FileNotFoundError: If auth file doesn't exist.
|
||||
ValueError: If auth file is invalid.
|
||||
"""
|
||||
auth_file = Path.home() / ".codex" / "auth.json"
|
||||
from .providers.codex_auth import default_codex_auth_file
|
||||
|
||||
auth_file = default_codex_auth_file()
|
||||
|
||||
if not auth_file.exists():
|
||||
raise FileNotFoundError(
|
||||
@@ -1134,18 +1293,40 @@ class LLMProvider:
|
||||
@classmethod
|
||||
def from_env(cls) -> "LLMProvider":
|
||||
"""Create provider from environment variables using config.py constants."""
|
||||
# Read every field straight from the environment. The constructor no longer
|
||||
# resolves global-config fallbacks, so this factory must supply them — and it
|
||||
# does so without building the full HindsightConfig, keeping from_env() a
|
||||
# lightweight env-only loader (see test_llm_provider_from_env_keeps_lightweight_loader).
|
||||
from ..config import (
|
||||
DEFAULT_LLM_GROQ_SERVICE_TIER,
|
||||
DEFAULT_LLM_OPENAI_SERVICE_TIER,
|
||||
DEFAULT_LLM_PROMPT_CACHE_ENABLED,
|
||||
DEFAULT_LLM_PROVIDER,
|
||||
DEFAULT_LLM_REASONING_EFFORT,
|
||||
DEFAULT_LLM_TIMEOUT,
|
||||
ENV_LLM_API_KEY,
|
||||
ENV_LLM_BASE_URL,
|
||||
ENV_LLM_BEDROCK_SERVICE_TIER,
|
||||
ENV_LLM_DEFAULT_HEADERS,
|
||||
ENV_LLM_EXTRA_BODY,
|
||||
ENV_LLM_GEMINI_SAFETY_SETTINGS,
|
||||
ENV_LLM_GEMINI_SERVICE_TIER,
|
||||
ENV_LLM_GROQ_SERVICE_TIER,
|
||||
ENV_LLM_LITELLMROUTER_CONFIG,
|
||||
ENV_LLM_MODEL,
|
||||
ENV_LLM_OLLAMA_NUM_CTX,
|
||||
ENV_LLM_OPENAI_SERVICE_TIER,
|
||||
ENV_LLM_PROMPT_CACHE_ENABLED,
|
||||
ENV_LLM_PROVIDER,
|
||||
ENV_LLM_REASONING_EFFORT,
|
||||
ENV_LLM_TIMEOUT,
|
||||
ENV_LLM_VERTEXAI_PROJECT_ID,
|
||||
ENV_LLM_VERTEXAI_REGION,
|
||||
ENV_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY,
|
||||
_get_default_model_for_provider,
|
||||
_parse_llm_router_config,
|
||||
_parse_optional_positive_int,
|
||||
parse_gemini_service_tier,
|
||||
)
|
||||
|
||||
provider = os.getenv(ENV_LLM_PROVIDER, DEFAULT_LLM_PROVIDER)
|
||||
@@ -1162,6 +1343,14 @@ class LLMProvider:
|
||||
model = os.getenv(ENV_LLM_MODEL) or _get_default_model_for_provider(provider)
|
||||
extra_body = json.loads(os.getenv(ENV_LLM_EXTRA_BODY, "null"))
|
||||
default_headers = json.loads(os.getenv(ENV_LLM_DEFAULT_HEADERS, "null"))
|
||||
prompt_cache_enabled = os.getenv(
|
||||
ENV_LLM_PROMPT_CACHE_ENABLED, str(DEFAULT_LLM_PROMPT_CACHE_ENABLED)
|
||||
).lower() in (
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
"on",
|
||||
)
|
||||
|
||||
return cls(
|
||||
provider=provider,
|
||||
@@ -1171,7 +1360,22 @@ class LLMProvider:
|
||||
reasoning_effort=os.getenv(ENV_LLM_REASONING_EFFORT, DEFAULT_LLM_REASONING_EFFORT),
|
||||
extra_body=extra_body,
|
||||
default_headers=default_headers,
|
||||
groq_service_tier=os.getenv(ENV_LLM_GROQ_SERVICE_TIER, DEFAULT_LLM_GROQ_SERVICE_TIER),
|
||||
openai_service_tier=os.getenv(ENV_LLM_OPENAI_SERVICE_TIER, DEFAULT_LLM_OPENAI_SERVICE_TIER),
|
||||
bedrock_service_tier=os.getenv(ENV_LLM_BEDROCK_SERVICE_TIER) or None,
|
||||
gemini_service_tier=(
|
||||
parse_gemini_service_tier(os.getenv(ENV_LLM_GEMINI_SERVICE_TIER))
|
||||
if provider.lower() == "gemini"
|
||||
else None
|
||||
),
|
||||
gemini_safety_settings=json.loads(os.getenv(ENV_LLM_GEMINI_SAFETY_SETTINGS, "null")),
|
||||
prompt_cache_enabled=prompt_cache_enabled,
|
||||
ollama_num_ctx=_parse_optional_positive_int(ENV_LLM_OLLAMA_NUM_CTX, os.getenv(ENV_LLM_OLLAMA_NUM_CTX)),
|
||||
litellmrouter_config=_parse_llm_router_config(ENV_LLM_LITELLMROUTER_CONFIG),
|
||||
vertexai_project_id=os.getenv(ENV_LLM_VERTEXAI_PROJECT_ID) or None,
|
||||
vertexai_region=os.getenv(ENV_LLM_VERTEXAI_REGION) or None,
|
||||
vertexai_service_account_key=os.getenv(ENV_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY) or None,
|
||||
timeout=float(os.getenv(ENV_LLM_TIMEOUT, str(DEFAULT_LLM_TIMEOUT))),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -11,13 +11,20 @@ from one place, so we don't spawn a separate ``asyncio`` task per concern:
|
||||
consolidation operation failed terminally and left them with
|
||||
``consolidated_at IS NULL AND consolidation_failed_at IS NULL`` and nothing to
|
||||
re-trigger them.
|
||||
- **Scheduled mental model refresh** (configurable check cadence, default 60s):
|
||||
refresh mental models whose ``trigger.refresh_cron`` schedule is due, but only
|
||||
when the model is stale (new memories in its scope since its last refresh), so
|
||||
a scheduled tick never burns an LLM call to regenerate identical content. The
|
||||
per-model schedule lives in the cron expression; this loop only decides when to
|
||||
*check*.
|
||||
|
||||
The loop wakes on a short fixed tick and runs each job when its own
|
||||
``last_run + interval`` is due (run-at-start, then on interval), so adding jobs
|
||||
with different cadences doesn't burst CPU. Cross-tenant discovery goes through
|
||||
server-side PL/pgSQL routines (``public.schemas_with_expired_rows`` and
|
||||
``public.banks_needing_consolidation``) — one round-trip each — instead of a
|
||||
per-schema query storm, which matters at thousands of tenants.
|
||||
server-side PL/pgSQL routines (``schemas_with_expired_rows`` and
|
||||
``banks_needing_consolidation``, in the configured schema — see ``fq_routine``) —
|
||||
one round-trip each — instead of a per-schema query storm, which matters at
|
||||
thousands of tenants.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -25,12 +32,14 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
from collections.abc import Coroutine
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from ..config import HindsightConfig, get_config
|
||||
from ..models import RequestContext
|
||||
from .db_utils import acquire_with_retry
|
||||
from .schema import _is_oracle
|
||||
from .schema import _is_oracle, fq_routine, fq_table, fq_table_explicit
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .memory_engine import MemoryEngine
|
||||
@@ -41,6 +50,10 @@ logger = logging.getLogger(__name__)
|
||||
_TICK_SECONDS = 60
|
||||
# Retention sweeps are not time-sensitive; hourly matches the previous per-sweep cadence.
|
||||
_RETENTION_INTERVAL_SECONDS = 3600
|
||||
# Operation cleanup deletes one bounded batch per schema per run, so its cadence
|
||||
# sets the drain rate for a backlog. Kept at one-per-tick (the value it used while
|
||||
# it rode the worker's poll loop) so throughput is unchanged by the move.
|
||||
_OPERATION_CLEANUP_INTERVAL_SECONDS = 60
|
||||
|
||||
|
||||
class MaintenanceLoop:
|
||||
@@ -89,9 +102,14 @@ class MaintenanceLoop:
|
||||
def _any_job_enabled() -> bool:
|
||||
cfg = get_config()
|
||||
reconcile_on = cfg.consolidation_reconcile_interval_seconds > 0
|
||||
audit_on = cfg.audit_log_enabled and cfg.audit_log_retention_days > 0
|
||||
# Not gated on audit_log_enabled: that is per-bank overridable, so rows
|
||||
# can exist even when the deployment default is off. Retention is driven
|
||||
# purely by the (server-level) window.
|
||||
audit_on = cfg.audit_log_retention_days > 0
|
||||
llm_on = cfg.llm_trace_enabled and cfg.llm_trace_retention_days > 0
|
||||
return reconcile_on or audit_on or llm_on
|
||||
mm_refresh_on = cfg.mental_model_refresh_tick_seconds > 0
|
||||
op_cleanup_on = cfg.operation_retention_days > 0
|
||||
return reconcile_on or audit_on or llm_on or mm_refresh_on or op_cleanup_on
|
||||
|
||||
# ── loop ───────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -118,17 +136,37 @@ class MaintenanceLoop:
|
||||
async def _tick(self) -> None:
|
||||
cfg = get_config()
|
||||
if self._is_due("retention", _RETENTION_INTERVAL_SECONDS):
|
||||
await self._run_retention(cfg)
|
||||
await self._run_timed("retention", self._run_retention(cfg))
|
||||
interval = cfg.consolidation_reconcile_interval_seconds
|
||||
if interval > 0 and self._is_due("reconcile", interval):
|
||||
await self._run_reconcile()
|
||||
await self._run_timed("consolidation reconcile", self._run_reconcile())
|
||||
mm_interval = cfg.mental_model_refresh_tick_seconds
|
||||
if mm_interval > 0 and self._is_due("mm_refresh", mm_interval):
|
||||
await self._run_timed("scheduled mental model refresh", self._run_scheduled_mm_refresh())
|
||||
if cfg.operation_retention_days > 0 and self._is_due("operation_cleanup", _OPERATION_CLEANUP_INTERVAL_SECONDS):
|
||||
await self._run_timed("operation cleanup", self._run_operation_cleanup(cfg))
|
||||
|
||||
async def _run_timed(self, name: str, coro: Coroutine[Any, Any, None]) -> None:
|
||||
"""Run a maintenance job and emit one timing line for it.
|
||||
|
||||
Each job keeps its own summary log (counts of work done); this adds a
|
||||
single, uniform line per run so the cost of every sweep is observable.
|
||||
"""
|
||||
start = time.monotonic()
|
||||
try:
|
||||
await coro
|
||||
finally:
|
||||
logger.info(f"Maintenance: {name} took {time.monotonic() - start:.3f}s")
|
||||
|
||||
# ── retention ──────────────────────────────────────────────────────────
|
||||
|
||||
async def _run_retention(self, cfg: HindsightConfig) -> None:
|
||||
# Retention days are static server-level config, so one global cutoff
|
||||
# applies to every tenant schema (the routine sweeps them all).
|
||||
if cfg.audit_log_enabled and cfg.audit_log_retention_days > 0:
|
||||
# Not gated on audit_log_enabled: it is per-bank overridable, so a bank
|
||||
# may be writing audit rows while the deployment default is off. Gating
|
||||
# the purge on the global flag would let those rows accumulate forever.
|
||||
if cfg.audit_log_retention_days > 0:
|
||||
await self._purge_expired("audit_log", "started_at", cfg.audit_log_retention_days)
|
||||
if cfg.llm_trace_enabled and cfg.llm_trace_retention_days > 0:
|
||||
await self._purge_expired("llm_requests", "started_at", cfg.llm_trace_retention_days)
|
||||
@@ -139,7 +177,7 @@ class MaintenanceLoop:
|
||||
try:
|
||||
async with acquire_with_retry(backend, max_retries=1) as conn:
|
||||
rows = await conn.fetch(
|
||||
"SELECT * FROM public.schemas_with_expired_rows($1, $2, $3)", table, ts_col, days
|
||||
f"SELECT * FROM {fq_routine('schemas_with_expired_rows')}($1, $2, $3)", table, ts_col, days
|
||||
)
|
||||
for row in rows:
|
||||
schema = row[0]
|
||||
@@ -154,6 +192,73 @@ class MaintenanceLoop:
|
||||
except Exception as e:
|
||||
logger.warning(f"Retention sweep failed for {table}: {e}")
|
||||
|
||||
# ── terminal operation cleanup ─────────────────────────────────────────
|
||||
|
||||
async def _run_operation_cleanup(self, cfg: HindsightConfig) -> None:
|
||||
"""Prune one bounded batch of expired terminal operations per tenant schema.
|
||||
|
||||
Previously this rode the worker's task-claiming loop, so it only fired
|
||||
when that loop happened to iterate and was interleaved with claiming. It
|
||||
is a periodic housekeeping sweep like the retention jobs above, so it
|
||||
belongs on the same schedule.
|
||||
|
||||
Discovery is one cross-tenant round-trip (``schemas_with_expired_operations``)
|
||||
rather than a connection + prune transaction per tenant; pending and
|
||||
processing rows are never prunable, so a schema holding only in-flight
|
||||
work is correctly reported as having nothing to do.
|
||||
"""
|
||||
engine = self._engine
|
||||
backend = engine._backend
|
||||
try:
|
||||
async with acquire_with_retry(backend, max_retries=1) as conn:
|
||||
rows = await conn.fetch(
|
||||
f"SELECT * FROM {fq_routine('schemas_with_expired_operations')}($1)",
|
||||
cfg.operation_retention_days,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Operation cleanup discovery failed: {e}")
|
||||
return
|
||||
if not rows:
|
||||
return
|
||||
|
||||
# Prune only schemas the deployment actually serves. The routine reports
|
||||
# every schema owning an async_operations table, including ones tenant
|
||||
# discovery doesn't claim.
|
||||
try:
|
||||
tenants = await engine._tenant_extension.list_tenants()
|
||||
except Exception as e:
|
||||
logger.warning(f"Operation cleanup tenant discovery failed: {e}")
|
||||
return
|
||||
known = {t.schema for t in tenants} | {get_config().database_schema}
|
||||
|
||||
from .memory_engine import _current_schema
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=cfg.operation_retention_days)
|
||||
pruned = 0
|
||||
for row in rows:
|
||||
schema = row[0]
|
||||
if schema not in known:
|
||||
continue
|
||||
# Oracle resolves unqualified names from a context-bound session
|
||||
# schema; on PostgreSQL this is harmless and fq_table stays explicit.
|
||||
token = _current_schema.set(schema)
|
||||
try:
|
||||
table = fq_table_explicit("async_operations", schema)
|
||||
async with acquire_with_retry(backend, max_retries=1) as conn:
|
||||
async with conn.transaction():
|
||||
deleted = await backend.ops.prune_terminal_operations(
|
||||
conn, table, cutoff, batch_size=cfg.operation_cleanup_batch_size
|
||||
)
|
||||
if deleted:
|
||||
pruned += deleted
|
||||
logger.info(f"Operation cleanup pruned {deleted} expired terminal operations from {schema}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Operation cleanup failed for schema {schema}: {e}")
|
||||
finally:
|
||||
_current_schema.reset(token)
|
||||
if pruned:
|
||||
logger.info(f"Operation cleanup: pruned {pruned} operation(s) total")
|
||||
|
||||
# ── consolidation reconcile ──────────────────────────────────────────────
|
||||
|
||||
async def _run_reconcile(self) -> None:
|
||||
@@ -161,7 +266,9 @@ class MaintenanceLoop:
|
||||
engine = self._engine
|
||||
try:
|
||||
async with acquire_with_retry(engine._backend, max_retries=1) as conn:
|
||||
rows = await conn.fetch("SELECT schema_name, bank_id FROM public.banks_needing_consolidation()")
|
||||
rows = await conn.fetch(
|
||||
f"SELECT schema_name, bank_id FROM {fq_routine('banks_needing_consolidation')}()"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Consolidation reconcile discovery failed: {e}")
|
||||
return
|
||||
@@ -212,3 +319,112 @@ class MaintenanceLoop:
|
||||
f"Consolidation reconcile: scheduled {submitted} bank(s)"
|
||||
+ (f", skipped {skipped_unknown} in unrecognized schema(s)" if skipped_unknown else "")
|
||||
)
|
||||
|
||||
# ── scheduled mental model refresh ───────────────────────────────────────
|
||||
|
||||
async def _run_scheduled_mm_refresh(self) -> None:
|
||||
"""Refresh mental models whose ``trigger.refresh_cron`` is due.
|
||||
|
||||
Discovery (the set of cron-scheduled models, minus any with an in-flight
|
||||
refresh) is one cross-tenant round-trip via
|
||||
``mental_models_with_cron()``. Cron *due-ness* is evaluated here in
|
||||
Python — a scheduled fire has elapsed when the most recent cron boundary at
|
||||
or before now is later than ``last_refreshed_at`` — because cron arithmetic
|
||||
isn't expressible in plain SQL. Each due model is refreshed only when it is
|
||||
actually stale, so a schedule that fires while nothing changed costs a
|
||||
cheap staleness query, not an LLM call.
|
||||
"""
|
||||
engine = self._engine
|
||||
try:
|
||||
async with acquire_with_retry(engine._backend, max_retries=1) as conn:
|
||||
rows = await conn.fetch(
|
||||
"SELECT schema_name, bank_id, mental_model_id, refresh_cron, last_refreshed_at "
|
||||
f"FROM {fq_routine('mental_models_with_cron')}()"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Scheduled mental model refresh discovery failed: {e}")
|
||||
return
|
||||
if not rows:
|
||||
return
|
||||
|
||||
from croniter import croniter
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
due = []
|
||||
for row in rows:
|
||||
cron = row["refresh_cron"]
|
||||
last = row["last_refreshed_at"]
|
||||
try:
|
||||
prev_fire = croniter(cron, now).get_prev(datetime)
|
||||
except (ValueError, KeyError) as e:
|
||||
logger.warning(
|
||||
f"Scheduled mental model refresh: skipping invalid cron {cron!r} for "
|
||||
f"{row['schema_name']}/{row['mental_model_id']}: {e}"
|
||||
)
|
||||
continue
|
||||
if last is None or prev_fire > last:
|
||||
due.append(row)
|
||||
if not due:
|
||||
return
|
||||
|
||||
# Only enqueue into schemas the worker actually polls (tenant discovery),
|
||||
# otherwise the op would never be claimed. The tenant_id (when provided)
|
||||
# lets config resolution honor tenant-level overrides.
|
||||
try:
|
||||
tenants = await engine._tenant_extension.list_tenants()
|
||||
except Exception as e:
|
||||
logger.warning(f"Scheduled mental model refresh tenant discovery failed: {e}")
|
||||
return
|
||||
tenant_by_schema = {t.schema: t for t in tenants}
|
||||
default_schema = get_config().database_schema
|
||||
|
||||
from .memory_engine import _current_schema
|
||||
|
||||
submitted = 0
|
||||
skipped_unknown = 0
|
||||
skipped_fresh = 0
|
||||
for row in due:
|
||||
schema = row["schema_name"]
|
||||
bank_id = row["bank_id"]
|
||||
mm_id = row["mental_model_id"]
|
||||
tenant = tenant_by_schema.get(schema)
|
||||
if tenant is None and schema != default_schema:
|
||||
skipped_unknown += 1
|
||||
continue
|
||||
tenant_id = tenant.tenant_id if tenant else None
|
||||
token = _current_schema.set(schema)
|
||||
try:
|
||||
context = RequestContext(internal=True, tenant_id=tenant_id)
|
||||
# Skip if nothing in the model's scope changed since its last
|
||||
# refresh — a scheduled refresh must not regenerate identical
|
||||
# content. compute_mental_model_is_stale needs the model's tags +
|
||||
# trigger, which the discovery routine doesn't return, so re-read
|
||||
# the row under the bank's schema context.
|
||||
async with acquire_with_retry(engine._backend, max_retries=1) as conn:
|
||||
mm_row = await conn.fetchrow(
|
||||
f"SELECT id, tags, trigger, last_refreshed_at FROM {fq_table('mental_models')} "
|
||||
"WHERE bank_id = $1 AND id = $2",
|
||||
bank_id,
|
||||
mm_id,
|
||||
)
|
||||
if mm_row is None:
|
||||
continue
|
||||
is_stale = await engine.compute_mental_model_is_stale(conn, bank_id, mm_row)
|
||||
if not is_stale:
|
||||
skipped_fresh += 1
|
||||
continue
|
||||
await engine.submit_async_refresh_mental_model(
|
||||
bank_id=bank_id, mental_model_id=mm_id, request_context=context
|
||||
)
|
||||
submitted += 1
|
||||
except Exception as e:
|
||||
logger.warning(f"Scheduled mental model refresh failed for {mm_id} in {schema}: {e}")
|
||||
finally:
|
||||
_current_schema.reset(token)
|
||||
|
||||
if submitted or skipped_unknown or skipped_fresh:
|
||||
logger.info(
|
||||
f"Scheduled mental model refresh: scheduled {submitted} model(s)"
|
||||
+ (f", {skipped_fresh} up-to-date" if skipped_fresh else "")
|
||||
+ (f", skipped {skipped_unknown} in unrecognized schema(s)" if skipped_unknown else "")
|
||||
)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,204 @@
|
||||
"""Multi-LLM routing: failover and (weighted) round-robin across N providers.
|
||||
|
||||
``MultiLLMProvider`` wraps an ordered list of :class:`LLMProvider` members and a
|
||||
:class:`~hindsight_api.config.LLMStrategyConfig`, exposing the same public surface
|
||||
as a single ``LLMProvider`` so it drops into every existing call path (including
|
||||
``with_config()`` / ``ConfiguredLLMProvider``).
|
||||
|
||||
Member 0 is the **primary** (the operation's unindexed/base LLM); members 1..N are
|
||||
the indexed extras (``HINDSIGHT_API_<OP>LLM_<n>_*``). Each member keeps its own
|
||||
internal retry budget, so we only advance to the next member after a member has
|
||||
exhausted its retries and raised.
|
||||
|
||||
Strategies:
|
||||
- ``failover``: try members in declared order ``[0..N]``.
|
||||
- ``round-robin``: rotate the starting member per request (optionally weighted),
|
||||
then fall through the remaining members on error.
|
||||
|
||||
Batch retain and any direct ``_provider_impl`` access operate on the **primary
|
||||
member only** (via attribute passthrough) — failover/round-robin apply to the
|
||||
interactive ``call`` / ``call_with_tools`` paths.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from ..config import LLM_STRATEGY_FAILOVER, LLMStrategyConfig
|
||||
from .llm_wrapper import LLMProvider, OutputTooLongError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .llm_wrapper import ConfiguredLLMProvider, LLMToolCallResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _should_failover(exc: BaseException) -> bool:
|
||||
"""Whether ``exc`` from one member should trigger a try on the next member.
|
||||
|
||||
Generic ``Exception`` instances (network errors, provider 5xx, timeouts after
|
||||
a member's own retries) fail over. ``OutputTooLongError`` is propagated — a
|
||||
different provider won't fit an over-length output either. ``CancelledError``,
|
||||
``KeyboardInterrupt`` and ``SystemExit`` are ``BaseException`` (not
|
||||
``Exception``) and therefore propagate unchanged.
|
||||
"""
|
||||
if isinstance(exc, OutputTooLongError):
|
||||
return False
|
||||
return isinstance(exc, Exception)
|
||||
|
||||
|
||||
class _WeightedRoundRobin:
|
||||
"""Smooth weighted round-robin scheduler (nginx SWRR).
|
||||
|
||||
Produces a starting member index per request such that, over time, member
|
||||
``i`` is chosen in proportion to ``weights[i]`` while keeping selections
|
||||
interleaved rather than bursty. Uniform weights degrade to plain round-robin.
|
||||
The tiny selection critical section is mutex-guarded so concurrent callers
|
||||
don't corrupt the running totals (they may still interleave, which only
|
||||
affects distribution, never correctness).
|
||||
"""
|
||||
|
||||
def __init__(self, weights: list[int]) -> None:
|
||||
self._weights = list(weights)
|
||||
self._current = [0] * len(weights)
|
||||
self._total = sum(weights)
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def next(self) -> int:
|
||||
with self._lock:
|
||||
best = 0
|
||||
for i, w in enumerate(self._weights):
|
||||
self._current[i] += w
|
||||
if self._current[i] > self._current[best]:
|
||||
best = i
|
||||
self._current[best] -= self._total
|
||||
return best
|
||||
|
||||
|
||||
class MultiLLMProvider:
|
||||
"""Route LLM calls across multiple members per a failover / round-robin strategy."""
|
||||
|
||||
def __init__(self, members: list[LLMProvider], strategy: LLMStrategyConfig) -> None:
|
||||
if not members:
|
||||
raise ValueError("MultiLLMProvider requires at least one member")
|
||||
self._members = members
|
||||
self._strategy = strategy
|
||||
|
||||
weights = strategy.weights or [1] * len(members)
|
||||
if len(weights) != len(members):
|
||||
raise ValueError(
|
||||
f"LLM strategy 'weights' has {len(weights)} entries but the chain has "
|
||||
f"{len(members)} members (primary + indexed); they must match."
|
||||
)
|
||||
self._scheduler = _WeightedRoundRobin(weights)
|
||||
|
||||
# ── routing ────────────────────────────────────────────────────────────────
|
||||
|
||||
def _member_order(self) -> list[int]:
|
||||
"""Indices to try, in order, for one request."""
|
||||
n = len(self._members)
|
||||
if self._strategy.mode == LLM_STRATEGY_FAILOVER:
|
||||
return list(range(n))
|
||||
start = self._scheduler.next()
|
||||
return [(start + i) % n for i in range(n)]
|
||||
|
||||
async def _dispatch(self, method_name: str, **kwargs: Any) -> Any:
|
||||
last_exc: BaseException | None = None
|
||||
order = self._member_order()
|
||||
for position, idx in enumerate(order):
|
||||
member = self._members[idx]
|
||||
try:
|
||||
return await getattr(member, method_name)(**kwargs)
|
||||
except BaseException as e: # noqa: BLE001 - re-raised unless it should fail over
|
||||
if not _should_failover(e):
|
||||
raise
|
||||
last_exc = e
|
||||
remaining = len(order) - position - 1
|
||||
logger.warning(
|
||||
"LLM member %d (%s/%s) failed on %s: %s%s",
|
||||
idx,
|
||||
member.provider,
|
||||
member.model,
|
||||
method_name,
|
||||
e,
|
||||
f"; trying next member ({remaining} left)" if remaining else "; no members left",
|
||||
)
|
||||
# All members failed; surface the last error (loop ran at least once).
|
||||
assert last_exc is not None
|
||||
raise last_exc
|
||||
|
||||
async def call(self, messages: list[dict[str, Any]], **kwargs: Any) -> Any:
|
||||
return await self._dispatch("call", messages=messages, **kwargs)
|
||||
|
||||
async def call_with_tools(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
**kwargs: Any,
|
||||
) -> "LLMToolCallResult":
|
||||
return await self._dispatch("call_with_tools", messages=messages, tools=tools, **kwargs)
|
||||
|
||||
# ── lifecycle ────────────────────────────────────────────────────────────────
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""Strictly verify the primary; soft-verify the rest (warn, don't fail).
|
||||
|
||||
A failover member being unreachable at startup must not block the server —
|
||||
it may come back before it's needed. The primary is the steady-state path,
|
||||
so its failure is still surfaced (the caller already wraps this in a
|
||||
warn-only try/except at startup).
|
||||
"""
|
||||
await self._members[0].verify_connection()
|
||||
for member in self._members[1:]:
|
||||
try:
|
||||
await member.verify_connection()
|
||||
except Exception as e: # noqa: BLE001 - soft verification
|
||||
logger.warning(
|
||||
"Failover LLM member %s/%s failed connection verification: %s. "
|
||||
"It will be tried at request time if the primary fails.",
|
||||
member.provider,
|
||||
member.model,
|
||||
e,
|
||||
)
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
for member in self._members:
|
||||
await member.cleanup()
|
||||
|
||||
def with_config(
|
||||
self,
|
||||
config: Any,
|
||||
*,
|
||||
bank_id: str | None = None,
|
||||
operation: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> "ConfiguredLLMProvider":
|
||||
"""Mirror ``LLMProvider.with_config`` so the strategy runs inside the
|
||||
per-operation configured wrapper (gemini-safety + trace contextvars wrap
|
||||
every member call)."""
|
||||
from .llm_trace import LLMTraceContext
|
||||
from .llm_wrapper import ConfiguredLLMProvider
|
||||
|
||||
trace_ctx = None
|
||||
if bank_id is not None or operation is not None or metadata:
|
||||
trace_ctx = LLMTraceContext(
|
||||
bank_id=bank_id,
|
||||
operation=operation,
|
||||
metadata=dict(metadata or {}),
|
||||
trace_id=str(uuid.uuid4()),
|
||||
operation_span_id=str(uuid.uuid4()),
|
||||
)
|
||||
return ConfiguredLLMProvider(self, config.llm_gemini_safety_settings, trace_ctx)
|
||||
|
||||
# ── attribute passthrough ────────────────────────────────────────────────────
|
||||
|
||||
@property
|
||||
def members(self) -> list[LLMProvider]:
|
||||
return self._members
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
# Anything not defined here (provider, model, api_key, base_url,
|
||||
# _provider_impl, mock helpers, batch helpers, ...) delegates to the
|
||||
# primary member so existing call sites keep working unchanged.
|
||||
return getattr(object.__getattribute__(self, "_members")[0], name)
|
||||
@@ -142,3 +142,21 @@ class RefreshMentalModelMetadata:
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
"""Convert to dict for JSON serialization."""
|
||||
return asdict(self)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RefreshMentalModelOutcomeMetadata:
|
||||
"""Machine-readable outcome metadata for a completed refresh_mental_model operation.
|
||||
|
||||
Refresh parity with RetainOutcomeMetadata (#2605): lets a monitoring layer
|
||||
distinguish "refreshed with real content" from "refreshed empty" by reading
|
||||
result_metadata alone, without a follow-up content fetch.
|
||||
"""
|
||||
|
||||
content_len: int
|
||||
populated_content: bool
|
||||
based_on_counts: dict[str, int] = field(default_factory=dict)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
"""Convert to dict for JSON serialization."""
|
||||
return asdict(self)
|
||||
|
||||
@@ -3,43 +3,138 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from hindsight_api.config import DEFAULT_FILE_PARSER_MARKITDOWN_OCR_PROMPT
|
||||
|
||||
from .base import FileParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from markitdown import StreamInfo
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Extensions whose markitdown converters decode the raw bytes as text. markitdown
|
||||
# samples only the first chunk for charset detection, so a UTF-8 file with a long
|
||||
# ASCII-only prefix is mis-detected as ASCII; the JSON/ipynb converter then crashes
|
||||
# decoding the first multibyte byte. Passing an explicit UTF-8 hint when the bytes
|
||||
# are valid UTF-8 sidesteps the faulty detection without affecting other encodings.
|
||||
_TEXT_EXTENSIONS = {
|
||||
".json",
|
||||
".jsonl",
|
||||
".ipynb",
|
||||
".txt",
|
||||
".text",
|
||||
".md",
|
||||
".markdown",
|
||||
".csv",
|
||||
".html",
|
||||
".htm",
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MarkitdownOcrOptions:
|
||||
"""OpenAI-compatible OCR options passed through to MarkItDown."""
|
||||
|
||||
# Keep this typed as object so the OpenAI SDK import stays lazy for non-OCR users.
|
||||
llm_client: object
|
||||
llm_model: str
|
||||
llm_prompt: str
|
||||
|
||||
|
||||
class MarkitdownParser(FileParser):
|
||||
"""
|
||||
Markitdown file parser.
|
||||
|
||||
Uses Microsoft's markitdown library to convert various file formats
|
||||
to markdown including PDF, Office docs, images (via OCR), audio, HTML.
|
||||
to markdown including PDF, Office docs, images with optional OCR,
|
||||
audio, HTML.
|
||||
|
||||
Supported formats:
|
||||
- PDF (.pdf)
|
||||
- Word (.docx, .doc)
|
||||
- PowerPoint (.pptx, .ppt)
|
||||
- Excel (.xlsx, .xls)
|
||||
- Images (.jpg, .jpeg, .png) - with OCR
|
||||
- Images (.jpg, .jpeg, .png) - optional OCR
|
||||
- HTML (.html, .htm)
|
||||
- Text (.txt, .md)
|
||||
- Audio (.mp3, .wav) - with transcription
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
ocr_enabled: bool = False,
|
||||
ocr_api_key: str | None = None,
|
||||
ocr_base_url: str | None = None,
|
||||
ocr_model: str | None = None,
|
||||
ocr_prompt: str | None = None,
|
||||
):
|
||||
"""Initialize markitdown parser."""
|
||||
# Lazy import to avoid requiring markitdown for all users
|
||||
try:
|
||||
from markitdown import MarkItDown
|
||||
|
||||
self._markitdown = MarkItDown()
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"markitdown package is required for file parsing. Install with: pip install markitdown"
|
||||
) from e
|
||||
|
||||
self._ocr_enabled = ocr_enabled
|
||||
if ocr_enabled:
|
||||
ocr_options = self._build_ocr_options(
|
||||
api_key=ocr_api_key,
|
||||
base_url=ocr_base_url,
|
||||
model=ocr_model,
|
||||
prompt=ocr_prompt,
|
||||
)
|
||||
self._markitdown = MarkItDown(
|
||||
llm_client=ocr_options.llm_client,
|
||||
llm_model=ocr_options.llm_model,
|
||||
llm_prompt=ocr_options.llm_prompt,
|
||||
)
|
||||
else:
|
||||
self._markitdown = MarkItDown()
|
||||
|
||||
def _build_ocr_options(
|
||||
self,
|
||||
*,
|
||||
api_key: str | None,
|
||||
base_url: str | None,
|
||||
model: str | None,
|
||||
prompt: str | None,
|
||||
) -> MarkitdownOcrOptions:
|
||||
"""Build MarkItDown options for OpenAI-compatible image OCR."""
|
||||
if not model or not model.strip():
|
||||
raise ValueError(
|
||||
"Markitdown OCR is enabled but no model is configured. "
|
||||
"Set HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_MODEL to an OpenAI-compatible OCR/vision model "
|
||||
"with image-input support."
|
||||
)
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"Markitdown OCR is enabled but no API key is configured. "
|
||||
"Set HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_API_KEY."
|
||||
)
|
||||
if not base_url or not base_url.strip():
|
||||
raise ValueError(
|
||||
"Markitdown OCR is enabled but no base URL is configured. "
|
||||
"Set HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_BASE_URL to an OpenAI-compatible OCR/vision endpoint."
|
||||
)
|
||||
|
||||
try:
|
||||
from openai import OpenAI
|
||||
except ImportError as e:
|
||||
raise RuntimeError("openai package is required when Markitdown OCR is enabled.") from e
|
||||
|
||||
return MarkitdownOcrOptions(
|
||||
llm_client=OpenAI(api_key=api_key, base_url=base_url.strip()),
|
||||
llm_model=model.strip(),
|
||||
llm_prompt=prompt or DEFAULT_FILE_PARSER_MARKITDOWN_OCR_PROMPT,
|
||||
)
|
||||
|
||||
async def convert(self, file_data: bytes, filename: str) -> str:
|
||||
"""Parse file to markdown using markitdown."""
|
||||
# markitdown is synchronous, so we run it in executor to avoid blocking
|
||||
@@ -48,14 +143,22 @@ class MarkitdownParser(FileParser):
|
||||
|
||||
def _convert_sync(self, file_data: bytes, filename: str) -> str:
|
||||
"""Synchronous parsing (runs in thread pool)."""
|
||||
if self._is_image_file(filename) and not self._ocr_enabled:
|
||||
raise RuntimeError(
|
||||
"Image OCR is not enabled for the markitdown parser. "
|
||||
"Set HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_ENABLED=true and configure an OpenAI-compatible "
|
||||
"OCR/vision endpoint with image-input support, or choose an OCR-capable parser."
|
||||
)
|
||||
|
||||
# Write to temp file (markitdown requires file path)
|
||||
with tempfile.NamedTemporaryFile(suffix=Path(filename).suffix, delete=False) as tmp:
|
||||
tmp.write(file_data)
|
||||
tmp_path = tmp.name
|
||||
|
||||
try:
|
||||
# Parse using markitdown
|
||||
result = self._markitdown.convert(tmp_path)
|
||||
# Parse using markitdown, passing an explicit charset hint for text
|
||||
# files to avoid markitdown's sample-based (and crash-prone) detection.
|
||||
result = self._markitdown.convert(tmp_path, stream_info=self._utf8_stream_info(file_data, filename))
|
||||
|
||||
if not result or not result.text_content:
|
||||
raise RuntimeError(f"No content extracted from '{filename}'")
|
||||
@@ -73,6 +176,32 @@ class MarkitdownParser(FileParser):
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _utf8_stream_info(file_data: bytes, filename: str) -> "StreamInfo | None":
|
||||
"""Return a UTF-8 charset hint for text files that decode cleanly as UTF-8.
|
||||
|
||||
Returns None for binary files or non-UTF-8 text so markitdown falls back
|
||||
to its own detection.
|
||||
"""
|
||||
if Path(filename).suffix.lower() not in _TEXT_EXTENSIONS:
|
||||
return None
|
||||
try:
|
||||
# file_data may arrive as a non-``bytes`` buffer (e.g. a memoryview or
|
||||
# a native/Rust-backed buffer object) that has no ``.decode``; coerce
|
||||
# through the buffer protocol before the UTF-8 probe. The ``tmp.write``
|
||||
# in the caller already relies only on the same buffer protocol.
|
||||
bytes(file_data).decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
return None
|
||||
from markitdown import StreamInfo
|
||||
|
||||
return StreamInfo(charset="utf-8")
|
||||
|
||||
@staticmethod
|
||||
def _is_image_file(filename: str) -> bool:
|
||||
"""Return whether the file type needs OCR to extract useful text."""
|
||||
return Path(filename).suffix.lower() in {".jpg", ".jpeg", ".png"}
|
||||
|
||||
def supports(self, filename: str, content_type: str | None = None) -> bool:
|
||||
"""Check if markitdown supports this file type."""
|
||||
# Supported extensions (from markitdown docs)
|
||||
@@ -85,7 +214,7 @@ class MarkitdownParser(FileParser):
|
||||
".ppt",
|
||||
".xlsx",
|
||||
".xls",
|
||||
# Images (with OCR)
|
||||
# Images (optional OCR)
|
||||
".jpg",
|
||||
".jpeg",
|
||||
".png",
|
||||
|
||||
@@ -14,13 +14,64 @@ import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface
|
||||
from hindsight_api.engine.llm_interface import LLM_TOOL_CHOICE_AUTO, LLMInterface, LLMToolChoice
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.providers.llm_debug import dump_request_on_4xx
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _usage_from_anthropic_response(response: Any) -> LLMResponseUsage:
|
||||
"""Extract input/output/cached token counts from an Anthropic usage block."""
|
||||
usage = getattr(response, "usage", None)
|
||||
if not usage:
|
||||
return LLMResponseUsage()
|
||||
return LLMResponseUsage(
|
||||
input_tokens=usage.input_tokens or 0,
|
||||
output_tokens=usage.output_tokens or 0,
|
||||
cached_tokens=getattr(usage, "cache_read_input_tokens", 0) or 0,
|
||||
)
|
||||
|
||||
|
||||
_EPHEMERAL_CACHE = {"type": "ephemeral"}
|
||||
|
||||
|
||||
def _cached_system_blocks(system_prompt: str) -> list[dict[str, Any]]:
|
||||
"""Render the system prompt as a block list with a cache_control marker.
|
||||
|
||||
Anthropic prompt caching is a prefix match: marking the (single) system
|
||||
block caches tools + system together. The system prompt is stable per
|
||||
scope — fact extraction reuses it across every chunk, reflect and
|
||||
consolidation keep their stable instructions there — so repeat calls read
|
||||
it at ~10% of the base input price. Markers below the model's minimum
|
||||
cacheable prefix are silently ignored (no write premium), so marking is
|
||||
safe unconditionally. This is the "inline-marker provider" strategy that
|
||||
``LLMInterface.get_or_create_cached_prefix`` documents for Anthropic.
|
||||
"""
|
||||
return [{"type": "text", "text": system_prompt, "cache_control": _EPHEMERAL_CACHE}]
|
||||
|
||||
|
||||
def _mark_last_message_for_caching(messages: list[dict[str, Any]]) -> None:
|
||||
"""Add a cache_control marker to the final content block, in place.
|
||||
|
||||
Used on the multi-turn (tool-calling) path: the reflect agent loop resends
|
||||
the entire growing conversation each iteration, so this request's
|
||||
end-marker becomes the next iteration's cache read point. Together with
|
||||
the system marker this uses 2 of the 4 allowed breakpoints.
|
||||
"""
|
||||
if not messages:
|
||||
return
|
||||
last = messages[-1]
|
||||
content = last.get("content")
|
||||
if isinstance(content, str):
|
||||
if content.strip(): # the API rejects empty text blocks
|
||||
last["content"] = [{"type": "text", "text": content, "cache_control": _EPHEMERAL_CACHE}]
|
||||
elif isinstance(content, list) and content and isinstance(content[-1], dict):
|
||||
content[-1]["cache_control"] = _EPHEMERAL_CACHE
|
||||
|
||||
|
||||
class AnthropicLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using Anthropic's Claude models.
|
||||
@@ -136,7 +187,9 @@ class AnthropicLLM(LLMInterface):
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
skip_validation: Return raw JSON without Pydantic validation.
|
||||
strict_schema: Use strict JSON schema enforcement (not supported by Anthropic).
|
||||
strict_schema: Route structured output through a forced tool_use tool for
|
||||
native constrained decoding (issue #1002). When False, falls back to
|
||||
schema-in-prompt + JSON parse.
|
||||
return_usage: If True, return tuple (result, TokenUsage) instead of just result.
|
||||
|
||||
Returns:
|
||||
@@ -167,14 +220,21 @@ class AnthropicLLM(LLMInterface):
|
||||
else:
|
||||
anthropic_messages.append({"role": role, "content": content})
|
||||
|
||||
# Add JSON schema instruction if response_format is provided
|
||||
# Structured output: prefer Anthropic-native constrained decoding via a single
|
||||
# forced tool_use tool (strict_schema) over text-injecting the schema and
|
||||
# parsing the reply. Native constrained decoding guarantees schema-valid JSON,
|
||||
# eliminating the invalid-JSON retry storm (issue #1002). When strict_schema is
|
||||
# off we keep the text-inject + json.loads fallback for backward compatibility.
|
||||
schema = None
|
||||
use_forced_tool = False
|
||||
_tool_name = "structured_response"
|
||||
if response_format is not None and hasattr(response_format, "model_json_schema"):
|
||||
schema = response_format.model_json_schema()
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
if system_prompt:
|
||||
system_prompt += schema_msg
|
||||
if strict_schema:
|
||||
use_forced_tool = True
|
||||
else:
|
||||
system_prompt = schema_msg
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
system_prompt = (system_prompt + schema_msg) if system_prompt else schema_msg
|
||||
|
||||
# Prepare parameters
|
||||
call_params: dict[str, Any] = {
|
||||
@@ -184,7 +244,17 @@ class AnthropicLLM(LLMInterface):
|
||||
}
|
||||
|
||||
if system_prompt:
|
||||
call_params["system"] = system_prompt
|
||||
# One-shot calls share only the system prompt with each other, so
|
||||
# that is the sole cache breakpoint on this path.
|
||||
call_params["system"] = _cached_system_blocks(system_prompt)
|
||||
|
||||
if use_forced_tool:
|
||||
# Single tool whose input_schema IS the response schema; force the model to
|
||||
# emit it via tool_choice so the SDK does constrained decoding for us.
|
||||
call_params["tools"] = [
|
||||
{"name": _tool_name, "description": "Return the structured response.", "input_schema": schema}
|
||||
]
|
||||
call_params["tool_choice"] = {"type": "tool", "name": _tool_name}
|
||||
|
||||
if self._extra_body:
|
||||
call_params["extra_body"] = self._extra_body
|
||||
@@ -194,40 +264,61 @@ class AnthropicLLM(LLMInterface):
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.messages.create(**call_params)
|
||||
# Stash usage before parse/validate, which may raise locally
|
||||
# even though the provider charged for these tokens (#2387).
|
||||
stash_response_usage(_usage_from_anthropic_response(response))
|
||||
|
||||
# Anthropic response content is a list of blocks
|
||||
content = ""
|
||||
for block in response.content:
|
||||
if block.type == "text":
|
||||
content += block.text
|
||||
|
||||
if response_format is not None:
|
||||
# Models may wrap JSON in markdown code blocks
|
||||
clean_content = content
|
||||
if "```json" in content:
|
||||
clean_content = content.split("```json")[1].split("```")[0].strip()
|
||||
elif "```" in content:
|
||||
clean_content = content.split("```")[1].split("```")[0].strip()
|
||||
|
||||
try:
|
||||
json_data = json.loads(clean_content)
|
||||
except json.JSONDecodeError:
|
||||
# Fallback to parsing raw content if markdown stripping failed
|
||||
json_data = json.loads(content)
|
||||
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
else:
|
||||
result = response_format.model_validate(json_data)
|
||||
if use_forced_tool:
|
||||
# Forced tool_use → the validated args are already a dict; no parsing,
|
||||
# no markdown-strip, no JSON-decode retry possible.
|
||||
tool_input = None
|
||||
for block in response.content:
|
||||
if block.type == "tool_use" and block.name == _tool_name:
|
||||
tool_input = block.input or {}
|
||||
break
|
||||
if tool_input is None:
|
||||
# Model ignored the forced tool (rare, e.g. a gateway that drops
|
||||
# tool_choice). Fall back to text parse so we don't hard-fail; the
|
||||
# existing retry loop still covers genuine errors.
|
||||
content = "".join(b.text for b in response.content if b.type == "text")
|
||||
tool_input = json.loads(content)
|
||||
content = json.dumps(tool_input)
|
||||
result = tool_input if skip_validation else response_format.model_validate(tool_input)
|
||||
else:
|
||||
result = content
|
||||
# Anthropic response content is a list of blocks
|
||||
content = ""
|
||||
for block in response.content:
|
||||
if block.type == "text":
|
||||
content += block.text
|
||||
|
||||
if response_format is not None:
|
||||
# Models may wrap JSON in markdown code blocks
|
||||
clean_content = content
|
||||
if "```json" in content:
|
||||
clean_content = content.split("```json")[1].split("```")[0].strip()
|
||||
elif "```" in content:
|
||||
clean_content = content.split("```")[1].split("```")[0].strip()
|
||||
|
||||
try:
|
||||
json_data = json.loads(clean_content)
|
||||
except json.JSONDecodeError:
|
||||
# Fallback to parsing raw content if markdown stripping failed
|
||||
json_data = json.loads(content)
|
||||
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
else:
|
||||
result = response_format.model_validate(json_data)
|
||||
else:
|
||||
result = content
|
||||
|
||||
# Record metrics and log slow calls
|
||||
duration = time.time() - start_time
|
||||
input_tokens = response.usage.input_tokens or 0 if response.usage else 0
|
||||
output_tokens = response.usage.output_tokens or 0 if response.usage else 0
|
||||
response_usage = _usage_from_anthropic_response(response)
|
||||
input_tokens = response_usage.input_tokens
|
||||
output_tokens = response_usage.output_tokens
|
||||
total_tokens = input_tokens + output_tokens
|
||||
cached_tokens = getattr(response.usage, "cache_read_input_tokens", 0) or 0 if response.usage else 0
|
||||
cached_tokens = response_usage.cached_tokens
|
||||
|
||||
# Record LLM metrics
|
||||
metrics = get_metrics_collector()
|
||||
@@ -295,6 +386,9 @@ class AnthropicLLM(LLMInterface):
|
||||
logger.error(f"Anthropic auth error (HTTP {e.status_code}), not retrying: {str(e)}")
|
||||
raise
|
||||
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx.
|
||||
dump_request_on_4xx(scope=scope, provider=self.provider, model=self.model, err=e, request=call_params)
|
||||
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
# Check if it's a rate limit or server error
|
||||
@@ -329,7 +423,7 @@ class AnthropicLLM(LLMInterface):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make an LLM API call with tool/function calling support.
|
||||
@@ -399,6 +493,11 @@ class AnthropicLLM(LLMInterface):
|
||||
else:
|
||||
anthropic_messages.append({"role": role, "content": content})
|
||||
|
||||
# Multi-turn tool loop: cache the stable prefix (tools + system) via
|
||||
# the system marker, and the growing conversation via an end-marker
|
||||
# that the next iteration reads back.
|
||||
_mark_last_message_for_caching(anthropic_messages)
|
||||
|
||||
call_params: dict[str, Any] = {
|
||||
"model": self.model,
|
||||
"messages": anthropic_messages,
|
||||
@@ -406,7 +505,7 @@ class AnthropicLLM(LLMInterface):
|
||||
"max_tokens": max_completion_tokens or 4096,
|
||||
}
|
||||
if system_prompt:
|
||||
call_params["system"] = system_prompt
|
||||
call_params["system"] = _cached_system_blocks(system_prompt)
|
||||
|
||||
if self._extra_body:
|
||||
call_params["extra_body"] = self._extra_body
|
||||
@@ -415,6 +514,7 @@ class AnthropicLLM(LLMInterface):
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.messages.create(**call_params)
|
||||
stash_response_usage(_usage_from_anthropic_response(response))
|
||||
|
||||
# Extract content and tool calls
|
||||
content_parts = []
|
||||
@@ -481,6 +581,8 @@ class AnthropicLLM(LLMInterface):
|
||||
except (APIConnectionError, APIStatusError) as e:
|
||||
if isinstance(e, APIStatusError) and e.status_code in (401, 403):
|
||||
raise
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx.
|
||||
dump_request_on_4xx(scope=scope, provider=self.provider, model=self.model, err=e, request=call_params)
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
await asyncio.sleep(min(initial_backoff * (2**attempt), max_backoff))
|
||||
@@ -491,6 +593,217 @@ class AnthropicLLM(LLMInterface):
|
||||
raise last_exception
|
||||
raise RuntimeError("Anthropic tool call failed")
|
||||
|
||||
# ── Message Batches API (50% token discount) ─────────────────────────────
|
||||
|
||||
_BATCH_TOOL_NAME = "structured_response"
|
||||
|
||||
async def supports_batch_api(self) -> bool:
|
||||
"""Anthropic supports batch operations via the Message Batches API."""
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _map_batch_status(processing_status: str) -> str:
|
||||
"""Map Anthropic ``processing_status`` onto the OpenAI vocabulary.
|
||||
|
||||
The engine's poll loop breaks on "completed" and hard-fails on
|
||||
"failed"/"expired"/"cancelled"; anything else keeps polling. Anthropic
|
||||
batches only end as "ended" (per-request failures surface in the
|
||||
results, mirroring OpenAI's "completed"-with-errors semantics), so
|
||||
"ended" maps to "completed" and the non-terminal states pass through.
|
||||
"""
|
||||
return "completed" if processing_status == "ended" else processing_status
|
||||
|
||||
def _translate_batch_body(self, body: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Translate one OpenAI-shaped request body into Messages API params.
|
||||
|
||||
Mirrors the conversion rules of ``call()``: system messages fold into
|
||||
the ``system`` param; ``max_completion_tokens`` becomes ``max_tokens``
|
||||
(default 4096); ``temperature`` is dropped (the sync path never sends
|
||||
it either — current Claude models reject non-default sampling params);
|
||||
an OpenAI ``response_format`` json_schema becomes a single forced
|
||||
tool_use tool when strict (native constrained decoding, issue #1002),
|
||||
else the schema is injected into the system prompt.
|
||||
|
||||
The system prompt carries the same cache_control marker as the sync
|
||||
one-shot path (its sole breakpoint): every request in a retain batch
|
||||
shares the fact-extraction system prompt, so the first item's cache
|
||||
write serves the remaining items as best-effort reads — and the
|
||||
cache-read discount stacks with the 50% batch discount.
|
||||
"""
|
||||
system_prompt: str | None = None
|
||||
messages: list[dict[str, Any]] = []
|
||||
for msg in body.get("messages", []):
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
if role == "system":
|
||||
system_prompt = (system_prompt + "\n\n" + content) if system_prompt else content
|
||||
else:
|
||||
messages.append({"role": role, "content": content})
|
||||
|
||||
params: dict[str, Any] = {
|
||||
"model": body.get("model") or self.model,
|
||||
"messages": messages,
|
||||
"max_tokens": body.get("max_completion_tokens") or 4096,
|
||||
}
|
||||
|
||||
json_schema = (body.get("response_format") or {}).get("json_schema") or {}
|
||||
schema = json_schema.get("schema")
|
||||
if schema is not None:
|
||||
if json_schema.get("strict"):
|
||||
params["tools"] = [
|
||||
{
|
||||
"name": self._BATCH_TOOL_NAME,
|
||||
"description": "Return the structured response.",
|
||||
"input_schema": schema,
|
||||
}
|
||||
]
|
||||
params["tool_choice"] = {"type": "tool", "name": self._BATCH_TOOL_NAME}
|
||||
else:
|
||||
schema_msg = "\n\nYou must respond with valid JSON matching this schema:\n" + json.dumps(
|
||||
schema, indent=2, ensure_ascii=False
|
||||
)
|
||||
system_prompt = (system_prompt + schema_msg) if system_prompt else schema_msg
|
||||
|
||||
if system_prompt:
|
||||
params["system"] = _cached_system_blocks(system_prompt)
|
||||
|
||||
# Batch params ARE the raw Messages body, so operator-configured extra
|
||||
# body params merge directly (the sync path routes them through the
|
||||
# SDK's extra_body, which does the same merge server-side).
|
||||
if self._extra_body:
|
||||
params.update(self._extra_body)
|
||||
|
||||
return params
|
||||
|
||||
def _translate_batch_message(self, message: Any) -> dict[str, Any]:
|
||||
"""Render an Anthropic Message as the OpenAI response body the engine parses.
|
||||
|
||||
The engine reads ``choices[0].message.content`` (json.loads'ing it when
|
||||
a schema was requested) and sums ``usage`` under the OpenAI key names.
|
||||
Forced-tool responses carry their JSON in the tool_use block's input,
|
||||
so that is re-serialized as the content string.
|
||||
"""
|
||||
content = ""
|
||||
tool_input = None
|
||||
for block in message.content:
|
||||
if block.type == "tool_use" and block.name == self._BATCH_TOOL_NAME:
|
||||
tool_input = block.input or {}
|
||||
elif block.type == "text":
|
||||
content += block.text
|
||||
if tool_input is not None:
|
||||
content = json.dumps(tool_input, ensure_ascii=False)
|
||||
|
||||
usage = getattr(message, "usage", None)
|
||||
input_tokens = (usage.input_tokens or 0) if usage else 0
|
||||
output_tokens = (usage.output_tokens or 0) if usage else 0
|
||||
|
||||
return {
|
||||
"choices": [
|
||||
{
|
||||
"message": {"role": "assistant", "content": content},
|
||||
"finish_reason": getattr(message, "stop_reason", None),
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": input_tokens,
|
||||
"completion_tokens": output_tokens,
|
||||
"total_tokens": input_tokens + output_tokens,
|
||||
},
|
||||
}
|
||||
|
||||
async def submit_batch(
|
||||
self,
|
||||
requests: list[dict[str, Any]],
|
||||
endpoint: str = "/v1/chat/completions",
|
||||
completion_window: str = "24h",
|
||||
) -> dict[str, Any]:
|
||||
"""Submit a batch of requests to the Message Batches API.
|
||||
|
||||
Accepts the engine's OpenAI-JSONL-shaped entries. ``endpoint`` and
|
||||
``completion_window`` belong to that shared shape and have no Anthropic
|
||||
equivalent (batches always resolve within 24 hours); both are ignored.
|
||||
"""
|
||||
batch_requests = [
|
||||
{
|
||||
"custom_id": req["custom_id"],
|
||||
"params": self._translate_batch_body(req.get("body") or {}),
|
||||
}
|
||||
for req in requests
|
||||
]
|
||||
|
||||
logger.info(f"Submitting Anthropic message batch with {len(batch_requests)} requests")
|
||||
batch = await self._client.messages.batches.create(requests=batch_requests)
|
||||
logger.info(f"Anthropic batch submitted: {batch.id}, status={batch.processing_status}")
|
||||
|
||||
return {
|
||||
"batch_id": batch.id,
|
||||
"status": self._map_batch_status(batch.processing_status),
|
||||
"created_at": batch.created_at,
|
||||
"request_count": len(batch_requests),
|
||||
}
|
||||
|
||||
async def get_batch_status(self, batch_id: str) -> dict[str, Any]:
|
||||
"""Get batch status in the shape the engine's poll loop expects."""
|
||||
batch = await self._client.messages.batches.retrieve(batch_id)
|
||||
|
||||
counts = batch.request_counts
|
||||
processing = getattr(counts, "processing", 0) or 0
|
||||
succeeded = getattr(counts, "succeeded", 0) or 0
|
||||
errored = getattr(counts, "errored", 0) or 0
|
||||
canceled = getattr(counts, "canceled", 0) or 0
|
||||
expired = getattr(counts, "expired", 0) or 0
|
||||
resolved = succeeded + errored + canceled + expired
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"batch_id": batch.id,
|
||||
"status": self._map_batch_status(batch.processing_status),
|
||||
"created_at": batch.created_at,
|
||||
"request_counts": {
|
||||
"total": processing + resolved,
|
||||
"completed": resolved,
|
||||
"failed": errored,
|
||||
},
|
||||
}
|
||||
|
||||
ended_at = getattr(batch, "ended_at", None)
|
||||
if ended_at:
|
||||
result["completed_at"] = ended_at
|
||||
|
||||
return result
|
||||
|
||||
async def retrieve_batch_results(self, batch_id: str) -> list[dict[str, Any]]:
|
||||
"""Retrieve completed batch results, translated to the OpenAI shape.
|
||||
|
||||
Succeeded entries become ``{"custom_id", "response": {"body": ...}}``;
|
||||
errored/canceled/expired entries become ``{"custom_id", "error": ...}``
|
||||
so the engine's per-result error handling applies unchanged.
|
||||
"""
|
||||
batch = await self._client.messages.batches.retrieve(batch_id)
|
||||
if batch.processing_status != "ended":
|
||||
raise ValueError(f"Batch {batch_id} is not completed yet (status: {batch.processing_status})")
|
||||
|
||||
decoder = await self._client.messages.batches.results(batch_id)
|
||||
results: list[dict[str, Any]] = []
|
||||
async for entry in decoder:
|
||||
outcome = entry.result
|
||||
if outcome.type == "succeeded":
|
||||
results.append(
|
||||
{
|
||||
"custom_id": entry.custom_id,
|
||||
"response": {"body": self._translate_batch_message(outcome.message)},
|
||||
}
|
||||
)
|
||||
else:
|
||||
error = getattr(outcome, "error", None)
|
||||
if error is not None:
|
||||
detail = f"{getattr(error, 'type', 'error')}: {getattr(error, 'message', error)}"
|
||||
else:
|
||||
detail = f"batch request {outcome.type}"
|
||||
results.append({"custom_id": entry.custom_id, "error": detail})
|
||||
|
||||
logger.info(f"Retrieved {len(results)} results for Anthropic batch {batch_id}")
|
||||
return results
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
"""Clean up resources (close Anthropic client connections)."""
|
||||
if hasattr(self, "_client") and self._client:
|
||||
|
||||
@@ -15,7 +15,8 @@ from typing import Any
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface
|
||||
from hindsight_api.engine.llm_interface import LLM_TOOL_CHOICE_AUTO, LLMInterface, LLMToolChoice, LLMToolChoiceMode
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
|
||||
@@ -48,6 +49,20 @@ def _get_isolated_claude_env() -> dict[str, str]:
|
||||
return _isolated_claude_env
|
||||
|
||||
|
||||
def _result_error_detail(message: Any) -> str:
|
||||
"""Build an actionable error string from an ``is_error`` ResultMessage.
|
||||
|
||||
The CLI can report a failure with ``is_error=True`` while ``subtype``
|
||||
still reads ``"success"``, putting the real detail in ``result`` (e.g.
|
||||
quota exhaustion: ``You've hit your weekly limit · resets ...`` with
|
||||
``api_error_status: 429``). The SDK's own fallback exception surfaces
|
||||
only the subtype, producing the misleading "Claude Code returned an
|
||||
error result: success" (issue #2702) — so prefer ``result``.
|
||||
"""
|
||||
detail = (message.result or "").strip() or message.subtype or "unknown error"
|
||||
return f"Claude Code reported an error: {detail}"
|
||||
|
||||
|
||||
class ClaudeCodeLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using Claude Code authentication.
|
||||
@@ -118,12 +133,14 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
Raises:
|
||||
RuntimeError: If the connection test fails.
|
||||
"""
|
||||
from ...config import get_config
|
||||
|
||||
try:
|
||||
test_messages = [{"role": "user", "content": "test"}]
|
||||
await self.call(
|
||||
messages=test_messages,
|
||||
max_completion_tokens=10,
|
||||
temperature=0.0,
|
||||
temperature=get_config().llm_temperature_verification,
|
||||
scope="verification",
|
||||
max_retries=0,
|
||||
)
|
||||
@@ -173,6 +190,7 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
from claude_agent_sdk import ( # type: ignore[unresolved-import]
|
||||
AssistantMessage,
|
||||
ClaudeAgentOptions,
|
||||
ResultMessage,
|
||||
TextBlock,
|
||||
query,
|
||||
)
|
||||
@@ -206,9 +224,19 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
user_content += schema_instruction
|
||||
|
||||
# Configure SDK options
|
||||
#
|
||||
# tools=[] is required here for the same reason call_with_tools() below
|
||||
# already sets it: with `tools` left at its default (None -> full
|
||||
# "claude_code" built-in preset), allowed_tools=[] alone does not stop
|
||||
# the CLI from loading the full built-in toolset and deferring into
|
||||
# ToolSearch before answering, which burns the single max_turns=1
|
||||
# budget on a tool-deferral step instead of a text response. Without
|
||||
# this, single-turn calls intermittently fail with "Reached maximum
|
||||
# number of turns (1)" even though the prompt itself needs no tools.
|
||||
options = ClaudeAgentOptions(
|
||||
system_prompt=system_prompt if system_prompt else None,
|
||||
max_turns=1, # Single-turn for API-style interactions
|
||||
tools=[], # Disable built-in tools so nothing forces a ToolSearch deferral
|
||||
allowed_tools=[], # Disable tools for standard LLM calls
|
||||
env=_get_isolated_claude_env(),
|
||||
)
|
||||
@@ -225,6 +253,21 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
for block in message.content:
|
||||
if isinstance(block, TextBlock):
|
||||
full_text += block.text
|
||||
elif isinstance(message, ResultMessage) and message.is_error:
|
||||
# Surface the CLI's actual error text (e.g. quota
|
||||
# exhaustion) instead of the SDK's subtype-based
|
||||
# fallback exception (issue #2702).
|
||||
raise RuntimeError(_result_error_detail(message))
|
||||
|
||||
# The Claude Agent SDK doesn't report exact counts; stash the same
|
||||
# char/4 estimate the success path traces so a later parse/validate
|
||||
# failure records consistent (estimated) tokens, not zero (#2387).
|
||||
stash_response_usage(
|
||||
LLMResponseUsage(
|
||||
input_tokens=sum(len(m.get("content", "")) for m in messages) // 4,
|
||||
output_tokens=len(full_text) // 4,
|
||||
)
|
||||
)
|
||||
|
||||
# Handle structured output
|
||||
if response_format is not None:
|
||||
@@ -349,7 +392,7 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make an LLM API call with tool/function calling support using Claude Agent SDK.
|
||||
@@ -367,7 +410,7 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or specific function dict.
|
||||
tool_choice: Canonical tool-selection policy.
|
||||
- "auto": Model decides whether to call tools (default)
|
||||
- "required": Model must call at least one tool
|
||||
- "none": Model must not call any tools
|
||||
@@ -380,6 +423,7 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
AssistantMessage,
|
||||
ClaudeAgentOptions,
|
||||
ClaudeSDKClient,
|
||||
ResultMessage,
|
||||
SdkMcpTool,
|
||||
TextBlock,
|
||||
ToolUseBlock,
|
||||
@@ -460,30 +504,27 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
mcp_servers_config = {"hindsight_tools": mcp_server} if sdk_tools else {}
|
||||
|
||||
# Process tool_choice
|
||||
if isinstance(tool_choice, dict) and tool_choice.get("type") == "function":
|
||||
if tool_choice.mode is LLMToolChoiceMode.NAMED:
|
||||
# Force a specific tool: filter allowed_tools to only that tool and add instruction
|
||||
forced_name = tool_choice.get("function", {}).get("name")
|
||||
if forced_name:
|
||||
# Filter to only the forced tool (with MCP prefix)
|
||||
forced_tool_mcp_name = f"mcp__hindsight_tools__{forced_name}"
|
||||
if forced_tool_mcp_name in allowed_tool_names:
|
||||
allowed_tool_names = [forced_tool_mcp_name]
|
||||
# Add strong instruction to system prompt
|
||||
force_instruction = (
|
||||
f"\n\nIMPORTANT: You MUST call the '{forced_name}' tool. Do not respond with text only."
|
||||
)
|
||||
system_prompt += force_instruction
|
||||
logger.debug(f"Claude Code: Forcing tool call to '{forced_name}'")
|
||||
else:
|
||||
logger.warning(f"Claude Code: Forced tool '{forced_name}' not found in available tools")
|
||||
elif tool_choice == "required":
|
||||
forced_name = tool_choice.selected_function_name
|
||||
forced_tool_mcp_name = f"mcp__hindsight_tools__{forced_name}"
|
||||
if forced_tool_mcp_name in allowed_tool_names:
|
||||
allowed_tool_names = [forced_tool_mcp_name]
|
||||
force_instruction = (
|
||||
f"\n\nIMPORTANT: You MUST call the '{forced_name}' tool. Do not respond with text only."
|
||||
)
|
||||
system_prompt += force_instruction
|
||||
logger.debug(f"Claude Code: Forcing tool call to '{forced_name}'")
|
||||
else:
|
||||
logger.warning(f"Claude Code: Forced tool '{forced_name}' not found in available tools")
|
||||
elif tool_choice.mode is LLMToolChoiceMode.REQUIRED:
|
||||
# Must call at least one tool
|
||||
tool_instruction = (
|
||||
"\n\nIMPORTANT: You MUST call at least one of the available tools. Do not respond with text only."
|
||||
)
|
||||
system_prompt += tool_instruction
|
||||
logger.debug("Claude Code: Tool call required")
|
||||
elif tool_choice == "none":
|
||||
elif tool_choice.mode is LLMToolChoiceMode.NONE:
|
||||
# No tools should be called - disable all tools
|
||||
allowed_tool_names = []
|
||||
mcp_servers_config = {}
|
||||
@@ -519,6 +560,9 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
|
||||
# Receive response
|
||||
async for message in client.receive_response():
|
||||
if isinstance(message, ResultMessage) and message.is_error:
|
||||
# Surface the CLI's actual error text (issue #2702).
|
||||
raise RuntimeError(_result_error_detail(message))
|
||||
if isinstance(message, AssistantMessage):
|
||||
for block in message.content:
|
||||
if isinstance(block, TextBlock):
|
||||
|
||||
@@ -19,6 +19,7 @@ from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
@@ -31,6 +32,11 @@ from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
try:
|
||||
import fcntl
|
||||
except ImportError: # pragma: no cover - Windows
|
||||
fcntl = None # type: ignore[assignment]
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -58,6 +64,63 @@ _CODEX_TOKEN_REFRESH_SKEW_SECONDS = 60
|
||||
_CODEX_TERMINAL_REFRESH_ERROR_CODES = frozenset(
|
||||
{"refresh_token_expired", "refresh_token_reused", "refresh_token_invalidated"}
|
||||
)
|
||||
_CODEX_AUTH_LOCK_TIMEOUT_SECONDS = 20.0
|
||||
_CODEX_AUTH_LOCKS_GUARD = threading.Lock()
|
||||
_CODEX_AUTH_LOCKS: dict[Path, threading.Lock] = {}
|
||||
|
||||
|
||||
def default_codex_auth_file() -> Path:
|
||||
"""Return the path to Codex's ``auth.json``.
|
||||
|
||||
Honors the ``CODEX_HOME`` environment variable — the same variable the
|
||||
canonical ``@openai/codex`` CLI uses to relocate its config/credentials
|
||||
directory — and falls back to ``~/.codex`` when it is unset or empty.
|
||||
|
||||
Resolved lazily on each call (rather than cached at import time) so that
|
||||
the environment is read at the point of use.
|
||||
"""
|
||||
codex_home = os.environ.get("CODEX_HOME")
|
||||
if codex_home:
|
||||
return Path(codex_home) / "auth.json"
|
||||
return Path.home() / ".codex" / "auth.json"
|
||||
|
||||
|
||||
def _path_scoped_lock(auth_file: Path) -> threading.Lock:
|
||||
key = auth_file.expanduser().resolve(strict=False)
|
||||
with _CODEX_AUTH_LOCKS_GUARD:
|
||||
lock = _CODEX_AUTH_LOCKS.get(key)
|
||||
if lock is None:
|
||||
lock = threading.Lock()
|
||||
_CODEX_AUTH_LOCKS[key] = lock
|
||||
return lock
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _codex_auth_lock(auth_file: Path, timeout_seconds: float = _CODEX_AUTH_LOCK_TIMEOUT_SECONDS):
|
||||
"""Cross-process advisory lock for one Codex auth store."""
|
||||
with _path_scoped_lock(auth_file):
|
||||
if fcntl is None: # pragma: no cover - Windows
|
||||
logger.debug("fcntl unavailable; Codex refresh proceeds without a cross-process lock.")
|
||||
yield
|
||||
return
|
||||
|
||||
lock_path = auth_file.with_suffix(".lock")
|
||||
lock_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(lock_path, "a+") as lock_file:
|
||||
deadline = time.monotonic() + max(1.0, timeout_seconds)
|
||||
while True:
|
||||
try:
|
||||
fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||
break
|
||||
except (BlockingIOError, OSError):
|
||||
if time.monotonic() >= deadline:
|
||||
raise TimeoutError("Timed out waiting for the Codex auth store lock") from None
|
||||
time.sleep(0.05)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
with contextlib.suppress(OSError):
|
||||
fcntl.flock(lock_file.fileno(), fcntl.LOCK_UN)
|
||||
|
||||
|
||||
class CodexRefreshExpiredError(RuntimeError):
|
||||
@@ -86,7 +149,7 @@ class CodexAuthManager:
|
||||
The OAuth refresh token. May be ``None`` when the auth file omits it;
|
||||
the provider still works as a one-shot loader in that case.
|
||||
auth_file:
|
||||
Path to ``~/.codex/auth.json``. Used for re-reading the refresh token
|
||||
Path to the Codex ``auth.json``. Used for re-reading the refresh token
|
||||
on demand and for atomic persistence of rotated credentials.
|
||||
"""
|
||||
|
||||
@@ -115,7 +178,8 @@ class CodexAuthManager:
|
||||
Parameters
|
||||
----------
|
||||
auth_file:
|
||||
Defaults to ``~/.codex/auth.json``.
|
||||
Defaults to ``$CODEX_HOME/auth.json`` (or ``~/.codex/auth.json``
|
||||
when ``CODEX_HOME`` is unset).
|
||||
|
||||
Raises
|
||||
------
|
||||
@@ -126,7 +190,7 @@ class CodexAuthManager:
|
||||
``auth_mode``.
|
||||
"""
|
||||
if auth_file is None:
|
||||
auth_file = Path.home() / ".codex" / "auth.json"
|
||||
auth_file = default_codex_auth_file()
|
||||
|
||||
if not auth_file.exists():
|
||||
raise FileNotFoundError(f"Codex auth file not found: {auth_file}. Run 'codex auth login' to authenticate.")
|
||||
@@ -175,6 +239,34 @@ class CodexAuthManager:
|
||||
return None
|
||||
return data.get("tokens", {}).get("refresh_token")
|
||||
|
||||
@staticmethod
|
||||
def _load_tokens_from_file(auth_file: Path) -> dict[str, Any] | None:
|
||||
try:
|
||||
with open(auth_file) as f:
|
||||
data = json.load(f)
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return None
|
||||
tokens = data.get("tokens")
|
||||
return tokens if isinstance(tokens, dict) else None
|
||||
|
||||
def _adopt_tokens(self, tokens: dict[str, Any]) -> bool:
|
||||
"""Adopt a newer on-disk Codex token set if present."""
|
||||
access_token = tokens.get("access_token")
|
||||
refresh_token = tokens.get("refresh_token")
|
||||
account_id = tokens.get("account_id")
|
||||
|
||||
changed = False
|
||||
if isinstance(access_token, str) and access_token and access_token != self.access_token:
|
||||
self.access_token = access_token
|
||||
changed = True
|
||||
if isinstance(refresh_token, str) and refresh_token and refresh_token != self.refresh_token:
|
||||
self.refresh_token = refresh_token
|
||||
changed = True
|
||||
if isinstance(account_id, str) and account_id and account_id != self.account_id:
|
||||
self.account_id = account_id
|
||||
changed = True
|
||||
return changed
|
||||
|
||||
@staticmethod
|
||||
def _decode_jwt_exp_unixtime(token: str) -> int | None:
|
||||
"""Return the JWT ``exp`` claim as a unix timestamp, or None on parse failure.
|
||||
@@ -208,6 +300,11 @@ class CodexAuthManager:
|
||||
return False
|
||||
return exp <= int(time.time()) + skew_seconds
|
||||
|
||||
def _token_is_fresh_with_known_expiry(self, skew_seconds: int = _CODEX_TOKEN_REFRESH_SKEW_SECONDS) -> bool:
|
||||
"""True only when the cached token has a known expiry outside the skew window."""
|
||||
exp = self._decode_jwt_exp_unixtime(self.access_token)
|
||||
return exp is not None and exp > int(time.time()) + skew_seconds
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Persistence
|
||||
# ------------------------------------------------------------------
|
||||
@@ -322,78 +419,93 @@ class CodexAuthManager:
|
||||
if not self._token_is_stale():
|
||||
return
|
||||
|
||||
if not self.refresh_token:
|
||||
raise RuntimeError(
|
||||
"Codex access_token is expired but no refresh_token is available. "
|
||||
"Run 'codex auth login' to re-authenticate."
|
||||
)
|
||||
with _codex_auth_lock(self._auth_file):
|
||||
disk_tokens = self._load_tokens_from_file(self._auth_file)
|
||||
if disk_tokens and self._adopt_tokens(disk_tokens):
|
||||
if force or self._token_is_fresh_with_known_expiry():
|
||||
return
|
||||
|
||||
log_reason = f" ({reason})" if reason else ""
|
||||
logger.info(f"Refreshing Codex OAuth access_token{log_reason}")
|
||||
|
||||
request_body = {
|
||||
"client_id": _CODEX_CLIENT_ID,
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": self.refresh_token,
|
||||
}
|
||||
try:
|
||||
response = self._http_client.post(
|
||||
_CODEX_REFRESH_TOKEN_URL,
|
||||
json=request_body,
|
||||
headers={"Content-Type": "application/json"},
|
||||
timeout=30.0,
|
||||
)
|
||||
except httpx.RequestError as e:
|
||||
raise RuntimeError(f"Codex OAuth refresh network error: {type(e).__name__}") from e
|
||||
|
||||
if response.status_code == 401:
|
||||
error_code = self._extract_oauth_error_code(response)
|
||||
if error_code in _CODEX_TERMINAL_REFRESH_ERROR_CODES:
|
||||
raise CodexRefreshExpiredError(
|
||||
f"Codex refresh_token is permanently invalid (error.code={error_code}). "
|
||||
if not self.refresh_token:
|
||||
raise RuntimeError(
|
||||
"Codex access_token is expired but no refresh_token is available. "
|
||||
"Run 'codex auth login' to re-authenticate."
|
||||
)
|
||||
raise CodexRefreshExpiredError(
|
||||
f"Codex OAuth refresh returned 401 with unrecognized error code "
|
||||
f"({error_code or 'none'}). Run 'codex auth login' to re-authenticate."
|
||||
)
|
||||
|
||||
if response.status_code >= 400:
|
||||
raise RuntimeError(f"Codex OAuth refresh failed with HTTP {response.status_code}")
|
||||
log_reason = f" ({reason})" if reason else ""
|
||||
logger.info(f"Refreshing Codex OAuth access_token{log_reason}")
|
||||
|
||||
try:
|
||||
body = response.json()
|
||||
except json.JSONDecodeError as e:
|
||||
raise RuntimeError(f"Codex OAuth refresh returned non-JSON body: {e}") from e
|
||||
request_access_token = self.access_token
|
||||
request_refresh_token = self.refresh_token
|
||||
request_body = {
|
||||
"client_id": _CODEX_CLIENT_ID,
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": request_refresh_token,
|
||||
}
|
||||
try:
|
||||
response = self._http_client.post(
|
||||
_CODEX_REFRESH_TOKEN_URL,
|
||||
json=request_body,
|
||||
headers={"Content-Type": "application/json"},
|
||||
timeout=30.0,
|
||||
)
|
||||
except httpx.RequestError as e:
|
||||
raise RuntimeError(f"Codex OAuth refresh network error: {type(e).__name__}") from e
|
||||
|
||||
new_access = body.get("access_token")
|
||||
if not new_access:
|
||||
raise RuntimeError("Codex OAuth refresh returned no access_token")
|
||||
if response.status_code == 401:
|
||||
error_code = self._extract_oauth_error_code(response)
|
||||
disk_tokens = self._load_tokens_from_file(self._auth_file)
|
||||
if disk_tokens and (
|
||||
disk_tokens.get("access_token") != request_access_token
|
||||
or disk_tokens.get("refresh_token") != request_refresh_token
|
||||
):
|
||||
self._adopt_tokens(disk_tokens)
|
||||
return
|
||||
if error_code in _CODEX_TERMINAL_REFRESH_ERROR_CODES:
|
||||
raise CodexRefreshExpiredError(
|
||||
f"Codex refresh_token is permanently invalid (error.code={error_code}). "
|
||||
"Run 'codex auth login' to re-authenticate."
|
||||
)
|
||||
raise CodexRefreshExpiredError(
|
||||
f"Codex OAuth refresh returned 401 with unrecognized error code "
|
||||
f"({error_code or 'none'}). Run 'codex auth login' to re-authenticate."
|
||||
)
|
||||
|
||||
new_refresh = body.get("refresh_token") or self.refresh_token
|
||||
new_id_token = body.get("id_token")
|
||||
if response.status_code >= 400:
|
||||
raise RuntimeError(f"Codex OAuth refresh failed with HTTP {response.status_code}")
|
||||
|
||||
# Update in-memory state first so waiters see fresh credentials
|
||||
# immediately, even if disk write fails.
|
||||
self.access_token = new_access
|
||||
self.refresh_token = new_refresh
|
||||
try:
|
||||
body = response.json()
|
||||
except json.JSONDecodeError as e:
|
||||
raise RuntimeError(f"Codex OAuth refresh returned non-JSON body: {e}") from e
|
||||
|
||||
persisted: dict[str, Any] = {
|
||||
"access_token": new_access,
|
||||
"refresh_token": new_refresh,
|
||||
}
|
||||
if new_id_token:
|
||||
persisted["id_token"] = new_id_token
|
||||
new_access = body.get("access_token")
|
||||
if not new_access:
|
||||
raise RuntimeError("Codex OAuth refresh returned no access_token")
|
||||
|
||||
try:
|
||||
self._persist_auth_atomic(persisted)
|
||||
except OSError as e:
|
||||
logger.warning(
|
||||
f"Codex OAuth refresh succeeded but persisting auth.json failed: {type(e).__name__}. "
|
||||
"In-memory credentials are up to date; on-disk file is stale."
|
||||
)
|
||||
new_refresh = body.get("refresh_token") or self.refresh_token
|
||||
new_id_token = body.get("id_token")
|
||||
|
||||
logger.info("Codex OAuth access_token refreshed successfully")
|
||||
# Update in-memory state first so waiters see fresh credentials
|
||||
# immediately, even if disk write fails.
|
||||
self.access_token = new_access
|
||||
self.refresh_token = new_refresh
|
||||
|
||||
persisted: dict[str, Any] = {
|
||||
"access_token": new_access,
|
||||
"refresh_token": new_refresh,
|
||||
}
|
||||
if new_id_token:
|
||||
persisted["id_token"] = new_id_token
|
||||
|
||||
try:
|
||||
self._persist_auth_atomic(persisted)
|
||||
except OSError as e:
|
||||
logger.warning(
|
||||
f"Codex OAuth refresh succeeded but persisting auth.json failed: {type(e).__name__}. "
|
||||
"In-memory credentials are up to date; on-disk file is stale."
|
||||
)
|
||||
|
||||
logger.info("Codex OAuth access_token refreshed successfully")
|
||||
|
||||
def ensure_fresh_token(self) -> None:
|
||||
"""Proactively refresh the access_token if it is near or past expiry.
|
||||
|
||||
@@ -2,8 +2,9 @@
|
||||
OpenAI Codex LLM provider using ChatGPT Plus/Pro OAuth authentication.
|
||||
|
||||
This provider enables using ChatGPT Plus/Pro subscriptions for API calls
|
||||
without separate OpenAI Platform API credits. It uses OAuth tokens from
|
||||
~/.codex/auth.json and communicates with the ChatGPT backend API.
|
||||
without separate OpenAI Platform API credits. It uses OAuth tokens from the
|
||||
Codex ``auth.json`` (``$CODEX_HOME/auth.json``, or ``~/.codex/auth.json`` when
|
||||
``CODEX_HOME`` is unset) and communicates with the ChatGPT backend API.
|
||||
|
||||
Tokens are refreshed automatically: the provider decodes the access_token
|
||||
JWT's ``exp`` claim and proactively refreshes via
|
||||
@@ -24,8 +25,11 @@ from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface
|
||||
from hindsight_api.engine.llm_interface import LLM_TOOL_CHOICE_AUTO, LLMInterface, LLMToolChoice, LLMToolChoiceMode
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.providers.llm_debug import dump_request_on_4xx
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.engine.structured_output import strict_json_schema
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
|
||||
from .codex_auth import (
|
||||
@@ -35,6 +39,7 @@ from .codex_auth import (
|
||||
_CODEX_TOKEN_REFRESH_SKEW_SECONDS,
|
||||
CodexAuthManager,
|
||||
CodexRefreshExpiredError,
|
||||
default_codex_auth_file,
|
||||
)
|
||||
|
||||
# Re-export for backward compatibility (tests import from this module).
|
||||
@@ -50,19 +55,73 @@ __all__ = [
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Newer Codex models are gated on the first-party client identity; the previous
|
||||
# browser-shaped User-Agent returned "Model not found" for Luna (#2643).
|
||||
# Use a neutral version because Hindsight must not claim a specific Codex release.
|
||||
_CODEX_ORIGINATOR = "codex_cli_rs"
|
||||
_CODEX_USER_AGENT = "codex_cli_rs/0.0.0 (Hindsight)"
|
||||
|
||||
# Name of the single forced function tool used to carry structured output when
|
||||
# strict_schema is on. The Codex backend speaks the OpenAI Responses API, so a
|
||||
# forced function call gives us constrained decoding straight into the response
|
||||
# schema — no prompt-injected schema, no raw json.loads on free-form model text,
|
||||
# no invalid-\escape retry storm (issue #2504, same class as #1002 / #2339).
|
||||
_STRUCTURED_TOOL_NAME = "structured_response"
|
||||
|
||||
# Valid JSON string escape characters (the char that may follow a backslash).
|
||||
_VALID_JSON_ESCAPE_CHARS = set('"\\/bfnrtu')
|
||||
|
||||
|
||||
def _repair_invalid_json_escapes(text: str) -> str:
|
||||
"""Best-effort repair of invalid ``\\escape`` sequences in a JSON string.
|
||||
|
||||
Escape-heavy content (code, serial/CLI commands, Windows paths, regexes)
|
||||
makes weaker models emit backslashes that aren't valid JSON escapes (e.g.
|
||||
``\\d``, ``\\s``, ``C:\\Users``), so ``json.loads`` fails deterministically
|
||||
and every retry re-fails the same way (issue #2504). This doubles any
|
||||
backslash that isn't part of a valid escape so the payload parses. It is a
|
||||
lenient fallback only — the strict_schema forced-tool path is the real fix.
|
||||
"""
|
||||
result: list[str] = []
|
||||
i = 0
|
||||
n = len(text)
|
||||
while i < n:
|
||||
ch = text[i]
|
||||
if ch == "\\" and i + 1 < n:
|
||||
nxt = text[i + 1]
|
||||
if nxt in _VALID_JSON_ESCAPE_CHARS:
|
||||
# Preserve the valid escape (both chars) verbatim.
|
||||
result.append(ch)
|
||||
result.append(nxt)
|
||||
i += 2
|
||||
continue
|
||||
# Invalid escape: escape the lone backslash so JSON parses.
|
||||
result.append("\\\\")
|
||||
i += 1
|
||||
continue
|
||||
if ch == "\\" and i + 1 == n:
|
||||
# Trailing lone backslash — escape it.
|
||||
result.append("\\\\")
|
||||
i += 1
|
||||
continue
|
||||
result.append(ch)
|
||||
i += 1
|
||||
return "".join(result)
|
||||
|
||||
|
||||
class CodexLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using OpenAI Codex OAuth authentication.
|
||||
|
||||
Authenticates using ChatGPT Plus/Pro credentials stored in ~/.codex/auth.json
|
||||
and makes API calls to chatgpt.com/backend-api/codex/responses.
|
||||
Authenticates using ChatGPT Plus/Pro credentials stored in the Codex
|
||||
``auth.json`` (honoring ``CODEX_HOME``, default ``~/.codex``) and makes API
|
||||
calls to chatgpt.com/backend-api/codex/responses.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str, # Will be ignored, reads from ~/.codex/auth.json
|
||||
api_key: str, # Will be ignored, reads from the Codex auth.json (CODEX_HOME or ~/.codex)
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
@@ -81,12 +140,14 @@ class CodexLLM(LLMInterface):
|
||||
refresh_token = self._load_codex_refresh_token()
|
||||
logger.info(f"Loaded Codex OAuth credentials for account: {account_id}")
|
||||
except Exception as e:
|
||||
auth_file = default_codex_auth_file()
|
||||
raise RuntimeError(
|
||||
f"Failed to load Codex OAuth credentials from ~/.codex/auth.json: {e}\n\n"
|
||||
f"Failed to load Codex OAuth credentials from {auth_file}: {e}\n\n"
|
||||
"To set up Codex authentication:\n"
|
||||
"1. Install Codex CLI: npm install -g @openai/codex\n"
|
||||
"2. Login: codex auth login\n"
|
||||
"3. Verify: ls ~/.codex/auth.json\n\n"
|
||||
f"3. Verify: ls {auth_file}\n\n"
|
||||
"(Set CODEX_HOME to use a credentials directory other than ~/.codex.)\n\n"
|
||||
"Or use a different provider (openai, anthropic, gemini) with API keys."
|
||||
) from e
|
||||
|
||||
@@ -94,7 +155,7 @@ class CodexLLM(LLMInterface):
|
||||
access_token=access_token,
|
||||
account_id=account_id,
|
||||
refresh_token=refresh_token,
|
||||
auth_file=Path.home() / ".codex" / "auth.json",
|
||||
auth_file=default_codex_auth_file(),
|
||||
)
|
||||
|
||||
# Use ChatGPT backend API endpoint. Codex auth is tied to
|
||||
@@ -111,8 +172,8 @@ class CodexLLM(LLMInterface):
|
||||
if self.model.startswith("openai/"):
|
||||
self.model = self.model[len("openai/") :]
|
||||
|
||||
# Map reasoning effort to Codex reasoning summary format
|
||||
# Codex supports: "auto", "concise", "detailed"
|
||||
# Reasoning summary controls presentation separately from the backend's
|
||||
# reasoning effort, which is sent unchanged in each request payload.
|
||||
self.reasoning_summary = self._map_reasoning_effort(reasoning_effort)
|
||||
|
||||
# HTTP client for SSE streaming
|
||||
@@ -134,6 +195,18 @@ class CodexLLM(LLMInterface):
|
||||
def account_id(self) -> str:
|
||||
return self._auth_manager.account_id
|
||||
|
||||
def _build_request_headers(self) -> httpx.Headers:
|
||||
return httpx.Headers(
|
||||
{
|
||||
"Authorization": f"Bearer {self.access_token}",
|
||||
"Content-Type": "application/json",
|
||||
"OpenAI-Account-ID": self.account_id,
|
||||
"User-Agent": _CODEX_USER_AGENT,
|
||||
"Origin": "https://chatgpt.com",
|
||||
"originator": _CODEX_ORIGINATOR,
|
||||
}
|
||||
)
|
||||
|
||||
@property
|
||||
def refresh_token(self) -> str | None:
|
||||
return self._auth_manager.refresh_token
|
||||
@@ -156,7 +229,7 @@ class CodexLLM(LLMInterface):
|
||||
|
||||
def _load_codex_auth(self) -> tuple[str, str]:
|
||||
"""
|
||||
Load OAuth credentials from ~/.codex/auth.json.
|
||||
Load OAuth credentials from the Codex ``auth.json`` (CODEX_HOME or ~/.codex).
|
||||
|
||||
Returns:
|
||||
Tuple of (access_token, account_id).
|
||||
@@ -165,7 +238,7 @@ class CodexLLM(LLMInterface):
|
||||
FileNotFoundError: If auth file doesn't exist.
|
||||
ValueError: If auth file is invalid.
|
||||
"""
|
||||
auth_file = Path.home() / ".codex" / "auth.json"
|
||||
auth_file = default_codex_auth_file()
|
||||
|
||||
if not auth_file.exists():
|
||||
raise FileNotFoundError(
|
||||
@@ -197,9 +270,7 @@ class CodexLLM(LLMInterface):
|
||||
pre- and post-``__init__`` because it does not depend on
|
||||
``_auth_manager`` being constructed yet.
|
||||
"""
|
||||
auth_file = (
|
||||
self._auth_manager._auth_file if hasattr(self, "_auth_manager") else Path.home() / ".codex" / "auth.json"
|
||||
)
|
||||
auth_file = self._auth_manager._auth_file if hasattr(self, "_auth_manager") else default_codex_auth_file()
|
||||
return CodexAuthManager.load_refresh_token_from_file(auth_file)
|
||||
|
||||
@staticmethod
|
||||
@@ -272,32 +343,6 @@ class CodexLLM(LLMInterface):
|
||||
}
|
||||
return mapping.get(effort.lower(), "auto")
|
||||
|
||||
def _normalize_tool_choice(self, tool_choice: str | dict[str, Any]) -> str | dict[str, Any]:
|
||||
"""Normalize forced function tool choice for the Codex Responses API.
|
||||
|
||||
Older agent paths may still pass OpenAI chat-completions style named
|
||||
tool choice payloads such as:
|
||||
|
||||
{"type": "function", "function": {"name": "recall"}}
|
||||
|
||||
Codex Responses expects the named function at the top level instead:
|
||||
|
||||
{"type": "function", "name": "recall"}
|
||||
"""
|
||||
if not isinstance(tool_choice, dict):
|
||||
return tool_choice
|
||||
if str(tool_choice.get("type") or "").strip() != "function":
|
||||
return tool_choice
|
||||
function_payload = tool_choice.get("function")
|
||||
if isinstance(function_payload, dict):
|
||||
function_name = str(function_payload.get("name") or "").strip()
|
||||
if function_name:
|
||||
return {"type": "function", "name": function_name}
|
||||
function_name = str(tool_choice.get("name") or "").strip()
|
||||
if function_name:
|
||||
return {"type": "function", "name": function_name}
|
||||
return tool_choice
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""Verify Codex connection by making a simple test call."""
|
||||
try:
|
||||
@@ -332,7 +377,18 @@ class CodexLLM(LLMInterface):
|
||||
strict_schema: bool = False,
|
||||
return_usage: bool = False,
|
||||
) -> Any:
|
||||
"""Make API call to Codex backend with SSE streaming."""
|
||||
"""Make API call to Codex backend with SSE streaming.
|
||||
|
||||
Args:
|
||||
strict_schema: Route structured output through a single forced
|
||||
function tool (constrained decoding) instead of prompt-injecting
|
||||
the schema and parsing free-form text. The Codex backend speaks
|
||||
the OpenAI Responses API, so the forced function call emits the
|
||||
response schema directly as tool arguments — eliminating the
|
||||
invalid-``\\escape`` retry storm (issue #2504). When False, falls
|
||||
back to schema-in-prompt + JSON parse, now hardened with a lenient
|
||||
invalid-escape repair before giving up.
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
# Proactively refresh the OAuth access_token if it's near expiry.
|
||||
@@ -357,11 +413,22 @@ class CodexLLM(LLMInterface):
|
||||
else:
|
||||
user_messages.append(msg)
|
||||
|
||||
# Add JSON schema instruction if response_format is provided
|
||||
# Structured output: prefer a single forced function tool (constrained
|
||||
# decoding) over text-injecting the schema and parsing the reply. The
|
||||
# forced tool guarantees schema-shaped JSON in the tool arguments,
|
||||
# eliminating the invalid-\escape retry storm (issue #2504). When
|
||||
# strict_schema is off we keep the schema-in-prompt + json.loads
|
||||
# fallback (now hardened with a lenient escape repair) for callers that
|
||||
# can't force tools.
|
||||
schema = None
|
||||
use_forced_tool = False
|
||||
if response_format is not None and hasattr(response_format, "model_json_schema"):
|
||||
schema = response_format.model_json_schema()
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
system_instruction += schema_msg
|
||||
schema = strict_json_schema(response_format) if strict_schema else response_format.model_json_schema()
|
||||
if strict_schema:
|
||||
use_forced_tool = True
|
||||
else:
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
system_instruction += schema_msg
|
||||
|
||||
# gpt-5.2-codex only supports "detailed" reasoning summary
|
||||
reasoning_summary = "detailed" if "5.2" in self.model else self.reasoning_summary
|
||||
@@ -381,20 +448,28 @@ class CodexLLM(LLMInterface):
|
||||
"tools": [],
|
||||
"tool_choice": "auto",
|
||||
"parallel_tool_calls": True,
|
||||
"reasoning": {"summary": reasoning_summary},
|
||||
"reasoning": {"effort": self.reasoning_effort, "summary": reasoning_summary},
|
||||
"store": False, # Codex uses stateless mode
|
||||
"stream": True, # SSE streaming
|
||||
"include": ["reasoning.encrypted_content"],
|
||||
"prompt_cache_key": str(uuid.uuid4()),
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.access_token}",
|
||||
"Content-Type": "application/json",
|
||||
"OpenAI-Account-ID": self.account_id,
|
||||
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7)",
|
||||
"Origin": "https://chatgpt.com",
|
||||
}
|
||||
if use_forced_tool and schema is not None:
|
||||
# Single function tool whose parameters ARE the response schema;
|
||||
# force it via tool_choice so the backend does constrained decoding.
|
||||
payload["tools"] = [
|
||||
{
|
||||
"type": "function",
|
||||
"name": _STRUCTURED_TOOL_NAME,
|
||||
"description": "Return the structured response.",
|
||||
"parameters": schema,
|
||||
}
|
||||
]
|
||||
payload["tool_choice"] = {"type": "function", "name": _STRUCTURED_TOOL_NAME}
|
||||
payload["parallel_tool_calls"] = False
|
||||
|
||||
headers = self._build_request_headers()
|
||||
|
||||
url = f"{self.base_url}/codex/responses"
|
||||
|
||||
@@ -408,11 +483,49 @@ class CodexLLM(LLMInterface):
|
||||
response = await self._client.post(url, json=payload, headers=headers, timeout=120.0)
|
||||
response.raise_for_status()
|
||||
|
||||
# Parse SSE stream
|
||||
content = await self._parse_sse_stream(response)
|
||||
# Forced-tool path: read structured output from the function-call
|
||||
# arguments (already a JSON string in a dedicated channel) rather
|
||||
# than from free-form assistant text.
|
||||
if use_forced_tool:
|
||||
text_content, tool_calls = await self._parse_sse_tool_stream(response)
|
||||
content = text_content or ""
|
||||
else:
|
||||
tool_calls = []
|
||||
content = await self._parse_sse_stream(response)
|
||||
|
||||
# Codex SSE carries no usage block; stash the same char/4 estimate
|
||||
# the success path traces so a later parse/validate failure records
|
||||
# consistent (estimated) token counts rather than zero (#2387).
|
||||
stash_response_usage(
|
||||
LLMResponseUsage(
|
||||
input_tokens=sum(len(m.get("content", "")) for m in messages) // 4,
|
||||
output_tokens=len(content) // 4,
|
||||
)
|
||||
)
|
||||
|
||||
# Handle structured output
|
||||
if response_format is not None:
|
||||
if use_forced_tool:
|
||||
tool_input = None
|
||||
for tc in tool_calls:
|
||||
if tc.name == _STRUCTURED_TOOL_NAME:
|
||||
tool_input = tc.arguments if isinstance(tc.arguments, dict) else None
|
||||
break
|
||||
if tool_input is None:
|
||||
# Model ignored the forced tool (rare — e.g. a gateway that
|
||||
# drops tool_choice). Retry so we don't hard-fail.
|
||||
logger.warning(
|
||||
f"Codex forced structured tool missing from response "
|
||||
f"(attempt {attempt + 1}/{max_retries + 1})"
|
||||
)
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
attempt += 1
|
||||
continue
|
||||
raise RuntimeError("Codex did not return the forced structured_response tool call")
|
||||
content = json.dumps(tool_input)
|
||||
result = tool_input if skip_validation else response_format.model_validate(tool_input)
|
||||
elif response_format is not None:
|
||||
# Models may wrap JSON in markdown
|
||||
clean_content = content
|
||||
if "```json" in content:
|
||||
@@ -423,13 +536,20 @@ class CodexLLM(LLMInterface):
|
||||
try:
|
||||
json_data = json.loads(clean_content)
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f"Codex JSON parse error (attempt {attempt + 1}/{max_retries + 1}): {e}")
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
attempt += 1
|
||||
continue
|
||||
raise
|
||||
# Escape-heavy content deterministically re-fails every
|
||||
# retry (issue #2504). Try a lenient invalid-escape repair
|
||||
# before burning a retry / re-raising.
|
||||
try:
|
||||
json_data = json.loads(_repair_invalid_json_escapes(clean_content))
|
||||
logger.info("Codex JSON parsed after repairing invalid escape sequences")
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(f"Codex JSON parse error (attempt {attempt + 1}/{max_retries + 1}): {e}")
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
attempt += 1
|
||||
continue
|
||||
raise
|
||||
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
@@ -528,6 +648,9 @@ class CodexLLM(LLMInterface):
|
||||
"Run 'codex auth login' to re-authenticate."
|
||||
) from e
|
||||
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx.
|
||||
dump_request_on_4xx(scope=scope, provider=self.provider, model=self.model, err=e, request=payload)
|
||||
|
||||
# Log the actual error message from the API
|
||||
error_detail = e.response.text[:500] if hasattr(e.response, "text") else str(e)
|
||||
|
||||
@@ -623,7 +746,7 @@ class CodexLLM(LLMInterface):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make API call with tool calling support.
|
||||
@@ -640,7 +763,7 @@ class CodexLLM(LLMInterface):
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or a specific function.
|
||||
tool_choice: Canonical tool-selection policy.
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
@@ -702,22 +825,20 @@ class CodexLLM(LLMInterface):
|
||||
"instructions": system_instruction,
|
||||
"input": user_messages,
|
||||
"tools": codex_tools,
|
||||
"tool_choice": self._normalize_tool_choice(tool_choice),
|
||||
"tool_choice": (
|
||||
{"type": "function", "name": tool_choice.selected_function_name}
|
||||
if tool_choice.mode is LLMToolChoiceMode.NAMED
|
||||
else tool_choice.mode.value
|
||||
),
|
||||
"parallel_tool_calls": True,
|
||||
"reasoning": {"summary": reasoning_summary},
|
||||
"reasoning": {"effort": self.reasoning_effort, "summary": reasoning_summary},
|
||||
"store": False,
|
||||
"stream": True,
|
||||
"include": ["reasoning.encrypted_content"],
|
||||
"prompt_cache_key": str(uuid.uuid4()),
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.access_token}",
|
||||
"Content-Type": "application/json",
|
||||
"OpenAI-Account-ID": self.account_id,
|
||||
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7)",
|
||||
"Origin": "https://chatgpt.com",
|
||||
}
|
||||
headers = self._build_request_headers()
|
||||
|
||||
url = f"{self.base_url}/codex/responses"
|
||||
|
||||
@@ -814,6 +935,8 @@ class CodexLLM(LLMInterface):
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx.
|
||||
dump_request_on_4xx(scope=scope, provider=self.provider, model=self.model, err=e, request=payload)
|
||||
logger.error(f"Codex tool call error: {e}")
|
||||
raise
|
||||
|
||||
@@ -858,8 +981,13 @@ class CodexLLM(LLMInterface):
|
||||
try:
|
||||
arguments = json.loads(arguments_str)
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(f"Failed to parse tool arguments: {arguments_str}")
|
||||
arguments = {}
|
||||
# Escape-heavy content can emit invalid \escape
|
||||
# sequences (issue #2504); repair before giving up.
|
||||
try:
|
||||
arguments = json.loads(_repair_invalid_json_escapes(arguments_str))
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(f"Failed to parse tool arguments: {arguments_str}")
|
||||
arguments = {}
|
||||
|
||||
tool_calls.append(
|
||||
LLMToolCall(
|
||||
|
||||
@@ -56,6 +56,14 @@ _DEFAULT_REFRESH_MARGIN_SECONDS = 5 * 60
|
||||
# to None and callers proceed uncached, rather than stalling the whole batch.
|
||||
_DEFAULT_CREATE_TIMEOUT_SECONDS = 30.0
|
||||
|
||||
# TTL for the per-step reflect caches created by ``create_incremental``. These
|
||||
# live only for the duration of one reflect (seconds), so the TTL is just a
|
||||
# storage backstop in case the explicit ``delete_session`` at reflect end is
|
||||
# missed (crash / event-loop teardown). Short so orphaned caches age out fast —
|
||||
# storage is billed per token-hour, so a 5-minute cap keeps the cost of a leaked
|
||||
# cache negligible.
|
||||
_DEFAULT_INCREMENTAL_TTL_SECONDS = 5 * 60
|
||||
|
||||
|
||||
@dataclass
|
||||
class _CacheEntry:
|
||||
@@ -92,6 +100,10 @@ class GeminiCacheManager:
|
||||
self._create_timeout_seconds = create_timeout_seconds
|
||||
self._entries: dict[str, _CacheEntry] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
# session_id -> CachedContent names created via ``create_incremental``.
|
||||
# A reflect creates a fresh rolling cache per step under one session id;
|
||||
# ``delete_session`` tears them all down when the reflect finishes.
|
||||
self._sessions: dict[str, list[str]] = {}
|
||||
|
||||
@staticmethod
|
||||
def fingerprint(
|
||||
@@ -229,18 +241,99 @@ class GeminiCacheManager:
|
||||
if entry.name == name:
|
||||
self._entries.pop(key, None)
|
||||
|
||||
async def create_incremental(
|
||||
self,
|
||||
*,
|
||||
session_id: str,
|
||||
model: str,
|
||||
system_instruction: str,
|
||||
contents: list[Any],
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
) -> str | None:
|
||||
"""Create a fresh CachedContent holding ``system + tools + contents`` and
|
||||
track it under ``session_id`` for later teardown.
|
||||
|
||||
Unlike ``get_or_create``, this does NOT deduplicate by fingerprint: each
|
||||
step of a reflect grows the conversation prefix, so every call is a
|
||||
distinct, single-use cache. The reflect loop creates one per step (each
|
||||
covering the previous step's full input) and reuses it for exactly the
|
||||
next model turn, then supersedes it. All caches for the session are
|
||||
deleted by ``delete_session`` when the reflect ends; the short TTL is
|
||||
only a backstop.
|
||||
|
||||
Returns the cache resource name, or ``None`` when caching is disabled,
|
||||
the prefix is below the model minimum, or the create otherwise fails —
|
||||
callers MUST fall back to an uncached call in that case.
|
||||
"""
|
||||
try:
|
||||
name = await self._create_cache(
|
||||
model=model,
|
||||
system_instruction=system_instruction,
|
||||
tools=tools,
|
||||
contents=contents,
|
||||
ttl_seconds=_DEFAULT_INCREMENTAL_TTL_SECONDS,
|
||||
)
|
||||
except _CacheNotEligible as e:
|
||||
logger.debug(
|
||||
"GeminiCacheManager: incremental prefix not eligible (model=%s, reason=%s) — caller falls back",
|
||||
model,
|
||||
e,
|
||||
)
|
||||
return None
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"GeminiCacheManager: failed to create incremental cache (model=%s); caller falls back",
|
||||
model,
|
||||
)
|
||||
return None
|
||||
if name is not None:
|
||||
self._sessions.setdefault(session_id, []).append(name)
|
||||
return name
|
||||
|
||||
async def delete(self, name: str) -> None:
|
||||
"""Best-effort server-side delete of a single CachedContent.
|
||||
|
||||
Swallows all errors: a failed delete just means the cache ages out on
|
||||
its TTL. Also drops any matching in-process entry.
|
||||
"""
|
||||
self.invalidate(name)
|
||||
try:
|
||||
await self._client.aio.caches.delete(name=name)
|
||||
except Exception:
|
||||
logger.debug("GeminiCacheManager: delete of cache %s failed (will age out on TTL)", name, exc_info=True)
|
||||
|
||||
async def delete_session(self, session_id: str) -> None:
|
||||
"""Delete every CachedContent created for ``session_id`` (reflect teardown).
|
||||
|
||||
Deletes concurrently and best-effort — a reflect must never fail because
|
||||
a cache couldn't be torn down; the short TTL is the backstop.
|
||||
"""
|
||||
names = self._sessions.pop(session_id, [])
|
||||
if not names:
|
||||
return
|
||||
await asyncio.gather(*(self.delete(n) for n in names), return_exceptions=True)
|
||||
|
||||
async def _create_cache(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
system_instruction: str,
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
contents: list[Any] | None = None,
|
||||
ttl_seconds: int | None = None,
|
||||
) -> str | None:
|
||||
"""Wrap ``client.aio.caches.create`` with the config we want.
|
||||
|
||||
The SDK surface differs slightly across google-genai versions;
|
||||
this implementation targets the >=1.0.0 line where caches live
|
||||
under ``client.aio.caches``.
|
||||
|
||||
``contents`` (already-converted ``genai_types.Content`` turns) is
|
||||
appended after the system_instruction/tools so the cache can hold a
|
||||
growing multi-turn conversation prefix, not just the static prefix —
|
||||
this is what the step-by-step reflect cache relies on. ``ttl_seconds``
|
||||
overrides the manager default (used to give per-step reflect caches a
|
||||
short backstop TTL).
|
||||
"""
|
||||
# Lazy import so this module doesn't require the SDK at import time.
|
||||
from google.genai import types as genai_types
|
||||
@@ -254,8 +347,10 @@ class GeminiCacheManager:
|
||||
# still part of the fingerprint so a schema change keys a fresh cache.
|
||||
config_kwargs: dict[str, Any] = {
|
||||
"system_instruction": system_instruction,
|
||||
"ttl": f"{self._ttl_seconds}s",
|
||||
"ttl": f"{ttl_seconds if ttl_seconds is not None else self._ttl_seconds}s",
|
||||
}
|
||||
if contents:
|
||||
config_kwargs["contents"] = contents
|
||||
if tools:
|
||||
# OpenAI-style {"function": {...}} entries must be converted to
|
||||
# Gemini's Tool/FunctionDeclaration shape before caching.
|
||||
|
||||
@@ -13,14 +13,17 @@ import json
|
||||
import logging
|
||||
import time
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from google import genai
|
||||
from google.genai import errors as genai_errors
|
||||
from google.genai import types as genai_types
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface
|
||||
from hindsight_api.engine.llm_interface import LLM_TOOL_CHOICE_AUTO, LLMInterface, LLMToolChoice, LLMToolChoiceMode
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.llm_wrapper import parse_llm_json
|
||||
from hindsight_api.engine.providers.llm_debug import dump_request_on_4xx
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
from hindsight_api.worker.stage import set_stage
|
||||
@@ -50,6 +53,111 @@ def _to_int(value: Any) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
def _usage_from_gemini_response(response: Any) -> LLMResponseUsage:
|
||||
"""Extract prompt/candidate/cached token counts from a Gemini usage_metadata block."""
|
||||
usage = getattr(response, "usage_metadata", None)
|
||||
if not usage:
|
||||
return LLMResponseUsage()
|
||||
return LLMResponseUsage(
|
||||
input_tokens=usage.prompt_token_count or 0,
|
||||
output_tokens=usage.candidates_token_count or 0,
|
||||
cached_tokens=getattr(usage, "cached_content_token_count", 0) or 0,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _GeminiConversation:
|
||||
"""A message list converted to Gemini's request shape."""
|
||||
|
||||
system_instruction: str | None
|
||||
contents: list["genai_types.Content"]
|
||||
|
||||
|
||||
def _convert_messages_to_gemini(msg_list: list[dict[str, Any]]) -> _GeminiConversation:
|
||||
"""Convert OpenAI-style messages to a Gemini (system_instruction, contents) pair.
|
||||
|
||||
Shared by ``call_with_tools`` (request body) and the incremental cache
|
||||
builder so a cached prefix and the live request serialise turns identically —
|
||||
any drift would fingerprint differently and defeat the cache. Consecutive
|
||||
``role="tool"`` messages are grouped into a single ``user`` Content with
|
||||
multiple FunctionResponse parts, matching Gemini's multi-turn requirement.
|
||||
"""
|
||||
system_instruction: str | None = None
|
||||
gemini_contents: list[genai_types.Content] = []
|
||||
pending_tool_names_by_call_id: dict[str, str] = {}
|
||||
i = 0
|
||||
while i < len(msg_list):
|
||||
msg = msg_list[i]
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role != "tool" and pending_tool_names_by_call_id:
|
||||
missing_ids = ", ".join(sorted(pending_tool_names_by_call_id))
|
||||
raise ValueError(f"Gemini assistant tool calls require results before the next message: {missing_ids}")
|
||||
|
||||
if role == "system":
|
||||
system_instruction = (system_instruction + "\n\n" + content) if system_instruction else content
|
||||
i += 1
|
||||
elif role == "tool":
|
||||
parts = []
|
||||
while i < len(msg_list) and msg_list[i].get("role") == "tool":
|
||||
tool_msg = msg_list[i]
|
||||
tool_content = tool_msg.get("content", "")
|
||||
tool_call_id = tool_msg["tool_call_id"]
|
||||
tool_name = pending_tool_names_by_call_id.pop(tool_call_id, None)
|
||||
if tool_name is None:
|
||||
raise ValueError(f"Gemini tool result references unknown tool_call_id {tool_call_id!r}")
|
||||
parts.append(
|
||||
genai_types.Part(
|
||||
function_response=genai_types.FunctionResponse(
|
||||
name=tool_name,
|
||||
response={"result": tool_content},
|
||||
)
|
||||
)
|
||||
)
|
||||
i += 1
|
||||
if pending_tool_names_by_call_id:
|
||||
missing_ids = ", ".join(sorted(pending_tool_names_by_call_id))
|
||||
raise ValueError(f"Gemini assistant tool calls are missing results: {missing_ids}")
|
||||
gemini_contents.append(genai_types.Content(role="user", parts=parts))
|
||||
elif role == "assistant":
|
||||
tool_calls_in_msg = msg.get("tool_calls", [])
|
||||
if tool_calls_in_msg:
|
||||
parts = []
|
||||
if content:
|
||||
parts.append(genai_types.Part(text=content))
|
||||
for tc in tool_calls_in_msg:
|
||||
tool_call_id = tc["id"]
|
||||
fn = tc["function"]
|
||||
fn_name = fn["name"]
|
||||
if tool_call_id in pending_tool_names_by_call_id:
|
||||
raise ValueError(
|
||||
f"Gemini assistant tool call id {tool_call_id!r} must be unique within its turn"
|
||||
)
|
||||
pending_tool_names_by_call_id[tool_call_id] = fn_name
|
||||
fn_args_str = fn.get("arguments", "{}")
|
||||
fn_args = parse_llm_json(fn_args_str)
|
||||
thought_signature = tc.get("thought_signature")
|
||||
fc_kwargs: dict[str, Any] = {"name": fn_name, "args": fn_args}
|
||||
part_kwargs: dict[str, Any] = {"function_call": genai_types.FunctionCall(**fc_kwargs)}
|
||||
if thought_signature:
|
||||
part_kwargs["thought_signature"] = base64.b64decode(thought_signature)
|
||||
parts.append(genai_types.Part(**part_kwargs))
|
||||
gemini_contents.append(genai_types.Content(role="model", parts=parts))
|
||||
else:
|
||||
gemini_contents.append(genai_types.Content(role="model", parts=[genai_types.Part(text=content)]))
|
||||
i += 1
|
||||
else:
|
||||
gemini_contents.append(genai_types.Content(role="user", parts=[genai_types.Part(text=content)]))
|
||||
i += 1
|
||||
|
||||
if pending_tool_names_by_call_id:
|
||||
missing_ids = ", ".join(sorted(pending_tool_names_by_call_id))
|
||||
raise ValueError(f"Gemini assistant tool calls are missing results: {missing_ids}")
|
||||
|
||||
return _GeminiConversation(system_instruction=system_instruction, contents=gemini_contents)
|
||||
|
||||
|
||||
class GeminiLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider for Google Gemini and Vertex AI.
|
||||
@@ -76,6 +184,7 @@ class GeminiLLM(LLMInterface):
|
||||
|
||||
# Safety settings: None means use Gemini's defaults
|
||||
self._safety_settings: list | None = kwargs.get("gemini_safety_settings")
|
||||
self._service_tier: str | None = kwargs.get("gemini_service_tier")
|
||||
|
||||
# User-configured extra params merged into the GenerateContentConfig of
|
||||
# every call. Gemini's request body nests generation params, so we expose
|
||||
@@ -106,6 +215,16 @@ class GeminiLLM(LLMInterface):
|
||||
self._client = genai.Client(api_key=self.api_key)
|
||||
logger.info(f"Gemini API: model={self.model}")
|
||||
|
||||
def _apply_service_tier(self, config_kwargs: dict[str, Any]) -> None:
|
||||
if not self._service_tier:
|
||||
return
|
||||
|
||||
http_options = dict(config_kwargs.get("http_options") or {})
|
||||
extra_body = dict(http_options.get("extra_body") or {})
|
||||
extra_body.setdefault("service_tier", self._service_tier)
|
||||
http_options["extra_body"] = extra_body
|
||||
config_kwargs["http_options"] = http_options
|
||||
|
||||
def _init_vertexai(self, **kwargs: Any) -> None:
|
||||
"""Initialize Vertex AI client with project, region, and credentials."""
|
||||
# Extract Vertex AI config from kwargs
|
||||
@@ -247,16 +366,13 @@ class GeminiLLM(LLMInterface):
|
||||
else:
|
||||
gemini_contents.append(genai_types.Content(role="user", parts=[genai_types.Part(text=content)]))
|
||||
|
||||
# Add the JSON schema as a textual hint in the system_instruction (matching
|
||||
# the normal uncached path). Structured output is still enforced via
|
||||
# response_schema regardless; this is just guidance text.
|
||||
if response_format is not None and hasattr(response_format, "model_json_schema"):
|
||||
def _system_instruction_with_schema() -> str:
|
||||
schema = response_format.model_json_schema()
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
if system_instruction:
|
||||
system_instruction += schema_msg
|
||||
else:
|
||||
system_instruction = schema_msg
|
||||
schema_msg = (
|
||||
f"\n\nYou must respond with valid JSON matching this schema:\n"
|
||||
f"{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
)
|
||||
return (system_instruction + schema_msg) if system_instruction else schema_msg
|
||||
|
||||
# Apply safety settings: context var (per-request bank override) takes precedence over instance default
|
||||
effective_safety_settings = _safety_settings_ctx.get()
|
||||
@@ -273,11 +389,18 @@ class GeminiLLM(LLMInterface):
|
||||
def _build_generation_config(use_cache: bool) -> "genai_types.GenerateContentConfig | None":
|
||||
# Seed with user-configured extra params; explicit settings below win.
|
||||
config_kwargs: dict[str, Any] = dict(self._extra_body)
|
||||
self._apply_service_tier(config_kwargs)
|
||||
if use_cache:
|
||||
config_kwargs["cached_content"] = cached_prefix
|
||||
elif (
|
||||
use_schema_prompt_fallback
|
||||
and response_format is not None
|
||||
and hasattr(response_format, "model_json_schema")
|
||||
):
|
||||
config_kwargs["system_instruction"] = _system_instruction_with_schema()
|
||||
elif system_instruction:
|
||||
config_kwargs["system_instruction"] = system_instruction
|
||||
if response_format is not None:
|
||||
if response_format is not None and not use_schema_prompt_fallback:
|
||||
config_kwargs["response_mime_type"] = "application/json"
|
||||
config_kwargs["response_schema"] = response_format
|
||||
if temperature is not None:
|
||||
@@ -295,6 +418,7 @@ class GeminiLLM(LLMInterface):
|
||||
return genai_types.GenerateContentConfig(**config_kwargs) if config_kwargs else None
|
||||
|
||||
cache_active = using_cache
|
||||
use_schema_prompt_fallback = False
|
||||
generation_config = _build_generation_config(cache_active)
|
||||
|
||||
last_exception = None
|
||||
@@ -311,6 +435,9 @@ class GeminiLLM(LLMInterface):
|
||||
),
|
||||
timeout=90.0, # Safety net for network hangs; valid slow responses are <90s
|
||||
)
|
||||
# Stash usage before parse/validate, which may raise locally
|
||||
# even though the provider charged for these tokens (#2387).
|
||||
stash_response_usage(_usage_from_gemini_response(response))
|
||||
|
||||
content = response.text
|
||||
|
||||
@@ -412,12 +539,26 @@ class GeminiLLM(LLMInterface):
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=input_tokens + output_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
return result, token_usage
|
||||
return result
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
last_exception = e
|
||||
if (
|
||||
attempt < max_retries
|
||||
and response_format is not None
|
||||
and hasattr(response_format, "model_json_schema")
|
||||
and not cache_active
|
||||
and not use_schema_prompt_fallback
|
||||
):
|
||||
logger.warning("Gemini returned invalid JSON, retrying with prompt-side schema guidance...")
|
||||
cache_active = False
|
||||
use_schema_prompt_fallback = True
|
||||
generation_config = _build_generation_config(cache_active)
|
||||
await asyncio.sleep(min(initial_backoff * (2**attempt), max_backoff))
|
||||
continue
|
||||
if attempt < max_retries:
|
||||
logger.warning("Gemini returned invalid JSON, retrying...")
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
@@ -433,6 +574,17 @@ class GeminiLLM(LLMInterface):
|
||||
logger.error(f"Gemini auth error (HTTP {e.code}), not retrying: {str(e)}")
|
||||
raise
|
||||
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx, captured
|
||||
# before the cache-drop retry below rebuilds the config so we see what failed.
|
||||
dump_request_on_4xx(
|
||||
scope=scope,
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
err=e,
|
||||
request=generation_config,
|
||||
messages=gemini_contents,
|
||||
)
|
||||
|
||||
# Cached-request safety net: a stale/invalid/expired CachedContent
|
||||
# (or an incompatibility like cache + tool_config) surfaces as a 400.
|
||||
# Retrying the same cached request can't recover, so on the first
|
||||
@@ -479,8 +631,9 @@ class GeminiLLM(LLMInterface):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
cached_prefix: str | None = None,
|
||||
cached_prefix_message_count: int = 0,
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make a Gemini/VertexAI API call with tool/function calling support.
|
||||
@@ -494,15 +647,22 @@ class GeminiLLM(LLMInterface):
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools (Gemini uses "auto" only).
|
||||
tool_choice: Canonical tool-selection policy.
|
||||
cached_prefix: Optional CachedContent resource name (from
|
||||
``GeminiCacheManager.get_or_create`` with ``tools=...``). When
|
||||
set, the system_instruction and tool definitions are assumed
|
||||
``GeminiCacheManager.get_or_create`` or ``create_incremental``).
|
||||
When set, the system_instruction and tool definitions are assumed
|
||||
to live in the cache; this call will skip resending them and
|
||||
the cached prefix is billed at the cached-input rate. The
|
||||
``tools`` argument is still required (the caller may pass
|
||||
an empty list when the cache holds them) so existing call
|
||||
sites don't break.
|
||||
cached_prefix_message_count: Number of leading ``messages`` already
|
||||
baked into ``cached_prefix`` (the step-by-step reflect cache holds
|
||||
a growing conversation prefix, not just system+tools). Only the
|
||||
messages AFTER this index are sent as request contents — the rest
|
||||
come from the cache and bill at the cached rate. 0 means the cache
|
||||
holds only the static prefix (system+tools), so the full
|
||||
conversation is still sent (legacy behaviour).
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
@@ -510,86 +670,45 @@ class GeminiLLM(LLMInterface):
|
||||
start_time = time.time()
|
||||
using_cache = cached_prefix is not None
|
||||
|
||||
# Convert tools to Gemini format. When the cache is in use, the
|
||||
# tool definitions are baked into the CachedContent at create time
|
||||
# and the SDK rejects re-sending them alongside ``cached_content``.
|
||||
# Convert tools to Gemini format. While the cache is in use the tool
|
||||
# definitions live in the CachedContent and the SDK rejects re-sending
|
||||
# them alongside ``cached_content`` (see ``_build_tools_config``), but we
|
||||
# still build them unconditionally so the cached-call-failed fallback —
|
||||
# which drops the cache and re-sends prefix + tools inline — has real
|
||||
# tools to send rather than an empty list.
|
||||
gemini_tools = []
|
||||
if not using_cache:
|
||||
for tool in tools:
|
||||
func = tool.get("function", {})
|
||||
gemini_tools.append(
|
||||
genai_types.Tool(
|
||||
function_declarations=[
|
||||
genai_types.FunctionDeclaration(
|
||||
name=func.get("name", ""),
|
||||
description=func.get("description", ""),
|
||||
parameters=func.get("parameters"),
|
||||
)
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
# Convert messages
|
||||
system_instruction = None
|
||||
gemini_contents = []
|
||||
msg_list = list(messages)
|
||||
i = 0
|
||||
while i < len(msg_list):
|
||||
msg = msg_list[i]
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == "system":
|
||||
# Always capture system_instruction. _build_tools_config omits it
|
||||
# (and tools) from the request while the cache carries the prefix,
|
||||
# but it must be available so the cached-call-failed safety net can
|
||||
# re-send the prefix + tools inline.
|
||||
system_instruction = (system_instruction + "\n\n" + content) if system_instruction else content
|
||||
i += 1
|
||||
elif role == "tool":
|
||||
# Gemini requires ALL tool responses for a given model turn to be grouped
|
||||
# into a single Content with multiple FunctionResponse parts.
|
||||
# Consecutive role="tool" messages correspond to one model turn's tool calls.
|
||||
parts = []
|
||||
while i < len(msg_list) and msg_list[i].get("role") == "tool":
|
||||
tool_msg = msg_list[i]
|
||||
tool_content = tool_msg.get("content", "")
|
||||
parts.append(
|
||||
genai_types.Part(
|
||||
function_response=genai_types.FunctionResponse(
|
||||
name=tool_msg.get("name", ""),
|
||||
response={"result": tool_content},
|
||||
)
|
||||
for tool in tools:
|
||||
func = tool.get("function", {})
|
||||
gemini_tools.append(
|
||||
genai_types.Tool(
|
||||
function_declarations=[
|
||||
genai_types.FunctionDeclaration(
|
||||
name=func.get("name", ""),
|
||||
description=func.get("description", ""),
|
||||
parameters=func.get("parameters"),
|
||||
)
|
||||
)
|
||||
i += 1
|
||||
gemini_contents.append(genai_types.Content(role="user", parts=parts))
|
||||
elif role == "assistant":
|
||||
tool_calls_in_msg = msg.get("tool_calls", [])
|
||||
if tool_calls_in_msg:
|
||||
# Convert OpenAI-style tool_calls to Gemini function_call parts
|
||||
# This is required for proper multi-turn conversation history
|
||||
parts = []
|
||||
if content:
|
||||
parts.append(genai_types.Part(text=content))
|
||||
for tc in tool_calls_in_msg:
|
||||
fn = tc.get("function", {})
|
||||
fn_name = fn.get("name", "")
|
||||
fn_args_str = fn.get("arguments", "{}")
|
||||
fn_args = parse_llm_json(fn_args_str)
|
||||
thought_signature = tc.get("thought_signature")
|
||||
fc_kwargs: dict[str, Any] = {"name": fn_name, "args": fn_args}
|
||||
part_kwargs: dict[str, Any] = {"function_call": genai_types.FunctionCall(**fc_kwargs)}
|
||||
if thought_signature:
|
||||
part_kwargs["thought_signature"] = base64.b64decode(thought_signature)
|
||||
parts.append(genai_types.Part(**part_kwargs))
|
||||
gemini_contents.append(genai_types.Content(role="model", parts=parts))
|
||||
else:
|
||||
gemini_contents.append(genai_types.Content(role="model", parts=[genai_types.Part(text=content)]))
|
||||
i += 1
|
||||
else:
|
||||
gemini_contents.append(genai_types.Content(role="user", parts=[genai_types.Part(text=content)]))
|
||||
i += 1
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
# Convert messages. ``system_instruction`` and the FULL contents are always
|
||||
# computed: _build_tools_config omits system/tools from the request while
|
||||
# the cache carries the prefix, but the cached-call-failed safety net must
|
||||
# be able to re-send the whole prefix + tools inline.
|
||||
converted = _convert_messages_to_gemini(list(messages))
|
||||
system_instruction = converted.system_instruction
|
||||
full_contents = converted.contents
|
||||
|
||||
# Step-by-step reflect cache: when the cache already holds the first
|
||||
# ``cached_prefix_message_count`` messages, send ONLY the newer turns as
|
||||
# request contents — the cached prefix supplies the rest at the cached
|
||||
# rate. The split is always at a whole-turn boundary (the reflect loop
|
||||
# advances the cache one completed turn at a time), so slicing the raw
|
||||
# messages before conversion never splits a grouped tool turn.
|
||||
if using_cache and cached_prefix_message_count > 0:
|
||||
delta_contents = _convert_messages_to_gemini(list(messages)[cached_prefix_message_count:]).contents
|
||||
else:
|
||||
delta_contents = full_contents
|
||||
|
||||
# Apply safety settings: context var (per-request bank override) takes precedence over instance default
|
||||
effective_safety_settings = _safety_settings_ctx.get()
|
||||
@@ -604,6 +723,7 @@ class GeminiLLM(LLMInterface):
|
||||
def _build_tools_config(use_cache: bool) -> "genai_types.GenerateContentConfig":
|
||||
# Seed with user-configured extra params; explicit settings below win.
|
||||
config_kwargs: dict[str, Any] = dict(self._extra_body)
|
||||
self._apply_service_tier(config_kwargs)
|
||||
if use_cache:
|
||||
config_kwargs["cached_content"] = cached_prefix
|
||||
else:
|
||||
@@ -618,22 +738,20 @@ class GeminiLLM(LLMInterface):
|
||||
config_kwargs["max_output_tokens"] = max_completion_tokens
|
||||
|
||||
# Map OpenAI-style tool_choice to Gemini FunctionCallingConfig
|
||||
if tool_choice == "required":
|
||||
if tool_choice.mode is LLMToolChoiceMode.REQUIRED:
|
||||
config_kwargs["tool_config"] = genai_types.ToolConfig(
|
||||
function_calling_config=genai_types.FunctionCallingConfig(
|
||||
mode="ANY",
|
||||
)
|
||||
)
|
||||
elif isinstance(tool_choice, dict) and tool_choice.get("type") == "function":
|
||||
fn_name = tool_choice.get("function", {}).get("name")
|
||||
if fn_name:
|
||||
config_kwargs["tool_config"] = genai_types.ToolConfig(
|
||||
function_calling_config=genai_types.FunctionCallingConfig(
|
||||
mode="ANY",
|
||||
allowed_function_names=[fn_name],
|
||||
)
|
||||
elif tool_choice.mode is LLMToolChoiceMode.NAMED:
|
||||
config_kwargs["tool_config"] = genai_types.ToolConfig(
|
||||
function_calling_config=genai_types.FunctionCallingConfig(
|
||||
mode="ANY",
|
||||
allowed_function_names=[tool_choice.selected_function_name],
|
||||
)
|
||||
elif tool_choice == "none":
|
||||
)
|
||||
elif tool_choice.mode is LLMToolChoiceMode.NONE:
|
||||
config_kwargs["tool_config"] = genai_types.ToolConfig(
|
||||
function_calling_config=genai_types.FunctionCallingConfig(mode="NONE")
|
||||
)
|
||||
@@ -654,14 +772,19 @@ class GeminiLLM(LLMInterface):
|
||||
if attempt > 0:
|
||||
set_stage(f"llm.gemini.tools.attempt={attempt + 1}/{max_retries + 1}")
|
||||
try:
|
||||
# With the cache active, send only the un-cached tail (delta);
|
||||
# on the uncached fallback path send the full conversation so the
|
||||
# re-inlined system+tools prefix has its whole context.
|
||||
active_contents = delta_contents if cache_active else full_contents
|
||||
response = await asyncio.wait_for(
|
||||
self._client.aio.models.generate_content(
|
||||
model=self.model,
|
||||
contents=gemini_contents,
|
||||
contents=active_contents,
|
||||
config=config,
|
||||
),
|
||||
timeout=90.0, # Safety net for network hangs; valid slow responses are <90s
|
||||
)
|
||||
stash_response_usage(_usage_from_gemini_response(response))
|
||||
|
||||
# Extract content and tool calls
|
||||
content = None
|
||||
@@ -749,6 +872,8 @@ class GeminiLLM(LLMInterface):
|
||||
finish_reason=finish_reason,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cached_tokens=cached_input_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
|
||||
except genai_errors.APIError as e:
|
||||
@@ -757,6 +882,17 @@ class GeminiLLM(LLMInterface):
|
||||
logger.error(f"Gemini auth error (HTTP {e.code}), not retrying: {str(e)}")
|
||||
raise
|
||||
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx, captured
|
||||
# before the cache-drop retry below rebuilds the config so we see what failed.
|
||||
dump_request_on_4xx(
|
||||
scope=scope,
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
err=e,
|
||||
request=config,
|
||||
messages=active_contents,
|
||||
)
|
||||
|
||||
# Cached-request safety net (see ``call``): a stale/invalid cache or
|
||||
# a cache+tool_config conflict surfaces as a 400. Drop the cache,
|
||||
# invalidate it for later operations, and retry THIS call inline
|
||||
@@ -833,6 +969,56 @@ class GeminiLLM(LLMInterface):
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
# ── Step-by-step incremental prompt caching (reflect tool loop) ──────────
|
||||
|
||||
def supports_incremental_prompt_cache(self) -> bool:
|
||||
"""True when explicit caching is on — the reflect loop can then roll a
|
||||
per-step CachedContent that grows with the conversation."""
|
||||
return self._prompt_cache_enabled
|
||||
|
||||
def _ensure_cache_manager(self) -> Any:
|
||||
if self._cache_manager is None:
|
||||
from hindsight_api.engine.providers.gemini_cache import GeminiCacheManager
|
||||
|
||||
self._cache_manager = GeminiCacheManager(self._client)
|
||||
return self._cache_manager
|
||||
|
||||
async def create_incremental_cache(
|
||||
self,
|
||||
*,
|
||||
session_id: str,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
) -> str | None:
|
||||
"""Cache ``system + tools + messages`` as a conversation prefix and return
|
||||
its resource name (or ``None`` — caller falls back to an uncached call).
|
||||
|
||||
The reflect loop calls this once per step with the growing message list so
|
||||
each step's cache entirely contains the previous step's input; the next
|
||||
model turn then references it and re-sends only its own delta. Caches are
|
||||
tracked under ``session_id`` and torn down by ``delete_cache_session``.
|
||||
"""
|
||||
if not self._prompt_cache_enabled or self._client is None:
|
||||
return None
|
||||
converted = _convert_messages_to_gemini(list(messages))
|
||||
return await self._ensure_cache_manager().create_incremental(
|
||||
session_id=session_id,
|
||||
model=self.model,
|
||||
system_instruction=converted.system_instruction or "",
|
||||
contents=converted.contents,
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
async def delete_cached_prefix(self, name: str) -> None:
|
||||
"""Best-effort delete of a single CachedContent (superseded reflect step)."""
|
||||
if self._cache_manager is not None:
|
||||
await self._cache_manager.delete(name)
|
||||
|
||||
async def delete_cache_session(self, session_id: str) -> None:
|
||||
"""Tear down every CachedContent created for a reflect session."""
|
||||
if self._cache_manager is not None:
|
||||
await self._cache_manager.delete_session(session_id)
|
||||
|
||||
# ── Batch API (Gemini API only — not Vertex AI) ─────────────────────────
|
||||
#
|
||||
# Google's Gemini Batch API gives a flat 50% discount on input + output
|
||||
@@ -980,7 +1166,7 @@ class GeminiLLM(LLMInterface):
|
||||
Mirrors the synchronous ``call`` path: system messages become
|
||||
``systemInstruction``; a ``response_format`` json_schema forces JSON
|
||||
output (``responseMimeType``), appends the schema as a textual hint, and
|
||||
grammar-enforces via ``responseJsonSchema`` when ``strict`` is set.
|
||||
grammar-enforces via ``responseJsonSchema`` whenever a schema is present.
|
||||
"""
|
||||
system_texts: list[str] = []
|
||||
contents: list[dict[str, Any]] = []
|
||||
@@ -1009,8 +1195,13 @@ class GeminiLLM(LLMInterface):
|
||||
system_texts.append(
|
||||
"You must respond with valid JSON matching this schema:\n" + json.dumps(schema, ensure_ascii=False)
|
||||
)
|
||||
if json_schema.get("strict"):
|
||||
generation_config["responseJsonSchema"] = schema
|
||||
# #2699: Gemini always grammar-enforces structured output via its native
|
||||
# response_schema (``strict`` is an OpenAI concept, meaningless here). Set
|
||||
# the native schema whenever one is present so the batch path mirrors the
|
||||
# interactive path; otherwise batch requests at default config
|
||||
# (HINDSIGHT_API_LLM_STRICT_SCHEMA=False) get only a textual hint and
|
||||
# intermittently emit malformed JSON, losing every fact in the chunk.
|
||||
generation_config["responseJsonSchema"] = schema
|
||||
|
||||
request: dict[str, Any] = {"contents": contents}
|
||||
if system_texts:
|
||||
|
||||
@@ -15,17 +15,47 @@ is handled automatically by LiteLLM.
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface, OutputTooLongError
|
||||
from litellm.exceptions import Timeout as LiteLLMTimeout
|
||||
|
||||
from hindsight_api.config import DEFAULT_LLM_TIMEOUT, ENV_LLM_TIMEOUT
|
||||
from hindsight_api.engine.llm_interface import (
|
||||
LLM_TOOL_CHOICE_AUTO,
|
||||
LLMInterface,
|
||||
LLMToolChoice,
|
||||
LLMToolChoiceMode,
|
||||
OutputTooLongError,
|
||||
)
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.llm_wrapper import parse_llm_json
|
||||
from hindsight_api.engine.providers.llm_debug import dump_request_on_4xx
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.engine.structured_output import strict_json_schema
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
from hindsight_api.worker.stage import set_stage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _usage_from_litellm_response(response: Any) -> LLMResponseUsage:
|
||||
"""Extract prompt/completion/cached token counts from a LiteLLM (OpenAI-shaped) usage block."""
|
||||
usage = getattr(response, "usage", None)
|
||||
if not usage:
|
||||
return LLMResponseUsage()
|
||||
cached_tokens = 0
|
||||
details = getattr(usage, "prompt_tokens_details", None)
|
||||
if details:
|
||||
cached_tokens = getattr(details, "cached_tokens", 0) or 0
|
||||
return LLMResponseUsage(
|
||||
input_tokens=getattr(usage, "prompt_tokens", 0) or 0,
|
||||
output_tokens=getattr(usage, "completion_tokens", 0) or 0,
|
||||
cached_tokens=cached_tokens,
|
||||
)
|
||||
|
||||
|
||||
class LiteLLMLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using the LiteLLM SDK for universal model support.
|
||||
@@ -47,13 +77,16 @@ class LiteLLMLLM(LLMInterface):
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
timeout: float = 300.0,
|
||||
timeout: float | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
bedrock_service_tier: str | None = None,
|
||||
default_headers: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(provider, api_key, base_url, model, reasoning_effort, **kwargs)
|
||||
self.timeout = timeout
|
||||
# ``None`` falls back to HINDSIGHT_API_LLM_TIMEOUT, then DEFAULT_LLM_TIMEOUT — never None,
|
||||
# so the hard ``asyncio.wait_for`` backstop in ``call`` is always bounded.
|
||||
self.timeout = timeout if timeout is not None else float(os.getenv(ENV_LLM_TIMEOUT, str(DEFAULT_LLM_TIMEOUT)))
|
||||
self._litellm: Any = None
|
||||
# User-configured extra params merged as top-level kwargs into every
|
||||
# completion call so LiteLLM normalizes them per-provider (e.g. maps
|
||||
@@ -61,6 +94,13 @@ class LiteLLMLLM(LLMInterface):
|
||||
# drops any the target model rejects (litellm.drop_params=True below).
|
||||
# Sourced from llm_extra_body (env: HINDSIGHT_API_LLM_EXTRA_BODY).
|
||||
self._extra_body: dict[str, Any] = extra_body or {}
|
||||
# Operator-configured default headers forwarded to litellm.acompletion as
|
||||
# ``extra_headers`` (used by deployments routing through proxies / request-
|
||||
# tracing middleware). Mirrors the Anthropic provider's default_headers
|
||||
# wiring. Sourced from llm_default_headers (env: HINDSIGHT_API_LLM_DEFAULT_HEADERS).
|
||||
# Copied so a caller-owned dict can't be mutated through us, and a fresh
|
||||
# copy is handed to each call below to avoid cross-request contamination.
|
||||
self._default_headers: dict[str, Any] = dict(default_headers or {})
|
||||
self.bedrock_service_tier = bedrock_service_tier
|
||||
|
||||
try:
|
||||
@@ -77,12 +117,14 @@ class LiteLLMLLM(LLMInterface):
|
||||
raise RuntimeError("LiteLLM SDK not installed. Run: uv add litellm or pip install litellm") from e
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
from ...config import get_config
|
||||
|
||||
try:
|
||||
test_messages = [{"role": "user", "content": "test"}]
|
||||
await self.call(
|
||||
messages=test_messages,
|
||||
max_completion_tokens=50,
|
||||
temperature=0.0,
|
||||
temperature=get_config().llm_temperature_verification,
|
||||
scope="verification",
|
||||
max_retries=0,
|
||||
)
|
||||
@@ -121,6 +163,13 @@ class LiteLLMLLM(LLMInterface):
|
||||
for key, value in self._extra_body.items():
|
||||
kwargs.setdefault(key, value)
|
||||
|
||||
# Forward operator-configured default headers as ``extra_headers`` so they
|
||||
# reach the provider behind LiteLLM (proxies / request-tracing middleware).
|
||||
# ``setdefault`` keeps any explicit per-call ``extra_headers`` authoritative;
|
||||
# a per-call copy prevents LiteLLM/downstream from mutating the stored dict.
|
||||
if self._default_headers:
|
||||
kwargs.setdefault("extra_headers", dict(self._default_headers))
|
||||
|
||||
# Bedrock service tier: flex (50% cheaper), priority, or reserved
|
||||
if self.model.startswith("bedrock/") and self.bedrock_service_tier is not None:
|
||||
kwargs["service_tier"] = self.bedrock_service_tier
|
||||
@@ -193,7 +242,7 @@ class LiteLLMLLM(LLMInterface):
|
||||
|
||||
# Add JSON schema response format if provided
|
||||
if response_format is not None and hasattr(response_format, "model_json_schema"):
|
||||
schema = response_format.model_json_schema()
|
||||
schema = strict_json_schema(response_format) if strict_schema else response_format.model_json_schema()
|
||||
call_kwargs["response_format"] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
@@ -209,7 +258,14 @@ class LiteLLMLLM(LLMInterface):
|
||||
if attempt > 0:
|
||||
set_stage(f"llm.{self._stage_label}.{scope}.attempt={attempt + 1}/{max_retries + 1}")
|
||||
try:
|
||||
response = await self._acompletion(**call_kwargs)
|
||||
response = await asyncio.wait_for(
|
||||
self._acompletion(**call_kwargs),
|
||||
timeout=self.timeout,
|
||||
)
|
||||
# Stash usage before the length check and parse/validate below,
|
||||
# which may raise locally even though the provider charged for
|
||||
# these tokens (#2387).
|
||||
stash_response_usage(_usage_from_litellm_response(response))
|
||||
|
||||
content = response.choices[0].message.content or ""
|
||||
finish_reason = response.choices[0].finish_reason
|
||||
@@ -230,7 +286,17 @@ class LiteLLMLLM(LLMInterface):
|
||||
try:
|
||||
json_data = json.loads(clean_content)
|
||||
except json.JSONDecodeError:
|
||||
json_data = json.loads(content)
|
||||
try:
|
||||
json_data = json.loads(content)
|
||||
except json.JSONDecodeError:
|
||||
if attempt < max_retries:
|
||||
# Prefer a clean re-roll first — a fresh generation
|
||||
# usually beats repairing a malformed one.
|
||||
raise
|
||||
# Retry budget spent: structural repair as a last
|
||||
# resort (#2547/#2544). Raises again if unrecoverable,
|
||||
# which the outer handler surfaces loudly.
|
||||
json_data = parse_llm_json(content)
|
||||
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
@@ -240,8 +306,9 @@ class LiteLLMLLM(LLMInterface):
|
||||
result = content
|
||||
|
||||
# Extract usage
|
||||
input_tokens = getattr(response.usage, "prompt_tokens", 0) or 0
|
||||
output_tokens = getattr(response.usage, "completion_tokens", 0) or 0
|
||||
response_usage = _usage_from_litellm_response(response)
|
||||
input_tokens = response_usage.input_tokens
|
||||
output_tokens = response_usage.output_tokens
|
||||
total_tokens = input_tokens + output_tokens
|
||||
|
||||
# Record metrics
|
||||
@@ -304,6 +371,25 @@ class LiteLLMLLM(LLMInterface):
|
||||
logger.error(f"LiteLLM returned invalid JSON after {max_retries + 1} attempts")
|
||||
raise
|
||||
|
||||
except (TimeoutError, asyncio.TimeoutError, LiteLLMTimeout) as e:
|
||||
# litellm/httpx don't always honor their own ``timeout=`` (e.g. a connection held
|
||||
# open with no token progress), so ``wait_for`` is the hard cap that cancels a hung
|
||||
# call regardless — otherwise one straggler pins a worker slot and stalls its gather.
|
||||
last_exception = e
|
||||
exc_name = type(e).__name__
|
||||
if attempt < max_retries:
|
||||
logger.warning(
|
||||
f"LiteLLM call exceeded timeout={self.timeout}s ({exc_name}, scope={scope}), retrying..."
|
||||
)
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
logger.error(
|
||||
f"LiteLLM call timed out after {self.timeout}s on {attempt + 1} attempts "
|
||||
f"({exc_name}, scope={scope})"
|
||||
)
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
error_str = str(e).lower()
|
||||
# Fast fail on auth errors
|
||||
@@ -311,6 +397,9 @@ class LiteLLMLLM(LLMInterface):
|
||||
logger.error(f"LiteLLM auth error, not retrying: {e}")
|
||||
raise
|
||||
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx.
|
||||
dump_request_on_4xx(scope=scope, provider=self.provider, model=self.model, err=e, request=call_kwargs)
|
||||
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
# Retry on rate limits, connection errors, server errors
|
||||
@@ -341,20 +430,38 @@ class LiteLLMLLM(LLMInterface):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
) -> LLMToolCallResult:
|
||||
start_time = time.time()
|
||||
|
||||
call_kwargs = self._build_common_kwargs(messages, max_completion_tokens, temperature)
|
||||
call_kwargs["tools"] = tools
|
||||
call_kwargs["tool_choice"] = tool_choice
|
||||
call_kwargs["tool_choice"] = (
|
||||
{
|
||||
"type": "function",
|
||||
"function": {"name": tool_choice.selected_function_name},
|
||||
}
|
||||
if tool_choice.mode is LLMToolChoiceMode.NAMED
|
||||
else tool_choice.mode.value
|
||||
)
|
||||
|
||||
last_exception = None
|
||||
for attempt in range(max_retries + 1):
|
||||
if attempt > 0:
|
||||
set_stage(f"llm.{self._stage_label}.tools.attempt={attempt + 1}/{max_retries + 1}")
|
||||
try:
|
||||
response = await self._acompletion(**call_kwargs)
|
||||
response = await asyncio.wait_for(
|
||||
self._acompletion(**call_kwargs),
|
||||
timeout=self.timeout,
|
||||
)
|
||||
# Stash usage before the tool-call argument parse below, which
|
||||
# can raise json.JSONDecodeError locally even though the provider
|
||||
# already billed for these tokens; without this the error trace
|
||||
# records 0/0 tokens (#2387). Mirrors call() and the anthropic/
|
||||
# gemini call_with_tools paths so the litellm tool path (and the
|
||||
# LiteLLMRouterLLM subclass that inherits this method) completes
|
||||
# the #2396 usage-on-error coverage.
|
||||
stash_response_usage(_usage_from_litellm_response(response))
|
||||
|
||||
message = response.choices[0].message
|
||||
content = message.content
|
||||
@@ -424,11 +531,31 @@ class LiteLLMLLM(LLMInterface):
|
||||
output_tokens=output_tokens,
|
||||
)
|
||||
|
||||
except (TimeoutError, asyncio.TimeoutError, LiteLLMTimeout) as e:
|
||||
# See ``call`` — hard cap so a hung completion cannot block
|
||||
# forever and pin a worker slot / concurrency permit.
|
||||
last_exception = e
|
||||
exc_name = type(e).__name__
|
||||
if attempt < max_retries:
|
||||
logger.warning(
|
||||
f"LiteLLM tool call exceeded timeout={self.timeout}s ({exc_name}, scope={scope}), retrying..."
|
||||
)
|
||||
await asyncio.sleep(min(initial_backoff * (2**attempt), max_backoff))
|
||||
continue
|
||||
logger.error(
|
||||
f"LiteLLM tool call timed out after {self.timeout}s on {attempt + 1} attempts "
|
||||
f"({exc_name}, scope={scope})"
|
||||
)
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
error_str = str(e).lower()
|
||||
if "401" in error_str or "403" in error_str or "unauthorized" in error_str:
|
||||
raise
|
||||
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx.
|
||||
dump_request_on_4xx(scope=scope, provider=self.provider, model=self.model, err=e, request=call_kwargs)
|
||||
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
is_retryable = any(
|
||||
|
||||
@@ -67,7 +67,7 @@ class LiteLLMRouterLLM(LiteLLMLLM):
|
||||
model: str,
|
||||
config: dict[str, Any],
|
||||
reasoning_effort: str = "low",
|
||||
timeout: float = 300.0,
|
||||
timeout: float | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(
|
||||
@@ -146,16 +146,28 @@ class LiteLLMRouterLLM(LiteLLMLLM):
|
||||
kwargs["max_completion_tokens"] = self._cap_max_completion_tokens(max_completion_tokens)
|
||||
if temperature is not None:
|
||||
kwargs["temperature"] = temperature
|
||||
|
||||
# Forward operator-configured default headers as ``extra_headers`` so they
|
||||
# reach the provider behind the Router (proxies / request-tracing middleware).
|
||||
# This override deliberately omits api_key/base_url/extra_body (those live in
|
||||
# the per-deployment Router config), but headers are a cross-cutting operator
|
||||
# concern, so we inject them here too — mirroring the base provider.
|
||||
# ``setdefault`` keeps any explicit per-call ``extra_headers`` authoritative;
|
||||
# a per-call copy prevents LiteLLM/downstream from mutating the stored dict.
|
||||
if self._default_headers:
|
||||
kwargs.setdefault("extra_headers", dict(self._default_headers))
|
||||
return kwargs
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
from hindsight_api.engine.llm_interface import OutputTooLongError
|
||||
|
||||
from ...config import get_config
|
||||
|
||||
try:
|
||||
await self.call(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
max_completion_tokens=50,
|
||||
temperature=0.0,
|
||||
temperature=get_config().llm_temperature_verification,
|
||||
scope="verification",
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
@@ -22,7 +22,7 @@ import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface
|
||||
from hindsight_api.engine.llm_interface import LLM_TOOL_CHOICE_AUTO, LLMInterface, LLMToolChoice
|
||||
from hindsight_api.engine.response_models import LLMToolCallResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -394,7 +394,7 @@ class LlamaCppLLM(LLMInterface):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
) -> LLMToolCallResult:
|
||||
"""Delegate tool calls to the OpenAI-compatible API."""
|
||||
await self._ensure_initialized()
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
"""Opt-in diagnostic: dump the exact request behind an LLM 4xx rejection.
|
||||
|
||||
Some ``400 INVALID_ARGUMENT`` / ``400 Bad Request`` rejections of structured-output
|
||||
calls are not reproducible by reconstructing the request after the fact — the failing
|
||||
factor lives in the request as it was actually assembled at runtime. Reconstructed
|
||||
replays of the same inputs return ``200``, so the only reliable way to see what the
|
||||
model rejected is to capture the real request at the moment it fails.
|
||||
|
||||
This helper is provider-agnostic. Every provider's error handler calls
|
||||
``dump_request_on_4xx`` with whatever it assembled — a Pydantic config
|
||||
(google-genai ``GenerateContentConfig``), a kwargs dict (OpenAI / Anthropic /
|
||||
LiteLLM ``**call_params``), etc. — plus the raised error. The helper self-gates:
|
||||
it is a no-op unless the ``llm_debug_dump_4xx`` config flag
|
||||
(``HINDSIGHT_API_LLM_DEBUG_DUMP_4XX``) is enabled AND the error carries a 4xx
|
||||
status, so callers can drop one unconditional call into each ``except`` block.
|
||||
|
||||
Safety / scope:
|
||||
- Off by default — the config flag is unset in normal operation.
|
||||
- The serialized config omits message bodies (the ``messages``/``contents``/``input``
|
||||
keys are stripped); message previews are length-capped, so an enabled dump can't
|
||||
flood logs or spill large bodies.
|
||||
- Never raises — diagnostics must not break the request path (falls back to ``repr``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Top-level request keys whose values are message bodies. Stripped from the config
|
||||
# view so the dump never spills large user content — previews are logged separately.
|
||||
_CONTENT_KEYS = ("messages", "contents", "input")
|
||||
|
||||
_PREVIEW_CHARS = 1500
|
||||
_CONFIG_REPR_CAP = 8000
|
||||
_ERR_CAP = 200
|
||||
|
||||
|
||||
def _enabled() -> bool:
|
||||
from hindsight_api.config import get_config
|
||||
|
||||
return bool(get_config().llm_debug_dump_4xx)
|
||||
|
||||
|
||||
def status_code_of(err: Any) -> int | None:
|
||||
"""Best-effort HTTP status of a provider error, across SDK error shapes.
|
||||
|
||||
OpenAI/Anthropic expose ``status_code``; google-genai uses ``code``; some wrap the
|
||||
status on a ``response``. Returns None when no integer status is discoverable.
|
||||
"""
|
||||
for attr in ("status_code", "code", "http_status"):
|
||||
value = getattr(err, attr, None)
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
response = getattr(err, "response", None)
|
||||
if response is not None:
|
||||
value = getattr(response, "status_code", None)
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _serialize_config(request: Any) -> str:
|
||||
"""Render the request config to a string without message bodies, never raising."""
|
||||
try:
|
||||
if request is None:
|
||||
return "null"
|
||||
# Pydantic models (google-genai GenerateContentConfig, SDK params objects).
|
||||
dump = getattr(request, "model_dump_json", None)
|
||||
if callable(dump):
|
||||
return dump(exclude_none=True)
|
||||
if isinstance(request, dict):
|
||||
view = {k: v for k, v in request.items() if k not in _CONTENT_KEYS}
|
||||
return json.dumps(view, ensure_ascii=False, default=str)
|
||||
return repr(request)[:_CONFIG_REPR_CAP]
|
||||
except Exception:
|
||||
return repr(request)[:_CONFIG_REPR_CAP]
|
||||
|
||||
|
||||
@dataclass
|
||||
class _MessagePreview:
|
||||
"""A message rendered for the dump: role + extracted text (not yet length-capped)."""
|
||||
|
||||
role: str
|
||||
text: str
|
||||
|
||||
|
||||
def _message_preview(msg: Any) -> _MessagePreview:
|
||||
"""Extract role + text from a message across dict and provider-object shapes."""
|
||||
# OpenAI / Anthropic dict: {"role": ..., "content": str | list[block]}
|
||||
if isinstance(msg, dict):
|
||||
role = str(msg.get("role", "?"))
|
||||
content = msg.get("content")
|
||||
if isinstance(content, str):
|
||||
return _MessagePreview(role, content)
|
||||
if isinstance(content, list):
|
||||
text = ""
|
||||
for block in content:
|
||||
if isinstance(block, dict):
|
||||
text += block.get("text") or ""
|
||||
else:
|
||||
text += getattr(block, "text", "") or ""
|
||||
return _MessagePreview(role, text)
|
||||
return _MessagePreview(role, "" if content is None else str(content))
|
||||
# google-genai Content: role + parts[].text
|
||||
role = str(getattr(msg, "role", "?"))
|
||||
text = ""
|
||||
for part in getattr(msg, "parts", None) or []:
|
||||
text += getattr(part, "text", None) or ""
|
||||
if not text:
|
||||
text = getattr(msg, "content", "") or ""
|
||||
return _MessagePreview(role, text)
|
||||
|
||||
|
||||
def _resolve_messages(request: Any, messages: Any) -> Any:
|
||||
"""Where per-message previews come from: explicit ``messages``, else inside ``request``."""
|
||||
if messages is not None:
|
||||
return messages
|
||||
if isinstance(request, dict):
|
||||
for key in _CONTENT_KEYS:
|
||||
if key in request:
|
||||
return request[key]
|
||||
return []
|
||||
|
||||
|
||||
def dump_request_on_4xx(
|
||||
*,
|
||||
scope: str,
|
||||
provider: str,
|
||||
model: str,
|
||||
err: Any,
|
||||
request: Any = None,
|
||||
messages: Any = None,
|
||||
) -> None:
|
||||
"""Log the exact request behind an LLM 4xx when the diagnostic is enabled.
|
||||
|
||||
No-op unless ``HINDSIGHT_API_LLM_DEBUG_DUMP_4XX`` is truthy and ``err`` carries a
|
||||
4xx status. ``request`` is whatever the provider assembled (a Pydantic config, a
|
||||
kwargs dict, ...); ``messages`` overrides where the per-message previews come from
|
||||
(defaults to the message list found inside ``request``).
|
||||
"""
|
||||
if not _enabled():
|
||||
return
|
||||
code = status_code_of(err)
|
||||
if code is None or not (400 <= code < 500):
|
||||
return
|
||||
try:
|
||||
cfg_repr = _serialize_config(request)
|
||||
summary = []
|
||||
for msg in _resolve_messages(request, messages) or []:
|
||||
m = _message_preview(msg)
|
||||
summary.append({"role": m.role, "chars": len(m.text), "preview": m.text[:_PREVIEW_CHARS]})
|
||||
logger.error(
|
||||
"[LLM_4XX_DUMP] provider=%s model=%s scope=%s code=%s err=%s config=%s contents=%s",
|
||||
provider,
|
||||
model,
|
||||
scope,
|
||||
code,
|
||||
str(err)[:_ERR_CAP],
|
||||
cfg_repr,
|
||||
json.dumps(summary, ensure_ascii=False),
|
||||
)
|
||||
except Exception as dump_exc: # never let diagnostics break the request path
|
||||
logger.warning("[LLM_4XX_DUMP] failed to serialize rejected request: %s", dump_exc)
|
||||
@@ -9,7 +9,7 @@ import logging
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from ..llm_interface import LLMInterface
|
||||
from ..llm_interface import LLM_TOOL_CHOICE_AUTO, LLMInterface, LLMToolChoice
|
||||
from ..response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -101,7 +101,7 @@ class MockLLM(LLMInterface):
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
response_format: Optional Pydantic model for structured output.
|
||||
max_completion_tokens: Not used in mock.
|
||||
temperature: Not used in mock.
|
||||
temperature: Recorded on the call record for test assertions.
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Not used in mock.
|
||||
initial_backoff: Not used in mock.
|
||||
@@ -123,6 +123,9 @@ class MockLLM(LLMInterface):
|
||||
if response_format and hasattr(response_format, "__name__")
|
||||
else str(response_format),
|
||||
"scope": scope,
|
||||
# Record the temperature so tests can assert per-operation temperature
|
||||
# wiring (None means the parameter was omitted from the call).
|
||||
"temperature": temperature,
|
||||
}
|
||||
self._mock_calls.append(call_record)
|
||||
logger.debug(f"Mock LLM call recorded: scope={scope}, model={self.model}")
|
||||
@@ -197,7 +200,7 @@ class MockLLM(LLMInterface):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make a mock LLM API call with tool/function calling support.
|
||||
@@ -208,7 +211,7 @@ class MockLLM(LLMInterface):
|
||||
messages: List of message dicts. Can include tool results with role='tool'.
|
||||
tools: List of tool definitions in OpenAI format.
|
||||
max_completion_tokens: Not used in mock.
|
||||
temperature: Not used in mock.
|
||||
temperature: Recorded on the call record for test assertions.
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Not used in mock.
|
||||
initial_backoff: Not used in mock.
|
||||
@@ -225,6 +228,9 @@ class MockLLM(LLMInterface):
|
||||
"messages": messages,
|
||||
"tools": [t.get("function", {}).get("name") for t in tools],
|
||||
"scope": scope,
|
||||
# Record the temperature so tests can assert per-operation temperature
|
||||
# wiring (None means the parameter was omitted from the call).
|
||||
"temperature": temperature,
|
||||
}
|
||||
self._mock_calls.append(call_record)
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ it raises a clear error instead of a confusing connection failure.
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from ..llm_interface import LLMInterface
|
||||
from ..llm_interface import LLM_TOOL_CHOICE_AUTO, LLMInterface, LLMToolChoice
|
||||
from ..response_models import LLMToolCallResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -65,7 +65,7 @@ class NoneLLM(LLMInterface):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
) -> LLMToolCallResult:
|
||||
"""Raise LLMNotAvailableError — no LLM is configured."""
|
||||
raise LLMNotAvailableError(
|
||||
|
||||
@@ -26,6 +26,8 @@ import logging
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from email.utils import parsedate_to_datetime
|
||||
from typing import Any
|
||||
from urllib.parse import parse_qs, urlparse, urlunparse
|
||||
|
||||
@@ -34,8 +36,18 @@ from openai import APIConnectionError, APIStatusError, AsyncOpenAI, LengthFinish
|
||||
|
||||
from hindsight_api.config import DEFAULT_LLM_TIMEOUT, ENV_LLM_TIMEOUT
|
||||
from hindsight_api.engine.bank_attribution import apply_bank_attribution
|
||||
from hindsight_api.engine.llm_interface import LLMInterface, OutputTooLongError
|
||||
from hindsight_api.engine.llm_interface import (
|
||||
LLM_TOOL_CHOICE_AUTO,
|
||||
LLMInterface,
|
||||
LLMToolChoice,
|
||||
LLMToolChoiceMode,
|
||||
OutputTooLongError,
|
||||
ProviderRateLimitResetError,
|
||||
)
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.providers.llm_debug import dump_request_on_4xx
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.engine.structured_output import strict_json_schema
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
from hindsight_api.worker.stage import set_stage
|
||||
|
||||
@@ -44,17 +56,57 @@ logger = logging.getLogger(__name__)
|
||||
# Seed applied to every Groq request for deterministic behavior
|
||||
DEFAULT_LLM_SEED = 4242
|
||||
JSON_MODE_USER_HINT = "Return valid json only."
|
||||
DEFAULT_VERIFICATION_MAX_COMPLETION_TOKENS = 512
|
||||
|
||||
# Self-hosted OpenAI-compatible servers that advertise tool_choice="required"
|
||||
|
||||
def _validate_ollama_num_ctx(value: Any) -> int | None:
|
||||
"""Validate a native Ollama context-window override."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError(f"ollama_num_ctx must be a positive integer, got {value!r}")
|
||||
if value < 1:
|
||||
raise ValueError(f"ollama_num_ctx must be >= 1, got {value}")
|
||||
return value
|
||||
|
||||
|
||||
# Provider implementations that advertise tool_choice="required"
|
||||
# but silently ignore it: instead of forcing a tool call they return
|
||||
# finish_reason "stop"/"tool_calls" with an EMPTY tool_calls array and no error.
|
||||
# Reflect's agent loop then sees no tool call, runs synthesis with no retrieval,
|
||||
# and answers "I don't have information" even when the bank holds the answer.
|
||||
# See issues #1563 (LM Studio), #1179 (LM Studio + Qwen), #1877 (vLLM with
|
||||
# --enable-auto-tool-choice). llama-server (the "llamacpp" provider) honors
|
||||
# "required" correctly and is intentionally excluded (#1179).
|
||||
# --enable-auto-tool-choice). The generic OpenAI provider is intentionally not
|
||||
# inferred from its URL: custom OpenAI-compatible endpoints can implement the
|
||||
# required-tool contract, and silently downgrading them changes request semantics.
|
||||
# llama-server (the "llamacpp" provider) honors "required" correctly and is
|
||||
# intentionally excluded (#1179).
|
||||
_TOOL_CHOICE_REQUIRED_UNSUPPORTED_PROVIDERS = frozenset({"lmstudio", "ollama"})
|
||||
|
||||
# Local providers whose OpenAI-compatible surface always lives under a `/v1`
|
||||
# path (LM Studio: http://localhost:1234/v1, Ollama: http://localhost:11434/v1).
|
||||
# For these we know the exact endpoint shape, so a bare host base URL can be
|
||||
# normalized safely. Cloud/proxy endpoints are left untouched — their path is
|
||||
# provider-specific and must be supplied verbatim.
|
||||
_V1_PATH_LOCAL_PROVIDERS = frozenset({"lmstudio", "ollama"})
|
||||
|
||||
|
||||
def _ensure_v1_base_url(base_url: str) -> str:
|
||||
"""Append the OpenAI-compatible ``/v1`` prefix to a bare local base URL.
|
||||
|
||||
LM Studio's server UI advertises its address as ``http://localhost:1234``,
|
||||
so users commonly set ``HINDSIGHT_API_LLM_BASE_URL`` to that bare host. The
|
||||
OpenAI SDK then POSTs to ``<host>/chat/completions`` and LM Studio rejects it
|
||||
with ``Unexpected endpoint or method`` — its OpenAI-compatible routes live
|
||||
under ``/v1``. Only a base URL with no meaningful path (bare host or a lone
|
||||
trailing slash) is rewritten; anything with an explicit path (e.g. a reverse
|
||||
proxy mount or an already-correct ``/v1``) is returned unchanged. See #2922.
|
||||
"""
|
||||
parsed = urlparse(base_url)
|
||||
if parsed.path.strip("/"):
|
||||
return base_url
|
||||
return urlunparse(parsed._replace(path="/v1"))
|
||||
|
||||
|
||||
class ProviderResponseError(RuntimeError):
|
||||
"""Raised when a provider returns a success response without usable content."""
|
||||
@@ -64,23 +116,111 @@ class ProviderResponseError(RuntimeError):
|
||||
self.retryable = retryable
|
||||
|
||||
|
||||
def _is_json(text: str) -> bool:
|
||||
"""True if ``text`` parses as a JSON value."""
|
||||
try:
|
||||
json.loads(text)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _outer_json_span(content: str) -> str | None:
|
||||
"""Return the outermost ``{...}`` / ``[...]`` span if it parses as JSON, else None.
|
||||
|
||||
Fallback for responses where fences are partial/absent or the model wrapped
|
||||
the JSON in surrounding prose. Only returned when it is valid JSON so callers
|
||||
never receive a worse candidate than the raw content.
|
||||
"""
|
||||
starts = [i for i in (content.find("{"), content.find("[")) if i >= 0]
|
||||
ends = [i for i in (content.rfind("}"), content.rfind("]")) if i >= 0]
|
||||
if not starts or not ends:
|
||||
return None
|
||||
start, end = min(starts), max(ends)
|
||||
if end <= start:
|
||||
return None
|
||||
candidate = content[start : end + 1].strip()
|
||||
return candidate if _is_json(candidate) else None
|
||||
|
||||
|
||||
def _strip_code_fences(content: str) -> str:
|
||||
"""Strip markdown code fences from LLM response if present.
|
||||
|
||||
Many LLM providers (MiniMax, some Ollama models, Claude via proxies)
|
||||
wrap JSON responses in ```json ... ``` fences even when json_object
|
||||
response format is requested. This strips the fences while preserving
|
||||
the JSON content inside. Returns the original content unchanged if
|
||||
no fences are detected.
|
||||
response format is requested. Fences are detected by line (a closing
|
||||
``` must sit alone on its line) so triple-backticks *inside* JSON string
|
||||
values do not truncate the payload. When the stripped candidate is not
|
||||
valid JSON (partial fence, prose-wrapped output, truncated response), fall
|
||||
back to the outermost parseable JSON span. Returns the original content
|
||||
unchanged if no better candidate is found.
|
||||
"""
|
||||
if "```" not in content:
|
||||
return content
|
||||
try:
|
||||
if "```json" in content:
|
||||
return content.split("```json")[1].split("```")[0].strip()
|
||||
return content.split("```")[1].split("```")[0].strip()
|
||||
except (IndexError, ValueError):
|
||||
return content
|
||||
candidate = content
|
||||
if "```" in content:
|
||||
lines = content.split("\n")
|
||||
# Find first line that starts a code fence (``` optionally followed by language)
|
||||
fence_start = next((i for i, line in enumerate(lines) if line.startswith("```")), None)
|
||||
if fence_start is not None:
|
||||
# Find matching closing fence (``` alone or with trailing whitespace)
|
||||
fence_end = next(
|
||||
(j for j in range(fence_start + 1, len(lines)) if lines[j].strip() == "```"),
|
||||
None,
|
||||
)
|
||||
if fence_end is not None:
|
||||
candidate = "\n".join(lines[fence_start + 1 : fence_end]).strip()
|
||||
|
||||
if _is_json(candidate):
|
||||
return candidate
|
||||
|
||||
# Fence stripping did not yield valid JSON — try to recover the outer JSON span.
|
||||
span = _outer_json_span(content)
|
||||
if span is not None:
|
||||
return span
|
||||
|
||||
return candidate
|
||||
|
||||
|
||||
# Reasoning/thinking tags emitted by extended-thinking models. Some providers
|
||||
# (e.g. MiniMax-M3) leak the chain-of-thought wrapped in these tags into the
|
||||
# response body instead of a separate reasoning_content field. Each entry is
|
||||
# (open_tag, close_tag); the open tag also matches when the close tag is missing
|
||||
# (truncated output) so a dangling block is removed to end-of-string.
|
||||
_REASONING_TAG_PAIRS: tuple[tuple[str, str], ...] = (
|
||||
("<think>", "</think>"),
|
||||
("<thinking>", "</thinking>"),
|
||||
("<thought>", "</thought>"),
|
||||
("<reasoning>", "</reasoning>"),
|
||||
("|startthink|", "|endthink|"),
|
||||
)
|
||||
|
||||
|
||||
def _strip_reasoning_tags(text: str) -> str:
|
||||
"""Strip extended-thinking/reasoning blocks from an LLM response.
|
||||
|
||||
Removes the full set of tag styles emitted by reasoning models:
|
||||
``<think>``, ``<thinking>``, ``<thought>``, ``<reasoning>`` and the
|
||||
``|startthink|...|endthink|`` markers. Both the structured (JSON) path and
|
||||
the free-form path must call this — otherwise a non-structured response
|
||||
(e.g. a mental-model markdown blob from MiniMax-M3) leaks the raw
|
||||
``<think>...</think>`` verbatim into stored memories.
|
||||
|
||||
Handles two cases:
|
||||
1. Closed blocks: ``<think>...</think>`` removed wherever they appear.
|
||||
2. Unclosed blocks: a dangling ``<think>`` with no closing tag (model output
|
||||
truncated mid-thought) is removed from the open tag to end-of-string.
|
||||
|
||||
Returns the input unchanged (modulo surrounding whitespace) when no tags are
|
||||
present.
|
||||
"""
|
||||
if not text:
|
||||
return text
|
||||
for open_tag, close_tag in _REASONING_TAG_PAIRS:
|
||||
open_re = re.escape(open_tag)
|
||||
close_re = re.escape(close_tag)
|
||||
# Closed blocks first, then any remaining unclosed (truncated) block.
|
||||
text = re.sub(rf"{open_re}.*?{close_re}", "", text, flags=re.DOTALL)
|
||||
text = re.sub(rf"{open_re}.*", "", text, flags=re.DOTALL)
|
||||
return text.strip()
|
||||
|
||||
|
||||
def _response_get(response: Any, key: str, default: Any = None) -> Any:
|
||||
@@ -187,6 +327,21 @@ def _content_or_error(response: Any, *, provider: str, model: str, scope: str) -
|
||||
return content, choice
|
||||
|
||||
|
||||
def _usage_from_openai_response(response: Any) -> LLMResponseUsage:
|
||||
"""Extract prompt/completion/cached token counts from an OpenAI-shaped usage block."""
|
||||
usage = getattr(response, "usage", None)
|
||||
input_tokens = (usage.prompt_tokens or 0) if usage else 0
|
||||
output_tokens = (usage.completion_tokens or 0) if usage else 0
|
||||
cached_tokens = 0
|
||||
if usage and getattr(usage, "prompt_tokens_details", None):
|
||||
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
|
||||
return LLMResponseUsage(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
)
|
||||
|
||||
|
||||
def _ensure_json_word_in_user_message(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
"""Some OpenAI-compatible gateways require 'json' in a user message for json_object mode."""
|
||||
|
||||
@@ -234,6 +389,122 @@ def _summarize_status_error(e: APIStatusError, body_max: int = 400) -> str:
|
||||
return f"HTTP {e.status_code}: {body_str or '<no body>'}"
|
||||
|
||||
|
||||
_RATE_LIMIT_RESET_AT_RE = re.compile(
|
||||
r"\breset at\s+"
|
||||
r"(?P<reset_at>\d{4}-\d{2}-\d{2}[ T]\d{2}:\d{2}:\d{2}(?:\s*(?:Z|[+-]\d{2}:?\d{2}))?)",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
_RATE_LIMIT_WINDOW_RE = re.compile(
|
||||
r"\b(?:for|in)\s+(?P<amount>\d+)\s*(?P<unit>second|minute|hour|day)s?\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _status_error_body_text(e: APIStatusError) -> str:
|
||||
body: Any = getattr(e, "body", None)
|
||||
if body is None:
|
||||
try:
|
||||
body = e.response.text
|
||||
except Exception:
|
||||
body = None
|
||||
if isinstance(body, (dict, list)):
|
||||
try:
|
||||
return json.dumps(body, default=str, ensure_ascii=False)
|
||||
except Exception:
|
||||
return str(body)
|
||||
return str(body or "").strip()
|
||||
|
||||
|
||||
def _parse_retry_after_header(value: str | None, now: datetime) -> datetime | None:
|
||||
if not value:
|
||||
return None
|
||||
raw = value.strip()
|
||||
try:
|
||||
seconds = float(raw)
|
||||
except ValueError:
|
||||
seconds = -1.0
|
||||
if seconds >= 0:
|
||||
return now + timedelta(seconds=seconds)
|
||||
|
||||
try:
|
||||
parsed = parsedate_to_datetime(raw)
|
||||
except (TypeError, ValueError, IndexError, OverflowError):
|
||||
return None
|
||||
if parsed.tzinfo is None:
|
||||
parsed = parsed.replace(tzinfo=UTC)
|
||||
return parsed.astimezone(UTC)
|
||||
|
||||
|
||||
def _parse_reset_at_datetime(value: str) -> datetime | None:
|
||||
raw = value.strip().replace(" ", "T")
|
||||
if raw.endswith("Z"):
|
||||
raw = f"{raw[:-1]}+00:00"
|
||||
elif re.search(r"[+-]\d{4}$", raw):
|
||||
raw = f"{raw[:-2]}:{raw[-2:]}"
|
||||
try:
|
||||
parsed = datetime.fromisoformat(raw)
|
||||
except ValueError:
|
||||
return None
|
||||
if parsed.tzinfo is None:
|
||||
# Some providers (z.ai included) return a wall-clock reset timestamp
|
||||
# without a zone. Interpret it in the host's local zone so logs, status
|
||||
# pages, and the queued next_retry_at describe the same operator-facing
|
||||
# clock instead of silently shifting by UTC offset.
|
||||
parsed = parsed.astimezone()
|
||||
return parsed.astimezone(UTC)
|
||||
|
||||
|
||||
def _rate_limit_retry_at(e: APIStatusError) -> datetime | None:
|
||||
now = datetime.now(UTC)
|
||||
response = getattr(e, "response", None)
|
||||
headers = getattr(response, "headers", None)
|
||||
if headers is not None:
|
||||
retry_at = _parse_retry_after_header(headers.get("retry-after") or headers.get("Retry-After"), now)
|
||||
if retry_at is not None and retry_at > now:
|
||||
return retry_at
|
||||
|
||||
body_text = _status_error_body_text(e)
|
||||
reset_match = _RATE_LIMIT_RESET_AT_RE.search(body_text)
|
||||
if reset_match:
|
||||
retry_at = _parse_reset_at_datetime(reset_match.group("reset_at"))
|
||||
if retry_at is not None and retry_at > now:
|
||||
return retry_at
|
||||
|
||||
window_match = _RATE_LIMIT_WINDOW_RE.search(body_text)
|
||||
if not window_match:
|
||||
return None
|
||||
amount = int(window_match.group("amount"))
|
||||
unit = window_match.group("unit").lower()
|
||||
if unit == "second":
|
||||
seconds = amount
|
||||
elif unit == "minute":
|
||||
seconds = amount * 60
|
||||
elif unit == "hour":
|
||||
seconds = amount * 3600
|
||||
else:
|
||||
seconds = amount * 86400
|
||||
return now + timedelta(seconds=seconds)
|
||||
|
||||
|
||||
def _raise_provider_quota_defer(
|
||||
e: APIStatusError, *, provider: str, model: str, scope: str, max_backoff: float
|
||||
) -> None:
|
||||
if e.status_code != 429:
|
||||
return
|
||||
retry_at = _rate_limit_retry_at(e)
|
||||
if retry_at is None:
|
||||
return
|
||||
if (retry_at - datetime.now(UTC)).total_seconds() <= max_backoff:
|
||||
return
|
||||
summary = _summarize_status_error(e)
|
||||
raise ProviderRateLimitResetError(
|
||||
retry_at=retry_at,
|
||||
message=(
|
||||
f"Provider quota exhausted ({provider}/{model}, scope={scope}); retry at {retry_at.isoformat()}: {summary}"
|
||||
),
|
||||
) from e
|
||||
|
||||
|
||||
class OpenAICompatibleLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider for OpenAI-compatible APIs.
|
||||
@@ -258,6 +529,8 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
timeout: float | None = None,
|
||||
groq_service_tier: str | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
*,
|
||||
ollama_num_ctx: int | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
@@ -269,9 +542,11 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
base_url: Base URL for the API (uses defaults for groq/ollama/lmstudio if empty).
|
||||
model: Model name.
|
||||
reasoning_effort: Reasoning effort level for supported models ("low", "medium", "high").
|
||||
timeout: Request timeout in seconds (uses env var or 300s default).
|
||||
timeout: Request timeout in seconds (uses env var or 120s default).
|
||||
groq_service_tier: Groq service tier ("on_demand", "flex", "auto").
|
||||
extra_body: Extra body params merged into every API call.
|
||||
ollama_num_ctx: Native Ollama context window override. None lets Ollama use
|
||||
the model/server default.
|
||||
**kwargs: Additional provider-specific parameters.
|
||||
"""
|
||||
super().__init__(provider, api_key, base_url, model, reasoning_effort, **kwargs)
|
||||
@@ -288,8 +563,10 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
"deepseek",
|
||||
"volcano",
|
||||
"openrouter",
|
||||
"requesty",
|
||||
"zai",
|
||||
"opencode-go",
|
||||
"atlas",
|
||||
"fireworks",
|
||||
]
|
||||
if self.provider not in valid_providers:
|
||||
@@ -311,15 +588,24 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
self.base_url = "https://api.deepseek.com"
|
||||
elif self.provider == "openrouter":
|
||||
self.base_url = "https://openrouter.ai/api/v1"
|
||||
elif self.provider == "requesty":
|
||||
self.base_url = "https://router.requesty.ai/v1"
|
||||
elif self.provider == "zai":
|
||||
self.base_url = "https://api.z.ai/api/coding/paas/v4"
|
||||
elif self.provider == "opencode-go":
|
||||
self.base_url = "https://opencode.ai/zen/go/v1"
|
||||
elif self.provider == "atlas":
|
||||
self.base_url = "https://api.atlascloud.ai/v1"
|
||||
elif self.provider == "fireworks":
|
||||
# OpenAI-compatible inference host (online path). The batch API
|
||||
# lives on a separate control-plane host — see FireworksLLM.
|
||||
self.base_url = "https://api.fireworks.ai/inference/v1"
|
||||
|
||||
# Normalize bare local base URLs (e.g. a user pasting the address shown
|
||||
# in the LM Studio UI) so the OpenAI SDK targets the `/v1` routes. See #2922.
|
||||
if self.provider in _V1_PATH_LOCAL_PROVIDERS and self.base_url:
|
||||
self.base_url = _ensure_v1_base_url(self.base_url)
|
||||
|
||||
# For ollama/lmstudio, use dummy key if not provided
|
||||
if self.provider in ("ollama", "lmstudio") and not self.api_key:
|
||||
self.api_key = "local"
|
||||
@@ -333,8 +619,10 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
"minimax",
|
||||
"deepseek",
|
||||
"openrouter",
|
||||
"requesty",
|
||||
"zai",
|
||||
"opencode-go",
|
||||
"atlas",
|
||||
"ollama-cloud",
|
||||
)
|
||||
and not self.api_key
|
||||
@@ -344,6 +632,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
# Service tier configuration (from config, not env vars)
|
||||
self.groq_service_tier = groq_service_tier
|
||||
self.openai_service_tier = kwargs.get("openai_service_tier")
|
||||
self.ollama_num_ctx = _validate_ollama_num_ctx(ollama_num_ctx)
|
||||
# User-configured extra body params (merged into every API call)
|
||||
self._config_extra_body = extra_body or {}
|
||||
|
||||
@@ -374,17 +663,17 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
def _drops_tool_choice_required(self) -> bool:
|
||||
"""Whether this endpoint silently ignores ``tool_choice="required"``.
|
||||
|
||||
True for self-hosted OpenAI-compatible servers known to return an empty
|
||||
tool_calls array for "required" instead of forcing a call (#1563/#1179/
|
||||
#1877). Covers LM Studio / Ollama directly, plus any server reached via
|
||||
the generic "openai" provider with a custom ``base_url`` (e.g. a local
|
||||
vLLM endpoint). The real OpenAI API (no base_url override) honors
|
||||
"required", and cloud providers keep their own default base_urls, so both
|
||||
are left untouched.
|
||||
Only explicitly identified provider implementations are classified as
|
||||
unsupported. A custom base URL does not identify endpoint capabilities:
|
||||
an OpenAI-compatible endpoint may correctly enforce required tool calls,
|
||||
and replacing ``required`` with ``auto`` would violate the caller's named
|
||||
tool choice after the tools list has been narrowed.
|
||||
"""
|
||||
if self.provider in _TOOL_CHOICE_REQUIRED_UNSUPPORTED_PROVIDERS:
|
||||
return True
|
||||
return self.provider == "openai" and bool(self.base_url)
|
||||
return self.provider in _TOOL_CHOICE_REQUIRED_UNSUPPORTED_PROVIDERS
|
||||
|
||||
def _verification_max_completion_tokens(self) -> int:
|
||||
"""Return the startup verification budget for OpenAI-compatible gateways."""
|
||||
return DEFAULT_VERIFICATION_MAX_COMPLETION_TOKENS
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""
|
||||
@@ -397,7 +686,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
logger.info(f"Verifying connection: {self.provider}/{self.model}")
|
||||
await self.call(
|
||||
messages=[{"role": "user", "content": "Say 'ok'"}],
|
||||
max_completion_tokens=100,
|
||||
max_completion_tokens=self._verification_max_completion_tokens(),
|
||||
max_retries=2,
|
||||
initial_backoff=0.5,
|
||||
max_backoff=2.0,
|
||||
@@ -459,6 +748,11 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
# use the widely-supported max_tokens
|
||||
return "max_tokens"
|
||||
|
||||
def _apply_provider_extra_body_defaults(self, extra_body: dict[str, Any]) -> None:
|
||||
"""Apply provider-specific extra_body defaults while preserving user overrides."""
|
||||
if self.provider == "minimax":
|
||||
extra_body.setdefault("thinking", {"type": "disabled"})
|
||||
|
||||
async def call(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
@@ -546,6 +840,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
|
||||
# Provider-specific parameters
|
||||
extra_body: dict[str, Any] = {**self._config_extra_body}
|
||||
self._apply_provider_extra_body_defaults(extra_body)
|
||||
if self.provider == "groq":
|
||||
call_params["seed"] = DEFAULT_LLM_SEED
|
||||
# Add service_tier if configured
|
||||
@@ -561,7 +856,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
if response_format is not None:
|
||||
schema = None
|
||||
if hasattr(response_format, "model_json_schema"):
|
||||
schema = response_format.model_json_schema()
|
||||
schema = strict_json_schema(response_format) if strict_schema else response_format.model_json_schema()
|
||||
|
||||
if strict_schema and schema is not None:
|
||||
# Use OpenAI's strict JSON schema enforcement
|
||||
@@ -609,6 +904,9 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
try:
|
||||
if response_format is not None:
|
||||
response = await self._client.chat.completions.create(**call_params)
|
||||
# Stash usage before parse/validate, which may raise locally
|
||||
# even though the provider charged for these tokens (#2387).
|
||||
stash_response_usage(_usage_from_openai_response(response))
|
||||
|
||||
content, first_choice = _content_or_error(
|
||||
response,
|
||||
@@ -617,15 +915,10 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
scope=scope,
|
||||
)
|
||||
|
||||
# Strip reasoning model thinking tags
|
||||
# Strip reasoning model thinking tags (closed and unclosed).
|
||||
# Supports: <think>, <thinking>, <thought>, <reasoning>, |startthink|/|endthink|
|
||||
original_len = len(content)
|
||||
content = re.sub(r"<think>.*?</think>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"<thinking>.*?</thinking>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"<thought>.*?</thought>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"<reasoning>.*?</reasoning>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"\|startthink\|.*?\|endthink\|", "", content, flags=re.DOTALL)
|
||||
content = content.strip()
|
||||
content = _strip_reasoning_tags(content)
|
||||
if len(content) < original_len:
|
||||
logger.debug(f"Stripped {original_len - len(content)} chars of reasoning tokens")
|
||||
|
||||
@@ -667,6 +960,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
result = response_format.model_validate(json_data)
|
||||
else:
|
||||
response = await self._client.chat.completions.create(**call_params)
|
||||
stash_response_usage(_usage_from_openai_response(response))
|
||||
result, first_choice = _content_or_error(
|
||||
response,
|
||||
provider=self.provider,
|
||||
@@ -674,17 +968,37 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
scope=scope,
|
||||
)
|
||||
|
||||
# Free-form (non-structured) output also leaks reasoning tags:
|
||||
# reasoning models like MiniMax-M3 wrap their chain-of-thought
|
||||
# in <think>...</think> in the response body. Without this strip
|
||||
# a mental-model markdown blob is stored verbatim with the raw
|
||||
# thinking tags. Mirrors the structured-output path above.
|
||||
result = _strip_reasoning_tags(result)
|
||||
|
||||
# Record token usage metrics
|
||||
duration = time.time() - start_time
|
||||
usage = response.usage
|
||||
input_tokens = usage.prompt_tokens or 0 if usage else 0
|
||||
output_tokens = usage.completion_tokens or 0 if usage else 0
|
||||
response_usage = _usage_from_openai_response(response)
|
||||
input_tokens = response_usage.input_tokens
|
||||
output_tokens = response_usage.output_tokens
|
||||
total_tokens = usage.total_tokens or 0 if usage else 0
|
||||
cached_tokens = 0
|
||||
if usage and getattr(usage, "prompt_tokens_details", None):
|
||||
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
|
||||
cached_tokens = response_usage.cached_tokens
|
||||
thoughts_tokens = 0
|
||||
if usage and getattr(usage, "completion_tokens_details", None):
|
||||
thoughts_tokens = getattr(usage.completion_tokens_details, "reasoning_tokens", 0) or 0
|
||||
# OpenAI-compatible providers fold reasoning tokens into
|
||||
# ``completion_tokens`` (and thus ``total_tokens``), but the
|
||||
# TokenUsage contract — and the Gemini provider — treat
|
||||
# ``output_tokens``/``total_tokens`` as visible-only, surfacing
|
||||
# reasoning separately in ``thoughts_tokens``. Subtract so the
|
||||
# two fields don't double-count reasoning (cost over-attribution).
|
||||
if thoughts_tokens:
|
||||
output_tokens = max(0, output_tokens - thoughts_tokens)
|
||||
total_tokens = max(0, total_tokens - thoughts_tokens)
|
||||
|
||||
# Record LLM metrics
|
||||
# Record LLM metrics. ``output_tokens`` is visible-only by now, so
|
||||
# ``thoughts_tokens`` has to be recorded alongside it or the reasoning
|
||||
# half of the billed output reaches no counter at all.
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
@@ -694,6 +1008,8 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
cached_input_tokens=cached_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
|
||||
# Record trace span
|
||||
@@ -731,6 +1047,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
return result, token_usage
|
||||
return result
|
||||
@@ -761,6 +1078,13 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
logger.error(f"Auth error (HTTP {e.status_code}), not retrying: {str(e)}")
|
||||
raise
|
||||
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx.
|
||||
dump_request_on_4xx(scope=scope, provider=self.provider, model=self.model, err=e, request=call_params)
|
||||
|
||||
_raise_provider_quota_defer(
|
||||
e, provider=self.provider, model=self.model, scope=scope, max_backoff=max_backoff
|
||||
)
|
||||
|
||||
# Handle tool_use_failed error - model outputted in tool call format
|
||||
if e.status_code == 400 and response_format is not None:
|
||||
try:
|
||||
@@ -814,7 +1138,6 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
f"scope={scope}): {_summarize_status_error(e)}"
|
||||
)
|
||||
raise
|
||||
|
||||
except ProviderResponseError as e:
|
||||
last_exception = e
|
||||
if e.retryable and attempt < max_retries:
|
||||
@@ -848,7 +1171,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
tool_choice: LLMToolChoice = LLM_TOOL_CHOICE_AUTO,
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make an LLM API call with tool/function calling support.
|
||||
@@ -862,51 +1185,43 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or specific function.
|
||||
tool_choice: Canonical tool-selection policy.
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
request_tool_choice: str | dict[str, Any] | None = tool_choice
|
||||
|
||||
# Normalize named tool_choice dicts to "required" + filter tools.
|
||||
# Some providers (e.g. LM Studio, Ollama) reject the OpenAI named format
|
||||
# {"type": "function", "function": {"name": "..."}}. The semantics are
|
||||
# identical to tool_choice="required" with the tools list restricted to
|
||||
# just the requested tool, so we apply that transformation where supported.
|
||||
if isinstance(request_tool_choice, dict) and request_tool_choice.get("type") == "function":
|
||||
forced_name = request_tool_choice.get("function", {}).get("name")
|
||||
if forced_name:
|
||||
filtered = [t for t in tools if t.get("function", {}).get("name") == forced_name]
|
||||
if filtered:
|
||||
tools = filtered
|
||||
request_tool_choice = "required"
|
||||
request_tool_choice: str | None
|
||||
if tool_choice.mode is LLMToolChoiceMode.NAMED:
|
||||
forced_name = tool_choice.selected_function_name
|
||||
filtered = [tool for tool in tools if tool.get("function", {}).get("name") == forced_name]
|
||||
if len(filtered) != 1:
|
||||
raise ValueError(
|
||||
f"Named tool_choice must reference exactly one declared tool; "
|
||||
f"found {len(filtered)} definitions for {forced_name!r}"
|
||||
)
|
||||
tools = filtered
|
||||
request_tool_choice = LLMToolChoiceMode.REQUIRED.value
|
||||
elif tool_choice.mode is LLMToolChoiceMode.AUTO:
|
||||
request_tool_choice = None
|
||||
else:
|
||||
request_tool_choice = tool_choice.mode.value
|
||||
|
||||
# DeepSeek accepts tool calls but rejects explicit required/named
|
||||
# tool_choice values. The tools list has already been narrowed for
|
||||
# forced calls, so omitting tool_choice preserves the practical behavior.
|
||||
if "deepseek" in self.model.lower() and request_tool_choice != "auto":
|
||||
if "deepseek" in self.model.lower() and tool_choice.mode is not LLMToolChoiceMode.AUTO:
|
||||
request_tool_choice = None
|
||||
|
||||
# "auto" is the OpenAI API default — omitting tool_choice is semantically
|
||||
# identical. Some providers (e.g. DeepSeek's reasoner pathway, which
|
||||
# deepseek-v4-flash falls into when thinking mode is enabled) reject the
|
||||
# parameter outright, returning HTTP 400 even for value "auto". Sending it
|
||||
# only when the caller asks for a non-default behaviour avoids those 400s
|
||||
# without changing semantics for compliant providers.
|
||||
if request_tool_choice == "auto":
|
||||
request_tool_choice = None
|
||||
|
||||
# vLLM (--enable-auto-tool-choice), LM Studio, Ollama and similar
|
||||
# self-hosted servers silently drop tool_choice="required", returning an
|
||||
# empty tool_calls array instead of forcing a call (#1563/#1179/#1877).
|
||||
# LM Studio and Ollama silently drop tool_choice="required", returning an
|
||||
# empty tool_calls array instead of forcing a call (#1563/#1179).
|
||||
# Downgrade to auto (None) so the model still gets to call a tool. Named
|
||||
# tool_choice dicts were already normalized to "required" + a single
|
||||
# filtered tool above, so the call stays practically forced even under
|
||||
# auto. The real OpenAI API honors "required" and is left untouched.
|
||||
if request_tool_choice == "required" and self._drops_tool_choice_required():
|
||||
# auto. Generic OpenAI-compatible endpoints retain the canonical
|
||||
# ``required`` contract regardless of whether they use a custom base URL.
|
||||
if request_tool_choice == LLMToolChoiceMode.REQUIRED.value and self._drops_tool_choice_required():
|
||||
request_tool_choice = None
|
||||
|
||||
# DeepSeek tool-call replies can carry provider-specific reasoning_content.
|
||||
@@ -943,6 +1258,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
|
||||
# Provider-specific parameters
|
||||
extra_body: dict[str, Any] = {**self._config_extra_body}
|
||||
self._apply_provider_extra_body_defaults(extra_body)
|
||||
if self.provider == "groq":
|
||||
call_params["seed"] = DEFAULT_LLM_SEED
|
||||
if extra_body:
|
||||
@@ -978,7 +1294,20 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
usage = response.usage
|
||||
input_tokens = usage.prompt_tokens or 0 if usage else 0
|
||||
output_tokens = usage.completion_tokens or 0 if usage else 0
|
||||
cached_tokens = 0
|
||||
if usage and getattr(usage, "prompt_tokens_details", None):
|
||||
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
|
||||
thoughts_tokens = 0
|
||||
if usage and getattr(usage, "completion_tokens_details", None):
|
||||
thoughts_tokens = getattr(usage.completion_tokens_details, "reasoning_tokens", 0) or 0
|
||||
# See ``call()``: OpenAI-compatible ``completion_tokens`` includes
|
||||
# reasoning, so make ``output_tokens`` visible-only to avoid
|
||||
# double-counting it against ``thoughts_tokens``.
|
||||
if thoughts_tokens:
|
||||
output_tokens = max(0, output_tokens - thoughts_tokens)
|
||||
|
||||
# See ``call()``: record the reasoning and cached counts too, so no
|
||||
# billed token is dropped from the metrics counters.
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
@@ -988,6 +1317,8 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
cached_input_tokens=cached_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
|
||||
# Record OpenTelemetry span
|
||||
@@ -1020,6 +1351,8 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
finish_reason=finish_reason,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
|
||||
except APIConnectionError as e:
|
||||
@@ -1047,6 +1380,14 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
f"not retrying: {_summarize_status_error(e)}"
|
||||
)
|
||||
raise
|
||||
|
||||
# Diagnostic dump (opt-in) of the exact request behind any 4xx.
|
||||
dump_request_on_4xx(scope=scope, provider=self.provider, model=self.model, err=e, request=call_params)
|
||||
|
||||
_raise_provider_quota_defer(
|
||||
e, provider=self.provider, model=self.model, scope=scope, max_backoff=max_backoff
|
||||
)
|
||||
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
logger.warning(
|
||||
@@ -1060,7 +1401,6 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
f"({self.provider}/{self.model}, scope={scope}): {_summarize_status_error(e)}"
|
||||
)
|
||||
raise
|
||||
|
||||
except Exception:
|
||||
raise
|
||||
|
||||
@@ -1115,9 +1455,10 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
|
||||
# Add optional parameters with optimized defaults for Ollama
|
||||
options: dict[str, Any] = {
|
||||
"num_ctx": 16384, # 16k context window for larger prompts
|
||||
"num_batch": 512, # Optimal batch size for prompt processing
|
||||
}
|
||||
if self.ollama_num_ctx is not None:
|
||||
options["num_ctx"] = self.ollama_num_ctx
|
||||
if max_completion_tokens:
|
||||
options["num_predict"] = max_completion_tokens
|
||||
if temperature is not None:
|
||||
|
||||
@@ -6,6 +6,7 @@ structured information like temporal constraints.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
@@ -19,6 +20,103 @@ from hindsight_api.engine.temporal_periods import (
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# dateparser.search_dates over-matches: short common words that happen to be
|
||||
# weekday/month abbreviations in *some* language ("we"/"me"/"did" -> a weekday,
|
||||
# "do" -> Sunday) come back as bogus dates. When such a false positive appears
|
||||
# *before* the real date in the query, taking the first match (or a hard-coded
|
||||
# blacklist of such words) silently produces a wrong temporal window — worse
|
||||
# than none, because the constraint is non-null so nothing downstream can tell
|
||||
# extraction failed. See issue #2768.
|
||||
#
|
||||
# Instead of blacklisting words one at a time (a moving target — every short
|
||||
# word dateparser resolves is a new instance of the same bug), we score each
|
||||
# match by the date signal it actually carries and keep only matches with a
|
||||
# real signal, preferring the strongest. A bare weekday abbreviation carries no
|
||||
# day/month/year and scores zero, so it is rejected regardless of language or
|
||||
# dateparser version.
|
||||
_TOKEN_RE = re.compile(r"[a-z0-9]+")
|
||||
_MONTH_WORDS = {
|
||||
"january",
|
||||
"february",
|
||||
"march",
|
||||
"april",
|
||||
"may",
|
||||
"june",
|
||||
"july",
|
||||
"august",
|
||||
"september",
|
||||
"october",
|
||||
"november",
|
||||
"december",
|
||||
}
|
||||
_RELATIVE_WORDS = {"today", "yesterday", "tomorrow", "tonight", "now"}
|
||||
_WEEKDAY_WORDS = {
|
||||
"monday",
|
||||
"tuesday",
|
||||
"wednesday",
|
||||
"thursday",
|
||||
"friday",
|
||||
"saturday",
|
||||
"sunday",
|
||||
}
|
||||
_PERIOD_WORDS = {
|
||||
"last",
|
||||
"next",
|
||||
"this",
|
||||
"past",
|
||||
"coming",
|
||||
"ago",
|
||||
"week",
|
||||
"weeks",
|
||||
"month",
|
||||
"months",
|
||||
"year",
|
||||
"years",
|
||||
"day",
|
||||
"days",
|
||||
"hour",
|
||||
"hours",
|
||||
"minute",
|
||||
"minutes",
|
||||
"quarter",
|
||||
"decade",
|
||||
"century",
|
||||
"weekend",
|
||||
"morning",
|
||||
"afternoon",
|
||||
"evening",
|
||||
"night",
|
||||
"noon",
|
||||
"midnight",
|
||||
}
|
||||
|
||||
|
||||
def _date_match_score(text: str) -> int:
|
||||
"""Score how strong a temporal signal a matched span carries.
|
||||
|
||||
A score of 0 means the span is a bare token with no explicit date content
|
||||
(the false-positive class from issue #2768) and should be rejected. Higher
|
||||
scores mean a stronger, less ambiguous date reference. A digit is the
|
||||
strongest signal (day/year/ISO date); an explicit English month/relative
|
||||
word next; weekday names and period words weakest but still explicit.
|
||||
"""
|
||||
tokens = _TOKEN_RE.findall(text.lower())
|
||||
if not tokens:
|
||||
return 0
|
||||
score = 0
|
||||
if any(any(ch.isdigit() for ch in tok) for tok in tokens):
|
||||
score += 100
|
||||
token_set = set(tokens)
|
||||
if token_set & _MONTH_WORDS:
|
||||
score += 50
|
||||
if token_set & _RELATIVE_WORDS:
|
||||
score += 50
|
||||
if token_set & _WEEKDAY_WORDS:
|
||||
score += 30
|
||||
if token_set & _PERIOD_WORDS:
|
||||
score += 20
|
||||
return score
|
||||
|
||||
|
||||
class TemporalConstraint(BaseModel):
|
||||
"""
|
||||
@@ -164,20 +262,23 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
if not results:
|
||||
return QueryAnalysis(temporal_constraint=None)
|
||||
|
||||
# Filter out false positives (common words parsed as dates)
|
||||
false_positives = {"do", "may", "march", "will", "can", "sat", "sun", "mon", "tue", "wed", "thu", "fri"}
|
||||
valid_results = [
|
||||
(text, date)
|
||||
# Score each match by the date signal it carries and keep only those
|
||||
# with a real signal, rejecting bare weekday/month-abbreviation false
|
||||
# positives ("we"/"me"/"did"). Prefer the strongest match, breaking ties
|
||||
# by longest span, so an explicit date ("in May", "2026-06-10") always
|
||||
# beats an earlier weak word regardless of position. See issue #2768.
|
||||
scored_results = [
|
||||
(_date_match_score(text), len(text), date)
|
||||
for text, date in results
|
||||
if (text.lower() not in false_positives or len(text) > 3)
|
||||
and not is_embedded_cjk_dateparser_match(query, text)
|
||||
if not is_embedded_cjk_dateparser_match(query, text)
|
||||
]
|
||||
scored_results = [entry for entry in scored_results if entry[0] > 0]
|
||||
|
||||
if not valid_results:
|
||||
if not scored_results:
|
||||
return QueryAnalysis(temporal_constraint=None)
|
||||
|
||||
# Use the first valid date found
|
||||
_, parsed_date = valid_results[0]
|
||||
# Highest signal score wins; ties broken by the longest matched span.
|
||||
_, _, parsed_date = max(scored_results, key=lambda entry: (entry[0], entry[1]))
|
||||
|
||||
# Create constraint for single day
|
||||
start_date = parsed_date.replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
|
||||
@@ -15,7 +15,8 @@ import time
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable
|
||||
|
||||
from ...config import get_config
|
||||
from .models import DirectiveInfo, LLMCall, ReflectAgentResult, TokenUsageSummary, ToolCall
|
||||
from ..llm_interface import LLM_TOOL_CHOICE_AUTO, LLMToolChoice
|
||||
from .models import DirectiveInfo, LLMCall, ReflectAgentResult, StructuredOutputResult, TokenUsageSummary, ToolCall
|
||||
from .prompts import (
|
||||
_extract_directive_rules,
|
||||
build_final_prompt,
|
||||
@@ -49,6 +50,11 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_MAX_ITERATIONS = 10
|
||||
|
||||
# Fallback answer when the LLM returns nothing usable. Consumers that need to
|
||||
# tell a real answer from this placeholder (e.g. refresh outcome metadata's
|
||||
# populated_content) compare against this constant rather than the literal.
|
||||
NO_ANSWER_TEXT = "No answer provided."
|
||||
|
||||
|
||||
def _normalize_tool_name(name: str) -> str:
|
||||
"""Normalize tool name from various LLM output formats.
|
||||
@@ -90,12 +96,87 @@ _LEAKED_JSON_SUFFIX = re.compile(
|
||||
r'\s*```(?:json)?\s*\{[^}]*(?:"(?:observation_ids|memory_ids|mental_model_ids)"|\})\s*```\s*$',
|
||||
re.DOTALL | re.IGNORECASE,
|
||||
)
|
||||
_LEAKED_JSON_OBJECT = re.compile(
|
||||
r'\s*\{[^{]*"(?:observation_ids|memory_ids|mental_model_ids|answer)"[^}]*\}\s*$', re.DOTALL
|
||||
)
|
||||
_TRAILING_IDS_PATTERN = re.compile(
|
||||
r"\s*(?:observation_ids|memory_ids|mental_model_ids)\s*[=:]\s*\[.*?\]\s*$", re.DOTALL | re.IGNORECASE
|
||||
)
|
||||
_JSON_CODE_FENCE_PATTERN = re.compile(r"^\s*```(?:json)?\s*(\{.*\})\s*```\s*$", re.DOTALL | re.IGNORECASE)
|
||||
|
||||
_DONE_ARGUMENT_KEYS = frozenset(
|
||||
{
|
||||
"answer",
|
||||
"directive_compliance",
|
||||
"memory_ids",
|
||||
"mental_model_ids",
|
||||
"observation_ids",
|
||||
"model_ids",
|
||||
}
|
||||
)
|
||||
_DONE_ARGUMENT_MARKER_KEYS = _DONE_ARGUMENT_KEYS - {"answer"}
|
||||
_LEAKED_JSON_ID_KEYS = frozenset({"memory_ids", "mental_model_ids", "observation_ids", "model_ids"})
|
||||
|
||||
|
||||
def _unwrap_leaked_done_arguments(text: str) -> str | None:
|
||||
"""Return the answer when a done tool call was rendered as JSON text.
|
||||
|
||||
Some providers leak the done tool's argument object instead of surfacing it
|
||||
as a native tool call, e.g. {"answer": "...", "memory_ids": [...]}. Only
|
||||
unwrap objects that match the done argument shape so normal JSON answers
|
||||
stay intact.
|
||||
"""
|
||||
candidate = text.strip()
|
||||
if not candidate:
|
||||
return None
|
||||
|
||||
fenced = _JSON_CODE_FENCE_PATTERN.match(candidate)
|
||||
if fenced:
|
||||
candidate = fenced.group(1).strip()
|
||||
|
||||
try:
|
||||
payload = json.loads(candidate)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
return None
|
||||
answer = payload.get("answer")
|
||||
if not isinstance(answer, str) or not answer.strip():
|
||||
return None
|
||||
|
||||
keys = set(payload)
|
||||
if not keys.intersection(_DONE_ARGUMENT_MARKER_KEYS):
|
||||
return None
|
||||
if not keys.issubset(_DONE_ARGUMENT_KEYS):
|
||||
return None
|
||||
|
||||
for key in ("memory_ids", "mental_model_ids", "observation_ids", "model_ids"):
|
||||
value = payload.get(key)
|
||||
if value is not None and not isinstance(value, list):
|
||||
return None
|
||||
|
||||
return answer.strip()
|
||||
|
||||
|
||||
def _strip_trailing_id_json_object(text: str) -> str:
|
||||
stripped = text.rstrip()
|
||||
if not stripped.endswith("}"):
|
||||
return text.strip()
|
||||
|
||||
start = stripped.rfind("{")
|
||||
if start < 0:
|
||||
return text.strip()
|
||||
|
||||
try:
|
||||
payload = json.loads(stripped[start:])
|
||||
except json.JSONDecodeError:
|
||||
return text.strip()
|
||||
|
||||
if not isinstance(payload, dict) or not payload:
|
||||
return text.strip()
|
||||
keys = set(payload)
|
||||
if not keys.issubset(_LEAKED_JSON_ID_KEYS):
|
||||
return text.strip()
|
||||
|
||||
return stripped[:start].strip()
|
||||
|
||||
|
||||
def _clean_answer_text(text: str) -> str:
|
||||
@@ -104,6 +185,10 @@ def _clean_answer_text(text: str) -> str:
|
||||
Some LLMs output the done() call as text instead of a proper tool call.
|
||||
This strips out patterns like: done({"answer": "...", ...})
|
||||
"""
|
||||
unwrapped = _unwrap_leaked_done_arguments(text)
|
||||
if unwrapped is not None:
|
||||
return unwrapped
|
||||
|
||||
# Remove done() call pattern from the end of the text
|
||||
cleaned = _DONE_CALL_PATTERN.sub("", text).strip()
|
||||
return cleaned if cleaned else text
|
||||
@@ -122,13 +207,17 @@ def _clean_done_answer(text: str) -> str:
|
||||
if not text:
|
||||
return text
|
||||
|
||||
unwrapped = _unwrap_leaked_done_arguments(text)
|
||||
if unwrapped is not None:
|
||||
return unwrapped
|
||||
|
||||
cleaned = text
|
||||
|
||||
# Remove leaked JSON in code blocks at the end
|
||||
cleaned = _LEAKED_JSON_SUFFIX.sub("", cleaned).strip()
|
||||
|
||||
# Remove leaked raw JSON objects at the end
|
||||
cleaned = _LEAKED_JSON_OBJECT.sub("", cleaned).strip()
|
||||
cleaned = _strip_trailing_id_json_object(cleaned)
|
||||
|
||||
# Remove trailing ID patterns
|
||||
cleaned = _TRAILING_IDS_PATTERN.sub("", cleaned).strip()
|
||||
@@ -141,7 +230,8 @@ async def _generate_structured_output(
|
||||
response_schema: dict,
|
||||
llm_config: "LLMProvider",
|
||||
reflect_id: str,
|
||||
) -> tuple[dict[str, Any] | None, int, int]:
|
||||
max_tokens: int | None = None,
|
||||
) -> StructuredOutputResult:
|
||||
"""Generate structured output from an answer using the provided JSON schema.
|
||||
|
||||
Args:
|
||||
@@ -149,10 +239,14 @@ async def _generate_structured_output(
|
||||
response_schema: JSON Schema for the expected output structure
|
||||
llm_config: LLM provider for making the extraction call
|
||||
reflect_id: Reflect ID for logging
|
||||
max_tokens: Output-token budget for the extraction call, mirroring the
|
||||
plain reflect calls (omitted when None); without it, reasoning /
|
||||
preamble models can exhaust the provider default before emitting any
|
||||
JSON (finish_reason=length, empty content -> issue #2431)
|
||||
|
||||
Returns:
|
||||
Tuple of (structured_output, input_tokens, output_tokens).
|
||||
structured_output is None if generation fails.
|
||||
A StructuredOutputResult carrying the structured output (None if
|
||||
generation fails) and the call's token usage.
|
||||
"""
|
||||
try:
|
||||
from typing import Any as TypingAny
|
||||
@@ -186,7 +280,7 @@ async def _generate_structured_output(
|
||||
|
||||
if not fields:
|
||||
logger.warning(f"[REFLECT {reflect_id}] No fields found in response_schema, skipping structured output")
|
||||
return None, 0, 0
|
||||
return StructuredOutputResult()
|
||||
|
||||
DynamicModel = create_model("StructuredResponse", **fields)
|
||||
|
||||
@@ -239,6 +333,11 @@ OUTPUT:"""
|
||||
],
|
||||
response_format=DynamicModel,
|
||||
scope="reflect_structured",
|
||||
strict_schema=get_config().llm_strict_schema_reflect,
|
||||
max_completion_tokens=max_tokens,
|
||||
max_retries=1,
|
||||
initial_backoff=0.25,
|
||||
max_backoff=1.0,
|
||||
skip_validation=True, # We'll handle the dict ourselves
|
||||
return_usage=True,
|
||||
)
|
||||
@@ -259,11 +358,17 @@ OUTPUT:"""
|
||||
logger.warning(f"[REFLECT {reflect_id}] Required field '{field_name}' is empty in structured output")
|
||||
|
||||
logger.info(f"[REFLECT {reflect_id}] Generated structured output with {len(structured_output)} fields")
|
||||
return structured_output, usage.input_tokens, usage.output_tokens
|
||||
return StructuredOutputResult(
|
||||
structured_output=structured_output,
|
||||
input_tokens=usage.input_tokens,
|
||||
output_tokens=usage.output_tokens,
|
||||
cached_tokens=usage.cached_tokens,
|
||||
thoughts_tokens=usage.thoughts_tokens,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"[REFLECT {reflect_id}] Failed to generate structured output: {e}")
|
||||
return None, 0, 0
|
||||
return StructuredOutputResult()
|
||||
|
||||
|
||||
def _count_messages_tokens(messages: list[dict[str, Any]]) -> int:
|
||||
@@ -321,7 +426,104 @@ def _all_mental_models_are_usable_and_fresh(tool_output: dict[str, Any]) -> bool
|
||||
return True
|
||||
|
||||
|
||||
# Detached cache-teardown tasks. asyncio holds only weak references to tasks, so
|
||||
# a fire-and-forget task can be garbage-collected mid-flight — keep a strong
|
||||
# reference here until it finishes.
|
||||
_cache_cleanup_tasks: set[asyncio.Task] = set()
|
||||
|
||||
|
||||
def _spawn_cache_cleanup(
|
||||
provider_impl: Any,
|
||||
session_id: str,
|
||||
cache_tasks: list[asyncio.Task],
|
||||
reflect_id: str,
|
||||
) -> None:
|
||||
"""Delete a reflect's ephemeral context caches in the background.
|
||||
|
||||
The per-reflect caches are dead the moment the reflect returns — nothing ever
|
||||
reuses them — so the caller must not wait on teardown: draining the in-flight
|
||||
create plus the delete round-trips would add latency to every single answer.
|
||||
Detach it instead. The short cache TTL is the backstop if the process dies
|
||||
before the task runs.
|
||||
"""
|
||||
|
||||
async def _cleanup() -> None:
|
||||
try:
|
||||
# Let any overlapped create land first, so its cache is registered in
|
||||
# the session and actually gets deleted rather than lingering to TTL.
|
||||
if cache_tasks:
|
||||
await asyncio.gather(*cache_tasks, return_exceptions=True)
|
||||
await provider_impl.delete_cache_session(session_id)
|
||||
except Exception:
|
||||
logger.debug("[REFLECT %s] cache session teardown failed (will age out on TTL)", reflect_id)
|
||||
|
||||
try:
|
||||
task = asyncio.create_task(_cleanup())
|
||||
except RuntimeError:
|
||||
# No running loop to detach onto (not expected in the server); TTL cleans up.
|
||||
return
|
||||
_cache_cleanup_tasks.add(task)
|
||||
task.add_done_callback(_cache_cleanup_tasks.discard)
|
||||
|
||||
|
||||
async def run_reflect_agent(
|
||||
llm_config: "LLMProvider",
|
||||
bank_id: str,
|
||||
query: str,
|
||||
bank_profile: dict[str, Any],
|
||||
search_mental_models_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
search_observations_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
recall_fn: Callable[[str, int, int], Awaitable[dict[str, Any]]],
|
||||
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
|
||||
**kwargs: Any,
|
||||
) -> ReflectAgentResult:
|
||||
"""Public entrypoint: runs the agent loop and tears down any per-step context
|
||||
caches it created.
|
||||
|
||||
The step-by-step caches (Gemini ``CachedContent``) are ephemeral — scoped to
|
||||
exactly one reflect and never reused after it — so teardown is scheduled on
|
||||
every exit path (answer, error, cancellation) but runs **detached**: the
|
||||
caller gets its answer without waiting on the delete round-trips. The short
|
||||
cache TTL is the backstop if the teardown never runs; the delete is
|
||||
best-effort and never allowed to fail a reflect.
|
||||
"""
|
||||
reflect_id = f"{bank_id[:8]}-{int(time.time() * 1000) % 100000}"
|
||||
provider_impl = getattr(llm_config, "_provider_impl", None)
|
||||
# Reflect step-by-step caching needs the provider to support it AND the
|
||||
# dedicated reflect flag (on by default; distinct from the global prompt-cache
|
||||
# switch so it can be turned off for reflect alone).
|
||||
incremental_caching = (
|
||||
provider_impl is not None
|
||||
and provider_impl.supports_incremental_prompt_cache()
|
||||
and get_config().reflect_prompt_cache_enabled
|
||||
)
|
||||
cache_session_id = f"reflect:{reflect_id}"
|
||||
# In-flight cache-create tasks (scheduled to overlap tool execution). Awaited
|
||||
# before teardown so every created cache is tracked and deleted — no orphans.
|
||||
cache_tasks: list[asyncio.Task] = []
|
||||
try:
|
||||
return await _run_reflect_agent_inner(
|
||||
llm_config,
|
||||
bank_id,
|
||||
query,
|
||||
bank_profile,
|
||||
search_mental_models_fn,
|
||||
search_observations_fn,
|
||||
recall_fn,
|
||||
expand_fn,
|
||||
reflect_id=reflect_id,
|
||||
provider_impl=provider_impl,
|
||||
incremental_caching=incremental_caching,
|
||||
cache_session_id=cache_session_id,
|
||||
cache_tasks=cache_tasks,
|
||||
**kwargs,
|
||||
)
|
||||
finally:
|
||||
if incremental_caching and provider_impl is not None:
|
||||
_spawn_cache_cleanup(provider_impl, cache_session_id, cache_tasks, reflect_id)
|
||||
|
||||
|
||||
async def _run_reflect_agent_inner(
|
||||
llm_config: "LLMProvider",
|
||||
bank_id: str,
|
||||
query: str,
|
||||
@@ -342,6 +544,13 @@ async def run_reflect_agent(
|
||||
max_context_tokens: int = 100_000,
|
||||
llm_output_language: str | None = None,
|
||||
cancel_check: Callable[[], None] | None = None,
|
||||
store_document_text: bool = True,
|
||||
*,
|
||||
reflect_id: str,
|
||||
provider_impl: Any,
|
||||
incremental_caching: bool,
|
||||
cache_session_id: str,
|
||||
cache_tasks: list[asyncio.Task],
|
||||
) -> ReflectAgentResult:
|
||||
"""
|
||||
Execute the reflect agent loop using native tool calling.
|
||||
@@ -369,7 +578,6 @@ async def run_reflect_agent(
|
||||
Returns:
|
||||
ReflectAgentResult with final answer and metadata
|
||||
"""
|
||||
reflect_id = f"{bank_id[:8]}-{int(time.time() * 1000) % 100000}"
|
||||
start_time = time.time()
|
||||
|
||||
# Build directives_applied for the trace
|
||||
@@ -380,8 +588,8 @@ async def run_reflect_agent(
|
||||
|
||||
# Get tools for this agent (with directive compliance field if directives exist).
|
||||
# The expand tool only reads back raw source text (chunks/documents), so it is
|
||||
# useless and excluded when document text storage is disabled.
|
||||
include_expand = get_config().store_document_text
|
||||
# useless and excluded when document text storage is disabled (per bank).
|
||||
include_expand = store_document_text
|
||||
tools = get_reflect_tools(
|
||||
directive_rules=directive_rules,
|
||||
include_mental_models=has_mental_models,
|
||||
@@ -406,27 +614,68 @@ async def run_reflect_agent(
|
||||
{"role": "user", "content": query},
|
||||
]
|
||||
|
||||
# Opt into context caching for the agentic tool loop. The system
|
||||
# prompt and tool definitions are stable for the duration of this
|
||||
# reflect call (and across reflects against the same bank), so
|
||||
# caching them once and reusing across every iteration of the loop
|
||||
# collapses the dominant input cost — the prefix repeated on every
|
||||
# turn. ``get_or_create_cached_prefix`` returns None when caching is
|
||||
# disabled, unsupported, or the prefix is too small; the
|
||||
# ``call_with_tools`` invocation below transparently falls back to
|
||||
# the uncached path in that case.
|
||||
cached_prefix_name: str | None = None
|
||||
provider_impl = getattr(llm_config, "_provider_impl", None)
|
||||
if provider_impl is not None and provider_impl.supports_prompt_caching():
|
||||
# Step-by-step context caching for the agentic tool loop.
|
||||
#
|
||||
# Caching only the static system+tools prefix wins little here: it's dwarfed
|
||||
# by the tool results (recall/observations) that get re-sent on every turn.
|
||||
# Instead we roll a cache forward one step at a time — after each turn the
|
||||
# cache is extended to cover that turn's FULL input, so the next ``auto`` turn
|
||||
# reuses the entire prior conversation at the cached rate and sends only its
|
||||
# own new tool results as the delta. Each new tool payload is therefore billed
|
||||
# at full price exactly once (the turn it's produced), then cached thereafter.
|
||||
#
|
||||
# The cache create for turn N+1 covers turn N's input, which is fully known the
|
||||
# moment turn N's LLM call returns — so we kick it off as a background task that
|
||||
# runs CONCURRENTLY with turn N's tool execution (``_schedule_cache``) and only
|
||||
# await it (``_resolve_pending_cache``) right before the next ``auto`` call,
|
||||
# hiding the create latency behind work we'd do anyway.
|
||||
#
|
||||
# ``rolling_cache_boundary`` is the number of leading ``messages`` baked into
|
||||
# the adopted ``rolling_cache_name``. ``incremental_caching`` is False for
|
||||
# providers/config without explicit caching, so every branch below is a no-op.
|
||||
rolling_cache_name: str | None = None
|
||||
rolling_cache_boundary = 0
|
||||
pending_cache_task: asyncio.Task | None = None
|
||||
pending_cache_boundary = 0
|
||||
|
||||
async def _resolve_pending_cache() -> None:
|
||||
"""Adopt the overlapped next-cache once it's ready as the rolling cache.
|
||||
|
||||
Best-effort: a failed/``None`` create just leaves the previous (smaller)
|
||||
cache in place, so the next call sends a larger delta but stays correct.
|
||||
"""
|
||||
nonlocal rolling_cache_name, rolling_cache_boundary, pending_cache_task
|
||||
if pending_cache_task is None:
|
||||
return
|
||||
task = pending_cache_task
|
||||
pending_cache_task = None
|
||||
try:
|
||||
cached_prefix_name = await provider_impl.get_or_create_cached_prefix(
|
||||
system_instruction=system_prompt,
|
||||
tools=tools,
|
||||
new_name = await task
|
||||
except Exception:
|
||||
new_name = None
|
||||
if new_name is not None:
|
||||
rolling_cache_name = new_name
|
||||
rolling_cache_boundary = pending_cache_boundary
|
||||
|
||||
def _schedule_cache(upto: int) -> None:
|
||||
"""Start building the cache covering ``messages[:upto]`` in the background
|
||||
so it overlaps the tool execution that follows this turn."""
|
||||
nonlocal pending_cache_task, pending_cache_boundary
|
||||
# ``messages[:upto]`` is snapshotted now, so appends during tool execution
|
||||
# can't change what gets cached. ``ensure_future`` raises if the provider
|
||||
# didn't return a coroutine (e.g. a test double) — caching is a soft
|
||||
# optimisation and must never break a reflect, so swallow and skip.
|
||||
try:
|
||||
task = asyncio.ensure_future(
|
||||
provider_impl.create_incremental_cache(
|
||||
session_id=cache_session_id, messages=messages[:upto], tools=tools
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
# Caching is a soft optimisation; never let a cache-side
|
||||
# error block a reflect.
|
||||
cached_prefix_name = None
|
||||
return
|
||||
pending_cache_boundary = upto
|
||||
pending_cache_task = task
|
||||
cache_tasks.append(task)
|
||||
|
||||
# Tracking
|
||||
total_tools_called = 0
|
||||
@@ -435,9 +684,14 @@ async def run_reflect_agent(
|
||||
llm_trace: list[dict[str, Any]] = []
|
||||
context_history: list[dict[str, Any]] = [] # For final prompt fallback
|
||||
|
||||
# Token usage tracking - accumulate across all LLM calls
|
||||
# Token usage tracking - accumulate across all LLM calls.
|
||||
# cached_tokens and thoughts_tokens are surfaced for cost attribution
|
||||
# and prompt-cache tuning. Both are subsets of (or parallel to) the
|
||||
# input/output counts and are NOT double-counted in total_tokens.
|
||||
total_input_tokens = 0
|
||||
total_output_tokens = 0
|
||||
total_cached_tokens = 0
|
||||
total_thoughts_tokens = 0
|
||||
|
||||
# Track available IDs for validation (prevents hallucinated citations)
|
||||
available_memory_ids: set[str] = set()
|
||||
@@ -460,6 +714,8 @@ async def run_reflect_agent(
|
||||
input_tokens=total_input_tokens,
|
||||
output_tokens=total_output_tokens,
|
||||
total_tokens=total_input_tokens + total_output_tokens,
|
||||
cached_tokens=total_cached_tokens,
|
||||
thoughts_tokens=total_thoughts_tokens,
|
||||
)
|
||||
|
||||
def _log_completion(answer: str, iterations: int, forced: bool = False):
|
||||
@@ -526,6 +782,8 @@ async def run_reflect_agent(
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
total_cached_tokens += getattr(usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
@@ -539,11 +797,12 @@ async def run_reflect_agent(
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id, max_tokens)
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
@@ -588,6 +847,8 @@ async def run_reflect_agent(
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
total_cached_tokens += getattr(usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
@@ -600,11 +861,12 @@ async def run_reflect_agent(
|
||||
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id, max_tokens)
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
@@ -634,12 +896,35 @@ async def run_reflect_agent(
|
||||
|
||||
if stop_forcing_from_iteration is not None and iteration >= stop_forcing_from_iteration:
|
||||
# A fresh mental model already short-circuited the forced path.
|
||||
iter_tool_choice: str | dict = "auto"
|
||||
iter_tool_choice = LLM_TOOL_CHOICE_AUTO
|
||||
elif iteration < len(forced_sequence):
|
||||
iter_tool_choice = {"type": "function", "function": {"name": forced_sequence[iteration]}}
|
||||
iter_tool_choice = LLMToolChoice.named(forced_sequence[iteration])
|
||||
else:
|
||||
iter_tool_choice = "auto"
|
||||
iter_tool_choice = LLM_TOOL_CHOICE_AUTO
|
||||
|
||||
# Will the NEXT turn be an ``auto`` turn (the only kind that references a
|
||||
# cache)? The cache we schedule this turn covers this turn's input and is
|
||||
# used by the next turn, so we only bother building it when the next turn
|
||||
# can use it — skipping the wasted creates between two forced turns.
|
||||
next_iter = iteration + 1
|
||||
if stop_forcing_from_iteration is not None and next_iter >= stop_forcing_from_iteration:
|
||||
next_is_auto = True
|
||||
elif next_iter < len(forced_sequence):
|
||||
next_is_auto = False
|
||||
else:
|
||||
next_is_auto = True
|
||||
|
||||
# Before an ``auto`` turn, adopt the cache that was being built in the
|
||||
# background during the previous turn's tool execution. It covers that
|
||||
# turn's full input, so THIS call reuses the entire prior conversation at
|
||||
# the cached rate and sends only the turns appended since. Forced turns
|
||||
# can't use a cache (Gemini rejects ``cached_content`` + ``tool_config``),
|
||||
# but the cache still advances underneath them, so the first ``auto`` turn
|
||||
# inherits a cache covering all the forced results.
|
||||
if incremental_caching and iter_tool_choice is LLM_TOOL_CHOICE_AUTO:
|
||||
await _resolve_pending_cache()
|
||||
|
||||
call_msg_count = len(messages)
|
||||
try:
|
||||
ct_kwargs: dict[str, Any] = dict(
|
||||
messages=messages,
|
||||
@@ -647,20 +932,16 @@ async def run_reflect_agent(
|
||||
scope="reflect_tool_call",
|
||||
tool_choice=iter_tool_choice,
|
||||
)
|
||||
# Gemini rejects ``cached_content`` alongside a per-request
|
||||
# ``tool_config`` (forced tool choice): "CachedContent can not be used
|
||||
# with GenerateContent request setting system_instruction, tools or
|
||||
# tool_config." The forced-sequence iterations set tool_config, so only
|
||||
# the ``auto`` iterations can reference the cache; forced iterations send
|
||||
# the prefix inline. The cache (tools + system prompt) is identical
|
||||
# either way, so this just limits *which* iterations are billed cached.
|
||||
if cached_prefix_name is not None and iter_tool_choice == "auto":
|
||||
ct_kwargs["cached_prefix"] = cached_prefix_name
|
||||
if incremental_caching and iter_tool_choice is LLM_TOOL_CHOICE_AUTO and rolling_cache_name is not None:
|
||||
ct_kwargs["cached_prefix"] = rolling_cache_name
|
||||
ct_kwargs["cached_prefix_message_count"] = rolling_cache_boundary
|
||||
result = await llm_config.call_with_tools(**ct_kwargs)
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
consecutive_errors = 0
|
||||
total_input_tokens += result.input_tokens
|
||||
total_output_tokens += result.output_tokens
|
||||
total_cached_tokens += getattr(result, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(result, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": f"agent_{iteration + 1}",
|
||||
@@ -709,6 +990,8 @@ async def run_reflect_agent(
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
total_cached_tokens += getattr(usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
@@ -722,11 +1005,12 @@ async def run_reflect_agent(
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id, max_tokens)
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
@@ -783,6 +1067,8 @@ async def run_reflect_agent(
|
||||
)
|
||||
total_input_tokens += rewrite_usage.input_tokens
|
||||
total_output_tokens += rewrite_usage.output_tokens
|
||||
total_cached_tokens += getattr(rewrite_usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(rewrite_usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final_rewrite",
|
||||
@@ -796,11 +1082,14 @@ async def run_reflect_agent(
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
struct = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id, max_tokens
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
|
||||
_log_completion(answer, iteration + 1)
|
||||
return ReflectAgentResult(
|
||||
@@ -835,6 +1124,8 @@ async def run_reflect_agent(
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
total_cached_tokens += getattr(usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
@@ -848,11 +1139,12 @@ async def run_reflect_agent(
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id, max_tokens)
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
@@ -885,7 +1177,6 @@ async def run_reflect_agent(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": done_call.id,
|
||||
"name": done_call.name, # Required by Gemini
|
||||
"content": json.dumps(
|
||||
{
|
||||
"error": "You must search for information first. Use search_mental_models(), search_observations(), or recall() before providing your final answer."
|
||||
@@ -919,6 +1210,7 @@ async def run_reflect_agent(
|
||||
directives_applied=directives_applied,
|
||||
llm_config=llm_config,
|
||||
response_schema=response_schema,
|
||||
max_tokens=max_tokens,
|
||||
)
|
||||
|
||||
# Execute other tools in parallel (exclude done tool in all its format variants)
|
||||
@@ -950,7 +1242,6 @@ async def run_reflect_agent(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tc.id,
|
||||
"name": tc.name,
|
||||
"content": json.dumps(
|
||||
{
|
||||
"error": f"Tool '{_normalize_tool_name(tc.name)}' is not available. Use only the tools provided to you."
|
||||
@@ -962,6 +1253,16 @@ async def run_reflect_agent(
|
||||
|
||||
other_tools = allowed_tools
|
||||
|
||||
# Kick off the next-turn cache (covering THIS call's input) so it
|
||||
# builds concurrently with the tool execution below — hiding the
|
||||
# create latency. Only schedule when the next turn is ``auto`` (the
|
||||
# only kind that references it); the next turn's pre-call resolve then
|
||||
# adopts it. Resolve any prior in-flight create first so we don't drop
|
||||
# its handle.
|
||||
if incremental_caching and next_is_auto:
|
||||
await _resolve_pending_cache()
|
||||
_schedule_cache(call_msg_count)
|
||||
|
||||
# Execute tools in parallel
|
||||
tool_tasks = [
|
||||
_execute_tool_with_timing(
|
||||
@@ -1044,7 +1345,6 @@ async def run_reflect_agent(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tc.id,
|
||||
"name": tc.name, # Required by Gemini
|
||||
"content": json.dumps(output, default=str, ensure_ascii=False),
|
||||
}
|
||||
)
|
||||
@@ -1128,6 +1428,7 @@ async def _process_done_tool(
|
||||
directives_applied: list[DirectiveInfo],
|
||||
llm_config: "LLMProvider | None" = None,
|
||||
response_schema: dict | None = None,
|
||||
max_tokens: int | None = None,
|
||||
) -> ReflectAgentResult:
|
||||
"""Process the done tool call and return the result."""
|
||||
args = done_call.arguments
|
||||
@@ -1136,7 +1437,46 @@ async def _process_done_tool(
|
||||
raw_answer = args.get("answer", "").strip()
|
||||
answer = _clean_done_answer(raw_answer) if raw_answer else ""
|
||||
if not answer:
|
||||
answer = "No answer provided."
|
||||
answer = NO_ANSWER_TEXT
|
||||
|
||||
final_usage = usage
|
||||
if llm_config and max_tokens is not None and count_cl100k_tokens(answer) > max_tokens:
|
||||
rewrite_start = time.time()
|
||||
rewritten, rewrite_usage = await llm_config.call(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": (
|
||||
"Rewrite the user's text so it fits within the requested token budget. "
|
||||
"Preserve the key facts and structure; drop lower-priority detail. "
|
||||
"Respond with the rewritten text only, no preamble."
|
||||
),
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"Target budget: {max_tokens} tokens.\n\nText to rewrite:\n{answer}",
|
||||
},
|
||||
],
|
||||
scope="reflect",
|
||||
max_completion_tokens=max_tokens,
|
||||
return_usage=True,
|
||||
)
|
||||
answer = _clean_answer_text(rewritten.strip())
|
||||
final_usage = TokenUsageSummary(
|
||||
input_tokens=usage.input_tokens + rewrite_usage.input_tokens,
|
||||
output_tokens=usage.output_tokens + rewrite_usage.output_tokens,
|
||||
total_tokens=usage.total_tokens + rewrite_usage.input_tokens + rewrite_usage.output_tokens,
|
||||
cached_tokens=usage.cached_tokens + (getattr(rewrite_usage, "cached_tokens", 0) or 0),
|
||||
thoughts_tokens=usage.thoughts_tokens + (getattr(rewrite_usage, "thoughts_tokens", 0) or 0),
|
||||
)
|
||||
llm_trace.append(
|
||||
LLMCall(
|
||||
scope="final_rewrite",
|
||||
duration_ms=int((time.time() - rewrite_start) * 1000),
|
||||
input_tokens=rewrite_usage.input_tokens,
|
||||
output_tokens=rewrite_usage.output_tokens,
|
||||
)
|
||||
)
|
||||
|
||||
# Validate IDs (only include IDs that were actually retrieved)
|
||||
used_memory_ids = [mid for mid in (args.get("memory_ids") or []) if mid in available_memory_ids]
|
||||
@@ -1145,16 +1485,16 @@ async def _process_done_tool(
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
final_usage = usage
|
||||
if response_schema and llm_config and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id, max_tokens)
|
||||
structured_output = struct.structured_output
|
||||
# Add structured output tokens to usage
|
||||
final_usage = TokenUsageSummary(
|
||||
input_tokens=usage.input_tokens + struct_in,
|
||||
output_tokens=usage.output_tokens + struct_out,
|
||||
total_tokens=usage.total_tokens + struct_in + struct_out,
|
||||
input_tokens=final_usage.input_tokens + struct.input_tokens,
|
||||
output_tokens=final_usage.output_tokens + struct.output_tokens,
|
||||
total_tokens=final_usage.total_tokens + struct.input_tokens + struct.output_tokens,
|
||||
cached_tokens=final_usage.cached_tokens + struct.cached_tokens,
|
||||
thoughts_tokens=final_usage.thoughts_tokens + struct.thoughts_tokens,
|
||||
)
|
||||
|
||||
log_completion(answer, iterations)
|
||||
@@ -1269,22 +1609,35 @@ async def _execute_tool(
|
||||
query = args.get("query")
|
||||
if not query:
|
||||
return {"error": "search_mental_models requires a query parameter"}
|
||||
max_results = int(args.get("max_results") or 5)
|
||||
max_results, error = _parse_tool_int_arg_or_error(args, "max_results", default=5)
|
||||
if error:
|
||||
return {"error": error}
|
||||
return await search_mental_models_fn(query, max_results)
|
||||
|
||||
elif tool_name == "search_observations":
|
||||
query = args.get("query")
|
||||
if not query:
|
||||
return {"error": "search_observations requires a query parameter"}
|
||||
max_tokens = max(int(args.get("max_tokens") or 5000), 1000) # Default 5000, min 1000
|
||||
max_tokens, error = _parse_tool_int_arg_or_error(args, "max_tokens", default=5000, minimum=1000)
|
||||
if error:
|
||||
return {"error": error}
|
||||
return await search_observations_fn(query, max_tokens)
|
||||
|
||||
elif tool_name == "recall":
|
||||
query = args.get("query")
|
||||
if not query:
|
||||
return {"error": "recall requires a query parameter"}
|
||||
max_tokens = max(int(args.get("max_tokens") or 2048), 1000) # Default 2048, min 1000
|
||||
max_chunk_tokens = max(int(args.get("max_chunk_tokens") or 1000), 1000) # Always enabled, min 1000
|
||||
max_tokens, error = _parse_tool_int_arg_or_error(args, "max_tokens", default=2048, minimum=1000)
|
||||
if error:
|
||||
return {"error": error}
|
||||
max_chunk_tokens, error = _parse_tool_int_arg_or_error(
|
||||
args,
|
||||
"max_chunk_tokens",
|
||||
default=1000,
|
||||
minimum=1000,
|
||||
)
|
||||
if error:
|
||||
return {"error": error}
|
||||
return await recall_fn(query, max_tokens, max_chunk_tokens)
|
||||
|
||||
elif tool_name == "expand":
|
||||
@@ -1298,23 +1651,63 @@ async def _execute_tool(
|
||||
return {"error": f"Unknown tool: {tool_name}"}
|
||||
|
||||
|
||||
_NULLISH_TOOL_INT_STRINGS = {"", "none", "null"}
|
||||
|
||||
|
||||
def _parse_tool_int_arg(args: dict[str, Any], key: str, *, default: int, minimum: int | None = None) -> int:
|
||||
raw_value = args.get(key)
|
||||
if not raw_value:
|
||||
value = default
|
||||
elif isinstance(raw_value, str) and raw_value.strip().lower() in _NULLISH_TOOL_INT_STRINGS:
|
||||
value = default
|
||||
else:
|
||||
value = int(raw_value)
|
||||
if minimum is None:
|
||||
return value
|
||||
return max(value, minimum)
|
||||
|
||||
|
||||
def _parse_tool_int_arg_or_error(
|
||||
args: dict[str, Any],
|
||||
key: str,
|
||||
*,
|
||||
default: int,
|
||||
minimum: int | None = None,
|
||||
) -> tuple[int, str | None]:
|
||||
try:
|
||||
return _parse_tool_int_arg(args, key, default=default, minimum=minimum), None
|
||||
except (OverflowError, TypeError, ValueError):
|
||||
return default, f"{key} must be an integer or null-like value"
|
||||
|
||||
|
||||
def _summarize_tool_int_arg(args: dict[str, Any], key: str, *, default: int, minimum: int | None = None) -> str:
|
||||
try:
|
||||
return str(_parse_tool_int_arg(args, key, default=default, minimum=minimum))
|
||||
except (OverflowError, TypeError, ValueError):
|
||||
return f"invalid:{args.get(key)!r}"
|
||||
|
||||
|
||||
def _summarize_tool_query(args: dict[str, Any]) -> str:
|
||||
query = args.get("query") or ""
|
||||
if not isinstance(query, str):
|
||||
query = str(query)
|
||||
return f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
|
||||
|
||||
|
||||
def _summarize_input(tool_name: str, args: dict[str, Any]) -> str:
|
||||
"""Create a summary of tool input for logging, showing all params."""
|
||||
if tool_name == "search_mental_models":
|
||||
query = args.get("query", "")
|
||||
query_preview = f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
|
||||
max_results = int(args.get("max_results") or 5)
|
||||
query_preview = _summarize_tool_query(args)
|
||||
max_results = _summarize_tool_int_arg(args, "max_results", default=5)
|
||||
return f"(query={query_preview}, max_results={max_results})"
|
||||
elif tool_name == "search_observations":
|
||||
query = args.get("query", "")
|
||||
query_preview = f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
|
||||
max_tokens = max(int(args.get("max_tokens") or 5000), 1000)
|
||||
query_preview = _summarize_tool_query(args)
|
||||
max_tokens = _summarize_tool_int_arg(args, "max_tokens", default=5000, minimum=1000)
|
||||
return f"(query={query_preview}, max_tokens={max_tokens})"
|
||||
elif tool_name == "recall":
|
||||
query = args.get("query", "")
|
||||
query_preview = f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
|
||||
max_tokens = max(int(args.get("max_tokens") or 2048), 1000)
|
||||
max_chunk_tokens = max(int(args.get("max_chunk_tokens") or 1000), 1000)
|
||||
query_preview = _summarize_tool_query(args)
|
||||
max_tokens = _summarize_tool_int_arg(args, "max_tokens", default=2048, minimum=1000)
|
||||
max_chunk_tokens = _summarize_tool_int_arg(args, "max_chunk_tokens", default=1000, minimum=1000)
|
||||
return f"(query={query_preview}, max_tokens={max_tokens}, max_chunk_tokens={max_chunk_tokens})"
|
||||
elif tool_name == "expand":
|
||||
memory_ids = args.get("memory_ids", [])
|
||||
|
||||
@@ -78,9 +78,32 @@ class DirectiveInfo(BaseModel):
|
||||
class TokenUsageSummary(BaseModel):
|
||||
"""Total token usage across all LLM calls."""
|
||||
|
||||
input_tokens: int = Field(default=0, description="Total input tokens used")
|
||||
output_tokens: int = Field(default=0, description="Total output tokens used")
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output)")
|
||||
input_tokens: int = Field(default=0, description="Total input tokens used (includes any cached prefix tokens)")
|
||||
output_tokens: int = Field(default=0, description="Total visible output tokens used (excludes reasoning/thoughts)")
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output, excludes thoughts)")
|
||||
cached_tokens: int = Field(
|
||||
default=0,
|
||||
description="Cached/cache-read prompt tokens summed across calls. Subset of input_tokens.",
|
||||
)
|
||||
thoughts_tokens: int = Field(
|
||||
default=0,
|
||||
description=(
|
||||
"Reasoning/thinking tokens summed across calls. Billed at the output rate by some providers "
|
||||
"but not part of visible output."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class StructuredOutputResult(BaseModel):
|
||||
"""Result of structured-output generation, including token usage for the call."""
|
||||
|
||||
structured_output: dict[str, Any] | None = Field(
|
||||
default=None, description="Generated structured output, or None if generation failed"
|
||||
)
|
||||
input_tokens: int = Field(default=0, description="Input tokens used")
|
||||
output_tokens: int = Field(default=0, description="Visible output tokens used")
|
||||
cached_tokens: int = Field(default=0, description="Cached prefix tokens. Subset of input_tokens.")
|
||||
thoughts_tokens: int = Field(default=0, description="Reasoning/thinking tokens, when reported by the provider")
|
||||
|
||||
|
||||
class ReflectAgentResult(BaseModel):
|
||||
|
||||
@@ -177,20 +177,17 @@ _HEADING_RX = re.compile(r"^(#{1,6})\s+(.+?)\s*$")
|
||||
_BULLET_RX = re.compile(r"^\s*[-*+]\s+(.*)$")
|
||||
_ORDERED_RX = re.compile(r"^\s*\d+[.)]\s+(.*)$")
|
||||
_FENCE_RX = re.compile(r"^```([A-Za-z0-9_+-]*)\s*$")
|
||||
|
||||
|
||||
def _strip_separators(lines: list[str]) -> list[str]:
|
||||
"""Drop horizontal-rule lines (`---`, `***`) used as section separators.
|
||||
|
||||
Our renderer never emits these, but LLM output frequently includes them
|
||||
between sections; treating them as blank lines avoids parsing them as
|
||||
paragraphs.
|
||||
"""
|
||||
return ["" if re.fullmatch(r"\s*([-*_])\1{2,}\s*", line) else line for line in lines]
|
||||
_SEPARATOR_RX = re.compile(r"\s*([-*_])\1{2,}\s*")
|
||||
|
||||
|
||||
def _split_blocks(lines: list[str]) -> list[list[str]]:
|
||||
"""Group consecutive non-blank lines into block chunks."""
|
||||
"""Group consecutive non-blank lines into block chunks.
|
||||
|
||||
Horizontal-rule lines (`---`, `***`) count as blank. Our renderer never
|
||||
emits these, but LLM output frequently includes them between sections;
|
||||
treating them as blank avoids parsing them as paragraphs. Inside a fence
|
||||
they are code, not a separator, so they are kept verbatim.
|
||||
"""
|
||||
chunks: list[list[str]] = []
|
||||
current: list[str] = []
|
||||
in_fence = False
|
||||
@@ -202,7 +199,7 @@ def _split_blocks(lines: list[str]) -> list[list[str]]:
|
||||
if in_fence:
|
||||
current.append(line)
|
||||
continue
|
||||
if line.strip() == "":
|
||||
if line.strip() == "" or _SEPARATOR_RX.fullmatch(line):
|
||||
if current:
|
||||
chunks.append(current)
|
||||
current = []
|
||||
@@ -250,8 +247,7 @@ def parse_markdown(markdown: str) -> StructuredDocument:
|
||||
so we never silently drop user content. Section IDs are unique slugs of
|
||||
their headings.
|
||||
"""
|
||||
raw_lines = (markdown or "").splitlines()
|
||||
lines = _strip_separators(raw_lines)
|
||||
lines = (markdown or "").splitlines()
|
||||
|
||||
sections: list[Section] = []
|
||||
used_ids: set[str] = set()
|
||||
|
||||
@@ -328,18 +328,21 @@ async def tool_expand(
|
||||
if not memory_ids:
|
||||
return {"error": "memory_ids is required and must not be empty"}
|
||||
|
||||
# Validate and convert UUIDs
|
||||
valid_uuids: list[uuid.UUID] = []
|
||||
# Validate and convert UUIDs. Each id keeps a handle on its own UUID: a list of
|
||||
# only the valid ones no longer lines up with memory_ids once one id is invalid.
|
||||
uuid_by_id: dict[str, uuid.UUID] = {}
|
||||
errors: dict[str, str] = {}
|
||||
for mid in memory_ids:
|
||||
try:
|
||||
valid_uuids.append(uuid.UUID(mid))
|
||||
uuid_by_id[mid] = uuid.UUID(mid)
|
||||
except ValueError:
|
||||
errors[mid] = f"Invalid memory_id format: {mid}"
|
||||
|
||||
if not valid_uuids:
|
||||
if not uuid_by_id:
|
||||
return {"error": "No valid memory IDs provided", "details": errors}
|
||||
|
||||
valid_uuids = list(uuid_by_id.values())
|
||||
|
||||
# Batch fetch all memory units
|
||||
memories = await conn.fetch(
|
||||
f"""
|
||||
@@ -395,12 +398,12 @@ async def tool_expand(
|
||||
|
||||
# Build results
|
||||
results: list[dict[str, Any]] = []
|
||||
for mid, mem_uuid in zip(memory_ids, valid_uuids):
|
||||
for mid in memory_ids:
|
||||
if mid in errors:
|
||||
results.append({"memory_id": mid, "error": errors[mid]})
|
||||
continue
|
||||
|
||||
memory = memory_map.get(mem_uuid)
|
||||
memory = memory_map.get(uuid_by_id[mid])
|
||||
if not memory:
|
||||
results.append({"memory_id": mid, "error": f"Memory not found: {mid}"})
|
||||
continue
|
||||
|
||||
@@ -31,8 +31,20 @@ class LLMToolCallResult(BaseModel):
|
||||
content: str | None = Field(default=None, description="Text content if any")
|
||||
tool_calls: list[LLMToolCall] = Field(default_factory=list, description="Tool calls requested by the LLM")
|
||||
finish_reason: str | None = Field(default=None, description="Reason the LLM stopped: 'stop', 'tool_calls', etc.")
|
||||
input_tokens: int = Field(default=0, description="Input tokens used in this call")
|
||||
output_tokens: int = Field(default=0, description="Output tokens used in this call")
|
||||
input_tokens: int = Field(
|
||||
default=0,
|
||||
description="Input tokens used in this call (includes any cached prefix tokens reported by the provider)",
|
||||
)
|
||||
output_tokens: int = Field(
|
||||
default=0, description="Visible output tokens used in this call (excludes reasoning/thoughts)"
|
||||
)
|
||||
cached_tokens: int = Field(
|
||||
default=0, description="Cached prefix tokens, when reported by the provider. Subset of input_tokens."
|
||||
)
|
||||
thoughts_tokens: int = Field(
|
||||
default=0,
|
||||
description="Reasoning/thinking tokens. Billed at the output rate by some providers but not part of visible output.",
|
||||
)
|
||||
|
||||
|
||||
class ToolCallTrace(BaseModel):
|
||||
@@ -91,9 +103,18 @@ class TokenUsage(BaseModel):
|
||||
)
|
||||
|
||||
input_tokens: int = Field(default=0, description="Number of input/prompt tokens consumed")
|
||||
output_tokens: int = Field(default=0, description="Number of output/completion tokens generated")
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output)")
|
||||
output_tokens: int = Field(
|
||||
default=0, description="Number of visible output/completion tokens generated (excludes reasoning/thoughts)"
|
||||
)
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output, excludes thoughts)")
|
||||
cached_tokens: int = Field(default=0, description="Cached/cache-read prompt tokens, when reported by the provider")
|
||||
thoughts_tokens: int = Field(
|
||||
default=0,
|
||||
description=(
|
||||
"Reasoning/thinking tokens generated by the model. Billed at the output rate by some providers "
|
||||
"(e.g. Gemini 2.5+ family) but not surfaced in the visible response."
|
||||
),
|
||||
)
|
||||
|
||||
def __add__(self, other: "TokenUsage") -> "TokenUsage":
|
||||
"""Allow aggregating token usage from multiple calls."""
|
||||
@@ -102,6 +123,7 @@ class TokenUsage(BaseModel):
|
||||
output_tokens=self.output_tokens + other.output_tokens,
|
||||
total_tokens=self.total_tokens + other.total_tokens,
|
||||
cached_tokens=self.cached_tokens + other.cached_tokens,
|
||||
thoughts_tokens=self.thoughts_tokens + other.thoughts_tokens,
|
||||
)
|
||||
|
||||
|
||||
@@ -150,6 +172,47 @@ class DispositionTraits(BaseModel):
|
||||
model_config = ConfigDict(json_schema_extra={"example": {"skepticism": 3, "literalism": 3, "empathy": 3}})
|
||||
|
||||
|
||||
class RecallScores(BaseModel):
|
||||
"""Per-result recall scores from different stages of the pipeline.
|
||||
|
||||
``final`` is the value results are ranked by. The others are diagnostic and
|
||||
can be filtered on via the recall ``min_scores`` request parameter. ``semantic``
|
||||
and ``keyword`` are the raw per-strategy retrieval scores (``None`` when that
|
||||
strategy did not surface this result); ``reranker`` is the cross-encoder's
|
||||
normalized relevance.
|
||||
"""
|
||||
|
||||
final: float = Field(description="Final ranking score (combined reranker + recency/temporal/proof boosts)")
|
||||
reranker: float | None = Field(
|
||||
default=None,
|
||||
description="Cross-encoder relevance, normalized 0-1. None when the reranker is a passthrough (rrf/interleave modes).",
|
||||
)
|
||||
semantic: float | None = Field(
|
||||
default=None, description="Vector cosine similarity (0-1). None if this result was not surfaced semantically."
|
||||
)
|
||||
keyword: float | None = Field(
|
||||
default=None,
|
||||
description="Keyword/full-text (BM25) score (>= 0, unbounded). None if this result was not surfaced by keyword search.",
|
||||
)
|
||||
|
||||
|
||||
class MinScores(BaseModel):
|
||||
"""Optional per-stage score floors for recall (all inclusive, AND-ed).
|
||||
|
||||
``semantic`` and ``keyword`` are **retrieval-level** cutoffs pushed into the SQL
|
||||
arms (overriding the global ``semantic_min_similarity`` / ``bm25_min_score``
|
||||
config for this request), so they prune weak matches before fusion. ``reranker``
|
||||
and ``final`` are **post-query** filters applied to the scored results after
|
||||
reranking. Any field left None imposes no floor; all-None (the default) means
|
||||
no score filtering.
|
||||
"""
|
||||
|
||||
semantic: float | None = Field(default=None, description="Retrieval-level: minimum vector similarity (0-1).")
|
||||
keyword: float | None = Field(default=None, description="Retrieval-level: minimum keyword/full-text (BM25) score.")
|
||||
reranker: float | None = Field(default=None, description="Post-query: minimum normalized reranker score (0-1).")
|
||||
final: float | None = Field(default=None, description="Post-query: minimum final ranking score.")
|
||||
|
||||
|
||||
class MemoryFact(BaseModel):
|
||||
"""
|
||||
A single memory fact returned by search or think operations.
|
||||
@@ -180,7 +243,7 @@ class MemoryFact(BaseModel):
|
||||
|
||||
id: str = Field(description="Unique identifier for the memory fact")
|
||||
text: str = Field(description="The actual text content of the memory")
|
||||
fact_type: str = Field(description="Type of fact: 'world', 'experience', 'opinion', or 'observation'")
|
||||
fact_type: str = Field(description="Type of fact: 'world', 'experience', or 'observation'")
|
||||
entities: list[str] | None = Field(None, description="Entity names mentioned in this fact")
|
||||
context: str | None = Field(None, description="Additional context for the memory")
|
||||
occurred_start: str | None = Field(None, description="ISO format date when the event started occurring")
|
||||
@@ -192,13 +255,20 @@ class MemoryFact(BaseModel):
|
||||
@field_validator("metadata", mode="before")
|
||||
@classmethod
|
||||
def parse_metadata(cls, v: Any) -> dict[str, str] | None:
|
||||
"""Parse metadata from JSON string if needed (asyncpg may return JSONB as str)."""
|
||||
"""Parse metadata from JSON string if needed (asyncpg may return JSONB as str).
|
||||
|
||||
Also coerces non-string dict values (e.g., integer IDs stored in JSONB)
|
||||
to strings, preventing ValidationError when consolidation encounters
|
||||
metadata like {"original_id": 348} instead of {"original_id": "348"}.
|
||||
"""
|
||||
if v is None:
|
||||
return None
|
||||
if isinstance(v, str):
|
||||
import json
|
||||
|
||||
return json.loads(v)
|
||||
v = json.loads(v)
|
||||
if isinstance(v, dict):
|
||||
return {str(k): str(val) for k, val in v.items()}
|
||||
return v
|
||||
|
||||
chunk_id: str | None = Field(
|
||||
@@ -209,6 +279,10 @@ class MemoryFact(BaseModel):
|
||||
None,
|
||||
description="IDs of source facts this observation was derived from (observation type only, when source_facts is enabled)",
|
||||
)
|
||||
scores: RecallScores | None = Field(
|
||||
None,
|
||||
description="Recall scores from each pipeline stage (final/reranker/semantic/keyword). Not returned for source facts.",
|
||||
)
|
||||
|
||||
|
||||
class ChunkInfo(BaseModel):
|
||||
@@ -307,7 +381,8 @@ class ReflectResult(BaseModel):
|
||||
],
|
||||
"experience": [],
|
||||
"opinion": [],
|
||||
"mental_models": [],
|
||||
"observation": [],
|
||||
"mental-models": [],
|
||||
"directives": [
|
||||
{
|
||||
"id": "directive-123",
|
||||
@@ -324,7 +399,7 @@ class ReflectResult(BaseModel):
|
||||
|
||||
text: str = Field(description="The formulated answer text")
|
||||
based_on: dict[str, Any] = Field(
|
||||
description="Facts used to formulate the answer, organized by type (world, experience, mental_models, directives)"
|
||||
description="Facts used to formulate the answer, organized by type (world, experience, observation, mental-models, directives)"
|
||||
)
|
||||
structured_output: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
|
||||
@@ -12,7 +12,7 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from ..._vector_index import index_using_clause, uses_per_bank_vector_indexes
|
||||
from ...config import get_config
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..db_utils import acquire_with_retry, retry_with_backoff
|
||||
from ..memory_engine import fq_table, get_current_schema
|
||||
from ..response_models import DispositionTraits
|
||||
|
||||
@@ -188,8 +188,19 @@ async def get_or_create_bank_profile(pool, bank_id: str) -> BankProfileResult:
|
||||
or rolls back atomically with the caller's write), use
|
||||
``get_or_create_bank_profile_on_conn`` instead.
|
||||
"""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
return await get_or_create_bank_profile_on_conn(conn, bank_id, ops=pool.ops)
|
||||
|
||||
# A fresh bank builds its per-(bank, fact_type) partial vector indexes with
|
||||
# a plain CREATE INDEX (it must — this runs inside the bank-create tx, and
|
||||
# CONCURRENTLY cannot). That CREATE takes a ShareLock on the shared
|
||||
# memory_units table, which can deadlock with concurrent writers. The build
|
||||
# is idempotent (INSERT ... ON CONFLICT + CREATE INDEX IF NOT EXISTS), so a
|
||||
# transient deadlock (40P01 / ORA-00060) is safe to retry as a whole tx.
|
||||
async def _create() -> BankProfileResult:
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
return await get_or_create_bank_profile_on_conn(conn, bank_id, ops=pool.ops)
|
||||
|
||||
return await retry_with_backoff(_create)
|
||||
|
||||
|
||||
async def get_or_create_bank_profile_on_conn(conn, bank_id: str, *, ops) -> BankProfileResult:
|
||||
|
||||
@@ -8,7 +8,7 @@ import hashlib
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
from ...config import get_config
|
||||
from ...config import _get_raw_config
|
||||
from ..memory_engine import fq_table
|
||||
from .types import ChunkMetadata
|
||||
|
||||
@@ -64,14 +64,64 @@ async def delete_chunks_by_ids(conn, chunk_ids: list[str]) -> None:
|
||||
"""
|
||||
if not chunk_ids:
|
||||
return
|
||||
|
||||
# PostgreSQL's FK cascade deletes child memory_links in executor-chosen
|
||||
# order. Concurrent chunk deletes for the same bank can then lock overlapping
|
||||
# memory_links in opposite orders and deadlock. Delete links explicitly in a
|
||||
# total order before deleting chunks so every writer takes row locks the same
|
||||
# way; the FK cascade still handles anything inserted later in this txn.
|
||||
await conn.execute(
|
||||
f"DELETE FROM {fq_table('chunks')} WHERE chunk_id = ANY($1::text[])",
|
||||
f"""
|
||||
WITH target_units AS MATERIALIZED (
|
||||
SELECT id
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE chunk_id = ANY($1::text[])
|
||||
),
|
||||
ordered_links AS MATERIALIZED (
|
||||
SELECT ml.ctid
|
||||
FROM {fq_table("memory_links")} ml
|
||||
WHERE EXISTS (
|
||||
SELECT 1
|
||||
FROM target_units tu
|
||||
WHERE tu.id = ml.from_unit_id OR tu.id = ml.to_unit_id
|
||||
)
|
||||
ORDER BY
|
||||
LEAST(ml.from_unit_id, ml.to_unit_id),
|
||||
GREATEST(ml.from_unit_id, ml.to_unit_id),
|
||||
ml.link_type,
|
||||
COALESCE(ml.entity_id, '00000000-0000-0000-0000-000000000000'::uuid)
|
||||
FOR UPDATE OF ml
|
||||
)
|
||||
DELETE FROM {fq_table("memory_links")} ml
|
||||
USING ordered_links ol
|
||||
WHERE ml.ctid = ol.ctid
|
||||
""",
|
||||
chunk_ids,
|
||||
)
|
||||
await conn.execute(
|
||||
f"""
|
||||
WITH ordered_chunks AS MATERIALIZED (
|
||||
SELECT chunk_id
|
||||
FROM {fq_table("chunks")}
|
||||
WHERE chunk_id = ANY($1::text[])
|
||||
ORDER BY chunk_id
|
||||
FOR UPDATE
|
||||
)
|
||||
DELETE FROM {fq_table("chunks")} c
|
||||
USING ordered_chunks oc
|
||||
WHERE c.chunk_id = oc.chunk_id
|
||||
""",
|
||||
chunk_ids,
|
||||
)
|
||||
|
||||
|
||||
async def store_chunks_batch(
|
||||
conn, bank_id: str, document_id: str, chunks: list[ChunkMetadata], ops=None
|
||||
conn,
|
||||
bank_id: str,
|
||||
document_id: str,
|
||||
chunks: list[ChunkMetadata],
|
||||
ops=None,
|
||||
store_document_text: bool | None = None,
|
||||
) -> dict[int, str]:
|
||||
"""
|
||||
Store document chunks in the database.
|
||||
@@ -82,6 +132,9 @@ async def store_chunks_batch(
|
||||
document_id: Document identifier
|
||||
chunks: List of ChunkMetadata objects
|
||||
ops: DataAccessOps instance (from backend.ops)
|
||||
store_document_text: Whether to persist raw chunk text. When ``None``,
|
||||
falls back to the server-level default; callers on the retain path
|
||||
pass the per-bank resolved value.
|
||||
|
||||
Returns:
|
||||
Dictionary mapping global chunk index to chunk_id
|
||||
@@ -92,7 +145,9 @@ async def store_chunks_batch(
|
||||
# When document text storage is disabled, persist empty chunk_text (the
|
||||
# column is NOT NULL) while still computing content_hash from the real text
|
||||
# so delta-retain dedup is unaffected.
|
||||
store_text = get_config().store_document_text
|
||||
# Fallback to the raw global default (not get_config(), which guards
|
||||
# bank-configurable fields); the retain path always passes the resolved value.
|
||||
store_text = store_document_text if store_document_text is not None else _get_raw_config().store_document_text
|
||||
|
||||
# Prepare chunk data for batch insert
|
||||
chunk_ids = []
|
||||
|
||||
@@ -7,7 +7,7 @@ Handles entity extraction and resolution for stored facts.
|
||||
import logging
|
||||
|
||||
from . import link_utils
|
||||
from .types import ProcessedFact
|
||||
from .types import EntityResolutionResult, ProcessedFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -58,7 +58,7 @@ async def resolve_entities(
|
||||
log_buffer: list[str] = None,
|
||||
user_entities_per_content: dict[int, list[dict]] = None,
|
||||
entity_labels: list | None = None,
|
||||
) -> tuple[list[str], list[tuple], dict[str, list[str]]]:
|
||||
) -> EntityResolutionResult:
|
||||
"""
|
||||
Phase 1: Resolve entity names to canonical IDs (read-heavy).
|
||||
|
||||
@@ -76,10 +76,10 @@ async def resolve_entities(
|
||||
entity_labels: Optional entity label taxonomy
|
||||
|
||||
Returns:
|
||||
Tuple of (resolved_entity_ids, entity_to_unit, unit_to_entity_ids).
|
||||
EntityResolutionResult with the resolved identities and unit mappings.
|
||||
"""
|
||||
if not unit_ids or not facts:
|
||||
return [], [], {}
|
||||
return EntityResolutionResult(resolved_entities=[], entity_to_unit=[], unit_to_entity_ids={})
|
||||
|
||||
if len(unit_ids) != len(facts):
|
||||
raise ValueError(f"Mismatch between unit_ids ({len(unit_ids)}) and facts ({len(facts)})")
|
||||
|
||||
@@ -14,9 +14,11 @@ from typing import Any, Literal, cast
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, create_model, field_validator
|
||||
|
||||
from ..llm_wrapper import LLMConfig, OutputTooLongError, sanitize_llm_output
|
||||
from ..llm_interface import ProviderRateLimitResetError
|
||||
from ..llm_wrapper import LLMConfig, OutputTooLongError, parse_llm_json, sanitize_llm_output
|
||||
from ..operation_metadata import RetainExtractionErrors
|
||||
from ..response_models import TokenUsage
|
||||
from ..structured_output import strict_json_schema
|
||||
from .entity_labels import (
|
||||
EntityLabelsConfig,
|
||||
MapField,
|
||||
@@ -31,7 +33,7 @@ def _extract_map_entities(
|
||||
entity_obj: dict,
|
||||
fields: dict[str, MapField],
|
||||
prefix: str,
|
||||
validated_entities: "list[Entity]",
|
||||
validated_entities: list[str],
|
||||
existing_texts_lower: set[str],
|
||||
) -> None:
|
||||
"""Recursively extract key:field:value entity strings from a map entity dict."""
|
||||
@@ -58,7 +60,7 @@ def _extract_map_entities(
|
||||
continue
|
||||
label_str = f"{prefix}{field_name}:{v.strip()}"
|
||||
if label_str.lower() not in existing_texts_lower:
|
||||
validated_entities.append(Entity(text=label_str))
|
||||
validated_entities.append(label_str)
|
||||
existing_texts_lower.add(label_str.lower())
|
||||
else:
|
||||
# text or value — single string
|
||||
@@ -66,7 +68,7 @@ def _extract_map_entities(
|
||||
continue
|
||||
label_str = f"{prefix}{field_name}:{field_val.strip()}"
|
||||
if label_str.lower() not in existing_texts_lower:
|
||||
validated_entities.append(Entity(text=label_str))
|
||||
validated_entities.append(label_str)
|
||||
existing_texts_lower.add(label_str.lower())
|
||||
|
||||
|
||||
@@ -113,12 +115,33 @@ def _sanitize_text(text: str | None) -> str | None:
|
||||
return sanitize_llm_output(text)
|
||||
|
||||
|
||||
class Entity(BaseModel):
|
||||
"""An entity extracted from text."""
|
||||
def _coerce_entity_strings(v: Any) -> Any:
|
||||
"""
|
||||
Normalize the LLM's `entities` field to a plain list of strings.
|
||||
|
||||
text: str = Field(
|
||||
description="The specific, named entity as it appears in the fact. Must be a proper noun or specific identifier."
|
||||
)
|
||||
The schema previously asked for `Entity` objects ({"text": "..."}) while the
|
||||
prompt's few-shot examples taught a flat string array. Models that followed
|
||||
the examples literally returned strings, and the entities were silently
|
||||
dropped — none were ever persisted (#2749). The `Entity` wrapper carried no
|
||||
information beyond the string, so it was removed rather than taught to the
|
||||
prompt; the object form is still unwrapped here for models that learned it
|
||||
and for in-flight batch jobs.
|
||||
|
||||
Returns non-list input untouched so pydantic reports the type error itself.
|
||||
"""
|
||||
if v is None:
|
||||
return []
|
||||
if not isinstance(v, list):
|
||||
return v
|
||||
coerced = []
|
||||
for item in v:
|
||||
if isinstance(item, dict):
|
||||
text = item.get("text")
|
||||
if isinstance(text, str):
|
||||
coerced.append(text)
|
||||
else:
|
||||
coerced.append(item)
|
||||
return coerced
|
||||
|
||||
|
||||
class Fact(BaseModel):
|
||||
@@ -143,7 +166,7 @@ class Fact(BaseModel):
|
||||
)
|
||||
|
||||
# Optional structured data
|
||||
entities: list[Entity] | None = None
|
||||
entities: list[str] | None = None
|
||||
causal_relations: list["CausalRelation"] | None = None
|
||||
|
||||
|
||||
@@ -192,9 +215,11 @@ class ExtractedFact(BaseModel):
|
||||
occurred_start: str | None = Field(default=None, description="ISO timestamp for events")
|
||||
occurred_end: str | None = Field(default=None, description="ISO timestamp for event end")
|
||||
fact_type: Literal["world", "assistant"] = Field(
|
||||
description="'world' = objective/external facts. 'assistant' = first-person actions, experiences, or observations by the speaker."
|
||||
description="'world' = objective/external facts, including user preferences, rules, corrections, and constraints even when stated during a conversation. 'assistant' = actions, experiences, or observations the assistant/agent actually performed."
|
||||
)
|
||||
entities: list[str] = Field(
|
||||
default_factory=list, description='People, places, concepts - plain strings, e.g. ["Alice", "Kubernetes"]'
|
||||
)
|
||||
entities: list[Entity] | None = Field(default=None, description="People, places, concepts")
|
||||
causal_relations: list[FactCausalRelation] | None = Field(
|
||||
default=None, description="Links to previous facts (target_index < this fact's index)"
|
||||
)
|
||||
@@ -202,10 +227,7 @@ class ExtractedFact(BaseModel):
|
||||
@field_validator("entities", mode="before")
|
||||
@classmethod
|
||||
def ensure_entities_list(cls, v):
|
||||
"""Ensure entities is always a list (convert None to empty list)."""
|
||||
if v is None:
|
||||
return []
|
||||
return v
|
||||
return _coerce_entity_strings(v)
|
||||
|
||||
def build_fact_text(self) -> str:
|
||||
"""Combine all dimensions into a single comprehensive fact string."""
|
||||
@@ -231,6 +253,55 @@ class FactExtractionResponse(BaseModel):
|
||||
facts: list[ExtractedFact] = Field(description="List of extracted factual statements")
|
||||
|
||||
|
||||
def _split_chunk_for_output_retry(chunk: str) -> tuple[str, str] | None:
|
||||
"""Split an oversized extraction chunk without corrupting structured input."""
|
||||
stripped = chunk.strip()
|
||||
if len(stripped) <= 1:
|
||||
return None
|
||||
|
||||
try:
|
||||
parsed = json.loads(stripped)
|
||||
except (TypeError, ValueError, json.JSONDecodeError):
|
||||
parsed = None
|
||||
|
||||
if isinstance(parsed, list):
|
||||
if len(parsed) >= 2:
|
||||
mid = len(parsed) // 2
|
||||
return json.dumps(parsed[:mid]), json.dumps(parsed[mid:])
|
||||
|
||||
if len(parsed) == 1 and isinstance(parsed[0], dict):
|
||||
turn = parsed[0]
|
||||
content = turn.get("content")
|
||||
if isinstance(content, str) and len(content) > 1:
|
||||
cut = len(content) // 2
|
||||
first_turn = dict(turn)
|
||||
second_turn = dict(turn)
|
||||
first_turn["content"] = content[:cut]
|
||||
second_turn["content"] = content[cut:]
|
||||
return json.dumps([first_turn]), json.dumps([second_turn])
|
||||
|
||||
return None
|
||||
|
||||
# Split plain text at the midpoint, preferring sentence boundaries nearby.
|
||||
mid_point = len(stripped) // 2
|
||||
search_range = int(len(stripped) * 0.2)
|
||||
search_start = max(0, mid_point - search_range)
|
||||
search_end = min(len(stripped), mid_point + search_range)
|
||||
|
||||
best_split = mid_point
|
||||
for ending in [". ", "! ", "? ", "\n\n"]:
|
||||
pos = stripped.rfind(ending, search_start, search_end)
|
||||
if pos != -1:
|
||||
best_split = pos + len(ending)
|
||||
break
|
||||
|
||||
first_half = stripped[:best_split].strip()
|
||||
second_half = stripped[best_split:].strip()
|
||||
if not first_half or not second_half or first_half == stripped or second_half == stripped:
|
||||
return None
|
||||
return first_half, second_half
|
||||
|
||||
|
||||
class ExtractedFactVerbose(BaseModel):
|
||||
"""A single extracted fact with verbose field descriptions for detailed extraction."""
|
||||
|
||||
@@ -295,12 +366,12 @@ class ExtractedFactVerbose(BaseModel):
|
||||
)
|
||||
|
||||
fact_type: Literal["world", "assistant"] = Field(
|
||||
description="'world' = objective/external facts about other people, events, general knowledge. 'assistant' = first-person actions, experiences, or observations by the speaker (e.g., 'I changed X', 'I discovered Y')."
|
||||
description="'world' = objective/external facts about the user, other people, events, general knowledge, preferences, rules, corrections, or constraints. 'assistant' = actions, experiences, or observations the assistant/agent actually performed (e.g., 'I changed X', 'I discovered Y')."
|
||||
)
|
||||
|
||||
entities: list[Entity] | None = Field(
|
||||
default=None,
|
||||
description="Named entities, objects, AND abstract concepts from the fact. Include: people names, organizations, places, significant objects (e.g., 'coffee maker', 'car'), AND abstract concepts/themes (e.g., 'friendship', 'career growth', 'loss', 'celebration'). Extract anything that could help link related facts together.",
|
||||
entities: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="Named entities, objects, AND abstract concepts from the fact, as plain strings (e.g. [\"Alice\", \"friendship\"]). Include: people names, organizations, places, significant objects (e.g., 'coffee maker', 'car'), AND abstract concepts/themes (e.g., 'friendship', 'career growth', 'loss', 'celebration'). Extract anything that could help link related facts together.",
|
||||
)
|
||||
|
||||
causal_relations: list[FactCausalRelation] | None = Field(
|
||||
@@ -312,9 +383,7 @@ class ExtractedFactVerbose(BaseModel):
|
||||
@field_validator("entities", mode="before")
|
||||
@classmethod
|
||||
def ensure_entities_list(cls, v):
|
||||
if v is None:
|
||||
return []
|
||||
return v
|
||||
return _coerce_entity_strings(v)
|
||||
|
||||
|
||||
class FactExtractionResponseVerbose(BaseModel):
|
||||
@@ -345,19 +414,17 @@ class ExtractedFactNoCausal(BaseModel):
|
||||
occurred_start: str | None = Field(default=None, description="WHEN the event happened (ISO timestamp).")
|
||||
occurred_end: str | None = Field(default=None, description="WHEN the event ended (ISO timestamp).")
|
||||
fact_type: Literal["world", "assistant"] = Field(
|
||||
description="'world' = about the user/others. 'assistant' = experience with assistant."
|
||||
description="'world' = about the user/others, including user preferences, rules, corrections, and constraints. 'assistant' = actions or experiences the assistant/agent actually performed."
|
||||
)
|
||||
entities: list[Entity] | None = Field(
|
||||
default=None,
|
||||
description="Named entities, objects, and concepts from the fact.",
|
||||
entities: list[str] = Field(
|
||||
default_factory=list,
|
||||
description='Named entities, objects, and concepts from the fact, as plain strings (e.g. ["Alice", "Kubernetes"]).',
|
||||
)
|
||||
|
||||
@field_validator("entities", mode="before")
|
||||
@classmethod
|
||||
def ensure_entities_list(cls, v):
|
||||
if v is None:
|
||||
return []
|
||||
return v
|
||||
return _coerce_entity_strings(v)
|
||||
|
||||
|
||||
class FactExtractionResponseNoCausal(BaseModel):
|
||||
@@ -389,14 +456,14 @@ class VerbatimExtractedFact(BaseModel):
|
||||
fact_type: Literal["world", "assistant"] = Field(
|
||||
description="'world' = objective/external facts. 'assistant' = first-person actions, experiences, or observations by the speaker."
|
||||
)
|
||||
entities: list[Entity] | None = Field(default=None, description="People, places, concepts")
|
||||
entities: list[str] = Field(
|
||||
default_factory=list, description='People, places, concepts - plain strings, e.g. ["Alice", "Kubernetes"]'
|
||||
)
|
||||
|
||||
@field_validator("entities", mode="before")
|
||||
@classmethod
|
||||
def ensure_entities_list(cls, v):
|
||||
if v is None:
|
||||
return []
|
||||
return v
|
||||
return _coerce_entity_strings(v)
|
||||
|
||||
|
||||
class VerbatimFactExtractionResponse(BaseModel):
|
||||
@@ -451,6 +518,11 @@ def chunk_text(text: str, max_chars: int, structured_chunk_size: int | None = No
|
||||
``structured_chunk_size``. When unset, that limit defaults to ``max_chars``.
|
||||
For plain text, uses sentence-aware splitting.
|
||||
|
||||
The result is idempotent: re-chunking any chunk this returns yields that chunk
|
||||
unchanged. The streaming retain pipeline pre-chunks each document once and then
|
||||
re-chunks every piece during extraction; if a piece re-split, its sub-chunks
|
||||
would inherit one chunk_index and collide on ``chunk_id`` (issue #2301).
|
||||
|
||||
Args:
|
||||
text: Input text to chunk (plain text, JSON conversation, or JSONL)
|
||||
max_chars: Target maximum characters per chunk
|
||||
@@ -469,11 +541,23 @@ def chunk_text(text: str, max_chars: int, structured_chunk_size: int | None = No
|
||||
# Try to parse as JSON conversation array
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
if isinstance(parsed, list) and all(isinstance(turn, dict) for turn in parsed):
|
||||
# This looks like a conversation - chunk at turn boundaries
|
||||
return _chunk_conversation(parsed, max_chars, structured_limit)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
pass
|
||||
parsed = None
|
||||
|
||||
if isinstance(parsed, list) and all(isinstance(turn, dict) for turn in parsed):
|
||||
# This looks like a conversation - chunk at turn boundaries
|
||||
return _chunk_conversation(parsed, max_chars, structured_limit)
|
||||
|
||||
if isinstance(parsed, dict):
|
||||
# A single JSON object — e.g. one JSONL line handed back to the extractor
|
||||
# after the producer already pre-chunked it. It is one structured unit:
|
||||
# keep it whole up to the structured limit, else split it as text within
|
||||
# the chunk budget. Without this, a lone object (one line, so _chunk_jsonl
|
||||
# declines) would fall through to plain-text splitting and re-split a chunk
|
||||
# the producer deliberately kept whole — breaking idempotency (issue #2301).
|
||||
if len(text) <= structured_limit:
|
||||
return [text]
|
||||
return _split_oversized_unit(text, max_chars)
|
||||
|
||||
# Try to parse as JSONL (newline-delimited JSON objects, e.g. session logs)
|
||||
jsonl_chunks = _chunk_jsonl(text, max_chars, structured_limit)
|
||||
@@ -515,10 +599,12 @@ def _chunk_conversation(turns: list[dict], max_chars: int, structured_limit: int
|
||||
turn_size = turn_unit_size + 1 # +1 for comma
|
||||
|
||||
# A turn too large to keep whole even alone: flush, then split it as
|
||||
# text so no chunk runs far over budget (the extractor won't re-chunk).
|
||||
# text. Fragment within min(structured_limit, max_chars) so no fragment
|
||||
# exceeds the chunk budget — otherwise a downstream re-chunk would split
|
||||
# it again and collide on chunk_id (issue #2301).
|
||||
if turn_unit_size > structured_limit:
|
||||
_flush()
|
||||
chunks.extend(_split_oversized_unit(turn_json, structured_limit))
|
||||
chunks.extend(_split_oversized_unit(turn_json, min(structured_limit, max_chars)))
|
||||
continue
|
||||
|
||||
# If adding this turn would exceed limit and we have turns, save current chunk
|
||||
@@ -581,10 +667,12 @@ def _chunk_jsonl(text: str, max_chars: int, structured_limit: int) -> list[str]
|
||||
line_size = len(line) + 1 # +1 for the joining newline
|
||||
|
||||
# A line too large to keep whole even alone: flush, then split it as
|
||||
# text so no chunk runs far over budget (the extractor won't re-chunk).
|
||||
# text. Fragment within min(structured_limit, max_chars) so no fragment
|
||||
# exceeds the chunk budget — otherwise a downstream re-chunk would split
|
||||
# it again and collide on chunk_id (issue #2301).
|
||||
if line_unit_size > structured_limit:
|
||||
_flush()
|
||||
chunks.extend(_split_oversized_unit(line, structured_limit))
|
||||
chunks.extend(_split_oversized_unit(line, min(structured_limit, max_chars)))
|
||||
continue
|
||||
|
||||
# If adding this line would exceed the limit and we have lines, flush.
|
||||
@@ -641,8 +729,8 @@ fact_kind:
|
||||
- "conversation": Ongoing state, preference, trait (no dates)
|
||||
|
||||
fact_type:
|
||||
- "world": About other people, external events, general knowledge, objective facts
|
||||
- "assistant": First-person actions, experiences, or observations by the speaker/author (e.g., "I changed X", "I discovered Y", "I debugged Z"). Also includes interactions with the user (requests, recommendations). If the narrator describes something they did, tried, learned, or decided — use "assistant".
|
||||
- "world": Objective/external facts, including the user's preferences, rules, corrections, constraints, plans, traits, or context. These stay "world" even when the user states them during an assistant interaction (e.g., "User prefers browser_navigate over web_search", "User corrected the project deadline").
|
||||
- "assistant": Actions, experiences, or observations the assistant/agent actually performed (e.g., "I changed X", "I discovered Y", "I debugged Z"). Use this for the assistant/agent doing, trying, learning, deciding, recommending, or responding — not merely for user facts mentioned in conversation.
|
||||
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
TEMPORAL HANDLING
|
||||
@@ -659,6 +747,11 @@ Use "Event Date" from input as reference for relative dates.
|
||||
ENTITIES
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
ALWAYS return "entities" as an array of plain strings — never objects, never null.
|
||||
Correct: entities=["Alice", "Kubernetes", "CKA"]
|
||||
Wrong: entities as an array of objects with a "text" key ← never use this form
|
||||
Use an empty array [] only when the fact truly names nothing.
|
||||
|
||||
Include: people names, organizations, places, key objects, abstract concepts (career, friendship, etc.)
|
||||
Always include "user" when fact is about the user.{examples}"""
|
||||
|
||||
@@ -744,7 +837,7 @@ RULES:
|
||||
- Extract all entities (people, places, organizations, objects, concepts).
|
||||
- Extract temporal information (occurred_start, occurred_end, fact_kind, when).
|
||||
- Extract location (where) and people (who).
|
||||
- fact_type: use "world" unless the content is clearly an interaction with the assistant."""
|
||||
- fact_type: use "world" for user preferences, rules, corrections, constraints, traits, and other objective facts, even when stated during an assistant interaction. Use "assistant" only for actions or experiences the assistant/agent actually performed."""
|
||||
|
||||
VERBATIM_FACT_EXTRACTION_PROMPT = _BASE_FACT_EXTRACTION_PROMPT.format(
|
||||
retain_mission_section="{retain_mission_section}",
|
||||
@@ -845,8 +938,8 @@ For CONVERSATIONS (fact_kind="conversation"):
|
||||
FACT TYPE
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
- **world**: User's life, other people, events (would exist without this conversation)
|
||||
- **assistant**: Interactions with assistant (requests, recommendations, help)
|
||||
- **world**: User's life, preferences, rules, corrections, constraints, other people, and events (facts that would exist without this conversation)
|
||||
- **assistant**: Actions or experiences the assistant/agent actually performed while helping the user (requests, recommendations, help)
|
||||
⚠️ CRITICAL for assistant facts: ALWAYS capture the user's request/question in the fact!
|
||||
Include: what the user asked, what problem they wanted solved, what context they provided
|
||||
|
||||
@@ -1076,8 +1169,8 @@ def _build_extraction_prompt_and_schema(config) -> tuple[str, type]:
|
||||
}
|
||||
if not free_form_entities:
|
||||
dynamic_fields["entities"] = (
|
||||
list[Entity] | None,
|
||||
Field(default=None, description="Leave empty — labels-only mode"),
|
||||
list[str],
|
||||
Field(default_factory=list, description="Leave empty — labels-only mode"),
|
||||
)
|
||||
# Inherit parent's required fields and add 'labels' so it appears in the JSON schema
|
||||
# required array (the base class json_schema_extra overrides required entirely)
|
||||
@@ -1181,9 +1274,17 @@ def _build_request_body(llm_config, config, prompt: str, user_message: str, resp
|
||||
request_body = {
|
||||
"model": llm_config.model,
|
||||
"messages": [{"role": "system", "content": prompt}, {"role": "user", "content": user_message}],
|
||||
"temperature": 0.1,
|
||||
}
|
||||
|
||||
# Honour the configured retain temperature. ``None`` omits the parameter
|
||||
# entirely (for models like Azure GPT-5.5 that reject explicit temperatures),
|
||||
# mirroring LLMProvider.call, which drops temperature when it is None. The
|
||||
# batch path builds the request body directly instead of going through
|
||||
# LLMProvider.call (#2469 only de-hardcoded the streaming path), so it must
|
||||
# apply the same rule here.
|
||||
if config.llm_temperature_retain is not None:
|
||||
request_body["temperature"] = config.llm_temperature_retain
|
||||
|
||||
# Add max_completion_tokens if configured
|
||||
if config.retain_max_completion_tokens:
|
||||
request_body["max_completion_tokens"] = config.retain_max_completion_tokens
|
||||
@@ -1193,19 +1294,32 @@ def _build_request_body(llm_config, config, prompt: str, user_message: str, resp
|
||||
request_body["service_tier"] = llm_config._provider_impl.openai_service_tier
|
||||
|
||||
# Add response_format (JSON schema). The batch path builds the request body
|
||||
# directly instead of going through LLMProvider.call(), so honour
|
||||
# HINDSIGHT_API_LLM_STRICT_SCHEMA here too: strict=True grammar-enforces the
|
||||
# output on capable backends rather than relying on the model to emit clean JSON.
|
||||
# directly instead of going through LLMProvider.call(), so resolve the
|
||||
# strict-schema flag here too: strict=True grammar-enforces the output on capable
|
||||
# backends rather than relying on the model to emit clean JSON. Reads the
|
||||
# retain-scoped field, which already folds in the global HINDSIGHT_API_LLM_STRICT_SCHEMA
|
||||
# fallback, so the batch and streaming paths can't disagree.
|
||||
if hasattr(response_schema, "model_json_schema"):
|
||||
schema = response_schema.model_json_schema()
|
||||
schema = (
|
||||
strict_json_schema(response_schema) if config.llm_strict_schema else response_schema.model_json_schema()
|
||||
)
|
||||
request_body["response_format"] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {"name": "facts", "schema": schema, "strict": config.llm_strict_schema},
|
||||
"json_schema": {"name": "facts", "schema": schema, "strict": config.llm_strict_schema_retain},
|
||||
}
|
||||
|
||||
return request_body
|
||||
|
||||
|
||||
def _coerce_fact_response(response: Any) -> dict[str, Any] | None:
|
||||
"""Accept the schema wrapper, or a recoverable top-level facts array."""
|
||||
if isinstance(response, dict):
|
||||
return response
|
||||
if isinstance(response, list) and all(isinstance(item, dict) for item in response):
|
||||
return {"facts": response}
|
||||
return None
|
||||
|
||||
|
||||
async def _extract_facts_from_chunk(
|
||||
chunk: str,
|
||||
chunk_index: int,
|
||||
@@ -1274,10 +1388,15 @@ async def _extract_facts_from_chunk(
|
||||
llm_max_retries = (
|
||||
config.retain_llm_max_retries if config.retain_llm_max_retries is not None else config.llm_max_retries
|
||||
)
|
||||
# OUTER content-validation attempts (re-prompts on malformed JSON). Follows the
|
||||
# same `N + 1` convention as the providers' transport-retry loops — N retries after
|
||||
# the initial request — so a zero budget still performs one request (#2731). The raw
|
||||
# budget is forwarded unchanged to llm_config.call(), which owns transport retries.
|
||||
outer_attempts = llm_max_retries + 1
|
||||
last_error: Exception | None = None
|
||||
|
||||
usage = TokenUsage() # Track cumulative usage across retries
|
||||
for attempt in range(llm_max_retries):
|
||||
for attempt in range(outer_attempts):
|
||||
try:
|
||||
initial_backoff = (
|
||||
config.retain_llm_initial_backoff
|
||||
@@ -1292,7 +1411,8 @@ async def _extract_facts_from_chunk(
|
||||
messages=[{"role": "system", "content": prompt}, {"role": "user", "content": user_message}],
|
||||
response_format=response_schema,
|
||||
scope="retain_extract_facts",
|
||||
temperature=0.1,
|
||||
temperature=config.llm_temperature_retain,
|
||||
strict_schema=config.llm_strict_schema_retain,
|
||||
max_completion_tokens=config.retain_max_completion_tokens,
|
||||
max_retries=llm_max_retries,
|
||||
initial_backoff=initial_backoff,
|
||||
@@ -1311,10 +1431,11 @@ async def _extract_facts_from_chunk(
|
||||
has_malformed_facts = False
|
||||
|
||||
# Handle malformed LLM responses
|
||||
if not isinstance(extraction_response_json, dict):
|
||||
if attempt < llm_max_retries - 1:
|
||||
coerced_response_json = _coerce_fact_response(extraction_response_json)
|
||||
if coerced_response_json is None:
|
||||
if attempt < outer_attempts - 1:
|
||||
logger.warning(
|
||||
f"LLM returned non-dict JSON on attempt {attempt + 1}/{llm_max_retries}: {type(extraction_response_json).__name__}. Retrying..."
|
||||
f"LLM returned non-dict JSON on attempt {attempt + 1}/{outer_attempts}: {type(extraction_response_json).__name__}. Retrying..."
|
||||
)
|
||||
continue
|
||||
else:
|
||||
@@ -1323,9 +1444,10 @@ async def _extract_facts_from_chunk(
|
||||
# worker's retry machinery and ultimately fails loudly — never
|
||||
# silently commit the document with 0 facts. See issue #1833.
|
||||
raise RuntimeError(
|
||||
f"Fact extraction failed: LLM returned non-dict JSON after {llm_max_retries} attempts "
|
||||
f"Fact extraction failed: LLM returned non-dict JSON after {outer_attempts} attempts "
|
||||
f"({type(extraction_response_json).__name__}). Raw: {str(extraction_response_json)[:500]}"
|
||||
)
|
||||
extraction_response_json = coerced_response_json
|
||||
|
||||
raw_facts = extraction_response_json.get("facts", [])
|
||||
|
||||
@@ -1421,21 +1543,9 @@ async def _extract_facts_from_chunk(
|
||||
elif fact_data.get("occurred_start"):
|
||||
fact_data["occurred_end"] = fact_data["occurred_start"]
|
||||
|
||||
# Add entities if present (validate as Entity objects)
|
||||
# LLM sometimes returns strings instead of {"text": "..."} format
|
||||
entities = get_value("entities")
|
||||
validated_entities = []
|
||||
if entities:
|
||||
# Validate and normalize each entity
|
||||
for ent in entities:
|
||||
if isinstance(ent, str):
|
||||
# Normalize string to Entity object
|
||||
validated_entities.append(Entity(text=ent))
|
||||
elif isinstance(ent, dict) and "text" in ent:
|
||||
try:
|
||||
validated_entities.append(Entity.model_validate(ent))
|
||||
except Exception as e:
|
||||
logger.warning(f"Invalid entity {ent}: {e}")
|
||||
# Entities are plain strings. Older prompts taught a {"text": ...}
|
||||
# object form, so keep unwrapping it for models that still emit it.
|
||||
validated_entities = _coerce_entity_strings(get_value("entities"))
|
||||
|
||||
# Post-process label entities from structured labels object
|
||||
entity_labels_raw = getattr(config, "entity_labels", None)
|
||||
@@ -1445,7 +1555,7 @@ async def _extract_facts_from_chunk(
|
||||
labels_lookup = build_labels_lookup(labels_cfg)
|
||||
labels_data = llm_fact.get("labels") or {}
|
||||
if isinstance(labels_data, dict):
|
||||
existing_texts_lower = {e.text.lower() for e in validated_entities}
|
||||
existing_texts_lower = {e.lower() for e in validated_entities}
|
||||
for group in labels_cfg.attributes:
|
||||
value = labels_data.get(group.key)
|
||||
if not value:
|
||||
@@ -1470,12 +1580,12 @@ async def _extract_facts_from_chunk(
|
||||
label_str = f"{group.key}:{v.strip()}"
|
||||
if group.type == "text":
|
||||
if label_str.lower() not in existing_texts_lower:
|
||||
validated_entities.append(Entity(text=label_str))
|
||||
validated_entities.append(label_str)
|
||||
existing_texts_lower.add(label_str.lower())
|
||||
elif (
|
||||
label_str.lower() in labels_lookup and label_str.lower() not in existing_texts_lower
|
||||
):
|
||||
validated_entities.append(Entity(text=label_str))
|
||||
validated_entities.append(label_str)
|
||||
existing_texts_lower.add(label_str.lower())
|
||||
else:
|
||||
logger.warning(f"Label '{label_str}' not in valid label values, skipping")
|
||||
@@ -1483,7 +1593,7 @@ async def _extract_facts_from_chunk(
|
||||
# In labels-only mode, keep only label entities
|
||||
if not free_form_entities:
|
||||
validated_entities = [
|
||||
e for e in validated_entities if is_label_entity(e.text, labels_cfg, labels_lookup)
|
||||
e for e in validated_entities if is_label_entity(e, labels_cfg, labels_lookup)
|
||||
]
|
||||
elif not free_form_entities:
|
||||
# No labels but free_form disabled: clear all entities
|
||||
@@ -1541,9 +1651,9 @@ async def _extract_facts_from_chunk(
|
||||
continue
|
||||
|
||||
# If we got malformed facts and haven't exhausted retries, try again
|
||||
if has_malformed_facts and len(chunk_facts) < len(raw_facts) * 0.8 and attempt < llm_max_retries - 1:
|
||||
if has_malformed_facts and len(chunk_facts) < len(raw_facts) * 0.8 and attempt < outer_attempts - 1:
|
||||
logger.warning(
|
||||
f"Got {len(raw_facts) - len(chunk_facts)} malformed facts out of {len(raw_facts)} on attempt {attempt + 1}/{llm_max_retries}. Retrying..."
|
||||
f"Got {len(raw_facts) - len(chunk_facts)} malformed facts out of {len(raw_facts)} on attempt {attempt + 1}/{outer_attempts}. Retrying..."
|
||||
)
|
||||
continue
|
||||
|
||||
@@ -1582,7 +1692,7 @@ async def _extract_facts_from_chunk(
|
||||
# If we exhausted all retries, raise the last error or a descriptive fallback
|
||||
if last_error is not None:
|
||||
raise last_error
|
||||
raise RuntimeError(f"Fact extraction failed after {llm_max_retries} attempts: LLM did not return valid JSON")
|
||||
raise RuntimeError(f"Fact extraction failed after {outer_attempts} attempts: LLM did not return valid JSON")
|
||||
|
||||
|
||||
async def _extract_facts_with_auto_split(
|
||||
@@ -1634,33 +1744,22 @@ async def _extract_facts_with_auto_split(
|
||||
metadata=metadata,
|
||||
)
|
||||
except OutputTooLongError:
|
||||
# Output exceeded token limits - split the chunk in half and retry
|
||||
# Output exceeded token limits - split the chunk and retry. Conversation
|
||||
# chunks are JSON arrays, so preserve array/turn boundaries when possible.
|
||||
logger.warning(
|
||||
f"Output too long for chunk {chunk_index + 1}/{total_chunks} "
|
||||
f"({len(chunk)} chars). Splitting in half and retrying..."
|
||||
f"({len(chunk)} chars). Splitting and retrying..."
|
||||
)
|
||||
|
||||
# Split at the midpoint, preferring sentence boundaries
|
||||
mid_point = len(chunk) // 2
|
||||
split_chunks = _split_chunk_for_output_retry(chunk)
|
||||
if split_chunks is None:
|
||||
logger.warning(
|
||||
f"Cannot make progress splitting chunk {chunk_index + 1}/{total_chunks} "
|
||||
f"({len(chunk)} chars); dropping this sub-chunk."
|
||||
)
|
||||
return [], TokenUsage()
|
||||
|
||||
# Try to find a sentence boundary near the midpoint
|
||||
# Look for ". ", "! ", "? " within 20% of midpoint
|
||||
search_range = int(len(chunk) * 0.2)
|
||||
search_start = max(0, mid_point - search_range)
|
||||
search_end = min(len(chunk), mid_point + search_range)
|
||||
|
||||
sentence_endings = [". ", "! ", "? ", "\n\n"]
|
||||
best_split = mid_point
|
||||
|
||||
for ending in sentence_endings:
|
||||
pos = chunk.rfind(ending, search_start, search_end)
|
||||
if pos != -1:
|
||||
best_split = pos + len(ending)
|
||||
break
|
||||
|
||||
# Split the chunk
|
||||
first_half = chunk[:best_split].strip()
|
||||
second_half = chunk[best_split:].strip()
|
||||
first_half, second_half = split_chunks
|
||||
|
||||
logger.info(
|
||||
f"Split chunk {chunk_index + 1} into two sub-chunks: {len(first_half)} chars and {len(second_half)} chars"
|
||||
@@ -1792,10 +1891,28 @@ async def extract_facts_from_text(
|
||||
total_usage = total_usage + chunk_usage
|
||||
|
||||
if failed_chunks:
|
||||
# Include the exception message — not just the type — so operators
|
||||
# can tell a structured-JSON parse failure apart from a rate limit
|
||||
# apart from a network 5xx, all of which can surface as the same
|
||||
# exception types. The error_message we propagate to the
|
||||
# async_operations row is the only inspection surface a worker-side
|
||||
# failure leaves behind, and a bare "chunk 0: RuntimeError" is not
|
||||
# actionable.
|
||||
failed_summary = ", ".join(f"chunk {idx}: {type(err).__name__}: {err}" for idx, err in failed_chunks[:5])
|
||||
quota_errors = [err for _, err in failed_chunks if isinstance(err, ProviderRateLimitResetError)]
|
||||
if quota_errors and len(quota_errors) == len(failed_chunks):
|
||||
retry_at = max(err.retry_at for err in quota_errors)
|
||||
raise ProviderRateLimitResetError(
|
||||
retry_at=retry_at,
|
||||
message=(
|
||||
f"Fact extraction deferred by provider quota: {len(failed_chunks)}/{len(chunks)} chunks failed. "
|
||||
f"First failures: {failed_summary}. Provider detail: {quota_errors[0]}"
|
||||
),
|
||||
) from quota_errors[0]
|
||||
|
||||
# Fail the entire retain — partial extraction is not acceptable.
|
||||
# All successfully extracted facts are discarded because the transaction
|
||||
# hasn't committed yet. The worker poller will retry the entire task.
|
||||
failed_summary = ", ".join(f"chunk {idx}: {type(err).__name__}" for idx, err in failed_chunks[:5])
|
||||
raise RuntimeError(
|
||||
f"Fact extraction failed: {len(failed_chunks)}/{len(chunks)} chunks failed. "
|
||||
f"First failures: {failed_summary}"
|
||||
@@ -2084,7 +2201,10 @@ async def extract_facts_from_contents_batch_api(
|
||||
content_str = message.get("content", "{}")
|
||||
|
||||
try:
|
||||
extraction_response_json = json.loads(content_str)
|
||||
# #2701: use the lenient parser (strips markdown fences, scrubs
|
||||
# embedded control chars) so recoverable batch responses — e.g.
|
||||
# transient Gemini quirks — aren't dropped along with all their facts.
|
||||
extraction_response_json = parse_llm_json(content_str)
|
||||
except json.JSONDecodeError as e:
|
||||
message = f"{custom_id}: failed to parse JSON: {e}"
|
||||
logger.error(message)
|
||||
@@ -2096,6 +2216,19 @@ async def extract_facts_from_contents_batch_api(
|
||||
)
|
||||
continue
|
||||
|
||||
response_type_name = type(extraction_response_json).__name__
|
||||
extraction_response_json = _coerce_fact_response(extraction_response_json)
|
||||
if extraction_response_json is None:
|
||||
message = f"{custom_id}: LLM returned non-dict JSON ({response_type_name})"
|
||||
logger.error(message)
|
||||
extraction_errors.add(message)
|
||||
chunks_metadata.append(
|
||||
ChunkMetadata(
|
||||
chunk_text=chunk_content, fact_count=0, content_index=content_index, chunk_index=chunk_idx
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Parse facts (reuse existing logic from _extract_facts_from_chunk)
|
||||
raw_facts = extraction_response_json.get("facts", [])
|
||||
chunk_facts = []
|
||||
@@ -2162,18 +2295,9 @@ async def extract_facts_from_contents_batch_api(
|
||||
elif fact_data.get("occurred_start"):
|
||||
fact_data["occurred_end"] = fact_data["occurred_start"]
|
||||
|
||||
# Entities
|
||||
entities = get_value("entities")
|
||||
validated_entities = []
|
||||
if entities:
|
||||
for ent in entities:
|
||||
if isinstance(ent, str):
|
||||
validated_entities.append(Entity(text=ent))
|
||||
elif isinstance(ent, dict) and "text" in ent:
|
||||
try:
|
||||
validated_entities.append(Entity.model_validate(ent))
|
||||
except Exception:
|
||||
pass
|
||||
# Entities are plain strings. Older prompts taught a {"text": ...}
|
||||
# object form, so keep unwrapping it for models that still emit it.
|
||||
validated_entities = _coerce_entity_strings(get_value("entities"))
|
||||
|
||||
# Post-process label entities from structured labels object
|
||||
entity_labels_raw = getattr(config, "entity_labels", None)
|
||||
@@ -2183,7 +2307,7 @@ async def extract_facts_from_contents_batch_api(
|
||||
labels_lookup_batch = build_labels_lookup(labels_cfg_batch)
|
||||
labels_data = llm_fact.get("labels") or {}
|
||||
if isinstance(labels_data, dict):
|
||||
existing_texts_lower = {e.text.lower() for e in validated_entities}
|
||||
existing_texts_lower = {e.lower() for e in validated_entities}
|
||||
for group in labels_cfg_batch.attributes:
|
||||
value = labels_data.get(group.key)
|
||||
if not value:
|
||||
@@ -2208,18 +2332,18 @@ async def extract_facts_from_contents_batch_api(
|
||||
label_str = f"{group.key}:{v.strip()}"
|
||||
if group.type == "text":
|
||||
if label_str.lower() not in existing_texts_lower:
|
||||
validated_entities.append(Entity(text=label_str))
|
||||
validated_entities.append(label_str)
|
||||
existing_texts_lower.add(label_str.lower())
|
||||
elif (
|
||||
label_str.lower() in labels_lookup_batch
|
||||
and label_str.lower() not in existing_texts_lower
|
||||
):
|
||||
validated_entities.append(Entity(text=label_str))
|
||||
validated_entities.append(label_str)
|
||||
existing_texts_lower.add(label_str.lower())
|
||||
|
||||
if not free_form_entities_batch:
|
||||
validated_entities = [
|
||||
e for e in validated_entities if is_label_entity(e.text, labels_cfg_batch, labels_lookup_batch)
|
||||
e for e in validated_entities if is_label_entity(e, labels_cfg_batch, labels_lookup_batch)
|
||||
]
|
||||
elif not free_form_entities_batch:
|
||||
validated_entities = []
|
||||
@@ -2301,15 +2425,18 @@ async def extract_facts_from_contents_batch_api(
|
||||
|
||||
for chunk_meta, chunk_facts in facts_by_chunk:
|
||||
content = contents[chunk_meta.content_index]
|
||||
extraction_group_start_idx = global_fact_idx
|
||||
|
||||
for fact_from_llm in chunk_facts:
|
||||
extracted_fact = ExtractedFactType(
|
||||
fact_text=fact_from_llm.fact,
|
||||
fact_type=fact_from_llm.fact_type,
|
||||
entities=[e.text for e in (fact_from_llm.entities or [])],
|
||||
entities=list(fact_from_llm.entities or []),
|
||||
occurred_start=_parse_datetime(fact_from_llm.occurred_start) if fact_from_llm.occurred_start else None,
|
||||
occurred_end=_parse_datetime(fact_from_llm.occurred_end) if fact_from_llm.occurred_end else None,
|
||||
causal_relations=_convert_causal_relations(fact_from_llm.causal_relations or [], global_fact_idx),
|
||||
causal_relations=_convert_causal_relations(
|
||||
fact_from_llm.causal_relations or [], extraction_group_start_idx, len(chunk_facts)
|
||||
),
|
||||
content_index=chunk_meta.content_index,
|
||||
chunk_index=chunk_meta.chunk_index,
|
||||
context=content.context,
|
||||
@@ -2493,40 +2620,37 @@ async def extract_facts_from_contents(
|
||||
fact_idx_in_content = 0
|
||||
for chunk_idx_in_content, (chunk_text, chunk_fact_count) in enumerate(chunks_from_llm):
|
||||
chunk_global_idx = chunk_start_idx + chunk_idx_in_content
|
||||
extraction_group_start_idx = global_fact_idx
|
||||
chunk_facts = facts_from_llm[fact_idx_in_content : fact_idx_in_content + chunk_fact_count]
|
||||
|
||||
for _ in range(chunk_fact_count):
|
||||
if fact_idx_in_content < len(facts_from_llm):
|
||||
fact_from_llm = facts_from_llm[fact_idx_in_content]
|
||||
for fact_from_llm in chunk_facts:
|
||||
# Convert Fact model from LLM to ExtractedFactType dataclass
|
||||
# mentioned_at is always the event_date (when the conversation/document occurred)
|
||||
extracted_fact = ExtractedFactType(
|
||||
fact_text=fact_from_llm.fact,
|
||||
fact_type=fact_from_llm.fact_type,
|
||||
entities=list(fact_from_llm.entities or []),
|
||||
# occurred_start/end: from LLM only, leave None if not provided
|
||||
occurred_start=_parse_datetime(fact_from_llm.occurred_start)
|
||||
if fact_from_llm.occurred_start
|
||||
else None,
|
||||
occurred_end=_parse_datetime(fact_from_llm.occurred_end) if fact_from_llm.occurred_end else None,
|
||||
causal_relations=_convert_causal_relations(
|
||||
fact_from_llm.causal_relations or [], extraction_group_start_idx, len(chunk_facts)
|
||||
),
|
||||
content_index=content_index,
|
||||
chunk_index=chunk_global_idx,
|
||||
context=content.context,
|
||||
# mentioned_at: always the event_date (when the conversation/document occurred)
|
||||
mentioned_at=content.event_date,
|
||||
metadata=content.metadata,
|
||||
tags=content.tags,
|
||||
observation_scopes=content.observation_scopes,
|
||||
)
|
||||
|
||||
# Convert Fact model from LLM to ExtractedFactType dataclass
|
||||
# mentioned_at is always the event_date (when the conversation/document occurred)
|
||||
extracted_fact = ExtractedFactType(
|
||||
fact_text=fact_from_llm.fact,
|
||||
fact_type=fact_from_llm.fact_type,
|
||||
entities=[e.text for e in (fact_from_llm.entities or [])],
|
||||
# occurred_start/end: from LLM only, leave None if not provided
|
||||
occurred_start=_parse_datetime(fact_from_llm.occurred_start)
|
||||
if fact_from_llm.occurred_start
|
||||
else None,
|
||||
occurred_end=_parse_datetime(fact_from_llm.occurred_end)
|
||||
if fact_from_llm.occurred_end
|
||||
else None,
|
||||
causal_relations=_convert_causal_relations(
|
||||
fact_from_llm.causal_relations or [], global_fact_idx
|
||||
),
|
||||
content_index=content_index,
|
||||
chunk_index=chunk_global_idx,
|
||||
context=content.context,
|
||||
# mentioned_at: always the event_date (when the conversation/document occurred)
|
||||
mentioned_at=content.event_date,
|
||||
metadata=content.metadata,
|
||||
tags=content.tags,
|
||||
observation_scopes=content.observation_scopes,
|
||||
)
|
||||
|
||||
extracted_facts.append(extracted_fact)
|
||||
global_fact_idx += 1
|
||||
fact_idx_in_content += 1
|
||||
extracted_facts.append(extracted_fact)
|
||||
global_fact_idx += 1
|
||||
fact_idx_in_content += 1
|
||||
|
||||
# Step 4: For verbatim mode, collapse to one fact per chunk with original text
|
||||
if config.retain_extraction_mode == "verbatim":
|
||||
@@ -2578,7 +2702,9 @@ def _parse_datetime(date_str: str):
|
||||
return None
|
||||
|
||||
|
||||
def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> list[CausalRelationType]:
|
||||
def _convert_causal_relations(
|
||||
relations_from_llm, extraction_group_start_idx: int, extraction_group_size: int
|
||||
) -> list[CausalRelationType]:
|
||||
"""
|
||||
Convert causal relations from LLM format to ExtractedFact format.
|
||||
|
||||
@@ -2586,9 +2712,16 @@ def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> list[C
|
||||
"""
|
||||
causal_relations = []
|
||||
for rel in relations_from_llm:
|
||||
target_fact_index = rel.target_fact_index
|
||||
if (
|
||||
not isinstance(target_fact_index, int)
|
||||
or isinstance(target_fact_index, bool)
|
||||
or not 0 <= target_fact_index < extraction_group_size
|
||||
):
|
||||
continue
|
||||
causal_relation = CausalRelationType(
|
||||
relation_type=rel.relation_type,
|
||||
target_fact_index=fact_start_idx + rel.target_fact_index,
|
||||
target_fact_index=extraction_group_start_idx + target_fact_index,
|
||||
)
|
||||
causal_relations.append(causal_relation)
|
||||
return causal_relations
|
||||
|
||||
@@ -9,7 +9,7 @@ import logging
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from ...config import get_config
|
||||
from ...config import _get_raw_config, get_config
|
||||
from ..memory_engine import fq_table
|
||||
from .bank_utils import DEFAULT_DISPOSITION, create_bank_vector_indexes
|
||||
from .fact_extraction import _sanitize_text
|
||||
@@ -271,6 +271,7 @@ async def handle_document_tracking(
|
||||
retain_params: dict | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
ops=None,
|
||||
store_document_text: bool | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Handle document tracking in the database (full-replace mode).
|
||||
@@ -358,6 +359,7 @@ async def handle_document_tracking(
|
||||
retain_params,
|
||||
document_tags,
|
||||
preserved_created_at=preserved_created_at,
|
||||
store_document_text=store_document_text,
|
||||
)
|
||||
|
||||
|
||||
@@ -368,6 +370,7 @@ async def upsert_document_metadata(
|
||||
combined_content: str,
|
||||
retain_params: dict | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
store_document_text: bool | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Update document metadata without deleting existing facts/chunks.
|
||||
@@ -380,7 +383,16 @@ async def upsert_document_metadata(
|
||||
combined_content = _sanitize_text(combined_content) or ""
|
||||
content_hash = hashlib.sha256(combined_content.encode()).hexdigest()
|
||||
|
||||
await _upsert_document_row(conn, bank_id, document_id, combined_content, content_hash, retain_params, document_tags)
|
||||
await _upsert_document_row(
|
||||
conn,
|
||||
bank_id,
|
||||
document_id,
|
||||
combined_content,
|
||||
content_hash,
|
||||
retain_params,
|
||||
document_tags,
|
||||
store_document_text=store_document_text,
|
||||
)
|
||||
|
||||
|
||||
async def _upsert_document_row(
|
||||
@@ -392,6 +404,7 @@ async def _upsert_document_row(
|
||||
retain_params: dict | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
preserved_created_at: datetime | None = None,
|
||||
store_document_text: bool | None = None,
|
||||
) -> None:
|
||||
"""Insert or update a document row.
|
||||
|
||||
@@ -403,8 +416,13 @@ async def _upsert_document_row(
|
||||
When ``store_document_text`` is disabled, the raw source text
|
||||
is dropped and ``original_text`` is stored as NULL. The ``content_hash`` is
|
||||
still computed from the real content so delta-retain dedup is unaffected.
|
||||
``store_document_text`` defaults to the server-level config when ``None``;
|
||||
the retain path passes the per-bank resolved value.
|
||||
"""
|
||||
original_text = combined_content if get_config().store_document_text else None
|
||||
# Fallback to the raw global default (not get_config(), which guards
|
||||
# bank-configurable fields); the retain path always passes the resolved value.
|
||||
store_text = store_document_text if store_document_text is not None else _get_raw_config().store_document_text
|
||||
original_text = combined_content if store_text else None
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, retain_params, tags, created_at, updated_at)
|
||||
|
||||
@@ -74,7 +74,9 @@ async def create_causal_links_batch(
|
||||
"""
|
||||
Create causal links between facts.
|
||||
|
||||
Links facts that have causal relationships (causes, enables, prevents).
|
||||
Retain writes the canonical ``caused_by`` relationship only. The database and
|
||||
retrieval paths also recognize historical causal types so imported and
|
||||
pre-existing memories remain traversable.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
@@ -90,22 +92,7 @@ async def create_causal_links_batch(
|
||||
if len(unit_ids) != len(facts):
|
||||
raise ValueError(f"Mismatch between unit_ids ({len(unit_ids)}) and facts ({len(facts)})")
|
||||
|
||||
# Extract causal relations in the format expected by link_utils
|
||||
# Format: List of lists, where each inner list is the causal relations for that fact
|
||||
causal_relations_per_fact = []
|
||||
for fact in facts:
|
||||
if fact.causal_relations:
|
||||
# Convert CausalRelation objects to dicts
|
||||
relations_dicts = [
|
||||
{
|
||||
"relation_type": rel.relation_type,
|
||||
"target_fact_index": rel.target_fact_index,
|
||||
}
|
||||
for rel in fact.causal_relations
|
||||
]
|
||||
causal_relations_per_fact.append(relations_dicts)
|
||||
else:
|
||||
causal_relations_per_fact.append([])
|
||||
causal_relations_per_fact = [fact.causal_relations or [] for fact in facts]
|
||||
|
||||
link_count = await link_utils.create_causal_links_batch(conn, bank_id, unit_ids, causal_relations_per_fact, ops=ops)
|
||||
|
||||
|
||||
@@ -7,7 +7,11 @@ import time
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
from ..._vector_index import ann_search_tuning_settings, configured_vector_extension
|
||||
from ..causal_links import CANONICAL_CAUSAL_LINK_TYPES, LEGACY_CAUSAL_LINK_TYPES
|
||||
from ..db.base import DatabaseConnection
|
||||
from ..db.ops import DataAccessOps
|
||||
from ..memory_engine import fq_table
|
||||
from .types import CausalRelation, EntityResolutionResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -300,7 +304,7 @@ async def resolve_entities_only(
|
||||
llm_entities: list[list[dict]],
|
||||
log_buffer: list[str] = None,
|
||||
entity_labels: list | None = None,
|
||||
) -> tuple[list[str], list[tuple], dict[str, list[str]]]:
|
||||
) -> EntityResolutionResult:
|
||||
"""
|
||||
Phase 1 of entity processing: resolve entity names to canonical IDs.
|
||||
|
||||
@@ -321,10 +325,10 @@ async def resolve_entities_only(
|
||||
entity_labels: Optional entity label taxonomy
|
||||
|
||||
Returns:
|
||||
Tuple of (resolved_entity_ids, entity_to_unit, unit_to_entity_ids) where:
|
||||
- resolved_entity_ids: list of entity IDs in same order as flattened entities
|
||||
- entity_to_unit: maps flat index to (unit_id, local_index, fact_date)
|
||||
- unit_to_entity_ids: maps unit_id to list of resolved entity IDs
|
||||
EntityResolutionResult carrying the resolved entity identities (id +
|
||||
stored canonical name, in flattened order), the flat-index → unit map,
|
||||
and the unit → entity-id map used to remap placeholder unit IDs in
|
||||
Phase 2.
|
||||
"""
|
||||
all_entities_flat, _all_entities, entity_to_unit = _prepare_entities_for_resolution(
|
||||
unit_ids, sentences, fact_dates, llm_entities, log_buffer
|
||||
@@ -332,10 +336,10 @@ async def resolve_entities_only(
|
||||
|
||||
if not all_entities_flat:
|
||||
_log(log_buffer, " [6.2] Entity resolution (batched): 0 entities", level="debug")
|
||||
return [], [], {}
|
||||
return EntityResolutionResult(resolved_entities=[], entity_to_unit=[], unit_to_entity_ids={})
|
||||
|
||||
step_start = time.time()
|
||||
resolved_entity_ids = await entity_resolver.resolve_entities_batch(
|
||||
resolved_entities = await entity_resolver.resolve_entities_batch(
|
||||
bank_id=bank_id,
|
||||
entities_data=all_entities_flat,
|
||||
context=context,
|
||||
@@ -354,7 +358,7 @@ async def resolve_entities_only(
|
||||
for idx, (unit_id, _local_idx, _fact_date) in enumerate(entity_to_unit):
|
||||
if unit_id not in unit_to_entity_ids:
|
||||
unit_to_entity_ids[unit_id] = []
|
||||
unit_to_entity_ids[unit_id].append(resolved_entity_ids[idx])
|
||||
unit_to_entity_ids[unit_id].append(resolved_entities[idx].entity_id)
|
||||
|
||||
_log(
|
||||
log_buffer,
|
||||
@@ -362,7 +366,11 @@ async def resolve_entities_only(
|
||||
level="debug",
|
||||
)
|
||||
|
||||
return resolved_entity_ids, entity_to_unit, unit_to_entity_ids
|
||||
return EntityResolutionResult(
|
||||
resolved_entities=resolved_entities,
|
||||
entity_to_unit=entity_to_unit,
|
||||
unit_to_entity_ids=unit_to_entity_ids,
|
||||
)
|
||||
|
||||
|
||||
async def create_temporal_links_batch_per_fact(
|
||||
@@ -658,7 +666,7 @@ def compute_semantic_links_within_batch(
|
||||
"""
|
||||
Compute semantic links between units within the same batch (no DB needed).
|
||||
|
||||
Uses numpy dot product on embeddings already in memory — instant.
|
||||
Uses cosine similarity on embeddings already in memory — instant.
|
||||
|
||||
Args:
|
||||
unit_ids: Unit IDs (real IDs from insert_facts_batch)
|
||||
@@ -675,15 +683,25 @@ def compute_semantic_links_within_batch(
|
||||
import numpy as np
|
||||
|
||||
links = []
|
||||
new_embeddings_matrix = np.array(embeddings)
|
||||
new_embeddings_matrix = np.asarray(embeddings, dtype=float)
|
||||
norms = np.linalg.norm(new_embeddings_matrix, axis=1)
|
||||
valid_embeddings = np.isfinite(new_embeddings_matrix).all(axis=1) & np.isfinite(norms) & (norms > 0)
|
||||
normalized_embeddings = np.zeros_like(new_embeddings_matrix)
|
||||
normalized_embeddings[valid_embeddings] = (
|
||||
new_embeddings_matrix[valid_embeddings] / norms[valid_embeddings, np.newaxis]
|
||||
)
|
||||
|
||||
for i, unit_id in enumerate(unit_ids):
|
||||
if not valid_embeddings[i]:
|
||||
continue
|
||||
|
||||
other_indices = [j for j in range(len(unit_ids)) if j != i]
|
||||
if not other_indices:
|
||||
continue
|
||||
|
||||
other_embeddings = new_embeddings_matrix[other_indices]
|
||||
similarities = np.dot(other_embeddings, new_embeddings_matrix[i])
|
||||
other_embeddings = normalized_embeddings[other_indices]
|
||||
similarities = np.dot(other_embeddings, normalized_embeddings[i])
|
||||
similarities[~valid_embeddings[other_indices]] = -np.inf
|
||||
|
||||
above_threshold = np.where(similarities >= threshold)[0]
|
||||
if len(above_threshold) > 0:
|
||||
@@ -771,28 +789,61 @@ async def create_semantic_links_batch(
|
||||
|
||||
|
||||
async def create_causal_links_batch(
|
||||
conn,
|
||||
conn: DatabaseConnection,
|
||||
bank_id: str,
|
||||
unit_ids: list[str],
|
||||
causal_relations_per_fact: list[list[dict]],
|
||||
ops=None,
|
||||
causal_relations_per_fact: list[list[CausalRelation]],
|
||||
ops: DataAccessOps | None = None,
|
||||
) -> int:
|
||||
"""
|
||||
Create causal links between facts based on LLM-extracted causal relationships.
|
||||
"""Create canonical causal links for the retain pipeline.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
unit_ids: List of unit IDs (in same order as causal_relations_per_fact)
|
||||
causal_relations_per_fact: List of causal relations for each fact.
|
||||
Each element is a list of dicts with:
|
||||
- target_fact_index: Index into unit_ids for the target fact
|
||||
- relation_type: "caused_by"
|
||||
Retain must only create the backward-looking ``caused_by`` form. Historical
|
||||
types are restored exclusively through ``restore_legacy_causal_links_batch``.
|
||||
"""
|
||||
return await _write_causal_links_batch(
|
||||
conn,
|
||||
bank_id,
|
||||
unit_ids,
|
||||
causal_relations_per_fact,
|
||||
CANONICAL_CAUSAL_LINK_TYPES,
|
||||
ops=ops,
|
||||
)
|
||||
|
||||
|
||||
async def restore_legacy_causal_links_batch(
|
||||
conn: DatabaseConnection,
|
||||
bank_id: str,
|
||||
unit_ids: list[str],
|
||||
causal_relations_per_fact: list[list[CausalRelation]],
|
||||
ops: DataAccessOps | None = None,
|
||||
) -> int:
|
||||
"""Restore historical causal links while importing a transfer archive.
|
||||
|
||||
This is deliberately separate from the retain writer: retrieval continues
|
||||
reading historical types, but only transfer import may create them.
|
||||
"""
|
||||
return await _write_causal_links_batch(
|
||||
conn,
|
||||
bank_id,
|
||||
unit_ids,
|
||||
causal_relations_per_fact,
|
||||
LEGACY_CAUSAL_LINK_TYPES,
|
||||
ops=ops,
|
||||
)
|
||||
|
||||
|
||||
async def _write_causal_links_batch(
|
||||
conn: DatabaseConnection,
|
||||
bank_id: str,
|
||||
unit_ids: list[str],
|
||||
causal_relations_per_fact: list[list[CausalRelation]],
|
||||
allowed_relation_types: frozenset[str],
|
||||
ops: DataAccessOps | None = None,
|
||||
) -> int:
|
||||
"""Write causal links after the caller has selected its allowed taxonomy.
|
||||
|
||||
Returns:
|
||||
Number of causal links created
|
||||
|
||||
Causal link type:
|
||||
- "caused_by": This fact was caused by the target fact
|
||||
"""
|
||||
if not unit_ids or not causal_relations_per_fact:
|
||||
return 0
|
||||
@@ -809,15 +860,13 @@ async def create_causal_links_batch(
|
||||
from_unit_id = unit_ids[fact_idx]
|
||||
|
||||
for relation in causal_relations:
|
||||
target_idx = relation["target_fact_index"]
|
||||
relation_type = relation["relation_type"]
|
||||
target_idx = relation.target_fact_index
|
||||
relation_type = relation.relation_type
|
||||
|
||||
# Validate relation_type - only "caused_by" is supported (DB constraint)
|
||||
valid_types = {"caused_by"}
|
||||
if relation_type not in valid_types:
|
||||
if relation_type not in allowed_relation_types:
|
||||
logger.error(
|
||||
f"Invalid relation_type '{relation_type}' (type: {type(relation_type).__name__}) "
|
||||
f"from fact {fact_idx}. Must be one of: {valid_types}. "
|
||||
f"from fact {fact_idx}. Must be one of: {allowed_relation_types}. "
|
||||
f"Relation data: {relation}"
|
||||
)
|
||||
continue
|
||||
|
||||
@@ -71,6 +71,25 @@ def _redact_document_body(body: str, config: Any) -> str:
|
||||
return apply_redaction(body).content
|
||||
|
||||
|
||||
def _is_strict_append_of_stored_document(
|
||||
stored_original_text: str | None,
|
||||
document_body_override: str | None,
|
||||
config: Any,
|
||||
) -> bool:
|
||||
"""Return whether an oversized document body strictly appends stored text.
|
||||
|
||||
``documents.original_text`` is sanitized and may also be Memory Defense
|
||||
redacted before persistence. Apply those same transformations to the
|
||||
complete incoming body before comparing it with the stored prefix.
|
||||
"""
|
||||
if stored_original_text is None or document_body_override is None:
|
||||
return False
|
||||
|
||||
redacted_body = _redact_document_body(document_body_override, config)
|
||||
sanitized_body = fact_extraction._sanitize_text(redacted_body) or ""
|
||||
return len(sanitized_body) > len(stored_original_text) and sanitized_body.startswith(stored_original_text)
|
||||
|
||||
|
||||
async def _fire_memory_defense_webhook(
|
||||
webhook_manager: Any,
|
||||
*,
|
||||
@@ -143,7 +162,7 @@ async def _fire_memory_defense_webhook(
|
||||
logger.warning("memory_defense webhook delivery failed", exc_info=True)
|
||||
|
||||
|
||||
def _audit_memory_defense(
|
||||
async def _audit_memory_defense(
|
||||
audit_logger: Any,
|
||||
*,
|
||||
bank_id: str,
|
||||
@@ -152,11 +171,15 @@ def _audit_memory_defense(
|
||||
) -> None:
|
||||
"""Write a fire-and-forget ``memory_defense`` audit entry for a non-allow decision.
|
||||
|
||||
No-op when audit logging is disabled (the logger gates on its own config).
|
||||
No-op when auditing is off for this bank. ``audit_log_enabled`` is per-bank
|
||||
overridable, so the decision must be awaited here rather than relying on the
|
||||
logger's synchronous allowlist check alone.
|
||||
The action taken (redact/block) and what matched live in the entry metadata.
|
||||
"""
|
||||
if audit_logger is None:
|
||||
return
|
||||
if not await audit_logger.should_log("memory_defense", bank_id):
|
||||
return
|
||||
from ..audit import AuditEntry
|
||||
|
||||
entry = AuditEntry(
|
||||
@@ -246,10 +269,12 @@ from . import (
|
||||
link_creation,
|
||||
)
|
||||
from .types import (
|
||||
CausalRelation,
|
||||
ChunkMetadata,
|
||||
EntityResolutionResult,
|
||||
ExtractedFact,
|
||||
Phase1Result,
|
||||
ProcessedFact,
|
||||
ResolvedEntity,
|
||||
RetainContent,
|
||||
RetainContentDict,
|
||||
)
|
||||
@@ -260,6 +285,15 @@ RetainOutboxCallback = Callable[[asyncpg.Connection], Awaitable[None]]
|
||||
RetainOutboxCallbackFactory = Callable[[list[RetainContentDict]], RetainOutboxCallback | None]
|
||||
|
||||
|
||||
@dataclass
|
||||
class _ProcessedFactBatch:
|
||||
"""Aligned survivors from converting extracted facts for storage."""
|
||||
|
||||
extracted_facts: list[ExtractedFact]
|
||||
processed_facts: list[ProcessedFact]
|
||||
retained_index_by_original: list[int | None]
|
||||
|
||||
|
||||
def _resolve_narrator(profile_name: str, bank_id: str) -> str | None:
|
||||
"""Resolve the narrator (memory owner) used to prime fact extraction.
|
||||
|
||||
@@ -344,7 +378,7 @@ async def _pre_resolve_phase1(
|
||||
embeddings = [fact.embedding for fact in processed_facts]
|
||||
|
||||
async with acquire_with_retry(pool) as resolve_conn:
|
||||
resolved_entity_ids, entity_to_unit, unit_to_entity_ids = await entity_processing.resolve_entities(
|
||||
entity_resolution = await entity_processing.resolve_entities(
|
||||
entity_resolver,
|
||||
resolve_conn,
|
||||
bank_id,
|
||||
@@ -366,11 +400,7 @@ async def _pre_resolve_phase1(
|
||||
)
|
||||
|
||||
return Phase1Result(
|
||||
entities=EntityResolutionResult(
|
||||
resolved_entity_ids=resolved_entity_ids,
|
||||
entity_to_unit=entity_to_unit,
|
||||
unit_to_entity_ids=unit_to_entity_ids,
|
||||
),
|
||||
entities=entity_resolution,
|
||||
semantic_ann_links=semantic_ann_links,
|
||||
)
|
||||
|
||||
@@ -421,7 +451,7 @@ async def _insert_facts_and_links(
|
||||
processed_facts: list[ProcessedFact],
|
||||
config,
|
||||
log_buffer: list[str],
|
||||
resolved_entity_ids: list[str],
|
||||
resolved_entities: list[ResolvedEntity],
|
||||
entity_to_unit: list[tuple],
|
||||
unit_to_entity_ids: dict[str, list[str]],
|
||||
semantic_ann_links: list[tuple],
|
||||
@@ -448,6 +478,7 @@ async def _insert_facts_and_links(
|
||||
# Entity resolution was done in Phase 1 (separate connection).
|
||||
# Remap placeholder IDs to actual unit IDs.
|
||||
step_start = time.time()
|
||||
resolved_entity_ids = [entity.entity_id for entity in resolved_entities]
|
||||
remapped_entity_to_unit, _remapped_unit_to_entity_ids, remapped_semantic = _remap_phase1_results(
|
||||
resolved_entity_ids, entity_to_unit, unit_to_entity_ids, semantic_ann_links or [], unit_ids
|
||||
)
|
||||
@@ -460,6 +491,10 @@ async def _insert_facts_and_links(
|
||||
(unit_id, resolved_entity_ids[idx], fact_date)
|
||||
for idx, (unit_id, _local_idx, fact_date) in enumerate(remapped_entity_to_unit)
|
||||
]
|
||||
# Lock/re-create the resolved parents on THIS transaction before linking,
|
||||
# closing the window where prune_orphan_entities could have deleted one
|
||||
# between Phase-1 resolution and this insert (#2662).
|
||||
await entity_resolver.reassert_entities_batch(bank_id, resolved_entities, conn=conn)
|
||||
await entity_resolver.link_units_to_entities_batch(unit_entity_pairs, conn=conn)
|
||||
log_buffer.append(f" Insert unit_entities: {len(unit_entity_pairs)} pairs in {time.time() - step_start:.3f}s")
|
||||
|
||||
@@ -549,9 +584,90 @@ async def _extract_and_embed(
|
||||
embeddings = await embedding_processing.generate_embeddings_batch(embeddings_model, augmented_texts)
|
||||
log_buffer.append(f" Generate embeddings: {len(embeddings)} embeddings in {time.time() - step_start:.3f}s")
|
||||
|
||||
processed_facts = [ProcessedFact.from_extracted_fact(ef, emb) for ef, emb in zip(extracted_facts, embeddings)]
|
||||
fact_batch = _process_extracted_facts(extracted_facts, embeddings)
|
||||
|
||||
return extracted_facts, processed_facts, chunks, usage
|
||||
return fact_batch.extracted_facts, fact_batch.processed_facts, chunks, usage
|
||||
|
||||
|
||||
def _remap_causal_relations(
|
||||
relations_per_fact: list[list[CausalRelation]],
|
||||
retained_index_by_original: list[int | None],
|
||||
) -> list[list[CausalRelation]]:
|
||||
"""Remap a causal relation matrix after facts have been filtered.
|
||||
|
||||
Both the source row and each ``target_fact_index`` use fact ordinals. A
|
||||
rejected source disappears with its row; a relation to a rejected target
|
||||
must disappear rather than silently pointing at the next surviving fact.
|
||||
"""
|
||||
remapped = [[] for retained_index in retained_index_by_original if retained_index is not None]
|
||||
for original_source, retained_source in enumerate(retained_index_by_original):
|
||||
if retained_source is None:
|
||||
continue
|
||||
for relation in relations_per_fact[original_source]:
|
||||
original_target = relation.target_fact_index
|
||||
retained_target = (
|
||||
retained_index_by_original[original_target]
|
||||
if 0 <= original_target < len(retained_index_by_original)
|
||||
else None
|
||||
)
|
||||
if retained_target is None:
|
||||
continue
|
||||
remapped[retained_source].append(
|
||||
CausalRelation(
|
||||
relation_type=relation.relation_type,
|
||||
target_fact_index=retained_target,
|
||||
)
|
||||
)
|
||||
return remapped
|
||||
|
||||
|
||||
def _process_extracted_facts(
|
||||
extracted_facts: list[ExtractedFact],
|
||||
embeddings: list[list[float]],
|
||||
) -> _ProcessedFactBatch:
|
||||
"""Process facts while preserving their positional relationships.
|
||||
|
||||
``ProcessedFact.from_extracted_fact`` may reject a degenerate fact. Keep
|
||||
the surviving extracted and processed facts in lockstep, and translate
|
||||
causal ordinals from the original extraction into that retained sequence.
|
||||
The returned index table is also used by transfer import for archive-only
|
||||
links and observation source references.
|
||||
"""
|
||||
if len(extracted_facts) != len(embeddings):
|
||||
raise ValueError(
|
||||
f"Extracted facts/embeddings length mismatch: {len(extracted_facts)} facts, {len(embeddings)} embeddings"
|
||||
)
|
||||
|
||||
retained_extracted: list[ExtractedFact] = []
|
||||
processed_facts: list[ProcessedFact] = []
|
||||
retained_index_by_original: list[int | None] = [None] * len(extracted_facts)
|
||||
|
||||
for original_index, (extracted_fact, embedding) in enumerate(zip(extracted_facts, embeddings, strict=True)):
|
||||
processed_fact = ProcessedFact.from_extracted_fact(extracted_fact, embedding)
|
||||
if processed_fact is None:
|
||||
continue
|
||||
retained_index_by_original[original_index] = len(processed_facts)
|
||||
retained_extracted.append(extracted_fact)
|
||||
processed_facts.append(processed_fact)
|
||||
|
||||
remapped_relations = _remap_causal_relations(
|
||||
[fact.causal_relations for fact in extracted_facts],
|
||||
retained_index_by_original,
|
||||
)
|
||||
for extracted_fact, processed_fact, causal_relations in zip(
|
||||
retained_extracted,
|
||||
processed_facts,
|
||||
remapped_relations,
|
||||
strict=True,
|
||||
):
|
||||
extracted_fact.causal_relations = causal_relations
|
||||
processed_fact.causal_relations = causal_relations
|
||||
|
||||
return _ProcessedFactBatch(
|
||||
extracted_facts=retained_extracted,
|
||||
processed_facts=processed_facts,
|
||||
retained_index_by_original=retained_index_by_original,
|
||||
)
|
||||
|
||||
|
||||
async def retain_batch(
|
||||
@@ -732,7 +848,7 @@ async def retain_batch(
|
||||
document_id=_item_doc_id,
|
||||
decision=_decision,
|
||||
)
|
||||
_audit_memory_defense(
|
||||
await _audit_memory_defense(
|
||||
audit_logger,
|
||||
bank_id=bank_id,
|
||||
document_id=_item_doc_id,
|
||||
@@ -831,9 +947,43 @@ async def retain_batch(
|
||||
first = contents_dicts[0]
|
||||
if first.get("context"):
|
||||
existing_content["context"] = first["context"]
|
||||
if first.get("event_date"):
|
||||
existing_content["event_date"] = first["event_date"]
|
||||
if first.get("metadata"):
|
||||
existing_content["metadata"] = first["metadata"]
|
||||
if first.get("observation_scopes") is not None:
|
||||
existing_content["observation_scopes"] = first["observation_scopes"]
|
||||
if first.get("tags"):
|
||||
existing_content["tags"] = first["tags"]
|
||||
contents_dicts = [existing_content, *contents_dicts]
|
||||
# Merge JSON arrays to keep original_text valid (#2409).
|
||||
# Without this, combined_content joins items with "\n", producing
|
||||
# "[...]\n[...]" which is not valid JSON. On the next append cycle
|
||||
# chunk_text() fails to parse it and falls through to sentence-
|
||||
# boundary text splitting, breaking speaker attribution.
|
||||
try:
|
||||
_merged = []
|
||||
for _item in contents_dicts:
|
||||
_parsed = json.loads(_item.get("content", ""))
|
||||
if isinstance(_parsed, list) and all(isinstance(_e, dict) for _e in _parsed):
|
||||
_merged.extend(_parsed)
|
||||
else:
|
||||
_merged = None
|
||||
break
|
||||
if _merged is not None:
|
||||
contents_dicts = [{"content": json.dumps(_merged, ensure_ascii=False)}]
|
||||
if first.get("context"):
|
||||
contents_dicts[0]["context"] = first["context"]
|
||||
if first.get("event_date"):
|
||||
contents_dicts[0]["event_date"] = first["event_date"]
|
||||
if first.get("metadata"):
|
||||
contents_dicts[0]["metadata"] = first["metadata"]
|
||||
if first.get("observation_scopes") is not None:
|
||||
contents_dicts[0]["observation_scopes"] = first["observation_scopes"]
|
||||
if first.get("tags"):
|
||||
contents_dicts[0]["tags"] = first["tags"]
|
||||
except (json.JSONDecodeError, ValueError, TypeError):
|
||||
pass
|
||||
# Rebuild contents list to match
|
||||
contents = _build_contents(contents_dicts, document_tags)
|
||||
log_buffer.append(
|
||||
@@ -1182,6 +1332,14 @@ async def _streaming_retain_batch(
|
||||
retain_params, merged_tags = _build_retain_params(contents_dicts, document_tags)
|
||||
# Track whether document tracking has been done (by the first batch)
|
||||
doc_tracking_done = [False]
|
||||
# Track whether the transactional-outbox callback has already fired inside a
|
||||
# batch write TXN. The in-TXN fire only runs on a final facts-bearing batch
|
||||
# (is_last=True); two success paths never reach it — a committed-chunk count
|
||||
# that lands exactly on a chunk_batch_size boundary (the sentinel drains an
|
||||
# empty batch), and a final batch that extracts zero facts (it returns before
|
||||
# the insert). A post-loop fallback fires the callback in those cases, so this
|
||||
# flag exists to guarantee the callback fires exactly once.
|
||||
outbox_fired = [False]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Producer-consumer pipeline: LLM extraction runs concurrently with DB writes
|
||||
@@ -1354,12 +1512,26 @@ async def _streaming_retain_batch(
|
||||
# from a single oversized item sharing one document_id — without it
|
||||
# each sub-batch restarts at 0 and their chunk_ids collide (#1888).
|
||||
doc_chunk_index = global_idx + chunk_index_offset
|
||||
for fact in extracted:
|
||||
fact_index_offset = len(batch_processed)
|
||||
for fact, processed_fact in zip(extracted, processed, strict=True):
|
||||
fact.content_index = content_idx_in_batch
|
||||
if fact.chunk_index is not None:
|
||||
fact.chunk_index = doc_chunk_index
|
||||
for pf in processed:
|
||||
pf.content_index = content_idx_in_batch
|
||||
processed_fact.content_index = content_idx_in_batch
|
||||
|
||||
# Each producer call extracts one chunk, so its causal ordinals
|
||||
# start at zero. Translate them into the combined consumer-batch
|
||||
# sequence before link creation; otherwise later chunks can point
|
||||
# at equally numbered facts from the first completed chunk.
|
||||
causal_relations = [
|
||||
CausalRelation(
|
||||
relation_type=relation.relation_type,
|
||||
target_fact_index=relation.target_fact_index + fact_index_offset,
|
||||
)
|
||||
for relation in processed_fact.causal_relations
|
||||
]
|
||||
fact.causal_relations = causal_relations
|
||||
processed_fact.causal_relations = causal_relations
|
||||
for cm in chunk_meta:
|
||||
cm.chunk_index = doc_chunk_index
|
||||
|
||||
@@ -1372,7 +1544,12 @@ async def _streaming_retain_batch(
|
||||
nonlocal total_usage
|
||||
total_usage = total_usage + batch_usage
|
||||
|
||||
if not batch_extracted:
|
||||
# ``batch_extracted`` contains only survivors after the degenerate-text
|
||||
# guard. Chunk metadata still records whether extraction originally
|
||||
# produced facts, so an all-rejected batch follows the normal write path
|
||||
# and preserves chunk/outbox behavior from before filtering was added.
|
||||
had_extracted_facts = bool(batch_extracted) or any(chunk.fact_count for chunk in batch_chunk_meta)
|
||||
if not had_extracted_facts:
|
||||
# Even with 0 facts, the first batch must still run document tracking
|
||||
# (cascade-delete + insert doc row) to establish ownership and prevent
|
||||
# concurrent requests from interleaving. Later batches can safely skip.
|
||||
@@ -1400,6 +1577,7 @@ async def _streaming_retain_batch(
|
||||
combined_content,
|
||||
retain_params,
|
||||
merged_tags,
|
||||
store_document_text=getattr(config, "store_document_text", True),
|
||||
)
|
||||
else:
|
||||
await fact_storage.handle_document_tracking(
|
||||
@@ -1411,6 +1589,7 @@ async def _streaming_retain_batch(
|
||||
retain_params,
|
||||
merged_tags,
|
||||
ops=pool.ops,
|
||||
store_document_text=getattr(config, "store_document_text", True),
|
||||
)
|
||||
doc_tracking_done[0] = True
|
||||
# Memory: combined_content has been persisted; release
|
||||
@@ -1502,6 +1681,7 @@ async def _streaming_retain_batch(
|
||||
combined_content,
|
||||
retain_params,
|
||||
merged_tags,
|
||||
store_document_text=getattr(config, "store_document_text", True),
|
||||
)
|
||||
log_buffer.append(
|
||||
f"[streaming] Document {effective_doc_id} updated "
|
||||
@@ -1517,6 +1697,7 @@ async def _streaming_retain_batch(
|
||||
retain_params,
|
||||
merged_tags,
|
||||
ops=pool.ops,
|
||||
store_document_text=getattr(config, "store_document_text", True),
|
||||
)
|
||||
log_buffer.append(f"[streaming] Document {effective_doc_id} tracked (full content)")
|
||||
doc_tracking_done[0] = True
|
||||
@@ -1543,7 +1724,12 @@ async def _streaming_retain_batch(
|
||||
chunk_id_map = {}
|
||||
if batch_chunk_meta:
|
||||
chunk_id_map = await chunk_storage.store_chunks_batch(
|
||||
conn, bank_id, effective_doc_id, batch_chunk_meta, ops=pool.ops
|
||||
conn,
|
||||
bank_id,
|
||||
effective_doc_id,
|
||||
batch_chunk_meta,
|
||||
ops=pool.ops,
|
||||
store_document_text=getattr(config, "store_document_text", True),
|
||||
)
|
||||
log_buffer.append(
|
||||
f" Store chunks: {len(batch_chunk_meta)} chunks in {time.time() - step_start:.3f}s"
|
||||
@@ -1568,7 +1754,7 @@ async def _streaming_retain_batch(
|
||||
batch_processed,
|
||||
config,
|
||||
log_buffer,
|
||||
resolved_entity_ids=phase1.entities.resolved_entity_ids,
|
||||
resolved_entities=phase1.entities.resolved_entities,
|
||||
entity_to_unit=phase1.entities.entity_to_unit,
|
||||
unit_to_entity_ids=phase1.entities.unit_to_entity_ids,
|
||||
semantic_ann_links=[],
|
||||
@@ -1579,6 +1765,12 @@ async def _streaming_retain_batch(
|
||||
|
||||
logger.info(f"[streaming] Phase 2 (write txn): {time.time() - p2_start:.3f}s")
|
||||
|
||||
# The write TXN above committed the transactional-outbox row in the
|
||||
# same transaction as this batch's facts. Record it so the post-loop
|
||||
# fallback doesn't queue a duplicate delivery.
|
||||
if is_last and outbox_callback is not None:
|
||||
outbox_fired[0] = True
|
||||
|
||||
# Best-effort: flush entity_cooccurrences and other deferred stats.
|
||||
try:
|
||||
await entity_resolver.flush_pending_stats()
|
||||
@@ -1615,8 +1807,19 @@ async def _streaming_retain_batch(
|
||||
# Check if facts are already committed (recovery from previous crash).
|
||||
# If so, skip extraction+writes and jump straight to final ANN pass.
|
||||
# ---------------------------------------------------------------------------
|
||||
# Only the call that starts a document at chunk 0 may take the whole-document
|
||||
# skip. When an oversized single item is split into several sequential
|
||||
# sub-batches that SHARE one document_id AND one operation_id (see
|
||||
# _split_contents_into_sub_batches), the first sub-batch commits its chunks
|
||||
# and stamps effective_doc_id into result_metadata.facts_committed_document_ids.
|
||||
# Without the offset gate, every later sub-batch (chunk_index_offset > 0) would
|
||||
# then see its own document already "committed" and skip extraction, dropping
|
||||
# all chunks past the first slice. A non-zero offset inherently means this call
|
||||
# continues a document another sub-batch already started, so it must always do
|
||||
# its work — crash-safety for those chunks still comes from the per-chunk hash
|
||||
# recovery (existing_chunk_hashes) below.
|
||||
facts_already_committed = False
|
||||
if operation_id:
|
||||
if operation_id and chunk_index_offset == 0:
|
||||
try:
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
row = await conn.fetchrow(
|
||||
@@ -1683,6 +1886,7 @@ async def _streaming_retain_batch(
|
||||
combined_content,
|
||||
retain_params,
|
||||
merged_tags,
|
||||
store_document_text=getattr(config, "store_document_text", True),
|
||||
)
|
||||
else:
|
||||
await fact_storage.handle_document_tracking(
|
||||
@@ -1694,6 +1898,7 @@ async def _streaming_retain_batch(
|
||||
retain_params,
|
||||
merged_tags,
|
||||
ops=pool.ops,
|
||||
store_document_text=getattr(config, "store_document_text", True),
|
||||
)
|
||||
doc_tracking_done[0] = True
|
||||
# Memory: combined_content has been persisted and won't be
|
||||
@@ -1701,6 +1906,21 @@ async def _streaming_retain_batch(
|
||||
combined_content = ""
|
||||
log_buffer.append(f"[streaming] Document {effective_doc_id} tracked (no facts extracted)")
|
||||
|
||||
# Transactional-outbox fallback. The in-TXN fire only runs on a final
|
||||
# facts-bearing batch (is_last=True). When the committed-chunk count lands
|
||||
# exactly on a chunk_batch_size boundary the sentinel drains an empty batch
|
||||
# and never marks one last; when the final batch extracts zero facts it
|
||||
# returns before the insert; and when every chunk is skipped as already
|
||||
# committed no batch runs at all. In each of those the retain still
|
||||
# succeeded, so the retain.completed delivery must be queued — exactly once,
|
||||
# in its own transaction (there is no batch TXN left to attach it to). Skip
|
||||
# it on a concurrent takeover: an aborted request must not emit completion.
|
||||
if outbox_callback is not None and not outbox_fired[0] and not pipeline_aborted[0]:
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
await outbox_callback(conn)
|
||||
outbox_fired[0] = True
|
||||
|
||||
# Mark facts as committed in operation metadata (crash recovery checkpoint)
|
||||
if operation_id and all_unit_ids:
|
||||
try:
|
||||
@@ -1868,12 +2088,26 @@ async def _try_delta_retain(
|
||||
# between this read and the write. The write TXN verifies the hash hasn't
|
||||
# changed; if it has, we fall back to streaming (which has full protection).
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
if document_body_override is not None:
|
||||
doc_row_at_load = await conn.fetchrow(
|
||||
f"SELECT content_hash, original_text FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2",
|
||||
effective_doc_id,
|
||||
bank_id,
|
||||
)
|
||||
doc_hash_at_load = doc_row_at_load["content_hash"] if doc_row_at_load else None
|
||||
original_text_at_load = doc_row_at_load["original_text"] if doc_row_at_load else None
|
||||
else:
|
||||
doc_hash_at_load = await conn.fetchval(
|
||||
f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2",
|
||||
effective_doc_id,
|
||||
bank_id,
|
||||
)
|
||||
original_text_at_load = None
|
||||
|
||||
# Load chunks after the document version. If a concurrent writer commits
|
||||
# between these reads, the hash precondition on metadata-only writes (or
|
||||
# the extraction freshness recheck below) forces a streaming fallback.
|
||||
existing_chunks = await chunk_storage.load_existing_chunks(conn, bank_id, effective_doc_id)
|
||||
doc_hash_at_load = await conn.fetchval(
|
||||
f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2",
|
||||
effective_doc_id,
|
||||
bank_id,
|
||||
)
|
||||
|
||||
if not existing_chunks:
|
||||
return None
|
||||
@@ -1905,6 +2139,30 @@ async def _try_delta_retain(
|
||||
)
|
||||
|
||||
if not unchanged_indices:
|
||||
if _is_strict_append_of_stored_document(
|
||||
original_text_at_load,
|
||||
document_body_override,
|
||||
config,
|
||||
):
|
||||
log_buffer.append(
|
||||
"[delta] First oversized slice has no stored chunk match, but "
|
||||
"the complete document strictly appends the stored source — "
|
||||
"preserving historical chunks and advancing document metadata"
|
||||
)
|
||||
return await _delta_metadata_only(
|
||||
pool,
|
||||
bank_id,
|
||||
contents_dicts,
|
||||
contents,
|
||||
effective_doc_id,
|
||||
document_tags,
|
||||
log_buffer,
|
||||
start_time,
|
||||
outbox_callback,
|
||||
document_body_override=document_body_override,
|
||||
config=config,
|
||||
expected_content_hash=doc_hash_at_load,
|
||||
)
|
||||
logger.info(f"Delta retain: no unchanged chunks for {effective_doc_id}, falling back to full retain")
|
||||
return None
|
||||
|
||||
@@ -1925,6 +2183,7 @@ async def _try_delta_retain(
|
||||
outbox_callback,
|
||||
document_body_override=document_body_override,
|
||||
config=config,
|
||||
expected_content_hash=doc_hash_at_load,
|
||||
)
|
||||
|
||||
# Build content items for only the changed/new chunks
|
||||
@@ -1943,6 +2202,7 @@ async def _try_delta_retain(
|
||||
outbox_callback,
|
||||
document_body_override=document_body_override,
|
||||
config=config,
|
||||
expected_content_hash=doc_hash_at_load,
|
||||
)
|
||||
|
||||
# Freshness recheck BEFORE the (expensive) LLM extraction.
|
||||
@@ -1995,6 +2255,7 @@ async def _try_delta_retain(
|
||||
outbox_callback,
|
||||
document_body_override=document_body_override,
|
||||
config=config,
|
||||
expected_content_hash=recheck_hash,
|
||||
)
|
||||
log_buffer.append(
|
||||
f"[delta] Recheck: {len(recheck.changed) + len(recheck.new) + len(recheck.removed)} chunks still differ — "
|
||||
@@ -2124,7 +2385,12 @@ async def _try_delta_retain(
|
||||
for cm in new_chunk_metadata
|
||||
]
|
||||
chunk_id_map = await chunk_storage.store_chunks_batch(
|
||||
conn, bank_id, effective_doc_id, remapped_chunks, ops=pool.ops
|
||||
conn,
|
||||
bank_id,
|
||||
effective_doc_id,
|
||||
remapped_chunks,
|
||||
ops=pool.ops,
|
||||
store_document_text=getattr(config, "store_document_text", True),
|
||||
)
|
||||
for chunk_idx, chunk_id in chunk_id_map.items():
|
||||
chunk_id_map_by_doc[(effective_doc_id, chunk_idx)] = chunk_id
|
||||
@@ -2153,7 +2419,7 @@ async def _try_delta_retain(
|
||||
processed_facts,
|
||||
config,
|
||||
log_buffer,
|
||||
resolved_entity_ids=phase1.entities.resolved_entity_ids,
|
||||
resolved_entities=phase1.entities.resolved_entities,
|
||||
entity_to_unit=phase1.entities.entity_to_unit,
|
||||
unit_to_entity_ids=phase1.entities.unit_to_entity_ids,
|
||||
semantic_ann_links=phase1.semantic_ann_links,
|
||||
@@ -2203,16 +2469,22 @@ async def _delta_metadata_only(
|
||||
*,
|
||||
document_body_override: str | None = None,
|
||||
config: Any = None,
|
||||
):
|
||||
expected_content_hash: str | None = None,
|
||||
) -> tuple[list[list[str]], TokenUsage, int] | None:
|
||||
"""Handle the case where no chunks changed — just update document metadata and tags."""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
# Lock the document row to serialize with concurrent retains
|
||||
await conn.fetchval(
|
||||
current_content_hash = await conn.fetchval(
|
||||
f"SELECT content_hash FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2 FOR UPDATE",
|
||||
document_id,
|
||||
bank_id,
|
||||
)
|
||||
if expected_content_hash is not None and current_content_hash != expected_content_hash:
|
||||
log_buffer.append(
|
||||
f"[delta] Document {document_id} changed before metadata update — falling back to full retain"
|
||||
)
|
||||
return None
|
||||
# When this sub-batch is a slice of an oversized item, write the
|
||||
# full original body (issue #1838) instead of just the slice.
|
||||
# Redact the override since it bypassed per-chunk screening.
|
||||
|
||||
@@ -5,11 +5,14 @@ These dataclasses provide type safety throughout the retain operation,
|
||||
from content input to fact storage.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import Literal, TypedDict
|
||||
from uuid import UUID
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class RetainContentDict(TypedDict, total=False):
|
||||
"""Type definition for content items in retain_batch_async.
|
||||
@@ -96,10 +99,12 @@ class CausalRelation:
|
||||
"""
|
||||
Causal relationship between facts.
|
||||
|
||||
Represents how one fact was caused by another.
|
||||
Retain emits only the backward-looking ``caused_by`` form. Transfer import
|
||||
reuses this structure to restore historical causal types without allowing
|
||||
normal retain writes to create them.
|
||||
"""
|
||||
|
||||
relation_type: str # "caused_by"
|
||||
relation_type: str # ``caused_by`` for retain; legacy types for transfer restore
|
||||
target_fact_index: int # Index of the target fact in the batch
|
||||
|
||||
|
||||
@@ -185,10 +190,47 @@ class ProcessedFact:
|
||||
"""Check if this fact was marked as a duplicate."""
|
||||
return self.unit_id is None
|
||||
|
||||
@staticmethod
|
||||
def _is_degenerate_text(text: str) -> bool:
|
||||
"""Check if fact text has zero information content.
|
||||
|
||||
Rejects empty strings, whitespace-only, single punctuation marks,
|
||||
and common LLM hallucination patterns that carry no semantic meaning.
|
||||
"""
|
||||
stripped = (text or "").strip()
|
||||
if not stripped:
|
||||
return True
|
||||
# Single or repeated punctuation patterns with no semantic content
|
||||
degenerate_patterns = {
|
||||
"...",
|
||||
"…",
|
||||
"-",
|
||||
"--",
|
||||
"---",
|
||||
".",
|
||||
"..",
|
||||
"•",
|
||||
"·",
|
||||
"*",
|
||||
"**",
|
||||
"***",
|
||||
"_,_",
|
||||
"_, _, _",
|
||||
}
|
||||
if stripped in degenerate_patterns:
|
||||
return True
|
||||
# Strings composed entirely of punctuation and whitespace
|
||||
if all(c in ".,;:!?-–—…\"'`´ \t\n\r" for c in stripped):
|
||||
return True
|
||||
# Very short text (<= 2 chars) that is only punctuation
|
||||
if len(stripped) <= 2 and all(not c.isalnum() for c in stripped):
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def from_extracted_fact(
|
||||
extracted_fact: "ExtractedFact", embedding: list[float], chunk_id: str | None = None
|
||||
) -> "ProcessedFact":
|
||||
) -> "ProcessedFact | None":
|
||||
"""
|
||||
Create ProcessedFact from ExtractedFact.
|
||||
|
||||
@@ -198,8 +240,17 @@ class ProcessedFact:
|
||||
chunk_id: Optional chunk ID
|
||||
|
||||
Returns:
|
||||
ProcessedFact ready for storage
|
||||
ProcessedFact ready for storage, or None if the fact text is degenerate
|
||||
(zero information content — punctuation-only, empty, etc.)
|
||||
"""
|
||||
fact_text = extracted_fact.fact_text or ""
|
||||
if ProcessedFact._is_degenerate_text(fact_text):
|
||||
logger.warning(
|
||||
f"Rejected degenerate fact text: type={extracted_fact.fact_type}, "
|
||||
f"text={fact_text[:80]!r}, entities={extracted_fact.entities}"
|
||||
)
|
||||
return None
|
||||
|
||||
# Use occurred dates only if explicitly provided by LLM
|
||||
occurred_start = extracted_fact.occurred_start
|
||||
occurred_end = extracted_fact.occurred_end
|
||||
@@ -209,7 +260,7 @@ class ProcessedFact:
|
||||
entities = [EntityRef(name=name) for name in extracted_fact.entities]
|
||||
|
||||
return ProcessedFact(
|
||||
fact_text=extracted_fact.fact_text,
|
||||
fact_text=fact_text,
|
||||
fact_type=extracted_fact.fact_type,
|
||||
embedding=embedding,
|
||||
occurred_start=occurred_start,
|
||||
@@ -226,19 +277,43 @@ class ProcessedFact:
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResolvedEntity:
|
||||
"""Identity of a resolved entity carried across the retain phase boundary.
|
||||
|
||||
``canonical_name`` is the value stored on the entity row (NOT the raw input
|
||||
mention), captured during Phase-1 resolution. It is threaded to Phase 2 so a
|
||||
parent pruned between phases can be re-created with its real name — the row
|
||||
is gone by then, so the name is otherwise unrecoverable (#2662).
|
||||
"""
|
||||
|
||||
entity_id: str
|
||||
canonical_name: str
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
# Callers pass UUID objects or strings; normalize once so downstream
|
||||
# comparisons, set membership, and SQL binds all see a plain str.
|
||||
self.entity_id = str(self.entity_id)
|
||||
|
||||
|
||||
@dataclass
|
||||
class EntityResolutionResult:
|
||||
"""
|
||||
Result of Phase 1 entity resolution.
|
||||
|
||||
Contains resolved entity IDs and the mapping data needed to remap
|
||||
Contains resolved entity identities and the mapping data needed to remap
|
||||
placeholder unit IDs to real IDs after fact insertion in Phase 2.
|
||||
"""
|
||||
|
||||
resolved_entity_ids: list[str]
|
||||
resolved_entities: list[ResolvedEntity]
|
||||
entity_to_unit: list[tuple]
|
||||
unit_to_entity_ids: dict[str, list[str]]
|
||||
|
||||
@property
|
||||
def resolved_entity_ids(self) -> list[str]:
|
||||
"""Entity IDs in flattened resolution order (used by link remapping)."""
|
||||
return [entity.entity_id for entity in self.resolved_entities]
|
||||
|
||||
|
||||
@dataclass
|
||||
class Phase1Result:
|
||||
|
||||
@@ -27,6 +27,23 @@ def fq_table(table_name: str) -> str:
|
||||
return f"{get_current_schema()}.{table_name}"
|
||||
|
||||
|
||||
def fq_routine(name: str) -> str:
|
||||
"""Schema-qualified name of a cross-tenant discovery routine.
|
||||
|
||||
These routines are database-global — each enumerates ``pg_class`` across every
|
||||
schema and dispatches per schema — so exactly one copy exists, installed into
|
||||
the configured schema by ``b6d2f8a4c1e7``. Calling it through the configured
|
||||
schema rather than a hardcoded ``public.`` is what makes a deployment living
|
||||
in a dedicated non-``public`` schema work (#2638).
|
||||
|
||||
Unlike :func:`fq_table` this ignores the per-request schema contextvar: the
|
||||
routines are deliberately cross-tenant, called from background loops that have
|
||||
no request context.
|
||||
"""
|
||||
schema = get_config().database_schema or "public"
|
||||
return '"' + schema.replace('"', '""') + '".' + name
|
||||
|
||||
|
||||
def fq_table_explicit(table: str, schema: str | None = None) -> str:
|
||||
"""Get fully-qualified table name with an explicit schema override.
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
Helper functions for hybrid search (semantic + BM25 + graph).
|
||||
"""
|
||||
|
||||
from .types import MergedCandidate, RetrievalResult
|
||||
from .types import ArmScores, MergedCandidate, RetrievalResult
|
||||
|
||||
|
||||
def cap_per_source(results: list[RetrievalResult], cap: int) -> list[RetrievalResult]:
|
||||
@@ -51,6 +51,7 @@ def reciprocal_rank_fusion(result_lists: list[list[RetrievalResult]], k: int = 6
|
||||
rrf_scores = {}
|
||||
source_ranks = {} # Track rank from each source for each doc_id
|
||||
all_retrievals = {} # Store the actual RetrievalResult (use first occurrence)
|
||||
arm_scores: dict[str, ArmScores] = {} # doc_id -> raw per-strategy scores across arms
|
||||
|
||||
source_names = ["semantic", "bm25", "graph", "temporal"]
|
||||
|
||||
@@ -79,17 +80,29 @@ def reciprocal_rank_fusion(result_lists: list[list[RetrievalResult]], k: int = 6
|
||||
if doc_id not in rrf_scores:
|
||||
rrf_scores[doc_id] = 0.0
|
||||
source_ranks[doc_id] = {}
|
||||
arm_scores[doc_id] = ArmScores()
|
||||
|
||||
rrf_scores[doc_id] += 1.0 / (k + rank)
|
||||
source_ranks[doc_id][f"{source_name}_rank"] = rank
|
||||
|
||||
# Capture this arm's raw score for the doc (the merged RetrievalResult
|
||||
# below keeps only the first arm's score, so record each arm here).
|
||||
if source_name == "semantic" and retrieval.similarity is not None:
|
||||
arm_scores[doc_id].semantic = retrieval.similarity
|
||||
elif source_name == "bm25" and retrieval.bm25_score is not None:
|
||||
arm_scores[doc_id].keyword = retrieval.bm25_score
|
||||
|
||||
# Combine into final results with metadata
|
||||
merged_results = []
|
||||
for rrf_rank, (doc_id, rrf_score) in enumerate(
|
||||
sorted(rrf_scores.items(), key=lambda x: x[1], reverse=True), start=1
|
||||
):
|
||||
merged_candidate = MergedCandidate(
|
||||
retrieval=all_retrievals[doc_id], rrf_score=rrf_score, rrf_rank=rrf_rank, source_ranks=source_ranks[doc_id]
|
||||
retrieval=all_retrievals[doc_id],
|
||||
rrf_score=rrf_score,
|
||||
rrf_rank=rrf_rank,
|
||||
source_ranks=source_ranks[doc_id],
|
||||
arm_scores=arm_scores[doc_id],
|
||||
)
|
||||
merged_results.append(merged_candidate)
|
||||
|
||||
@@ -118,6 +131,7 @@ def interleave_fusion(result_lists: list[list[RetrievalResult]]) -> list[MergedC
|
||||
source_names = ["semantic", "bm25", "graph", "temporal"]
|
||||
source_ranks: dict[str, dict[str, int]] = {}
|
||||
all_retrievals: dict[str, RetrievalResult] = {}
|
||||
arm_scores: dict[str, ArmScores] = {}
|
||||
|
||||
for source_idx, results in enumerate(result_lists):
|
||||
source_name = source_names[source_idx] if source_idx < len(source_names) else f"source_{source_idx}"
|
||||
@@ -129,6 +143,11 @@ def interleave_fusion(result_lists: list[list[RetrievalResult]]) -> list[MergedC
|
||||
doc_id = retrieval.id
|
||||
all_retrievals.setdefault(doc_id, retrieval)
|
||||
source_ranks.setdefault(doc_id, {})[f"{source_name}_rank"] = rank
|
||||
arm = arm_scores.setdefault(doc_id, ArmScores())
|
||||
if source_name == "semantic" and retrieval.similarity is not None:
|
||||
arm.semantic = retrieval.similarity
|
||||
elif source_name == "bm25" and retrieval.bm25_score is not None:
|
||||
arm.keyword = retrieval.bm25_score
|
||||
|
||||
# Round-robin pick across arms in priority order: all #1s, then all #2s, ...
|
||||
ordered_ids: list[str] = []
|
||||
@@ -151,6 +170,7 @@ def interleave_fusion(result_lists: list[list[RetrievalResult]]) -> list[MergedC
|
||||
rrf_score=float(n - pos),
|
||||
rrf_rank=pos + 1,
|
||||
source_ranks=source_ranks[doc_id],
|
||||
arm_scores=arm_scores[doc_id],
|
||||
)
|
||||
for pos, doc_id in enumerate(ordered_ids)
|
||||
]
|
||||
|
||||
@@ -40,8 +40,6 @@ class GraphRetriever(ABC):
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
query_text: str | None = None,
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # TypedAdjacency, optional pre-loaded graph
|
||||
tags: list[str] | None = None, # Visibility scope tags for filtering
|
||||
tags_match: TagsMatch = "any", # How to match tags: 'any' (OR) or 'all' (AND)
|
||||
@@ -59,8 +57,6 @@ class GraphRetriever(ABC):
|
||||
fact_type: Fact type to filter ('world', 'experience', 'observation')
|
||||
budget: Maximum number of nodes to explore/return
|
||||
query_text: Original query text (optional, for some strategies)
|
||||
semantic_seeds: Pre-computed semantic entry points (from semantic retrieval)
|
||||
temporal_seeds: Pre-computed temporal entry points (from temporal retrieval)
|
||||
adjacency: Pre-loaded typed adjacency graph (optional)
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
"""
|
||||
Link Expansion graph retrieval.
|
||||
|
||||
Expands from semantic/temporal seeds through three parallel, first-class signals
|
||||
stored in memory_links:
|
||||
Selects bounded semantic seeds, then expands through three parallel,
|
||||
first-class signals stored in memory_links:
|
||||
|
||||
1. Entity links — query-time self-join through unit_entities. Score = number of distinct
|
||||
shared entities between the seed set and each candidate, computed via
|
||||
@@ -127,8 +127,6 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
query_text: str | None = None,
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
@@ -146,8 +144,6 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
fact_type: Fact type to filter
|
||||
budget: Maximum results to return
|
||||
query_text: Original query text (unused)
|
||||
semantic_seeds: Pre-computed semantic entry points
|
||||
temporal_seeds: Pre-computed temporal entry points
|
||||
adjacency: Unused, kept for interface compatibility
|
||||
tags: Optional list of tags for visibility filtering
|
||||
|
||||
@@ -158,32 +154,28 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
timings = GraphRetrievalTimings(fact_type=fact_type)
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Find seeds if not provided
|
||||
if semantic_seeds:
|
||||
all_seeds = list(semantic_seeds)
|
||||
else:
|
||||
seeds_start = time.time()
|
||||
all_seeds = await _find_semantic_seeds(
|
||||
conn,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit=20,
|
||||
threshold=0.3,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
tag_groups=tag_groups,
|
||||
created_after=created_after,
|
||||
created_before=created_before,
|
||||
)
|
||||
timings.seeds_time = time.time() - seeds_start
|
||||
logger.debug(
|
||||
f"[LinkExpansion] Found {len(all_seeds)} semantic seeds for fact_type={fact_type} "
|
||||
f"(tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
if temporal_seeds:
|
||||
all_seeds.extend(temporal_seeds)
|
||||
# Graph traversal deliberately chooses its own bounded seeds. The semantic and temporal
|
||||
# retrieval arms have independent candidate limits and thresholds, so reusing their
|
||||
# results would silently change graph-retrieval recall behavior.
|
||||
seeds_start = time.time()
|
||||
all_seeds = await _find_semantic_seeds(
|
||||
conn,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit=20,
|
||||
threshold=0.3,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
tag_groups=tag_groups,
|
||||
created_after=created_after,
|
||||
created_before=created_before,
|
||||
)
|
||||
timings.seeds_time = time.time() - seeds_start
|
||||
logger.debug(
|
||||
f"[LinkExpansion] Found {len(all_seeds)} semantic seeds for fact_type={fact_type} "
|
||||
f"(tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
if not all_seeds:
|
||||
return [], timings
|
||||
@@ -243,16 +235,22 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
}
|
||||
|
||||
sorted_ids = sorted(score_map.keys(), key=lambda x: score_map[x], reverse=True)[:budget]
|
||||
rows = [row_map[fact_id] for fact_id in sorted_ids]
|
||||
|
||||
results = []
|
||||
for row in rows:
|
||||
for fact_id in sorted_ids:
|
||||
row = row_map[fact_id]
|
||||
result = RetrievalResult.from_db_row(dict(row))
|
||||
result.activation = row["score"]
|
||||
# ``activation`` is used to re-sort graph results after fact types are
|
||||
# combined. It must retain the final additive score rather than the
|
||||
# raw score from one signal, which would otherwise discard the other
|
||||
# signals and make the cross-fact-type order disagree with this order.
|
||||
result.activation = score_map[fact_id]
|
||||
results.append(result)
|
||||
|
||||
if tags:
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
# filter_results_by_tags is a no-op when no filter applies (tags falsy and not
|
||||
# the exact-empty/global scope), so call it unconditionally — gating on `if tags:`
|
||||
# would skip the untagged-only filter for tags=[] + tags_match="exact".
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
if tag_groups:
|
||||
results = filter_results_by_tag_groups(results, tag_groups)
|
||||
|
||||
@@ -16,6 +16,44 @@ _RECENCY_ALPHA: float = 0.2
|
||||
_TEMPORAL_ALPHA: float = 0.2
|
||||
_PROOF_COUNT_ALPHA: float = 0.1 # Conservative: max ±5% for evidence strength
|
||||
|
||||
# Recency decay: maps a memory's age (days) onto a freshness signal in [0, 1]
|
||||
# where 0.5 is neutral (no boost). The signal is then folded into the
|
||||
# multiplicative recency_boost via `1 + recency_alpha * (recency - 0.5)`.
|
||||
#
|
||||
# "linear" — straight line from 1.0 (today) to a floor of 0.1, reaching
|
||||
# the floor at `linear_window_days`. The historical default.
|
||||
# "exponential" — 0.5 ** (days_ago / halflife_days). The half-life is the age
|
||||
# at which the signal is exactly neutral (0.5): younger
|
||||
# memories are boosted, older ones penalised, with a smooth
|
||||
# asymptote toward 0 (no hard cutoff).
|
||||
# "none" — always neutral (0.5), disabling the recency boost entirely.
|
||||
# The validated set of names lives in config.RECENCY_DECAY_FUNCTIONS.
|
||||
_RECENCY_DECAY_FUNCTION: str = "linear"
|
||||
_RECENCY_DECAY_LINEAR_WINDOW_DAYS: float = 365.0
|
||||
_RECENCY_DECAY_HALFLIFE_DAYS: float = 90.0
|
||||
|
||||
|
||||
def compute_recency_decay(
|
||||
days_ago: float,
|
||||
function: str = _RECENCY_DECAY_FUNCTION,
|
||||
linear_window_days: float = _RECENCY_DECAY_LINEAR_WINDOW_DAYS,
|
||||
halflife_days: float = _RECENCY_DECAY_HALFLIFE_DAYS,
|
||||
) -> float:
|
||||
"""Map a memory's age in days to a freshness signal in [0, 1] (neutral 0.5).
|
||||
|
||||
Future-dated memories (negative ``days_ago``) clamp to the maximum freshness
|
||||
so they are never penalised. See ``RECENCY_DECAY_FUNCTIONS`` for the shapes.
|
||||
"""
|
||||
if function == "none":
|
||||
return 0.5
|
||||
if function == "exponential":
|
||||
if halflife_days <= 0:
|
||||
return 0.5
|
||||
return min(1.0, 0.5 ** (days_ago / halflife_days))
|
||||
# "linear" (default): straight decay to a 0.1 floor over the window.
|
||||
window = linear_window_days if linear_window_days > 0 else _RECENCY_DECAY_LINEAR_WINDOW_DAYS
|
||||
return max(0.1, min(1.0, 1.0 - (days_ago / window)))
|
||||
|
||||
|
||||
def apply_combined_scoring(
|
||||
scored_results: list[ScoredResult],
|
||||
@@ -24,6 +62,9 @@ def apply_combined_scoring(
|
||||
temporal_alpha: float = _TEMPORAL_ALPHA,
|
||||
proof_count_alpha: float = _PROOF_COUNT_ALPHA,
|
||||
is_passthrough_reranker: bool = False,
|
||||
recency_decay_function: str = _RECENCY_DECAY_FUNCTION,
|
||||
recency_decay_linear_window_days: float = _RECENCY_DECAY_LINEAR_WINDOW_DAYS,
|
||||
recency_decay_halflife_days: float = _RECENCY_DECAY_HALFLIFE_DAYS,
|
||||
) -> None:
|
||||
"""Apply combined scoring to a list of ScoredResults in-place.
|
||||
|
||||
@@ -57,6 +98,12 @@ def apply_combined_scoring(
|
||||
recency_alpha: Max relative recency adjustment (default 0.2 → ±10%).
|
||||
temporal_alpha: Max relative temporal adjustment (default 0.2 → ±10%).
|
||||
proof_count_alpha: Max relative proof count adjustment (default 0.1 → ±5%).
|
||||
recency_decay_function: Age→freshness curve — "linear" (default),
|
||||
"exponential", or "none". See compute_recency_decay.
|
||||
recency_decay_linear_window_days: Days over which the linear curve
|
||||
decays to its floor (default 365).
|
||||
recency_decay_halflife_days: For the exponential curve, the age at which
|
||||
the recency signal is neutral (0.5) (default 90).
|
||||
"""
|
||||
if now.tzinfo is None:
|
||||
now = now.replace(tzinfo=UTC)
|
||||
@@ -98,7 +145,8 @@ def apply_combined_scoring(
|
||||
sr.cross_encoder_score_normalized = 1.0 - (0.9 * new_rank / denom)
|
||||
|
||||
for sr in scored_results:
|
||||
# Recency: linear decay over 365 days → [0.1, 1.0]; neutral 0.5 if no date.
|
||||
# Recency: configurable decay (linear default; see compute_recency_decay)
|
||||
# → [0.0, 1.0]; neutral 0.5 if no date.
|
||||
# Use the unit's effective time (occurred_start, then mentioned_at, then
|
||||
# occurred_end) — the same COALESCE order as retrieval._coalesce_date — so a
|
||||
# memory that carries only a mentioned_at / occurred_end (e.g. conversation
|
||||
@@ -111,7 +159,12 @@ def apply_combined_scoring(
|
||||
if occurred.tzinfo is None:
|
||||
occurred = occurred.replace(tzinfo=UTC)
|
||||
days_ago = (now - occurred).total_seconds() / 86400
|
||||
sr.recency = max(0.1, min(1.0, 1.0 - (days_ago / 365)))
|
||||
sr.recency = compute_recency_decay(
|
||||
days_ago,
|
||||
recency_decay_function,
|
||||
recency_decay_linear_window_days,
|
||||
recency_decay_halflife_days,
|
||||
)
|
||||
|
||||
# Temporal proximity: meaningful only for temporal queries; neutral otherwise.
|
||||
sr.temporal = sr.retrieval.temporal_proximity if sr.retrieval.temporal_proximity is not None else 0.5
|
||||
@@ -124,6 +177,9 @@ def apply_combined_scoring(
|
||||
else:
|
||||
# Neutral baseline is precisely 0.5, ensuring neutral multiplier (1.0)
|
||||
proof_norm = 0.5
|
||||
# Surface the proof signal so the trace can show the proof_count_boost
|
||||
# factor (otherwise the reranked breakdown can't reconcile CE × boosts).
|
||||
sr.proof_norm = proof_norm
|
||||
|
||||
# RRF: kept at 0.0 for trace continuity but excluded from scoring.
|
||||
# RRF is batch-relative (min-max normalised) and redundant after reranking.
|
||||
|
||||
@@ -15,7 +15,7 @@ from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
from ...config import get_config
|
||||
from ...config import DEFAULT_BM25_MAX_QUERY_TERMS, get_config
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from ..sql import create_sql_dialect
|
||||
@@ -104,6 +104,8 @@ async def retrieve_semantic_bm25_combined(
|
||||
tag_groups: list[TagGroup] | None = None,
|
||||
created_after: datetime | None = None,
|
||||
created_before: datetime | None = None,
|
||||
min_semantic: float | None = None,
|
||||
min_keyword: float | None = None,
|
||||
) -> dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]]:
|
||||
"""
|
||||
Combined semantic + BM25 retrieval for multiple fact types in a single query.
|
||||
@@ -143,6 +145,12 @@ async def retrieve_semantic_bm25_combined(
|
||||
config = get_config()
|
||||
tokens = tokenize_query(query_text)
|
||||
|
||||
# Per-request retrieval-level score floors (recall min_scores.semantic / .keyword)
|
||||
# override the global config defaults for this query, pruning weak matches in
|
||||
# the SQL arms before fusion.
|
||||
sem_min = min_semantic if min_semantic is not None else config.semantic_min_similarity
|
||||
bm25_min = min_keyword if min_keyword is not None else config.bm25_min_score
|
||||
|
||||
# Over-fetch for HNSW approximation; semantic results trimmed to limit in Python.
|
||||
hnsw_fetch = max(limit * 5, 100)
|
||||
|
||||
@@ -203,7 +211,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
embedding_param="$1",
|
||||
bank_id_param="$2",
|
||||
fetch_limit=hnsw_fetch,
|
||||
min_similarity=config.semantic_min_similarity,
|
||||
min_similarity=sem_min,
|
||||
tags_clause=tags_clause,
|
||||
groups_clause=groups_clause,
|
||||
extra_where=created_range_clause,
|
||||
@@ -214,7 +222,12 @@ async def retrieve_semantic_bm25_combined(
|
||||
# --- BM25 UNION ALL arms (one per fact_type, only when tokens present) ---
|
||||
if _include_bm25:
|
||||
text_ext = config.text_search_extension
|
||||
bm25_text_param: str = dialect.prepare_bm25_text(tokens, query_text, text_search_extension=text_ext)
|
||||
bm25_text_param: str = dialect.prepare_bm25_text(
|
||||
tokens,
|
||||
query_text,
|
||||
text_search_extension=text_ext,
|
||||
max_query_terms=getattr(config, "bm25_max_query_terms", DEFAULT_BM25_MAX_QUERY_TERMS),
|
||||
)
|
||||
for i, ft in enumerate(fact_types):
|
||||
arms.append(
|
||||
dialect.build_bm25_arm(
|
||||
@@ -229,7 +242,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
arm_index=i,
|
||||
text_search_extension=text_ext,
|
||||
bm25_language=config.text_search_extension_native_language,
|
||||
bm25_min_score=config.bm25_min_score,
|
||||
bm25_min_score=bm25_min,
|
||||
extra_where=created_range_clause,
|
||||
)
|
||||
)
|
||||
@@ -277,7 +290,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
embedding_param="$1",
|
||||
bank_id_param="$2",
|
||||
fetch_limit=hnsw_fetch,
|
||||
min_similarity=config.semantic_min_similarity,
|
||||
min_similarity=sem_min,
|
||||
tags_clause=fb_tags_clause,
|
||||
groups_clause=fb_groups_clause,
|
||||
extra_where=fb_created_clause,
|
||||
@@ -608,7 +621,7 @@ async def retrieve_temporal_combined(
|
||||
# bank_id on memory_units lets the planner use idx_memory_units_bank_fact_type.
|
||||
neighbors = await conn.fetch(
|
||||
f"""
|
||||
SELECT src.from_unit_id, mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.fact_type, mu.document_id, mu.chunk_id, mu.tags, mu.metadata,
|
||||
SELECT src.from_unit_id, mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.fact_type, mu.document_id, mu.chunk_id, mu.tags, mu.metadata, mu.proof_count,
|
||||
l.weight, l.link_type,
|
||||
1 - (mu.embedding <=> $1::vector) AS similarity
|
||||
FROM unnest($2::uuid[]) AS src(from_unit_id)
|
||||
@@ -706,6 +719,8 @@ async def retrieve_all_fact_types_parallel(
|
||||
tag_groups: list[TagGroup] | None = None,
|
||||
created_after: datetime | None = None,
|
||||
created_before: datetime | None = None,
|
||||
min_semantic: float | None = None,
|
||||
min_keyword: float | None = None,
|
||||
) -> MultiFactTypeRetrievalResult:
|
||||
"""
|
||||
Optimized retrieval for multiple fact types using batched queries.
|
||||
@@ -766,6 +781,8 @@ async def retrieve_all_fact_types_parallel(
|
||||
tag_groups=tag_groups,
|
||||
created_after=created_after,
|
||||
created_before=created_before,
|
||||
min_semantic=min_semantic,
|
||||
min_keyword=min_keyword,
|
||||
)
|
||||
semantic_bm25_time = time.time() - semantic_bm25_start
|
||||
|
||||
@@ -805,8 +822,6 @@ async def retrieve_all_fact_types_parallel(
|
||||
fact_type=ft,
|
||||
budget=thinking_budget,
|
||||
query_text=query_text,
|
||||
semantic_seeds=None,
|
||||
temporal_seeds=None,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
tag_groups=tag_groups,
|
||||
|
||||
@@ -14,6 +14,12 @@ AND matching (all/all_strict): Memory matches if ALL request tags are present in
|
||||
EXACT matching: Memory matches only if its tag set EQUALS the request tag set (order-
|
||||
independent). Used for observation "scope" filtering, where each observation lives
|
||||
under exactly one scope (its full tag set) and "scope [a]" must not match "[a, b]".
|
||||
An EMPTY request scope (no tags — ``[]`` or ``None``) is the global/untagged scope and
|
||||
matches only untagged memories — the scope that ``observation_scopes="shared"``
|
||||
consolidation writes to. This is the one mode where absent tags filter rather than
|
||||
meaning "no filter"; all other modes treat empty/absent tags as "no filtering". This
|
||||
mirrors the ``GET .../graph`` endpoint, where ``tags_match="exact"`` with no tags also
|
||||
selects the global scope.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -82,11 +88,16 @@ def build_tags_where_clause(
|
||||
>>> clause, params, next_offset = build_tags_where_clause(['user_a'], 3, 'mu.', 'any_strict')
|
||||
>>> print(clause) # "AND mu.tags IS NOT NULL AND mu.tags != '{}' AND mu.tags && $3"
|
||||
"""
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
|
||||
if match == "exact" and not tags:
|
||||
# Empty/absent scope = global/untagged: match only untagged rows. No bind param
|
||||
# needed (callers gate the param on truthy `tags`, so none is appended).
|
||||
return f"AND ({column} IS NULL OR {column} = '{{}}')", [], param_offset
|
||||
|
||||
if not tags:
|
||||
return "", [], param_offset
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
|
||||
if match == "exact":
|
||||
# Set equality (order-independent): superset AND subset. Untagged rows
|
||||
# (empty array) never satisfy `@>` of a non-empty scope, so they're excluded.
|
||||
@@ -126,11 +137,16 @@ def build_tags_where_clause_simple(
|
||||
Returns:
|
||||
SQL clause string or empty string.
|
||||
"""
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
|
||||
if match == "exact" and not tags:
|
||||
# Empty/absent scope = global/untagged: match only untagged rows. No bind param
|
||||
# needed (callers gate the param on truthy `tags`, so none is appended).
|
||||
return f"AND ({column} IS NULL OR {column} = '{{}}')"
|
||||
|
||||
if not tags:
|
||||
return ""
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
|
||||
if match == "exact":
|
||||
# Set equality (order-independent): superset AND subset. Untagged rows
|
||||
# (empty array) never satisfy `@>` of a non-empty scope, so they're excluded.
|
||||
@@ -164,6 +180,10 @@ def filter_results_by_tags(
|
||||
Returns:
|
||||
Filtered list of results.
|
||||
"""
|
||||
if match == "exact" and not tags:
|
||||
# Empty/absent scope = global/untagged: keep only untagged results.
|
||||
return [r for r in results if not getattr(r, "tags", None)]
|
||||
|
||||
if not tags:
|
||||
return results
|
||||
|
||||
@@ -267,6 +287,9 @@ def _build_group_clause(
|
||||
if isinstance(group, TagGroupLeaf):
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
if group.match == "exact":
|
||||
if len(group.tags) == 0:
|
||||
# Empty scope = global/untagged: match only untagged rows (no bind param).
|
||||
return f"({column} IS NULL OR {column} = '{{}}')", [], param_offset
|
||||
clause = f"({column} @> ${param_offset} AND {column} <@ ${param_offset})"
|
||||
return clause, [group.tags], param_offset + 1
|
||||
operator, include_untagged = _parse_tags_match(group.match)
|
||||
@@ -369,6 +392,9 @@ def _match_group(result: object, group: TagGroup) -> bool:
|
||||
if isinstance(group, TagGroupLeaf):
|
||||
result_tags = getattr(result, "tags", None)
|
||||
is_untagged = result_tags is None or len(result_tags) == 0
|
||||
if group.match == "exact" and len(group.tags) == 0:
|
||||
# Empty scope = global/untagged: match only untagged results.
|
||||
return is_untagged
|
||||
_, include_untagged = _parse_tags_match(group.match)
|
||||
is_any_match = group.match in ("any", "any_strict")
|
||||
tags_set = set(group.tags)
|
||||
|
||||
@@ -5,6 +5,7 @@ Think operation utilities for formulating answers based on agent and world facts
|
||||
import logging
|
||||
from datetime import datetime
|
||||
|
||||
from ...config import get_config
|
||||
from ..response_models import DispositionTraits, MemoryFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -251,7 +252,7 @@ async def reflect(
|
||||
answer_text = await llm_config.call(
|
||||
messages=[{"role": "system", "content": system_message}, {"role": "user", "content": prompt}],
|
||||
scope="memory_think",
|
||||
temperature=0.9,
|
||||
temperature=get_config().llm_temperature_reflect,
|
||||
max_completion_tokens=1000,
|
||||
)
|
||||
|
||||
|
||||
@@ -392,7 +392,7 @@ class SearchTracer:
|
||||
|
||||
# Extract score components (only include non-None values)
|
||||
# Keys from ScoredResult.to_dict(): cross_encoder_score, cross_encoder_score_normalized,
|
||||
# rrf_normalized, temporal, recency, combined_score, weight
|
||||
# rrf_normalized, temporal, recency, proof_norm, combined_score, weight
|
||||
score_components = {}
|
||||
for key in [
|
||||
"cross_encoder_score",
|
||||
@@ -401,6 +401,7 @@ class SearchTracer:
|
||||
"rrf_normalized",
|
||||
"temporal",
|
||||
"recency",
|
||||
"proof_norm",
|
||||
"combined_score",
|
||||
]:
|
||||
if key in result and result[key] is not None:
|
||||
|
||||
@@ -22,7 +22,7 @@ class GraphRetrievalTimings:
|
||||
pattern_count: int = 0 # Number of patterns executed
|
||||
fusion: float = 0.0 # Time for RRF fusion
|
||||
fetch: float = 0.0 # Time to fetch memory unit details
|
||||
seeds_time: float = 0.0 # Time to find semantic seeds (if fallback used)
|
||||
seeds_time: float = 0.0 # Time spent selecting semantic graph seeds
|
||||
result_count: int = 0 # Number of results returned
|
||||
# Detailed per-hop timing: list of {hop, exec_time, uncached, load_time, edges_loaded, total_time}
|
||||
hop_details: list[dict] = field(default_factory=list)
|
||||
@@ -82,6 +82,20 @@ class RetrievalResult:
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ArmScores:
|
||||
"""Raw per-strategy retrieval scores for a single doc, aggregated across arms.
|
||||
|
||||
Fusion keeps only the first-seen RetrievalResult per doc, so its per-arm score
|
||||
fields reflect just one arm. This captures each arm's raw score for the same doc
|
||||
so the recall response can report them (and ``min_scores`` can filter on them).
|
||||
``None`` means the doc was not surfaced by that arm.
|
||||
"""
|
||||
|
||||
semantic: float | None = None # cosine similarity from the semantic arm
|
||||
keyword: float | None = None # BM25 / full-text score from the keyword arm
|
||||
|
||||
|
||||
@dataclass
|
||||
class MergedCandidate:
|
||||
"""
|
||||
@@ -97,6 +111,7 @@ class MergedCandidate:
|
||||
rrf_score: float
|
||||
rrf_rank: int = 0
|
||||
source_ranks: dict[str, int] = field(default_factory=dict) # method_name -> rank
|
||||
arm_scores: "ArmScores" = field(default_factory=lambda: ArmScores()) # raw per-strategy scores
|
||||
|
||||
@property
|
||||
def id(self) -> str:
|
||||
@@ -123,6 +138,7 @@ class ScoredResult:
|
||||
rrf_normalized: float = 0.0
|
||||
recency: float = 0.5
|
||||
temporal: float = 0.5
|
||||
proof_norm: float = 0.5 # log-normalized proof count (neutral 0.5); drives proof_count_boost
|
||||
|
||||
# Final combined score
|
||||
combined_score: float = 0.0
|
||||
@@ -179,6 +195,7 @@ class ScoredResult:
|
||||
result["rrf_normalized"] = self.rrf_normalized
|
||||
result["temporal"] = self.temporal
|
||||
result["recency"] = self.recency
|
||||
result["proof_norm"] = self.proof_norm
|
||||
result["combined_score"] = self.combined_score
|
||||
result["weight"] = self.weight
|
||||
result["activation"] = self.weight # Legacy field
|
||||
|
||||
@@ -300,19 +300,6 @@ class SQLDialect(ABC):
|
||||
"""FOR UPDATE SKIP LOCKED clause (same on both PG and Oracle)."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def advisory_lock(self, id_param: str) -> str:
|
||||
"""Advisory lock expression.
|
||||
|
||||
Args:
|
||||
id_param: Parameter placeholder for the lock ID.
|
||||
|
||||
Returns:
|
||||
PG: "pg_try_advisory_lock($1)"
|
||||
Oracle: "SELECT ... FOR UPDATE NOWAIT" equivalent.
|
||||
"""
|
||||
...
|
||||
|
||||
# -- UUID generation -------------------------------------------------
|
||||
|
||||
@abstractmethod
|
||||
@@ -449,6 +436,7 @@ class SQLDialect(ABC):
|
||||
query_text: str,
|
||||
*,
|
||||
text_search_extension: str = "native",
|
||||
max_query_terms: int | None = None,
|
||||
) -> str:
|
||||
"""Prepare the text parameter value for BM25 search.
|
||||
|
||||
@@ -459,6 +447,8 @@ class SQLDialect(ABC):
|
||||
tokens: Tokenized query words.
|
||||
query_text: Original query text.
|
||||
text_search_extension: Full-text search backend variant.
|
||||
max_query_terms: Optional backend-specific token cap. 0 or None
|
||||
leaves query terms uncapped.
|
||||
|
||||
Returns:
|
||||
Prepared text string to bind as the BM25 text parameter.
|
||||
|
||||
@@ -203,10 +203,6 @@ class OracleDialect(SQLDialect):
|
||||
def for_update_skip_locked(self) -> str:
|
||||
return "FOR UPDATE SKIP LOCKED"
|
||||
|
||||
def advisory_lock(self, id_param: str) -> str:
|
||||
# Oracle doesn't have advisory locks. Use SELECT FOR UPDATE NOWAIT on a lock row.
|
||||
return "SELECT 1 FROM dual FOR UPDATE NOWAIT"
|
||||
|
||||
# -- UUID generation -------------------------------------------------
|
||||
|
||||
def generate_uuid(self) -> str:
|
||||
@@ -303,6 +299,7 @@ class OracleDialect(SQLDialect):
|
||||
query_text: str,
|
||||
*,
|
||||
text_search_extension: str = "native",
|
||||
max_query_terms: int | None = None,
|
||||
) -> str:
|
||||
# Oracle Text: filter tokens with special chars, escape reserved words
|
||||
# with curly braces (e.g. "about" → "{about}"), and join with OR.
|
||||
|
||||
@@ -118,9 +118,6 @@ class PostgreSQLDialect(SQLDialect):
|
||||
def for_update_skip_locked(self) -> str:
|
||||
return "FOR UPDATE SKIP LOCKED"
|
||||
|
||||
def advisory_lock(self, id_param: str) -> str:
|
||||
return f"pg_try_advisory_lock({id_param})"
|
||||
|
||||
# -- UUID generation -------------------------------------------------
|
||||
|
||||
def generate_uuid(self) -> str:
|
||||
@@ -254,8 +251,11 @@ class PostgreSQLDialect(SQLDialect):
|
||||
query_text: str,
|
||||
*,
|
||||
text_search_extension: str = "native",
|
||||
max_query_terms: int | None = None,
|
||||
) -> str:
|
||||
if text_search_extension in ("vchord", "pg_textsearch", "pgroonga", "pg_search"):
|
||||
return query_text
|
||||
if max_query_terms is not None and max_query_terms > 0:
|
||||
tokens = tokens[:max_query_terms]
|
||||
# native tsvector: join tokens with OR operator
|
||||
return " | ".join(tokens)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user