Compare commits

...
52 Commits
Author SHA1 Message Date
Nicolò Boschi 6bd8aa26b0 docs: update 0.4.15 blog cover image 2026-03-03 15:42:08 +01:00
Nicolò Boschi 5b367aacce docs: add 0.4.15 release blog post and changelog 2026-03-03 15:39:17 +01:00
Nicolò Boschi 144e4c49d1 Release v0.4.15
- Update version to 0.4.15 in all components
- Regenerate OpenAPI spec and client SDKs
- Python packages: hindsight-api, hindsight-dev, hindsight-all, hindsight-litellm, hindsight-crewai, hindsight-pydantic-ai, hindsight-embed
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- OpenClaw integration: hindsight-integrations/openclaw
- AI SDK integration: hindsight-integrations/ai-sdk
- Chat SDK integration: hindsight-integrations/chat
- Helm chart
- Sync documentation to version-0.4
2026-03-03 15:03:42 +01:00
Nicolò Boschi 861295dd7c refactor: replace set_gemini_safety_settings() with LLMProvider.with_config() (#474)
* refactor: replace set_gemini_safety_settings() with LLMProvider.with_config()

Removes the fragile ContextVar-setter pattern where callers had to remember
to call set_gemini_safety_settings() at every operation entry point.

Instead, LLMProvider.with_config(resolved_config) returns a
ConfiguredLLMProvider wrapper that:
- injects per-bank settings (Gemini safety settings) on every call via
  token-based ContextVar set/reset — properly scoped, no leakage
- proxies all attribute access to the underlying provider via __getattr__
- requires zero changes to LLMInterface or any provider implementations

Call sites (retain, reflect, consolidation) now pass
llm_config.with_config(resolved_config) to sub-components instead of
setting a global context var and hoping nothing else runs in between.
This pattern also composes naturally with a future per-bank provider
factory: callers always receive something with a .call() method.

* fix: pass messages/tools as kwargs in ConfiguredLLMProvider to preserve class-level patch compatibility
2026-03-03 15:00:32 +01:00
Nicolò Boschi 15f4b8769b fix(ts-sdk): send null instead of undefined when includeEntities is false (#476)
* fix(ts-sdk): send null instead of undefined when includeEntities is false

When `includeEntities: false` was passed, the client serialized `entities`
as `undefined`, which is stripped from JSON. The API then applied its
default (`EntityIncludeOptions()` — enabled), silently ignoring the flag.

Fix: send `null` explicitly when `includeEntities === false` so the API
correctly interprets it as "disable entities".

chunks and source_facts are unaffected since their API defaults are null
(disabled), so omitting them from JSON produces the correct behaviour.

Also adds integration tests covering all three states of includeEntities.

* fix(ts-sdk): use toBeFalsy for null entity check in test
2026-03-03 14:54:48 +01:00
Nicolò Boschi 61bf428ba9 perf: fetch all recall chunks in a single query instead of batched while-loop (#475)
Replace the multi-round-trip while-loop in step 5.5 of recall_async with a
single WHERE chunk_id = ANY($1) query covering all candidate chunk IDs.
Token-budget accounting happens in Python after the single fetch.

Measured on a 97K-unit / 98M-link bank (budget=HIGH, include_chunks,
include_entities):
  p50:  1.209s → 0.611s  (−49%)
  mean: 1.534s → 0.772s  (−50%)
  p95:  3.366s → 2.316s  (−31%)

Also update recall_perf.py benchmark to use Budget.HIGH, include_chunks,
include_entities, and a realistic mixed fact_type distribution.
2026-03-03 14:52:48 +01:00
Nicolò Boschi 73ef99e7b1 feat: add configurable Gemini/Vertex AI safety settings (#473)
Adds per-bank configurable safety settings for Gemini/Vertex AI:
- New `HINDSIGHT_API_LLM_GEMINI_SAFETY_SETTINGS` env var (JSON array)
- Hierarchical config field so banks can override via Config API
- ContextVar pattern for zero-signature-change per-request override
- All 6 thresholds supported: UNSPECIFIED, OFF, BLOCK_NONE, BLOCK_LOW_AND_ABOVE, BLOCK_MEDIUM_AND_ABOVE, BLOCK_ONLY_HIGH
- UI: Models > Gemini/Vertex AI section with per-category threshold selectors and link to Google docs
- Graceful handling when bank_config_api feature is disabled
- 12 new tests covering config parsing, GeminiLLM behaviour, and context var override
2026-03-03 13:48:29 +01:00
Nicolò Boschi 7942f181c2 fix(performance): improve recall and retain performance on large banks (#469) 2026-03-03 13:35:22 +01:00
Anton EvseevandClaude Opus 4.6 5aff8e0c70 refactor(openclaw): replace console.log with debug() helper gated by plugin config (#456)
Replace ~73 console.log calls with a debug() helper that is silent by default.
Debug output is now controlled via plugin config param (debug: true) instead of
environment variables, making it easier for users to configure.

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-03-03 11:05:04 +01:00
DK09876andClaude Opus 4.6 e407f4bc55 feat: add extension hooks for root routing and error headers (#470)
* feat: add OAuth extension hooks for MCP authentication

Add extension points in core that allow cloud extensions to support
OAuth 2.1 (RFC 9728 / RFC 7591) for MCP server authentication:

- HttpExtension.get_root_router() for well-known endpoint mounting
- AuthenticationError.headers for WWW-Authenticate propagation
- MCP middleware forwards auth error headers to clients

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* docs: document get_root_router and AuthenticationError.headers

Add documentation for the new extension points introduced in the
OAuth extension hooks commit.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* Remove OAuth-specific wording from extension docs

Make the AuthenticationError headers example generic instead of
OAuth-specific, since these are general-purpose extension hooks.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

---------

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-03-03 10:50:03 +01:00
Ben 8138fa9002 blog: add CrewAI persistent memory post (#471)
* Add CrewAI persistent memory blog post

* Update blog: add image, remove full example and alternatives sections

* Add CrewAI blog hero image
2026-03-02 16:00:27 -05:00
Nicolò Boschi 1d70abfe85 feat: add tags filtering and q description fix for list documents API (#468)
* feat: add Pydantic AI integration to CI, release pipeline, and docs

- Add test-pydantic-ai-integration job to CI (test.yml)
- Add build, publish, and artifact steps to release workflow (release.yml)
- Add hindsight-integrations/pydantic-ai to release.sh version bumping
- Add Pydantic AI documentation page (sdks/integrations/pydantic-ai.md)
- Add Pydantic AI entry to sidebar with icon

* docs: remove Requirements section from pydantic-ai integration page

* feat: add tags filtering and fix offset pagination docs for list documents API

- Add `tags` and `tags_match` query params to GET /banks/{bank_id}/documents
- Supports any, all, any_strict, all_strict matching modes (default: any_strict)
- Fix `q` param description — it's a case-insensitive substring match on document ID only
- Add tests for offset pagination and all tags_match modes
- Regenerate OpenAPI spec and Python/TypeScript/Go clients
- Document the new filtering options in docs/developer/api/documents.mdx

* fix(cli): pass new tags/tags_match args to list_documents
2026-03-02 17:03:16 +01:00
Nicolò Boschi ecf609c8aa feat: add Pydantic AI integration to CI, release pipeline, and docs (#467)
* feat: add Pydantic AI integration to CI, release pipeline, and docs

- Add test-pydantic-ai-integration job to CI (test.yml)
- Add build, publish, and artifact steps to release workflow (release.yml)
- Add hindsight-integrations/pydantic-ai to release.sh version bumping
- Add Pydantic AI documentation page (sdks/integrations/pydantic-ai.md)
- Add Pydantic AI entry to sidebar with icon

* docs: remove Requirements section from pydantic-ai integration page
2026-03-02 15:40:37 +01:00
BenandClaude Opus 4.6 cab5a40f3a feat: add Pydantic AI integration for persistent agent memory (#441)
* feat: add Pydantic AI integration for persistent agent memory

Adds hindsight-pydantic-ai package providing Hindsight-backed memory
tools for Pydantic AI agents. Since Pydantic AI is async-native, tools
use the hindsight-client async API directly (no thread-pool compat layer).

- create_hindsight_tools(): factory returning retain/recall/reflect Tool instances
- memory_instructions(): auto-injects relevant memories via Agent instructions
- Global configure()/get_config()/reset_config() following existing integration pattern

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* doc: add README for Pydantic AI integration

Co-Authored-By: Claude Opus 4.6 <[email protected]>

---------

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-03-02 15:14:00 +01:00
Nicolò Boschi ab70da1ead docs: entities vs tags vs metadata (#466)
* docs: move entity labels detail to memory-banks, simplify retain overview

* docs: move entity labels blurb under entity-recognition section in retain

* docs: update metadata filtering FAQ to cover entity graph retrieval and entity labels tag option

* docs: enable TOC and fix missing separators in FAQ

* docs: add benchmarks leaderboard screenshot and link to models page

* docs: add 'Which model should I use?' FAQ entry with leaderboard screenshot

* docs: fix leaderboard description to cover retain, reflect, and observations
2026-03-02 14:47:40 +01:00
Nicolò Boschi 9b96becc5c feat: entity labels — optional, free_values, multi_value, UI polish (#450)
* feat: entity labels

* feat: entity labels — optional, free_values, multi_value, UI polish

Completes the entity labels system:

**Schema & extraction**
- Dynamic Pydantic Labels model per fact: each group becomes a typed
  field (Literal | None, list[Literal], str | None, or list[str])
- `optional: bool` flag per group — non-optional enum fields appear in
  JSON schema required array so structured-output providers enforce them
- `free_values: bool` flag per group — accepts any LLM-generated string
  instead of a predefined enum; example values shown as hints in prompt
- New `is_label_entity()` helper for labels-only mode filtering that
  handles both enum lookup and free_values key-prefix matching
- Sentinel rejection: "None"/"null"/"n/a" strings dropped in post-processing

**BM25 / dense retrieval**
- `text_signals` column on memory_units: entity names + date tokens for
  enriched BM25 indexing without polluting stored fact text
- Dense embedding includes occurred_end when it differs from occurred_start
- Alembic migration z1u2v3w4x5y6 (merge revision fixing two heads)

**UI (bank-config-view)**
- Shadcn Switch replaces custom Toggle for both entity-labels and observations
- Shadcn Checkbox for multi/optional/free_values per group
- Input heights bumped to h-8 throughout the editor
- "Label Groups" → "Entity Labels", "Free-form entities" → "Entities"
- Free-text groups show "Example hints" banner in values section

**Tests (45 unit + 3 LLM integration)**
- build_labels_model: single, multi, mixed, free_values optional/required/multi
- is_label_entity: enum match, free_values prefix match, no false positives
- Post-processing: null/absent/string-None/free_values/sentinels/multi-value
- Schema: labels in required, structured object, no labels when unconfigured
- LLM integration: single-value enum, multi-value enum, free_values retain

**Docs**
- retain.md: new Entity Labels section covering groups, flags, examples
- configuration.md: retain_free_form_entities env var + entity_labels note

* fix(tests): update hierarchical fields count for entity_labels additions

entity_labels and retain_free_form_entities are hierarchical fields,
bumping the expected count from 11 to 13.

* fix(migration): rename text_signals revision to avoid collision with main

Main branch claimed z1u2v3w4x5y6 for observation_scopes. Rename our
text_signals migration to a2b3c4d5e6f7, chaining after z1u2v3w4x5y6.

* refactor(entity-labels): simplify free_values — always str|None, no multi

- free_values groups always produce str | None (multi_value and optional
  flags are ignored for free text groups — always optional, never multi)
- Prompt section for free_values groups shows only key + description,
  no values list (users put examples in the description instead)
- UI: section title "Entities", toggle "Free Form Entities", replace
  per-group checkboxes with a type dropdown (Enum / Free text); only
  show multi checkbox and values list when type is Enum
- Update tests to reflect new behaviour

* refactor(entity-labels): replace free_values/multi_value booleans with type field

- LabelGroup now uses type: "value" | "multi-values" | "text" instead of
  free_values/multi_value boolean pair
- Backward-compat migration converts legacy dicts automatically
- Rename retain_free_form_entities → entities_allow_free_form throughout
- Update UI dropdown to show Single value / Multi-values / Free text
- Remove separate multi checkbox (captured by type selection)
- Update docs examples and configuration.md
- Update all tests to use new field names

* fix(migration): backfill observation_scopes column for DBs with swapped z1u2v3w4x5y6

Local DBs that had z1u2v3w4x5y6 applied when it referred to the old
text_signals migration (before it was renamed to a2b3c4d5e6f7) won't have
observation_scopes in their memory_units table. This migration adds the
column with IF NOT EXISTS so it's a no-op on clean installs.

* feat(entity-labels): add tag field to auto-populate memory unit tags from labels

When a LabelGroup has tag=True, extracted key:value entities for that group
are automatically written to the memory unit's tags array. This lets entity
labels double as tags, enabling immediate filtering via the existing
tags/tags_match API params with no extra infrastructure.

- Add tag: bool = False to LabelGroup
- _inject_label_tags() helper called in both sync and batch extraction paths
- UI: add Tag checkbox per label group row
- Docs: document the new tag field
- Tests: 4 new unit tests covering all tag injection paths

* style: ruff format migration file

* fix(migration): fix multiple alembic heads after rebase — point text_signals after nullable_event_date

* fix(clients): update timestamp field to use Timestamp wrapper type after timestamp=unset feature

* style: ruff format agent.py

* fix(docs): update Go quickstart example to use NullableTimestamp for timestamp field
2026-03-02 13:05:25 +01:00
Nicolò Boschi f903948a26 feat: support timestamp="unset" to retain content without a date (#465)
* feat: support timestamp="unset" to retain content without a date

When callers retain timeless content (e.g. fictional documents, static
reference material), passing timestamp="unset" now skips the utcnow()
default so mentioned_at is stored as NULL instead of an artificial date.

- HTTP: validate_timestamp recognises "unset" sentinel and threads it
  through api_retain as event_date=None (key present, value None), which
  the orchestrator distinguishes from key-absent (still defaults to now)
- Orchestrator: new branching logic separates "key absent" → utcnow()
  from "key present but None" → no date
- types.py: RetainContent.event_date and ProcessedFact.mentioned_at are
  now datetime | None; removed the unused _now_utc factory
- fact_extraction.py: all event_date params accept datetime | None;
  _build_user_message emits "Event Date: Unknown" when None; removed
  mentioned_at from the Fact LLM response model (LLM never sets it)
- embedding_processing: skip date suffix when fact_date is None
- entity_resolver: COALESCE(event_date, now()) for first_seen/last_seen
  so entities table NOT NULL constraint is preserved
- link_utils: skip temporal linking for units without event_date
- Migration aa2b3c4d5e6f: DROP NOT NULL on memory_units.event_date
- Tests: test_retain_no_timestamp and test_retain_omit_timestamp_defaults_to_now
- Docs + OpenAPI + TypeScript client updated

* refactor: replace _TIMESTAMP_UNKNOWN sentinel with plain string comparison

The sentinel object() was only needed to distinguish "unset" from None
at the boundary — but since the field type is datetime | str | None,
"unset" can pass through the validator unchanged and be compared directly.

* chore: regenerate OpenAPI spec and clients after timestamp type change

timestamp field is now datetime | str | None to accept the "unset" sentinel value.
2026-03-02 12:03:23 +01:00
Nicolò Boschi 77defd96e9 fix(reflect): prevent context_length_exceeded on large memory banks (#462)
* fix(reflect): prevent context_length_exceeded on large memory banks (#457)

The reflect agent's agentic loop accumulated tool-call messages across
iterations with no upper bound on token count, causing
context_length_exceeded errors on banks with 19K+ nodes.

Changes:
- Add proactive token-budget guard: before each call_with_tools, count
  accumulated message tokens via tiktoken; if >= max_context_tokens and
  evidence has been gathered, immediately synthesize from what was found
- Detect context-overflow errors specifically (_is_context_overflow_error)
  and skip the retry path — retrying after overflow only makes it worse
- Truncate context_history in build_final_prompt to a 60K-token budget
  so the fallback synthesis prompt itself cannot overflow
- Add HINDSIGHT_API_REFLECT_MAX_CONTEXT_TOKENS config (default 100000)
  wired through config.py → main.py → memory_engine → run_reflect_agent
- Tests: unit tests for helpers + mock-LLM behavior tests + an
  end-to-end integration test using a real LLM with max_context_tokens=1

* fix(reflect): derive final prompt context budget from max_context_tokens

Replace the hardcoded _FINAL_PROMPT_CONTEXT_BUDGET (60K tokens) with
a fraction of max_context_tokens (80%), so the fallback synthesis prompt
automatically scales with whatever context window is configured.
2026-03-02 12:03:12 +01:00
Nicolò Boschi c2876490df fix: resolve consolidation deadlock caused by zombie processing tasks on retry (#463)
* fix: resolve consolidation deadlock caused by zombie 'processing' tasks on retry

When a task failed and was rescheduled for retry, submit_task() only updated
task_payload without resetting status/worker_id/claimed_at. The task stayed
permanently in 'processing', blocking all future consolidation for that bank
via the NOT EXISTS guard in claim_batch().

Fix: remove the duplicate payload-based retry mechanism from execute_task().
Retryable failures now re-raise so the poller handles them via _retry_or_fail(),
which already correctly resets status='pending', worker_id=NULL, claimed_at=NULL
and uses the DB retry_count column as single source of truth.

Non-retryable tasks (file_convert_retain) continue to mark themselves failed
and return normally — no exception reaches the poller.

Tests: add regression tests for the retry path (status reset to pending) and
the max-retries exhaustion path (status set to failed).

* ci: re-trigger CI
2026-03-02 11:45:41 +01:00
Nicolò Boschi eaeaa1f24d fix(control-plane): observations count always showing 0 due to wrong field name (#464)
The BankStats interface used total_mental_models but the API returns
total_observations, causing the Observations card to always display 0.
2026-03-02 11:20:19 +01:00
Nicolò Boschi f6f1a7d889 fix: zeroentropy rerank URL missing /v1 prefix and MCP retain async_processing param (#460)
* fix: zeroentropy rerank URL missing /v1 prefix and MCP routing tests

- Fix ZeroEntropy reranker URL: /models/rerank -> /v1/models/rerank (#453)
- Fix test_mcp_routing tests: update assertions to use submit_async_retain
  instead of the non-existent async_processing=False/retain_batch_async pattern

* fix(openclaw): pass retainEveryNTurns through getPluginConfig and set it to 1 in tests

getPluginConfig was not forwarding retainEveryNTurns from the raw config,
so pluginConfig.retainEveryNTurns was always undefined (defaulting to 10).
The integration tests use retainEveryNTurns: 1 so retain fires every turn.
2026-03-02 10:32:01 +01:00
Nicolò Boschi ecb833f40d fix: resolve JSON serialization and logging exception propagation in claude_code_llm (#458, #459) (#461)
- Replace json.dumps(result) with result.model_dump_json() for Pydantic models to fix TypeError during consolidation
- Wrap record_llm_call tracing block in try/except so logging failures never propagate to retry handler
- Fix test_llm_provider.py to use _get_raw_config() for bank-configurable enable_observations field
2026-03-02 10:14:16 +01:00
Chris Bartholomew 5270aa5a6e Add bank-scoped validation to engine and HTTP handlers (#454)
* feat: add bank-scoped validation to engine methods and HTTP handlers

Add validate_bank_read/validate_bank_write hooks to all bank-scoped
engine methods so the operation validator can enforce per-bank API key
restrictions. Add OperationValidationError handling to HTTP handlers
and MCP tools to return proper 403 responses. Add allowed_bank_ids
field to RequestContext.

* Add OperationValidationError handling to mental model GET and DELETE endpoints
2026-03-02 09:55:21 +01:00
Fabio Scarsi ad1660b313 feat(openclaw): retain last n+2 turns every n turns (default n=10) (#452) 2026-03-02 09:44:52 +01:00
Nicolò Boschi 55af468187 feat: observation_scopes field to drive observations granularity (#447)
* feat: observation_scopes field to drive observations granularity

* fix(migration): make a2b3c4d5e6f7 a no-op to fix CI on fresh DB

The z1u2v3w4x5y6 migration already creates observation_scopes directly,
so the rename migration fails on fresh installs where observation_tags
never existed.

* chore: remove no-op migration a2b3c4d5e6f7

* feat: regenerate clients with observation_scopes field

- Add observation_scopes to OpenAPI spec and all generated clients
- Fix Rust build.rs to handle anyOf with >2 variants containing null
  (previously only handled 2-item anyOf, causing progenitor to panic
  on the observation_scopes union type)

* fix(rust): add observation_scopes: None to MemoryItem struct literals

* fix(api): add title to observation_scopes Field for deterministic client generation

Adding title="ObservationScopes" makes the inline anyOf schema use
the explicit name instead of deriving it from the field name, which
was non-deterministic between arm64 (macOS) and amd64 (CI) Docker.

Also fixes description: "each entity" -> "each tag".

* fix(scripts): use linux/amd64 Docker for client generation to ensure reproducibility

Both Python and Go client generation now use --platform linux/amd64
Docker, ensuring identical output on macOS arm64 (local) and Linux
amd64 (CI). Also switches Go from JAR+Java to Docker to eliminate
Java version variability.

* chore: update generated clients to API v0.4.14

* fix(test): add retry logic to test_retain_chinese_content to handle non-deterministic LLM output

* fix(test): mark test_retain_chinese_content as xfail due to non-deterministic LLM translation
2026-02-28 10:46:21 +01:00
Nicolò Boschi 2b5fb10dab doc: 0.4.14 (#449) 2026-02-27 15:41:08 +01:00
Nicolò Boschi 5443c18bfc doc: improvements (#448) 2026-02-27 14:59:11 +01:00
Nicolò Boschi 1c21c0c1a6 doc: changelog and blog post (#445)
* doc: changelog and blog post

* doc: changelog and blog post

* sync
2026-02-26 20:24:53 +01:00
Nicolò Boschi 145454533c Release v0.4.14
- Update version to 0.4.14 in all components
- Regenerate OpenAPI spec and client SDKs
- Python packages: hindsight-api, hindsight-dev, hindsight-all, hindsight-litellm, hindsight-crewai, hindsight-embed
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- OpenClaw integration: hindsight-integrations/openclaw
- AI SDK integration: hindsight-integrations/ai-sdk
- Chat SDK integration: hindsight-integrations/chat
- Helm chart
- Sync documentation to version-0.4
2026-02-26 18:03:14 +01:00
Nicolò Boschi 9aaa78b9a6 fix: doc build 2026-02-26 18:01:56 +01:00
Nicolò Boschi bfa09d1685 fix: doc build and add chat doc (#444)
* fix doc build and add chat doc

* fix doc build and add chat doc
2026-02-26 17:55:36 +01:00
Nicolò Boschi 6f5245ae58 chore: integrate chat with release (#443)
* integrate chat with release

* integrate chat with release
2026-02-26 17:38:51 +01:00
BenandClaude Opus 4.6 fed987f931 feat: add Chat SDK integration for persistent chat bot memory (#442)
Adds @vectorize-io/hindsight-chat, a wrapper for the Vercel Chat SDK
that gives any chat bot (Slack, Discord, Teams, etc.) long-term memory
via Hindsight. Includes withHindsightChat() handler wrapper with
auto-recall, auto-retain, and memoriesAsSystemPrompt() formatting.

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-26 17:07:47 +01:00
Anton EvseevandClaude Opus 4.6 8cd65b9896 fix: raise error when embedding dimensions exceed pgvector HNSW limit (#361)
Instead of silently skipping HNSW index creation for embeddings > 2000
dimensions, raise a RuntimeError with an actionable message suggesting
pgvectorscale/DiskANN as an alternative.

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-26 10:24:17 +01:00
Nicolò Boschi b813bd2728 doc: fix build 2026-02-25 16:58:06 +01:00
And#ocean 86d8ac08b1 fix(storage): use dynamic schema_getter in PostgreSQLFileStorage for multi-tenant (#440)
PostgreSQLFileStorage was initialized once at startup with a static
schema value. Since get_current_schema() returns the default schema at
init time, multi-tenant requests always queried the wrong schema,
causing "relation file_storage does not exist" errors.

Replace static schema with schema_getter callable (same pattern used
by BrokerTaskBackend since #208) so the schema is resolved dynamically
per-request via contextvars.
2026-02-25 15:19:08 +01:00
Nicolò Boschi 4b328a9cb3 feat: configure exposed mcp tools per bank (#439)
* feat: configure exposed mcp tools per bank

* fix: update configurable fields count to 11 after adding mcp_enabled_tools
2026-02-25 11:57:46 +01:00
Sense_wangandhaosenwang1018 f5b94d4b28 fix: catch ValueError instead of bare except in date parsing (#438)
The datetime.strptime() call can only raise ValueError on format
mismatch. Bare except catches KeyboardInterrupt and SystemExit,
which masks real errors.

Co-authored-by: haosenwang1018 <[email protected]>
2026-02-25 10:33:56 +01:00
Eliah RusinandClaude Opus 4.6 58f2de70fb fix: pass encoding_format="float" in LiteLLM embedding calls (#434)
DeepInfra rejects requests when encoding_format is null. LiteLLM sets
it to None by default, so we explicitly pass "float" — the only format
compatible with our list[list[float]] return type.

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-25 10:32:05 +01:00
Nicolò Boschi 0bb5ca4caf feat: filter graph memories with tags (#431)
* feat: filter graph memories with tags

* fix(cli): pass new q/tags/tags_match args to get_graph

* docs: use CodeSnippet for tags_match examples in recall.mdx
2026-02-25 10:31:40 +01:00
DK09876andClaude Opus 4.6 3ffec65090 feat: expand MCP tool surface area with 18 new tools and enhanced parameters (#435)
Add directives, memory browsing, documents, operations, tags, and bank
management tools to the MCP server. Expose previously hardcoded parameters
(budget, types, tags, response_schema, trigger) on retain, recall, reflect,
and mental model tools. Update docs for all new tools and parameters.

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-25 10:26:46 +01:00
Nicolò Boschi 0aa7c2b3a1 feat: batch observations consolidation (#430)
* feat: batch observations consolidation

* feat: batch observations consolidation

* docs: add CONSOLIDATION_LLM_BATCH_SIZE config flag documentation
2026-02-24 15:19:39 +01:00
Nicolò Boschi ac9a94ade3 fix: handle observations regeneration when memories get deleted (#429)
* fix: handle observations regeneration when memories get deleted

* feat: add clear_memory_observations endpoint and regenerate clients

- Add DELETE /banks/{id}/memories/{memory_id}/observations endpoint
- Add observations lifecycle/invalidation section to docs
- Regenerate OpenAPI spec and all clients (Python, TypeScript, Go, Rust)

* refactor: use dedicated response model for clear_memory_observations, remove code example from docs
2026-02-24 13:26:33 +01:00
Anton EvseevandClaude Opus 4.6 40b02645f4 fix(openclaw): pass auth token to health check endpoint (#427)
The checkExternalApiHealth function didn't include the Bearer token
in its requests. When the Hindsight API requires authentication
(HINDSIGHT_API_TENANT_API_KEY), health checks would fail with 401/403,
preventing plugin initialization.

Pass apiToken to all checkExternalApiHealth call sites and include
the Authorization header when a token is configured.

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-24 13:17:00 +01:00
Nicolò Boschi 5fddd9a79c feat: add reflect mode to LoComo benchmark and improve reflect agent (#428)
* feat: add reflect mode to LoComo benchmark and improve reflect agent

- Replace think mode with reflect mode in LoComo benchmark using reflect_async with Budget.HIGH
- Add --question-index CLI flag to run a single question by its index
- Track and display original question index in logs and visualizer
- Update visualizer to show reflect mode results

Reflect agent improvements:
- tool_recall: always fetch chunks (max_chunk_tokens=1000 min, non-optional)
- tool_search_observations: use include_source_facts=True instead of separate DB query
- Use model_dump() throughout to avoid manual error-prone dict conversion
- Enforce minimum 1000 tokens for max_tokens and max_chunk_tokens in _execute_tool
- Fix NoneType error when LLM passes null for mental_model_ids/observation_ids arrays
- Add non-conversational constraint to system prompt to prevent follow-up questions
- Fix recall_fn Callable type hint to include max_chunk_tokens parameter
- Fix main.py missing reranker_zeroentropy fields in HindsightConfig constructor

* fix: update tests for reflect tool API changes

- source_memory_ids -> source_fact_ids in test_search_observations (MemoryFact.model_dump() field name)
- Remove proof_count check (not in MemoryFact, was ObservationResult-specific)
- Remove max_results param from tool_recall call (no longer supported)
- Fix recall_result["count"] -> len(recall_result["memories"])
2026-02-24 09:48:23 +01:00
Nicolò Boschi 4d030707ad feat: enable bank config API by default (#426)
Change DEFAULT_ENABLE_BANK_CONFIG_API from false to true, update all docs,
error messages, and client docstrings to reflect the new default. Remove
explicit env var overrides in CI and tests that are no longer needed.
2026-02-24 08:52:42 +01:00
Chris Bartholomew 5fef54d501 Fix typos in README 2026-02-23 15:32:49 -05:00
Nicolò Boschi 2a32273226 feat: increase customization for reflect, retain and consolidation (#419) 2026-02-23 20:35:32 +01:00
Nicolò Boschi 87219b731d feat: include doc metadata in fact extraction (#424) 2026-02-23 11:31:02 +01:00
Nicolò Boschi 9f0c031df7 fix: improve memory footprint of recall (#423) 2026-02-23 11:30:16 +01:00
Chris Bartholomew 8b1a46585d Fix reflect based_on population and enforce full hierarchical retrieval (#421)
* Fix reflect based_on population and enforce full hierarchical retrieval

Problem 1: based_on field was incomplete
- search_observations results were never extracted into based_on, so
  observations used by the agent were invisible to callers
- search_mental_models and get_mental_model used non-existent fields
  (summary/description) instead of the actual content field, producing
  empty text in based_on entries
- A duplicate unreachable elif block for search_mental_models was dead
  code (the first identical condition always matched)

Problem 2: mental models could produce "I don't have information"
- When a bank has mental models, the agent's tool_choice forcing only
  covered iteration 0 (search_mental_models). Iterations 1+ were auto,
  allowing the LLM to short-circuit without ever searching observations
  or raw facts. Combined with the LOW budget prompt encouraging speed,
  this meant the agent would often stop after a single tool call.
- This created a self-reinforcing failure loop: if a mental model
  refresh produced "I don't have information" (e.g. due to the agent
  skipping recall), subsequent reflects would find that content and
  trust it, never searching deeper.

Fix: extend forced tool_choice to cover the full hierarchical retrieval
path before allowing auto mode:
- With mental models: search_mental_models(0) → search_observations(1)
  → recall(2) → auto(3+)
- Without mental models: search_observations(0) → recall(1) → auto(2+)

This matches the retrieval strategy documented in the system prompt and
ensures all three knowledge levels are always consulted. The agent still
has 2-3 auto iterations (with LOW budget, max_iterations=5) for
additional searches or calling done().

* Add Umami analytics tracking to docs site

Add conditional Umami script injection to docusaurus.config.ts and pass
UMAMI_URL/UMAMI_WEBSITE_ID env vars in the GitHub Pages deploy workflow.
The tracking script only loads when both env vars are set.
2026-02-23 10:13:19 +01:00
Eliah RusinandClaude Opus 4.6 172596751f feat: add ZeroEntropy reranker provider support (#420)
Add ZeroEntropy as a reranker provider using their Rerank API
(https://docs.zeroentropy.dev/models). Supports zerank-2 (flagship)
and zerank-2-small models via direct HTTP API calls with httpx (no
additional SDK dependency required).

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-23 10:12:32 +01:00
504 changed files with 29681 additions and 16577 deletions
+3
View File
@@ -31,6 +31,9 @@ jobs:
- run: npm ci --workspace=hindsight-docs
- run: uv run generate-llms-full
- run: npm run build --workspace=hindsight-docs
env:
UMAMI_URL: https://analytics.hindsight.vectorize.io
UMAMI_WEBSITE_ID: ${{ secrets.UMAMI_WEBSITE_ID }}
- uses: actions/upload-pages-artifact@v3
with:
path: hindsight-docs/build
+70 -1
View File
@@ -50,6 +50,10 @@ jobs:
working-directory: ./hindsight-integrations/crewai
run: uv build --out-dir dist
- name: Build hindsight-pydantic-ai
working-directory: ./hindsight-integrations/pydantic-ai
run: uv build --out-dir dist
# Publish in order (client and api first, then hindsight-all which depends on them)
- name: Publish hindsight-client to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
@@ -87,6 +91,12 @@ jobs:
packages-dir: ./hindsight-integrations/crewai/dist
skip-existing: true
- name: Publish hindsight-pydantic-ai to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: ./hindsight-integrations/pydantic-ai/dist
skip-existing: true
# Upload artifacts for GitHub release
- name: Upload artifacts
uses: actions/upload-artifact@v4
@@ -99,6 +109,7 @@ jobs:
hindsight-integrations/litellm/dist/*
hindsight-embed/dist/*
hindsight-integrations/crewai/dist/*
hindsight-integrations/pydantic-ai/dist/*
retention-days: 1
release-typescript-client:
@@ -248,6 +259,55 @@ jobs:
path: hindsight-integrations/ai-sdk/*.tgz
retention-days: 1
release-chat-integration:
runs-on: ubuntu-latest
environment: npm
steps:
- uses: actions/checkout@v4
- name: Set up Node.js
uses: actions/setup-node@v4
with:
node-version: '22'
registry-url: 'https://registry.npmjs.org'
- name: Install dependencies
working-directory: ./hindsight-integrations/chat
run: npm ci
- name: Build
working-directory: ./hindsight-integrations/chat
run: npm run build
- name: Publish to npm
working-directory: ./hindsight-integrations/chat
run: |
set +e
OUTPUT=$(npm publish --access public 2>&1)
EXIT_CODE=$?
echo "$OUTPUT"
if [ $EXIT_CODE -ne 0 ]; then
if echo "$OUTPUT" | grep -q "cannot publish over"; then
echo "Package version already published, skipping..."
exit 0
fi
exit $EXIT_CODE
fi
env:
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
- name: Pack for GitHub release
working-directory: ./hindsight-integrations/chat
run: npm pack
- name: Upload artifacts
uses: actions/upload-artifact@v4
with:
name: chat-integration
path: hindsight-integrations/chat/*.tgz
retention-days: 1
release-control-plane:
runs-on: ubuntu-latest
environment: npm
@@ -501,7 +561,7 @@ jobs:
create-github-release:
runs-on: ubuntu-latest
needs: [release-python-packages, release-typescript-client, release-openclaw-integration, release-ai-sdk-integration, release-control-plane, release-rust-cli, release-docker-images, release-helm-chart]
needs: [release-python-packages, release-typescript-client, release-openclaw-integration, release-ai-sdk-integration, release-chat-integration, release-control-plane, release-rust-cli, release-docker-images, release-helm-chart]
permissions:
contents: write
@@ -536,6 +596,12 @@ jobs:
name: ai-sdk-integration
path: ./artifacts/ai-sdk-integration
- name: Download Chat Integration
uses: actions/download-artifact@v4
with:
name: chat-integration
path: ./artifacts/chat-integration
- name: Download Control Plane
uses: actions/download-artifact@v4
with:
@@ -574,6 +640,7 @@ jobs:
cp artifacts/python-packages/hindsight-api/dist/* release-assets/ || true
cp artifacts/python-packages/hindsight/dist/* release-assets/ || true
cp artifacts/python-packages/hindsight-integrations/litellm/dist/* release-assets/ || true
cp artifacts/python-packages/hindsight-integrations/pydantic-ai/dist/* release-assets/ || true
cp artifacts/python-packages/hindsight-embed/dist/* release-assets/ || true
# TypeScript client
cp artifacts/typescript-client/*.tgz release-assets/ || true
@@ -581,6 +648,8 @@ jobs:
cp artifacts/openclaw-integration/*.tgz release-assets/ || true
# AI SDK Integration
cp artifacts/ai-sdk-integration/*.tgz release-assets/ || true
# Chat Integration
cp artifacts/chat-integration/*.tgz release-assets/ || true
# Control Plane
cp artifacts/control-plane/*.tgz release-assets/ || true
# Rust CLI binaries
+52
View File
@@ -97,6 +97,29 @@ jobs:
working-directory: ./hindsight-integrations/ai-sdk
run: npm run build
build-chat-integration:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Set up Node.js
uses: actions/setup-node@v4
with:
node-version: '22'
- name: Install dependencies
working-directory: ./hindsight-integrations/chat
run: npm ci
- name: Run tests
working-directory: ./hindsight-integrations/chat
run: npm test
- name: Build
working-directory: ./hindsight-integrations/chat
run: npm run build
build-control-plane:
runs-on: ubuntu-latest
@@ -1139,6 +1162,35 @@ jobs:
working-directory: ./hindsight-integrations/litellm
run: uv run pytest tests -v
test-pydantic-ai-integration:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Install uv
uses: astral-sh/setup-uv@v5
with:
enable-cache: true
prune-cache: false
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version-file: ".python-version"
- name: Build pydantic-ai integration
working-directory: ./hindsight-integrations/pydantic-ai
run: uv build
- name: Install dependencies
working-directory: ./hindsight-integrations/pydantic-ai
run: uv sync --frozen
- name: Run tests
working-directory: ./hindsight-integrations/pydantic-ai
run: uv run pytest tests -v
test-embed:
runs-on: ubuntu-latest
env:
+1 -1
View File
@@ -323,4 +323,4 @@ Optional (uses local models by default):
- `HINDSIGHT_API_EMBEDDINGS_PROVIDER`: local (default) or tei
- `HINDSIGHT_API_RERANKER_PROVIDER`: local (default) or tei
- `HINDSIGHT_API_DATABASE_URL`: External PostgreSQL (uses embedded pg0 by default)
- `HINDSIGHT_API_ENABLE_BANK_CONFIG_API`: Enable per-bank config API (default: false, disabled for security)
- `HINDSIGHT_API_ENABLE_BANK_CONFIG_API`: Enable per-bank config API (default: true)
+4 -2
View File
@@ -36,7 +36,7 @@ Hindsight is being used in production at Fortune 500 enterprises and by a growin
## Adding Hindsight to Your AI Agents
The easiest way use Hindsight with an existing agent is with the LLM Wrapper. You can add memory to your agent with 2 lines of code. That will swap your current LLM client out with the Hindsight wrapper. After that, memories will be stored and retrieved automatically as you make LLM calls.
The easiest way to use Hindsight with an existing agent is with the LLM Wrapper. You can add memory to your agent with 2 lines of code. That will swap your current LLM client out with the Hindsight wrapper. After that, memories will be stored and retrieved automatically as you make LLM calls.
If you need more control over how and when your agent stores and recalls memories, there's also a simple API you can integrate with using the SDKs or directly via HTTP.
@@ -181,7 +181,7 @@ Satisfying these requirements in Hindsight is straightforward. When new user inp
![Overview](./hindsight-docs/static/img/hindsight-overview.webp)
Most agent memory implementation rely on basic vector search or sometimes use a knowledge graph. Hindsight uses biomimetic data structures to organize agent memories in a way that is more like how human memory works:
Most agent memory implementations rely on basic vector search or sometimes use a knowledge graph. Hindsight uses biomimetic data structures to organize agent memories in a way that is more like how human memory works:
- **World:** Facts about the world ("The stove gets hot")
- **Experiences:** Agent's own experiences ("I touched the stove and it really hurt")
@@ -307,3 +307,5 @@ MIT — see [LICENSE](./LICENSE)
---
Built by [Vectorize.io](https://vectorize.io)
<img src="https://umami-pixel.chris-latimer.workers.dev/?id=a8b043e6-6964-454d-80df-69b69d3f0d50&host=github.com&url=/vectorize-io/hindsight" width="1" height="1" alt="" />
+2 -2
View File
@@ -97,7 +97,7 @@ fi
if [ "$ENABLE_CP" = "true" ]; then
echo "🎛️ Starting Control Plane..."
cd /app/control-plane
PORT=9999 node server.js &
PORT="${HINDSIGHT_CP_PORT:-9999}" node server.js &
CP_PID=$!
PIDS+=($CP_PID)
else
@@ -110,7 +110,7 @@ echo "✅ Hindsight is running!"
echo ""
echo "📍 Access:"
if [ "$ENABLE_CP" = "true" ]; then
echo " Control Plane: http://localhost:9999"
echo " Control Plane: http://localhost:${HINDSIGHT_CP_PORT:-9999}"
fi
if [ "$ENABLE_API" = "true" ]; then
echo " API: http://localhost:8888"
+2 -2
View File
@@ -2,8 +2,8 @@ apiVersion: v2
name: hindsight
description: Hindsight helm chart
type: application
version: 0.4.13
appVersion: "0.4.13"
version: 0.4.15
appVersion: "0.4.15"
keywords:
- ai
- memory
+1 -1
View File
@@ -46,4 +46,4 @@ __all__ = [
"RemoteTEICrossEncoder",
"LLMConfig",
]
__version__ = "0.4.13"
__version__ = "0.4.15"
@@ -0,0 +1,88 @@
"""Add text_signals column to memory_units for enriched BM25 indexing.
text_signals stores a denormalized space-separated string of entity names
(and future signals) to improve full-text search recall without polluting
the stored fact text.
- vchord: text_signals included in tokenize() at insert time
- native: search_vector GENERATED column regenerated to include text_signals
- pg_textsearch: no change (index only supports a single base column)
Revision ID: a2b3c4d5e6f7
Revises: z1u2v3w4x5y6
Create Date: 2026-02-28
"""
import os
from collections.abc import Sequence
from alembic import context, op
revision: str = "a2b3c4d5e6f7"
down_revision: str | Sequence[str] | None = "aa2b3c4d5e6f"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def _detect_text_search_extension() -> str:
return os.getenv("HINDSIGHT_API_TEXT_SEARCH_EXTENSION", "native").lower()
def upgrade() -> None:
schema = _get_schema_prefix()
table = f"{schema}memory_units"
text_search_ext = _detect_text_search_extension()
# Add text_signals column (nullable TEXT, populated at retain time)
op.execute(f"ALTER TABLE {table} ADD COLUMN IF NOT EXISTS text_signals TEXT")
if text_search_ext == "native":
# Native PostgreSQL: drop and recreate the GENERATED tsvector column to include text_signals
op.execute(f"ALTER TABLE {table} DROP COLUMN IF EXISTS search_vector")
op.execute(f"""
ALTER TABLE {table}
ADD COLUMN search_vector tsvector
GENERATED ALWAYS AS (
to_tsvector('english',
COALESCE(text, '') || ' ' ||
COALESCE(context, '') || ' ' ||
COALESCE(text_signals, '')
)
) STORED
""")
# Recreate GIN index (was dropped with the column)
op.execute(f"""
CREATE INDEX IF NOT EXISTS idx_memory_units_text_search
ON {table} USING gin(search_vector)
""")
# vchord: tokenize() call in fact_storage.py is updated to include text_signals at insert time
# pg_textsearch: no change — index operates on the base `text` column only
def downgrade() -> None:
schema = _get_schema_prefix()
table = f"{schema}memory_units"
text_search_ext = _detect_text_search_extension()
if text_search_ext == "native":
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_text_search")
op.execute(f"ALTER TABLE {table} DROP COLUMN IF EXISTS search_vector")
op.execute(f"""
ALTER TABLE {table}
ADD COLUMN search_vector tsvector
GENERATED ALWAYS AS (
to_tsvector('english', COALESCE(text, '') || ' ' || COALESCE(context, ''))
) STORED
""")
op.execute(f"""
CREATE INDEX idx_memory_units_text_search
ON {table} USING gin(search_vector)
""")
op.execute(f"ALTER TABLE {table} DROP COLUMN IF EXISTS text_signals")
@@ -0,0 +1,36 @@
"""Make event_date nullable in memory_units to support timestamp-free content
Revision ID: aa2b3c4d5e6f
Revises: z1u2v3w4x5y6
Create Date: 2026-03-02
When callers retain content without a timestamp (e.g. fictional documents, static text),
the event_date column should be allowed to be NULL rather than defaulting to utcnow().
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "aa2b3c4d5e6f"
down_revision: str | Sequence[str] | None = "z1u2v3w4x5y6"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
"""Get schema prefix for table names (required for multi-tenant support)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"ALTER TABLE {schema}memory_units ALTER COLUMN event_date DROP NOT NULL")
def downgrade() -> None:
schema = _get_schema_prefix()
# Backfill NULLs with now() before restoring the NOT NULL constraint
op.execute(f"UPDATE {schema}memory_units SET event_date = now() WHERE event_date IS NULL")
op.execute(f"ALTER TABLE {schema}memory_units ALTER COLUMN event_date SET NOT NULL")
@@ -0,0 +1,68 @@
"""Add partial indexes on memory_units temporal date fields for fast temporal retrieval
Revision ID: b3c4d5e6f7g8
Revises: c1a2b3d4e5f6
Create Date: 2026-03-02
The temporal retrieval entry-point query filters memory_units by occurred_start,
occurred_end, and mentioned_at using OR conditions. Without dedicated indexes the
planner falls back to a sequential scan of all bank rows after applying the
(bank_id, fact_type) index, then re-checks each date field.
These three partial indexes give the planner bitmap-index scan options for the
three most common date predicates, dramatically reducing the row set before any
embedding computation is required.
All indexes are created CONCURRENTLY so the migration does not block writes on
memory_units during production deployments. CONCURRENTLY requires running outside
a transaction block; see migrations.py for how this is handled safely.
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "b3c4d5e6f7g8"
down_revision: str | Sequence[str] | None = "c1a2b3d4e5f6"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
schema = _get_schema_prefix()
# Partial index on occurred_start (covers "occurred_start BETWEEN $4 AND $5")
op.execute("COMMIT")
op.execute(
f"CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_memory_units_bank_occurred_start "
f"ON {schema}memory_units(bank_id, fact_type, occurred_start) "
f"WHERE occurred_start IS NOT NULL"
)
# Partial index on occurred_end (covers "occurred_end BETWEEN $4 AND $5")
op.execute("COMMIT")
op.execute(
f"CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_memory_units_bank_occurred_end "
f"ON {schema}memory_units(bank_id, fact_type, occurred_end) "
f"WHERE occurred_end IS NOT NULL"
)
# Partial index on mentioned_at (covers "mentioned_at BETWEEN $4 AND $5")
op.execute("COMMIT")
op.execute(
f"CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_memory_units_bank_mentioned_at "
f"ON {schema}memory_units(bank_id, fact_type, mentioned_at) "
f"WHERE mentioned_at IS NOT NULL"
)
def downgrade() -> None:
schema = _get_schema_prefix()
op.execute("COMMIT")
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}idx_memory_units_bank_mentioned_at")
op.execute("COMMIT")
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}idx_memory_units_bank_occurred_end")
op.execute("COMMIT")
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}idx_memory_units_bank_occurred_start")
@@ -0,0 +1,34 @@
"""Backfill observation_scopes column if missing.
This migration ensures observation_scopes exists even on databases that had
revision z1u2v3w4x5y6 applied when it referred to the old text_signals migration
(before it was renamed to a2b3c4d5e6f7). The ADD COLUMN IF NOT EXISTS makes this
a no-op on databases that already have the column.
Revision ID: b4c5d6e7f8a9
Revises: a2b3c4d5e6f7
Create Date: 2026-03-02
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "b4c5d6e7f8a9"
down_revision: str | Sequence[str] | None = "a2b3c4d5e6f7"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"ALTER TABLE {schema}memory_units ADD COLUMN IF NOT EXISTS observation_scopes JSONB")
def downgrade() -> None:
pass # intentionally no-op — safe to leave the column in place
@@ -0,0 +1,46 @@
"""Enable pg_trgm extension and add GIN trigram index on entities.canonical_name
Revision ID: c1a2b3d4e5f6
Revises: b4c5d6e7f8a9
Create Date: 2026-03-02
Index is created CONCURRENTLY so the migration does not block writes on entities
during production deployments. CONCURRENTLY requires running outside a transaction
block; see migrations.py for how this is handled safely.
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "c1a2b3d4e5f6"
down_revision: str | Sequence[str] | None = "b4c5d6e7f8a9"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
# pg_trgm ships with every standard PostgreSQL installation as a contrib module.
# It enables fast similarity lookups via GIN indexes, used for entity name matching.
op.execute("CREATE EXTENSION IF NOT EXISTS pg_trgm")
schema = _get_schema_prefix()
# GIN index on canonical_name enables sub-millisecond trigram similarity queries
# (% operator, similarity()) instead of full-table scans across all bank entities.
op.execute("COMMIT")
op.execute(
f"CREATE INDEX CONCURRENTLY IF NOT EXISTS entities_canonical_name_trgm_idx "
f"ON {schema}entities USING GIN (canonical_name gin_trgm_ops)"
)
def downgrade() -> None:
schema = _get_schema_prefix()
op.execute("COMMIT")
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}entities_canonical_name_trgm_idx")
# Note: not dropping pg_trgm extension as other indexes may depend on it
@@ -0,0 +1,83 @@
"""Add covering and composite indexes to speed up link expansion graph retrieval.
Two indexes target the two bottlenecks identified by EXPLAIN ANALYZE on a 17M-row
memory_links table:
1. idx_memory_links_to_type_weight (to_unit_id, link_type, weight DESC)
The semantic incoming direction — finding facts that consider seeds as their
nearest neighbour — currently hits an expensive BitmapAnd of two separate
bitmap scans (to_unit_id bitmap ∩ link_type bitmap). A composite index
on (to_unit_id, link_type) turns this into a single index scan and reduces
latency from ~36 ms to < 5 ms per query.
2. idx_memory_links_entity_covering (from_unit_id) INCLUDE (to_unit_id, entity_id)
WHERE link_type = 'entity'
The entity co-occurrence expansion uses COUNT(DISTINCT ml.entity_id) and
joins on ml.to_unit_id. Without a covering index the planner must read
~2 500 heap pages to fetch entity_id and to_unit_id after the bitmap index
scan, adding ~230 ms of random I/O. INCLUDE adds those two columns to the
index leaf pages so the entire query can be served from the index (index-only
scan), eliminating the heap reads entirely.
Partial index (WHERE link_type = 'entity') keeps index size ~40 % smaller.
Both indexes are created with CONCURRENTLY so the migration does not block
concurrent reads or writes on memory_links. CONCURRENTLY requires running
outside a transaction block, so the migration emits an explicit COMMIT before
each statement and uses IF NOT EXISTS for idempotency.
Revision ID: d2e3f4a5b6c7
Revises: b3c4d5e6f7g8
Create Date: 2026-03-02
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "d2e3f4a5b6c7"
down_revision: str | Sequence[str] | None = "b3c4d5e6f7g8"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
schema = _get_schema_prefix()
# CREATE INDEX CONCURRENTLY cannot run inside a transaction block.
# Commit the current Alembic transaction, then issue each CONCURRENTLY
# statement in its own implicit autocommit transaction.
# IF NOT EXISTS makes each statement idempotent if the migration is retried.
# Index for the semantic *incoming* direction in link_expansion_retrieval.py.
# Replaces the BitmapAnd of idx_memory_links_to_unit ∩ idx_memory_links_link_type
# with a single composite index scan.
op.execute("COMMIT")
op.execute(
f"CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_memory_links_to_type_weight "
f"ON {schema}memory_links(to_unit_id, link_type, weight DESC)"
)
# Covering index for entity co-occurrence expansion.
# Enables an index-only scan: entity_id and to_unit_id are read from the
# index leaf pages instead of the heap, eliminating ~2 500 random heap-page
# reads per expansion query.
op.execute("COMMIT")
op.execute(
f"CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_memory_links_entity_covering "
f"ON {schema}memory_links(from_unit_id) "
f"INCLUDE (to_unit_id, entity_id) "
f"WHERE link_type = 'entity'"
)
def downgrade() -> None:
schema = _get_schema_prefix()
op.execute("COMMIT")
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}idx_memory_links_entity_covering")
op.execute("COMMIT")
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}idx_memory_links_to_type_weight")
@@ -0,0 +1,35 @@
"""Add observation_scopes column to memory_units table
Revision ID: z1u2v3w4x5y6
Revises: a1b2c3d4e5f6
Create Date: 2026-02-25
Adds observation_scopes JSONB column to memory_units to control how observations
are scoped during consolidation. Accepts "per_tag", "combined", or an explicit
list of tag-set lists for custom multi-pass consolidation.
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "z1u2v3w4x5y6"
down_revision: str | Sequence[str] | None = "a1b2c3d4e5f6"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
"""Get schema prefix for table names (required for multi-tenant support)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"ALTER TABLE {schema}memory_units ADD COLUMN IF NOT EXISTS observation_scopes JSONB")
def downgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS observation_scopes")
+325 -171
View File
@@ -383,7 +383,15 @@ class MemoryItem(BaseModel):
)
content: str
timestamp: datetime | None = None
timestamp: datetime | str | None = Field(
default=None,
description=(
"When the content occurred. "
"Accepts an ISO 8601 datetime string (e.g. '2024-01-15T10:30:00Z'), null/omitted (defaults to now), "
"or the special string 'unset' to explicitly store without any timestamp "
"(use this for timeless content such as fictional documents or static reference material)."
),
)
context: str | None = None
metadata: dict[str, str] | None = None
document_id: str | None = Field(default=None, description="Optional document ID for this memory item.")
@@ -395,6 +403,16 @@ class MemoryItem(BaseModel):
default=None,
description="Optional tags for visibility scoping. Memories with tags can be filtered during recall.",
)
observation_scopes: Literal["per_tag", "combined", "all_combinations"] | list[list[str]] | None = Field(
default=None,
title="ObservationScopes",
description=(
"How to scope observations during consolidation. "
"'per_tag' runs one consolidation pass per individual tag, creating separate observations for each tag. "
"'combined' (default) runs a single pass with all tags together. "
"A list of tag lists runs one pass per inner list, giving full control over which combinations to use."
),
)
@field_validator("timestamp", mode="before")
@classmethod
@@ -404,12 +422,14 @@ class MemoryItem(BaseModel):
if isinstance(v, datetime):
return v
if isinstance(v, str):
if v.lower() == "unset":
return "unset"
try:
# Try parsing as ISO format
return datetime.fromisoformat(v.replace("Z", "+00:00"))
except ValueError as e:
raise ValueError(
f"Invalid timestamp/event_date format: '{v}'. Expected ISO format like '2024-01-15T10:30:00' or '2024-01-15T10:30:00Z'"
f"Invalid timestamp/event_date format: '{v}'. Expected ISO format like '2024-01-15T10:30:00' or '2024-01-15T10:30:00Z', or the special value 'unset' to store without a timestamp."
) from e
raise ValueError(f"timestamp must be a string or datetime, got {type(v).__name__}")
@@ -879,18 +899,103 @@ class CreateBankRequest(BaseModel):
model_config = ConfigDict(
json_schema_extra={
"example": {
"name": "Alice",
"disposition": {"skepticism": 3, "literalism": 3, "empathy": 3},
"mission": "I am a PM helping my engineering team stay organized",
"retain_mission": "Always include technical decisions and architectural trade-offs. Ignore meeting logistics.",
"observations_mission": "Observations are stable facts about people and projects. Always include preferences and skills.",
}
}
)
name: str | None = None
disposition: DispositionTraits | None = None
mission: str | None = Field(default=None, description="The agent's mission")
# Deprecated: use mission instead
background: str | None = Field(default=None, description="Deprecated: use mission instead")
# Deprecated fields — kept for backwards compatibility only
name: str | None = Field(default=None, description="Deprecated: display label only, not advertised")
disposition: DispositionTraits | None = Field(
default=None, description="Deprecated: use update_bank_config instead"
)
disposition_skepticism: int | None = Field(
default=None, ge=1, le=5, description="Deprecated: use update_bank_config instead"
)
disposition_literalism: int | None = Field(
default=None, ge=1, le=5, description="Deprecated: use update_bank_config instead"
)
disposition_empathy: int | None = Field(
default=None, ge=1, le=5, description="Deprecated: use update_bank_config instead"
)
# Deprecated: use update_bank_config with reflect_mission instead
mission: str | None = Field(
default=None, description="Deprecated: use update_bank_config with reflect_mission instead"
)
# Deprecated alias for mission
background: str | None = Field(
default=None, description="Deprecated: use update_bank_config with reflect_mission instead"
)
# Reflect configuration
reflect_mission: str | None = Field(
default=None,
description="Mission/context for Reflect operations. Guides how Reflect interprets and uses memories.",
)
# Operational configuration (applied via config resolver)
retain_mission: str | None = Field(
default=None,
description="Steers what gets extracted during retain(). Injected alongside built-in extraction rules.",
)
retain_extraction_mode: str | None = Field(
default=None,
description="Fact extraction mode: 'concise' (default), 'verbose', or 'custom'.",
)
retain_custom_instructions: str | None = Field(
default=None,
description="Custom extraction prompt. Only active when retain_extraction_mode is 'custom'.",
)
retain_chunk_size: int | None = Field(
default=None,
description="Maximum token size for each content chunk during retain.",
)
enable_observations: bool | None = Field(
default=None,
description="Toggle automatic observation consolidation after retain().",
)
observations_mission: str | None = Field(
default=None,
description="Controls what gets synthesised into observations. Replaces built-in consolidation rules entirely.",
)
def get_config_updates(self) -> dict[str, Any]:
"""Return only the config fields that were explicitly set.
reflect_mission takes precedence over deprecated mission/background aliases.
Individual disposition_* fields take priority over the deprecated disposition dict.
"""
updates: dict[str, Any] = {}
# Resolve reflect mission: reflect_mission (new) > mission (deprecated) > background (deprecated)
resolved_reflect_mission = self.reflect_mission or self.mission or self.background
if resolved_reflect_mission is not None:
updates["reflect_mission"] = resolved_reflect_mission
# Disposition: individual fields take priority over legacy disposition dict
if self.disposition_skepticism is not None:
updates["disposition_skepticism"] = self.disposition_skepticism
elif self.disposition is not None:
updates["disposition_skepticism"] = self.disposition.skepticism
if self.disposition_literalism is not None:
updates["disposition_literalism"] = self.disposition_literalism
elif self.disposition is not None:
updates["disposition_literalism"] = self.disposition.literalism
if self.disposition_empathy is not None:
updates["disposition_empathy"] = self.disposition_empathy
elif self.disposition is not None:
updates["disposition_empathy"] = self.disposition.empathy
for field_name in (
"retain_mission",
"retain_extraction_mode",
"retain_custom_instructions",
"retain_chunk_size",
"enable_observations",
"observations_mission",
):
value = getattr(self, field_name)
if value is not None:
updates[field_name] = value
return updates
class BankConfigUpdate(BaseModel):
@@ -1150,6 +1255,14 @@ class DeleteResponse(BaseModel):
deleted_count: int | None = None
class ClearMemoryObservationsResponse(BaseModel):
"""Response model for clearing observations for a specific memory."""
model_config = ConfigDict(json_schema_extra={"example": {"deleted_count": 3}})
deleted_count: int
class BankStatsResponse(BaseModel):
"""Response model for bank statistics endpoint."""
@@ -1717,6 +1830,12 @@ def create_app(
app.include_router(extension_router, prefix="/ext", tags=["Extension"])
logging.info("HTTP extension router mounted at /ext/")
# Mount root router if provided (for well-known endpoints, etc.)
root_router = http_extension.get_root_router(memory)
if root_router:
app.include_router(root_router)
logging.info("HTTP extension root router mounted")
return app
@@ -1829,12 +1948,19 @@ def _register_routes(app: FastAPI):
bank_id: str,
type: str | None = None,
limit: int = 1000,
q: str | None = None,
tags: list[str] | None = Query(None),
tags_match: str = "all_strict",
request_context: RequestContext = Depends(get_request_context),
):
"""Get graph data from database, filtered by bank_id and optionally by type."""
try:
data = await app.state.memory.get_graph_data(bank_id, type, limit=limit, request_context=request_context)
data = await app.state.memory.get_graph_data(
bank_id, type, limit=limit, q=q, tags=tags, tags_match=tags_match, request_context=request_context
)
return data
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -1883,6 +2009,8 @@ def _register_routes(app: FastAPI):
request_context=request_context,
)
return data
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -1914,6 +2042,8 @@ def _register_routes(app: FastAPI):
if data is None:
raise HTTPException(status_code=404, detail=f"Memory unit '{memory_id}' not found")
return data
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2255,141 +2385,34 @@ def _register_routes(app: FastAPI):
):
"""Get statistics about memory nodes and links for a memory bank."""
try:
# Authenticate and set tenant schema
await app.state.memory._authenticate_tenant(request_context)
pool = await app.state.memory._get_pool()
async with acquire_with_retry(pool) as conn:
# Get node counts by fact_type
node_stats = await conn.fetch(
f"""
SELECT fact_type, COUNT(*) as count
FROM {fq_table("memory_units")}
WHERE bank_id = $1
GROUP BY fact_type
""",
bank_id,
)
# Get link counts by link_type
link_stats = await conn.fetch(
f"""
SELECT ml.link_type, COUNT(*) as count
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.from_unit_id = mu.id
WHERE mu.bank_id = $1
GROUP BY ml.link_type
""",
bank_id,
)
# Get link counts by fact_type (from nodes)
link_fact_type_stats = await conn.fetch(
f"""
SELECT mu.fact_type, COUNT(*) as count
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.from_unit_id = mu.id
WHERE mu.bank_id = $1
GROUP BY mu.fact_type
""",
bank_id,
)
# Get link counts by fact_type AND link_type
link_breakdown_stats = await conn.fetch(
f"""
SELECT mu.fact_type, ml.link_type, COUNT(*) as count
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.from_unit_id = mu.id
WHERE mu.bank_id = $1
GROUP BY mu.fact_type, ml.link_type
""",
bank_id,
)
# Get pending and failed operations counts
ops_stats = await conn.fetch(
f"""
SELECT status, COUNT(*) as count
FROM {fq_table("async_operations")}
WHERE bank_id = $1
GROUP BY status
""",
bank_id,
)
ops_by_status = {row["status"]: row["count"] for row in ops_stats}
pending_operations = ops_by_status.get("pending", 0)
failed_operations = ops_by_status.get("failed", 0)
# Get document count
doc_count_result = await conn.fetchrow(
f"""
SELECT COUNT(*) as count
FROM {fq_table("documents")}
WHERE bank_id = $1
""",
bank_id,
)
total_documents = doc_count_result["count"] if doc_count_result else 0
# Get consolidation stats from memory-level tracking
consolidation_stats = await conn.fetchrow(
f"""
SELECT
MAX(consolidated_at) as last_consolidated_at,
COUNT(*) FILTER (WHERE consolidated_at IS NULL AND fact_type IN ('experience', 'world')) as pending
FROM {fq_table("memory_units")}
WHERE bank_id = $1
""",
bank_id,
)
last_consolidated_at = consolidation_stats["last_consolidated_at"] if consolidation_stats else None
pending_consolidation = consolidation_stats["pending"] if consolidation_stats else 0
# Count total observations (consolidated knowledge)
observation_count_result = await conn.fetchrow(
f"""
SELECT COUNT(*) as count
FROM {fq_table("memory_units")}
WHERE bank_id = $1 AND fact_type = 'observation'
""",
bank_id,
)
total_observations = observation_count_result["count"] if observation_count_result else 0
# Format results
nodes_by_type = {row["fact_type"]: row["count"] for row in node_stats}
links_by_type = {row["link_type"]: row["count"] for row in link_stats}
links_by_fact_type = {row["fact_type"]: row["count"] for row in link_fact_type_stats}
# Build detailed breakdown: {fact_type: {link_type: count}}
links_breakdown = {}
for row in link_breakdown_stats:
fact_type = row["fact_type"]
link_type = row["link_type"]
count = row["count"]
if fact_type not in links_breakdown:
links_breakdown[fact_type] = {}
links_breakdown[fact_type][link_type] = count
total_nodes = sum(nodes_by_type.values())
total_links = sum(links_by_type.values())
return BankStatsResponse(
bank_id=bank_id,
total_nodes=total_nodes,
total_links=total_links,
total_documents=total_documents,
nodes_by_fact_type=nodes_by_type,
links_by_link_type=links_by_type,
links_by_fact_type=links_by_fact_type,
links_breakdown=links_breakdown,
pending_operations=pending_operations,
failed_operations=failed_operations,
last_consolidated_at=(last_consolidated_at.isoformat() if last_consolidated_at else None),
pending_consolidation=pending_consolidation,
total_observations=total_observations,
)
stats = await app.state.memory.get_bank_stats(bank_id, request_context=request_context)
nodes_by_type = stats["node_counts"]
links_by_type = stats["link_counts"]
links_by_fact_type = stats["link_counts_by_fact_type"]
links_breakdown: dict[str, dict[str, int]] = {}
for row in stats["link_breakdown"]:
ft = row["fact_type"]
if ft not in links_breakdown:
links_breakdown[ft] = {}
links_breakdown[ft][row["link_type"]] = row["count"]
ops = stats["operations"]
return BankStatsResponse(
bank_id=bank_id,
total_nodes=sum(nodes_by_type.values()),
total_links=sum(links_by_type.values()),
total_documents=stats["total_documents"],
nodes_by_fact_type=nodes_by_type,
links_by_link_type=links_by_type,
links_by_fact_type=links_by_fact_type,
links_breakdown=links_breakdown,
pending_operations=ops.get("pending", 0),
failed_operations=ops.get("failed", 0),
last_consolidated_at=stats["last_consolidated_at"],
pending_consolidation=stats["pending_consolidation"],
total_observations=stats["total_observations"],
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2424,6 +2447,8 @@ def _register_routes(app: FastAPI):
limit=data["limit"],
offset=data["offset"],
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2463,6 +2488,8 @@ def _register_routes(app: FastAPI):
for obs in entity["observations"]
],
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2524,6 +2551,8 @@ def _register_routes(app: FastAPI):
request_context=request_context,
)
return MentalModelListResponse(items=[MentalModelResponse(**m) for m in mental_models])
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2682,6 +2711,8 @@ def _register_routes(app: FastAPI):
if mental_model is None:
raise HTTPException(status_code=404, detail=f"Mental model '{mental_model_id}' not found")
return MentalModelResponse(**mental_model)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2713,6 +2744,8 @@ def _register_routes(app: FastAPI):
if not deleted:
raise HTTPException(status_code=404, detail=f"Mental model '{mental_model_id}' not found")
return {"status": "deleted"}
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2755,6 +2788,8 @@ def _register_routes(app: FastAPI):
request_context=request_context,
)
return DirectiveListResponse(items=[DirectiveResponse(**d) for d in directives])
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2787,6 +2822,8 @@ def _register_routes(app: FastAPI):
if directive is None:
raise HTTPException(status_code=404, detail=f"Directive '{directive_id}' not found")
return DirectiveResponse(**directive)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2823,6 +2860,8 @@ def _register_routes(app: FastAPI):
return DirectiveResponse(**directive)
except ValueError as e:
raise HTTPException(status_code=400, 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:
@@ -2861,6 +2900,8 @@ def _register_routes(app: FastAPI):
if directive is None:
raise HTTPException(status_code=404, detail=f"Directive '{directive_id}' not found")
return DirectiveResponse(**directive)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2892,6 +2933,8 @@ def _register_routes(app: FastAPI):
if not deleted:
raise HTTPException(status_code=404, detail=f"Directive '{directive_id}' not found")
return {"status": "deleted"}
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2911,7 +2954,13 @@ def _register_routes(app: FastAPI):
)
async def api_list_documents(
bank_id: str,
q: str | None = None,
q: str | None = Query(
None, description="Case-insensitive substring filter on document ID (e.g. 'report' matches 'report-2024')"
),
tags: list[str] | None = Query(None, description="Filter documents by tags"),
tags_match: str = Query(
"any_strict", description="How to match tags: 'any', 'all', 'any_strict', 'all_strict'"
),
limit: int = 100,
offset: int = 0,
request_context: RequestContext = Depends(get_request_context),
@@ -2921,15 +2970,25 @@ def _register_routes(app: FastAPI):
Args:
bank_id: Memory Bank ID (from path)
q: Search query (searches document ID and metadata)
q: Case-insensitive substring filter on document ID
tags: Filter documents by tags
tags_match: How to match tags (any, all, any_strict, all_strict)
limit: Maximum number of results (default: 100)
offset: Offset for pagination (default: 0)
"""
try:
data = await app.state.memory.list_documents(
bank_id=bank_id, search_query=q, limit=limit, offset=offset, request_context=request_context
bank_id=bank_id,
search_query=q,
tags=tags,
tags_match=tags_match,
limit=limit,
offset=offset,
request_context=request_context,
)
return data
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -2962,6 +3021,8 @@ def _register_routes(app: FastAPI):
if not document:
raise HTTPException(status_code=404, detail="Document not found")
return document
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3015,6 +3076,8 @@ def _register_routes(app: FastAPI):
request_context=request_context,
)
return data
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3044,6 +3107,8 @@ def _register_routes(app: FastAPI):
if not chunk:
raise HTTPException(status_code=404, detail="Chunk not found")
return chunk
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3088,6 +3153,8 @@ def _register_routes(app: FastAPI):
document_id=document_id,
memory_units_deleted=result["memory_units_deleted"],
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3124,6 +3191,8 @@ def _register_routes(app: FastAPI):
offset=offset,
operations=[OperationResponse(**op) for op in result["operations"]],
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3155,6 +3224,8 @@ def _register_routes(app: FastAPI):
result = await app.state.memory.get_operation_status(bank_id, operation_id, request_context=request_context)
return OperationStatusResponse(**result)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3187,6 +3258,8 @@ def _register_routes(app: FastAPI):
return CancelOperationResponse(**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:
@@ -3203,6 +3276,7 @@ def _register_routes(app: FastAPI):
description="Get disposition traits and mission for a memory bank. Auto-creates agent with defaults if not exists.",
operation_id="get_bank_profile",
tags=["Banks"],
deprecated=True,
)
async def api_get_bank_profile(bank_id: str, request_context: RequestContext = Depends(get_request_context)):
"""Get memory bank profile (disposition + mission)."""
@@ -3222,6 +3296,8 @@ def _register_routes(app: FastAPI):
mission=mission,
background=mission, # Backwards compat
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3238,6 +3314,7 @@ def _register_routes(app: FastAPI):
description="Update bank's disposition traits (skepticism, literalism, empathy)",
operation_id="update_bank_disposition",
tags=["Banks"],
deprecated=True,
)
async def api_update_bank_disposition(
bank_id: str, request: UpdateDispositionRequest, request_context: RequestContext = Depends(get_request_context)
@@ -3264,6 +3341,8 @@ def _register_routes(app: FastAPI):
mission=mission,
background=mission, # Backwards compat
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3292,6 +3371,8 @@ def _register_routes(app: FastAPI):
)
mission = result.get("mission") or ""
return BackgroundResponse(mission=mission, background=mission)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3317,21 +3398,18 @@ def _register_routes(app: FastAPI):
# Ensure bank exists by getting profile (auto-creates with defaults)
await app.state.memory.get_bank_profile(bank_id, request_context=request_context)
# Update name and/or mission if provided (support both mission and deprecated background)
mission_value = request.mission or request.background
if request.name is not None or mission_value is not None:
# Update name if provided (stored in DB for display only, deprecated)
if request.name is not None:
await app.state.memory.update_bank(
bank_id,
name=request.name,
mission=mission_value,
request_context=request_context,
)
# Update disposition if provided
if request.disposition is not None:
await app.state.memory.update_bank_disposition(
bank_id, request.disposition.model_dump(), request_context=request_context
)
# Apply all config overrides (includes reflect_mission, disposition, retain settings)
config_updates = request.get_config_updates()
if config_updates:
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)
@@ -3348,6 +3426,8 @@ def _register_routes(app: FastAPI):
mission=mission,
background=mission, # Backwards compat
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3373,21 +3453,18 @@ def _register_routes(app: FastAPI):
# Ensure bank exists
await app.state.memory.get_bank_profile(bank_id, request_context=request_context)
# Update name and/or mission if provided
mission_value = request.mission or request.background
if request.name is not None or mission_value is not None:
# Update name if provided (stored in DB for display only, deprecated)
if request.name is not None:
await app.state.memory.update_bank(
bank_id,
name=request.name,
mission=mission_value,
request_context=request_context,
)
# Update disposition if provided
if request.disposition is not None:
await app.state.memory.update_bank_disposition(
bank_id, request.disposition.model_dump(), request_context=request_context
)
# Apply all config overrides (includes reflect_mission, disposition, retain settings)
config_updates = request.get_config_updates()
if config_updates:
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)
@@ -3404,6 +3481,8 @@ def _register_routes(app: FastAPI):
mission=mission,
background=mission, # Backwards compat
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3433,6 +3512,8 @@ def _register_routes(app: FastAPI):
+ result.get("entities_deleted", 0)
+ result.get("documents_deleted", 0),
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3459,6 +3540,8 @@ def _register_routes(app: FastAPI):
message=f"Cleared {result.get('deleted_count', 0)} observations",
deleted_count=result.get("deleted_count", 0),
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3468,6 +3551,42 @@ def _register_routes(app: FastAPI):
logger.error(f"Error in DELETE /v1/default/banks/{bank_id}/observations: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.delete(
"/v1/default/banks/{bank_id}/memories/{memory_id}/observations",
response_model=ClearMemoryObservationsResponse,
summary="Clear observations for a memory",
description="Delete all observations derived from a specific memory and reset it for re-consolidation. "
"The memory itself is not deleted. A consolidation job is triggered automatically so the memory "
"will produce fresh observations on the next consolidation run.",
operation_id="clear_memory_observations",
tags=["Memory"],
)
async def api_clear_memory_observations(
bank_id: str,
memory_id: str,
request_context: RequestContext = Depends(get_request_context),
):
"""Clear all observations derived from a specific memory."""
try:
result = await app.state.memory.clear_observations_for_memory(
bank_id=bank_id,
memory_id=memory_id,
request_context=request_context,
)
return ClearMemoryObservationsResponse(deleted_count=result["deleted_count"])
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}/memories/{memory_id}/observations: {error_detail}"
)
raise HTTPException(status_code=500, detail=str(e))
@app.get(
"/v1/default/banks/{bank_id}/config",
response_model=BankConfigResponse,
@@ -3482,11 +3601,18 @@ def _register_routes(app: FastAPI):
if not get_config().enable_bank_config_api:
raise HTTPException(
status_code=404,
detail="Bank configuration API is disabled. Set HINDSIGHT_API_ENABLE_BANK_CONFIG_API=true to enable.",
detail="Bank configuration API is disabled. Set HINDSIGHT_API_ENABLE_BANK_CONFIG_API=true to re-enable.",
)
try:
# 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
ctx = BankReadContext(bank_id=bank_id, operation="get_bank_config", request_context=request_context)
await app.state.memory._validate_operation(
app.state.memory._operation_validator.validate_bank_read(ctx)
)
# Get resolved config from config resolver
config_dict = await app.state.memory._config_resolver.get_bank_config(bank_id, request_context)
@@ -3495,6 +3621,8 @@ def _register_routes(app: FastAPI):
bank_overrides = await app.state.memory._config_resolver._load_bank_config(bank_id)
return BankConfigResponse(bank_id=bank_id, config=config_dict, overrides=bank_overrides)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3520,11 +3648,18 @@ def _register_routes(app: FastAPI):
if not get_config().enable_bank_config_api:
raise HTTPException(
status_code=404,
detail="Bank configuration API is disabled. Set HINDSIGHT_API_ENABLE_BANK_CONFIG_API=true to enable.",
detail="Bank configuration API is disabled. Set HINDSIGHT_API_ENABLE_BANK_CONFIG_API=true to re-enable.",
)
try:
# 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
ctx = BankWriteContext(bank_id=bank_id, operation="update_bank_config", request_context=request_context)
await app.state.memory._validate_operation(
app.state.memory._operation_validator.validate_bank_write(ctx)
)
# Update config via config resolver (validates configurable fields and permissions)
await app.state.memory._config_resolver.update_bank_config(bank_id, request.updates, request_context)
@@ -3534,6 +3669,8 @@ def _register_routes(app: FastAPI):
bank_overrides = await app.state.memory._config_resolver._load_bank_config(bank_id)
return BankConfigResponse(bank_id=bank_id, config=config_dict, overrides=bank_overrides)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except ValueError as e:
# Validation error (e.g., trying to override static field)
raise HTTPException(status_code=400, detail=str(e))
@@ -3560,11 +3697,18 @@ def _register_routes(app: FastAPI):
if not get_config().enable_bank_config_api:
raise HTTPException(
status_code=404,
detail="Bank configuration API is disabled. Set HINDSIGHT_API_ENABLE_BANK_CONFIG_API=true to enable.",
detail="Bank configuration API is disabled. Set HINDSIGHT_API_ENABLE_BANK_CONFIG_API=true to re-enable.",
)
try:
# 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
ctx = BankWriteContext(bank_id=bank_id, operation="reset_bank_config", request_context=request_context)
await app.state.memory._validate_operation(
app.state.memory._operation_validator.validate_bank_write(ctx)
)
# Reset config via config resolver
await app.state.memory._config_resolver.reset_bank_config(bank_id)
@@ -3574,6 +3718,8 @@ def _register_routes(app: FastAPI):
bank_overrides = await app.state.memory._config_resolver._load_bank_config(bank_id)
return BankConfigResponse(bank_id=bank_id, config=config_dict, overrides=bank_overrides)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3599,6 +3745,8 @@ def _register_routes(app: FastAPI):
operation_id=result["operation_id"],
deduplicated=result.get("deduplicated", False),
)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -3644,7 +3792,9 @@ def _register_routes(app: FastAPI):
contents = []
for item in request.items:
content_dict = {"content": item.content}
if item.timestamp:
if item.timestamp == "unset":
content_dict["event_date"] = None
elif item.timestamp:
content_dict["event_date"] = item.timestamp
if item.context:
content_dict["context"] = item.context
@@ -3656,6 +3806,8 @@ def _register_routes(app: FastAPI):
content_dict["entities"] = [{"text": e.text, "type": e.type or "CONCEPT"} for e in item.entities]
if item.tags:
content_dict["tags"] = item.tags
if item.observation_scopes is not None:
content_dict["observation_scopes"] = item.observation_scopes
contents.append(content_dict)
if request.async_:
@@ -3884,6 +4036,8 @@ def _register_routes(app: FastAPI):
await app.state.memory.delete_bank(bank_id, fact_type=type, request_context=request_context)
return DeleteResponse(success=True)
except OperationValidationError as e:
raise HTTPException(status_code=e.status_code, detail=e.reason)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
+81 -15
View File
@@ -8,12 +8,48 @@ from contextvars import ContextVar
from fastmcp import FastMCP
from hindsight_api import MemoryEngine
from hindsight_api.config import _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
from hindsight_api.mcp_tools import MCPToolsConfig, register_mcp_tools
from hindsight_api.models import RequestContext
# All tools available in the system (explicit list — no wildcards)
_ALL_TOOLS: frozenset[str] = frozenset(
{
"retain",
"recall",
"reflect",
"list_banks",
"create_bank",
"list_mental_models",
"get_mental_model",
"create_mental_model",
"update_mental_model",
"delete_mental_model",
"refresh_mental_model",
"list_directives",
"create_directive",
"delete_directive",
"list_memories",
"get_memory",
"delete_memory",
"list_documents",
"get_document",
"delete_document",
"list_operations",
"get_operation",
"cancel_operation",
"list_tags",
"get_bank",
"get_bank_stats",
"update_bank",
"delete_bank",
"clear_memories",
}
)
# Configure logging from HINDSIGHT_API_LOG_LEVEL environment variable
_log_level_str = os.environ.get("HINDSIGHT_API_LOG_LEVEL", "info").lower()
_log_level_map = {
@@ -82,16 +118,11 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
"""
mcp = FastMCP("hindsight-mcp-server")
# Configure and register tools using shared module
config = MCPToolsConfig(
bank_id_resolver=get_current_bank_id,
api_key_resolver=get_current_api_key, # Propagate API key for tenant auth
tenant_id_resolver=get_current_tenant_id, # Propagate tenant_id for usage metering
api_key_id_resolver=get_current_api_key_id, # Propagate api_key_id for usage metering
include_bank_id_param=multi_bank,
tools=None
if multi_bank
else {
global_config = _get_raw_config()
# Tools available for this mode (multi-bank exposes all tools; single-bank excludes bank-management tools)
_SINGLE_BANK_TOOLS: frozenset[str] = frozenset(
{
"retain",
"recall",
"reflect",
@@ -101,8 +132,40 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
"update_mental_model",
"delete_mental_model",
"refresh_mental_model",
}, # Scoped tools for single-bank mode (excludes bank management: list_banks, create_bank)
retain_fire_and_forget=False, # HTTP MCP supports sync/async modes
"list_directives",
"create_directive",
"delete_directive",
"list_memories",
"get_memory",
"delete_memory",
"list_documents",
"get_document",
"delete_document",
"list_operations",
"get_operation",
"cancel_operation",
"list_tags",
"get_bank",
"update_bank",
"delete_bank",
"clear_memories",
}
)
base_tools: frozenset[str] | None = None if multi_bank else _SINGLE_BANK_TOOLS
# Apply global mcp_enabled_tools filter (env-level allowlist)
if global_config.mcp_enabled_tools is not None:
allowed = frozenset(global_config.mcp_enabled_tools)
base_tools = (base_tools if base_tools is not None else _ALL_TOOLS) & allowed
# Configure and register tools using shared module
config = MCPToolsConfig(
bank_id_resolver=get_current_bank_id,
api_key_resolver=get_current_api_key, # Propagate API key for tenant auth
tenant_id_resolver=get_current_tenant_id, # Propagate tenant_id for usage metering
api_key_id_resolver=get_current_api_key_id, # Propagate api_key_id for usage metering
include_bank_id_param=multi_bank,
tools=base_tools,
)
register_mcp_tools(mcp, memory, config)
@@ -268,7 +331,7 @@ class MCPMiddleware:
auth_tenant_id = auth_context.tenant_id
auth_api_key_id = auth_context.api_key_id
except AuthenticationError as e:
await self._send_error(send, 401, str(e))
await self._send_error(send, 401, str(e), extra_headers=e.headers)
return
# Set schema from tenant context so downstream DB queries use the correct schema
@@ -350,14 +413,17 @@ class MCPMiddleware:
if schema_token is not None:
_current_schema.reset(schema_token)
async def _send_error(self, send, status: int, message: str):
async def _send_error(self, send, status: int, message: str, extra_headers: dict[str, str] | None = None):
"""Send an error response."""
body = json.dumps({"error": message}).encode()
headers = [(b"content-type", b"application/json")]
for key, value in (extra_headers or {}).items():
headers.append((key.encode(), value.encode()))
await send(
{
"type": "http.response.start",
"status": status,
"headers": [(b"content-type", b"application/json")],
"headers": headers,
}
)
await send(
+108 -2
View File
@@ -218,6 +218,10 @@ ENV_RERANKER_MAX_CANDIDATES = "HINDSIGHT_API_RERANKER_MAX_CANDIDATES"
ENV_RERANKER_FLASHRANK_MODEL = "HINDSIGHT_API_RERANKER_FLASHRANK_MODEL"
ENV_RERANKER_FLASHRANK_CACHE_DIR = "HINDSIGHT_API_RERANKER_FLASHRANK_CACHE_DIR"
# ZeroEntropy configuration (reranker only)
ENV_RERANKER_ZEROENTROPY_API_KEY = "HINDSIGHT_API_RERANKER_ZEROENTROPY_API_KEY"
ENV_RERANKER_ZEROENTROPY_MODEL = "HINDSIGHT_API_RERANKER_ZEROENTROPY_MODEL"
ENV_VECTOR_EXTENSION = "HINDSIGHT_API_VECTOR_EXTENSION"
ENV_TEXT_SEARCH_EXTENSION = "HINDSIGHT_API_TEXT_SEARCH_EXTENSION"
@@ -228,6 +232,7 @@ ENV_LOG_LEVEL = "HINDSIGHT_API_LOG_LEVEL"
ENV_LOG_FORMAT = "HINDSIGHT_API_LOG_FORMAT"
ENV_WORKERS = "HINDSIGHT_API_WORKERS"
ENV_MCP_ENABLED = "HINDSIGHT_API_MCP_ENABLED"
ENV_MCP_ENABLED_TOOLS = "HINDSIGHT_API_MCP_ENABLED_TOOLS"
ENV_ENABLE_BANK_CONFIG_API = "HINDSIGHT_API_ENABLE_BANK_CONFIG_API"
ENV_GRAPH_RETRIEVER = "HINDSIGHT_API_GRAPH_RETRIEVER"
ENV_MPFP_TOP_K_NEIGHBORS = "HINDSIGHT_API_MPFP_TOP_K_NEIGHBORS"
@@ -247,13 +252,18 @@ ENV_LLM_VERTEXAI_PROJECT_ID = "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID"
ENV_LLM_VERTEXAI_REGION = "HINDSIGHT_API_LLM_VERTEXAI_REGION"
ENV_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY = "HINDSIGHT_API_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY"
# Gemini safety settings
ENV_LLM_GEMINI_SAFETY_SETTINGS = "HINDSIGHT_API_LLM_GEMINI_SAFETY_SETTINGS"
# Retain settings
ENV_RETAIN_MAX_COMPLETION_TOKENS = "HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS"
ENV_RETAIN_CHUNK_SIZE = "HINDSIGHT_API_RETAIN_CHUNK_SIZE"
ENV_RETAIN_EXTRACT_CAUSAL_LINKS = "HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS"
ENV_RETAIN_EXTRACTION_MODE = "HINDSIGHT_API_RETAIN_EXTRACTION_MODE"
ENV_RETAIN_MISSION = "HINDSIGHT_API_RETAIN_MISSION"
ENV_RETAIN_CUSTOM_INSTRUCTIONS = "HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS"
ENV_RETAIN_BATCH_TOKENS = "HINDSIGHT_API_RETAIN_BATCH_TOKENS"
ENV_RETAIN_ENTITY_LOOKUP = "HINDSIGHT_API_RETAIN_ENTITY_LOOKUP"
ENV_RETAIN_BATCH_ENABLED = "HINDSIGHT_API_RETAIN_BATCH_ENABLED"
ENV_RETAIN_BATCH_POLL_INTERVAL_SECONDS = "HINDSIGHT_API_RETAIN_BATCH_POLL_INTERVAL_SECONDS"
@@ -280,7 +290,9 @@ ENV_FILE_DELETE_AFTER_RETAIN = "HINDSIGHT_API_FILE_DELETE_AFTER_RETAIN"
# Observations settings (consolidated knowledge from facts)
ENV_ENABLE_OBSERVATIONS = "HINDSIGHT_API_ENABLE_OBSERVATIONS"
ENV_CONSOLIDATION_BATCH_SIZE = "HINDSIGHT_API_CONSOLIDATION_BATCH_SIZE"
ENV_CONSOLIDATION_LLM_BATCH_SIZE = "HINDSIGHT_API_CONSOLIDATION_LLM_BATCH_SIZE"
ENV_CONSOLIDATION_MAX_TOKENS = "HINDSIGHT_API_CONSOLIDATION_MAX_TOKENS"
ENV_OBSERVATIONS_MISSION = "HINDSIGHT_API_OBSERVATIONS_MISSION"
# Optimization flags
ENV_SKIP_LLM_VERIFICATION = "HINDSIGHT_API_SKIP_LLM_VERIFICATION"
@@ -306,6 +318,13 @@ ENV_WORKER_CONSOLIDATION_MAX_SLOTS = "HINDSIGHT_API_WORKER_CONSOLIDATION_MAX_SLO
# Reflect agent settings
ENV_REFLECT_MAX_ITERATIONS = "HINDSIGHT_API_REFLECT_MAX_ITERATIONS"
ENV_REFLECT_MAX_CONTEXT_TOKENS = "HINDSIGHT_API_REFLECT_MAX_CONTEXT_TOKENS"
ENV_REFLECT_MISSION = "HINDSIGHT_API_REFLECT_MISSION"
# Disposition settings
ENV_DISPOSITION_SKEPTICISM = "HINDSIGHT_API_DISPOSITION_SKEPTICISM"
ENV_DISPOSITION_LITERALISM = "HINDSIGHT_API_DISPOSITION_LITERALISM"
ENV_DISPOSITION_EMPATHY = "HINDSIGHT_API_DISPOSITION_EMPATHY"
# Default values
DEFAULT_DATABASE_URL = "pg0"
@@ -337,6 +356,9 @@ DEFAULT_LLM_VERTEXAI_PROJECT_ID = None # Required for Vertex AI
DEFAULT_LLM_VERTEXAI_REGION = "us-central1"
DEFAULT_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY = None # Optional, uses ADC if not set
# Gemini safety settings defaults
DEFAULT_LLM_GEMINI_SAFETY_SETTINGS = None # None = use Gemini default safety settings
DEFAULT_EMBEDDINGS_PROVIDER = "local"
DEFAULT_EMBEDDINGS_LOCAL_MODEL = "BAAI/bge-small-en-v1.5"
DEFAULT_EMBEDDINGS_LOCAL_FORCE_CPU = False # Force CPU mode for local embeddings (avoids MPS/XPC issues on macOS)
@@ -360,6 +382,8 @@ DEFAULT_RERANKER_FLASHRANK_CACHE_DIR = None # Use default cache directory
DEFAULT_EMBEDDINGS_COHERE_MODEL = "embed-english-v3.0"
DEFAULT_RERANKER_COHERE_MODEL = "rerank-english-v3.0"
DEFAULT_RERANKER_ZEROENTROPY_MODEL = "zerank-2"
# Vector extension (pgvector, vchord, or pgvectorscale)
DEFAULT_VECTOR_EXTENSION = "pgvector" # Options: "pgvector", "vchord", "pgvectorscale"
@@ -382,7 +406,8 @@ DEFAULT_LOG_LEVEL = "info"
DEFAULT_LOG_FORMAT = "text" # Options: "text", "json"
DEFAULT_WORKERS = 1
DEFAULT_MCP_ENABLED = True
DEFAULT_ENABLE_BANK_CONFIG_API = False # Disabled by default for security
DEFAULT_MCP_ENABLED_TOOLS: list[str] | None = None # None = all tools enabled
DEFAULT_ENABLE_BANK_CONFIG_API = True
DEFAULT_GRAPH_RETRIEVER = "link_expansion" # Options: "link_expansion", "mpfp", "bfs"
DEFAULT_MPFP_TOP_K_NEIGHBORS = 20 # Fan-out limit per node in MPFP graph traversal
DEFAULT_RECALL_MAX_CONCURRENT = 32 # Max concurrent recall operations per worker
@@ -395,8 +420,10 @@ DEFAULT_RETAIN_CHUNK_SIZE = 3000 # Max chars per chunk for fact extraction
DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS = True # Extract causal links between facts
DEFAULT_RETAIN_EXTRACTION_MODE = "concise" # Extraction mode: "concise", "verbose", or "custom"
RETAIN_EXTRACTION_MODES = ("concise", "verbose", "custom") # Allowed extraction modes
DEFAULT_RETAIN_MISSION = None # Declarative spec of what to retain (injected into any extraction mode)
DEFAULT_RETAIN_CUSTOM_INSTRUCTIONS = None # Custom extraction guidelines (only used when mode="custom")
DEFAULT_RETAIN_BATCH_TOKENS = 10_000 # ~40KB of text # Max chars per sub-batch for async retain auto-splitting
DEFAULT_RETAIN_ENTITY_LOOKUP = "trigram" # "full" or "trigram"
DEFAULT_RETAIN_BATCH_ENABLED = False # Use LLM Batch API for fact extraction (only when async=True)
DEFAULT_RETAIN_BATCH_POLL_INTERVAL_SECONDS = 60 # Batch API polling interval in seconds
@@ -411,7 +438,9 @@ DEFAULT_FILE_DELETE_AFTER_RETAIN = True # Delete file bytes after retain (saves
# Observations defaults (consolidated knowledge from facts)
DEFAULT_ENABLE_OBSERVATIONS = True # Observations enabled by default
DEFAULT_CONSOLIDATION_BATCH_SIZE = 50 # Memories to load per batch (internal memory optimization)
DEFAULT_CONSOLIDATION_MAX_TOKENS = 1024 # Max tokens for recall when finding related observations
DEFAULT_CONSOLIDATION_LLM_BATCH_SIZE = 8 # Facts per LLM call (1 = no batching; >1 = batch mode)
DEFAULT_CONSOLIDATION_MAX_TOKENS = 512 # Max tokens for recall when finding related observations
DEFAULT_OBSERVATIONS_MISSION = None # Declarative spec of what observations are for this bank
# Database migrations
DEFAULT_RUN_MIGRATIONS_ON_STARTUP = True
@@ -433,6 +462,12 @@ DEFAULT_WORKER_CONSOLIDATION_MAX_SLOTS = 2 # Max concurrent consolidation tasks
# Reflect agent settings
DEFAULT_REFLECT_MAX_ITERATIONS = 10 # Max tool call iterations before forcing response
DEFAULT_REFLECT_MAX_CONTEXT_TOKENS = 100_000 # Max accumulated context tokens before forcing final prompt
# Disposition defaults (None = not set, fall back to bank DB value or 3)
DEFAULT_DISPOSITION_SKEPTICISM = None
DEFAULT_DISPOSITION_LITERALISM = None
DEFAULT_DISPOSITION_EMPATHY = None
# OpenTelemetry tracing configuration
DEFAULT_OTEL_TRACES_ENABLED = False # Disabled by default for backward compatibility
@@ -538,6 +573,9 @@ class HindsightConfig:
llm_vertexai_region: str
llm_vertexai_service_account_key: str | None
# Gemini safety settings (None = use Gemini defaults; list of dicts with category/threshold)
llm_gemini_safety_settings: list | None
# Per-operation LLM configuration (None = use default LLM config)
retain_llm_provider: str | None
retain_llm_api_key: str | None
@@ -605,6 +643,8 @@ class HindsightConfig:
reranker_litellm_sdk_api_key: str | None
reranker_litellm_sdk_model: str
reranker_litellm_sdk_api_base: str | None
reranker_zeroentropy_api_key: str | None
reranker_zeroentropy_model: str
# Server
host: str
@@ -613,6 +653,7 @@ class HindsightConfig:
log_level: str
log_format: str
mcp_enabled: bool
mcp_enabled_tools: list[str] | None # None = all tools; explicit list = allowlist
enable_bank_config_api: bool
# Recall
@@ -627,10 +668,12 @@ class HindsightConfig:
retain_chunk_size: int
retain_extract_causal_links: bool
retain_extraction_mode: str
retain_mission: str | None
retain_custom_instructions: str | None
retain_batch_tokens: int
retain_batch_enabled: bool
retain_batch_poll_interval_seconds: int
retain_entity_lookup: str # "full" or "trigram"
# File storage (static - server-level only)
file_storage_type: str # "native" (PostgreSQL) or "s3" (S3-compatible)
@@ -655,7 +698,24 @@ class HindsightConfig:
# Observations settings (consolidated knowledge from facts)
enable_observations: bool
consolidation_batch_size: int
consolidation_llm_batch_size: int
consolidation_max_tokens: int
observations_mission: str | None
# Entity labels (controlled vocabulary of key:value classification labels extracted at retain time)
# List of label group dicts: [{key, description, type, optional, values: [{value, description}]}]
entity_labels: list | None
# Whether to extract regular named entities alongside entity labels (default: True)
# When False: only label entities are extracted (or no entities at all if no labels configured)
entities_allow_free_form: bool
# Reflect agent settings
reflect_mission: str | None
# Disposition settings (hierarchical - can be overridden per bank; None = fall back to DB)
disposition_skepticism: int | None
disposition_literalism: int | None
disposition_empathy: int | None
# Optimization flags
skip_llm_verification: bool
@@ -681,6 +741,7 @@ class HindsightConfig:
# Reflect agent settings
reflect_max_iterations: int
reflect_max_context_tokens: int
# OpenTelemetry tracing configuration
otel_traces_enabled: bool
@@ -721,12 +782,27 @@ class HindsightConfig:
# These fields are manually tagged as safe to expose and modify.
# Excludes credentials, infrastructure config, provider/model selection, and performance tuning.
_CONFIGURABLE_FIELDS = {
# MCP tool access control
"mcp_enabled_tools",
# Retention settings (behavioral)
"retain_chunk_size",
"retain_extraction_mode",
"retain_mission",
"retain_custom_instructions",
# Entity labels (controlled vocabulary for entity classification)
"entity_labels",
"entities_allow_free_form",
# Consolidation settings
"enable_observations",
"observations_mission",
# Reflect settings
"reflect_mission",
# Disposition settings
"disposition_skepticism",
"disposition_literalism",
"disposition_empathy",
# Gemini safety settings (controls content filtering for Gemini/VertexAI providers)
"llm_gemini_safety_settings",
}
@property
@@ -847,6 +923,8 @@ class HindsightConfig:
llm_vertexai_region=os.getenv(ENV_LLM_VERTEXAI_REGION, DEFAULT_LLM_VERTEXAI_REGION),
llm_vertexai_service_account_key=os.getenv(ENV_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY)
or DEFAULT_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY,
# Gemini safety settings (JSON-encoded list of {category, threshold} dicts)
llm_gemini_safety_settings=json.loads(os.getenv(ENV_LLM_GEMINI_SAFETY_SETTINGS, "null")),
# Per-operation LLM config (None = use default)
retain_llm_provider=os.getenv(ENV_RETAIN_LLM_PROVIDER) or None,
retain_llm_api_key=os.getenv(ENV_RETAIN_LLM_API_KEY) or None,
@@ -979,6 +1057,9 @@ class HindsightConfig:
reranker_litellm_sdk_api_key=os.getenv(ENV_RERANKER_LITELLM_SDK_API_KEY),
reranker_litellm_sdk_model=os.getenv(ENV_RERANKER_LITELLM_SDK_MODEL, DEFAULT_RERANKER_LITELLM_SDK_MODEL),
reranker_litellm_sdk_api_base=os.getenv(ENV_RERANKER_LITELLM_SDK_API_BASE) or None,
# ZeroEntropy reranker
reranker_zeroentropy_api_key=os.getenv(ENV_RERANKER_ZEROENTROPY_API_KEY),
reranker_zeroentropy_model=os.getenv(ENV_RERANKER_ZEROENTROPY_MODEL, DEFAULT_RERANKER_ZEROENTROPY_MODEL),
# Server
host=os.getenv(ENV_HOST, DEFAULT_HOST),
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
@@ -986,6 +1067,9 @@ class HindsightConfig:
log_level=os.getenv(ENV_LOG_LEVEL, DEFAULT_LOG_LEVEL),
log_format=os.getenv(ENV_LOG_FORMAT, DEFAULT_LOG_FORMAT).lower(),
mcp_enabled=os.getenv(ENV_MCP_ENABLED, str(DEFAULT_MCP_ENABLED)).lower() == "true",
mcp_enabled_tools=[t.strip() for t in os.getenv(ENV_MCP_ENABLED_TOOLS).split(",") if t.strip()]
if os.getenv(ENV_MCP_ENABLED_TOOLS)
else DEFAULT_MCP_ENABLED_TOOLS,
enable_bank_config_api=os.getenv(ENV_ENABLE_BANK_CONFIG_API, str(DEFAULT_ENABLE_BANK_CONFIG_API)).lower()
== "true",
# Recall
@@ -1013,8 +1097,10 @@ class HindsightConfig:
retain_extraction_mode=_validate_extraction_mode(
os.getenv(ENV_RETAIN_EXTRACTION_MODE, DEFAULT_RETAIN_EXTRACTION_MODE)
),
retain_mission=os.getenv(ENV_RETAIN_MISSION) or DEFAULT_RETAIN_MISSION,
retain_custom_instructions=os.getenv(ENV_RETAIN_CUSTOM_INSTRUCTIONS) or DEFAULT_RETAIN_CUSTOM_INSTRUCTIONS,
retain_batch_tokens=int(os.getenv(ENV_RETAIN_BATCH_TOKENS, str(DEFAULT_RETAIN_BATCH_TOKENS))),
retain_entity_lookup=os.getenv(ENV_RETAIN_ENTITY_LOOKUP, DEFAULT_RETAIN_ENTITY_LOOKUP),
retain_batch_enabled=os.getenv(ENV_RETAIN_BATCH_ENABLED, str(DEFAULT_RETAIN_BATCH_ENABLED)).lower()
== "true",
retain_batch_poll_interval_seconds=int(
@@ -1052,9 +1138,15 @@ class HindsightConfig:
consolidation_batch_size=int(
os.getenv(ENV_CONSOLIDATION_BATCH_SIZE, str(DEFAULT_CONSOLIDATION_BATCH_SIZE))
),
consolidation_llm_batch_size=int(
os.getenv(ENV_CONSOLIDATION_LLM_BATCH_SIZE, str(DEFAULT_CONSOLIDATION_LLM_BATCH_SIZE))
),
consolidation_max_tokens=int(
os.getenv(ENV_CONSOLIDATION_MAX_TOKENS, str(DEFAULT_CONSOLIDATION_MAX_TOKENS))
),
observations_mission=os.getenv(ENV_OBSERVATIONS_MISSION) or DEFAULT_OBSERVATIONS_MISSION,
entity_labels=None,
entities_allow_free_form=True,
# Database migrations
run_migrations_on_startup=os.getenv(ENV_RUN_MIGRATIONS_ON_STARTUP, "true").lower() == "true",
# Database connection pool
@@ -1074,6 +1166,20 @@ class HindsightConfig:
),
# Reflect agent settings
reflect_max_iterations=int(os.getenv(ENV_REFLECT_MAX_ITERATIONS, str(DEFAULT_REFLECT_MAX_ITERATIONS))),
reflect_max_context_tokens=int(
os.getenv(ENV_REFLECT_MAX_CONTEXT_TOKENS, str(DEFAULT_REFLECT_MAX_CONTEXT_TOKENS))
),
reflect_mission=os.getenv(ENV_REFLECT_MISSION) or None,
# Disposition settings (None = fall back to DB value)
disposition_skepticism=int(os.getenv(ENV_DISPOSITION_SKEPTICISM))
if os.getenv(ENV_DISPOSITION_SKEPTICISM)
else DEFAULT_DISPOSITION_SKEPTICISM,
disposition_literalism=int(os.getenv(ENV_DISPOSITION_LITERALISM))
if os.getenv(ENV_DISPOSITION_LITERALISM)
else DEFAULT_DISPOSITION_LITERALISM,
disposition_empathy=int(os.getenv(ENV_DISPOSITION_EMPATHY))
if os.getenv(ENV_DISPOSITION_EMPATHY)
else DEFAULT_DISPOSITION_EMPATHY,
# OpenTelemetry tracing configuration
otel_traces_enabled=os.getenv(ENV_OTEL_TRACES_ENABLED, str(DEFAULT_OTEL_TRACES_ENABLED)).lower()
in ("true", "1", "yes"),
File diff suppressed because it is too large Load Diff
@@ -1,84 +1,66 @@
"""Prompts for the consolidation engine."""
CONSOLIDATION_SYSTEM_PROMPT = """You are a memory consolidation system. Your job is to convert facts into durable knowledge (observations) and merge with existing knowledge when appropriate.
# Default mission when no bank-specific mission is set
_DEFAULT_MISSION = "Track every detail: names, numbers, dates, places, and relationships. Prefer specifics over abstractions, never generalise."
You must output a JSON object with an "actions" array. The "text" field within each action should use markdown formatting (headers, lists, bold, etc.) for clarity and readability.
# Processing rules — always present regardless of mission
_PROCESSING_RULES = """Processing rules (always apply):
- REDUNDANT: same info worded differently → UPDATE the existing observation.
- CONTRADICTION/UPDATE: capture both states with temporal markers ("used to X, now Y").
- RESOLVE REFERENCES: when a new fact provides a concrete value resolving a vague placeholder in an existing observation (e.g. "home country", "hometown", "birthplace", "native language", "her ex", "that city"), UPDATE the observation to embed the resolved value explicitly. Example: new fact says "grandma in Sweden" + existing observation says "moved from her home country" → update to "home country is Sweden".
- NEVER merge observations about different people or unrelated topics."""
## EXTRACT DURABLE KNOWLEDGE, NOT EPHEMERAL STATE
Facts often describe events or actions. Extract the DURABLE KNOWLEDGE implied by the fact, not the transient state.
# Data section — format placeholders {facts_text} and {observations_text} are substituted at call time
_BATCH_DATA_SECTION = """
NEW FACTS:
{facts_text}
Examples of extracting durable knowledge:
- "User moved to Room 203" -> "Room 203 exists" (location exists, not where user is now)
- "User visited Acme Corp at Room 105" -> "Acme Corp is located in Room 105"
- "User took the elevator to floor 3" -> "Floor 3 is accessible by elevator"
- "User met Sarah at the lobby" -> "Sarah can be found at the lobby"
DO NOT track current user position/state as knowledge - that changes constantly.
DO track permanent facts learned from the user's actions.
## PRESERVE SPECIFIC DETAILS
Keep names, locations, numbers, and other specifics. Do NOT:
- Abstract into general principles
- Generate business insights
- Make knowledge generic
GOOD examples:
- Fact: "John likes pizza" -> "John likes pizza"
- Fact: "Alice works at Google" -> "Alice works at Google"
BAD examples:
- "John likes pizza" -> "Understanding dietary preferences helps..." (TOO ABSTRACT)
- "User is at Room 203" -> "User is currently at Room 203" (EPHEMERAL STATE)
## MERGE RULES (when comparing to existing observations):
1. REDUNDANT: Same information worded differently → update existing
2. CONTRADICTION: Opposite information about same topic → update with temporal markers showing change
Example: "Alex used to love pizza but now hates it" OR "Alex's pizza preference changed from love to hate"
3. UPDATE: New state replacing old state → update showing the transition with "used to", "now", "changed from X to Y"
## CRITICAL RULES:
- NEVER merge facts about DIFFERENT people
- NEVER merge unrelated topics (food preferences vs work vs hobbies)
- When merging contradictions, the "text" field MUST capture BOTH states with temporal markers:
* Use "used to X, now Y" OR "changed from X to Y" OR "X but now Y"
* DO NOT just state the new fact - you MUST show the change
- Keep observations focused on ONE specific topic per person
- The "text" field MUST contain durable knowledge, not ephemeral state
- Do NOT include "tags" in output - tags are handled automatically"""
CONSOLIDATION_USER_PROMPT = """Analyze this new fact and consolidate into knowledge.
{mission_section}
NEW FACT: {fact_text}
EXISTING OBSERVATIONS (JSON array with source memories and dates):
EXISTING OBSERVATIONS (JSON array, pooled from recalls across all facts above):
{observations_text}
Each observation includes:
- id: unique identifier for updating
- text: the observation content
- proof_count: number of supporting memories
- tags: visibility scope (handled automatically)
- occurred_start/occurred_end: temporal range of source facts
- source_memories: array of supporting facts with their text and dates
Instructions:
1. Extract DURABLE KNOWLEDGE from the new fact (not ephemeral state)
2. Review source_memories in existing observations to understand evidence
3. Check dates to detect contradictions or updates
4. Compare with observations:
- Same topic → UPDATE with learning_id
- New topic → CREATE new observation
- Purely ephemeral → return empty actions list
Compare the facts against existing observations:
- Same topic as an existing observation → UPDATE it (observation_id + source_fact_ids)
- New topic with durable knowledge → CREATE a new observation (source_fact_ids)
- Cross-reference facts within the batch: a later fact may resolve a vague reference in an earlier one
- Purely ephemeral facts → omit them (no create/update needed)"""
Output a JSON object with an "actions" array (the "text" field should use markdown formatting for structure):
{{"actions": [
{{"action": "update", "learning_id": "uuid-from-observations", "text": "## Updated Knowledge\n\n**Key point**: details here\n\n- Supporting detail 1\n- Supporting detail 2", "reason": "..."}},
{{"action": "create", "text": "## New Durable Knowledge\n\nDescription with **emphasis** and proper structure", "reason": "..."}}
]}}
# Output format — JSON braces escaped as {{ }} so .format() leaves them literal
_BATCH_OUTPUT_FORMAT = """
Output a JSON object with three arrays.
Return {{"actions": []}} if fact contains no durable knowledge.
Example (showing the required UUID format for all IDs):
{{"creates": [{{"text": "Alice lives in Berlin", "source_fact_ids": ["a1b2c3d4-e5f6-7890-abcd-ef1234567890", "b2c3d4e5-f6a7-8901-bcde-f12345678901"]}}],
"updates": [{{"text": "Alice works at Acme Corp as a senior engineer", "observation_id": "c3d4e5f6-a7b8-9012-cdef-123456789012", "source_fact_ids": ["d4e5f6a7-b8c9-0123-defa-234567890123"]}}],
"deletes": [{{"observation_id": "e5f6a7b8-c9d0-1234-efab-345678901234"}}]}}
IMPORTANT: Format the "text" field with markdown for better readability:
- Use headers, lists, bold/italic, tables where appropriate
- CRITICAL: Add blank lines before and after block elements (tables, code blocks, lists)
- Ensure proper spacing for markdown to render correctly"""
Rules:
- "source_fact_ids": copy the EXACT UUID strings shown in brackets [uuid] from NEW FACTS — never use integers or positions.
- "observation_id": copy the EXACT "id" UUID string from EXISTING OBSERVATIONS.
- One create/update may reference multiple facts when they jointly support the observation.
- "deletes": only when an observation is directly superseded or contradicted by new facts.
- Do NOT include "tags" — handled automatically.
- Return {{"creates": [], "updates": [], "deletes": []}} if nothing durable is found."""
def build_batch_consolidation_prompt(observations_mission: str | None = None) -> str:
"""
Build the consolidation prompt for batch mode (multiple facts per LLM call).
The mission defines *what* to track (customisable per bank).
Processing rules and output format are always present regardless of mission.
"""
mission = observations_mission or _DEFAULT_MISSION
return (
"You are a memory consolidation system. Synthesize facts into observations "
"and merge with existing observations when appropriate.\n\n"
f"## MISSION\n{mission}\n\n"
f"{_PROCESSING_RULES}" + _BATCH_DATA_SECTION + _BATCH_OUTPUT_FORMAT
)
@@ -29,6 +29,7 @@ from ..config import (
DEFAULT_RERANKER_PROVIDER,
DEFAULT_RERANKER_TEI_BATCH_SIZE,
DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
DEFAULT_RERANKER_ZEROENTROPY_MODEL,
ENV_RERANKER_COHERE_API_KEY,
ENV_RERANKER_COHERE_MODEL,
ENV_RERANKER_FLASHRANK_CACHE_DIR,
@@ -42,6 +43,7 @@ from ..config import (
ENV_RERANKER_TEI_BATCH_SIZE,
ENV_RERANKER_TEI_MAX_CONCURRENT,
ENV_RERANKER_TEI_URL,
ENV_RERANKER_ZEROENTROPY_API_KEY,
)
logger = logging.getLogger(__name__)
@@ -556,6 +558,104 @@ class CohereCrossEncoder(CrossEncoderModel):
return all_scores
class ZeroEntropyCrossEncoder(CrossEncoderModel):
"""
ZeroEntropy cross-encoder implementation using the ZeroEntropy Rerank API.
Supports zerank-2 (flagship) and zerank-2-small models.
See: https://docs.zeroentropy.dev/models
"""
RERANK_URL = "https://api.zeroentropy.dev/v1/models/rerank"
def __init__(
self,
api_key: str,
model: str = DEFAULT_RERANKER_ZEROENTROPY_MODEL,
timeout: float = 60.0,
):
"""
Initialize ZeroEntropy cross-encoder client.
Args:
api_key: ZeroEntropy API key
model: ZeroEntropy rerank model name (default: zerank-2)
timeout: Request timeout in seconds (default: 60.0)
"""
self.api_key = api_key
self.model = model
self.timeout = timeout
self._async_client: httpx.AsyncClient | None = None
@property
def provider_name(self) -> str:
return "zeroentropy"
async def initialize(self) -> None:
"""Initialize the async HTTP client."""
if self._async_client is not None:
return
logger.info(f"Reranker: initializing ZeroEntropy provider with model {self.model}")
self._async_client = httpx.AsyncClient(
timeout=self.timeout,
headers={
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
},
)
logger.info("Reranker: ZeroEntropy provider initialized")
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using the ZeroEntropy Rerank API.
Args:
pairs: List of (query, document) tuples to score
Returns:
List of relevance scores
"""
if self._async_client is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
if not pairs:
return []
# Group pairs by query for efficient batching
query_groups: dict[str, list[tuple[int, str]]] = {}
for idx, (query, text) in enumerate(pairs):
if query not in query_groups:
query_groups[query] = []
query_groups[query].append((idx, text))
all_scores = [0.0] * len(pairs)
for query, indexed_texts in query_groups.items():
texts = [text for _, text in indexed_texts]
indices = [idx for idx, _ in indexed_texts]
response = await self._async_client.post(
self.RERANK_URL,
json={
"model": self.model,
"query": query,
"documents": texts,
"top_n": len(texts),
},
)
response.raise_for_status()
result = response.json()
# Map scores back to original positions
for item in result.get("results", []):
original_idx = item["index"]
score = item["relevance_score"]
all_scores[indices[original_idx]] = score
return all_scores
class RRFPassthroughCrossEncoder(CrossEncoderModel):
"""
Passthrough cross-encoder that preserves RRF scores without neural reranking.
@@ -1010,9 +1110,19 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
model=config.reranker_litellm_sdk_model,
api_base=config.reranker_litellm_sdk_api_base,
)
elif provider == "zeroentropy":
api_key = config.reranker_zeroentropy_api_key
if not api_key:
raise ValueError(
f"{ENV_RERANKER_ZEROENTROPY_API_KEY} is required when {ENV_RERANKER_PROVIDER} is 'zeroentropy'"
)
return ZeroEntropyCrossEncoder(
api_key=api_key,
model=config.reranker_zeroentropy_model,
)
elif provider == "rrf":
return RRFPassthroughCrossEncoder()
else:
raise ValueError(
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'flashrank', 'litellm', 'litellm-sdk', 'rrf'"
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'zeroentropy', 'flashrank', 'litellm', 'litellm-sdk', 'rrf'"
)
@@ -20,6 +20,7 @@ RETRYABLE_EXCEPTIONS = (
asyncpg.exceptions.InterfaceError,
asyncpg.exceptions.ConnectionDoesNotExistError,
asyncpg.exceptions.TooManyConnectionsError,
asyncpg.exceptions.DeadlockDetectedError,
OSError,
ConnectionError,
asyncio.TimeoutError,
@@ -794,6 +794,7 @@ class LiteLLMSDKEmbeddings(Embeddings):
"model": self.model,
"input": ["test"],
"api_key": self.api_key,
"encoding_format": "float",
}
if self.api_base:
embed_kwargs["api_base"] = self.api_base
@@ -840,6 +841,7 @@ class LiteLLMSDKEmbeddings(Embeddings):
"model": self.model,
"input": batch,
"api_key": self.api_key,
"encoding_format": "float",
}
if self.api_base:
embed_kwargs["api_base"] = self.api_base
@@ -5,6 +5,10 @@ Uses spaCy for entity extraction and implements resolution logic
to disambiguate entities across memory units.
"""
import asyncio
import logging
from collections import defaultdict
from dataclasses import dataclass, field
from datetime import UTC, datetime
from difflib import SequenceMatcher
@@ -12,6 +16,43 @@ import asyncpg
from .db_utils import acquire_with_retry
from .memory_engine import fq_table
from .retain.entity_labels import build_labels_lookup as _build_labels_lookup_from_config
logger = logging.getLogger(__name__)
@dataclass
class _EntityToCreate:
"""An entity that needs to be inserted (no matching candidate found)."""
idx: int
name: str
event_date: datetime | None
@dataclass
class _EntityStat:
"""Stat accumulation entry for a resolved entity (post-transaction update)."""
entity_id: str
event_date: datetime | None
@dataclass
class _EntityStatAgg:
"""Aggregated stats used when flushing pending updates."""
count: int = 0
max_date: datetime | None = None
@dataclass
class _CooccurrencePair:
"""A (entity_id_1, entity_id_2) pair observed in a retain batch (for post-txn flush)."""
entity_id_1: str
entity_id_2: str
# Load spaCy model (singleton)
_nlp = None
@@ -22,14 +63,95 @@ class EntityResolver:
Resolves entities to canonical IDs with disambiguation.
"""
def __init__(self, pool: asyncpg.Pool):
def __init__(self, pool: asyncpg.Pool, entity_lookup: str = "full"):
"""
Initialize entity resolver.
Args:
pool: asyncpg connection pool
entity_lookup: Lookup strategy — "full" loads all bank entities then
matches in Python; "trigram" uses pg_trgm GIN index to fetch only
similar candidates per entity name (much faster for large banks).
"""
self.pool = pool
self.entity_lookup = entity_lookup
# Keyed by asyncio task id so concurrent retain batches never mix their
# pending updates. flush_pending_stats() pops only the calling task's items.
self._pending_stats: dict[int, list[_EntityStat]] = {}
self._pending_cooccurrences: dict[int, list[_CooccurrencePair]] = {}
def _task_key(self) -> int:
"""Return a unique key for the current asyncio task (or 0 for non-task context)."""
task = asyncio.current_task()
return id(task) if task is not None else 0
async def flush_pending_stats(self) -> None:
"""
Flush accumulated entity stats and co-occurrence counts for the current task.
Must be called AFTER the retain transaction commits. Pops only the items
accumulated by the calling asyncio task so concurrent retain batches never
flush each other's uncommitted entity IDs.
"""
if self.pool is None:
return
key = self._task_key()
stats = self._pending_stats.pop(key, [])
cooccurrences = self._pending_cooccurrences.pop(key, [])
if not stats and not cooccurrences:
return
async with acquire_with_retry(self.pool) as conn:
if stats:
# Aggregate: sum counts and find max date per entity_id.
agg: dict[str, _EntityStatAgg] = defaultdict(_EntityStatAgg)
for s in stats:
entry = agg[s.entity_id]
entry.count += 1
if s.event_date is not None:
entry.max_date = s.event_date if entry.max_date is None else max(entry.max_date, s.event_date)
# Sort by entity_id so all concurrent workers acquire row locks in
# the same order — prevents circular lock dependencies (deadlocks).
rows = sorted((eid, a.count, a.max_date) for eid, a in agg.items())
await conn.executemany(
f"""
UPDATE {fq_table("entities")} SET
mention_count = mention_count + $2,
last_seen = GREATEST(last_seen, $3)
WHERE id = $1::uuid
""",
rows,
)
if cooccurrences:
# Aggregate: count occurrences per (entity_id_1, entity_id_2) pair.
coo_agg: dict[tuple[str, str], int] = {}
for c in cooccurrences:
pair = (c.entity_id_1, c.entity_id_2)
coo_agg[pair] = coo_agg.get(pair, 0) + 1
now = datetime.now(UTC)
# Sort by (entity_id_1, entity_id_2) for consistent lock ordering.
await conn.executemany(
f"""
INSERT INTO {fq_table("entity_cooccurrences")}
(entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
VALUES ($1, $2, $3, $4)
ON CONFLICT (entity_id_1, entity_id_2)
DO UPDATE SET
cooccurrence_count = {fq_table("entity_cooccurrences")}.cooccurrence_count + EXCLUDED.cooccurrence_count,
last_cooccurred = GREATEST({fq_table("entity_cooccurrences")}.last_cooccurred, EXCLUDED.last_cooccurred)
""",
sorted((e1, e2, count, now) for (e1, e2), count in coo_agg.items()),
)
@staticmethod
def _build_labels_lookup(entity_labels: list | None) -> set[str]:
"""Build a set of valid 'key:value' entity label strings for fast lookup."""
return _build_labels_lookup_from_config(entity_labels)
async def resolve_entities_batch(
self,
@@ -38,6 +160,7 @@ class EntityResolver:
context: str,
unit_event_date,
conn=None,
entity_labels: list | None = None,
) -> list[str]:
"""
Resolve multiple entities in batch (MUCH faster than sequential).
@@ -58,15 +181,34 @@ class EntityResolver:
if not entities_data:
return []
taxonomy_lookup = self._build_labels_lookup(entity_labels)
if conn is None:
async with acquire_with_retry(self.pool) as conn:
return await self._resolve_entities_batch_impl(conn, bank_id, entities_data, context, unit_event_date)
return await self._resolve_entities_batch_impl(
conn, bank_id, entities_data, context, unit_event_date, taxonomy_lookup
)
else:
return await self._resolve_entities_batch_impl(conn, bank_id, entities_data, context, unit_event_date)
return await self._resolve_entities_batch_impl(
conn, bank_id, entities_data, context, unit_event_date, taxonomy_lookup
)
async def _resolve_entities_batch_impl(
self, conn, bank_id: str, entities_data: list[dict], context: str, unit_event_date
self,
conn,
bank_id: str,
entities_data: list[dict],
context: str,
unit_event_date,
taxonomy_lookup: set[str] | None = None,
) -> list[str]:
if self.entity_lookup == "trigram":
return await self._resolve_entities_batch_trigram(conn, bank_id, entities_data, unit_event_date)
return await self._resolve_entities_batch_full(conn, bank_id, entities_data, unit_event_date)
async def _resolve_entities_batch_full(
self, conn, bank_id: str, entities_data: list[dict], unit_event_date
) -> list[str]:
"""Original strategy: load all bank entities then match in Python."""
# Query ALL candidates for this bank
all_entities = await conn.fetch(
f"""
@@ -130,10 +272,103 @@ class EntityResolver:
matching.append((ent_id, canonical_name, metadata, last_seen, mention_count))
all_candidates[entity_text] = matching
return await self._resolve_from_candidates(
conn, bank_id, entities_data, unit_event_date, all_candidates, cooccurrence_map
)
async def _resolve_entities_batch_trigram(
self, conn, bank_id: str, entities_data: list[dict], unit_event_date
) -> list[str]:
"""
Trigram strategy: fetch only similar candidates per entity name using pg_trgm.
Instead of loading all bank entities (O(N)), uses a GIN trigram index to fetch
only the small set of candidates that are textually similar to each input name.
Reduces DB data transfer from 165K rows to ~5-20 rows per entity.
"""
entity_texts = list(set(e["text"] for e in entities_data))
# Fetch candidates for all unique entity texts in a single batched query.
# The trigram % operator uses the GIN index; the substring conditions cover
# exact prefix/suffix matches that trigrams might miss at low similarity.
rows = await conn.fetch(
f"""
SELECT DISTINCT ON (e.id)
e.id, e.canonical_name, e.metadata, e.last_seen, e.mention_count,
q.query_text
FROM unnest($2::text[]) AS q(query_text)
JOIN {fq_table("entities")} e ON (
e.bank_id = $1
AND (
e.canonical_name % q.query_text
OR LOWER(e.canonical_name) LIKE '%' || LOWER(q.query_text) || '%'
OR LOWER(q.query_text) LIKE '%' || LOWER(e.canonical_name) || '%'
)
)
""",
bank_id,
entity_texts,
)
# Group candidates by query_text
all_candidates: dict[str, list] = {t: [] for t in entity_texts}
candidate_ids: set = set()
for row in rows:
query_text = row["query_text"]
all_candidates[query_text].append(
(row["id"], row["canonical_name"], row["metadata"], row["last_seen"], row["mention_count"])
)
candidate_ids.add(row["id"])
# Fetch co-occurrences only for the candidate entities (not all bank entities)
cooccurrence_map: dict[str, set[str]] = {}
if candidate_ids:
candidate_id_list = list(candidate_ids)
cooc_rows = await conn.fetch(
f"""
SELECT ec.entity_id_1, ec.entity_id_2
FROM {fq_table("entity_cooccurrences")} ec
WHERE ec.entity_id_1 = ANY($1::uuid[])
OR ec.entity_id_2 = ANY($1::uuid[])
""",
candidate_id_list,
)
# Build name lookup for co-occurrence mapping
id_to_name = {
row["id"]: row["canonical_name"].lower()
for cands in all_candidates.values()
for row in [{"id": c[0], "canonical_name": c[1]} for c in cands]
}
for row in cooc_rows:
eid1, eid2 = row["entity_id_1"], row["entity_id_2"]
if eid1 not in cooccurrence_map:
cooccurrence_map[eid1] = set()
if eid2 not in cooccurrence_map:
cooccurrence_map[eid2] = set()
if eid2 in id_to_name:
cooccurrence_map[eid1].add(id_to_name[eid2])
if eid1 in id_to_name:
cooccurrence_map[eid2].add(id_to_name[eid1])
return await self._resolve_from_candidates(
conn, bank_id, entities_data, unit_event_date, all_candidates, cooccurrence_map
)
async def _resolve_from_candidates(
self,
conn,
bank_id: str,
entities_data: list[dict],
unit_event_date,
all_candidates: dict[str, list],
cooccurrence_map: dict[str, set[str]],
) -> list[str]:
"""Shared scoring + upsert logic used by both lookup strategies."""
# Resolve each entity using pre-fetched candidates
entity_ids = [None] * len(entities_data)
entities_to_update = [] # (entity_id, event_date)
entities_to_create = [] # (idx, entity_data, event_date)
entities_to_update: list[_EntityStat] = []
entities_to_create: list[_EntityToCreate] = []
for idx, entity_data in enumerate(entities_data):
entity_text = entity_data["text"]
@@ -145,7 +380,7 @@ class EntityResolver:
if not candidates:
# Will create new entity
entities_to_create.append((idx, entity_data, entity_event_date))
entities_to_create.append(_EntityToCreate(idx=idx, name=entity_text, event_date=entity_event_date))
continue
# Score candidates
@@ -189,73 +424,83 @@ class EntityResolver:
if best_score > threshold:
entity_ids[idx] = best_candidate
entities_to_update.append((best_candidate, entity_event_date))
entities_to_update.append(_EntityStat(entity_id=best_candidate, event_date=entity_event_date))
else:
entities_to_create.append((idx, entity_data, entity_event_date))
entities_to_create.append(
_EntityToCreate(idx=idx, name=entity_data["text"], event_date=entity_event_date)
)
# Batch update existing entities
if entities_to_update:
await conn.executemany(
f"""
UPDATE {fq_table("entities")} SET
mention_count = mention_count + 1,
last_seen = $2
WHERE id = $1::uuid
""",
entities_to_update,
)
# Existing entities: IDs already known from the candidate SELECT above.
# No in-transaction UPDATE — mention_count/last_seen are stats deferred to
# flush_pending_stats() which the orchestrator calls after the transaction.
pending: list[_EntityStat] = list(entities_to_update)
# Batch create new entities using COPY + INSERT for maximum speed
# This handles duplicates via ON CONFLICT and returns all IDs
# New entities: INSERT with DO NOTHING to avoid row locks on concurrent races.
# ON CONFLICT DO NOTHING returns nothing for rows that conflicted; we handle
# that rare case with a fallback SELECT.
if entities_to_create:
# Group entities by canonical name (lowercase) to handle duplicates within batch
# For duplicates, we only insert once and reuse the ID, but track the count
unique_entities = {} # lowercase_name -> (entity_data, event_date, [indices])
for idx, entity_data, event_date in entities_to_create:
name_lower = entity_data["text"].lower()
if name_lower not in unique_entities:
unique_entities[name_lower] = (entity_data, event_date, [idx])
else:
# Same entity appears multiple times - add index to list
unique_entities[name_lower][2].append(idx)
# Group by lowercase name — deduplicate within the batch.
@dataclass
class _NameGroup:
name: str
event_date: datetime | None
indices: list[int] = field(default_factory=list)
# Batch insert unique entities and get their IDs
# Use a single query with unnest for speed
entity_names = []
entity_dates = []
entity_counts = [] # Track how many times each entity appears in this batch
indices_map = [] # Maps result index -> list of original indices
groups: dict[str, _NameGroup] = {}
for e in entities_to_create:
name_lower = e.name.lower()
if name_lower not in groups:
groups[name_lower] = _NameGroup(name=e.name, event_date=e.event_date)
groups[name_lower].indices.append(e.idx)
for name_lower, (entity_data, event_date, indices) in unique_entities.items():
entity_names.append(entity_data["text"])
entity_dates.append(event_date)
entity_counts.append(len(indices)) # Count of occurrences in this batch
indices_map.append(indices)
# Sort by lowercase name for deterministic ordering.
sorted_groups = sorted(groups.items())
entity_names = [g.name for _, g in sorted_groups]
entity_dates = [g.event_date for _, g in sorted_groups]
# Batch INSERT ... ON CONFLICT with RETURNING
# Uses the batch count for mention_count instead of always 1
rows = await conn.fetch(
# INSERT ... ON CONFLICT DO NOTHING — no row lock on already-existing entities.
inserted_rows = await conn.fetch(
f"""
INSERT INTO {fq_table("entities")} (bank_id, canonical_name, first_seen, last_seen, mention_count)
SELECT $1, name, event_date, event_date, cnt
FROM unnest($2::text[], $3::timestamptz[], $4::int[]) AS t(name, event_date, cnt)
SELECT $1, name, COALESCE(event_date, now()), COALESCE(event_date, now()), 1
FROM unnest($2::text[], $3::timestamptz[]) AS t(name, event_date)
ON CONFLICT (bank_id, LOWER(canonical_name))
DO UPDATE SET
mention_count = {fq_table("entities")}.mention_count + EXCLUDED.mention_count,
last_seen = EXCLUDED.last_seen
RETURNING id
DO NOTHING
RETURNING id, LOWER(canonical_name) AS name_lower
""",
bank_id,
entity_names,
entity_dates,
entity_counts,
)
id_by_name: dict[str, str] = {row["name_lower"]: row["id"] for row in inserted_rows}
# Map returned IDs back to original indices
for result_idx, row in enumerate(rows):
entity_id = row["id"]
for original_idx in indices_map[result_idx]:
entity_ids[original_idx] = entity_id
# Fallback SELECT for names that conflicted (another worker won the race).
missing = [n for n, _ in sorted_groups if n not in id_by_name]
if missing:
existing_rows = await conn.fetch(
f"""
SELECT id, LOWER(canonical_name) AS name_lower
FROM {fq_table("entities")}
WHERE bank_id = $1 AND LOWER(canonical_name) = ANY($2::text[])
""",
bank_id,
missing,
)
for row in existing_rows:
id_by_name[row["name_lower"]] = row["id"]
# Assign entity IDs back and queue for post-txn stats flush.
for name_lower, g in sorted_groups:
entity_id = id_by_name.get(name_lower)
if entity_id:
for original_idx in g.indices:
entity_ids[original_idx] = entity_id
pending.append(_EntityStat(entity_id=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
@@ -408,7 +653,7 @@ class EntityResolver:
entity_id = await conn.fetchval(
f"""
INSERT INTO {fq_table("entities")} (bank_id, canonical_name, first_seen, last_seen, mention_count)
VALUES ($1, $2, $3, $4, 1)
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,
@@ -541,19 +786,14 @@ class EntityResolver:
entity_id_1, entity_id_2 = entity_id_2, entity_id_1
cooccurrence_pairs.add((entity_id_1, entity_id_2))
# Batch update co-occurrences
# Accumulate co-occurrence pairs for post-transaction flush.
# The actual INSERT/UPDATE is deferred to flush_pending_stats() to avoid
# row-level lock contention (ON CONFLICT DO UPDATE inside a long transaction
# serialises concurrent writers on popular entity pairs).
if cooccurrence_pairs:
now = datetime.now(UTC)
await conn.executemany(
f"""
INSERT INTO {fq_table("entity_cooccurrences")} (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
VALUES ($1, $2, $3, $4)
ON CONFLICT (entity_id_1, entity_id_2)
DO UPDATE SET
cooccurrence_count = {fq_table("entity_cooccurrences")}.cooccurrence_count + 1,
last_cooccurred = EXCLUDED.last_cooccurred
""",
[(e1, e2, 1, now) for e1, e2 in cooccurrence_pairs],
key = self._task_key()
self._pending_cooccurrences.setdefault(key, []).extend(
_CooccurrencePair(entity_id_1=e1, entity_id_2=e2) for e1, e2 in cooccurrence_pairs
)
async def get_units_by_entity(self, entity_id: str, limit: int = 100) -> list[str]:
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from hindsight_api.engine.memory_engine import Budget
from hindsight_api.engine.response_models import RecallResult, ReflectResult
from hindsight_api.engine.search.tags import TagsMatch
from hindsight_api.models import RequestContext
@@ -337,6 +338,8 @@ class MemoryEngineInterface(ABC):
bank_id: str,
*,
search_query: str | None = None,
tags: list[str] | None = None,
tags_match: "TagsMatch" = "any_strict",
limit: int = 100,
offset: int = 0,
request_context: "RequestContext",
@@ -346,7 +349,9 @@ class MemoryEngineInterface(ABC):
Args:
bank_id: The memory bank ID.
search_query: Search query.
search_query: Case-insensitive substring filter on document ID.
tags: Filter by tags.
tags_match: How to match tags (any, all, any_strict, all_strict).
limit: Maximum results.
offset: Pagination offset.
request_context: Request context for authentication.
@@ -124,6 +124,7 @@ def create_llm_provider(
vertexai_project_id: str | None = None,
vertexai_region: str | None = None,
vertexai_credentials: Any = None,
gemini_safety_settings: list | None = None,
) -> Any: # Returns LLMInterface
"""
Factory function to create the appropriate LLM provider implementation.
@@ -192,6 +193,7 @@ def create_llm_provider(
vertexai_project_id=vertexai_project_id,
vertexai_region=vertexai_region,
vertexai_credentials=vertexai_credentials,
gemini_safety_settings=gemini_safety_settings,
)
elif provider_lower == "anthropic":
@@ -234,6 +236,7 @@ class LLMProvider:
reasoning_effort: str = "low",
groq_service_tier: str | None = None,
openai_service_tier: str | None = None,
gemini_safety_settings: list | None = None,
):
"""
Initialize LLM provider.
@@ -246,6 +249,7 @@ class LLMProvider:
reasoning_effort: Reasoning effort level for supported providers.
groq_service_tier: Groq service tier ("on_demand", "flex", "auto") - from config.
openai_service_tier: OpenAI service tier (None or "flex") - from config.
gemini_safety_settings: Safety settings for Gemini/VertexAI providers.
"""
self.provider = provider.lower()
self.api_key = api_key
@@ -255,6 +259,8 @@ class LLMProvider:
# Service tiers from hierarchical config (not env vars)
self.groq_service_tier = groq_service_tier
self.openai_service_tier = openai_service_tier
# Gemini safety settings (instance default; can be overridden per-request via context var)
self.gemini_safety_settings = gemini_safety_settings
# Validate provider
valid_providers = [
@@ -323,6 +329,18 @@ 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
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
# Create provider implementation using factory
self._provider_impl = create_llm_provider(
provider=self.provider,
@@ -335,6 +353,7 @@ class LLMProvider:
vertexai_project_id=vertexai_project_id,
vertexai_region=vertexai_region,
vertexai_credentials=vertexai_credentials,
gemini_safety_settings=self.gemini_safety_settings,
)
# Backward compatibility: Keep mock provider properties
@@ -503,6 +522,14 @@ class LLMProvider:
return result
def set_response_callback(self, fn: Any) -> None:
"""Set a callback invoked on each call() instead of the fixed mock response."""
if self.provider == "mock":
from .providers.mock_llm import MockLLM
if isinstance(self._provider_impl, MockLLM):
self._provider_impl.set_response_callback(fn)
def set_mock_response(self, response: Any) -> None:
"""Set the response to return from mock calls."""
# Backward compatibility: Store in both wrapper and provider implementation
@@ -595,6 +622,23 @@ class LLMProvider:
# SDK will automatically check for authentication when first used
# No need to verify here - let it fail gracefully on first call with helpful error
def with_config(self, config: Any) -> "ConfiguredLLMProvider":
"""
Return a configured wrapper for a specific bank operation.
The wrapper applies per-bank overrides (e.g. Gemini safety settings)
to every ``call()`` / ``call_with_tools()`` invocation without
changing the underlying provider or its long-lived client connection.
Args:
config: Resolved ``HindsightConfig`` for the current bank/request.
Returns:
A ``ConfiguredLLMProvider`` that delegates to this provider with
the supplied config applied.
"""
return ConfiguredLLMProvider(self, config.llm_gemini_safety_settings)
async def cleanup(self) -> None:
"""Clean up resources."""
pass
@@ -656,5 +700,58 @@ class LLMProvider:
return cls(provider=provider, api_key=api_key, base_url=base_url, model=model, reasoning_effort="high")
class ConfiguredLLMProvider:
"""
Thin wrapper around LLMProvider that applies bank-specific config to every call.
Obtained via ``LLMProvider.with_config(resolved_config)``. The wrapper
sets any provider-specific overrides (currently Gemini safety settings)
immediately before each call using a ContextVar token, then resets it
afterwards — so nesting is safe and the configuration cannot leak across
operations.
All attribute access falls through to the underlying provider so callers
that read ``llm.provider``, ``llm.model``, etc. continue to work without
any changes.
"""
def __init__(self, provider: "LLMProvider", gemini_safety_settings: list | None) -> None:
# Use object.__setattr__ to avoid triggering __getattr__
object.__setattr__(self, "_provider", provider)
object.__setattr__(self, "_gemini_safety_settings", gemini_safety_settings)
# ── attribute passthrough ──────────────────────────────────────────────────
def __getattr__(self, name: str) -> Any:
return getattr(object.__getattribute__(self, "_provider"), name)
# ── overridden call methods ────────────────────────────────────────────────
async def call(self, messages: list[dict[str, Any]], **kwargs: Any) -> Any:
from .providers.gemini_llm import _safety_settings_ctx
token = _safety_settings_ctx.set(object.__getattribute__(self, "_gemini_safety_settings"))
try:
return await object.__getattribute__(self, "_provider").call(messages=messages, **kwargs)
finally:
_safety_settings_ctx.reset(token)
async def call_with_tools(
self,
messages: list[dict[str, Any]],
tools: list[dict[str, Any]],
**kwargs: Any,
) -> "LLMToolCallResult":
from .providers.gemini_llm import _safety_settings_ctx
token = _safety_settings_ctx.set(object.__getattribute__(self, "_gemini_safety_settings"))
try:
return await object.__getattribute__(self, "_provider").call_with_tools(
messages=messages, tools=tools, **kwargs
)
finally:
_safety_settings_ctx.reset(token)
# Backwards compatibility alias
LLMConfig = LLMProvider
File diff suppressed because it is too large Load Diff
@@ -238,21 +238,24 @@ class ClaudeCodeLLM(LLMInterface):
)
# Record trace span
from hindsight_api.tracing import get_span_recorder
try:
from hindsight_api.tracing import get_span_recorder
span_recorder = get_span_recorder()
span_recorder.record_llm_call(
provider=self.provider,
model=self.model,
scope=scope,
messages=messages,
response_content=result if isinstance(result, str) else json.dumps(result),
input_tokens=estimated_input,
output_tokens=estimated_output,
duration=duration,
finish_reason=None,
error=None,
)
span_recorder = get_span_recorder()
span_recorder.record_llm_call(
provider=self.provider,
model=self.model,
scope=scope,
messages=messages,
response_content=result if isinstance(result, str) else result.model_dump_json(),
input_tokens=estimated_input,
output_tokens=estimated_output,
duration=duration,
finish_reason=None,
error=None,
)
except Exception:
pass # logging failure must never affect the operation
# Log slow calls
if duration > 10.0:
@@ -11,6 +11,7 @@ import json
import logging
import os
import time
from contextvars import ContextVar
from typing import Any
from google import genai
@@ -24,6 +25,12 @@ from hindsight_api.metrics import get_metrics_collector
logger = logging.getLogger(__name__)
# Per-request Gemini safety settings override.
# Set exclusively by ConfiguredLLMProvider.call() / call_with_tools() via token-based
# set/reset, so it is properly scoped to each individual LLM call and never leaks.
_safety_settings_ctx: ContextVar[list | None] = ContextVar("gemini_safety_settings", default=None)
# Vertex AI imports (optional)
try:
import google.auth
@@ -58,6 +65,9 @@ class GeminiLLM(LLMInterface):
self._client = None
self._is_vertexai = self.provider == "vertexai"
# Safety settings: None means use Gemini's defaults
self._safety_settings: list | None = kwargs.get("gemini_safety_settings")
if self._is_vertexai:
self._init_vertexai(**kwargs)
else:
@@ -216,6 +226,16 @@ class GeminiLLM(LLMInterface):
if temperature is not None:
config_kwargs["temperature"] = temperature
# Apply safety settings: context var (per-request bank override) takes precedence over instance default
effective_safety_settings = _safety_settings_ctx.get()
if effective_safety_settings is None:
effective_safety_settings = self._safety_settings
if effective_safety_settings is not None:
config_kwargs["safety_settings"] = [
genai_types.SafetySetting(category=s["category"], threshold=s["threshold"])
for s in effective_safety_settings
]
generation_config = genai_types.GenerateContentConfig(**config_kwargs) if config_kwargs else None
last_exception = None
@@ -489,6 +509,16 @@ class GeminiLLM(LLMInterface):
)
# "auto" is the default (no tool_config needed)
# Apply safety settings: context var (per-request bank override) takes precedence over instance default
effective_safety_settings = _safety_settings_ctx.get()
if effective_safety_settings is None:
effective_safety_settings = self._safety_settings
if effective_safety_settings is not None:
config_kwargs["safety_settings"] = [
genai_types.SafetySetting(category=s["category"], threshold=s["threshold"])
for s in effective_safety_settings
]
config = genai_types.GenerateContentConfig(**config_kwargs)
last_exception = None
@@ -6,6 +6,7 @@ without making actual API calls to external LLM services.
"""
import logging
from collections.abc import Callable
from typing import Any
from ..llm_interface import LLMInterface
@@ -66,6 +67,7 @@ class MockLLM(LLMInterface):
self._mock_calls: list[dict] = []
self._mock_response: Any = None
self._mock_exception: Exception | None = None
self._response_callback: Callable[[list[dict], str], Any] | None = None
async def verify_connection(self) -> None:
"""
@@ -147,7 +149,9 @@ class MockLLM(LLMInterface):
)
# Return mock response
if self._mock_response is not None:
if self._response_callback is not None:
result = self._response_callback(messages, scope)
elif self._mock_response is not None:
result = self._mock_response
elif response_format is not None:
# Try to create a minimal valid instance of the response format
@@ -214,7 +218,15 @@ class MockLLM(LLMInterface):
span_recorder = get_span_recorder()
if self._mock_response is not None:
if self._response_callback is not None:
cb_result = self._response_callback(messages, scope)
if isinstance(cb_result, LLMToolCallResult):
result = cb_result
else:
result = LLMToolCallResult(
content=str(cb_result) if cb_result is not None else "mock response", finish_reason="stop"
)
elif self._mock_response is not None:
if isinstance(self._mock_response, LLMToolCallResult):
result = self._mock_response
elif isinstance(self._mock_response, list):
@@ -258,6 +270,16 @@ class MockLLM(LLMInterface):
"""Clean up resources (no-op for mock provider)."""
pass
def set_response_callback(self, fn: Callable[[list[dict], str], Any]) -> None:
"""
Set a callback invoked on each call() instead of _mock_response.
The callback receives (messages, scope) and returns the response.
Useful for returning different responses per call (e.g., cycling
through a corpus in a benchmark).
"""
self._response_callback = fn
def set_mock_response(self, response: Any) -> None:
"""
Set the response to return from mock calls.
@@ -92,11 +92,17 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
self._search_dates = None
def load(self) -> None:
"""Load dateparser (lazy import)."""
"""Load dateparser and warm up internal data structures.
Triggers the real initialization cost (regex tables, timezone data) at
load time so the first actual recall doesn't pay the cold-start penalty.
"""
if self._search_dates is None:
from dateparser.search import search_dates
self._search_dates = search_dates
# Warm up: fire a dummy call to trigger lazy-loaded internal tables.
self._search_dates("today")
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
"""
@@ -14,6 +14,8 @@ import re
import time
from typing import TYPE_CHECKING, Any, Awaitable, Callable
import tiktoken
from .models import DirectiveInfo, LLMCall, ReflectAgentResult, TokenUsageSummary, ToolCall
from .prompts import FINAL_SYSTEM_PROMPT, _extract_directive_rules, build_final_prompt, build_system_prompt_for_tools
from .tools_schema import get_reflect_tools
@@ -259,6 +261,46 @@ OUTPUT:"""
return None, 0, 0
_TIKTOKEN_ENCODING = tiktoken.get_encoding("cl100k_base")
def _count_messages_tokens(messages: list[dict[str, Any]]) -> int:
"""Estimate the token count of the messages list using cl100k_base encoding."""
total = 0
for msg in messages:
content = msg.get("content") or ""
if isinstance(content, str):
total += len(_TIKTOKEN_ENCODING.encode(content))
elif isinstance(content, list):
for part in content:
if isinstance(part, dict) and isinstance(part.get("text"), str):
total += len(_TIKTOKEN_ENCODING.encode(part["text"]))
# Tool call arguments and results also count
for tc in msg.get("tool_calls") or []:
if isinstance(tc, dict):
func = tc.get("function", {})
total += len(_TIKTOKEN_ENCODING.encode(func.get("arguments", "")))
return total
def _is_context_overflow_error(exc: Exception) -> bool:
"""Return True if the exception signals the LLM context window was exceeded."""
msg = str(exc).lower()
return any(
phrase in msg
for phrase in (
"context_length_exceeded",
"context length exceeded",
"maximum context length",
"prompt_too_long",
"prompt is too long",
"resource_exhausted",
"input is too long",
"too many tokens",
)
)
async def run_reflect_agent(
llm_config: "LLMProvider",
bank_id: str,
@@ -266,7 +308,7 @@ async def run_reflect_agent(
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], Awaitable[dict[str, Any]]],
recall_fn: Callable[[str, int, int], Awaitable[dict[str, Any]]],
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
context: str | None = None,
max_iterations: int = DEFAULT_MAX_ITERATIONS,
@@ -275,6 +317,7 @@ async def run_reflect_agent(
directives: list[dict[str, Any]] | None = None,
has_mental_models: bool = False,
budget: str | None = None,
max_context_tokens: int = 100_000,
) -> ReflectAgentResult:
"""
Execute the reflect agent loop using native tool calling.
@@ -388,7 +431,9 @@ async def run_reflect_agent(
if is_last:
# Force text response on last iteration - no tools
prompt = build_final_prompt(query, context_history, bank_profile, context)
prompt = build_final_prompt(
query, context_history, bank_profile, context, max_context_tokens=max_context_tokens
)
llm_start = time.time()
response, usage = await llm_config.call(
messages=[
@@ -433,19 +478,78 @@ async def run_reflect_agent(
directives_applied=directives_applied,
)
# Proactive context-window guard: if accumulated messages would exceed the
# configured token budget, bail out early and synthesize from what we have.
estimated_tokens = _count_messages_tokens(messages)
if estimated_tokens >= max_context_tokens and (
bool(available_memory_ids) or bool(available_mental_model_ids) or bool(available_observation_ids)
):
logger.warning(
f"[REFLECT {reflect_id}] Context budget exceeded on iteration {iteration + 1}: "
f"~{estimated_tokens} tokens >= {max_context_tokens} limit. Forcing final synthesis."
)
prompt = build_final_prompt(
query, context_history, bank_profile, context, max_context_tokens=max_context_tokens
)
llm_start = time.time()
response, usage = await llm_config.call(
messages=[
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
{"role": "user", "content": prompt},
],
scope="reflect",
max_completion_tokens=max_tokens,
return_usage=True,
)
llm_duration = int((time.time() - llm_start) * 1000)
total_input_tokens += usage.input_tokens
total_output_tokens += usage.output_tokens
llm_trace.append(
{
"scope": "final",
"duration_ms": llm_duration,
"input_tokens": usage.input_tokens,
"output_tokens": usage.output_tokens,
}
)
answer = _clean_answer_text(response.strip())
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
_log_completion(answer, iteration + 1, forced=True)
return ReflectAgentResult(
text=answer,
structured_output=structured_output,
iterations=iteration + 1,
tools_called=total_tools_called,
tool_trace=tool_trace,
llm_trace=_get_llm_trace(),
usage=_get_usage(),
directives_applied=directives_applied,
)
# Call LLM with tools
llm_start = time.time()
# Determine tool_choice for this iteration.
# Force the full hierarchical retrieval path before allowing auto:
# With mental models:
# 0 → search_mental_models, 1+ → auto
# Without mental models, enforce a minimum retrieval path:
# 0 → search_mental_models, 1 → search_observations, 2 → recall, 3+ → auto
# Without mental models:
# 0 → search_observations, 1 → recall, 2+ → auto
if iteration == 0 and has_mental_models:
iter_tool_choice: str | dict = {"type": "function", "function": {"name": "search_mental_models"}}
elif iteration == 0:
iter_tool_choice = {"type": "function", "function": {"name": "search_observations"}}
elif iteration == 1 and not has_mental_models:
elif iteration == 1 and has_mental_models:
iter_tool_choice = {"type": "function", "function": {"name": "search_observations"}}
elif iteration == 1 or (iteration == 2 and has_mental_models):
iter_tool_choice = {"type": "function", "function": {"name": "recall"}}
else:
iter_tool_choice = "auto"
@@ -475,13 +579,22 @@ async def run_reflect_agent(
consecutive_errors += 1
logger.warning(f"[REFLECT {reflect_id}] LLM error on iteration {iteration + 1}: {e} ({err_duration}ms)")
llm_trace.append({"scope": f"agent_{iteration + 1}_err", "duration_ms": err_duration})
# Guardrail: If no evidence gathered yet, retry (but cap consecutive errors to avoid long hangs)
has_gathered_evidence = (
bool(available_memory_ids) or bool(available_mental_model_ids) or bool(available_observation_ids)
)
if not has_gathered_evidence and iteration < max_iterations - 1 and consecutive_errors < 2:
# Context overflow errors must never be retried — retrying would only make them worse.
# Skip straight to final synthesis with whatever evidence we have.
if _is_context_overflow_error(e):
logger.warning(
f"[REFLECT {reflect_id}] Context window exceeded on iteration {iteration + 1}, "
"forcing final synthesis from gathered evidence."
)
# For other errors: retry if no evidence yet (but cap consecutive errors to avoid long hangs)
elif not has_gathered_evidence and iteration < max_iterations - 1 and consecutive_errors < 2:
continue
prompt = build_final_prompt(query, context_history, bank_profile, context)
prompt = build_final_prompt(
query, context_history, bank_profile, context, max_context_tokens=max_context_tokens
)
llm_start = time.time()
response, usage = await llm_config.call(
messages=[
@@ -552,7 +665,9 @@ async def run_reflect_agent(
directives_applied=directives_applied,
)
# Empty response, force final
prompt = build_final_prompt(query, context_history, bank_profile, context)
prompt = build_final_prompt(
query, context_history, bank_profile, context, max_context_tokens=max_context_tokens
)
llm_start = time.time()
response, usage = await llm_config.call(
messages=[
@@ -816,9 +931,9 @@ async def _process_done_tool(
answer = "No answer provided."
# Validate IDs (only include IDs that were actually retrieved)
used_memory_ids = [mid for mid in args.get("memory_ids", []) if mid in available_memory_ids]
used_mental_model_ids = [mid for mid in args.get("mental_model_ids", []) if mid in available_mental_model_ids]
used_observation_ids = [oid for oid in args.get("observation_ids", []) if oid in available_observation_ids]
used_memory_ids = [mid for mid in (args.get("memory_ids") or []) if mid in available_memory_ids]
used_mental_model_ids = [mid for mid in (args.get("mental_model_ids") or []) if mid in available_mental_model_ids]
used_observation_ids = [oid for oid in (args.get("observation_ids") or []) if oid in available_observation_ids]
# Generate structured output if schema provided
structured_output = None
@@ -854,7 +969,7 @@ async def _execute_tool_with_timing(
tc: "LLMToolCall",
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], Awaitable[dict[str, Any]]],
recall_fn: Callable[[str, int, int], Awaitable[dict[str, Any]]],
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
) -> tuple[dict[str, Any], int]:
"""Execute a tool call and return result with timing."""
@@ -926,7 +1041,7 @@ async def _execute_tool(
args: 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], Awaitable[dict[str, Any]]],
recall_fn: Callable[[str, int, int], Awaitable[dict[str, Any]]],
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
) -> dict[str, Any]:
"""Execute a single tool by name."""
@@ -952,7 +1067,8 @@ async def _execute_tool(
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
return await recall_fn(query, max_tokens)
max_chunk_tokens = max(int(args.get("max_chunk_tokens") or 1000), 1000) # Always enabled, min 1000
return await recall_fn(query, max_tokens, max_chunk_tokens)
elif tool_name == "expand":
memory_ids = args.get("memory_ids", [])
@@ -980,9 +1096,9 @@ def _summarize_input(tool_name: str, args: dict[str, Any]) -> str:
elif tool_name == "recall":
query = args.get("query", "")
query_preview = f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
# Show actual value used (default 2048, min 1000)
max_tokens = max(int(args.get("max_tokens") or 2048), 1000)
return f"(query={query_preview}, max_tokens={max_tokens})"
max_chunk_tokens = max(int(args.get("max_chunk_tokens") or 1000), 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", [])
depth = args.get("depth", "chunk")
@@ -10,6 +10,14 @@ The reflect agent uses hierarchical retrieval:
import json
from typing import Any
import tiktoken
_TIKTOKEN_ENCODING = tiktoken.get_encoding("cl100k_base")
# Fraction of max_context_tokens reserved for tool results in the final synthesis prompt.
# The remainder covers the system prompt, question, bank context, and output tokens.
_FINAL_PROMPT_CONTEXT_FRACTION = 0.8
def _extract_directive_rules(directives: list[dict[str, Any]]) -> list[str]:
"""Extract directive rules as a list of strings."""
@@ -286,6 +294,7 @@ def build_system_prompt_for_tools(
"- Format for clarity and readability with proper spacing and hierarchy",
"- NEVER include memory IDs, UUIDs, or 'Memory references' in the answer text",
"- Put IDs ONLY in the memory_ids/mental_model_ids/observation_ids arrays, not in the answer",
"- CRITICAL: This is a NON-CONVERSATIONAL system. NEVER ask follow-up questions, offer further assistance, or suggest next steps. Your answer must be complete and self-contained. The user cannot reply.",
]
)
@@ -393,6 +402,7 @@ def build_final_prompt(
context_history: list[dict],
bank_profile: dict,
additional_context: str | None = None,
max_context_tokens: int = 100_000,
) -> str:
"""Build the final prompt when forcing a text response (no tools)."""
parts = []
@@ -422,18 +432,32 @@ def build_final_prompt(
if additional_context:
parts.append(f"\n## Additional Context\n{additional_context}")
# Tool call history
# Tool call history — include as many entries as fit within the token budget,
# preferring the most recent calls (they tend to be the most targeted).
if context_history:
parts.append("\n## Retrieved Data (synthesize and reason from this data)")
for entry in context_history:
token_budget = int(max_context_tokens * _FINAL_PROMPT_CONTEXT_FRACTION)
# Render entries newest-first, then reverse so the prompt reads chronologically.
rendered: list[str] = []
truncated = False
for entry in reversed(context_history):
tool = entry["tool"]
output = entry["output"]
# Format as proper JSON for LLM readability
try:
output_str = json.dumps(output, indent=2, default=str)
except (TypeError, ValueError):
output_str = str(output)
parts.append(f"\n### From {tool}:\n```json\n{output_str}\n```")
block = f"\n### From {tool}:\n```json\n{output_str}\n```"
block_tokens = len(_TIKTOKEN_ENCODING.encode(block))
if block_tokens > token_budget:
truncated = True
break
rendered.append(block)
token_budget -= block_tokens
for block in reversed(rendered):
parts.append(block)
if truncated:
parts.append("\n*Note: Some earlier tool results were omitted to stay within the context window.*")
else:
parts.append("\n## Retrieved Data\nNo data was retrieved.")
@@ -481,4 +505,6 @@ CRITICAL: Output ONLY the final synthesized answer. Do NOT include:
- Meta-commentary about what you're doing ("I'll search...", "Let me analyze...")
- Explanations of your reasoning process
- Descriptions of your approach
Just provide the direct answer with proper markdown formatting."""
Just provide the direct answer with proper markdown formatting.
CRITICAL: This is a NON-CONVERSATIONAL system. NEVER ask follow-up questions, offer to search again, suggest alternatives, or end with anything like "Would you like me to..." or "Let me know if...". The user cannot reply. Your answer must be complete and self-contained."""
@@ -129,7 +129,7 @@ async def tool_search_observations(
pending_consolidation: int = 0,
) -> dict[str, Any]:
"""
Search consolidated observations using recall with include_observations.
Search consolidated observations using recall with include_source_facts.
Observations are auto-generated from memories. Returns freshness info
so the agent knows if it should also verify with recall().
@@ -146,72 +146,24 @@ async def tool_search_observations(
pending_consolidation: Number of memories waiting to be consolidated
Returns:
Dict with matching observations including freshness info
Dict with matching observations including freshness info and source memories
"""
from ..memory_engine import fq_table
# Use recall to search observations (they come back in results field when fact_type=["observation"])
result = await memory_engine.recall_async(
bank_id=bank_id,
query=query,
fact_type=["observation"], # Only retrieve observations
max_tokens=max_tokens, # Token budget controls how many observations are returned
fact_type=["observation"],
max_tokens=max_tokens,
enable_trace=False,
request_context=request_context,
tags=tags,
tags_match=tags_match,
include_source_facts=True,
max_source_facts_tokens=-1, # No token limit — include all source facts
_connection_budget=1,
_quiet=True,
)
observations = []
# When fact_type=["observation"], results come back in `results` field as MemoryFact objects
# We need to fetch additional fields (proof_count, source_memory_ids) from the database
if result.results:
obs_ids = [m.id for m in result.results]
# Fetch proof_count and source_memory_ids for these observations
pool = await memory_engine._get_pool()
async with pool.acquire() as conn:
obs_rows = await conn.fetch(
f"""
SELECT id, proof_count, source_memory_ids
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
""",
obs_ids,
)
obs_data = {str(row["id"]): row for row in obs_rows}
for m in result.results:
# Get additional data from DB lookup
extra = obs_data.get(m.id, {})
proof_count = extra.get("proof_count", 1) if extra else 1
source_ids = extra.get("source_memory_ids", []) if extra else []
# Convert UUIDs to strings
source_memory_ids = [str(sid) for sid in (source_ids or [])]
# Determine staleness
is_stale = False
staleness_reason = None
if pending_consolidation > 0:
is_stale = True
staleness_reason = f"{pending_consolidation} memories pending consolidation"
observations.append(
{
"id": str(m.id),
"text": m.text,
"proof_count": proof_count,
"source_memory_ids": source_memory_ids,
"tags": m.tags or [],
"is_stale": is_stale,
"staleness_reason": staleness_reason,
}
)
# Return freshness info (more understandable than raw pending_consolidation count)
is_stale = pending_consolidation > 0
if pending_consolidation == 0:
freshness = "up_to_date"
elif pending_consolidation < 10:
@@ -221,8 +173,10 @@ async def tool_search_observations(
return {
"query": query,
"count": len(observations),
"observations": observations,
"count": len(result.results),
"observations": [m.model_dump() for m in result.results],
"source_facts": {k: v.model_dump() for k, v in (result.source_facts or {}).items()},
"is_stale": is_stale,
"freshness": freshness,
}
@@ -233,10 +187,10 @@ async def tool_recall(
query: str,
request_context: "RequestContext",
max_tokens: int = 2048,
max_results: int = 50,
tags: list[str] | None = None,
tags_match: str = "any",
connection_budget: int = 1,
max_chunk_tokens: int = 1000,
) -> dict[str, Any]:
"""
Search memories using TEMPR retrieval.
@@ -250,18 +204,19 @@ async def tool_recall(
query: Search query
request_context: Request context for authentication
max_tokens: Maximum tokens for results (default 2048)
max_results: Maximum number of results
tags: Filter by tags (includes untagged memories)
tags_match: How to match tags - "any" (OR), "all" (AND), or "exact"
connection_budget: Max DB connections for this recall (default 1 for internal ops)
max_chunk_tokens: Maximum tokens for raw source chunk text (default 1000, always included)
Returns:
Dict with list of matching memories
Dict with list of matching memories including raw chunk text
"""
include_chunks = True
result = await memory_engine.recall_async(
bank_id=bank_id,
query=query,
fact_type=["experience", "world"], # Exclude opinions and observations
fact_type=["experience", "world"],
max_tokens=max_tokens,
enable_trace=False,
request_context=request_context,
@@ -269,24 +224,14 @@ async def tool_recall(
tags_match=tags_match,
_connection_budget=connection_budget,
_quiet=True, # Suppress logging for internal operations
include_chunks=include_chunks,
max_chunk_tokens=max_chunk_tokens,
)
memories = []
for m in result.results[:max_results]:
memories.append(
{
"id": str(m.id),
"text": m.text,
"type": m.fact_type,
"entities": m.entities or [],
"occurred": m.occurred_start, # Already ISO format string
}
)
return {
"query": query,
"count": len(memories),
"memories": memories,
"memories": [m.model_dump() for m in result.results],
"chunks": {k: v.model_dump() for k, v in (result.chunks or {}).items()},
}
@@ -96,6 +96,10 @@ TOOL_RECALL = {
"type": "integer",
"description": "Optional limit on result size (default 2048). Use higher values for broader searches.",
},
"max_chunk_tokens": {
"type": "integer",
"description": "Maximum tokens for raw source chunk text included alongside each memory fact (default 1000, min 1000). Chunks provide the surrounding context the fact was extracted from. Increase for broader context.",
},
},
"required": ["reason", "query"],
},
@@ -5,7 +5,6 @@ This package contains modular components for the retain operation:
- types: Type definitions for retain pipeline
- fact_extraction: Extract facts from content
- embedding_processing: Augment texts and generate embeddings
- deduplication: Check for duplicate facts
- entity_processing: Process and resolve entities
- link_creation: Create temporal, semantic, entity, and causal links
- chunk_storage: Handle chunk storage
@@ -14,7 +13,6 @@ This package contains modular components for the retain operation:
from . import (
chunk_storage,
deduplication,
embedding_processing,
entity_processing,
fact_extraction,
@@ -35,7 +33,6 @@ __all__ = [
# Modules
"fact_extraction",
"embedding_processing",
"deduplication",
"entity_processing",
"link_creation",
"chunk_storage",
@@ -1,85 +0,0 @@
"""
Deduplication logic for retain pipeline.
Checks for duplicate facts using semantic similarity and temporal proximity.
"""
import logging
from collections import defaultdict
from datetime import UTC
from .types import ProcessedFact
logger = logging.getLogger(__name__)
async def check_duplicates_batch(conn, bank_id: str, facts: list[ProcessedFact], duplicate_checker_fn) -> list[bool]:
"""
Check which facts are duplicates using batched time-window queries.
Groups facts by 12-hour time buckets to efficiently check for duplicates
within a 24-hour window.
Args:
conn: Database connection
bank_id: Bank identifier
facts: List of ProcessedFact objects to check
duplicate_checker_fn: Async function(conn, bank_id, texts, embeddings, date, time_window_hours)
that returns List[bool] indicating duplicates
Returns:
List of boolean flags (same length as facts) indicating if each fact is a duplicate
"""
if not facts:
return []
# Group facts by event_date (rounded to 12-hour buckets) for efficient batching
time_buckets = defaultdict(list)
for idx, fact in enumerate(facts):
# Use occurred_start if available, otherwise use mentioned_at
# For deduplication purposes, we need a time reference
fact_date = fact.occurred_start if fact.occurred_start is not None else fact.mentioned_at
# Defensive: if both are None (shouldn't happen), use now()
if fact_date is None:
from datetime import datetime
fact_date = datetime.now(UTC)
# Round to 12-hour bucket to group similar times
bucket_key = fact_date.replace(hour=(fact_date.hour // 12) * 12, minute=0, second=0, microsecond=0)
time_buckets[bucket_key].append((idx, fact))
# Process each bucket in batch
all_is_duplicate = [False] * len(facts)
for bucket_date, bucket_items in time_buckets.items():
indices = [item[0] for item in bucket_items]
texts = [item[1].fact_text for item in bucket_items]
embeddings = [item[1].embedding for item in bucket_items]
# Check duplicates for this time bucket
dup_flags = await duplicate_checker_fn(conn, bank_id, texts, embeddings, bucket_date, time_window_hours=24)
# Map results back to original indices
for idx, is_dup in zip(indices, dup_flags):
all_is_duplicate[idx] = is_dup
return all_is_duplicate
def filter_duplicates(facts: list[ProcessedFact], is_duplicate_flags: list[bool]) -> list[ProcessedFact]:
"""
Filter out duplicate facts based on duplicate flags.
Args:
facts: List of ProcessedFact objects
is_duplicate_flags: Boolean flags indicating which facts are duplicates
Returns:
List of non-duplicate facts
"""
if len(facts) != len(is_duplicate_flags):
raise ValueError(f"Mismatch between facts ({len(facts)}) and flags ({len(is_duplicate_flags)})")
return [fact for fact, is_dup in zip(facts, is_duplicate_flags) if not is_dup]
@@ -27,11 +27,21 @@ def augment_texts_with_dates(facts: list[ExtractedFact], format_date_fn) -> list
"""
augmented_texts = []
for fact in facts:
# Use occurred_start as the representative date
# Use occurred_start as the representative date, fall back to mentioned_at
fact_date = fact.occurred_start or fact.mentioned_at
readable_date = format_date_fn(fact_date)
# Augment text with date for embedding (but store original text in DB)
augmented_text = f"{fact.fact_text} (happened in {readable_date})"
# Augment text with date and entity names for embedding (but store original text in DB)
# Entity names (including key:value labels) improve retrieval without polluting stored content
if fact_date is not None:
readable_date = format_date_fn(fact_date)
if fact.occurred_end and fact.occurred_end != fact.occurred_start:
readable_end = format_date_fn(fact.occurred_end)
augmented_text = f"{fact.fact_text} (happened from {readable_date} to {readable_end})"
else:
augmented_text = f"{fact.fact_text} (happened in {readable_date})"
else:
augmented_text = fact.fact_text
if fact.entities:
augmented_text = f"{augmented_text} [{', '.join(fact.entities)}]"
augmented_texts.append(augmented_text)
return augmented_texts
@@ -41,10 +41,9 @@ async def generate_embeddings_batch(embeddings_backend, texts: list[str]) -> lis
List of embeddings in same order as input texts
"""
try:
# Run embeddings in thread pool to avoid blocking event loop
loop = asyncio.get_event_loop()
embeddings = await loop.run_in_executor(
None, # Use default thread pool
None,
embeddings_backend.encode,
texts,
)
@@ -0,0 +1,194 @@
"""
Entity labels models and helpers for retain pipeline.
Defines a controlled vocabulary of key:value classification labels
(e.g., 'pedagogy:scaffolding', 'interest:active') that are extracted
at retain time and stored as entities.
"""
from typing import Literal
from pydantic import BaseModel, Field, create_model
class LabelValue(BaseModel):
"""A single allowed value for a label group."""
value: str
description: str = ""
class LabelGroup(BaseModel):
"""A label group (dimension) with its type and allowed values."""
key: str
description: str = ""
type: Literal["value", "multi-values", "text"] = "value"
optional: bool = True
tag: bool = False
values: list[LabelValue] = []
class EntityLabelsConfig(BaseModel):
"""Entity labels configuration for a bank (controlled vocabulary)."""
attributes: list[LabelGroup] = []
def parse_entity_labels(raw: dict | list | None) -> EntityLabelsConfig | None:
"""
Parse raw entity labels config into EntityLabelsConfig.
Accepts:
- None → returns None
- list → list of attribute dicts (each may use legacy free_values/multi_value or new type field)
- dict → {attributes: [...]}
Legacy migration (backward-compat):
- free_values=True → type="text"
- multi_value=True → type="multi-values"
- neither / free_values=False → type="value"
Args:
raw: Raw entity labels config from bank config
Returns:
EntityLabelsConfig or None if raw is None/empty
"""
if raw is None:
return None
if isinstance(raw, list):
if not raw:
return None
attributes = [LabelGroup.model_validate(_migrate_label_group(a)) for a in raw]
return EntityLabelsConfig(attributes=attributes)
if isinstance(raw, dict):
attrs_raw = raw.get("attributes", [])
if not attrs_raw:
return None
attributes = [LabelGroup.model_validate(_migrate_label_group(a)) for a in attrs_raw]
return EntityLabelsConfig(attributes=attributes)
return None
def _migrate_label_group(raw: dict) -> dict:
"""Migrate legacy free_values/multi_value fields to the new type field."""
if not isinstance(raw, dict) or "type" in raw:
return raw
patched = dict(raw)
if patched.get("free_values"):
patched["type"] = "text"
elif patched.get("multi_value"):
patched["type"] = "multi-values"
else:
patched["type"] = "value"
# Remove legacy keys so Pydantic doesn't error on unknown fields
patched.pop("free_values", None)
patched.pop("multi_value", None)
return patched
def build_labels_model(labels_cfg: EntityLabelsConfig) -> type[BaseModel] | None:
"""
Build a dynamic Pydantic model for structured label extraction.
Each LabelGroup becomes a typed field based on its type:
- type="text" → str | None (always optional)
- type="value", optional=True → Literal["v1","v2"] | None
- type="value", optional=False → Literal["v1","v2"] (required)
- type="multi-values" → list[Literal["v1","v2"]]
Args:
labels_cfg: Parsed EntityLabelsConfig
Returns:
Dynamic Pydantic model class, or None if no groups defined
"""
fields: dict = {}
for group in labels_cfg.attributes:
if not group.key:
continue
description = group.description or group.key
if group.type == "text":
# Free-form: any string value accepted, always optional
fields[group.key] = (str | None, Field(default=None, description=description))
else:
# Enum-constrained: must have defined values
if not group.values:
continue
values = tuple(v.value for v in group.values if v.value)
if not values:
continue
# Literal[("v1", "v2")] is equivalent to Literal["v1", "v2"] in Python 3.11+
literal_type = Literal[values] # type: ignore[valid-type]
if group.type == "multi-values":
fields[group.key] = (
list[literal_type], # type: ignore[valid-type]
Field(default_factory=list, description=description),
)
elif group.optional:
fields[group.key] = (
literal_type | None, # type: ignore[valid-type]
Field(default=None, description=description),
)
else:
fields[group.key] = (
literal_type, # type: ignore[valid-type]
Field(description=description),
)
if not fields:
return None
return create_model("Labels", **fields)
def is_label_entity(text: str, labels_cfg: EntityLabelsConfig, labels_lookup: set[str]) -> bool:
"""
Return True if entity text belongs to any configured label group.
For enum groups: checks the pre-built lookup set.
For text groups: checks that the text starts with a known key prefix.
"""
if text.lower() in labels_lookup:
return True
for group in labels_cfg.attributes:
if group.type == "text" and group.key and text.lower().startswith(f"{group.key.lower()}:"):
return True
return False
def build_labels_lookup(labels_cfg: EntityLabelsConfig | list | None) -> set[str]:
"""
Build a set of valid 'key:value' label strings (lowercase) for fast lookup.
Accepts either EntityLabelsConfig or raw list/None for backwards compatibility.
Args:
labels_cfg: EntityLabelsConfig, raw list of attribute dicts, or None
Returns:
Set of lowercase 'key:value' strings
"""
if labels_cfg is None:
return set()
# Accept raw list/dict for backwards compatibility
if not isinstance(labels_cfg, EntityLabelsConfig):
parsed = parse_entity_labels(labels_cfg)
if parsed is None:
return set()
labels_cfg = parsed
valid = set()
for group in labels_cfg.attributes:
if group.type == "text":
continue # No fixed vocabulary — all values accepted in post-processing
for v in group.values:
if group.key and v.value:
valid.add(f"{group.key}:{v.value}".lower())
return valid
@@ -20,6 +20,7 @@ async def process_entities_batch(
facts: list[ProcessedFact],
log_buffer: list[str] = None,
user_entities_per_content: dict[int, list[dict]] = None,
entity_labels: list | None = None,
) -> list[EntityLink]:
"""
Process entities for all facts and create entity links.
@@ -90,6 +91,7 @@ async def process_entities_batch(
fact_dates,
entities_per_fact,
log_buffer, # Pass log_buffer for detailed logging
entity_labels=entity_labels,
)
return entity_links
@@ -10,22 +10,32 @@ import json
import logging
import re
from datetime import datetime, timedelta
from typing import Literal
from typing import Literal, cast
from pydantic import BaseModel, ConfigDict, Field, field_validator
from pydantic import BaseModel, ConfigDict, Field, create_model, field_validator
from ...config import get_config
from ..llm_wrapper import LLMConfig, OutputTooLongError
from ..response_models import TokenUsage
from .entity_labels import (
EntityLabelsConfig,
build_labels_lookup,
build_labels_model,
is_label_entity,
parse_entity_labels,
)
def _infer_temporal_date(fact_text: str, event_date: datetime) -> str | None:
def _infer_temporal_date(fact_text: str, event_date: datetime | None) -> str | None:
"""
Infer a temporal date from fact text when LLM didn't provide occurred_start.
This is a fallback for when the LLM fails to extract temporal information
from relative time expressions like "last night", "yesterday", etc.
"""
if event_date is None:
return None
fact_lower = fact_text.lower()
# Map relative time expressions to day offsets
@@ -100,7 +110,6 @@ class Fact(BaseModel):
# Optional temporal fields
occurred_start: str | None = None
occurred_end: str | None = None
mentioned_at: str | None = None
# Optional location field
where: str | None = Field(
@@ -440,9 +449,7 @@ _BASE_FACT_EXTRACTION_PROMPT = """Extract SIGNIFICANT facts from text. Be SELECT
LANGUAGE: MANDATORY — Detect the language of the input text and produce ALL output in that EXACT same language. You are STRICTLY FORBIDDEN from translating or switching to any other language. Every single word of your output must be in the same language as the input. Do NOT output in a different language under any circumstance.
{fact_types_instruction}
{extraction_guidelines}
{retain_mission_section}{extraction_guidelines}
══════════════════════════════════════════════════════════════════════════
FACT FORMAT - BE CONCISE
@@ -549,16 +556,16 @@ about experiences ARE important to remember, even if they seem small (e.g., how
tasted, how someone looked, how loud music was). Extract these if they characterize
an experience or person."""
# Assembled concise prompt (backward compatible - exact same output as before)
# Assembled concise prompt
CONCISE_FACT_EXTRACTION_PROMPT = _BASE_FACT_EXTRACTION_PROMPT.format(
fact_types_instruction="{fact_types_instruction}",
retain_mission_section="{retain_mission_section}",
extraction_guidelines=_CONCISE_GUIDELINES,
examples=_CONCISE_EXAMPLES,
)
# Custom prompt uses same base but without examples
CUSTOM_FACT_EXTRACTION_PROMPT = _BASE_FACT_EXTRACTION_PROMPT.format(
fact_types_instruction="{fact_types_instruction}",
retain_mission_section="{retain_mission_section}",
extraction_guidelines="{custom_instructions}",
examples="", # No examples for custom mode
)
@@ -569,8 +576,6 @@ VERBOSE_FACT_EXTRACTION_PROMPT = """Extract facts from text into structured form
LANGUAGE: MANDATORY — Detect the language of the input text and produce ALL output in that EXACT same language. You are STRICTLY FORBIDDEN from translating or switching to any other language. Every single word of your output must be in the same language as the input. Do NOT output in a different language under any circumstance.
{fact_types_instruction}
══════════════════════════════════════════════════════════════════════════
FACT FORMAT - ALL FIVE DIMENSIONS REQUIRED - MAXIMUM VERBOSITY
══════════════════════════════════════════════════════════════════════════
@@ -694,59 +699,185 @@ Example: "Lost job → couldn't pay rent → moved apartment"
- Fact 2: Moved apartment, causal_relations: [{target_index: 1, relation_type: "caused_by"}]"""
def _build_labels_prompt_section(labels_cfg: EntityLabelsConfig | list | None, free_form_entities: bool = True) -> str:
"""Build the entity labels classification section for the extraction prompt."""
if labels_cfg is None:
return ""
# Accept raw list for backwards compatibility
if isinstance(labels_cfg, list):
if not labels_cfg:
return ""
labels_cfg = parse_entity_labels(labels_cfg)
if labels_cfg is None:
return ""
if not labels_cfg.attributes:
return ""
if free_form_entities:
entities_instruction = "Classify each fact using the structured 'labels' field below. Continue extracting regular named entities in the 'entities' field."
else:
entities_instruction = "Classify each fact using the structured 'labels' field below. Do NOT add regular named entities — labels-only mode."
lines = [
"\n\n══════════════════════════════════════════════════════════════════════════",
"ENTITY LABELS - CLASSIFICATION ATTRIBUTES",
"══════════════════════════════════════════════════════════════════════════",
"",
entities_instruction,
"",
"For each fact, fill the 'labels' object. Each field is a label group:",
"",
]
for attr in labels_cfg.attributes:
if attr.type == "text":
# Free-text: no predefined values — LLM writes any relevant string or null
lines.append(f"- {attr.key} (free text or null): {attr.description}")
else:
mode = "multi-value (list)" if attr.type == "multi-values" else "single value or null"
lines.append(f"- {attr.key} ({mode}): {attr.description}")
for v in attr.values:
desc = f"{v.description}" if v.description else ""
lines.append(f'"{v.value}"{desc}')
lines.append("")
lines.append("Only assign labels when clearly applicable. Leave null/empty if the fact does not match.")
return "\n".join(lines)
def _build_extraction_prompt_and_schema(config) -> tuple[str, type]:
"""
Build extraction prompt and response schema based on config.
When a taxonomy is configured, dynamically builds a Pydantic model with a
typed `taxonomy_entities` field using an Enum built from valid taxonomy values.
This enables JSON schema enforcement for structured outputs.
Returns:
Tuple of (prompt, response_schema)
"""
fact_types_instruction = "Extract ONLY 'world' and 'assistant' type facts."
extraction_mode = config.retain_extraction_mode
extract_causal_links = config.retain_extract_causal_links
# Build retain_mission section if set - injected before the mode-specific guidelines
retain_mission = getattr(config, "retain_mission", None)
if retain_mission:
retain_mission_section = (
f"══════════════════════════════════════════════════════════════════════════\n"
f"FOCUS — What to retain for this bank\n"
f"══════════════════════════════════════════════════════════════════════════\n\n"
f"{retain_mission}\n\n"
)
else:
retain_mission_section = ""
# Select base prompt based on extraction mode
if extraction_mode == "custom":
if not config.retain_custom_instructions:
base_prompt = CONCISE_FACT_EXTRACTION_PROMPT
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
prompt = base_prompt.format(
retain_mission_section=retain_mission_section,
)
else:
base_prompt = CUSTOM_FACT_EXTRACTION_PROMPT
prompt = base_prompt.format(
fact_types_instruction=fact_types_instruction,
retain_mission_section=retain_mission_section,
custom_instructions=config.retain_custom_instructions,
)
elif extraction_mode == "verbose":
base_prompt = VERBOSE_FACT_EXTRACTION_PROMPT
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
prompt = VERBOSE_FACT_EXTRACTION_PROMPT
else:
base_prompt = CONCISE_FACT_EXTRACTION_PROMPT
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
prompt = base_prompt.format(
retain_mission_section=retain_mission_section,
)
# Add causal relationships section if enabled
if extract_causal_links:
prompt = prompt + CAUSAL_RELATIONSHIPS_SECTION
response_schema = FactExtractionResponseVerbose if extraction_mode == "verbose" else FactExtractionResponse
base_fact_class = ExtractedFactVerbose if extraction_mode == "verbose" else ExtractedFact
base_response_class = FactExtractionResponseVerbose if extraction_mode == "verbose" else FactExtractionResponse
else:
response_schema = FactExtractionResponseNoCausal
base_fact_class = ExtractedFactNoCausal
base_response_class = FactExtractionResponseNoCausal
# Add entity labels section if configured and build dynamic schema
entity_labels_raw = getattr(config, "entity_labels", None)
labels_cfg = parse_entity_labels(entity_labels_raw)
free_form_entities = getattr(config, "entities_allow_free_form", True)
labels_section = _build_labels_prompt_section(labels_cfg, free_form_entities)
if labels_section:
prompt = prompt + labels_section
response_schema = base_response_class
if labels_cfg and labels_cfg.attributes:
LabelsModel = build_labels_model(labels_cfg)
if LabelsModel is not None:
dynamic_fields: dict = {
"labels": (
LabelsModel,
Field(
description="Classification labels for this fact. Fill each applicable field; leave others null/empty."
),
)
}
if not free_form_entities:
dynamic_fields["entities"] = (
list[Entity] | None,
Field(default=None, 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)
base_extra = base_fact_class.model_config.get("json_schema_extra")
base_required = cast(dict, base_extra).get("required", []) if isinstance(base_extra, dict) else []
DynamicFact = create_model(
"LabelsFact",
__base__=base_fact_class,
__config__=ConfigDict(
json_schema_mode="validation",
json_schema_extra={"required": [*base_required, "labels"]},
),
**dynamic_fields,
)
DynamicResponse = create_model("LabelsResponse", facts=(list[DynamicFact], ...)) # type: ignore[valid-type]
response_schema = DynamicResponse
return prompt, response_schema
def _build_user_message(chunk: str, chunk_index: int, total_chunks: int, event_date: datetime, context: str) -> str:
def _build_user_message(
chunk: str,
chunk_index: int,
total_chunks: int,
event_date: datetime | None,
context: str,
metadata: dict[str, str] | None = None,
) -> str:
"""Build user message for fact extraction."""
from .orchestrator import parse_datetime_flexible
sanitized_chunk = _sanitize_text(chunk)
sanitized_context = _sanitize_text(context) if context else "none"
event_date = parse_datetime_flexible(event_date)
event_date_formatted = event_date.strftime("%A, %B %d, %Y")
if event_date is not None:
event_date = parse_datetime_flexible(event_date)
event_date_str = f"{event_date.strftime('%A, %B %d, %Y')} ({event_date.isoformat()})"
else:
event_date_str = "Unknown"
metadata_section = ""
if metadata:
metadata_lines = "\n".join(f" {k}: {v}" for k, v in metadata.items())
metadata_section = f"\nMetadata:\n{metadata_lines}"
return f"""Extract facts from the following text chunk.
Chunk: {chunk_index + 1}/{total_chunks}
Event Date: {event_date_formatted} ({event_date.isoformat()})
Context: {sanitized_context}
Event Date: {event_date_str}
Context: {sanitized_context}{metadata_section}
Text:
{sanitized_chunk}"""
@@ -783,11 +914,12 @@ async def _extract_facts_from_chunk(
chunk: str,
chunk_index: int,
total_chunks: int,
event_date: datetime,
event_date: datetime | None,
context: str,
llm_config: "LLMConfig",
config,
agent_name: str = None,
metadata: dict[str, str] | None = None,
) -> tuple[list[dict[str, str]], TokenUsage]:
"""
Extract facts from a single chunk (internal helper for parallel processing).
@@ -809,7 +941,7 @@ async def _extract_facts_from_chunk(
extract_causal_links = config.retain_extract_causal_links
# Build user message using helper function
user_message = _build_user_message(chunk, chunk_index, total_chunks, event_date, context)
user_message = _build_user_message(chunk, chunk_index, total_chunks, event_date, context, metadata)
# Retry logic for JSON validation errors
max_retries = 2
@@ -968,9 +1100,9 @@ async def _extract_facts_from_chunk(
# 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
validated_entities = []
for ent in entities:
if isinstance(ent, str):
# Normalize string to Entity object
@@ -980,8 +1112,48 @@ async def _extract_facts_from_chunk(
validated_entities.append(Entity.model_validate(ent))
except Exception as e:
logger.warning(f"Invalid entity {ent}: {e}")
if validated_entities:
fact_data["entities"] = validated_entities
# Post-process label entities from structured labels object
entity_labels_raw = getattr(config, "entity_labels", None)
labels_cfg = parse_entity_labels(entity_labels_raw)
free_form_entities = getattr(config, "entities_allow_free_form", True)
if labels_cfg and labels_cfg.attributes:
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}
for group in labels_cfg.attributes:
value = labels_data.get(group.key)
if not value:
continue
values_list = value if isinstance(value, list) else [value]
for v in values_list:
if not isinstance(v, str) or not v.strip() or v.lower() in ("none", "null", "n/a"):
continue
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))
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))
existing_texts_lower.add(label_str.lower())
else:
logger.warning(f"Label '{label_str}' not in valid label values, skipping")
# 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)
]
elif not free_form_entities:
# No labels but free_form disabled: clear all entities
validated_entities = []
if validated_entities:
fact_data["entities"] = validated_entities
# Add per-fact causal relations (only if enabled in config)
if extract_causal_links:
@@ -1020,8 +1192,9 @@ async def _extract_facts_from_chunk(
if validated_relations:
fact_data["causal_relations"] = validated_relations
# Always set mentioned_at to the event_date (when the conversation/document occurred)
fact_data["mentioned_at"] = event_date.isoformat()
# Set mentioned_at to the event_date (when the conversation/document occurred),
# or None when the caller opted into no timestamp.
fact_data["mentioned_at"] = event_date.isoformat() if event_date is not None else None
# Build Fact model instance
try:
@@ -1084,11 +1257,12 @@ async def _extract_facts_with_auto_split(
chunk: str,
chunk_index: int,
total_chunks: int,
event_date: datetime,
event_date: datetime | None,
context: str,
llm_config: LLMConfig,
config,
agent_name: str = None,
metadata: dict[str, str] | None = None,
) -> tuple[list[dict[str, str]], TokenUsage]:
"""
Extract facts from a chunk with automatic splitting if output exceeds token limits.
@@ -1105,6 +1279,7 @@ async def _extract_facts_with_auto_split(
llm_config: LLM configuration to use
config: Resolved HindsightConfig for this bank
agent_name: Optional agent name (memory owner)
metadata: Optional document metadata key-value pairs
Returns:
Tuple of (facts list, token usage) extracted from the chunk (possibly from sub-chunks)
@@ -1124,6 +1299,7 @@ async def _extract_facts_with_auto_split(
llm_config=llm_config,
config=config,
agent_name=agent_name,
metadata=metadata,
)
except OutputTooLongError:
# Output exceeded token limits - split the chunk in half and retry
@@ -1169,6 +1345,7 @@ async def _extract_facts_with_auto_split(
llm_config=llm_config,
config=config,
agent_name=agent_name,
metadata=metadata,
),
_extract_facts_with_auto_split(
chunk=second_half,
@@ -1179,6 +1356,7 @@ async def _extract_facts_with_auto_split(
llm_config=llm_config,
config=config,
agent_name=agent_name,
metadata=metadata,
),
]
@@ -1198,11 +1376,12 @@ async def _extract_facts_with_auto_split(
async def extract_facts_from_text(
text: str,
event_date: datetime,
event_date: datetime | None,
llm_config: LLMConfig,
agent_name: str,
config,
context: str = "",
metadata: dict[str, str] | None = None,
) -> tuple[list[Fact], list[tuple[str, int]], TokenUsage]:
"""
Extract semantic facts from conversational or narrative text using LLM.
@@ -1220,6 +1399,7 @@ async def extract_facts_from_text(
agent_name: Agent name (memory owner)
config: Resolved HindsightConfig for this bank
context: Context about the conversation/document
metadata: Optional document metadata key-value pairs
Returns:
Tuple of (facts, chunks, usage) where:
@@ -1247,6 +1427,7 @@ async def extract_facts_from_text(
llm_config=llm_config,
config=config,
agent_name=agent_name,
metadata=metadata,
)
for i, chunk in enumerate(chunks)
]
@@ -1356,7 +1537,7 @@ async def extract_facts_from_contents_batch_api(
# Build user message using helper function
user_message = _build_user_message(
chunk, chunk_index_in_content, len(chunks), item.event_date, item.context
chunk, chunk_index_in_content, len(chunks), item.event_date, item.context, item.metadata or None
)
# Build request body using helper function
@@ -1568,8 +1749,8 @@ async def extract_facts_from_contents_batch_api(
# Entities
entities = get_value("entities")
validated_entities = []
if entities:
validated_entities = []
for ent in entities:
if isinstance(ent, str):
validated_entities.append(Entity(text=ent))
@@ -1578,8 +1759,45 @@ async def extract_facts_from_contents_batch_api(
validated_entities.append(Entity.model_validate(ent))
except Exception:
pass
if validated_entities:
fact_data["entities"] = validated_entities
# Post-process label entities from structured labels object
entity_labels_raw = getattr(config, "entity_labels", None)
labels_cfg_batch = parse_entity_labels(entity_labels_raw)
free_form_entities_batch = getattr(config, "entities_allow_free_form", True)
if labels_cfg_batch and labels_cfg_batch.attributes:
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}
for group in labels_cfg_batch.attributes:
value = labels_data.get(group.key)
if not value:
continue
values_list = value if isinstance(value, list) else [value]
for v in values_list:
if not isinstance(v, str) or not v.strip() or v.lower() in ("none", "null", "n/a"):
continue
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))
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))
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)
]
elif not free_form_entities_batch:
validated_entities = []
if validated_entities:
fact_data["entities"] = validated_entities
# Causal relations
if extract_causal_links:
@@ -1610,8 +1828,9 @@ async def extract_facts_from_contents_batch_api(
if validated_relations:
fact_data["causal_relations"] = validated_relations
# Always set mentioned_at
fact_data["mentioned_at"] = event_date.isoformat()
# Set mentioned_at to the event_date (when the conversation/document occurred),
# or None when the caller opted into no timestamp.
fact_data["mentioned_at"] = event_date.isoformat() if event_date is not None else None
try:
fact = Fact(fact=combined_text, fact_type=fact_type, **fact_data)
@@ -1670,6 +1889,7 @@ async def extract_facts_from_contents_batch_api(
mentioned_at=content.event_date,
metadata=content.metadata,
tags=content.tags,
observation_scopes=content.observation_scopes,
)
extracted_facts.append(extracted_fact)
@@ -1678,6 +1898,9 @@ async def extract_facts_from_contents_batch_api(
# Step 7: Add temporal offsets
_add_temporal_offsets(extracted_facts, contents)
# Step 8: Auto-tag facts from label groups with tag=True
_inject_label_tags(extracted_facts, config)
logger.info(f"Batch API extracted {len(extracted_facts)} facts from {len(all_chunks_info)} chunks")
return extracted_facts, chunks_metadata, total_usage
@@ -1736,6 +1959,7 @@ async def extract_facts_from_contents(
llm_config=llm_config,
agent_name=agent_name,
config=config,
metadata=item.metadata or None,
)
fact_extraction_tasks.append(task)
@@ -1799,6 +2023,7 @@ async def extract_facts_from_contents(
mentioned_at=content.event_date,
metadata=content.metadata,
tags=content.tags,
observation_scopes=content.observation_scopes,
)
extracted_facts.append(extracted_fact)
@@ -1808,6 +2033,9 @@ async def extract_facts_from_contents(
# Step 4: Add time offsets to preserve ordering within each content
_add_temporal_offsets(extracted_facts, contents)
# Step 5: Auto-tag facts from label groups with tag=True
_inject_label_tags(extracted_facts, config)
return extracted_facts, chunks_metadata, total_usage
@@ -1863,3 +2091,24 @@ def _add_temporal_offsets(facts: list[ExtractedFactType], contents: list[RetainC
fact.occurred_end = parse_datetime_flexible(fact.occurred_end) + offset
if fact.mentioned_at:
fact.mentioned_at = parse_datetime_flexible(fact.mentioned_at) + offset
def _inject_label_tags(facts: list[ExtractedFactType], config) -> None:
"""
For label groups with tag=True, add extracted key:value label entities
to each fact's tags list. Modifies facts in place.
This lets entity labels double as tags, enabling filtering via the
existing tags API without any extra query infrastructure.
"""
labels_cfg = parse_entity_labels(getattr(config, "entity_labels", None))
if not labels_cfg:
return
tag_group_keys = {g.key.lower() for g in labels_cfg.attributes if g.tag}
if not tag_group_keys:
return
for fact in facts:
label_tags = [e for e in fact.entities if ":" in e and e.split(":", 1)[0].lower() in tag_group_keys]
if label_tags:
existing = set(fact.tags)
fact.tags = fact.tags + [t for t in label_tags if t not in existing]
@@ -47,6 +47,8 @@ async def insert_facts_batch(
chunk_ids = []
document_ids = []
tags_list = []
observation_scopes_list = []
text_signals_list = []
for fact in facts:
fact_texts.append(_sanitize_text(fact.fact_text))
@@ -68,6 +70,19 @@ async def insert_facts_batch(
document_ids.append(fact.document_id if fact.document_id else document_id)
# Convert tags to JSON string for proper batch insertion (PostgreSQL unnest doesn't handle 2D arrays well)
tags_list.append(json.dumps(fact.tags if fact.tags else []))
# observation_scopes: stored as JSONB (string or 2D array), None if not provided
observation_scopes_list.append(
json.dumps(fact.observation_scopes) if fact.observation_scopes is not None else None
)
# Build text_signals: entity names + date tokens for enriched BM25 indexing
signal_parts = []
if fact.entities:
signal_parts.extend(e.name for e in fact.entities)
if fact.occurred_start:
signal_parts.append(fact.occurred_start.strftime("%B %-d %Y"))
if fact.occurred_end and fact.occurred_end != fact.occurred_start:
signal_parts.append(fact.occurred_end.strftime("%B %-d %Y"))
text_signals_list.append(" ".join(signal_parts) if signal_parts else None)
# Batch insert all facts
# Note: tags are passed as JSON strings and converted back to varchar[] via jsonb_array_elements_text + array_agg
@@ -75,16 +90,19 @@ async def insert_facts_batch(
config = get_config()
if config.text_search_extension == "vchord":
# VectorChord: manually tokenize and insert search_vector
# text_signals (entity names etc.) are included in the tokenize input for enriched BM25
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::float[], $11::jsonb[], $12::text[], $13::text[], $14::jsonb[]
$8::text[], $9::text[], $10::float[], $11::jsonb[], $12::text[], $13::text[], $14::jsonb[], $15::jsonb[], $16::text[]
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags_json)
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags_json,
observation_scopes_json, text_signals)
)
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags, search_vector)
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags,
observation_scopes, text_signals, search_vector)
SELECT
$1,
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
@@ -93,23 +111,30 @@ async def insert_facts_batch(
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
'{{}}'::varchar[]
),
tokenize(COALESCE(text, '') || ' ' || COALESCE(context, ''), 'llmlingua2')::bm25_catalog.bm25vector
observation_scopes_json,
text_signals,
tokenize(
COALESCE(text, '') || ' ' || COALESCE(context, '') || ' ' || COALESCE(text_signals, ''),
'llmlingua2'
)::bm25_catalog.bm25vector
FROM input_data
RETURNING id
"""
else: # native or pg_textsearch
# Native PostgreSQL: search_vector is GENERATED ALWAYS, don't include it
# Native PostgreSQL: search_vector is GENERATED ALWAYS (expression includes text_signals), don't include it
# pg_textsearch: indexes operate on base columns directly, don't populate search_vector
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::float[], $11::jsonb[], $12::text[], $13::text[], $14::jsonb[]
$8::text[], $9::text[], $10::float[], $11::jsonb[], $12::text[], $13::text[], $14::jsonb[], $15::jsonb[], $16::text[]
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags_json)
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags_json,
observation_scopes_json, text_signals)
)
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags)
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags,
observation_scopes, text_signals)
SELECT
$1,
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
@@ -117,7 +142,9 @@ async def insert_facts_batch(
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
"""
@@ -138,6 +165,8 @@ async def insert_facts_batch(
chunk_ids,
document_ids,
tags_list,
observation_scopes_list,
text_signals_list,
)
unit_ids = [str(row["id"]) for row in results]
@@ -47,6 +47,9 @@ def compute_temporal_links(
links = []
for unit_id, unit_event_date in new_units.items():
# Units without event_date can't form temporal links
if unit_event_date is None:
continue
# Normalize unit_event_date for consistent comparison
unit_event_date_norm = _normalize_datetime(unit_event_date)
@@ -96,7 +99,11 @@ def compute_temporal_query_bounds(
return None, None
# Normalize all dates to be timezone-aware to avoid comparison issues
all_dates = [_normalize_datetime(d) for d in new_units.values()]
# Filter out None values — units without event_date can't form temporal links
all_dates = [_normalize_datetime(d) for d in new_units.values() if d is not None]
if not all_dates:
return None, None
try:
min_date = min(all_dates) - timedelta(hours=time_window_hours)
@@ -143,6 +150,7 @@ async def extract_entities_batch_optimized(
fact_dates: list,
llm_entities: list[list[dict]],
log_buffer: list[str] = None,
entity_labels: list | None = None,
) -> list[tuple]:
"""
Process LLM-extracted entities for ALL facts in batch.
@@ -232,6 +240,7 @@ async def extract_entities_batch_optimized(
context=context,
unit_event_date=None, # Not used when per-entity dates provided
conn=conn, # Use main transaction connection
entity_labels=entity_labels,
)
_log(
@@ -432,20 +441,23 @@ async def create_temporal_links_batch_per_fact(
min_date, max_date = compute_temporal_query_bounds(new_units, time_window_hours)
fetch_neighbors_start = time_mod.time()
all_candidates = await conn.fetch(
f"""
SELECT id, event_date
FROM {fq_table("memory_units")}
WHERE bank_id = $1
AND event_date BETWEEN $2 AND $3
AND id::text != ALL($4)
ORDER BY event_date DESC
""",
bank_id,
min_date,
max_date,
unit_ids,
)
if min_date is not None and max_date is not None:
all_candidates = await conn.fetch(
f"""
SELECT id, event_date
FROM {fq_table("memory_units")}
WHERE bank_id = $1
AND event_date BETWEEN $2 AND $3
AND id::text != ALL($4)
ORDER BY event_date DESC
""",
bank_id,
min_date,
max_date,
unit_ids,
)
else:
all_candidates = []
_log(
log_buffer,
f" [7.2] Fetch {len(all_candidates)} candidate neighbors (1 query): {time_mod.time() - fetch_neighbors_start:.3f}s",
@@ -460,11 +472,15 @@ async def create_temporal_links_batch_per_fact(
# Convert new_units dict to candidate format for within-batch linking
new_unit_items = list(new_units.items())
for i, (unit_id, event_date) in enumerate(new_unit_items):
if event_date is None:
continue # Skip units without event_date for temporal linking
unit_event_date_norm = _normalize_datetime(event_date)
# Compare with other new units (only those after this one to avoid duplicates)
for j in range(i + 1, len(new_unit_items)):
other_id, other_event_date = new_unit_items[j]
if other_event_date is None:
continue # Skip units without event_date
other_event_date_norm = _normalize_datetime(other_event_date)
# Check if within time window
@@ -482,14 +498,13 @@ async def create_temporal_links_batch_per_fact(
# Batch inserts to avoid timeout on large batches
BATCH_SIZE = 1000
for batch_start in range(0, len(links), BATCH_SIZE):
batch = links[batch_start : batch_start + BATCH_SIZE]
await conn.executemany(
f"""
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
batch,
links[batch_start : batch_start + BATCH_SIZE],
)
_log(log_buffer, f" [7.4] Insert {len(links)} temporal links: {time_mod.time() - insert_start:.3f}s")
@@ -537,81 +552,45 @@ async def create_semantic_links_batch(
import numpy as np
# Fetch ALL existing units with embeddings in ONE query
fetch_start = time_mod.time()
all_existing = await conn.fetch(
f"""
SELECT id, embedding
FROM {fq_table("memory_units")}
WHERE bank_id = $1
AND embedding IS NOT NULL
AND id::text != ALL($2)
""",
bank_id,
unit_ids,
)
_log(
log_buffer,
f" [8.1] Fetch {len(all_existing)} existing embeddings (1 query): {time_mod.time() - fetch_start:.3f}s",
)
# Convert to numpy for vectorized similarity computation
compute_start = time_mod.time()
# Use pgvector ANN search (HNSW index) for each new unit instead of fetching
# all existing embeddings into Python. At large scale (100K+ units) the old
# approach would transfer 100K × 384 floats (~150 MB) per retain call; the
# ANN query completes in <5 ms and transfers only top_k rows.
ann_start = time_mod.time()
all_links = []
if all_existing:
# Convert existing embeddings to numpy array
existing_ids = [str(row["id"]) for row in all_existing]
# Stack embeddings as 2D array: (num_embeddings, embedding_dim)
embedding_arrays = []
for row in all_existing:
raw_emb = row["embedding"]
# Handle different pgvector formats
if isinstance(raw_emb, str):
# Parse string format: "[1.0, 2.0, ...]"
import json
# Build UUID exclude list once for all ANN queries
import uuid as uuid_mod
emb = np.array(json.loads(raw_emb), dtype=np.float32)
elif isinstance(raw_emb, (list, tuple)):
emb = np.array(raw_emb, dtype=np.float32)
else:
# Try direct conversion (works for numpy arrays, pgvector objects, etc.)
emb = np.array(raw_emb, dtype=np.float32)
exclude_uuids = [uuid_mod.UUID(uid) if isinstance(uid, str) else uid for uid in unit_ids]
# Ensure it's 1D
if emb.ndim != 1:
raise ValueError(f"Expected 1D embedding, got shape {emb.shape}")
embedding_arrays.append(emb)
for unit_id, new_embedding in zip(unit_ids, embeddings):
emb_str = str(list(new_embedding) if not isinstance(new_embedding, list) else new_embedding)
rows = await conn.fetch(
f"""
SELECT id::text,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND embedding IS NOT NULL
AND id != ALL($3::uuid[])
ORDER BY embedding <=> $1::vector
LIMIT $4
""",
emb_str,
bank_id,
exclude_uuids,
top_k,
)
for row in rows:
sim = float(min(1.0, max(0.0, row["similarity"])))
if sim >= threshold:
all_links.append((unit_id, str(row["id"]), "semantic", sim, None))
if not embedding_arrays:
existing_embeddings = np.array([])
elif len(embedding_arrays) == 1:
# Single embedding: reshape to (1, dim)
existing_embeddings = embedding_arrays[0].reshape(1, -1)
else:
# Multiple embeddings: vstack
existing_embeddings = np.vstack(embedding_arrays)
# For each new unit, compute similarities with ALL existing units
for unit_id, new_embedding in zip(unit_ids, embeddings):
new_emb_array = np.array(new_embedding)
# Compute cosine similarities (dot product for normalized vectors)
similarities = np.dot(existing_embeddings, new_emb_array)
# Find top-k above threshold
# Get indices of similarities above threshold
above_threshold = np.where(similarities >= threshold)[0]
if len(above_threshold) > 0:
# Sort by similarity (descending) and take top-k
sorted_indices = above_threshold[np.argsort(-similarities[above_threshold])][:top_k]
for idx in sorted_indices:
similar_id = existing_ids[idx]
# Clamp to [0, 1] to handle floating point precision issues
similarity = float(min(1.0, max(0.0, similarities[idx])))
all_links.append((unit_id, similar_id, "semantic", similarity, None))
_log(
log_buffer,
f" [8.1] ANN search for {len(unit_ids)} new units → {len(all_links)} candidate links: {time_mod.time() - ann_start:.3f}s",
)
# Also compute similarities WITHIN the new batch (new units to each other)
# Apply the same top_k limit per unit as we do for existing units
@@ -643,7 +622,7 @@ async def create_semantic_links_batch(
_log(
log_buffer,
f" [8.2] Compute similarities & generate {len(all_links)} semantic links: {time_mod.time() - compute_start:.3f}s",
f" [8.2] Within-batch similarities added {len(all_links)} total semantic links",
)
if all_links:
@@ -651,14 +630,13 @@ async def create_semantic_links_batch(
# Batch inserts to avoid timeout on large batches
BATCH_SIZE = 1000
for batch_start in range(0, len(all_links), BATCH_SIZE):
batch = all_links[batch_start : batch_start + BATCH_SIZE]
await conn.executemany(
f"""
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
batch,
all_links[batch_start : batch_start + BATCH_SIZE],
)
_log(
log_buffer, f" [8.3] Insert {len(all_links)} semantic links: {time_mod.time() - insert_start:.3f}s"
@@ -674,18 +652,18 @@ async def create_semantic_links_batch(
raise
async def insert_entity_links_batch(conn, links: list[EntityLink], chunk_size: int = 50000):
async def insert_entity_links_batch(conn, links: list[EntityLink], chunk_size: int = 5000):
"""
Insert all entity links using COPY to temp table + INSERT for maximum speed.
Insert all entity links using COPY to temp table + chunked INSERT for reliability.
Uses PostgreSQL COPY (via copy_records_to_table) for bulk loading,
then INSERT ... ON CONFLICT from temp table. This is the fastest
method for bulk inserts with conflict handling.
Uses PostgreSQL COPY (via copy_records_to_table) for bulk loading into a
temp table, then INSERT ... ON CONFLICT in chunks of chunk_size. Chunking
prevents single-query timeouts on very large tables (100M+ rows).
Args:
conn: Database connection
links: List of EntityLink objects
chunk_size: Number of rows per batch (default 50000)
chunk_size: Number of rows per INSERT chunk (default 5000)
"""
if not links:
return
@@ -694,10 +672,11 @@ async def insert_entity_links_batch(conn, links: list[EntityLink], chunk_size: i
total_start = time_mod.time()
# Create temp table for bulk loading
# Create temp table with serial for stable chunked access
create_start = time_mod.time()
await conn.execute("""
CREATE TEMP TABLE IF NOT EXISTS _temp_entity_links (
_row_num SERIAL,
from_unit_id uuid,
to_unit_id uuid,
link_type text,
@@ -714,9 +693,7 @@ async def insert_entity_links_batch(conn, links: list[EntityLink], chunk_size: i
# Convert EntityLink objects to tuples for COPY
convert_start = time_mod.time()
records = []
for link in links:
records.append((link.from_unit_id, link.to_unit_id, link.link_type, link.weight, link.entity_id))
records = [(link.from_unit_id, link.to_unit_id, link.link_type, link.weight, link.entity_id) for link in links]
logger.debug(f" [9.3] Convert {len(records)} records: {time_mod.time() - convert_start:.3f}s")
# Bulk load using COPY (fastest method)
@@ -728,15 +705,25 @@ async def insert_entity_links_batch(conn, links: list[EntityLink], chunk_size: i
)
logger.debug(f" [9.4] COPY {len(records)} records to temp table: {time_mod.time() - copy_start:.3f}s")
# Insert from temp table with ON CONFLICT (single query for all rows)
# Insert from temp table in chunks to avoid single-query timeouts on large tables
insert_start = time_mod.time()
await conn.execute(f"""
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
SELECT from_unit_id, to_unit_id, link_type, weight, entity_id
FROM _temp_entity_links
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""")
logger.debug(f" [9.5] INSERT from temp table: {time_mod.time() - insert_start:.3f}s")
total_rows = len(records)
chunks = 0
for chunk_start in range(0, total_rows, chunk_size):
chunk_end = chunk_start + chunk_size
await conn.execute(
f"""
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
SELECT from_unit_id, to_unit_id, link_type, weight, entity_id
FROM _temp_entity_links
WHERE _row_num > $1 AND _row_num <= $2
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
chunk_start,
chunk_end,
)
chunks += 1
logger.debug(f" [9.5] INSERT {total_rows} rows in {chunks} chunks: {time_mod.time() - insert_start:.3f}s")
logger.debug(f" [9.TOTAL] Entity links batch insert: {time_mod.time() - total_start:.3f}s")
@@ -55,7 +55,6 @@ def parse_datetime_flexible(value: Any) -> datetime:
from ..response_models import TokenUsage
from . import (
chunk_storage,
deduplication,
embedding_processing,
entity_processing,
fact_extraction,
@@ -73,7 +72,6 @@ async def retain_batch(
llm_config,
entity_resolver,
format_date_fn,
duplicate_checker_fn,
bank_id: str,
contents_dicts: list[RetainContentDict],
config,
@@ -94,7 +92,6 @@ async def retain_batch(
llm_config: LLM configuration for fact extraction
entity_resolver: Entity resolver for entity processing
format_date_fn: Function to format datetime to readable string
duplicate_checker_fn: Function to check for duplicate facts
bank_id: Bank identifier
contents_dicts: List of content dictionaries
config: Resolved HindsightConfig for this bank
@@ -128,12 +125,14 @@ async def retain_batch(
item_tags = item.get("tags", []) or []
merged_tags = list(set(item_tags + (document_tags or [])))
# Handle event_date: parse flexibly (handles both datetime objects and ISO strings)
event_date_value = item.get("event_date")
if event_date_value:
event_date_value = parse_datetime_flexible(event_date_value)
# Handle event_date: distinguish "not provided" (default to now) from
# "explicitly None" (caller opted into no timestamp).
if "event_date" in item and item["event_date"] is None:
event_date_value = None # Caller explicitly signalled "unknown date"
elif item.get("event_date"):
event_date_value = parse_datetime_flexible(item["event_date"])
else:
event_date_value = utcnow()
event_date_value = utcnow() # Backward-compatible default
content = RetainContent(
content=item["content"],
@@ -142,6 +141,7 @@ async def retain_batch(
metadata=item.get("metadata", {}),
entities=item.get("entities", []),
tags=merged_tags,
observation_scopes=item.get("observation_scopes"),
)
contents.append(content)
@@ -162,8 +162,6 @@ async def retain_batch(
docs_tracked = 0
async with acquire_with_retry(pool) as conn:
async with conn.transaction():
await fact_storage.ensure_bank_exists(conn, bank_id)
# Group contents by document_id (consistent with normal path)
contents_by_doc_early = defaultdict(list)
for idx, content_dict in enumerate(contents_dicts):
@@ -281,9 +279,6 @@ async def retain_batch(
# Step 4: Database transaction
async with acquire_with_retry(pool) as conn:
async with conn.transaction():
# Ensure bank exists
await fact_storage.ensure_bank_exists(conn, bank_id)
# Handle document tracking for all documents
step_start = time.time()
# Map None document_id to generated UUIDs
@@ -435,20 +430,7 @@ async def retain_batch(
actual_doc_id = document_id
processed_fact.document_id = actual_doc_id
# Deduplication
step_start = time.time()
is_duplicate_flags = await deduplication.check_duplicates_batch(
conn, bank_id, processed_facts, duplicate_checker_fn
)
log_buffer.append(
f"[4] Deduplication: {sum(is_duplicate_flags)} duplicates in {time.time() - step_start:.3f}s"
)
# Filter out duplicates
non_duplicate_facts = deduplication.filter_duplicates(processed_facts, is_duplicate_flags)
if not non_duplicate_facts:
return [[] for _ in contents], usage
non_duplicate_facts = processed_facts
# Insert facts (document_id is now stored per-fact)
step_start = time.time()
@@ -469,6 +451,7 @@ async def retain_batch(
non_duplicate_facts,
log_buffer,
user_entities_per_content=user_entities_per_content,
entity_labels=getattr(config, "entity_labels", None),
)
log_buffer.append(f"[6] Process entities: {len(entity_links)} links in {time.time() - step_start:.3f}s")
@@ -499,7 +482,11 @@ async def retain_batch(
log_buffer.append(f"[10] Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s")
# Map results back to original content items
result_unit_ids = _map_results_to_contents(contents, extracted_facts, is_duplicate_flags, unit_ids)
result_unit_ids = _map_results_to_contents(contents, extracted_facts, unit_ids)
# Flush entity stats (mention_count / last_seen) now that the transaction
# has committed. Uses a fresh pool connection — no locks held.
await entity_resolver.flush_pending_stats()
# Log final summary
total_time = time.time() - start_time
@@ -517,28 +504,20 @@ async def retain_batch(
def _map_results_to_contents(
contents: list[RetainContent],
extracted_facts: list[ExtractedFact],
is_duplicate_flags: list[bool],
unit_ids: list[str],
) -> list[list[str]]:
"""
Map created unit IDs back to original content items.
Accounts for duplicates when mapping back.
"""
result_unit_ids = []
filtered_idx = 0
# Group facts by content_index
facts_by_content = {i: [] for i in range(len(contents))}
"""Map created unit IDs back to original content items."""
facts_by_content: dict[int, list[int]] = {i: [] for i in range(len(contents))}
for i, fact in enumerate(extracted_facts):
facts_by_content[fact.content_index].append(i)
result_unit_ids = []
unit_idx = 0
for content_index in range(len(contents)):
content_unit_ids = []
for fact_idx in facts_by_content[content_index]:
if not is_duplicate_flags[fact_idx]:
content_unit_ids.append(unit_ids[filtered_idx])
filtered_idx += 1
for _ in facts_by_content[content_index]:
content_unit_ids.append(unit_ids[unit_idx])
unit_idx += 1
result_unit_ids.append(content_unit_ids)
return result_unit_ids
@@ -6,8 +6,8 @@ from content input to fact storage.
"""
from dataclasses import dataclass, field
from datetime import UTC, datetime
from typing import TypedDict
from datetime import datetime
from typing import Literal, TypedDict
from uuid import UUID
@@ -22,20 +22,21 @@ class RetainContentDict(TypedDict, total=False):
document_id: Document ID for this content item (optional)
entities: User-provided entities to merge with extracted entities (optional)
tags: Visibility scope tags for this content item (optional)
observation_scopes: How to scope observations for consolidation (optional).
"per_tag" runs one pass per individual tag; "combined" (default) runs a
single pass with all tags; a list[list[str]] specifies exact passes.
"""
content: str # Required
context: str
event_date: datetime
event_date: datetime | None
metadata: dict[str, str]
document_id: str
entities: list[dict[str, str]] # [{"text": "...", "type": "..."}]
tags: list[str] # Visibility scope tags
def _now_utc() -> datetime:
"""Factory function for default event_date."""
return datetime.now(UTC)
observation_scopes: (
Literal["per_tag", "combined", "all_combinations"] | list[list[str]]
) # Observation scopes for consolidation
@dataclass
@@ -48,10 +49,13 @@ class RetainContent:
content: str
context: str = ""
event_date: datetime = field(default_factory=_now_utc)
event_date: datetime | None = None
metadata: dict[str, str] = field(default_factory=dict)
entities: list[dict[str, str]] = field(default_factory=list) # User-provided entities
tags: list[str] = field(default_factory=list) # Visibility scope tags
observation_scopes: Literal["per_tag", "combined", "all_combinations"] | list[list[str]] | None = (
None # Observation scopes
)
@dataclass
@@ -117,6 +121,9 @@ class ExtractedFact:
mentioned_at: datetime | None = None
metadata: dict[str, str] = field(default_factory=dict)
tags: list[str] = field(default_factory=list) # Visibility scope tags
observation_scopes: Literal["per_tag", "combined", "all_combinations"] | list[list[str]] | None = (
None # Observation scopes
)
@dataclass
@@ -135,7 +142,7 @@ class ProcessedFact:
# Temporal data
occurred_start: datetime | None
occurred_end: datetime | None
mentioned_at: datetime
mentioned_at: datetime | None
# Context and metadata
context: str
@@ -165,6 +172,9 @@ class ProcessedFact:
# Visibility scope tags
tags: list[str] = field(default_factory=list)
# Observation scopes for consolidation
observation_scopes: Literal["per_tag", "combined", "all_combinations"] | list[list[str]] | None = None
@property
def is_duplicate(self) -> bool:
"""Check if this fact was marked as a duplicate."""
@@ -185,12 +195,10 @@ class ProcessedFact:
Returns:
ProcessedFact ready for storage
"""
from datetime import datetime
# Use occurred dates only if explicitly provided by LLM
occurred_start = extracted_fact.occurred_start
occurred_end = extracted_fact.occurred_end
mentioned_at = extracted_fact.mentioned_at or datetime.now(UTC)
mentioned_at = extracted_fact.mentioned_at # May be None when caller opted into no timestamp
# Convert entity strings to EntityRef objects
entities = [EntityRef(name=name) for name in extracted_fact.entities]
@@ -209,6 +217,7 @@ class ProcessedFact:
chunk_id=chunk_id,
content_index=extracted_fact.content_index,
tags=extracted_fact.tags,
observation_scopes=extracted_fact.observation_scopes,
)
@@ -162,7 +162,7 @@ class BFSGraphRetriever(GraphRetriever):
entry_points = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
mentioned_at, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
@@ -216,7 +216,7 @@ class BFSGraphRetriever(GraphRetriever):
neighbors = await conn.fetch(
f"""
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.occurred_end,
mu.mentioned_at, mu.embedding, mu.fact_type,
mu.mentioned_at, mu.fact_type,
mu.document_id, mu.chunk_id, mu.tags,
ml.weight, ml.link_type, ml.from_unit_id
FROM {fq_table("memory_links")} ml
@@ -1,18 +1,28 @@
"""
Link Expansion graph retrieval.
A simple, fast graph retrieval that expands from seeds via:
1. Entity links: Find facts sharing entities with seeds (filtered by entity frequency)
2. Causal links: Find facts causally linked to seeds (top-k by weight)
Expands from semantic/temporal seeds through three parallel, first-class signals
stored in memory_links:
Characteristics:
- 2-3 DB queries (seed finding + parallel entity/causal expansion)
- Sublinear: only touches connected facts via indexes
- No iteration, no propagation, no normalization
- Target: <100ms
1. Entity links — precomputed co-occurrence graph (created at retain time, bounded to
MAX_LINKS_PER_ENTITY per entity). Score = number of distinct shared
entities between the seed set and each candidate.
2. Semantic links — precomputed kNN graph (each new fact linked to its top-5 most
similar existing facts at insert time, similarity >= 0.7). Checked
in both directions since the graph is not symmetric. Score = weight.
3. Causal links — explicit causal chains (causes/caused_by/enables/prevents).
Score = weight + 1.0 (boosted as highest-quality signal).
All three signals are bounded at retain time, so no LATERAL fan-out caps are needed
at query time. Each expansion is a simple aggregation over a small result set.
For non-observation fact types the three expansions are issued as a single CTE query
(one roundtrip, one connection) with a `source` discriminator column so the Python
merge step can apply per-signal score transformations.
"""
import logging
import math
import time
from ..db_utils import acquire_with_retry
@@ -45,7 +55,7 @@ async def _find_semantic_seeds(
rows = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
mentioned_at, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
@@ -65,27 +75,23 @@ class LinkExpansionRetriever(GraphRetriever):
"""
Graph retrieval via direct link expansion from seeds.
Expands through entity co-occurrence and causal links in a single query.
Fast and simple alternative to MPFP.
Runs three expansions through precomputed memory_links: entity co-occurrence,
semantic kNN, and causal chains, all bounded at retain time.
For non-observation fact types the three expansions are issued as a single CTE
query (one roundtrip, one connection slot) with a `source` discriminator column.
The Python merge step applies per-signal score transformations.
"""
def __init__(
self,
max_entity_frequency: int = 500,
causal_weight_threshold: float = 0.3,
causal_limit_per_seed: int = 10,
):
"""
Initialize link expansion retriever.
Args:
max_entity_frequency: Skip entities appearing in more than this many facts
causal_weight_threshold: Minimum weight for causal links
causal_limit_per_seed: Max causal links to follow per seed
causal_weight_threshold: Minimum weight for causal links to follow.
"""
self.max_entity_frequency = max_entity_frequency
self.causal_weight_threshold = causal_weight_threshold
self.causal_limit_per_seed = causal_limit_per_seed
@property
def name(self) -> str:
@@ -110,7 +116,7 @@ class LinkExpansionRetriever(GraphRetriever):
Args:
pool: Database connection pool
query_embedding_str: Query embedding (unused, kept for interface)
query_embedding_str: Query embedding as string
bank_id: Memory bank ID
fact_type: Fact type to filter
budget: Maximum results to return
@@ -118,7 +124,7 @@ class LinkExpansionRetriever(GraphRetriever):
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 (OR matching)
tags: Optional list of tags for visibility filtering
Returns:
Tuple of (results, timings)
@@ -126,8 +132,6 @@ class LinkExpansionRetriever(GraphRetriever):
start_time = time.time()
timings = MPFPTimings(fact_type=fact_type)
# Use single connection for all queries to reduce pool pressure
# (queries are fast ~50ms each, connection acquisition is the bottleneck)
async with acquire_with_retry(pool) as conn:
# Find seeds if not provided
if semantic_seeds:
@@ -150,7 +154,6 @@ class LinkExpansionRetriever(GraphRetriever):
f"(tags={tags}, tags_match={tags_match})"
)
# Add temporal seeds if provided
if temporal_seeds:
all_seeds.extend(temporal_seeds)
@@ -160,223 +163,61 @@ class LinkExpansionRetriever(GraphRetriever):
seed_ids = list({s.id for s in all_seeds})
timings.pattern_count = len(seed_ids)
# Run entity and causal expansion sequentially on same connection
query_start = time.time()
# For observations, traverse through source_memory_ids to find entity connections.
# Observations don't have direct unit_entities - they inherit entities via their
# source world/experience facts.
#
# Path: observation → source_memory_ids → world fact → entities →
# ALL world facts with those entities → their observations (excluding seeds)
if fact_type == "observation":
# Debug: Check what source_memory_ids exist on seed observations
debug_sources = await conn.fetch(
f"""
SELECT id, source_memory_ids
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
""",
seed_ids,
)
source_ids_found = []
for row in debug_sources:
if row["source_memory_ids"]:
source_ids_found.extend(row["source_memory_ids"])
logger.debug(
f"[LinkExpansion] observation graph: {len(seed_ids)} seeds, "
f"{len(source_ids_found)} source_memory_ids found"
)
entity_rows = await conn.fetch(
f"""
WITH seed_sources AS (
-- Get source memory IDs from seed observations
SELECT DISTINCT unnest(source_memory_ids) AS source_id
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
AND source_memory_ids IS NOT NULL
),
source_entities AS (
-- Get entities from those source memories (filtered by frequency)
SELECT DISTINCT ue.entity_id
FROM seed_sources ss
JOIN {fq_table("unit_entities")} ue ON ss.source_id = ue.unit_id
JOIN {fq_table("entities")} e ON ue.entity_id = e.id
WHERE e.mention_count < $2
),
all_connected_sources AS (
-- Find ALL world facts sharing those entities (don't exclude seed sources)
-- The exclusion happens at the observation level, not the source level
SELECT DISTINCT other_ue.unit_id AS source_id
FROM source_entities se
JOIN {fq_table("unit_entities")} other_ue ON se.entity_id = other_ue.entity_id
)
-- Find observations derived from connected source memories
-- Only exclude the actual seed observations
SELECT
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
mu.occurred_end, mu.mentioned_at, mu.embedding,
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
COUNT(DISTINCT cs.source_id)::float AS score
FROM all_connected_sources cs
JOIN {fq_table("memory_units")} mu
ON mu.source_memory_ids @> ARRAY[cs.source_id]
WHERE mu.fact_type = 'observation'
AND mu.id != ALL($1::uuid[])
GROUP BY mu.id
ORDER BY score DESC
LIMIT $3
""",
seed_ids,
self.max_entity_frequency,
budget,
)
logger.debug(f"[LinkExpansion] observation graph: found {len(entity_rows)} connected observations")
entity_rows, semantic_rows, causal_rows = await self._expand_observations(conn, seed_ids, budget)
else:
# For world/experience facts, use direct entity lookup
entity_rows = await conn.fetch(
f"""
SELECT
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
mu.occurred_end, mu.mentioned_at, mu.embedding,
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
COUNT(*)::float AS score
FROM {fq_table("unit_entities")} seed_ue
JOIN {fq_table("entities")} e ON seed_ue.entity_id = e.id
JOIN {fq_table("unit_entities")} other_ue ON seed_ue.entity_id = other_ue.entity_id
JOIN {fq_table("memory_units")} mu ON other_ue.unit_id = mu.id
WHERE seed_ue.unit_id = ANY($1::uuid[])
AND e.mention_count < $2
AND mu.id != ALL($1::uuid[])
AND mu.fact_type = $3
GROUP BY mu.id
ORDER BY score DESC
LIMIT $4
""",
seed_ids,
self.max_entity_frequency,
fact_type,
budget,
)
causal_rows = await conn.fetch(
f"""
SELECT DISTINCT ON (mu.id)
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
mu.occurred_end, mu.mentioned_at, mu.embedding,
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
ml.weight + 1.0 AS score
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.link_type IN ('causes', 'caused_by', 'enables', 'prevents')
AND ml.weight >= $2
AND mu.fact_type = $3
ORDER BY mu.id, ml.weight DESC
LIMIT $4
""",
seed_ids,
self.causal_weight_threshold,
fact_type,
budget,
)
# Fallback: semantic/temporal/entity links from memory_links table
# These are secondary to entity links (via unit_entities) and causal links
# Weight is halved (0.5x) to prioritize primary link types
# Check both directions: seeds -> others AND others -> seeds
fallback_rows = await conn.fetch(
f"""
WITH outgoing AS (
-- Links FROM seeds TO other facts
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
mu.occurred_end, mu.mentioned_at, mu.embedding,
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
ml.weight
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.link_type IN ('semantic', 'temporal', 'entity')
AND ml.weight >= $2
AND mu.fact_type = $3
AND mu.id != ALL($1::uuid[])
),
incoming AS (
-- Links FROM other facts TO seeds (reverse direction)
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
mu.occurred_end, mu.mentioned_at, mu.embedding,
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
ml.weight
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.from_unit_id = mu.id
WHERE ml.to_unit_id = ANY($1::uuid[])
AND ml.link_type IN ('semantic', 'temporal', 'entity')
AND ml.weight >= $2
AND mu.fact_type = $3
AND mu.id != ALL($1::uuid[])
),
combined AS (
SELECT * FROM outgoing
UNION ALL
SELECT * FROM incoming
)
SELECT DISTINCT ON (id)
id, text, context, event_date, occurred_start,
occurred_end, mentioned_at, embedding,
fact_type, document_id, chunk_id, tags,
(MAX(weight) * 0.5) AS score
FROM combined
GROUP BY id, text, context, event_date, occurred_start,
occurred_end, mentioned_at, embedding,
fact_type, document_id, chunk_id, tags
ORDER BY id, score DESC
LIMIT $4
""",
seed_ids,
self.causal_weight_threshold,
fact_type,
budget,
)
entity_rows, semantic_rows, causal_rows = await self._expand_combined(conn, seed_ids, fact_type, budget)
timings.edge_load_time = time.time() - query_start
timings.db_queries = 3
timings.edge_count = len(entity_rows) + len(causal_rows) + len(fallback_rows)
timings.db_queries = 1
timings.edge_count = len(entity_rows) + len(semantic_rows) + len(causal_rows)
# Merge results, taking max score per fact
# Priority: entity links (unit_entities) > causal links > fallback links
score_map: dict[str, float] = {}
# Merge results with additive intra-score: entity + semantic + causal ∈ [0, 3].
#
# Entity score: tanh(count × 0.5) maps shared-entity count to [0, 1]:
# 1 entity → 0.46, 2 → 0.76, 3 → 0.91, 4 → 0.96 (saturates naturally)
# Semantic score: similarity weight, already ∈ [0.7, 1.0].
# Causal score: link weight, already ∈ [0, 1].
#
# Facts appearing in multiple signals accumulate higher scores, rewarding
# convergent evidence. The outer RRF uses rank position from this sorted list.
entity_scores: dict[str, float] = {}
semantic_scores: dict[str, float] = {}
causal_scores: dict[str, float] = {}
row_map: dict[str, dict] = {}
for row in entity_rows:
fact_id = str(row["id"])
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
entity_scores[fact_id] = math.tanh(row["score"] * 0.5)
row_map[fact_id] = dict(row)
for row in semantic_rows:
fact_id = str(row["id"])
semantic_scores[fact_id] = max(semantic_scores.get(fact_id, 0.0), row["score"])
row_map.setdefault(fact_id, dict(row))
for row in causal_rows:
fact_id = str(row["id"])
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
if fact_id not in row_map:
row_map[fact_id] = dict(row)
causal_scores[fact_id] = max(causal_scores.get(fact_id, 0.0), row["score"])
row_map.setdefault(fact_id, dict(row))
for row in fallback_rows:
fact_id = str(row["id"])
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
if fact_id not in row_map:
row_map[fact_id] = dict(row)
all_ids = set(entity_scores) | set(semantic_scores) | set(causal_scores)
score_map = {
fid: entity_scores.get(fid, 0.0) + semantic_scores.get(fid, 0.0) + causal_scores.get(fid, 0.0)
for fid in all_ids
}
# Sort by score and limit
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]
# Convert to results
results = []
for row in rows:
result = RetrievalResult.from_db_row(dict(row))
result.activation = row["score"]
results.append(result)
# Apply tags filtering (graph expansion may reach untagged memories)
if tags:
results = filter_results_by_tags(results, tags, match=tags_match)
@@ -389,3 +230,253 @@ class LinkExpansionRetriever(GraphRetriever):
)
return results, timings
async def _expand_combined(
self,
conn,
seed_ids: list,
fact_type: str,
budget: int,
) -> tuple[list, list, list]:
"""
Single-roundtrip CTE query combining entity, semantic, and causal expansions.
Uses a `source` discriminator column so the caller can apply per-signal
score transformations. The three CTEs share one connection slot — important
for asyncpg which does not allow concurrent queries on the same connection.
Index coverage (requires migration d2e3f4a5b6c7):
entity: idx_memory_links_entity_covering (from_unit_id) INCLUDE (to_unit_id, entity_id)
WHERE link_type = 'entity' → index-only scan, no heap reads
semantic incoming:
idx_memory_links_to_type_weight (to_unit_id, link_type, weight DESC)
→ replaces costly BitmapAnd of two separate scans
"""
ml = fq_table("memory_links")
mu = fq_table("memory_units")
all_rows = await conn.fetch(
f"""
WITH entity_expanded AS (
-- Entity co-occurrence: seeds → their precomputed entity-link neighbors.
-- Score = distinct shared entities (bounded at retain time to
-- MAX_LINKS_PER_ENTITY=50). GROUP BY mu.id is sufficient because mu.id
-- is the primary key and functionally determines all other mu columns.
SELECT
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,
COUNT(DISTINCT ml.entity_id)::float AS score,
'entity'::text AS source
FROM {ml} ml
JOIN {mu} mu ON mu.id = ml.to_unit_id
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.link_type = 'entity'
AND mu.fact_type = $2
AND mu.id != ALL($1::uuid[])
GROUP BY mu.id
ORDER BY score DESC
LIMIT $3
),
semantic_expanded AS (
-- Semantic kNN: both outgoing (seeds → their kNN at insert time) and
-- incoming (facts inserted after seeds that found seeds as kNN).
-- Score = max similarity weight across both directions.
SELECT
id, text, context, event_date, occurred_start,
occurred_end, mentioned_at,
fact_type, document_id, chunk_id, tags,
MAX(weight) AS score,
'semantic'::text AS source
FROM (
SELECT
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,
ml.weight
FROM {ml} ml
JOIN {mu} mu ON mu.id = ml.to_unit_id
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.link_type = 'semantic'
AND mu.fact_type = $2
AND mu.id != ALL($1::uuid[])
UNION ALL
SELECT
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,
ml.weight
FROM {ml} ml
JOIN {mu} mu ON mu.id = ml.from_unit_id
WHERE ml.to_unit_id = ANY($1::uuid[])
AND ml.link_type = 'semantic'
AND mu.fact_type = $2
AND mu.id != ALL($1::uuid[])
) sem_raw
GROUP BY id, text, context, event_date, occurred_start,
occurred_end, mentioned_at,
fact_type, document_id, chunk_id, tags
ORDER BY score DESC
LIMIT $3
),
causal_expanded AS (
-- Causal chains: explicit causes/enables/prevents links from seeds.
-- DISTINCT ON handles the case where a seed has multiple causal links
-- to the same target; best weight wins.
SELECT DISTINCT ON (mu.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,
ml.weight AS score,
'causal'::text AS source
FROM {ml} ml
JOIN {mu} mu ON ml.to_unit_id = mu.id
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.link_type IN ('causes', 'caused_by', 'enables', 'prevents')
AND ml.weight >= $4
AND mu.fact_type = $2
ORDER BY mu.id, ml.weight DESC
LIMIT $3
)
SELECT * FROM entity_expanded
UNION ALL
SELECT * FROM semantic_expanded
UNION ALL
SELECT * FROM causal_expanded
""",
seed_ids,
fact_type,
budget,
self.causal_weight_threshold,
)
entity_rows = [r for r in all_rows if r["source"] == "entity"]
semantic_rows = [r for r in all_rows if r["source"] == "semantic"]
causal_rows = [r for r in all_rows if r["source"] == "causal"]
return entity_rows, semantic_rows, causal_rows
async def _expand_observations(
self,
conn,
seed_ids: list,
budget: int,
) -> tuple[list, list, list]:
"""
Observation-specific expansion.
Observations don't have direct entity links in memory_links (they're created
by consolidation, not retain). Instead, traverse source_memory_ids → world
facts → entities → other world facts → their observations.
Semantic and causal expansions run as a second combined CTE query.
"""
source_ids_found: list = []
if logger.isEnabledFor(logging.DEBUG):
debug_rows = await conn.fetch(
f"""
SELECT id, source_memory_ids
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
""",
seed_ids,
)
for row in debug_rows:
if row["source_memory_ids"]:
source_ids_found.extend(row["source_memory_ids"])
logger.debug(
f"[LinkExpansion] observation graph: {len(seed_ids)} seeds, "
f"{len(source_ids_found)} source_memory_ids found"
)
entity_rows = await conn.fetch(
f"""
WITH seed_sources AS (
SELECT DISTINCT unnest(source_memory_ids) AS source_id
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
AND source_memory_ids IS NOT NULL
),
source_entities AS (
SELECT DISTINCT ue.entity_id
FROM seed_sources ss
JOIN {fq_table("unit_entities")} ue ON ss.source_id = ue.unit_id
),
all_connected_sources AS (
SELECT DISTINCT other_ue.unit_id AS source_id
FROM source_entities se
JOIN {fq_table("unit_entities")} other_ue ON se.entity_id = other_ue.entity_id
)
SELECT
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,
COUNT(DISTINCT cs.source_id)::float AS score
FROM all_connected_sources cs
JOIN {fq_table("memory_units")} mu
ON mu.source_memory_ids @> ARRAY[cs.source_id]
WHERE mu.fact_type = 'observation'
AND mu.id != ALL($1::uuid[])
GROUP BY mu.id
ORDER BY score DESC
LIMIT $2
""",
seed_ids,
budget,
)
logger.debug(f"[LinkExpansion] observation graph: found {len(entity_rows)} connected observations")
# Semantic + causal for observations in one query
ml = fq_table("memory_links")
mu = fq_table("memory_units")
sem_causal_rows = await conn.fetch(
f"""
WITH semantic_expanded AS (
SELECT
id, text, context, event_date, occurred_start,
occurred_end, mentioned_at,
fact_type, document_id, chunk_id, tags,
MAX(weight) AS score,
'semantic'::text AS source
FROM (
SELECT 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, ml.weight
FROM {ml} ml JOIN {mu} mu ON mu.id = ml.to_unit_id
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.link_type = 'semantic' AND mu.fact_type = 'observation'
AND mu.id != ALL($1::uuid[])
UNION ALL
SELECT 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, ml.weight
FROM {ml} ml JOIN {mu} mu ON mu.id = ml.from_unit_id
WHERE ml.to_unit_id = ANY($1::uuid[])
AND ml.link_type = 'semantic' AND mu.fact_type = 'observation'
AND mu.id != ALL($1::uuid[])
) sem_raw
GROUP BY id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, fact_type, document_id, chunk_id, tags
ORDER BY score DESC LIMIT $2
),
causal_expanded AS (
SELECT DISTINCT ON (mu.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, ml.weight AS score, 'causal'::text AS source
FROM {ml} ml JOIN {mu} mu ON ml.to_unit_id = mu.id
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.link_type IN ('causes', 'caused_by', 'enables', 'prevents')
AND ml.weight >= $3 AND mu.fact_type = 'observation'
ORDER BY mu.id, ml.weight DESC LIMIT $2
)
SELECT * FROM semantic_expanded
UNION ALL
SELECT * FROM causal_expanded
""",
seed_ids,
budget,
self.causal_weight_threshold,
)
semantic_rows = [r for r in sem_causal_rows if r["source"] == "semantic"]
causal_rows = [r for r in sem_causal_rows if r["source"] == "causal"]
return entity_rows, semantic_rows, causal_rows
@@ -449,7 +449,7 @@ async def fetch_memory_units_by_ids(
rows = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, embedding, fact_type, document_id, chunk_id, tags
mentioned_at, fact_type, document_id, chunk_id, tags
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
AND fact_type = $2
@@ -127,7 +127,7 @@ async def retrieve_semantic_bm25_combined(
results = await conn.fetch(
f"""
WITH semantic_ranked AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity,
NULL::float AS bm25_score,
'semantic' AS source,
@@ -139,7 +139,7 @@ async def retrieve_semantic_bm25_combined(
AND (1 - (embedding <=> $1::vector)) >= 0.3
{tags_clause}
)
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags,
similarity, bm25_score, source
FROM semantic_ranked
WHERE rn <= $4
@@ -194,7 +194,7 @@ async def retrieve_semantic_bm25_combined(
# Single query template with backend-specific parts injected
query = f"""
WITH semantic_ranked AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity,
NULL::float AS bm25_score,
'semantic' AS source,
@@ -207,7 +207,7 @@ async def retrieve_semantic_bm25_combined(
{tags_clause}
),
bm25_ranked AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags,
NULL::float AS similarity,
{bm25_score_expr} AS bm25_score,
'bm25' AS source,
@@ -219,12 +219,12 @@ async def retrieve_semantic_bm25_combined(
{tags_clause}
),
semantic AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags,
similarity, bm25_score, source
FROM semantic_ranked WHERE rn <= $4
),
bm25 AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags,
similarity, bm25_score, source
FROM bm25_ranked WHERE rn <= $4
)
@@ -297,13 +297,20 @@ async def retrieve_temporal_combined(
if tags:
params.append(tags)
# Batch query: Get entry points for ALL fact types at once with window function
# Two-phase entry point query:
# Phase 1 (date_ranked): rank by date only — no embedding computation — for all units in
# the temporal window. This lets the planner use date indexes for filtering.
# Phase 2 (sim_ranked): join back to memory_units for only the top-50-per-type candidates
# and compute embedding similarity for that small set (≤ 50 × len(fact_types) rows).
# This avoids computing embedding distances for potentially thousands of date-range rows.
entry_points = await conn.fetch(
f"""
WITH ranked_entries AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity,
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, embedding <=> $1::vector) AS rn
WITH date_ranked AS MATERIALIZED (
SELECT id, fact_type,
ROW_NUMBER() OVER (
PARTITION BY fact_type
ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC NULLS LAST
) AS rn
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND fact_type = ANY($3)
@@ -318,12 +325,20 @@ async def retrieve_temporal_combined(
OR
(occurred_end IS NOT NULL AND occurred_end BETWEEN $4 AND $5)
)
AND (1 - (embedding <=> $1::vector)) >= $6
{tags_clause}
),
sim_ranked AS (
SELECT 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,
1 - (mu.embedding <=> $1::vector) AS similarity,
ROW_NUMBER() OVER (PARTITION BY mu.fact_type ORDER BY mu.embedding <=> $1::vector) AS sim_rn
FROM date_ranked dr
JOIN {fq_table("memory_units")} mu ON mu.id = dr.id
WHERE dr.rn <= 50
AND (1 - (mu.embedding <=> $1::vector)) >= $6
)
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags, similarity
FROM ranked_entries
WHERE rn <= 10
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, fact_type, document_id, chunk_id, tags, similarity
FROM sim_ranked
WHERE sim_rn <= 10
""",
*params,
)
@@ -387,34 +402,52 @@ async def retrieve_temporal_combined(
frontier = list(node_scores.keys())
budget_remaining = budget - len(ft_entry_points)
batch_size = 20
# Per-source neighbor limit: lets the planner use the composite index
# (from_unit_id, link_type, weight DESC) with early termination, avoiding
# a full scan of all links from all source nodes before sorting.
per_source_limit = 10
# Safety cap on BFS iterations to prevent runaway spreading in dense graphs.
max_iterations = 5
iteration = 0
# Build tags clause for spreading (use param 6 since 1-5 are used)
spreading_tags_clause = build_tags_where_clause_simple(tags, 6, table_alias="mu.", match=tags_match)
# Build tags clause for spreading (use param 7 since 1-6 are used)
spreading_tags_clause = build_tags_where_clause_simple(tags, 7, table_alias="mu.", match=tags_match)
while frontier and budget_remaining > 0:
while frontier and budget_remaining > 0 and iteration < max_iterations:
iteration += 1
batch_ids = frontier[:batch_size]
frontier = frontier[batch_size:]
spreading_params = [query_emb_str, batch_ids, ft, semantic_threshold, batch_size * 10]
# $1=query_emb, $2=batch_ids, $3=fact_type, $4=threshold, $5=per_source_limit, $6=bank_id, $7=tags
spreading_params = [query_emb_str, batch_ids, ft, semantic_threshold, per_source_limit, bank_id]
if tags:
spreading_params.append(tags)
# LATERAL join: for each source node, fetch top-K neighbors by weight using
# the existing idx_memory_links_from_type_weight index with early-exit semantics.
# This avoids scanning all temporal links from all source nodes before sorting.
# bank_id on memory_units lets the planner use idx_memory_units_bank_fact_type.
neighbors = await conn.fetch(
f"""
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
ml.weight, ml.link_type, ml.from_unit_id,
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,
l.weight, l.link_type,
1 - (mu.embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
WHERE ml.from_unit_id = ANY($2::uuid[])
AND ml.link_type IN ('temporal', 'causes', 'caused_by', 'enables', 'prevents')
AND ml.weight >= 0.1
FROM unnest($2::uuid[]) AS src(from_unit_id)
CROSS JOIN LATERAL (
SELECT ml.to_unit_id, ml.weight, ml.link_type
FROM {fq_table("memory_links")} ml
WHERE ml.from_unit_id = src.from_unit_id
AND ml.link_type IN ('temporal', 'causes', 'caused_by', 'enables', 'prevents')
AND ml.weight >= 0.1
ORDER BY ml.weight DESC
LIMIT $5
) l
JOIN {fq_table("memory_units")} mu ON mu.id = l.to_unit_id
WHERE mu.bank_id = $6
AND mu.fact_type = $3
AND mu.embedding IS NOT NULL
AND (1 - (mu.embedding <=> $1::vector)) >= $4
{spreading_tags_clause}
ORDER BY ml.weight DESC
LIMIT $5
""",
*spreading_params,
)
@@ -46,7 +46,6 @@ class RetrievalResult:
mentioned_at: datetime | None = None
document_id: str | None = None
chunk_id: str | None = None
embedding: list[float] | None = None
tags: list[str] | None = None # Visibility scope tags
# Retrieval-specific scores (only one will be set depending on retrieval method)
@@ -70,7 +69,6 @@ class RetrievalResult:
mentioned_at=row.get("mentioned_at"),
document_id=row.get("document_id"),
chunk_id=row.get("chunk_id"),
embedding=row.get("embedding"),
tags=row.get("tags"),
similarity=row.get("similarity"),
bm25_score=row.get("bm25_score"),
@@ -154,7 +152,6 @@ class ScoredResult:
"mentioned_at": self.retrieval.mentioned_at,
"document_id": self.retrieval.document_id,
"chunk_id": self.retrieval.chunk_id,
"embedding": self.retrieval.embedding,
"tags": self.retrieval.tags,
"semantic_similarity": self.retrieval.similarity,
"bm25_score": self.retrieval.bm25_score,
@@ -12,6 +12,7 @@ def create_file_storage(
storage_type: str,
pool_getter: Callable | None = None,
schema: str | None = None,
schema_getter: Callable | None = None,
**kwargs,
) -> FileStorage:
"""
@@ -20,7 +21,8 @@ def create_file_storage(
Args:
storage_type: "native" (PostgreSQL BYTEA) or "s3" (S3-compatible object storage)
pool_getter: Database pool getter (required for native)
schema: Database schema (for native multi-tenant)
schema: Static database schema (for native single-tenant)
schema_getter: Callable returning current schema at query time (for native multi-tenant)
**kwargs: Additional args passed to storage backend
Returns:
@@ -32,7 +34,7 @@ def create_file_storage(
if storage_type == "native":
if not pool_getter:
raise ValueError("pool_getter required for native (PostgreSQL) storage")
return PostgreSQLFileStorage(pool_getter=pool_getter, schema=schema)
return PostgreSQLFileStorage(pool_getter=pool_getter, schema=schema, schema_getter=schema_getter)
elif storage_type == "s3":
from ...config import get_config
from .s3 import S3FileStorage
@@ -40,16 +40,30 @@ class PostgreSQLFileStorage(FileStorage):
For production/scale, consider S3FileStorage instead.
"""
def __init__(self, pool_getter: Callable[[], "asyncpg.Pool"], schema: str | None = None):
def __init__(
self,
pool_getter: Callable[[], "asyncpg.Pool"],
schema: str | None = None,
schema_getter: Callable[[], str] | None = None,
):
"""
Initialize PostgreSQL file storage.
Args:
pool_getter: Function that returns asyncpg connection pool
schema: Database schema (for multi-tenant support)
schema: Static database schema (fallback for single-tenant / tests)
schema_getter: Callable returning current schema at query time (for multi-tenant)
"""
self._pool_getter = pool_getter
self._schema = schema
self._static_schema = schema
self._schema_getter = schema_getter
@property
def _schema(self) -> str | None:
"""Resolve schema dynamically per-request when schema_getter is provided."""
if self._schema_getter:
return self._schema_getter()
return self._static_schema
async def store(
self,
@@ -22,6 +22,11 @@ from hindsight_api.extensions.http import HttpExtension
from hindsight_api.extensions.loader import load_extension
from hindsight_api.extensions.mcp import MCPExtension
from hindsight_api.extensions.operation_validator import (
# Bank Management operations
BankListContext,
BankListResult,
BankReadContext,
BankWriteContext,
# Consolidation operation
ConsolidateContext,
ConsolidateResult,
@@ -70,6 +75,11 @@ __all__ = [
"RetainContext",
"RetainResult",
"ValidationResult",
# Operation Validator - Bank Management
"BankListContext",
"BankListResult",
"BankReadContext",
"BankWriteContext",
# Operation Validator - Consolidation
"ConsolidateContext",
"ConsolidateResult",
@@ -87,3 +87,15 @@ class HttpExtension(Extension, ABC):
```
"""
pass
def get_root_router(self, memory: "MemoryEngine") -> APIRouter | None:
"""
Return a FastAPI router with endpoints mounted at the app root.
Unlike get_router() which is mounted at /ext/, this router is mounted
directly on the application root. Use for well-known endpoints or other
paths that must be at specific locations.
Returns None by default (no root routes). Override to provide root-level routes.
"""
return None
@@ -200,6 +200,44 @@ class ConsolidateResult:
error: str | None = None
# =============================================================================
# Bank Management Contexts
# =============================================================================
@dataclass
class BankReadContext:
"""Context for a bank read operation validation (pre-operation)."""
bank_id: str
operation: str # "get_bank_profile", "get_bank_stats"
request_context: "RequestContext"
@dataclass
class BankWriteContext:
"""Context for a bank write operation validation (pre-operation)."""
bank_id: str
operation: str # "delete_bank", "update_bank", "update_bank_disposition", "set_bank_mission", "merge_bank_mission", "clear_observations", "clear_observations_for_memory"
request_context: "RequestContext"
@dataclass
class BankListContext:
"""Context for filtering the bank list (post-query)."""
banks: list[dict]
request_context: "RequestContext"
@dataclass
class BankListResult:
"""Result of filtering the bank list."""
banks: list[dict]
# =============================================================================
# Mental Model Contexts
# =============================================================================
@@ -535,3 +573,63 @@ class OperationValidatorExtension(Extension, ABC):
- error: Error message (if failed)
"""
pass
# =========================================================================
# Bank Management - Validation hooks (optional - override to implement)
# =========================================================================
async def validate_bank_read(self, ctx: BankReadContext) -> ValidationResult:
"""
Validate a bank read operation before execution.
Override to implement custom validation logic for bank reads
(get_bank_profile, get_bank_stats).
Args:
ctx: Context containing:
- bank_id: Bank identifier
- operation: Operation name
- request_context: Request context with auth info
Returns:
ValidationResult indicating whether the operation is allowed.
"""
return ValidationResult.accept()
async def validate_bank_write(self, ctx: BankWriteContext) -> ValidationResult:
"""
Validate a bank write operation before execution.
Override to implement custom validation logic for bank writes
(delete_bank, update_bank, update_bank_disposition, set_bank_mission,
merge_bank_mission, clear_observations, clear_observations_for_memory).
Args:
ctx: Context containing:
- bank_id: Bank identifier
- operation: Operation name
- request_context: Request context with auth info
Returns:
ValidationResult indicating whether the operation is allowed.
"""
return ValidationResult.accept()
async def filter_bank_list(self, ctx: BankListContext) -> BankListResult:
"""
Filter the bank list after querying.
Unlike validate_* methods, this is a post-query filter that narrows results
rather than a gate that blocks the operation.
Override to implement custom filtering (e.g., restrict to allowed banks).
Args:
ctx: Context containing:
- banks: List of bank dicts from the database
- request_context: Request context with auth info
Returns:
BankListResult with the filtered list of banks.
"""
return BankListResult(banks=ctx.banks)
@@ -11,8 +11,9 @@ from hindsight_api.models import RequestContext
class AuthenticationError(Exception):
"""Raised when authentication fails."""
def __init__(self, reason: str):
def __init__(self, reason: str, headers: dict[str, str] | None = None):
self.reason = reason
self.headers = headers or {}
super().__init__(f"Authentication failed: {reason}")
+15
View File
@@ -171,6 +171,7 @@ def main():
llm_vertexai_project_id=config.llm_vertexai_project_id,
llm_vertexai_region=config.llm_vertexai_region,
llm_vertexai_service_account_key=config.llm_vertexai_service_account_key,
llm_gemini_safety_settings=config.llm_gemini_safety_settings,
retain_llm_provider=config.retain_llm_provider,
retain_llm_api_key=config.retain_llm_api_key,
retain_llm_model=config.retain_llm_model,
@@ -231,12 +232,15 @@ def main():
reranker_litellm_sdk_api_key=config.reranker_litellm_sdk_api_key,
reranker_litellm_sdk_model=config.reranker_litellm_sdk_model,
reranker_litellm_sdk_api_base=config.reranker_litellm_sdk_api_base,
reranker_zeroentropy_api_key=config.reranker_zeroentropy_api_key,
reranker_zeroentropy_model=config.reranker_zeroentropy_model,
host=args.host,
port=args.port,
base_path=config.base_path,
log_level=args.log_level,
log_format=config.log_format,
mcp_enabled=config.mcp_enabled,
mcp_enabled_tools=config.mcp_enabled_tools,
enable_bank_config_api=config.enable_bank_config_api,
graph_retriever=config.graph_retriever,
mpfp_top_k_neighbors=config.mpfp_top_k_neighbors,
@@ -246,8 +250,10 @@ def main():
retain_chunk_size=config.retain_chunk_size,
retain_extract_causal_links=config.retain_extract_causal_links,
retain_extraction_mode=config.retain_extraction_mode,
retain_mission=config.retain_mission,
retain_custom_instructions=config.retain_custom_instructions,
retain_batch_tokens=config.retain_batch_tokens,
retain_entity_lookup=config.retain_entity_lookup,
retain_batch_enabled=config.retain_batch_enabled,
retain_batch_poll_interval_seconds=config.retain_batch_poll_interval_seconds,
file_storage_type=config.file_storage_type,
@@ -270,7 +276,11 @@ def main():
file_delete_after_retain=config.file_delete_after_retain,
enable_observations=config.enable_observations,
consolidation_batch_size=config.consolidation_batch_size,
consolidation_llm_batch_size=config.consolidation_llm_batch_size,
consolidation_max_tokens=config.consolidation_max_tokens,
observations_mission=config.observations_mission,
entity_labels=config.entity_labels,
entities_allow_free_form=config.entities_allow_free_form,
skip_llm_verification=config.skip_llm_verification,
lazy_reranker=config.lazy_reranker,
run_migrations_on_startup=config.run_migrations_on_startup,
@@ -286,6 +296,11 @@ def main():
worker_max_slots=config.worker_max_slots,
worker_consolidation_max_slots=config.worker_consolidation_max_slots,
reflect_max_iterations=config.reflect_max_iterations,
reflect_max_context_tokens=config.reflect_max_context_tokens,
reflect_mission=config.reflect_mission,
disposition_skepticism=config.disposition_skepticism,
disposition_literalism=config.disposition_literalism,
disposition_empathy=config.disposition_empathy,
mental_model_refresh_concurrency=config.mental_model_refresh_concurrency,
otel_traces_enabled=config.otel_traces_enabled,
otel_exporter_otlp_endpoint=config.otel_exporter_otlp_endpoint,
File diff suppressed because it is too large Load Diff
+63 -4
View File
@@ -18,6 +18,7 @@ No alembic.ini required - all configuration is done programmatically.
import hashlib
import logging
import os
import time
from pathlib import Path
from alembic import command
@@ -220,13 +221,40 @@ def run_migrations(
lock_id = _get_schema_lock_id(schema) if schema else MIGRATION_LOCK_ID
schema_name = schema or "public"
# Use PostgreSQL advisory lock to coordinate between distributed workers
# Use PostgreSQL advisory lock to coordinate between distributed workers.
#
# IMPORTANT: We must avoid holding an open transaction on the advisory-lock
# connection while CREATE INDEX CONCURRENTLY runs inside a migration.
# CONCURRENTLY waits for ALL active transactions to finish before the index
# becomes valid. If the advisory-lock connection (or any waiting worker's
# connection) holds an open transaction, CONCURRENTLY deadlocks:
# - migration worker waits for other workers' transactions to close
# - other workers wait for the advisory lock to be released
#
# Fix:
# 1. Use pg_try_advisory_lock (non-blocking) in a poll loop instead of
# blocking pg_advisory_lock, so we can COMMIT the transaction between
# retries. Between retries the connection holds no open transaction.
# 2. After acquiring the lock, COMMIT the transaction on the advisory-lock
# connection itself before running migrations. pg_advisory_lock is
# session-level, so the lock survives the COMMIT.
engine = create_engine(database_url)
with engine.connect() as conn:
# pg_advisory_lock blocks until the lock is acquired
# The lock is automatically released when the connection closes
logger.debug(f"Acquiring migration advisory lock for schema '{schema_name}' (id={lock_id})...")
conn.execute(text(f"SELECT pg_advisory_lock({lock_id})"))
while True:
acquired = conn.execute(text(f"SELECT pg_try_advisory_lock({lock_id})")).scalar()
if acquired:
break
# Commit the transaction so this connection holds no open snapshot
# while waiting. This prevents blocking CREATE INDEX CONCURRENTLY
# that may be running in the migration worker.
conn.commit()
time.sleep(0.5)
# Commit AFTER acquiring the lock too. pg_advisory_lock is session-level
# and survives the COMMIT, but the open transaction on this connection
# would otherwise block any CREATE INDEX CONCURRENTLY in the migration.
conn.commit()
logger.debug("Migration advisory lock acquired")
try:
@@ -347,6 +375,13 @@ def run_migrations(
"Please install it with: CREATE EXTENSION vectorscale CASCADE;"
) from e
# Commit any pending transaction on the advisory-lock connection
# before running migrations. Some code paths above (e.g., the
# pgvector extension check) may have started a transaction via
# SQLAlchemy's autobegin. If we leave it open, CREATE INDEX
# CONCURRENTLY inside a migration will deadlock waiting for it.
conn.commit()
# Run migrations while holding the lock
_run_migrations_internal(database_url, script_location, schema=schema)
finally:
@@ -565,6 +600,12 @@ def ensure_embedding_dimension(
)
logger.info(f"Created vchordrq index for {required_dimension}-dimensional embeddings")
else: # pgvector
if required_dimension > 2000:
raise RuntimeError(
f"Embedding dimension {required_dimension} exceeds pgvector HNSW index limit of 2000. "
f"Use an embedding model with <= 2000 dimensions, or switch to a vector extension "
f"that supports higher dimensions (e.g., pgvectorscale/DiskANN)."
)
conn.execute(
text(f"""
CREATE INDEX IF NOT EXISTS idx_memory_units_embedding_hnsw
@@ -750,6 +791,24 @@ def ensure_vector_extension(
""")
)
else: # pgvector
# Check embedding dimension — pgvector HNSW indexes only support up to 2000 dims
embed_dim = conn.execute(
text("""
SELECT atttypmod
FROM pg_attribute a
JOIN pg_class c ON a.attrelid = c.oid
JOIN pg_namespace n ON c.relnamespace = n.oid
WHERE n.nspname = :schema AND c.relname = :table_name AND a.attname = 'embedding'
"""),
{"schema": schema_name, "table_name": table_name},
).scalar()
if embed_dim and embed_dim > 2000:
raise RuntimeError(
f"Embedding dimension {embed_dim} on {table_name} exceeds pgvector HNSW index limit of 2000. "
f"Use an embedding model with <= 2000 dimensions, or switch to a vector extension "
f"that supports higher dimensions (e.g., pgvectorscale/DiskANN)."
)
logger.info(f"Creating HNSW index on {table_name}")
conn.execute(
text(f"""
+1
View File
@@ -22,6 +22,7 @@ class RequestContext:
tenant_id: str | None = None # Tenant identifier (set by extension after auth)
internal: bool = False # True for background/internal operations (skips extension auth)
user_initiated: bool = False # True for async operations that originated from a user request
allowed_bank_ids: list[str] | None = None # None = unrestricted (all banks)
from pgvector.sqlalchemy import Vector
+8 -7
View File
@@ -376,11 +376,13 @@ class WorkerPoller:
del self._in_flight_by_type[operation_type]
async def _execute_task_inner(self, task: ClaimedTask):
"""Inner task execution with error handling.
"""Inner task execution with retry/fail handling.
Note: The executor (MemoryEngine.execute_task) handles status marking internally
(marking operations as completed/failed and handling retries). This method should
NOT override those status updates.
Retryable task failures are re-raised by the executor (MemoryEngine.execute_task)
and handled here via _retry_or_fail, which resets status='pending' (or marks as
'failed' after max retries). Non-retryable failures (e.g., file_convert_retain) are
handled by the executor internally it marks the operation as failed and returns
normally, so no exception reaches here.
"""
task_type = task.task_dict.get("type", "unknown")
bank_id = task.task_dict.get("bank_id", "unknown")
@@ -393,10 +395,9 @@ class WorkerPoller:
await self._executor(task.task_dict)
logger.debug(f"Task {task.operation_id} execution finished")
except Exception as e:
# The executor should handle its own errors, but if an unexpected exception
# propagates (e.g., from schema setup), log it as a warning
logger.error(f"Task {task.operation_id} raised unexpected exception: {e}")
logger.error(f"Task {task.operation_id} failed: {e}")
traceback.print_exc()
await self._retry_or_fail(task.operation_id, str(e), task.schema)
async def recover_own_tasks(self) -> int:
"""
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "hindsight-api"
version = "0.4.13"
version = "0.4.15"
description = "Hindsight: Agent Memory That Works Like Human Memory"
readme = "README.md"
requires-python = ">=3.11"
+9 -9
View File
@@ -27,9 +27,9 @@ class TestAgentProfile:
assert "disposition" in profile
disposition = profile["disposition"]
assert disposition.skepticism == 3
assert disposition.literalism == 3
assert disposition.empathy == 3
assert disposition["skepticism"] == 3
assert disposition["literalism"] == 3
assert disposition["empathy"] == 3
@pytest.mark.asyncio
async def test_update_agent_disposition(self, memory: MemoryEngine, request_context):
@@ -37,7 +37,7 @@ class TestAgentProfile:
bank_id = unique_agent_id("test_profile_update")
profile = await memory.get_bank_profile(bank_id, request_context=request_context)
assert profile["disposition"].skepticism == 3
assert profile["disposition"]["skepticism"] == 3
new_disposition = {
"skepticism": 5,
@@ -48,9 +48,9 @@ class TestAgentProfile:
updated_profile = await memory.get_bank_profile(bank_id, request_context=request_context)
disposition = updated_profile["disposition"]
assert disposition.skepticism == new_disposition["skepticism"]
assert disposition.literalism == new_disposition["literalism"]
assert disposition.empathy == new_disposition["empathy"]
assert disposition["skepticism"] == new_disposition["skepticism"]
assert disposition["literalism"] == new_disposition["literalism"]
assert disposition["empathy"] == new_disposition["empathy"]
@pytest.mark.asyncio
async def test_list_agents(self, memory: MemoryEngine, request_context):
@@ -104,8 +104,8 @@ class TestAgentEndpoint:
final_profile = await memory.get_bank_profile(bank_id, request_context=request_context)
assert final_profile["disposition"].skepticism == 4
assert final_profile["disposition"].literalism == 5
assert final_profile["disposition"]["skepticism"] == 4
assert final_profile["disposition"]["literalism"] == 5
class TestAgentDispositionIntegration:
@@ -14,6 +14,7 @@ async def test_submit_async_retain_includes_document_tags_in_task_payload():
engine = MemoryEngine.__new__(MemoryEngine)
engine._initialized = True
engine._authenticate_tenant = AsyncMock()
engine._operation_validator = None
engine._submit_async_operation = AsyncMock(return_value={"operation_id": "op-1"})
# Mock the pool and connection for parent operation creation
+337 -10
View File
@@ -1435,22 +1435,20 @@ class TestObservationDrillDown:
assert result["count"] > 0, "Expected at least one observation"
# Verify source_memory_ids and proof_count are present
# Verify source_fact_ids is present (MemoryFact field name for source memories)
obs = result["observations"][0]
assert "source_memory_ids" in obs, "Observation should have source_memory_ids"
assert "proof_count" in obs, "Observation should have proof_count"
assert obs["proof_count"] >= 1, "proof_count should be at least 1"
assert "source_fact_ids" in obs, "Observation should have source_fact_ids"
# If source_memory_ids exist, verify they can be used with expand
if obs["source_memory_ids"]:
assert len(obs["source_memory_ids"]) >= 1, "Should have at least one source memory"
# If source_fact_ids exist, verify they can be used with expand
if obs["source_fact_ids"]:
assert len(obs["source_fact_ids"]) >= 1, "Should have at least one source memory"
# Use expand tool to get source memory details
async with memory._pool.acquire() as conn:
expand_result = await tool_expand(
conn=conn,
bank_id=bank_id,
memory_ids=obs["source_memory_ids"][:2], # Take first 2
memory_ids=obs["source_fact_ids"][:2], # Take first 2
depth="chunk",
)
@@ -1717,11 +1715,10 @@ class TestHierarchicalRetrieval:
query="What was the quarterly revenue?",
request_context=request_context,
max_tokens=2048,
max_results=10,
)
# Should have raw facts with specific numbers
assert recall_result["count"] >= 1, "Recall should find the raw facts"
assert len(recall_result["memories"]) >= 1, "Recall should find the raw facts"
# Check that we get the actual numbers from the original memories
all_memory_text = " ".join([m["text"] for m in recall_result["memories"]])
@@ -1990,3 +1987,333 @@ class TestMentalModelRefreshAfterConsolidation:
# Cleanup
await memory.delete_bank(bank_id, request_context=request_context)
def test_consolidation_prompt_default():
"""Test that the default consolidation prompt contains the built-in mission and processing rules."""
from hindsight_api.engine.consolidation.prompts import build_batch_consolidation_prompt
prompt = build_batch_consolidation_prompt()
assert "temporal markers" in prompt
assert "RESOLVE REFERENCES" in prompt
assert "{facts_text}" in prompt
assert "{observations_text}" in prompt
def test_consolidation_prompt_observations_mission():
"""Test that observations_mission replaces the default mission but keeps processing rules."""
from hindsight_api.engine.consolidation.prompts import build_batch_consolidation_prompt
spec = "Observations are weekly summaries of sprint outcomes and team dynamics."
prompt = build_batch_consolidation_prompt(observations_mission=spec)
# Spec is injected
assert spec in prompt
# Processing rules and output format always remain
assert "RESOLVE REFERENCES" in prompt
assert "creates" in prompt
assert "updates" in prompt
assert "{facts_text}" in prompt
assert "{observations_text}" in prompt
# Renders cleanly
rendered = prompt.format(facts_text="Alice fixed a bug.", observations_text="[]")
assert "{facts_text}" not in rendered
assert spec in rendered
def test_observations_mission_config():
"""Test that observations_mission is loaded from env and exposed as configurable."""
import os
from hindsight_api.config import HindsightConfig, _get_raw_config, clear_config_cache
original = os.getenv("HINDSIGHT_API_OBSERVATIONS_MISSION")
try:
os.environ["HINDSIGHT_API_OBSERVATIONS_MISSION"] = "Weekly sprint summaries only."
clear_config_cache()
config = _get_raw_config()
assert config.observations_mission == "Weekly sprint summaries only."
assert "observations_mission" in HindsightConfig.get_configurable_fields()
finally:
if original is None:
os.environ.pop("HINDSIGHT_API_OBSERVATIONS_MISSION", None)
else:
os.environ["HINDSIGHT_API_OBSERVATIONS_MISSION"] = original
clear_config_cache()
@pytest.mark.asyncio
async def test_consolidation_with_observations_mission(memory: "MemoryEngine", request_context):
"""Test that observations_mission is used during consolidation without errors."""
import os
from hindsight_api.config import _get_raw_config, clear_config_cache
original = os.getenv("HINDSIGHT_API_OBSERVATIONS_MISSION")
try:
os.environ["HINDSIGHT_API_OBSERVATIONS_MISSION"] = (
"Observations are summaries of programming language usage patterns."
)
clear_config_cache()
config = _get_raw_config()
bank_id = f"test-obs-spec-{uuid.uuid4().hex[:8]}"
original_global_config = memory._config_resolver._global_config
memory._config_resolver._global_config = config
try:
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
await memory.retain_async(
bank_id=bank_id,
content="Alice uses Python for data analysis and loves its simplicity.",
request_context=request_context,
)
async with memory._pool.acquire() as conn:
observations = await conn.fetch(
"SELECT id, text, fact_type FROM memory_units WHERE bank_id = $1 AND fact_type = 'observation'",
bank_id,
)
assert isinstance(observations, list)
finally:
memory._config_resolver._global_config = original_global_config
await memory.delete_bank(bank_id, request_context=request_context)
finally:
if original is None:
os.environ.pop("HINDSIGHT_API_OBSERVATIONS_MISSION", None)
else:
os.environ["HINDSIGHT_API_OBSERVATIONS_MISSION"] = original
clear_config_cache()
@pytest.mark.asyncio
async def test_observation_scopes_explicit_multi_pass(memory: MemoryEngine, request_context):
"""Test that observation_scopes with an explicit list triggers separate consolidation passes.
A single memory stored with observation_scopes=[["user:alice"], ["teacher:ben"]]
must produce:
- At least one observation with tags containing ONLY "user:alice" (not "teacher:ben")
- At least one observation with tags containing ONLY "teacher:ben" (not "user:alice")
The two tag scopes must remain isolated no observation should carry both tags,
which would indicate the scopes were incorrectly merged.
"""
bank_id = f"test-obs-scopes-explicit-{uuid.uuid4().hex[:8]}"
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
# Retain a memory with two explicit observation scopes
await memory.retain_batch_async(
bank_id=bank_id,
contents=[
{
"content": "Alice, a student, worked hard in the lesson with teacher Ben.",
"observation_scopes": [["user:alice"], ["teacher:ben"]],
}
],
request_context=request_context,
)
async with memory._pool.acquire() as conn:
observations = await conn.fetch(
"""
SELECT id, text, tags
FROM memory_units
WHERE bank_id = $1 AND fact_type = 'observation'
ORDER BY created_at
""",
bank_id,
)
try:
# Must have at least 2 observations (one per tag scope)
assert len(observations) >= 2, (
f"Expected at least 2 observations (one per tag scope), got {len(observations)}: "
+ str([dict(o) for o in observations])
)
tag_sets = [set(obs["tags"] or []) for obs in observations]
# There must be at least one observation scoped to user:alice only
alice_only = [ts for ts in tag_sets if "user:alice" in ts and "teacher:ben" not in ts]
assert alice_only, (
f"Expected an observation scoped to 'user:alice' only, got tag sets: {tag_sets}"
)
# There must be at least one observation scoped to teacher:ben only
ben_only = [ts for ts in tag_sets if "teacher:ben" in ts and "user:alice" not in ts]
assert ben_only, (
f"Expected an observation scoped to 'teacher:ben' only, got tag sets: {tag_sets}"
)
# No observation should carry both tags (scopes must not be merged)
both = [ts for ts in tag_sets if "user:alice" in ts and "teacher:ben" in ts]
assert not both, (
f"Found observation(s) with both tags — scopes were incorrectly merged: {both}"
)
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_observation_scopes_per_tag(memory: MemoryEngine, request_context):
"""Test that observation_scopes='per_tag' derives one pass per individual tag.
A memory with tags=["user:alice", "teacher:ben"] and observation_scopes="per_tag"
must produce isolated observations one scoped to "user:alice" and one to "teacher:ben".
"""
bank_id = f"test-obs-scopes-pertag-{uuid.uuid4().hex[:8]}"
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
await memory.retain_batch_async(
bank_id=bank_id,
contents=[
{
"content": "Alice, a student, worked hard in the lesson with teacher Ben.",
"tags": ["user:alice", "teacher:ben"],
"observation_scopes": "per_tag",
}
],
request_context=request_context,
)
async with memory._pool.acquire() as conn:
observations = await conn.fetch(
"""
SELECT id, text, tags
FROM memory_units
WHERE bank_id = $1 AND fact_type = 'observation'
ORDER BY created_at
""",
bank_id,
)
try:
assert len(observations) >= 2, (
f"Expected at least 2 observations (one per tag), got {len(observations)}: "
+ str([dict(o) for o in observations])
)
tag_sets = [set(obs["tags"] or []) for obs in observations]
alice_only = [ts for ts in tag_sets if "user:alice" in ts and "teacher:ben" not in ts]
assert alice_only, f"Expected an observation scoped to 'user:alice' only, got: {tag_sets}"
ben_only = [ts for ts in tag_sets if "teacher:ben" in ts and "user:alice" not in ts]
assert ben_only, f"Expected an observation scoped to 'teacher:ben' only, got: {tag_sets}"
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_observation_scopes_combined(memory: MemoryEngine, request_context):
"""Test that observation_scopes='combined' produces a single observation with all tags.
A memory with tags=["user:alice", "teacher:ben"] and observation_scopes="combined"
must produce at least one observation that carries both tags together, and no
observation scoped to only one of them.
"""
bank_id = f"test-obs-scopes-combined-{uuid.uuid4().hex[:8]}"
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
await memory.retain_batch_async(
bank_id=bank_id,
contents=[
{
"content": "Alice, a student, worked hard in the lesson with teacher Ben.",
"tags": ["user:alice", "teacher:ben"],
"observation_scopes": "combined",
}
],
request_context=request_context,
)
async with memory._pool.acquire() as conn:
observations = await conn.fetch(
"""
SELECT id, text, tags
FROM memory_units
WHERE bank_id = $1 AND fact_type = 'observation'
ORDER BY created_at
""",
bank_id,
)
try:
assert len(observations) >= 1, (
"Expected at least 1 observation, got 0"
)
tag_sets = [set(obs["tags"] or []) for obs in observations]
# All observations must carry both tags (combined scope)
combined = [ts for ts in tag_sets if "user:alice" in ts and "teacher:ben" in ts]
assert combined, f"Expected at least one observation with both tags, got: {tag_sets}"
# No observation should be scoped to only one tag
alice_only = [ts for ts in tag_sets if "user:alice" in ts and "teacher:ben" not in ts]
assert not alice_only, f"Expected no alice-only observation in combined mode, got: {tag_sets}"
ben_only = [ts for ts in tag_sets if "teacher:ben" in ts and "user:alice" not in ts]
assert not ben_only, f"Expected no ben-only observation in combined mode, got: {tag_sets}"
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_observation_scopes_all_combinations(memory: MemoryEngine, request_context):
"""Test that observation_scopes='all_combinations' generates passes for every tag subset.
A memory with tags=["user:alice", "teacher:ben"] and observation_scopes="all_combinations"
must produce observations covering all subsets: ["user:alice"], ["teacher:ben"], and
["user:alice", "teacher:ben"].
"""
bank_id = f"test-obs-scopes-allcombos-{uuid.uuid4().hex[:8]}"
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
await memory.retain_batch_async(
bank_id=bank_id,
contents=[
{
"content": "Alice, a student, worked hard in the lesson with teacher Ben.",
"tags": ["user:alice", "teacher:ben"],
"observation_scopes": "all_combinations",
}
],
request_context=request_context,
)
async with memory._pool.acquire() as conn:
observations = await conn.fetch(
"""
SELECT id, text, tags
FROM memory_units
WHERE bank_id = $1 AND fact_type = 'observation'
ORDER BY created_at
""",
bank_id,
)
try:
# With 2 tags there are 3 subsets: {alice}, {ben}, {alice, ben}
assert len(observations) >= 3, (
f"Expected at least 3 observations (one per subset), got {len(observations)}: "
+ str([dict(o) for o in observations])
)
tag_sets = [set(obs["tags"] or []) for obs in observations]
alice_only = [ts for ts in tag_sets if "user:alice" in ts and "teacher:ben" not in ts]
assert alice_only, f"Expected an observation scoped to 'user:alice' only, got: {tag_sets}"
ben_only = [ts for ts in tag_sets if "teacher:ben" in ts and "user:alice" not in ts]
assert ben_only, f"Expected an observation scoped to 'teacher:ben' only, got: {tag_sets}"
combined = [ts for ts in tag_sets if "user:alice" in ts and "teacher:ben" in ts]
assert combined, f"Expected an observation scoped to both tags, got: {tag_sets}"
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@@ -15,7 +15,7 @@ import pytest
from sqlalchemy import create_engine, text
from hindsight_api import MemoryEngine, RequestContext
from hindsight_api.engine.cross_encoder import CohereCrossEncoder, LocalSTCrossEncoder
from hindsight_api.engine.cross_encoder import CohereCrossEncoder, LocalSTCrossEncoder, ZeroEntropyCrossEncoder
from hindsight_api.engine.embeddings import CohereEmbeddings, LocalSTEmbeddings, OpenAIEmbeddings
from hindsight_api.engine.query_analyzer import DateparserQueryAnalyzer
from hindsight_api.engine.task_backend import SyncTaskBackend
@@ -98,9 +98,7 @@ def get_row_count(db_url: str, schema: str = "public") -> int:
"""Get the number of rows with embeddings in memory_units."""
engine = create_engine(db_url)
with engine.connect() as conn:
return conn.execute(
text(f"SELECT COUNT(*) FROM {schema}.memory_units WHERE embedding IS NOT NULL")
).scalar()
return conn.execute(text(f"SELECT COUNT(*) FROM {schema}.memory_units WHERE embedding IS NOT NULL")).scalar()
def insert_test_embedding(db_url: str, schema: str, dimension: int):
@@ -610,3 +608,59 @@ class TestCohereIntegration:
await memory.close()
except Exception:
pass
# =============================================================================
# ZeroEntropy Reranker Tests
# =============================================================================
def has_zeroentropy_api_key() -> bool:
"""Check if ZeroEntropy API key is available."""
return bool(os.environ.get("ZEROENTROPY_API_KEY"))
def get_zeroentropy_api_key() -> str:
"""Get ZeroEntropy API key from environment."""
return os.environ.get("ZEROENTROPY_API_KEY", "")
@pytest.fixture(scope="module")
def zeroentropy_cross_encoder():
"""Create ZeroEntropy cross-encoder instance."""
if not has_zeroentropy_api_key():
pytest.skip("ZeroEntropy API key not available (set ZEROENTROPY_API_KEY)")
cross_encoder = ZeroEntropyCrossEncoder(
api_key=get_zeroentropy_api_key(),
model="zerank-2",
)
loop = asyncio.new_event_loop()
try:
loop.run_until_complete(cross_encoder.initialize())
finally:
loop.close()
return cross_encoder
class TestZeroEntropyCrossEncoder:
"""Tests for ZeroEntropy cross-encoder/reranker."""
def test_zeroentropy_cross_encoder_initialization(self, zeroentropy_cross_encoder):
"""Test that ZeroEntropy cross-encoder initializes correctly."""
assert zeroentropy_cross_encoder.provider_name == "zeroentropy"
@pytest.mark.asyncio
async def test_zeroentropy_cross_encoder_predict(self, zeroentropy_cross_encoder):
"""Test that ZeroEntropy cross-encoder can score pairs."""
pairs = [
("What is the capital of France?", "Paris is the capital of France."),
("What is the capital of France?", "The Eiffel Tower is in Paris."),
("What is the capital of France?", "Python is a programming language."),
]
scores = await zeroentropy_cross_encoder.predict(pairs)
assert len(scores) == 3
assert all(isinstance(s, float) for s in scores)
# The first result should be most relevant
assert scores[0] > scores[2], "Direct answer should score higher than unrelated text"
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,61 @@
"""
Unit tests for metadata inclusion in fact extraction LLM prompt.
"""
from datetime import datetime
from hindsight_api.engine.retain.fact_extraction import _build_user_message
def test_build_user_message_includes_metadata():
"""Metadata key-value pairs should appear in the user message."""
event_date = datetime(2024, 6, 15, 12, 0, 0)
metadata = {"title": "Q2 Planning Doc", "source": "confluence", "author": "Alice"}
msg = _build_user_message(
chunk="Some content.",
chunk_index=0,
total_chunks=1,
event_date=event_date,
context="planning meeting",
metadata=metadata,
)
assert "title" in msg
assert "Q2 Planning Doc" in msg
assert "source" in msg
assert "confluence" in msg
assert "author" in msg
assert "Alice" in msg
def test_build_user_message_no_metadata():
"""When metadata is empty, the message should still be valid and not include a metadata section."""
event_date = datetime(2024, 6, 15, 12, 0, 0)
msg = _build_user_message(
chunk="Some content.",
chunk_index=0,
total_chunks=1,
event_date=event_date,
context="planning meeting",
metadata={},
)
assert "Some content." in msg
assert "Metadata:" not in msg
def test_build_user_message_without_metadata_arg():
"""Calling without metadata (default) should behave the same as empty metadata."""
event_date = datetime(2024, 6, 15, 12, 0, 0)
msg = _build_user_message(
chunk="Some content.",
chunk_index=0,
total_chunks=1,
event_date=event_date,
context="none",
)
assert "Some content." in msg
assert "Metadata:" not in msg
@@ -0,0 +1,340 @@
"""
Tests for Gemini safety settings feature.
Verifies that:
- Safety settings are read from env var and stored on GeminiLLM instances
- Settings are applied to GenerateContentConfig in call() and call_with_tools()
- The context variable override allows per-bank settings at request time
- None (unset) means Gemini's default safety settings are used (no override)
"""
import os
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
pytest.importorskip("google.genai")
SAMPLE_SAFETY_SETTINGS = [
{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_SEXUALLY_EXPLICIT", "threshold": "BLOCK_NONE"},
{"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "BLOCK_NONE"},
]
# ─── Config / env var parsing ─────────────────────────────────────────────────
def test_gemini_safety_settings_parsed_from_env():
"""Safety settings JSON from env var is parsed into HindsightConfig."""
import json
from hindsight_api.config import ENV_LLM_GEMINI_SAFETY_SETTINGS, HindsightConfig, clear_config_cache
settings_json = json.dumps(SAMPLE_SAFETY_SETTINGS)
with patch.dict(os.environ, {ENV_LLM_GEMINI_SAFETY_SETTINGS: settings_json}, clear=False):
clear_config_cache()
config = HindsightConfig.from_env()
assert config.llm_gemini_safety_settings == SAMPLE_SAFETY_SETTINGS
clear_config_cache()
def test_gemini_safety_settings_default_is_none():
"""When env var is not set, llm_gemini_safety_settings defaults to None."""
from hindsight_api.config import ENV_LLM_GEMINI_SAFETY_SETTINGS, HindsightConfig, clear_config_cache
env = {k: v for k, v in os.environ.items() if k != ENV_LLM_GEMINI_SAFETY_SETTINGS}
with patch.dict(os.environ, env, clear=True):
clear_config_cache()
config = HindsightConfig.from_env()
assert config.llm_gemini_safety_settings is None
clear_config_cache()
def test_gemini_safety_settings_is_configurable_field():
"""llm_gemini_safety_settings appears in configurable (per-bank) fields."""
from hindsight_api.config import HindsightConfig
assert "llm_gemini_safety_settings" in HindsightConfig.get_configurable_fields()
def test_gemini_safety_settings_not_in_credential_fields():
"""llm_gemini_safety_settings is NOT a credential — it is safe to expose via API."""
from hindsight_api.config import HindsightConfig
assert "llm_gemini_safety_settings" not in HindsightConfig.get_credential_fields()
# ─── GeminiLLM instance ───────────────────────────────────────────────────────
def _make_gemini_provider(safety_settings=None):
"""Return a GeminiLLM instance with a mocked genai.Client."""
with patch("google.genai.Client") as mock_client_cls:
mock_client_cls.return_value = MagicMock()
from hindsight_api.engine.providers.gemini_llm import GeminiLLM
provider = GeminiLLM(
provider="gemini",
api_key="fake-api-key",
base_url="",
model="gemini-2.5-flash",
gemini_safety_settings=safety_settings,
)
# Replace client with a fresh mock so we can inspect calls
provider._client = MagicMock()
return provider
def test_gemini_llm_stores_safety_settings():
"""GeminiLLM stores safety settings passed at construction."""
provider = _make_gemini_provider(safety_settings=SAMPLE_SAFETY_SETTINGS)
assert provider._safety_settings == SAMPLE_SAFETY_SETTINGS
def test_gemini_llm_no_safety_settings_is_none():
"""GeminiLLM._safety_settings is None when not provided."""
provider = _make_gemini_provider(safety_settings=None)
assert provider._safety_settings is None
# ─── call() applies safety settings ──────────────────────────────────────────
@pytest.mark.asyncio
async def test_call_applies_safety_settings():
"""call() includes safety_settings in GenerateContentConfig when configured."""
from google.genai import types as genai_types
provider = _make_gemini_provider(safety_settings=SAMPLE_SAFETY_SETTINGS)
# Build a fake successful response
fake_response = MagicMock()
fake_response.text = "hello"
fake_response.candidates = [MagicMock(finish_reason="STOP")]
fake_response.usage_metadata = MagicMock(prompt_token_count=5, candidates_token_count=2)
provider._client.aio.models.generate_content = AsyncMock(return_value=fake_response)
await provider.call(
messages=[{"role": "user", "content": "hi"}],
scope="test",
)
# Inspect the config passed to generate_content
call_args = provider._client.aio.models.generate_content.call_args
config_arg = call_args.kwargs.get("config") or call_args.args[0] if call_args.args else None
# config may be in kwargs or positional; grab from kwargs
config_arg = call_args.kwargs.get("config")
assert config_arg is not None, "GenerateContentConfig should have been passed"
assert hasattr(config_arg, "safety_settings"), "Config should have safety_settings"
assert config_arg.safety_settings is not None
categories = [s.category.value if hasattr(s.category, "value") else str(s.category) for s in config_arg.safety_settings]
assert "HARM_CATEGORY_HARASSMENT" in categories
assert "HARM_CATEGORY_HATE_SPEECH" in categories
assert "HARM_CATEGORY_SEXUALLY_EXPLICIT" in categories
assert "HARM_CATEGORY_DANGEROUS_CONTENT" in categories
thresholds = [s.threshold.value if hasattr(s.threshold, "value") else str(s.threshold) for s in config_arg.safety_settings]
assert all(t == "BLOCK_NONE" for t in thresholds)
@pytest.mark.asyncio
async def test_call_no_safety_settings_omits_key():
"""call() does NOT add safety_settings to GenerateContentConfig when none configured."""
provider = _make_gemini_provider(safety_settings=None)
fake_response = MagicMock()
fake_response.text = "hello"
fake_response.candidates = [MagicMock(finish_reason="STOP")]
fake_response.usage_metadata = MagicMock(prompt_token_count=5, candidates_token_count=2)
provider._client.aio.models.generate_content = AsyncMock(return_value=fake_response)
await provider.call(
messages=[{"role": "user", "content": "hi"}],
scope="test",
)
call_args = provider._client.aio.models.generate_content.call_args
config_arg = call_args.kwargs.get("config")
# When no safety settings, config is either None or lacks safety_settings
if config_arg is not None:
assert not hasattr(config_arg, "safety_settings") or config_arg.safety_settings is None
# ─── call_with_tools() applies safety settings ────────────────────────────────
@pytest.mark.asyncio
async def test_call_with_tools_applies_safety_settings():
"""call_with_tools() includes safety_settings in GenerateContentConfig."""
provider = _make_gemini_provider(safety_settings=SAMPLE_SAFETY_SETTINGS)
# Build a fake tool-use response (no tool calls, just text)
fake_part = MagicMock()
fake_part.text = "answer"
fake_part.function_call = None
fake_candidate = MagicMock()
fake_candidate.content = MagicMock(parts=[fake_part])
fake_response = MagicMock()
fake_response.candidates = [fake_candidate]
fake_response.usage_metadata = MagicMock(prompt_token_count=5, candidates_token_count=3)
provider._client.aio.models.generate_content = AsyncMock(return_value=fake_response)
tools = [
{
"type": "function",
"function": {
"name": "test_tool",
"description": "A test tool",
"parameters": {"type": "object", "properties": {}, "required": []},
},
}
]
await provider.call_with_tools(
messages=[{"role": "user", "content": "hi"}],
tools=tools,
scope="test",
)
call_args = provider._client.aio.models.generate_content.call_args
config_arg = call_args.kwargs.get("config")
assert config_arg is not None
assert config_arg.safety_settings is not None
categories = [s.category.value if hasattr(s.category, "value") else str(s.category) for s in config_arg.safety_settings]
assert "HARM_CATEGORY_HARASSMENT" in categories
# ─── with_config() override ───────────────────────────────────────────────────
def _make_llm_provider(safety_settings=None):
"""Return an LLMProvider (wrapping GeminiLLM) with a mocked genai.Client."""
with patch("google.genai.Client") as mock_client_cls:
mock_client_cls.return_value = MagicMock()
from hindsight_api.engine.llm_wrapper import LLMProvider
provider = LLMProvider(
provider="gemini",
api_key="fake-api-key",
base_url="",
model="gemini-2.5-flash",
gemini_safety_settings=safety_settings,
)
# Replace the underlying Gemini client with a fresh mock
provider._provider_impl._client = MagicMock()
return provider
def _fake_response():
r = MagicMock()
r.text = "hello"
r.candidates = [MagicMock(finish_reason="STOP")]
r.usage_metadata = MagicMock(prompt_token_count=5, candidates_token_count=2)
return r
def _make_config(safety_settings):
"""Return a minimal config-like object with llm_gemini_safety_settings."""
cfg = MagicMock()
cfg.llm_gemini_safety_settings = safety_settings
return cfg
@pytest.mark.asyncio
async def test_with_config_overrides_instance_settings():
"""with_config() settings take precedence over the provider instance defaults."""
instance_settings = [{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_ONLY_HIGH"}]
override_settings = [{"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_NONE"}]
provider = _make_llm_provider(safety_settings=instance_settings)
provider._provider_impl._client.aio.models.generate_content = AsyncMock(return_value=_fake_response())
configured = provider.with_config(_make_config(override_settings))
await configured.call(messages=[{"role": "user", "content": "hi"}], scope="test")
config_arg = provider._provider_impl._client.aio.models.generate_content.call_args.kwargs.get("config")
assert config_arg is not None
categories = [s.category.value if hasattr(s.category, "value") else str(s.category) for s in config_arg.safety_settings]
# Should use override_settings (HATE_SPEECH), not instance_settings (HARASSMENT)
assert "HARM_CATEGORY_HATE_SPEECH" in categories
assert "HARM_CATEGORY_HARASSMENT" not in categories
@pytest.mark.asyncio
async def test_with_config_none_falls_back_to_instance():
"""When with_config() supplies None, the instance default is used."""
instance_settings = [{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"}]
provider = _make_llm_provider(safety_settings=instance_settings)
provider._provider_impl._client.aio.models.generate_content = AsyncMock(return_value=_fake_response())
configured = provider.with_config(_make_config(None))
await configured.call(messages=[{"role": "user", "content": "hi"}], scope="test")
config_arg = provider._provider_impl._client.aio.models.generate_content.call_args.kwargs.get("config")
assert config_arg is not None
categories = [s.category.value if hasattr(s.category, "value") else str(s.category) for s in config_arg.safety_settings]
assert "HARM_CATEGORY_HARASSMENT" in categories
@pytest.mark.asyncio
async def test_with_config_resets_after_call():
"""The ContextVar is properly reset after a with_config() call (no leakage)."""
from hindsight_api.engine.providers.gemini_llm import _safety_settings_ctx
settings = [{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"}]
provider = _make_llm_provider(safety_settings=None)
provider._provider_impl._client.aio.models.generate_content = AsyncMock(return_value=_fake_response())
before = _safety_settings_ctx.get()
configured = provider.with_config(_make_config(settings))
await configured.call(messages=[{"role": "user", "content": "hi"}], scope="test")
after = _safety_settings_ctx.get()
assert after == before # ContextVar restored to its original value
# ─── LLMProvider reads safety settings from config ────────────────────────────
def test_llm_provider_reads_safety_settings_from_config():
"""LLMProvider reads llm_gemini_safety_settings from global config for Gemini provider."""
import json
from hindsight_api.config import ENV_LLM_GEMINI_SAFETY_SETTINGS, clear_config_cache
settings_json = json.dumps(SAMPLE_SAFETY_SETTINGS)
env_overrides = {
"HINDSIGHT_API_LLM_PROVIDER": "gemini",
"HINDSIGHT_API_LLM_API_KEY": "fake-key",
ENV_LLM_GEMINI_SAFETY_SETTINGS: settings_json,
}
with patch.dict(os.environ, env_overrides, clear=False):
clear_config_cache()
with patch("google.genai.Client") as mock_client_cls:
mock_client_cls.return_value = MagicMock()
from hindsight_api.engine.llm_wrapper import LLMProvider
provider = LLMProvider(
provider="gemini",
api_key="fake-key",
base_url="",
model="gemini-2.5-flash",
)
assert provider.gemini_safety_settings == SAMPLE_SAFETY_SETTINGS
clear_config_cache()
+171
View File
@@ -0,0 +1,171 @@
"""
Tests for server-side filtering in the graph API endpoint.
Verifies that q (text search) and tags filters work correctly
when passed as query parameters to GET /v1/default/banks/{bank_id}/graph.
"""
from datetime import datetime
import httpx
import pytest
import pytest_asyncio
from hindsight_api.api import create_app
@pytest_asyncio.fixture
async def api_client(memory):
"""Create an async test client for the FastAPI app."""
app = create_app(memory, initialize_memory=False)
transport = httpx.ASGITransport(app=app)
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
yield client
@pytest.fixture
def test_bank_id():
"""Provide a unique bank ID for this test run."""
return f"graph_filter_test_{datetime.now().timestamp()}"
@pytest.mark.asyncio
async def test_graph_no_filter_returns_all(api_client, test_bank_id):
"""Without filters the graph endpoint returns all memories."""
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Alice loves hiking in the mountains.", "tags": ["user_alice"]},
{"content": "Bob enjoys swimming at the beach.", "tags": ["user_bob"]},
]
},
)
assert response.status_code == 200
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/graph")
assert response.status_code == 200
data = response.json()
assert "table_rows" in data
texts = [row["text"] for row in data["table_rows"]]
assert any("Alice" in t for t in texts)
assert any("Bob" in t for t in texts)
@pytest.mark.asyncio
async def test_graph_q_filter_returns_matching(api_client, test_bank_id):
"""The q parameter filters memories by text content."""
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Alice loves hiking in the mountains."},
{"content": "Bob enjoys swimming at the beach."},
]
},
)
assert response.status_code == 200
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/graph", params={"q": "Alice"})
assert response.status_code == 200
data = response.json()
texts = [row["text"] for row in data["table_rows"]]
assert all("Alice" in t or "alice" in t.lower() for t in texts), (
f"Expected only Alice memories, got: {texts}"
)
assert not any("Bob" in t for t in texts)
@pytest.mark.asyncio
async def test_graph_q_filter_case_insensitive(api_client, test_bank_id):
"""The q filter is case-insensitive."""
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Alice loves hiking in the mountains."},
{"content": "Bob enjoys swimming at the beach."},
]
},
)
assert response.status_code == 200
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/graph", params={"q": "alice"})
assert response.status_code == 200
data = response.json()
texts = [row["text"] for row in data["table_rows"]]
assert any("Alice" in t for t in texts)
assert not any("Bob" in t for t in texts)
@pytest.mark.asyncio
async def test_graph_tags_filter_returns_matching(api_client, test_bank_id):
"""The tags parameter filters memories to only those with matching tags."""
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Alice loves hiking.", "tags": ["user_alice"]},
{"content": "Bob enjoys swimming.", "tags": ["user_bob"]},
]
},
)
assert response.status_code == 200
response = await api_client.get(
f"/v1/default/banks/{test_bank_id}/graph",
params={"tags": "user_alice", "tags_match": "all_strict"},
)
assert response.status_code == 200
data = response.json()
texts = [row["text"] for row in data["table_rows"]]
assert any("Alice" in t for t in texts)
assert not any("Bob" in t for t in texts)
@pytest.mark.asyncio
async def test_graph_q_and_tags_filter_combined(api_client, test_bank_id):
"""Combining q and tags filters applies both server-side."""
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Alice loves hiking.", "tags": ["user_alice"]},
{"content": "Alice also loves coding.", "tags": ["user_alice"]},
{"content": "Bob enjoys swimming.", "tags": ["user_bob"]},
]
},
)
assert response.status_code == 200
response = await api_client.get(
f"/v1/default/banks/{test_bank_id}/graph",
params={"q": "hiking", "tags": "user_alice", "tags_match": "all_strict"},
)
assert response.status_code == 200
data = response.json()
texts = [row["text"] for row in data["table_rows"]]
assert any("hiking" in t.lower() for t in texts)
assert not any("coding" in t.lower() for t in texts)
assert not any("Bob" in t for t in texts)
@pytest.mark.asyncio
async def test_graph_q_filter_empty_results(api_client, test_bank_id):
"""The q filter returns empty results when no memory matches."""
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Alice loves hiking."},
]
},
)
assert response.status_code == 200
response = await api_client.get(
f"/v1/default/banks/{test_bank_id}/graph",
params={"q": "zzznomatchzzz"},
)
assert response.status_code == 200
data = response.json()
assert data["table_rows"] == []
@@ -15,9 +15,6 @@ from hindsight_api.config_resolver import ConfigResolver
from hindsight_api.extensions.tenant import TenantExtension
from hindsight_api.models import RequestContext
# Enable bank config API for all tests in this module
os.environ["HINDSIGHT_API_ENABLE_BANK_CONFIG_API"] = "true"
class MockTenantExtension(TenantExtension):
"""Mock tenant extension for testing tenant-level config."""
@@ -74,12 +71,22 @@ async def test_hierarchical_fields_categorization():
# Verify configurable fields include behavioral settings (safe to modify)
assert "retain_extraction_mode" in configurable
assert "enable_observations" in configurable
assert "retain_chunk_size" in configurable
assert "retain_mission" in configurable
assert "retain_custom_instructions" in configurable
assert "retain_chunk_size" in configurable
assert "enable_observations" in configurable
assert "observations_mission" in configurable
assert "reflect_mission" in configurable
assert "disposition_skepticism" in configurable
assert "disposition_literalism" in configurable
assert "disposition_empathy" in configurable
# Verify count is correct (only 4 fields)
assert len(configurable) == 4
# Verify entity labels fields are included
assert "entities_allow_free_form" in configurable
assert "entity_labels" in configurable
# Verify count is correct
assert len(configurable) == 14
# Verify credential fields (NEVER exposed)
assert "llm_api_key" in credentials
+174
View File
@@ -0,0 +1,174 @@
"""
Tests for list_documents pagination and tags filtering.
"""
from datetime import datetime, timezone
import pytest
async def _retain_doc(memory, bank_id, document_id, tags, request_context):
"""Helper to retain a document with given tags. Uses gibberish content to avoid LLM
fact extraction (documents are persisted even with zero facts)."""
await memory.retain_batch_async(
bank_id=bank_id,
contents=[{"content": f"xyzabc123 !@# $$$ {document_id}"}],
document_id=document_id,
document_tags=tags or None,
request_context=request_context,
)
@pytest.mark.asyncio
async def test_list_documents_offset_pagination(memory, request_context):
"""offset parameter returns the correct slice of documents."""
bank_id = f"test_list_docs_offset_{datetime.now(timezone.utc).timestamp()}"
try:
for i in range(4):
await _retain_doc(memory, bank_id, f"doc-{i:02d}", [], request_context)
# All documents, ordered by created_at DESC → doc-03, doc-02, doc-01, doc-00
all_docs = await memory.list_documents(
bank_id=bank_id, limit=10, offset=0, request_context=request_context
)
assert all_docs["total"] == 4
assert len(all_docs["items"]) == 4
all_ids = [d["id"] for d in all_docs["items"]]
# offset=2 should skip the first two and return the remaining two
page2 = await memory.list_documents(
bank_id=bank_id, limit=10, offset=2, request_context=request_context
)
assert page2["total"] == 4 # total is always the full count
assert len(page2["items"]) == 2
assert [d["id"] for d in page2["items"]] == all_ids[2:]
# offset beyond total returns empty items but correct total
beyond = await memory.list_documents(
bank_id=bank_id, limit=10, offset=10, request_context=request_context
)
assert beyond["total"] == 4
assert beyond["items"] == []
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_list_documents_tags_filter_any_strict(memory, request_context):
"""tags filter with any_strict returns only tagged documents that match."""
bank_id = f"test_list_docs_tags_{datetime.now(timezone.utc).timestamp()}"
try:
await _retain_doc(memory, bank_id, "doc-alpha", ["team-a"], request_context)
await _retain_doc(memory, bank_id, "doc-beta", ["team-b"], request_context)
await _retain_doc(memory, bank_id, "doc-both", ["team-a", "team-b"], request_context)
await _retain_doc(memory, bank_id, "doc-untagged", [], request_context)
# any_strict: only docs with at least one of the given tags, untagged excluded
result = await memory.list_documents(
bank_id=bank_id,
tags=["team-a"],
tags_match="any_strict",
request_context=request_context,
)
ids = {d["id"] for d in result["items"]}
assert ids == {"doc-alpha", "doc-both"}
assert result["total"] == 2
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_list_documents_tags_filter_any_includes_untagged(memory, request_context):
"""tags filter with 'any' mode includes untagged documents."""
bank_id = f"test_list_docs_tags_any_{datetime.now(timezone.utc).timestamp()}"
try:
await _retain_doc(memory, bank_id, "doc-tagged", ["team-a"], request_context)
await _retain_doc(memory, bank_id, "doc-other", ["team-b"], request_context)
await _retain_doc(memory, bank_id, "doc-untagged", [], request_context)
result = await memory.list_documents(
bank_id=bank_id,
tags=["team-a"],
tags_match="any",
request_context=request_context,
)
ids = {d["id"] for d in result["items"]}
# "any" includes untagged + matching tagged
assert "doc-tagged" in ids
assert "doc-untagged" in ids
assert "doc-other" not in ids
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_list_documents_tags_filter_all_strict(memory, request_context):
"""tags filter with all_strict returns only docs that have ALL the specified tags."""
bank_id = f"test_list_docs_tags_all_{datetime.now(timezone.utc).timestamp()}"
try:
await _retain_doc(memory, bank_id, "doc-a-only", ["team-a"], request_context)
await _retain_doc(memory, bank_id, "doc-a-and-b", ["team-a", "team-b"], request_context)
await _retain_doc(memory, bank_id, "doc-untagged", [], request_context)
result = await memory.list_documents(
bank_id=bank_id,
tags=["team-a", "team-b"],
tags_match="all_strict",
request_context=request_context,
)
ids = {d["id"] for d in result["items"]}
assert ids == {"doc-a-and-b"}
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_list_documents_no_tags_filter_returns_all(memory, request_context):
"""When no tags filter is specified, all documents are returned."""
bank_id = f"test_list_docs_no_tags_{datetime.now(timezone.utc).timestamp()}"
try:
await _retain_doc(memory, bank_id, "doc-tagged", ["team-a"], request_context)
await _retain_doc(memory, bank_id, "doc-untagged", [], request_context)
result = await memory.list_documents(
bank_id=bank_id,
tags=None,
request_context=request_context,
)
ids = {d["id"] for d in result["items"]}
assert ids == {"doc-tagged", "doc-untagged"}
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_list_documents_tags_and_search_query_combined(memory, request_context):
"""tags filter and q (search_query) can be combined."""
bank_id = f"test_list_docs_tags_q_{datetime.now(timezone.utc).timestamp()}"
try:
await _retain_doc(memory, bank_id, "report-2024", ["team-a"], request_context)
await _retain_doc(memory, bank_id, "report-2025", ["team-b"], request_context)
await _retain_doc(memory, bank_id, "summary-2024", ["team-a"], request_context)
result = await memory.list_documents(
bank_id=bank_id,
search_query="report",
tags=["team-a"],
tags_match="any_strict",
request_context=request_context,
)
ids = {d["id"] for d in result["items"]}
assert ids == {"report-2024"}
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@@ -85,6 +85,7 @@ class TestLiteLLMSDKEmbeddings:
model="cohere/embed-english-v3.0",
input=["test"],
api_key="test_key",
encoding_format="float",
)
async def test_initialization_missing_package(self):
@@ -137,6 +138,7 @@ class TestLiteLLMSDKEmbeddings:
model="cohere/embed-english-v3.0",
input=["Hello world"],
api_key="test_key",
encoding_format="float",
)
async def test_encode_multiple_texts(self, embeddings, mock_litellm):
+2 -2
View File
@@ -332,8 +332,8 @@ async def test_llm_provider_consolidation(memory_no_llm_verify, request_context,
test_bank_id = f"llm_test_consolidation_{provider}_{model}_{datetime.now().timestamp()}"
# Enable observations for this bank
from hindsight_api.config import get_config
config = get_config()
from hindsight_api.config import _get_raw_config
config = _get_raw_config()
original_value = config.enable_observations
config.enable_observations = True
+2 -2
View File
@@ -165,5 +165,5 @@ class TestMCPExtensionIntegration:
assert "create_bank" in tools
# Extension tool also present
assert "test_extension_tool" in tools
# At least 11 core + 1 extension = 12 tools (may grow as new tools are added)
assert len(tools) >= 12
# At least 29 core + 1 extension = 30 tools (may grow as new tools are added)
assert len(tools) >= 30
+74 -12
View File
@@ -46,16 +46,15 @@ async def test_mcp_tools_use_context_bank_id(mock_memory):
assert "retain" in tools
assert "recall" in tools
# Test retain with bank_id from context (use async_processing=False for synchronous test)
token = _current_bank_id.set("context-bank-id")
try:
retain_tool = tools["retain"]
result = await retain_tool.fn(content="test content", context="test_context", async_processing=False)
assert "successfully" in result.lower()
result = await retain_tool.fn(content="test content", context="test_context")
assert result["status"] == "accepted"
# Verify the memory was called with the context bank_id
mock_memory.retain_batch_async.assert_called_once()
call_kwargs = mock_memory.retain_batch_async.call_args.kwargs
mock_memory.submit_async_retain.assert_called_once()
call_kwargs = mock_memory.submit_async_retain.call_args.kwargs
assert call_kwargs["bank_id"] == "context-bank-id"
finally:
_current_bank_id.reset(token)
@@ -133,12 +132,12 @@ async def test_mcp_tools_propagate_api_key(mock_memory):
api_key_token = _current_api_key.set("test-bearer-token")
try:
retain_tool = tools["retain"]
result = await retain_tool.fn(content="test content", context="test_context", async_processing=False)
assert "successfully" in result.lower()
result = await retain_tool.fn(content="test content", context="test_context")
assert result["status"] == "accepted"
# Verify the memory was called with request_context containing api_key
mock_memory.retain_batch_async.assert_called_once()
call_kwargs = mock_memory.retain_batch_async.call_args.kwargs
mock_memory.submit_async_retain.assert_called_once()
call_kwargs = mock_memory.submit_async_retain.call_args.kwargs
assert call_kwargs["request_context"].api_key == "test-bearer-token"
finally:
_current_bank_id.reset(bank_token)
@@ -200,11 +199,11 @@ async def test_mcp_tools_propagate_tenant_id_and_api_key_id(mock_memory):
key_id_token = _current_api_key_id.set("key-uuid-456")
try:
retain_tool = tools["retain"]
await retain_tool.fn(content="test content", context="test_context", async_processing=False)
await retain_tool.fn(content="test content", context="test_context")
# Verify the RequestContext passed to memory engine has all auth fields
mock_memory.retain_batch_async.assert_called_once()
request_context = mock_memory.retain_batch_async.call_args.kwargs["request_context"]
mock_memory.submit_async_retain.assert_called_once()
request_context = mock_memory.submit_async_retain.call_args.kwargs["request_context"]
assert request_context.api_key == "hsk_test_key"
assert request_context.tenant_id == "org-billing-123"
assert request_context.api_key_id == "key-uuid-456"
@@ -352,6 +351,69 @@ async def test_middleware_handles_both_endpoints(mock_memory):
assert "create_bank" not in single_bank_tools
def test_global_mcp_enabled_tools_filter_restricts_registered_tools(mock_memory):
"""Test that global mcp_enabled_tools env setting restricts which tools are registered."""
from unittest.mock import MagicMock, patch
from hindsight_api.api.mcp import create_mcp_server
mock_cfg = MagicMock()
mock_cfg.mcp_enabled_tools = ["retain", "recall"]
with patch("hindsight_api.api.mcp._get_raw_config", return_value=mock_cfg):
mcp_server = create_mcp_server(mock_memory, multi_bank=True)
tools = mcp_server._tool_manager._tools
assert "retain" in tools
assert "recall" in tools
assert "reflect" not in tools
assert "list_banks" not in tools
assert "create_bank" not in tools
assert "list_mental_models" not in tools
def test_global_mcp_enabled_tools_none_exposes_all_tools(mock_memory):
"""Test that mcp_enabled_tools=None (default) exposes all tools."""
from unittest.mock import MagicMock, patch
from hindsight_api.api.mcp import create_mcp_server
mock_cfg = MagicMock()
mock_cfg.mcp_enabled_tools = None
with patch("hindsight_api.api.mcp._get_raw_config", return_value=mock_cfg):
mcp_server = create_mcp_server(mock_memory, multi_bank=True)
tools = mcp_server._tool_manager._tools
assert "retain" in tools
assert "recall" in tools
assert "reflect" in tools
assert "list_banks" in tools
assert "create_bank" in tools
def test_global_mcp_enabled_tools_intersects_with_single_bank_mode(mock_memory):
"""Test that global filter intersects with single-bank mode tool set.
list_banks is in the global allowlist but NOT in single-bank mode, so it
should be absent from the final registered set.
"""
from unittest.mock import MagicMock, patch
from hindsight_api.api.mcp import create_mcp_server
mock_cfg = MagicMock()
mock_cfg.mcp_enabled_tools = ["retain", "recall", "list_banks"]
with patch("hindsight_api.api.mcp._get_raw_config", return_value=mock_cfg):
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
tools = mcp_server._tool_manager._tools
assert "retain" in tools
assert "recall" in tools
assert "list_banks" not in tools # single-bank mode excludes it regardless
@pytest.mark.asyncio
async def test_routing_logic_from_url_path():
"""Test that routing correctly selects server based on URL structure.
+715 -2
View File
@@ -77,8 +77,9 @@ class TestBuildContentDict:
@pytest.fixture
def mock_memory():
"""Create a mock MemoryEngine with mental model methods."""
"""Create a mock MemoryEngine with all MCP tool methods."""
memory = MagicMock()
# Mental model methods
memory.list_mental_models = AsyncMock(
return_value=[
{"id": "mm-1", "name": "Coding Prefs", "source_query": "coding preferences?", "content": "Prefers Python"},
@@ -104,6 +105,41 @@ def mock_memory():
}
)
memory.delete_mental_model = AsyncMock(return_value=True)
# Retain/recall/reflect
memory.retain_batch_async = AsyncMock()
memory.submit_async_retain = AsyncMock(return_value={"operation_id": "op-retain"})
memory.recall_async = AsyncMock(return_value=MagicMock(model_dump_json=lambda indent=None: '{"results": []}', model_dump=lambda: {"results": []}))
memory.reflect_async = AsyncMock(return_value=MagicMock(model_dump_json=lambda indent=None: '{"text": "reflection"}', model_dump=lambda: {"text": "reflection"}, structured_output=None))
# Directive methods
memory.list_directives = AsyncMock(return_value=[{"id": "dir-1", "name": "Be concise", "content": "Keep responses short"}])
memory.create_directive = AsyncMock(return_value={"id": "dir-new", "name": "Test", "content": "Test content"})
memory.delete_directive = AsyncMock(return_value=True)
# Memory browsing methods
memory.list_memory_units = AsyncMock(return_value={"items": [{"id": "mem-1", "content": "Test"}], "total": 1})
memory.get_memory_unit = AsyncMock(return_value={"id": "mem-1", "content": "Test memory"})
memory.delete_memory_unit = AsyncMock(return_value={"deleted_count": 1})
# Document methods
memory.list_documents = AsyncMock(return_value={"items": [{"id": "doc-1", "name": "Test Doc"}], "total": 1})
memory.get_document = AsyncMock(return_value={"id": "doc-1", "name": "Test Doc"})
memory.delete_document = AsyncMock(return_value={"deleted_memories": 5})
# Operation methods
memory.list_operations = AsyncMock(return_value={"items": [{"id": "op-1", "status": "completed"}]})
memory.get_operation_status = AsyncMock(return_value={"id": "op-1", "status": "completed", "progress": 100})
memory.cancel_operation = AsyncMock(return_value={"id": "op-1", "status": "cancelled"})
# Tags & bank methods
memory.list_tags = AsyncMock(return_value={"items": ["tag1", "tag2"], "total": 2})
memory.get_bank_profile = AsyncMock(return_value={"id": "test-bank", "name": "Test Bank", "mission": "Testing"})
memory.get_bank_stats = AsyncMock(return_value={"nodes": 100, "links": 50})
memory.update_bank = AsyncMock(return_value={"id": "test-bank", "name": "Updated"})
memory.delete_bank = AsyncMock(return_value={"deleted_memories": 10, "deleted_entities": 5})
memory.list_banks = AsyncMock(return_value=[])
return memory
@@ -211,7 +247,7 @@ class TestMentalModelToolRegistration:
assert request_context.api_key == "test-api-key"
def test_mental_model_tools_in_default_set(self):
"""Mental model tools should be in the default tools set when config.tools is None."""
"""All tools should be in the default tools set when config.tools is None."""
from fastmcp import FastMCP
memory = MagicMock()
@@ -229,6 +265,21 @@ class TestMentalModelToolRegistration:
memory.submit_async_refresh_mental_model = AsyncMock()
memory.update_mental_model = AsyncMock()
memory.delete_mental_model = AsyncMock()
memory.list_directives = AsyncMock(return_value=[])
memory.create_directive = AsyncMock()
memory.delete_directive = AsyncMock()
memory.list_memory_units = AsyncMock(return_value={})
memory.get_memory_unit = AsyncMock()
memory.delete_memory_unit = AsyncMock()
memory.list_documents = AsyncMock(return_value={})
memory.get_document = AsyncMock()
memory.delete_document = AsyncMock()
memory.list_operations = AsyncMock(return_value={})
memory.get_operation_status = AsyncMock()
memory.cancel_operation = AsyncMock()
memory.list_tags = AsyncMock(return_value={})
memory.get_bank_stats = AsyncMock(return_value={})
memory.delete_bank = AsyncMock(return_value={})
mcp = FastMCP("test", stateless_http=True)
config = MCPToolsConfig(
@@ -241,6 +292,18 @@ class TestMentalModelToolRegistration:
assert "list_mental_models" in tools
assert "create_mental_model" in tools
assert "refresh_mental_model" in tools
# New tools
assert "list_directives" in tools
assert "list_memories" in tools
assert "list_documents" in tools
assert "list_operations" in tools
assert "list_tags" in tools
assert "get_bank" in tools
assert "get_bank_stats" in tools
assert "update_bank" in tools
assert "delete_bank" in tools
assert "clear_memories" in tools
assert len(tools) == 29
@pytest.fixture
@@ -644,3 +707,653 @@ class TestMentalModelInputValidation:
result = await _tools(mcp_server_single_bank)["get_mental_model"].fn(mental_model_id="missing")
assert isinstance(result, dict)
assert "fixed-bank" in result["error"]
# =========================================================================
# New Parameter Tests for Existing Tools
# =========================================================================
def _make_mcp_server(mock_memory, tools, include_bank_id=True):
"""Helper to create an MCP server with specific tools."""
from fastmcp import FastMCP
mcp = FastMCP("test", stateless_http=True)
config = MCPToolsConfig(
bank_id_resolver=lambda: "test-bank",
include_bank_id_param=include_bank_id,
tools=tools,
)
register_mcp_tools(mcp, mock_memory, config)
return mcp
@pytest.mark.asyncio
class TestRetainNewParams:
"""Tests for new retain parameters: tags, metadata, document_id."""
async def test_retain_with_tags(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"retain"})
await _tools(mcp)["retain"].fn(content="test", tags=["user:123", "project:alpha"])
call_args = mock_memory.submit_async_retain.call_args
contents = call_args.kwargs["contents"]
assert contents[0]["tags"] == ["user:123", "project:alpha"]
async def test_retain_with_metadata(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"retain"})
await _tools(mcp)["retain"].fn(content="test", metadata={"source": "slack"})
call_args = mock_memory.submit_async_retain.call_args
contents = call_args.kwargs["contents"]
assert contents[0]["metadata"] == {"source": "slack"}
async def test_retain_with_document_id(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"retain"})
await _tools(mcp)["retain"].fn(content="test", document_id="doc-1")
call_args = mock_memory.submit_async_retain.call_args
contents = call_args.kwargs["contents"]
assert contents[0]["document_id"] == "doc-1"
async def test_retain_without_new_params_backward_compat(self, mock_memory):
"""Existing behavior preserved when new params not provided."""
mcp = _make_mcp_server(mock_memory, {"retain"})
await _tools(mcp)["retain"].fn(content="test")
call_args = mock_memory.submit_async_retain.call_args
contents = call_args.kwargs["contents"]
assert "tags" not in contents[0]
assert "metadata" not in contents[0]
assert "document_id" not in contents[0]
@pytest.mark.asyncio
class TestRecallNewParams:
"""Tests for new recall parameters: budget, types, tags, tags_match, query_timestamp."""
async def test_recall_default_budget_high(self, mock_memory):
"""Default budget should be HIGH (backward compat)."""
from hindsight_api.engine.memory_engine import Budget
mcp = _make_mcp_server(mock_memory, {"recall"})
await _tools(mcp)["recall"].fn(query="test")
call_kwargs = mock_memory.recall_async.call_args.kwargs
assert call_kwargs["budget"] == Budget.HIGH
async def test_recall_budget_low(self, mock_memory):
from hindsight_api.engine.memory_engine import Budget
mcp = _make_mcp_server(mock_memory, {"recall"})
await _tools(mcp)["recall"].fn(query="test", budget="low")
call_kwargs = mock_memory.recall_async.call_args.kwargs
assert call_kwargs["budget"] == Budget.LOW
async def test_recall_with_types(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"recall"})
await _tools(mcp)["recall"].fn(query="test", types=["world"])
call_kwargs = mock_memory.recall_async.call_args.kwargs
assert call_kwargs["fact_type"] == ["world"]
async def test_recall_default_types_all(self, mock_memory):
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
mcp = _make_mcp_server(mock_memory, {"recall"})
await _tools(mcp)["recall"].fn(query="test")
call_kwargs = mock_memory.recall_async.call_args.kwargs
assert call_kwargs["fact_type"] == list(VALID_RECALL_FACT_TYPES)
async def test_recall_with_tags(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"recall"})
await _tools(mcp)["recall"].fn(query="test", tags=["project:x"])
call_kwargs = mock_memory.recall_async.call_args.kwargs
assert call_kwargs["tags"] == ["project:x"]
assert call_kwargs["tags_match"] == "any"
async def test_recall_with_query_timestamp(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"recall"})
await _tools(mcp)["recall"].fn(query="test", query_timestamp="2024-01-01T00:00:00Z")
call_kwargs = mock_memory.recall_async.call_args.kwargs
assert call_kwargs["question_date"] == datetime(2024, 1, 1, 0, 0, 0, tzinfo=timezone.utc)
@pytest.mark.asyncio
class TestReflectNewParams:
"""Tests for new reflect parameters: max_tokens, response_schema, tags, tags_match."""
async def test_reflect_with_max_tokens(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"reflect"})
await _tools(mcp)["reflect"].fn(query="test", max_tokens=2048)
call_kwargs = mock_memory.reflect_async.call_args.kwargs
assert call_kwargs["max_tokens"] == 2048
async def test_reflect_with_response_schema(self, mock_memory):
schema = {"type": "object", "properties": {"answer": {"type": "string"}}}
mock_memory.reflect_async = AsyncMock(
return_value=MagicMock(
model_dump_json=lambda indent=None: '{"text": "reflection"}',
model_dump=lambda: {"text": "reflection"},
structured_output={"answer": "yes"},
)
)
mcp = _make_mcp_server(mock_memory, {"reflect"})
result = await _tools(mcp)["reflect"].fn(query="test", response_schema=schema)
call_kwargs = mock_memory.reflect_async.call_args.kwargs
assert call_kwargs["response_schema"] == schema
# Multi-bank returns JSON string
import json
parsed = json.loads(result)
assert parsed["structured_output"] == {"answer": "yes"}
async def test_reflect_with_tags(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"reflect"})
await _tools(mcp)["reflect"].fn(query="test", tags=["scope:work"], tags_match="all")
call_kwargs = mock_memory.reflect_async.call_args.kwargs
assert call_kwargs["tags"] == ["scope:work"]
assert call_kwargs["tags_match"] == "all"
async def test_reflect_without_tags_no_tags_in_kwargs(self, mock_memory):
"""When tags not provided, they should not be passed to engine."""
mcp = _make_mcp_server(mock_memory, {"reflect"})
await _tools(mcp)["reflect"].fn(query="test")
call_kwargs = mock_memory.reflect_async.call_args.kwargs
assert "tags" not in call_kwargs
@pytest.mark.asyncio
class TestMentalModelTrigger:
"""Tests for trigger_refresh_after_consolidation on create/update mental model."""
async def test_create_with_trigger(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"create_mental_model"})
await _tools(mcp)["create_mental_model"].fn(
name="Test", source_query="query", trigger_refresh_after_consolidation=True
)
call_kwargs = mock_memory.create_mental_model.call_args.kwargs
assert call_kwargs["trigger"] == {"refresh_after_consolidation": True}
async def test_create_default_trigger_false(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"create_mental_model"})
await _tools(mcp)["create_mental_model"].fn(name="Test", source_query="query")
call_kwargs = mock_memory.create_mental_model.call_args.kwargs
assert call_kwargs["trigger"] == {"refresh_after_consolidation": False}
async def test_update_with_trigger(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"update_mental_model"})
await _tools(mcp)["update_mental_model"].fn(
mental_model_id="mm-1", trigger_refresh_after_consolidation=True
)
call_kwargs = mock_memory.update_mental_model.call_args.kwargs
assert call_kwargs["trigger"] == {"refresh_after_consolidation": True}
async def test_update_without_trigger_no_trigger_in_kwargs(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"update_mental_model"})
await _tools(mcp)["update_mental_model"].fn(mental_model_id="mm-1", name="New Name")
call_kwargs = mock_memory.update_mental_model.call_args.kwargs
assert "trigger" not in call_kwargs
# =========================================================================
# Directive Tool Tests
# =========================================================================
@pytest.mark.asyncio
class TestDirectiveTools:
async def test_list_directives_multi_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_directives"}, include_bank_id=True)
result = await _tools(mcp)["list_directives"].fn()
assert '"dir-1"' in result
mock_memory.list_directives.assert_called_once()
assert mock_memory.list_directives.call_args[0][0] == "test-bank"
async def test_list_directives_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_directives"}, include_bank_id=False)
result = await _tools(mcp)["list_directives"].fn()
assert isinstance(result, dict)
assert len(result["items"]) == 1
async def test_create_directive(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"create_directive"}, include_bank_id=True)
result = await _tools(mcp)["create_directive"].fn(name="Test", content="Be concise", priority=5)
assert '"dir-new"' in result
call_args = mock_memory.create_directive.call_args
assert call_args[0][0] == "test-bank"
assert call_args.kwargs["name"] == "Test"
assert call_args.kwargs["content"] == "Be concise"
assert call_args.kwargs["priority"] == 5
async def test_delete_directive(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"delete_directive"}, include_bank_id=True)
result = await _tools(mcp)["delete_directive"].fn(directive_id="dir-1")
assert '"deleted"' in result
assert mock_memory.delete_directive.call_args[0][1] == "dir-1"
async def test_delete_directive_not_found(self, mock_memory):
mock_memory.delete_directive.return_value = False
mcp = _make_mcp_server(mock_memory, {"delete_directive"}, include_bank_id=True)
result = await _tools(mcp)["delete_directive"].fn(directive_id="missing")
assert "not found" in result
# =========================================================================
# Memory Browsing Tool Tests
# =========================================================================
@pytest.mark.asyncio
class TestMemoryBrowsingTools:
async def test_list_memories_default(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_memories"}, include_bank_id=True)
result = await _tools(mcp)["list_memories"].fn()
assert '"mem-1"' in result
call_kwargs = mock_memory.list_memory_units.call_args.kwargs
assert call_kwargs["limit"] == 100
assert call_kwargs["offset"] == 0
async def test_list_memories_with_filters(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_memories"}, include_bank_id=True)
await _tools(mcp)["list_memories"].fn(type="world", q="test query", limit=50)
call_kwargs = mock_memory.list_memory_units.call_args.kwargs
assert call_kwargs["fact_type"] == "world"
assert call_kwargs["search_query"] == "test query"
assert call_kwargs["limit"] == 50
async def test_get_memory(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"get_memory"}, include_bank_id=True)
result = await _tools(mcp)["get_memory"].fn(memory_id="mem-1")
assert '"mem-1"' in result
async def test_get_memory_not_found(self, mock_memory):
mock_memory.get_memory_unit.return_value = None
mcp = _make_mcp_server(mock_memory, {"get_memory"}, include_bank_id=True)
result = await _tools(mcp)["get_memory"].fn(memory_id="missing")
assert "not found" in result
async def test_delete_memory(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"delete_memory"}, include_bank_id=True)
result = await _tools(mcp)["delete_memory"].fn(memory_id="mem-1")
assert '"deleted"' in result
assert mock_memory.delete_memory_unit.call_args.kwargs["unit_id"] == "mem-1"
async def test_list_memories_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_memories"}, include_bank_id=False)
result = await _tools(mcp)["list_memories"].fn()
assert isinstance(result, dict)
# =========================================================================
# Document Tool Tests
# =========================================================================
@pytest.mark.asyncio
class TestDocumentTools:
async def test_list_documents(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_documents"}, include_bank_id=True)
result = await _tools(mcp)["list_documents"].fn()
assert '"doc-1"' in result
async def test_get_document(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"get_document"}, include_bank_id=True)
result = await _tools(mcp)["get_document"].fn(document_id="doc-1")
assert '"doc-1"' in result
async def test_get_document_not_found(self, mock_memory):
mock_memory.get_document.return_value = None
mcp = _make_mcp_server(mock_memory, {"get_document"}, include_bank_id=True)
result = await _tools(mcp)["get_document"].fn(document_id="missing")
assert "not found" in result
async def test_delete_document(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"delete_document"}, include_bank_id=True)
result = await _tools(mcp)["delete_document"].fn(document_id="doc-1")
assert '"deleted"' in result
async def test_list_documents_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_documents"}, include_bank_id=False)
result = await _tools(mcp)["list_documents"].fn()
assert isinstance(result, dict)
# =========================================================================
# Operation Tool Tests
# =========================================================================
@pytest.mark.asyncio
class TestOperationTools:
async def test_list_operations(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_operations"}, include_bank_id=True)
result = await _tools(mcp)["list_operations"].fn()
assert '"op-1"' in result
async def test_list_operations_with_status(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_operations"}, include_bank_id=True)
await _tools(mcp)["list_operations"].fn(status="completed", limit=10)
call_kwargs = mock_memory.list_operations.call_args.kwargs
assert call_kwargs["status"] == "completed"
assert call_kwargs["limit"] == 10
async def test_get_operation(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"get_operation"}, include_bank_id=True)
result = await _tools(mcp)["get_operation"].fn(operation_id="op-1")
assert '"op-1"' in result
async def test_cancel_operation(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"cancel_operation"}, include_bank_id=True)
result = await _tools(mcp)["cancel_operation"].fn(operation_id="op-1")
assert '"cancelled"' in result
async def test_list_operations_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_operations"}, include_bank_id=False)
result = await _tools(mcp)["list_operations"].fn()
assert isinstance(result, dict)
# =========================================================================
# Tags & Bank Tool Tests
# =========================================================================
@pytest.mark.asyncio
class TestTagsAndBankTools:
async def test_list_tags(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_tags"}, include_bank_id=True)
result = await _tools(mcp)["list_tags"].fn(q="project:*", limit=50)
call_kwargs = mock_memory.list_tags.call_args.kwargs
assert call_kwargs["pattern"] == "project:*"
assert call_kwargs["limit"] == 50
async def test_get_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"get_bank"}, include_bank_id=True)
result = await _tools(mcp)["get_bank"].fn()
assert '"test-bank"' in result or "test-bank" in result
async def test_get_bank_stats(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"get_bank_stats"}, include_bank_id=True)
result = await _tools(mcp)["get_bank_stats"].fn()
assert "100" in result # nodes count
async def test_update_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"update_bank"}, include_bank_id=True)
result = await _tools(mcp)["update_bank"].fn(name="New Name", mission="New Mission")
call_kwargs = mock_memory.update_bank.call_args.kwargs
assert call_kwargs["name"] == "New Name"
assert call_kwargs["mission"] == "New Mission"
async def test_delete_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"delete_bank"}, include_bank_id=True)
result = await _tools(mcp)["delete_bank"].fn()
assert '"deleted"' in result
mock_memory.delete_bank.assert_called_once()
async def test_clear_memories(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"clear_memories"}, include_bank_id=True)
result = await _tools(mcp)["clear_memories"].fn()
assert '"cleared"' in result
mock_memory.delete_bank.assert_called_once()
async def test_clear_memories_with_type_filter(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"clear_memories"}, include_bank_id=True)
await _tools(mcp)["clear_memories"].fn(type="world")
call_kwargs = mock_memory.delete_bank.call_args.kwargs
assert call_kwargs["fact_type"] == "world"
async def test_list_tags_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"list_tags"}, include_bank_id=False)
result = await _tools(mcp)["list_tags"].fn()
assert isinstance(result, dict)
async def test_get_bank_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"get_bank"}, include_bank_id=False)
result = await _tools(mcp)["get_bank"].fn()
assert isinstance(result, dict)
async def test_delete_bank_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"delete_bank"}, include_bank_id=False)
result = await _tools(mcp)["delete_bank"].fn()
assert isinstance(result, dict)
assert result["status"] == "deleted"
async def test_clear_memories_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"clear_memories"}, include_bank_id=False)
result = await _tools(mcp)["clear_memories"].fn()
assert isinstance(result, dict)
assert result["status"] == "cleared"
# =========================================================================
# Additional Error Handling & Edge Case Tests
# =========================================================================
@pytest.mark.asyncio
class TestOperationErrorHandling:
"""Error handling tests for operation tools."""
async def test_get_operation_engine_error(self, mock_memory):
mock_memory.get_operation_status.side_effect = RuntimeError("Operation not found")
mcp = _make_mcp_server(mock_memory, {"get_operation"}, include_bank_id=True)
result = await _tools(mcp)["get_operation"].fn(operation_id="missing")
assert "error" in result
assert "Operation not found" in result
async def test_get_operation_engine_error_single_bank(self, mock_memory):
mock_memory.get_operation_status.side_effect = RuntimeError("Operation not found")
mcp = _make_mcp_server(mock_memory, {"get_operation"}, include_bank_id=False)
result = await _tools(mcp)["get_operation"].fn(operation_id="missing")
assert isinstance(result, dict)
assert "Operation not found" in result["error"]
async def test_cancel_operation_engine_error(self, mock_memory):
mock_memory.cancel_operation.side_effect = RuntimeError("Cannot cancel completed operation")
mcp = _make_mcp_server(mock_memory, {"cancel_operation"}, include_bank_id=True)
result = await _tools(mcp)["cancel_operation"].fn(operation_id="op-done")
assert "error" in result
assert "Cannot cancel" in result
async def test_cancel_operation_engine_error_single_bank(self, mock_memory):
mock_memory.cancel_operation.side_effect = RuntimeError("Cannot cancel")
mcp = _make_mcp_server(mock_memory, {"cancel_operation"}, include_bank_id=False)
result = await _tools(mcp)["cancel_operation"].fn(operation_id="op-done")
assert isinstance(result, dict)
assert "Cannot cancel" in result["error"]
@pytest.mark.asyncio
class TestDeleteErrorHandling:
"""Error handling tests for delete operations."""
async def test_delete_memory_engine_error(self, mock_memory):
mock_memory.delete_memory_unit.side_effect = RuntimeError("DB error")
mcp = _make_mcp_server(mock_memory, {"delete_memory"}, include_bank_id=True)
result = await _tools(mcp)["delete_memory"].fn(memory_id="mem-1")
assert "error" in result
assert "DB error" in result
async def test_delete_memory_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"delete_memory"}, include_bank_id=False)
result = await _tools(mcp)["delete_memory"].fn(memory_id="mem-1")
assert isinstance(result, dict)
assert result["status"] == "deleted"
async def test_delete_document_engine_error(self, mock_memory):
mock_memory.delete_document.side_effect = RuntimeError("DB error")
mcp = _make_mcp_server(mock_memory, {"delete_document"}, include_bank_id=True)
result = await _tools(mcp)["delete_document"].fn(document_id="doc-1")
assert "error" in result
assert "DB error" in result
async def test_delete_document_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"delete_document"}, include_bank_id=False)
result = await _tools(mcp)["delete_document"].fn(document_id="doc-1")
assert isinstance(result, dict)
assert result["status"] == "deleted"
@pytest.mark.asyncio
class TestUpdateBankVariants:
"""Additional tests for update_bank tool."""
async def test_update_bank_single_bank(self, mock_memory):
mcp = _make_mcp_server(mock_memory, {"update_bank"}, include_bank_id=False)
result = await _tools(mcp)["update_bank"].fn(name="New Name")
assert isinstance(result, dict)
call_kwargs = mock_memory.update_bank.call_args.kwargs
assert call_kwargs["name"] == "New Name"
async def test_update_bank_engine_error(self, mock_memory):
mock_memory.update_bank.side_effect = RuntimeError("DB error")
mcp = _make_mcp_server(mock_memory, {"update_bank"}, include_bank_id=True)
result = await _tools(mcp)["update_bank"].fn(name="X")
assert "error" in result
async def test_get_bank_stats_engine_error(self, mock_memory):
mock_memory.get_bank_stats.side_effect = RuntimeError("DB error")
mcp = _make_mcp_server(mock_memory, {"get_bank_stats"}, include_bank_id=True)
result = await _tools(mcp)["get_bank_stats"].fn()
assert "error" in result
@pytest.mark.asyncio
class TestEmptyListReturns:
"""Tests that empty lists are handled gracefully."""
async def test_list_memories_empty(self, mock_memory):
mock_memory.list_memory_units.return_value = {"items": [], "total": 0}
mcp = _make_mcp_server(mock_memory, {"list_memories"}, include_bank_id=True)
result = await _tools(mcp)["list_memories"].fn()
assert '"items": []' in result or "[]" in result
async def test_list_documents_empty(self, mock_memory):
mock_memory.list_documents.return_value = {"items": [], "total": 0}
mcp = _make_mcp_server(mock_memory, {"list_documents"}, include_bank_id=True)
result = await _tools(mcp)["list_documents"].fn()
assert '"items": []' in result or "[]" in result
async def test_list_operations_empty(self, mock_memory):
mock_memory.list_operations.return_value = {"items": []}
mcp = _make_mcp_server(mock_memory, {"list_operations"}, include_bank_id=True)
result = await _tools(mcp)["list_operations"].fn()
assert '"items": []' in result or "[]" in result
async def test_list_directives_empty(self, mock_memory):
mock_memory.list_directives.return_value = []
mcp = _make_mcp_server(mock_memory, {"list_directives"}, include_bank_id=True)
result = await _tools(mcp)["list_directives"].fn()
assert "[]" in result
async def test_list_tags_empty(self, mock_memory):
mock_memory.list_tags.return_value = {"items": [], "total": 0}
mcp = _make_mcp_server(mock_memory, {"list_tags"}, include_bank_id=True)
result = await _tools(mcp)["list_tags"].fn()
assert '"items": []' in result or "[]" in result
# =========================================================================
# Bank-Level Tool Filtering Tests
# =========================================================================
@pytest.fixture
def mock_memory_with_resolver():
"""Create a mock MemoryEngine with config resolver for bank filtering tests."""
memory = MagicMock()
memory.retain_batch_async = AsyncMock()
memory.recall_async = AsyncMock(
return_value=MagicMock(
model_dump_json=lambda indent=None: '{"results": []}',
model_dump=lambda: {"results": []},
)
)
memory._config_resolver = MagicMock()
memory._config_resolver.get_bank_config = AsyncMock(return_value={})
return memory
class TestBankToolFiltering:
"""Tests for bank-level mcp_enabled_tools filtering via _apply_bank_tool_filtering."""
@pytest.mark.asyncio
async def test_disallowed_tool_raises_error(self, mock_memory_with_resolver):
"""Tool not in bank's mcp_enabled_tools list is hidden from get_tools()."""
from fastmcp import FastMCP
mock_memory_with_resolver._config_resolver.get_bank_config = AsyncMock(
return_value={"mcp_enabled_tools": ["retain"]}
)
mcp = FastMCP("test")
config = MCPToolsConfig(
bank_id_resolver=lambda: "test-bank",
include_bank_id_param=False,
tools={"retain", "recall"},
)
register_mcp_tools(mcp, mock_memory_with_resolver, config)
# Both tools are registered in the manager's internal dict
assert "recall" in mcp._tool_manager._tools
# But get_tools() (used by tools/list and tools/call) filters it out
visible = await mcp._tool_manager.get_tools()
assert "retain" in visible
assert "recall" not in visible
@pytest.mark.asyncio
async def test_allowed_tool_remains_visible(self, mock_memory_with_resolver):
"""Tool in bank's mcp_enabled_tools list stays visible in get_tools()."""
from fastmcp import FastMCP
mock_memory_with_resolver._config_resolver.get_bank_config = AsyncMock(
return_value={"mcp_enabled_tools": ["retain", "recall"]}
)
mcp = FastMCP("test")
config = MCPToolsConfig(
bank_id_resolver=lambda: "test-bank",
include_bank_id_param=False,
tools={"retain", "recall"},
)
register_mcp_tools(mcp, mock_memory_with_resolver, config)
visible = await mcp._tool_manager.get_tools()
assert "retain" in visible
assert "recall" in visible
@pytest.mark.asyncio
async def test_no_filter_when_mcp_enabled_tools_absent(self, mock_memory_with_resolver):
"""When bank config has no mcp_enabled_tools key, all tools remain visible."""
from fastmcp import FastMCP
mock_memory_with_resolver._config_resolver.get_bank_config = AsyncMock(return_value={})
mcp = FastMCP("test")
config = MCPToolsConfig(
bank_id_resolver=lambda: "test-bank",
include_bank_id_param=False,
tools={"retain", "recall"},
)
register_mcp_tools(mcp, mock_memory_with_resolver, config)
visible = await mcp._tool_manager.get_tools()
assert "retain" in visible
assert "recall" in visible
@pytest.mark.asyncio
async def test_filter_skipped_when_no_bank_id(self, mock_memory_with_resolver):
"""When bank_id resolver returns None, config is not fetched and all tools are visible."""
from fastmcp import FastMCP
mock_memory_with_resolver._config_resolver.get_bank_config = AsyncMock(
return_value={"mcp_enabled_tools": ["retain"]} # Would block recall
)
mcp = FastMCP("test")
config = MCPToolsConfig(
bank_id_resolver=lambda: None, # No bank_id context
include_bank_id_param=False,
tools={"retain", "recall"},
)
register_mcp_tools(mcp, mock_memory_with_resolver, config)
visible = await mcp._tool_manager.get_tools()
# Filter bypassed — config resolver was never consulted, all tools visible
assert "recall" in visible
mock_memory_with_resolver._config_resolver.get_bank_config.assert_not_called()
+73 -52
View File
@@ -15,6 +15,10 @@ logger = logging.getLogger(__name__)
@pytest.mark.asyncio
@pytest.mark.xfail(
strict=False,
reason="Gemini sometimes consistently translates Chinese content to English despite instructions",
)
async def test_retain_chinese_content(memory, request_context):
"""
Test that retain correctly extracts facts from Chinese content
@@ -24,70 +28,87 @@ async def test_retain_chinese_content(memory, request_context):
1. Facts are extracted from Chinese text
2. The extracted facts contain Chinese characters
3. Entity names are preserved in Chinese
Note: LLM fact extraction is non-deterministic and may sometimes translate
content to English despite instructions. We retry up to 3 times.
"""
bank_id = f"test_chinese_retain_{datetime.now(timezone.utc).timestamp()}"
max_retries = 3
last_error = None
try:
# Chinese content about a person and their activities
chinese_content = """
张伟是一位资深软件工程师在腾讯工作了五年他专门研究分布式系统
并领导了公司微服务架构的开发他以编写干净文档完善的代码而闻名
for attempt in range(max_retries):
bank_id = f"test_chinese_retain_{datetime.now(timezone.utc).timestamp()}_{attempt}"
李明上个月加入团队担任初级开发人员他正在学习React和Node.js
李明很有热情在代码审查中提出很好的问题他最近完成了他的第一个功能
这是一个用户认证流程
try:
# Chinese content about a person and their activities
chinese_content = """
张伟是一位资深软件工程师在腾讯工作了五年他专门研究分布式系统
并领导了公司微服务架构的开发他以编写干净文档完善的代码而闻名
团队使用Kubernetes进行容器编排并部署到阿里云他们遵循敏捷方法论
采用两周冲刺周期合并前必须进行代码审查
"""
李明上个月加入团队担任初级开发人员他正在学习React和Node.js
李明很有热情在代码审查中提出很好的问题他最近完成了他的第一个功能
这是一个用户认证流程
# Retain the Chinese content
unit_ids = await memory.retain_async(
bank_id=bank_id,
content=chinese_content,
context="团队概述", # Chinese context
event_date=datetime(2024, 1, 15, tzinfo=timezone.utc),
request_context=request_context,
)
团队使用Kubernetes进行容器编排并部署到阿里云他们遵循敏捷方法论
采用两周冲刺周期合并前必须进行代码审查
"""
logger.info(f"Retained {len(unit_ids)} facts from Chinese content")
assert len(unit_ids) > 0, "Should have extracted and stored facts from Chinese content"
# Retain the Chinese content
unit_ids = await memory.retain_async(
bank_id=bank_id,
content=chinese_content,
context="团队概述", # Chinese context
event_date=datetime(2024, 1, 15, tzinfo=timezone.utc),
request_context=request_context,
)
# Recall the facts with a Chinese query
result = await memory.recall_async(
bank_id=bank_id,
query="告诉我关于张伟的信息", # "Tell me about Zhang Wei"
budget=Budget.MID,
max_tokens=1000,
fact_type=["world"],
request_context=request_context,
)
logger.info(f"Retained {len(unit_ids)} facts from Chinese content (attempt {attempt + 1})")
assert len(unit_ids) > 0, "Should have extracted and stored facts from Chinese content"
logger.info(f"Recalled {len(result.results)} facts")
assert len(result.results) > 0, "Should recall facts about Zhang Wei"
# Recall the facts with a Chinese query
result = await memory.recall_async(
bank_id=bank_id,
query="告诉我关于张伟的信息", # "Tell me about Zhang Wei"
budget=Budget.MID,
max_tokens=1000,
fact_type=["world"],
request_context=request_context,
)
# Verify that the facts contain Chinese characters
# At least one fact should mention 张伟 (Zhang Wei) or related Chinese content
chinese_facts_found = 0
for fact in result.results:
logger.info(f"Fact: {fact.text[:100]}...")
# Check for common Chinese characters or the name
if any(
char in fact.text
for char in ["", "", "腾讯", "软件", "工程师", "分布式", "系统", "代码"]
):
chinese_facts_found += 1
logger.info(f"Recalled {len(result.results)} facts")
assert len(result.results) > 0, "Should recall facts about Zhang Wei"
logger.info(f"Found {chinese_facts_found} facts with Chinese content")
assert chinese_facts_found > 0, (
f"Expected facts to contain Chinese characters, but none found. "
f"Facts: {[f.text for f in result.results]}"
)
# Verify that the facts contain Chinese characters
# At least one fact should mention 张伟 (Zhang Wei) or related Chinese content
chinese_facts_found = 0
for fact in result.results:
logger.info(f"Fact: {fact.text[:100]}...")
# Check for common Chinese characters or the name
if any(
char in fact.text
for char in ["", "", "腾讯", "软件", "工程师", "分布式", "系统", "代码"]
):
chinese_facts_found += 1
logger.info("Chinese retain test passed - facts preserved in Chinese")
logger.info(f"Found {chinese_facts_found} facts with Chinese content")
assert chinese_facts_found > 0, (
f"Expected facts to contain Chinese characters, but none found. "
f"Facts: {[f.text for f in result.results]}"
)
finally:
await memory.delete_bank(bank_id, request_context=request_context)
logger.info("Chinese retain test passed - facts preserved in Chinese")
return # Test passed
except AssertionError as e:
last_error = e
if attempt < max_retries - 1:
logger.warning(f"Attempt {attempt + 1} failed: {e}. Retrying...")
else:
raise e
finally:
try:
await memory.delete_bank(bank_id, request_context=request_context)
except Exception:
pass
@pytest.mark.asyncio
@@ -0,0 +1,457 @@
"""
Tests for observation invalidation when source memories are deleted.
These tests verify that:
1. Observations are deleted (not just updated) when their source memories are removed
2. Remaining source memories are reset for re-consolidation (consolidated_at=NULL)
3. The clear_observations_for_memory method correctly clears observations and
resets the target memory itself for re-consolidation
4. delete_bank(fact_type=...) also cleans up affected observations
"""
import uuid
from unittest.mock import AsyncMock, patch
import pytest
from hindsight_api import RequestContext
from hindsight_api.engine.memory_engine import MemoryEngine
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
async def _insert_memory(conn, bank_id: str, text: str, fact_type: str = "experience") -> uuid.UUID:
"""Insert a memory unit directly, bypassing LLM retain pipeline."""
mem_id = uuid.uuid4()
await conn.execute(
"""
INSERT INTO memory_units (id, bank_id, text, fact_type, event_date, created_at, updated_at, consolidated_at)
VALUES ($1, $2, $3, $4, NOW(), NOW(), NOW(), NOW())
""",
mem_id,
bank_id,
text,
fact_type,
)
return mem_id
async def _insert_observation(
conn, bank_id: str, text: str, source_memory_ids: list[uuid.UUID]
) -> uuid.UUID:
"""Insert an observation unit directly."""
obs_id = uuid.uuid4()
await conn.execute(
"""
INSERT INTO memory_units (
id, bank_id, text, fact_type, event_date, source_memory_ids, proof_count, created_at, updated_at
) VALUES ($1, $2, $3, 'observation', NOW(), $4, $5, NOW(), NOW())
""",
obs_id,
bank_id,
text,
source_memory_ids,
len(source_memory_ids),
)
return obs_id
async def _get_observation_ids(conn, bank_id: str) -> list[str]:
rows = await conn.fetch(
"SELECT id FROM memory_units WHERE bank_id = $1 AND fact_type = 'observation'",
bank_id,
)
return [str(r["id"]) for r in rows]
async def _get_consolidated_at(conn, memory_id: uuid.UUID):
return await conn.fetchval(
"SELECT consolidated_at FROM memory_units WHERE id = $1",
memory_id,
)
async def _ensure_bank(memory: MemoryEngine, bank_id: str, request_context: RequestContext):
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
# ---------------------------------------------------------------------------
# Tests: delete_memory_unit
# ---------------------------------------------------------------------------
class TestDeleteMemoryUnitObservationCleanup:
@pytest.mark.asyncio
async def test_deleting_source_memory_removes_observation(
self, memory: MemoryEngine, request_context: RequestContext
):
"""Deleting a source memory removes observations derived from it."""
bank_id = f"test-invalidate-del-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
m1 = await _insert_memory(conn, bank_id, "Alice loves hiking.")
m2 = await _insert_memory(conn, bank_id, "Alice goes hiking every weekend.")
obs_id = await _insert_observation(conn, bank_id, "Alice enjoys hiking regularly.", [m1, m2])
await memory.delete_memory_unit(str(m1), request_context=request_context)
async with pool.acquire() as conn:
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs_id) not in obs_ids, "Observation should have been deleted"
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_deleting_source_memory_resets_remaining_source_consolidated_at(
self, memory: MemoryEngine, request_context: RequestContext
):
"""After deleting a source memory, remaining source memories are reset for re-consolidation."""
bank_id = f"test-invalidate-reset-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
m1 = await _insert_memory(conn, bank_id, "Alice loves hiking.")
m2 = await _insert_memory(conn, bank_id, "Alice goes hiking every weekend.")
await _insert_observation(conn, bank_id, "Alice enjoys hiking regularly.", [m1, m2])
# Verify m2 starts with consolidated_at set
assert await _get_consolidated_at(conn, m2) is not None
# Patch out consolidation so it doesn't re-set consolidated_at before we can check it
with patch.object(memory, "submit_async_consolidation", new=AsyncMock()):
await memory.delete_memory_unit(str(m1), request_context=request_context)
async with pool.acquire() as conn:
# m2 should have consolidated_at reset to NULL
consolidated_at = await _get_consolidated_at(conn, m2)
assert consolidated_at is None, "Remaining source memory should be reset for re-consolidation"
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_deleting_non_source_memory_leaves_observations_intact(
self, memory: MemoryEngine, request_context: RequestContext
):
"""Deleting a memory that is not a source of any observation leaves observations unchanged."""
bank_id = f"test-invalidate-noop-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
m1 = await _insert_memory(conn, bank_id, "Alice loves hiking.")
m2 = await _insert_memory(conn, bank_id, "Alice goes hiking every weekend.")
unrelated = await _insert_memory(conn, bank_id, "Bob likes cycling.")
obs_id = await _insert_observation(conn, bank_id, "Alice enjoys hiking regularly.", [m1, m2])
await memory.delete_memory_unit(str(unrelated), request_context=request_context)
async with pool.acquire() as conn:
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs_id) in obs_ids, "Observation should remain untouched"
# m1 and m2 should still be consolidated
assert await _get_consolidated_at(conn, m1) is not None
assert await _get_consolidated_at(conn, m2) is not None
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_deleting_sole_source_memory_removes_observation_no_remaining_reset(
self, memory: MemoryEngine, request_context: RequestContext
):
"""When an observation has only one source and it's deleted, observation is removed with no remaining memories to reset."""
bank_id = f"test-invalidate-sole-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
m1 = await _insert_memory(conn, bank_id, "Alice loves hiking.")
obs_id = await _insert_observation(conn, bank_id, "Alice enjoys hiking.", [m1])
await memory.delete_memory_unit(str(m1), request_context=request_context)
async with pool.acquire() as conn:
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs_id) not in obs_ids, "Observation should have been deleted"
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_deleting_observation_type_memory_does_not_trigger_invalidation(
self, memory: MemoryEngine, request_context: RequestContext
):
"""Deleting a memory with fact_type='observation' directly does not trigger invalidation logic."""
bank_id = f"test-invalidate-obstype-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
m1 = await _insert_memory(conn, bank_id, "Alice loves hiking.")
obs_id = await _insert_observation(conn, bank_id, "Alice enjoys hiking.", [m1])
# Delete the observation directly (not the source memory)
await memory.delete_memory_unit(str(obs_id), request_context=request_context)
async with pool.acquire() as conn:
# Source memory should still be consolidated (not reset)
assert await _get_consolidated_at(conn, m1) is not None
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs_id) not in obs_ids
await memory.delete_bank(bank_id, request_context=request_context)
# ---------------------------------------------------------------------------
# Tests: delete_document
# ---------------------------------------------------------------------------
class TestDeleteDocumentObservationCleanup:
@pytest.mark.asyncio
async def test_deleting_document_removes_observations(
self, memory: MemoryEngine, request_context: RequestContext
):
"""Deleting a document removes observations derived from its memory units."""
bank_id = f"test-invalidate-doc-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
# Create a document and attach memories to it
async with pool.acquire() as conn:
doc_id = str(uuid.uuid4()) # documents.id is TEXT
await conn.execute(
"""
INSERT INTO documents (id, bank_id, original_text, content_hash, created_at, updated_at)
VALUES ($1, $2, 'some doc', 'hash123', NOW(), NOW())
""",
doc_id,
bank_id,
)
m1 = uuid.uuid4()
m2 = uuid.uuid4()
for mem_id, text in [(m1, "Alice loves hiking."), (m2, "Alice goes hiking every weekend.")]:
await conn.execute(
"""
INSERT INTO memory_units (id, bank_id, text, fact_type, event_date, document_id, created_at, updated_at, consolidated_at)
VALUES ($1, $2, $3, 'experience', NOW(), $4, NOW(), NOW(), NOW())
""",
mem_id,
bank_id,
text,
doc_id,
)
# Standalone memory (not in document)
m3 = await _insert_memory(conn, bank_id, "Alice is an avid outdoor person.")
# Observation referencing both doc memories and the standalone memory
obs_id = await _insert_observation(
conn, bank_id, "Alice enjoys outdoor activities.", [m1, m2, m3]
)
# Patch out consolidation so it doesn't re-set consolidated_at before we can check it
with patch.object(memory, "submit_async_consolidation", new=AsyncMock()):
await memory.delete_document(str(doc_id), bank_id, request_context=request_context)
async with pool.acquire() as conn:
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs_id) not in obs_ids, "Observation should have been deleted"
# m3 (remaining source) should be reset for re-consolidation
consolidated_at = await _get_consolidated_at(conn, m3)
assert consolidated_at is None, "Remaining source memory should be reset"
await memory.delete_bank(bank_id, request_context=request_context)
# ---------------------------------------------------------------------------
# Tests: delete_bank with fact_type filter
# ---------------------------------------------------------------------------
class TestDeleteBankByTypeObservationCleanup:
@pytest.mark.asyncio
async def test_clearing_experience_memories_removes_affected_observations(
self, memory: MemoryEngine, request_context: RequestContext
):
"""Clearing all experience memories removes observations sourced from them."""
bank_id = f"test-invalidate-banktype-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
exp1 = await _insert_memory(conn, bank_id, "Alice went hiking last week.", "experience")
world1 = await _insert_memory(conn, bank_id, "Alice is a hiker.", "world")
obs_id = await _insert_observation(
conn, bank_id, "Alice is a regular hiker.", [exp1, world1]
)
# Patch out consolidation so it doesn't re-set consolidated_at before we can check it
with patch.object(memory, "submit_async_consolidation", new=AsyncMock()):
await memory.delete_bank(bank_id, fact_type="experience", request_context=request_context)
async with pool.acquire() as conn:
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs_id) not in obs_ids, "Observation should have been deleted"
# world1 (remaining source) should be reset for re-consolidation
consolidated_at = await _get_consolidated_at(conn, world1)
assert consolidated_at is None, "World memory should be reset for re-consolidation"
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_clearing_unrelated_type_leaves_observations_intact(
self, memory: MemoryEngine, request_context: RequestContext
):
"""Clearing memories of a type that is not a source of any observation leaves observations untouched."""
bank_id = f"test-invalidate-banktype-noop-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
world1 = await _insert_memory(conn, bank_id, "Alice is a hiker.", "world")
obs_id = await _insert_observation(conn, bank_id, "Alice is a regular hiker.", [world1])
# Deleting 'experience' type should not affect observations sourced only from 'world'
await memory.delete_bank(bank_id, fact_type="experience", request_context=request_context)
async with pool.acquire() as conn:
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs_id) in obs_ids, "Observation should remain untouched"
await memory.delete_bank(bank_id, request_context=request_context)
# ---------------------------------------------------------------------------
# Tests: clear_observations_for_memory
# ---------------------------------------------------------------------------
class TestClearObservationsForMemory:
@pytest.mark.asyncio
async def test_clears_observations_and_resets_all_source_memories(
self, memory: MemoryEngine, request_context: RequestContext
):
"""Clearing observations for a memory deletes them and resets all related source memories."""
bank_id = f"test-clear-obs-mem-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
m1 = await _insert_memory(conn, bank_id, "Alice loves hiking.")
m2 = await _insert_memory(conn, bank_id, "Alice hikes every weekend.")
obs_id = await _insert_observation(conn, bank_id, "Alice is an avid hiker.", [m1, m2])
# Patch out consolidation so it doesn't re-set consolidated_at before we can check it
with patch.object(memory, "submit_async_consolidation", new=AsyncMock()):
result = await memory.clear_observations_for_memory(
bank_id, str(m1), request_context=request_context
)
assert result["deleted_count"] == 1
async with pool.acquire() as conn:
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs_id) not in obs_ids, "Observation should be deleted"
# Both m1 (target) and m2 (remaining source) should be reset
assert await _get_consolidated_at(conn, m1) is None, "Target memory should be reset"
assert await _get_consolidated_at(conn, m2) is None, "Remaining source should be reset"
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_no_observations_returns_zero(
self, memory: MemoryEngine, request_context: RequestContext
):
"""Returns 0 when the memory has no associated observations."""
bank_id = f"test-clear-obs-noop-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
m1 = await _insert_memory(conn, bank_id, "Alice loves hiking.")
result = await memory.clear_observations_for_memory(
bank_id, str(m1), request_context=request_context
)
assert result["deleted_count"] == 0
async with pool.acquire() as conn:
# Memory should still be consolidated (no observations were cleared)
assert await _get_consolidated_at(conn, m1) is not None
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_only_clears_observations_referencing_target_memory(
self, memory: MemoryEngine, request_context: RequestContext
):
"""Clearing observations for m1 does not affect observations that only reference m2."""
bank_id = f"test-clear-obs-selective-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
m1 = await _insert_memory(conn, bank_id, "Alice loves hiking.")
m2 = await _insert_memory(conn, bank_id, "Alice hikes every weekend.")
m3 = await _insert_memory(conn, bank_id, "Alice climbed a mountain.")
obs1_id = await _insert_observation(conn, bank_id, "Alice is an avid hiker.", [m1, m2])
obs2_id = await _insert_observation(conn, bank_id, "Alice is a mountaineer.", [m3])
result = await memory.clear_observations_for_memory(
bank_id, str(m1), request_context=request_context
)
assert result["deleted_count"] == 1
async with pool.acquire() as conn:
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs1_id) not in obs_ids, "obs1 (references m1) should be deleted"
assert str(obs2_id) in obs_ids, "obs2 (does not reference m1) should remain"
# m3 should still be consolidated
assert await _get_consolidated_at(conn, m3) is not None
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_multiple_observations_for_same_memory_all_cleared(
self, memory: MemoryEngine, request_context: RequestContext
):
"""All observations referencing the target memory are cleared in one call."""
bank_id = f"test-clear-obs-multi-{uuid.uuid4().hex[:8]}"
await _ensure_bank(memory, bank_id, request_context)
pool = await memory._get_pool()
async with pool.acquire() as conn:
m1 = await _insert_memory(conn, bank_id, "Alice loves hiking.")
m2 = await _insert_memory(conn, bank_id, "Alice hikes every weekend.")
obs1_id = await _insert_observation(conn, bank_id, "Alice hikes often.", [m1])
obs2_id = await _insert_observation(conn, bank_id, "Alice is outdoorsy.", [m1, m2])
# Patch out consolidation so it doesn't re-set consolidated_at before we can check it
with patch.object(memory, "submit_async_consolidation", new=AsyncMock()):
result = await memory.clear_observations_for_memory(
bank_id, str(m1), request_context=request_context
)
assert result["deleted_count"] == 2
async with pool.acquire() as conn:
obs_ids = await _get_observation_ids(conn, bank_id)
assert str(obs1_id) not in obs_ids
assert str(obs2_id) not in obs_ids
# m1 and m2 should both be reset
assert await _get_consolidated_at(conn, m1) is None
assert await _get_consolidated_at(conn, m2) is None
await memory.delete_bank(bank_id, request_context=request_context)
+196 -3
View File
@@ -7,14 +7,17 @@ These tests verify:
3. Recovery from tool execution errors
"""
from unittest.mock import AsyncMock, MagicMock
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from hindsight_api.engine.reflect.agent import (
_normalize_tool_name,
_is_done_tool,
_clean_answer_text,
_clean_done_answer,
_count_messages_tokens,
_is_context_overflow_error,
_is_done_tool,
_normalize_tool_name,
run_reflect_agent,
)
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
@@ -412,3 +415,193 @@ class TestReflectAgentMocked:
# Should have a result even if no memories found
assert result is not None
assert result.iterations == 3
class TestContextOverflowHelpers:
"""Unit tests for context-overflow detection helpers."""
def test_count_messages_tokens_basic(self):
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "What is the capital of France?"},
]
count = _count_messages_tokens(messages)
assert count > 0
# Rough sanity check: ~10 tokens for each message
assert count < 100
def test_count_messages_tokens_with_tool_result(self):
"""A large tool result should substantially increase the count."""
small_messages = [{"role": "user", "content": "hi"}]
large_messages = [
{"role": "user", "content": "hi"},
{
"role": "tool",
"tool_call_id": "x",
"name": "recall",
"content": '{"memories": [' + ', '.join([f'{{"id": "m{i}", "content": "A long memory fact about some topic that goes on and on."}}' for i in range(50)]) + ']}',
},
]
small = _count_messages_tokens(small_messages)
large = _count_messages_tokens(large_messages)
assert large > small + 200
def test_is_context_overflow_error_openai(self):
assert _is_context_overflow_error(Exception("context_length_exceeded: too many tokens"))
assert _is_context_overflow_error(Exception("This model's maximum context length is 128000 tokens. However, your messages resulted in 142164 tokens."))
def test_is_context_overflow_error_anthropic(self):
assert _is_context_overflow_error(Exception("prompt_too_long"))
assert _is_context_overflow_error(Exception("prompt is too long for this model"))
def test_is_context_overflow_error_gemini(self):
assert _is_context_overflow_error(Exception("RESOURCE_EXHAUSTED: quota exceeded"))
def test_is_context_overflow_error_generic(self):
assert _is_context_overflow_error(Exception("input is too long to process"))
assert _is_context_overflow_error(Exception("too many tokens in the request"))
def test_is_context_overflow_error_unrelated(self):
assert not _is_context_overflow_error(Exception("connection timeout"))
assert not _is_context_overflow_error(Exception("rate limit exceeded"))
assert not _is_context_overflow_error(ValueError("invalid argument"))
class TestContextOverflowBehavior:
"""Test that the reflect agent handles context overflow gracefully."""
@pytest.fixture
def mock_llm(self):
llm = MagicMock()
llm.call_with_tools = AsyncMock()
llm.call = AsyncMock(
return_value=("Synthesized answer from gathered evidence.", TokenUsage(input_tokens=50, output_tokens=20, total_tokens=70))
)
return llm
@pytest.fixture
def mock_functions_with_large_output(self):
"""Mock functions that return a large enough payload to exceed a tiny token budget."""
large_memories = [
{"id": f"mem-{i}", "content": f"Memory fact number {i}: " + "A" * 200}
for i in range(20)
]
return {
"search_mental_models_fn": AsyncMock(return_value={"mental_models": []}),
"search_observations_fn": AsyncMock(return_value={"observations": []}),
"recall_fn": AsyncMock(return_value={"memories": large_memories}),
"expand_fn": AsyncMock(return_value={"memories": []}),
}
@pytest.mark.asyncio
async def test_proactive_guard_fires_when_budget_exceeded(self, mock_llm, mock_functions_with_large_output):
"""When token count exceeds max_context_tokens after a tool call, the agent
should immediately synthesize from gathered evidence instead of making
another LLM call that would overflow."""
# First call: LLM calls recall (forced by iter 0 with no mental models)
mock_llm.call_with_tools.return_value = LLMToolCallResult(
tool_calls=[LLMToolCall(id="1", name="recall", arguments={"query": "test"})],
finish_reason="tool_calls",
)
# Set a tiny token budget — the recall result alone will blow past it
result = await run_reflect_agent(
llm_config=mock_llm,
bank_id="test-bank",
query="What do you know?",
bank_profile={"name": "Test", "mission": "Testing"},
max_context_tokens=100,
**mock_functions_with_large_output,
)
assert result.text == "Synthesized answer from gathered evidence."
# call_with_tools was called once (for the forced recall), then the guard
# kicked in — no further tool-call iterations
assert mock_llm.call_with_tools.call_count == 1
# llm.call() was invoked to generate the final synthesis
mock_llm.call.assert_called_once()
@pytest.mark.asyncio
async def test_context_overflow_error_skips_retry(self, mock_llm, mock_functions_with_large_output):
"""A context_length_exceeded error from the LLM should NOT be retried —
it should immediately fall back to final synthesis."""
mock_llm.call_with_tools.side_effect = Exception(
"context_length_exceeded: messages resulted in 150000 tokens."
)
result = await run_reflect_agent(
llm_config=mock_llm,
bank_id="test-bank",
query="What do you know?",
bank_profile={"name": "Test", "mission": "Testing"},
max_iterations=5,
**mock_functions_with_large_output,
)
assert result is not None
# Should have attempted only 1 iteration (no retry on overflow error)
assert mock_llm.call_with_tools.call_count == 1
# Final synthesis was called
mock_llm.call.assert_called_once()
class TestContextOverflowIntegration:
"""Integration test: real LLM with a very small max_context_tokens.
The agent will make one real LLM call (forced tool choice), receive a large
tool result that exceeds the tiny budget, then synthesize from it via a second
real LLM call all without raising a context_length_exceeded error.
"""
@pytest.mark.asyncio
async def test_reflect_completes_with_tiny_context_budget(self, memory, request_context):
"""End-to-end: reflect on a bank with max_context_tokens=1 (tiny budget).
Setting max_context_tokens=1 guarantees the proactive guard fires as soon
as the first tool result is received and evidence is available.
The result must be a non-empty string with no exception raised.
"""
import uuid
from unittest.mock import patch
bank_id = f"test-ctx-overflow-{uuid.uuid4().hex[:8]}"
try:
# Retain a handful of facts so the recall tool has something to return
await memory.retain_async(
bank_id=bank_id,
content="Alice is a software engineer who enjoys hiking on weekends.",
request_context=request_context,
)
await memory.retain_async(
bank_id=bank_id,
content="Bob is a designer who loves cooking Italian food.",
request_context=request_context,
)
# Patch get_config where memory_engine uses it, injecting a tiny
# max_context_tokens. Everything else delegates to the real config.
real_config = memory._get_raw_config() if hasattr(memory, "_get_raw_config") else None
from hindsight_api.config import get_config as _real_get_config
class _TinyContextProxy:
"""Forwards all attribute access to the real config proxy except
reflect_max_context_tokens which is forced to 1."""
_real = _real_get_config()
def __getattr__(self, name: str):
if name == "reflect_max_context_tokens":
return 1
return getattr(self._real, name)
with patch("hindsight_api.engine.memory_engine.get_config", return_value=_TinyContextProxy()):
result = await memory.reflect_async(
bank_id=bank_id,
query="Tell me about the people you know.",
request_context=request_context,
)
assert result.text, "reflect must return a non-empty answer"
assert result.usage.total_tokens > 0
finally:
await memory.delete_bank(bank_id, request_context=request_context)
+180
View File
@@ -591,6 +591,125 @@ async def test_mentioned_at_from_context_string(memory, request_context):
await memory.delete_bank(bank_id, request_context=request_context)
# ============================================================
# No Timestamp Tests
# ============================================================
@pytest.mark.asyncio
async def test_retain_no_timestamp(memory, request_context):
"""
Test retaining content with explicit "no timestamp" sentinel.
When event_date=None is passed explicitly in the dict (i.e. caller opted into
no timestamp), mentioned_at should be NULL in the DB rather than defaulting to now().
"""
bank_id = f"test_no_timestamp_{datetime.now(timezone.utc).timestamp()}"
try:
# Use retain_batch_async with explicit event_date=None key to signal "no timestamp"
unit_ids_list = await memory.retain_batch_async(
bank_id=bank_id,
contents=[
{
"content": "The capital of France is Paris. The Eiffel Tower is located in Paris.",
"context": "general knowledge",
"event_date": None, # Explicit sentinel: no timestamp
}
],
request_context=request_context,
)
assert len(unit_ids_list) > 0, "Should create at least one batch result"
unit_ids = unit_ids_list[0]
assert len(unit_ids) > 0, "Should have extracted and stored facts"
# Recall the facts
result = await memory.recall_async(
bank_id=bank_id,
query="Where is the Eiffel Tower?",
budget=Budget.LOW,
max_tokens=500,
fact_type=["world"],
request_context=request_context,
)
assert len(result.results) > 0, "Should recall the stored fact"
# All temporal fields should be None for temporally agnostic content
for fact in result.results:
assert fact.mentioned_at is None, (
f"mentioned_at should be None for no-timestamp content, got {fact.mentioned_at}"
)
print(f"\n✓ Test passed: mentioned_at is None for {len(result.results)} fact(s)")
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_retain_omit_timestamp_defaults_to_now(memory, request_context):
"""
Backward-compatibility regression test: omitting event_date still stores a real datetime.
When event_date is absent from the content dict (key not present), the orchestrator
should default to utcnow() preserving existing behavior.
"""
bank_id = f"test_default_timestamp_{datetime.now(timezone.utc).timestamp()}"
before = datetime.now(timezone.utc)
try:
# Omit event_date entirely — should default to now()
unit_ids_list = await memory.retain_batch_async(
bank_id=bank_id,
contents=[
{
"content": "Alice is a software engineer who loves Python.",
"context": "profile",
# event_date intentionally omitted
}
],
request_context=request_context,
)
after = datetime.now(timezone.utc)
assert len(unit_ids_list) > 0
unit_ids = unit_ids_list[0]
assert len(unit_ids) > 0, "Should have extracted and stored facts"
# Recall and verify mentioned_at is a real datetime close to now
result = await memory.recall_async(
bank_id=bank_id,
query="Who is Alice?",
budget=Budget.LOW,
max_tokens=500,
fact_type=["world"],
request_context=request_context,
)
assert len(result.results) > 0, "Should recall the fact"
fact = result.results[0]
assert fact.mentioned_at is not None, "mentioned_at should be set when event_date is omitted"
if isinstance(fact.mentioned_at, str):
mentioned_dt = datetime.fromisoformat(fact.mentioned_at.replace("Z", "+00:00"))
else:
mentioned_dt = fact.mentioned_at
# Should be within 60s of when we ran the test
assert before <= mentioned_dt <= after + timedelta(seconds=60), (
f"mentioned_at {mentioned_dt} should be close to now ({before} {after})"
)
print(f"\n✓ Test passed: mentioned_at={mentioned_dt} is a real datetime (backward compat)")
finally:
await memory.delete_bank(bank_id, request_context=request_context)
# ============================================================
# Context Tracking Tests
# ============================================================
@@ -2256,3 +2375,64 @@ async def test_retain_batch_with_per_item_tags_on_document(memory, request_conte
finally:
await memory.delete_bank(bank_id, request_context=request_context)
print(f"\n=== Cleaned up bank: {bank_id} ===")
def test_retain_mission_injected_into_prompt():
"""Test that retain_mission is injected as a FOCUS section into any extraction mode."""
from unittest.mock import MagicMock
from hindsight_api.engine.retain.fact_extraction import _build_extraction_prompt_and_schema
spec = "Focus on technical decisions and architecture choices only."
# Test with concise mode
config = MagicMock()
config.retain_extraction_mode = "concise"
config.retain_mission = spec
config.retain_custom_instructions = None
config.retain_extract_causal_links = False
prompt, _ = _build_extraction_prompt_and_schema(config)
assert spec in prompt
assert "FOCUS" in prompt
# retain_mission is present regardless of extraction mode (verbose has its own template, no spec injection)
config.retain_extraction_mode = "verbose"
prompt_verbose, _ = _build_extraction_prompt_and_schema(config)
# verbose uses its own template - spec not injected there
assert spec not in prompt_verbose
def test_retain_mission_absent_when_not_set():
"""Test that no FOCUS section appears when retain_mission is not set."""
from unittest.mock import MagicMock
from hindsight_api.engine.retain.fact_extraction import _build_extraction_prompt_and_schema
config = MagicMock()
config.retain_extraction_mode = "concise"
config.retain_mission = None
config.retain_custom_instructions = None
config.retain_extract_causal_links = False
prompt, _ = _build_extraction_prompt_and_schema(config)
assert "FOCUS" not in prompt
assert "retain_mission_section" not in prompt
def test_retain_mission_config_loaded_from_env():
"""Test that retain_mission is loaded from env and is a configurable field."""
import os
from hindsight_api.config import HindsightConfig, _get_raw_config, clear_config_cache
original = os.getenv("HINDSIGHT_API_RETAIN_MISSION")
try:
os.environ["HINDSIGHT_API_RETAIN_MISSION"] = "Only technical decisions."
clear_config_cache()
config = _get_raw_config()
assert config.retain_mission == "Only technical decisions."
assert "retain_mission" in HindsightConfig.get_configurable_fields()
finally:
if original is None:
os.environ.pop("HINDSIGHT_API_RETAIN_MISSION", None)
else:
os.environ["HINDSIGHT_API_RETAIN_MISSION"] = original
clear_config_cache()
@@ -233,6 +233,35 @@ class TestFilterResultsByTags:
assert len(filtered) == 1
assert filtered[0].tags == ["a", "b", "c"] # Has a, b, AND c
def test_all_strict_superset_observation_matches_incoming_memory_tags(self):
"""
Consolidation scenario: an incoming memory with tags ['user:bob', 'session:id1']
uses all_strict matching to find existing observations.
An observation tagged ['user:bob', 'session:id1', 'place:online'] IS matched
because it contains all of the incoming memory's tags (superset).
This is NOT exact matching an observation with extra tags is still a valid match.
"""
# Incoming memory tags (e.g. from a new retain call)
incoming_tags = ["user:bob", "session:id1"]
# Candidate observations with different tag sets
exact_match = MockResult(["user:bob", "session:id1"])
superset_match = MockResult(["session:id1", "user:bob", "place:online"])
different_user = MockResult(["user:alice", "session:id1"])
missing_session = MockResult(["user:bob"])
results = [exact_match, superset_match, different_user, missing_session]
filtered = filter_results_by_tags(results, incoming_tags, match="all_strict")
# Both exact_match and superset_match have all incoming tags → both match
assert len(filtered) == 2
assert exact_match in filtered
assert superset_match in filtered
# different_user and missing_session are excluded because they lack at least one tag
assert different_user not in filtered
assert missing_session not in filtered
# ============================================================================
# Integration Tests for tags in retain/recall/reflect
@@ -890,3 +919,40 @@ async def test_list_tags_ordered_by_count(api_client):
# common (3) should come before medium (2) which should come before rare (1)
assert tags.index("common") < tags.index("medium")
assert tags.index("medium") < tags.index("rare")
@pytest.mark.asyncio
async def test_list_memories_includes_tags(api_client, test_bank_id):
"""Test that list memories endpoint returns tags for each memory unit.
Regression test: tags were previously omitted from the SELECT query in
list_memory_units, causing the memory dialog in the UI to show no tags
even when memories had been stored with tags.
"""
tags = ["user_alice", "session_xyz", "project_alpha", "team_eng", "env_prod", "region_us"]
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{
"content": "Alice is a senior engineer on the platform team.",
"tags": tags,
}
]
},
)
assert response.status_code == 200
# List memories and verify all tags are returned
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/memories/list")
assert response.status_code == 200
result = response.json()
assert result["total"] > 0
memory_item = next((item for item in result["items"] if "Alice" in item["text"]), None)
assert memory_item is not None, "Should find the stored memory"
assert "tags" in memory_item, "Memory item must include a 'tags' field"
assert set(memory_item["tags"]) == set(tags), (
f"All {len(tags)} tags should be returned, got: {memory_item['tags']}"
)
+84 -25
View File
@@ -48,11 +48,11 @@ async def pool(pg0_db_url):
@pytest_asyncio.fixture
async def clean_operations(pool):
"""Clean up async_operations table before and after tests."""
# Clean before test
await pool.execute("DELETE FROM async_operations WHERE bank_id LIKE 'test-worker-%'")
# Clean before test - covers both 'test-worker-' and 'test_worker_recovery' patterns
await pool.execute("DELETE FROM async_operations WHERE bank_id LIKE 'test-worker-%' OR bank_id LIKE 'test_worker_%'")
yield
# Clean after test
await pool.execute("DELETE FROM async_operations WHERE bank_id LIKE 'test-worker-%'")
await pool.execute("DELETE FROM async_operations WHERE bank_id LIKE 'test-worker-%' OR bank_id LIKE 'test_worker_%'")
class TestBrokerTaskBackend:
@@ -268,24 +268,26 @@ class TestWorkerPoller:
assert row["completed_at"] is not None
@pytest.mark.asyncio
async def test_executor_exception_does_not_crash_poller(self, pool, clean_operations):
"""Test that unexpected exceptions from executor are caught and don't crash the poller.
async def test_executor_exception_triggers_retry(self, pool, clean_operations):
"""Test that exceptions from the executor trigger _retry_or_fail (not a crash).
If the executor raises an unexpected exception (which MemoryEngine.execute_task should NOT do,
but could happen from schema setup or other infrastructure issues), the poller should catch it
gracefully. Status remains 'processing' since neither executor nor poller handled it.
When the executor re-raises an exception (as MemoryEngine.execute_task does for
retryable task failures), the poller calls _retry_or_fail, which resets the task
back to 'pending' and increments retry_count so it can be reclaimed.
This is the fix for the consolidation deadlock: previously submit_task was called
with only a task_payload update, leaving status='processing' forever.
"""
from hindsight_api.worker import WorkerPoller
from hindsight_api.worker.poller import ClaimedTask
# Create a pending task
bank_id = f"test-worker-{uuid.uuid4().hex[:8]}"
op_id = uuid.uuid4()
payload = json.dumps({"type": "test_task", "operation_id": str(op_id), "bank_id": bank_id})
payload = json.dumps({"type": "consolidation", "operation_id": str(op_id), "bank_id": bank_id})
await pool.execute(
"""
INSERT INTO async_operations (operation_id, bank_id, operation_type, status, task_payload, worker_id)
VALUES ($1, $2, 'test', 'processing', $3::jsonb, 'test-worker-1')
INSERT INTO async_operations (operation_id, bank_id, operation_type, status, task_payload, worker_id, claimed_at)
VALUES ($1, $2, 'consolidation', 'processing', $3::jsonb, 'test-worker-1', now())
""",
op_id,
bank_id,
@@ -293,43 +295,100 @@ class TestWorkerPoller:
)
async def failing_executor(task_dict):
raise ValueError("Unexpected infrastructure failure")
raise ValueError("TimeoutError during recall")
poller = WorkerPoller(
pool=pool,
worker_id="test-worker-1",
executor=failing_executor,
max_retries=3,
)
# Execute - should catch exception without crashing
task_dict = json.loads(payload)
claimed_task = ClaimedTask(operation_id=str(op_id), task_dict=task_dict, schema=None)
await poller.execute_task(claimed_task)
# Wait for background task to complete
completed = await poller.wait_for_active_tasks(timeout=5.0)
assert completed, "Task did not complete within timeout"
# Status stays 'processing' since the poller no longer manages status
# Task must be reset to 'pending' with worker_id/claimed_at cleared — not left as
# 'processing', which would cause a permanent deadlock via the NOT EXISTS guard.
row = await pool.fetchrow(
"SELECT status FROM async_operations WHERE operation_id = $1",
"SELECT status, worker_id, claimed_at, retry_count FROM async_operations WHERE operation_id = $1",
op_id,
)
assert row["status"] == "processing"
assert row["status"] == "pending", (
f"REGRESSION: Task status is '{row['status']}' instead of 'pending'. "
"A task stuck in 'processing' after a retry causes a consolidation deadlock."
)
assert row["worker_id"] is None, "worker_id must be cleared on retry"
assert row["claimed_at"] is None, "claimed_at must be cleared on retry"
assert row["retry_count"] == 1
@pytest.mark.asyncio
async def test_executor_exception_marks_failed_after_max_retries(self, pool, clean_operations):
"""Test that a task is permanently marked 'failed' once retry_count hits max_retries.
After max_retries exhaustion the task must NOT be reset to 'pending' it should
be marked 'failed' with an error message so it stops consuming retry budget.
"""
from hindsight_api.worker import WorkerPoller
from hindsight_api.worker.poller import ClaimedTask
max_retries = 3
bank_id = f"test-worker-{uuid.uuid4().hex[:8]}"
op_id = uuid.uuid4()
payload = json.dumps({"type": "consolidation", "operation_id": str(op_id), "bank_id": bank_id})
# Insert with retry_count already at the limit
await pool.execute(
"""
INSERT INTO async_operations (operation_id, bank_id, operation_type, status, task_payload, worker_id, claimed_at, retry_count)
VALUES ($1, $2, 'consolidation', 'processing', $3::jsonb, 'test-worker-1', now(), $4)
""",
op_id,
bank_id,
payload,
max_retries,
)
async def failing_executor(task_dict):
raise ValueError("Still failing after all retries")
poller = WorkerPoller(
pool=pool,
worker_id="test-worker-1",
executor=failing_executor,
max_retries=max_retries,
)
task_dict = json.loads(payload)
claimed_task = ClaimedTask(operation_id=str(op_id), task_dict=task_dict, schema=None)
await poller.execute_task(claimed_task)
completed = await poller.wait_for_active_tasks(timeout=5.0)
assert completed, "Task did not complete within timeout"
row = await pool.fetchrow(
"SELECT status, error_message, retry_count FROM async_operations WHERE operation_id = $1",
op_id,
)
assert row["status"] == "failed", (
f"Expected 'failed' after max retries, got '{row['status']}'"
)
assert row["error_message"] is not None
assert "Max retries" in row["error_message"]
assert row["retry_count"] == max_retries # not incremented further
@pytest.mark.asyncio
async def test_executor_failed_status_not_overridden(self, pool, clean_operations):
"""REGRESSION TEST: Verify poller does NOT overwrite executor's 'failed' status to 'completed'.
This test catches the bug where the poller always called _mark_completed() after executor
returned, overwriting the 'failed' status that the executor had already set.
Scenario:
1. Executor catches an internal error and marks the operation as 'failed' in the DB
2. Executor returns normally (does NOT re-raise) - this is how MemoryEngine.execute_task works
This test covers the non-retryable failure path (e.g., file_convert_retain):
1. Executor catches an internal error, marks the operation as 'failed' in the DB
2. Executor returns normally (does NOT re-raise) so no exception reaches the poller
3. The poller must NOT overwrite the 'failed' status to 'completed'
With the old buggy code, this test would FAIL (status would be 'completed').
Retryable failures re-raise instead (see test_executor_exception_triggers_retry).
"""
from hindsight_api.worker import WorkerPoller
from hindsight_api.worker.poller import ClaimedTask
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "hindsight-cli"
version = "0.4.13"
version = "0.4.15"
edition = "2021"
authors = ["Hindsight Team"]
description = "A beautiful CLI for Hindsight - semantic memory system"
+5 -1
View File
@@ -127,6 +127,7 @@ impl ApiClient {
mission: None,
background: None,
disposition: None,
..Default::default()
};
let response = self.client.create_or_update_bank(agent_id, None, &request).await?;
Ok(response.into_inner())
@@ -299,6 +300,8 @@ impl ApiClient {
offset.map(|o| o as i64),
q,
None,
None,
None,
).await?;
Ok(response.into_inner())
})
@@ -436,6 +439,7 @@ impl ApiClient {
mission: Some(mission.to_string()),
background: None,
disposition: None,
..Default::default()
};
let response = self.client.update_bank(bank_id, None, &request).await?;
Ok(response.into_inner())
@@ -450,7 +454,7 @@ impl ApiClient {
_verbose: bool,
) -> Result<types::GraphDataResponse> {
self.runtime.block_on(async {
let response = self.client.get_graph(bank_id, limit, type_filter, None).await?;
let response = self.client.get_graph(bank_id, limit, type_filter, None, None, None, None).await?;
Ok(response.into_inner())
})
}
+2
View File
@@ -293,6 +293,7 @@ pub fn create(
mission: mission_text,
background: None,
disposition,
..Default::default()
};
let response = client.create_bank(bank_id, &request, verbose);
@@ -356,6 +357,7 @@ pub fn update(
mission: mission_text,
background: None,
disposition,
..Default::default()
};
let response = client.update_bank(bank_id, &request, verbose);
+1
View File
@@ -387,6 +387,7 @@ pub fn retain(
document_id: Some(doc_id.clone()),
entities: None,
tags: None,
observation_scopes: None,
};
let request = RetainRequest {
+1 -1
View File
@@ -67,7 +67,7 @@ fn format_error_message(err: &anyhow::Error, api_url: &str) -> String {
"Bank configuration API is disabled".bright_red().bold(),
"API URL:".bright_yellow(),
api_url.bright_white(),
"This feature is disabled by default for security.".bright_yellow(),
"This feature has been disabled on the server.".bright_yellow(),
"To enable, set HINDSIGHT_API_ENABLE_BANK_CONFIG_API=true on the API server".bright_white(),
"Note:".bright_cyan(),
"This allows per-bank LLM configuration overrides via API".bright_white()
+188 -11
View File
@@ -7,7 +7,7 @@ info:
name: Apache 2.0
url: https://www.apache.org/licenses/LICENSE-2.0.html
title: Hindsight HTTP API
version: 0.4.13
version: 0.4.15
servers:
- url: /
paths:
@@ -83,6 +83,34 @@ paths:
title: Limit
type: integer
style: form
- explode: true
in: query
name: q
required: false
schema:
nullable: true
type: string
style: form
- explode: true
in: query
name: tags
required: false
schema:
items:
nullable: true
type: string
nullable: true
type: array
style: form
- explode: true
in: query
name: tags_match
required: false
schema:
default: all_strict
title: Tags Match
type: string
style: form
- explode: false
in: header
name: authorization
@@ -1145,7 +1173,9 @@ paths:
title: Bank Id
type: string
style: simple
- explode: true
- description: Case-insensitive substring filter on document ID (e.g. 'report'
matches 'report-2024')
explode: true
in: query
name: q
required: false
@@ -1153,6 +1183,28 @@ paths:
nullable: true
type: string
style: form
- description: Filter documents by tags
explode: true
in: query
name: tags
required: false
schema:
items:
type: string
nullable: true
type: array
style: form
- description: "How to match tags: 'any', 'all', 'any_strict', 'all_strict'"
explode: true
in: query
name: tags_match
required: false
schema:
default: any_strict
description: "How to match tags: 'any', 'all', 'any_strict', 'all_strict'"
title: Tags Match
type: string
style: form
- explode: true
in: query
name: limit
@@ -1564,6 +1616,7 @@ paths:
- Operations
/v1/default/banks/{bank_id}/profile:
get:
deprecated: true
description: Get disposition traits and mission for a memory bank. Auto-creates
agent with defaults if not exists.
operationId: get_bank_profile
@@ -1601,6 +1654,7 @@ paths:
tags:
- Banks
put:
deprecated: true
description: "Update bank's disposition traits (skepticism, literalism, empathy)"
operationId: update_bank_disposition
parameters:
@@ -1850,6 +1904,54 @@ paths:
summary: Clear all observations
tags:
- Banks
/v1/default/banks/{bank_id}/memories/{memory_id}/observations:
delete:
description: Delete all observations derived from a specific memory and reset
it for re-consolidation. The memory itself is not deleted. A consolidation
job is triggered automatically so the memory will produce fresh observations
on the next consolidation run.
operationId: clear_memory_observations
parameters:
- explode: false
in: path
name: bank_id
required: true
schema:
title: Bank Id
type: string
style: simple
- explode: false
in: path
name: memory_id
required: true
schema:
title: Memory Id
type: string
style: simple
- explode: false
in: header
name: authorization
required: false
schema:
nullable: true
type: string
style: simple
responses:
"200":
content:
application/json:
schema:
$ref: '#/components/schemas/ClearMemoryObservationsResponse'
description: Successful Response
"422":
content:
application/json:
schema:
$ref: '#/components/schemas/HTTPValidationError'
description: Validation Error
summary: Clear observations for a memory
tags:
- Memory
/v1/default/banks/{bank_id}/config:
delete:
description: Reset bank configuration to defaults by removing all bank-specific
@@ -2595,6 +2697,17 @@ components:
- created_at
- document_id
title: ChunkResponse
ClearMemoryObservationsResponse:
description: Response model for clearing observations for a specific memory.
example:
deleted_count: 3
properties:
deleted_count:
title: Deleted Count
type: integer
required:
- deleted_count
title: ClearMemoryObservationsResponse
ConsolidationResponse:
description: Response model for consolidation trigger endpoint.
example:
@@ -2616,24 +2729,58 @@ components:
CreateBankRequest:
description: Request model for creating/updating a bank.
example:
disposition:
empathy: 3
literalism: 3
skepticism: 3
mission: I am a PM helping my engineering team stay organized
name: Alice
observations_mission: Observations are stable facts about people and projects.
Always include preferences and skills.
retain_mission: Always include technical decisions and architectural trade-offs.
Ignore meeting logistics.
properties:
name:
nullable: true
type: string
disposition:
$ref: '#/components/schemas/DispositionTraits'
disposition_skepticism:
maximum: 5.0
minimum: 1.0
nullable: true
type: integer
disposition_literalism:
maximum: 5.0
minimum: 1.0
nullable: true
type: integer
disposition_empathy:
maximum: 5.0
minimum: 1.0
nullable: true
type: integer
mission:
nullable: true
type: string
background:
nullable: true
type: string
reflect_mission:
nullable: true
type: string
retain_mission:
nullable: true
type: string
retain_extraction_mode:
nullable: true
type: string
retain_custom_instructions:
nullable: true
type: string
retain_chunk_size:
nullable: true
type: integer
enable_observations:
nullable: true
type: boolean
observations_mission:
nullable: true
type: string
title: CreateBankRequest
CreateDirectiveRequest:
description: Request model for creating a directive.
@@ -3354,9 +3501,7 @@ components:
title: Content
type: string
timestamp:
format: date-time
nullable: true
type: string
$ref: '#/components/schemas/Timestamp'
context:
nullable: true
type: string
@@ -3377,6 +3522,8 @@ components:
type: string
nullable: true
type: array
observation_scopes:
$ref: '#/components/schemas/ObservationScopes'
required:
- content
title: MemoryItem
@@ -4334,6 +4481,36 @@ components:
- api_version
- features
title: VersionResponse
Timestamp:
anyOf:
- format: date-time
type: string
- type: string
description: "When the content occurred. Accepts an ISO 8601 datetime string\
\ (e.g. '2024-01-15T10:30:00Z'), null/omitted (defaults to now), or the special\
\ string 'unset' to explicitly store without any timestamp (use this for timeless\
\ content such as fictional documents or static reference material)."
nullable: true
title: Timestamp
ObservationScopes:
anyOf:
- enum:
- per_tag
- combined
- all_combinations
type: string
- items:
items:
type: string
type: array
type: array
description: "How to scope observations during consolidation. 'per_tag' runs\
\ one consolidation pass per individual tag, creating separate observations\
\ for each tag. 'combined' (default) runs a single pass with all tags together.\
\ A list of tag lists runs one pass per inner list, giving full control over\
\ which combinations to use."
nullable: true
title: ObservationScopes
ValidationError_loc_inner:
anyOf:
- type: string
+7 -1
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
@@ -804,6 +804,8 @@ Get disposition traits and mission for a memory bank. Auto-creates agent with de
@param ctx context.Context - for authentication, logging, cancellation, deadlines, tracing, etc. Passed from http.Request or context.Background().
@param bankId
@return ApiGetBankProfileRequest
Deprecated
*/
func (a *BanksAPIService) GetBankProfile(ctx context.Context, bankId string) ApiGetBankProfileRequest {
return ApiGetBankProfileRequest{
@@ -815,6 +817,7 @@ func (a *BanksAPIService) GetBankProfile(ctx context.Context, bankId string) Api
// Execute executes the request
// @return BankProfileResponse
// Deprecated
func (a *BanksAPIService) GetBankProfileExecute(r ApiGetBankProfileRequest) (*BankProfileResponse, *http.Response, error) {
var (
localVarHTTPMethod = http.MethodGet
@@ -1560,6 +1563,8 @@ Update bank's disposition traits (skepticism, literalism, empathy)
@param ctx context.Context - for authentication, logging, cancellation, deadlines, tracing, etc. Passed from http.Request or context.Background().
@param bankId
@return ApiUpdateBankDispositionRequest
Deprecated
*/
func (a *BanksAPIService) UpdateBankDisposition(ctx context.Context, bankId string) ApiUpdateBankDispositionRequest {
return ApiUpdateBankDispositionRequest{
@@ -1571,6 +1576,7 @@ func (a *BanksAPIService) UpdateBankDisposition(ctx context.Context, bankId stri
// Execute executes the request
// @return BankProfileResponse
// Deprecated
func (a *BanksAPIService) UpdateBankDispositionExecute(r ApiUpdateBankDispositionRequest) (*BankProfileResponse, *http.Response, error) {
var (
localVarHTTPMethod = http.MethodPut
+1 -1
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
+34 -1
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
@@ -17,6 +17,7 @@ import (
"net/http"
"net/url"
"strings"
"reflect"
)
@@ -409,16 +410,31 @@ type ApiListDocumentsRequest struct {
ApiService *DocumentsAPIService
bankId string
q *string
tags *[]string
tagsMatch *string
limit *int32
offset *int32
authorization *string
}
// Case-insensitive substring filter on document ID (e.g. &#39;report&#39; matches &#39;report-2024&#39;)
func (r ApiListDocumentsRequest) Q(q string) ApiListDocumentsRequest {
r.q = &q
return r
}
// Filter documents by tags
func (r ApiListDocumentsRequest) Tags(tags []string) ApiListDocumentsRequest {
r.tags = &tags
return r
}
// How to match tags: &#39;any&#39;, &#39;all&#39;, &#39;any_strict&#39;, &#39;all_strict&#39;
func (r ApiListDocumentsRequest) TagsMatch(tagsMatch string) ApiListDocumentsRequest {
r.tagsMatch = &tagsMatch
return r
}
func (r ApiListDocumentsRequest) Limit(limit int32) ApiListDocumentsRequest {
r.limit = &limit
return r
@@ -480,6 +496,23 @@ func (a *DocumentsAPIService) ListDocumentsExecute(r ApiListDocumentsRequest) (*
if r.q != nil {
parameterAddToHeaderOrQuery(localVarQueryParams, "q", r.q, "form", "")
}
if r.tags != nil {
t := *r.tags
if reflect.TypeOf(t).Kind() == reflect.Slice {
s := reflect.ValueOf(t)
for i := 0; i < s.Len(); i++ {
parameterAddToHeaderOrQuery(localVarQueryParams, "tags", s.Index(i).Interface(), "form", "multi")
}
} else {
parameterAddToHeaderOrQuery(localVarQueryParams, "tags", t, "form", "multi")
}
}
if r.tagsMatch != nil {
parameterAddToHeaderOrQuery(localVarQueryParams, "tags_match", r.tagsMatch, "form", "")
} else {
var defaultValue string = "any_strict"
r.tagsMatch = &defaultValue
}
if r.limit != nil {
parameterAddToHeaderOrQuery(localVarQueryParams, "limit", r.limit, "form", "")
} else {
+1 -1
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
+1 -1
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
+166 -1
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
@@ -17,6 +17,7 @@ import (
"net/http"
"net/url"
"strings"
"reflect"
)
@@ -155,12 +156,141 @@ func (a *MemoryAPIService) ClearBankMemoriesExecute(r ApiClearBankMemoriesReques
return localVarReturnValue, localVarHTTPResponse, nil
}
type ApiClearMemoryObservationsRequest struct {
ctx context.Context
ApiService *MemoryAPIService
bankId string
memoryId string
authorization *string
}
func (r ApiClearMemoryObservationsRequest) Authorization(authorization string) ApiClearMemoryObservationsRequest {
r.authorization = &authorization
return r
}
func (r ApiClearMemoryObservationsRequest) Execute() (*ClearMemoryObservationsResponse, *http.Response, error) {
return r.ApiService.ClearMemoryObservationsExecute(r)
}
/*
ClearMemoryObservations Clear observations for a memory
Delete all observations derived from a specific memory and reset it for re-consolidation. The memory itself is not deleted. A consolidation job is triggered automatically so the memory will produce fresh observations on the next consolidation run.
@param ctx context.Context - for authentication, logging, cancellation, deadlines, tracing, etc. Passed from http.Request or context.Background().
@param bankId
@param memoryId
@return ApiClearMemoryObservationsRequest
*/
func (a *MemoryAPIService) ClearMemoryObservations(ctx context.Context, bankId string, memoryId string) ApiClearMemoryObservationsRequest {
return ApiClearMemoryObservationsRequest{
ApiService: a,
ctx: ctx,
bankId: bankId,
memoryId: memoryId,
}
}
// Execute executes the request
// @return ClearMemoryObservationsResponse
func (a *MemoryAPIService) ClearMemoryObservationsExecute(r ApiClearMemoryObservationsRequest) (*ClearMemoryObservationsResponse, *http.Response, error) {
var (
localVarHTTPMethod = http.MethodDelete
localVarPostBody interface{}
formFiles []formFile
localVarReturnValue *ClearMemoryObservationsResponse
)
localBasePath, err := a.client.cfg.ServerURLWithContext(r.ctx, "MemoryAPIService.ClearMemoryObservations")
if err != nil {
return localVarReturnValue, nil, &GenericOpenAPIError{error: err.Error()}
}
localVarPath := localBasePath + "/v1/default/banks/{bank_id}/memories/{memory_id}/observations"
localVarPath = strings.Replace(localVarPath, "{"+"bank_id"+"}", url.PathEscape(parameterValueToString(r.bankId, "bankId")), -1)
localVarPath = strings.Replace(localVarPath, "{"+"memory_id"+"}", url.PathEscape(parameterValueToString(r.memoryId, "memoryId")), -1)
localVarHeaderParams := make(map[string]string)
localVarQueryParams := url.Values{}
localVarFormParams := url.Values{}
// to determine the Content-Type header
localVarHTTPContentTypes := []string{}
// set Content-Type header
localVarHTTPContentType := selectHeaderContentType(localVarHTTPContentTypes)
if localVarHTTPContentType != "" {
localVarHeaderParams["Content-Type"] = localVarHTTPContentType
}
// to determine the Accept header
localVarHTTPHeaderAccepts := []string{"application/json"}
// set Accept header
localVarHTTPHeaderAccept := selectHeaderAccept(localVarHTTPHeaderAccepts)
if localVarHTTPHeaderAccept != "" {
localVarHeaderParams["Accept"] = localVarHTTPHeaderAccept
}
if r.authorization != nil {
parameterAddToHeaderOrQuery(localVarHeaderParams, "authorization", r.authorization, "simple", "")
}
req, err := a.client.prepareRequest(r.ctx, localVarPath, localVarHTTPMethod, localVarPostBody, localVarHeaderParams, localVarQueryParams, localVarFormParams, formFiles)
if err != nil {
return localVarReturnValue, nil, err
}
localVarHTTPResponse, err := a.client.callAPI(req)
if err != nil || localVarHTTPResponse == nil {
return localVarReturnValue, localVarHTTPResponse, err
}
localVarBody, err := io.ReadAll(localVarHTTPResponse.Body)
localVarHTTPResponse.Body.Close()
localVarHTTPResponse.Body = io.NopCloser(bytes.NewBuffer(localVarBody))
if err != nil {
return localVarReturnValue, localVarHTTPResponse, err
}
if localVarHTTPResponse.StatusCode >= 300 {
newErr := &GenericOpenAPIError{
body: localVarBody,
error: localVarHTTPResponse.Status,
}
if localVarHTTPResponse.StatusCode == 422 {
var v HTTPValidationError
err = a.client.decode(&v, localVarBody, localVarHTTPResponse.Header.Get("Content-Type"))
if err != nil {
newErr.error = err.Error()
return localVarReturnValue, localVarHTTPResponse, newErr
}
newErr.error = formatErrorMessage(localVarHTTPResponse.Status, &v)
newErr.model = v
}
return localVarReturnValue, localVarHTTPResponse, newErr
}
err = a.client.decode(&localVarReturnValue, localVarBody, localVarHTTPResponse.Header.Get("Content-Type"))
if err != nil {
newErr := &GenericOpenAPIError{
body: localVarBody,
error: err.Error(),
}
return localVarReturnValue, localVarHTTPResponse, newErr
}
return localVarReturnValue, localVarHTTPResponse, nil
}
type ApiGetGraphRequest struct {
ctx context.Context
ApiService *MemoryAPIService
bankId string
type_ *string
limit *int32
q *string
tags *[]*string
tagsMatch *string
authorization *string
}
@@ -174,6 +304,21 @@ func (r ApiGetGraphRequest) Limit(limit int32) ApiGetGraphRequest {
return r
}
func (r ApiGetGraphRequest) Q(q string) ApiGetGraphRequest {
r.q = &q
return r
}
func (r ApiGetGraphRequest) Tags(tags []*string) ApiGetGraphRequest {
r.tags = &tags
return r
}
func (r ApiGetGraphRequest) TagsMatch(tagsMatch string) ApiGetGraphRequest {
r.tagsMatch = &tagsMatch
return r
}
func (r ApiGetGraphRequest) Authorization(authorization string) ApiGetGraphRequest {
r.authorization = &authorization
return r
@@ -231,6 +376,26 @@ func (a *MemoryAPIService) GetGraphExecute(r ApiGetGraphRequest) (*GraphDataResp
var defaultValue int32 = 1000
r.limit = &defaultValue
}
if r.q != nil {
parameterAddToHeaderOrQuery(localVarQueryParams, "q", r.q, "form", "")
}
if r.tags != nil {
t := *r.tags
if reflect.TypeOf(t).Kind() == reflect.Slice {
s := reflect.ValueOf(t)
for i := 0; i < s.Len(); i++ {
parameterAddToHeaderOrQuery(localVarQueryParams, "tags", s.Index(i).Interface(), "form", "multi")
}
} else {
parameterAddToHeaderOrQuery(localVarQueryParams, "tags", t, "form", "multi")
}
}
if r.tagsMatch != nil {
parameterAddToHeaderOrQuery(localVarQueryParams, "tags_match", r.tagsMatch, "form", "")
} else {
var defaultValue string = "all_strict"
r.tagsMatch = &defaultValue
}
// to determine the Content-Type header
localVarHTTPContentTypes := []string{}
+1 -1
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
+1 -1
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
+1 -1
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
+2 -2
View File
@@ -3,7 +3,7 @@ Hindsight HTTP API
HTTP API for Hindsight
API version: 0.4.13
API version: 0.4.15
*/
// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
@@ -41,7 +41,7 @@ var (
queryDescape = strings.NewReplacer( "%5B", "[", "%5D", "]" )
)
// APIClient manages communication with the Hindsight HTTP API API v0.4.13
// APIClient manages communication with the Hindsight HTTP API API v0.4.15
// In most cases there should be only one, shared, APIClient.
type APIClient struct {
cfg *Configuration

Some files were not shown because too many files have changed in this diff Show More