Compare commits

...
54 Commits
Author SHA1 Message Date
Nicolò Boschi 017e624669 fixes 2026-01-29 16:43:33 +01:00
Nicolò Boschi 3ba9a7ce26 fix(doc): improve docs versioning and release 2026-01-29 14:45:16 +01:00
Nicolò Boschi 885b01d4f8 fix(doc): improve docs versioning and release 2026-01-29 14:21:48 +01:00
Nicolò Boschi e1e9496224 fix: hindsight-embed on macos crashes 2026-01-29 13:29:29 +01:00
Nicolò Boschi 16fc05f6b1 fix: hindsight-embed on macos crashes 2026-01-29 13:13:38 +01:00
Nicolò Boschi b8f06a09fb Release v0.4.1
- Update version to 0.4.1 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-all, hindsight-litellm, hindsight-embed
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2026-01-29 11:25:28 +01:00
Nicolò Boschi b43ef98686 feat: consolidation performance benchmark and optimization (#227) 2026-01-29 11:24:15 +01:00
Nicolò Boschi f17703fb37 doc: hide next version (#226) 2026-01-29 08:46:21 +01:00
Nicolò Boschi cfcc23c152 fix: /version endpoint return wrong version (#224)
* fix: /version endpoint return wrong version

* chore: update OpenAPI spec with correct version example
2026-01-29 08:40:01 +01:00
Chris Latimer 7300d5be4b README video 2026-01-28 19:26:12 -07:00
Chris Latimer 81c82d9b93 README tweak 2026-01-28 19:21:15 -07:00
Chris Latimer 7551e65e55 Updated video in readme 2026-01-28 14:40:44 -07:00
DK09876andClaude Opus 4.5 94cc0a1270 fix: search_mental_models uuid type mismatch after text id migration (#225)
The mental_models.id column was changed from UUID to TEXT in migration
u6p7q8r9s0t1, but the exclude_ids filter in search_mental_models still
cast the parameter as ::uuid[]. This caused every search_mental_models
call during reflect to fail with "operator does not exist: text <> uuid",
forcing the reflect agent to waste all 5 iterations on retries and
producing degraded mental model content.

Co-authored-by: Claude Opus 4.5 <[email protected]>
2026-01-28 20:04:19 +01:00
Nicolò BoschiandClaude Sonnet 4.5 67c47881cb fix: add defensive error handling to PyTorch device detection (#221)
* fix: include correct __version__ in python packages

* fix(embed): force CPU mode for local models in daemon to prevent XPC crashes

Adds HINDSIGHT_API_EMBEDDINGS_LOCAL_FORCE_CPU and HINDSIGHT_API_RERANKER_LOCAL_FORCE_CPU
environment variables to force CPU-only operation for local sentence-transformer models.

This prevents XPC_ERROR_CONNECTION_INVALID crashes on macOS when running in daemon mode.
The issue occurs because PyTorch's MPS (Metal Performance Shaders) backend has unstable
XPC connections in background processes, leading to C++ assertion failures that Python
exception handlers cannot catch.

Changes:
- config.py: Add ENV_*_FORCE_CPU constants and config dataclass fields
- embeddings.py: Add force_cpu parameter to LocalSTEmbeddings constructor
- cross_encoder.py: Add force_cpu parameter to LocalSTCrossEncoder constructor
- main.py: Set force CPU env vars in daemon mode, add fields to config constructor

The daemon mode automatically enables force CPU for both embeddings and reranker,
while normal mode allows hardware acceleration (GPU/MPS) as before.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude Sonnet 4.5 <[email protected]>

* fix: add defensive error handling to PyTorch device detection

Wraps all PyTorch device detection code (torch.cuda.is_available()
and torch.backends.mps.is_available()) in try-except blocks that
gracefully fall back to CPU if any errors occur.

This complements PR #218's force_cpu configuration by ensuring the
code works reliably in all environments without configuration:
- CI environments with CPU-only PyTorch builds
- Systems without proper GPU/MPS support
- Partial or misconfigured PyTorch installations

The defensive approach prevents startup failures while still taking
advantage of GPU/MPS acceleration when available and force_cpu is
not explicitly set.

Changes:
- embeddings.py: Added try-except in initialize() and _reinitialize_model_sync()
- cross_encoder.py: Added try-except in initialize() and _reinitialize_model_sync()

* refactor: use get_config() for embeddings and reranker force_cpu

Changes create_embeddings_from_env() and create_cross_encoder_from_env()
to read configuration via get_config() instead of directly accessing
os.environ. This ensures consistency across the codebase and properly
respects the force_cpu configuration set by daemon mode.

Changes:
- embeddings.py: Use config.embeddings_local_model and config.embeddings_local_force_cpu
- cross_encoder.py: Use config.reranker_local_model and config.reranker_local_force_cpu
- Both: Use get_config() for provider, tei_url, and other config fields
- Note: Some fields not in config (like max_concurrent for local reranker) still read from os.environ

This fixes the issue where force_cpu was read inconsistently from environment
variables instead of using the centralized config system.

* test: clear config cache in test_create_from_env

Fixes test failure caused by cached config not picking up
environment variable changes in test. The test now calls
clear_config_cache() before and after patching os.environ
to ensure the factory function reads the test's env vars.

* refactor: add reranker_local_max_concurrent to config system

Adds reranker_local_max_concurrent to HindsightConfig dataclass
and removes the workaround in create_cross_encoder_from_env() that
was reading it directly from os.environ.

Changes:
- config.py: Add reranker_local_max_concurrent field to dataclass and from_env()
- main.py: Add reranker_local_max_concurrent to manual config constructor
- cross_encoder.py: Use config.reranker_local_max_concurrent instead of os.environ

This completes the refactoring to use the centralized config system
for all reranker configuration.

---------

Co-authored-by: Claude Sonnet 4.5 <[email protected]>
2026-01-28 18:14:54 +01:00
Nicolò Boschi 2b72e1fd68 feat: support different default pg schema (#222)
* feat: support different default pg schema

* feat: support different default pg schema
2026-01-28 18:14:44 +01:00
Nicolò BoschiandChris Latimer d2b797fff8 doc: improve readme (#223)
* README updates

* Add captions to video

* Use cases and new banner

---------

Co-authored-by: Chris Latimer <[email protected]>
2026-01-28 17:55:20 +01:00
Nicolò Boschi fccbdfef16 fix: include correct __version__ in python packages (#218)
Updates:
- hindsight-api/hindsight_api/__init__.py: bump __version__ to 0.4.0
- scripts/release.sh: add logic to update __version__ in Python __init__.py files during release
2026-01-28 17:25:17 +01:00
Nicolò Boschi 20f2b92069 doc: release notes for 0.4.0 (#217)
* doc: release notes for 0.4.0

* doc: release notes for 0.4.0

* doc: release notes for 0.4.0

* doc: release notes for 0.4.0
2026-01-28 16:54:05 +01:00
Nicolò Boschi 1bf90358c3 doc: add blog (#201)
* doc: introduce mental models blog post

Write blog post introducing Mental Models in Hindsight 0.4.0:
- Evolution from observations and opinions
- How mental models work (consolidation, evidence tracking)
- Breaking changes and migration path
- Environment variable to enable (experimental)
- Agentic reflect explanation

* updates

* Update 2026-01-26-learning-capabilities.md

* fix: doc build issues

- Add missing code snippets for versioned docs (recall-opinions-only, recall-include-entities, bank-background)
- Fix broken links by using relative paths for version compatibility
- Update blog post title to sentence case
- Clear versions.json since v0.3 versioned docs don't exist yet
- Enable INCLUDE_CURRENT_VERSION in build script

* fix: update doc links after rebase

- Fix blog post to link to correct pages (/developer/api/mental-models and /developer/observations)
- Fix CLI docs to link to /api-reference instead of /api

* feat: add directives section to blog post

- Update intro to mention three layers of knowledge
- Add concise Directives section for compliance/guardrails
- Add directives to resources section
- Keep focus on learning capabilities (observations and mental models)

* fix: revert intro to focus on learning capabilities only

Directives are a separate feature for compliance/guardrails, not a learning capability. The blog post is about observations and mental models.
2026-01-28 15:42:14 +01:00
Nicolò Boschi 2118d0a7cd Release v0.4.0
- Update version to 0.4.0 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-all, hindsight-litellm, hindsight-embed
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2026-01-28 15:04:43 +01:00
Nicolò Boschi e5fc6eedb6 fix(embed): daemon process XPC connection crash on macos (#215)
* fix(embed): daemon process XPC connection crash on macos

* other fix
2026-01-28 14:52:31 +01:00
Nicolò Boschi bb0e0316a7 fix: graph endpoint not showing links for observations (#214) 2026-01-28 14:51:25 +01:00
Nicolò Boschi 3172e99cab feat: add custom extraction prompt (#213)
* feat: add custom extraction prompt

* feat: add custom extraction prompt

* test
2026-01-28 13:54:52 +01:00
Nicolò BoschiandClaude Sonnet 4.5 1c9a7a0d5e chore: cleanup benchmarks runner with old flags (#212)
* chore: cleanup benchmarks runner with old flags

* fix tests

* fix: observations rely on source_memory_ids, no link copying

Observations no longer copy any memory_links from their source facts.
Instead, retrieval uses source_memory_ids to traverse:
- Entity connections: observation → source_memory_ids → unit_entities
- Semantic similarity: observations have their own embeddings
- Temporal proximity: observations have their own temporal fields

This avoids data duplication and fixes bidirectionality issues with
entity links being copied to observations.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude Sonnet 4.5 <[email protected]>

* test: update consolidation test for source_memory_ids behavior

Updated test_consolidation_creates_memory_links to test_consolidation_uses_source_memory_ids
to reflect the new behavior where observations use source_memory_ids instead of memory_links
for traversal.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude Sonnet 4.5 <[email protected]>

---------

Co-authored-by: Claude Sonnet 4.5 <[email protected]>
2026-01-28 13:22:48 +01:00
Nicolò Boschi 90e370ef35 fix: misc fixes for observations and mental models (#209)
* fix: misc fixes for observations and mental models

* feat: improve graph retrieval for observations

- Update LinkExpansionRetriever to traverse through source_memory_ids
  for observation entity connections (avoiding data duplication)
- Remove entity link copy from world facts to observations in consolidator
- Add tests for link expansion graph retrieval
- Add directives_applied field to ReflectResult
- Include user's other changes (CLI, docs, client updates)

* fix: CI test failures

- Add mental_model_id parameter to create_mental_model function
- Fix ToolCallTrace not including reason field from ToolCall
- Improve test_link_expansion_observation_graph_retrieval to wait for consolidation with retry

* chore: reduce link expansion log verbosity

* Revert "chore: reduce link expansion log verbosity"

This reverts commit 3ce759391cead1012157785fa78fef16ef9bfe3b.

* feat: add semantic/temporal/entity links as fallback in graph retrieval

- Add fallback query for semantic, temporal, and entity links from memory_links
- Check both directions (outgoing and incoming links)
- Weight fallback results at 0.5x to prioritize entity links via unit_entities
- Fixes graph retrieval returning 0 when data has cross-cluster temporal connections

* fix: enable observations fixture for link expansion test

- Add enable_observations fixture to ensure observations are created
- Increase wait time from 10 to 30 seconds for CI reliability
2026-01-27 15:37:57 +01:00
Nicolò Boschi 084242a6dd chore: drop dead code (#210) 2026-01-27 15:03:25 +01:00
Chris Bartholomew 83f44c4b41 fix: multi-tenant schema context for worker task execution (#208)
Background tasks (async retain, consolidation, reflections) fail in
multi-tenant deployments because the worker executes tasks without
setting the tenant schema context. This causes two failures:

1. The cancellation check in execute_task queries public.async_operations
   instead of the tenant's schema, finds no row, and skips the task as
   "cancelled" — even though it wasn't.

2. Even if that were fixed, _authenticate_tenant would throw
   AuthenticationError because background tasks have no API key.

Changes:
- Poller passes task.schema into task_dict so execute_task can set it
- execute_task sets _current_schema before the cancellation check
- Task handlers use RequestContext(internal=True) to signal background ops
- _authenticate_tenant skips extension auth for internal requests when
  schema is already set
- BrokerTaskBackend uses schema_getter for dynamic schema resolution
  when submitting tasks and waiting for results
- Pass tenant_extension to WorkerPoller in create_app
2026-01-27 12:28:47 +01:00
Chris Bartholomew 7bdb8fc2e3 fix: include tags, created_at, proof_count in graph table_rows (#207)
The graph endpoint's table_rows response was missing three fields that
the control plane UI expects:
- tags: memory unit tags (shown in Tags column)
- created_at: creation timestamp (shown in Created column for mental models)
- proof_count: source memory count (shown in Sources column for mental models)

All three columns exist on the memory_units table but were not being
selected or included in the response.
2026-01-27 09:54:07 +01:00
Nicolò Boschi 5b52a84fff chore: internal renames (#204)
This commit renames the terminology across the entire codebase:
- "mental models" (fact_type='mental_model' in memory_units) → "observations"
- "reflections" table (stored reflect responses) → "mental_models"

Changes include:
- Database migration to rename tables, indexes, and constraints
- API endpoints: /reflections → /mental-models, /mental-models → /observations
- Config: ENABLE_MENTAL_MODELS → ENABLE_OBSERVATIONS
- Response models and Pydantic classes
- Reflect agent tools and prompts
- Control plane UI and routes
- Documentation and examples
- Regenerated OpenAPI spec and client SDKs (Python, TypeScript)
- Rust CLI: reflection commands → mental-model commands
- LiteLLM: updated fact_types documentation
2026-01-27 09:53:28 +01:00
Nicolò Boschi f3c5a9c1c2 feat(litellm): support tags and mission in litellm package (#202) 2026-01-26 20:37:39 +01:00
Nicolò Boschi 5832b907c6 fix(ui): reflections based on don't show up all contents (#203) 2026-01-26 18:43:47 +01:00
Nicolò Boschi 50fa2ed090 ci: add upgrade tests (#200) 2026-01-26 15:25:30 +01:00
Nicolò Boschi 522b71aab8 doc: mental models (#199)
* doc: mental models

* doc: mental models
2026-01-26 14:27:08 +01:00
Nicolò Boschi 31b5c5845d chore: versioned docs (#198) 2026-01-26 11:21:13 +01:00
c0ca9b027e Fix: Pass api_key to Hindsight client in litellm integration (#193)
* Fix: Pass api_key to Hindsight client in litellm integration

The recall(), reflect(), and retain() wrapper functions were creating
Hindsight client instances without passing the api_key from the config.
This caused 401 Unauthorized errors when using hindsight-litellm with
authenticated Hindsight API servers.

Also added api_key parameter to:
- HindsightOpenAI and HindsightAnthropic wrapper classes
- wrap_openai() and wrap_anthropic() functions

* Add sensible defaults for simpler API usage

Make it easier to get started with hindsight-litellm by providing
sensible defaults:

- Default API URL: https://api.hindsight.vectorize.io (production)
- Default bank_id: "default"
- Read api_key from HINDSIGHT_API_KEY environment variable

Now users can simply do:

    client = wrap_openai(OpenAI())

With just the HINDSIGHT_API_KEY env var set, and it works.

Also adds comprehensive unit tests for the new defaults behavior.

* Fix test using non-existent 'enabled' parameter in configure()

The test was calling configure(enabled=False) but configure() doesn't
have an enabled parameter. Changed to test is_configured() returns False
when reset_config() has been called.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

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

* Fix: rename 'background' parameter to 'mission' in Python client create_bank()

The parameter was named 'background' but the internal code used 'mission',
causing undefined variable errors. The tests also expected 'mission'.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

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

---------

Co-authored-by: Nicolò Boschi <[email protected]>
Co-authored-by: Claude Opus 4.5 <[email protected]>
2026-01-26 10:49:06 +01:00
1d4879a206 feat(litellm): async retain, reflect support, and API cleanup (#167)
* feat(litellm): async retain with sync option, fix client session cleanup

- Add sync parameter to retain() for blocking vs background operation
- Default to async retain (sync=False) for better performance
- Add get_pending_retain_errors() to check async failures
- Fix "Unclosed client session" warnings by properly closing clients
- Fix "Timeout context manager" asyncio errors by creating fresh clients
- Each API call now creates and closes its own client (aiohttp limitation)
- Add _get_client() and _close_client() helpers for consistent handling
- Update recall(), reflect(), _retain_sync() and _inject_memories()

🤖 Generated with [Claude Code](https://claude.com/claude-code)

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

* feat(litellm): add reflect support and require explicit hindsight_query

- Make hindsight_query required when inject_memories=True to enforce
  intentional memory queries (no automatic last-user-message fallback)
- Add reflect_context parameter for shaping LLM reasoning in reflect
- Add reflect_response_schema for structured JSON output from reflect
- Add _reflect_sync() and _reflect_async() methods in callbacks
- Update wrappers.py to support response_schema in reflect/areflect

This improves the developer experience by making memory injection
explicit and adds full reflect API support through the integration.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

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

* feat(litellm): rename recall_budget to budget, add per-call reflect context

- Rename `recall_budget` parameter to `budget` for consistency with API
- Add `hindsight_reflect_context` kwarg for per-call reflect context override
- Fix reflect() to not pass None values for optional parameters

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

* docs(litellm): update README for new API structure and features

- Document configure() vs set_defaults() separation
- Add hindsight_query requirement when inject_memories=True
- Document async retain (sync=False default) and get_pending_retain_errors()
- Add hindsight_reflect_context per-call override documentation
- Document budget parameter (renamed from recall_budget)
- Add reflect_context and reflect_response_schema options
- Update all code examples to use new API structure
- Add new functions to API Reference table

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

* test(litellm): update tests for new configure/set_defaults API

- Update tests to use separate configure() and set_defaults() calls
- Fix test assertions to check config vs defaults appropriately
- Add tests for legacy parameter backwards compatibility
- Add new TestSetDefaults test class
- Fix _format_memories test call signature (settings, config order)

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

* feat: add set_bank_mission(), deprecate set_bank_background()

- Add mission parameter to hindsight_client.create_bank()
- Add set_bank_mission() function to hindsight_litellm
- Deprecate set_bank_background() with DeprecationWarning
- Update _create_or_update_bank() to support mission parameter
- Update README and docstrings to document the new API

The 'background' field has been deprecated in the Hindsight API in favor
of 'mission' which is used for mental model generation.

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

* Remove deprecated background parameter and legacy configure() parameters

- Remove set_bank_background() in favor of set_bank_mission()
- Remove background parameter from _create_or_update_bank()
- Remove background parameter from hindsight_client.create_bank()
- Remove legacy parameters from configure() (bank_id, document_id, budget, etc.)
- These have been replaced by the set_defaults() API
- Remove legacy test cases for deprecated parameters

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

* fix: update tests and docs to use mission instead of background

The create_bank() parameter was renamed from background to mission.
Update all tests and doc examples to use the new parameter name.

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

---------

Co-authored-by: Claude Opus 4.5 <[email protected]>
Co-authored-by: Nicolò Boschi <[email protected]>
2026-01-26 10:07:32 +01:00
Nicolò Boschi 8e39cb7bc8 fix: improve mental model consolidation (#197)
* fix: improve mental model consolidation

* fix skill names

* fixes

* fix: add missing list_tenants to test mocks and update CLI for async refresh
2026-01-26 09:54:25 +01:00
Nicolò Boschi b378f6852f feat(mcp): add timestamp to retain (#190)
* feat(mcp): add timestamp to retain

* ci
2026-01-23 16:00:43 +01:00
Nicolò Boschi 9c2df9d89f fix skill names 2026-01-23 11:03:12 +01:00
Nicolò Boschi ec2231799e feat: support for npx add-skill (#191)
* feat: support for npx add-skill

* skills
2026-01-23 11:01:53 +01:00
Phạm Gia Linh aebef9408b feat(python-sdk): add tags filtering support to high-level client (#186)
Add tags and tags_match parameters to recall/reflect methods for
filtering
memories by visibility scope. Also add tags support to retain methods.

Changes:
- recall()/arecall(): add tags, tags_match parameters
- reflect()/areflect(): add tags, tags_match parameters
- retain()/aretain(): add tags parameter
- retain_batch()/aretain_batch(): add document_tags parameter
- Add TestTags test class with 7 tests
2026-01-23 10:02:05 +01:00
Chris Bartholomew 66abad61b8 Fix Gemini tool response format by including function name (#187)
Gemini requires the 'name' field in tool/function response messages,
while OpenAI infers it from tool_call_id. Without it, Gemini returns:
  'function_response.name: Name cannot be empty'

Added 'name' field to both tool result messages in the reflect agent.
2026-01-23 07:38:51 +01:00
Nicolò Boschi 9db64ecda3 feat: revisit mental models, directives and reflections (#179)
* chore: run benchmarks with reflect mode

* chore: run benchmarks with reflect mode

* fixes

* new mm

* bunch of fixes

* initial commit

* fixes

* fixes

* fixes

* fix: sometimes memories gets extracted in the wrong language
2026-01-22 17:13:16 +01:00
Nicolò Boschi ddaa5f5f1b fix: simplify pytorch model initialization to prevent meta tensor issues (#185)
Remove device_map from model_kwargs as it conflicts with CrossEncoder's
internal .to(device) call. The low_cpu_mem_usage=False setting alone is
sufficient to prevent lazy loading (meta tensors).
2026-01-22 16:25:28 +01:00
Nicolò Boschi 87d4a36509 fix: sometimes memories gets extracted in the wrong language (#184) 2026-01-22 14:21:59 +01:00
Nicolò Boschi 0bf85a3435 fix: improve pytorch model initialization to prevent meta tensor issues (#180)
* fix: prevent meta tensor issues when accelerate is installed without GPU

When accelerate is installed but no GPU is available, transformers can
incorrectly use lazy loading (meta tensors) which fails when
sentence-transformers tries to move the model to a device.

The fix checks hardware and installed packages to determine the right
loading strategy:
- GPU available: device=None, device_map=None (auto-detect GPU)
- No GPU + accelerate: device='cpu', device_map='cpu' (force CPU loading)
- No GPU + no accelerate: device='cpu', device_map=None (normal CPU)

🤖 Generated with [Claude Code](https://claude.com/claude-code)

* fix: add filelock for model initialization in parallel tests

When pytest-xdist runs multiple workers in parallel, they all try to
load models from the HuggingFace cache simultaneously, causing race
conditions and intermittent meta tensor errors.

Added filelock around embeddings and cross_encoder initialization in
conftest.py, similar to how pg0 database setup is serialized. Models
are now pre-initialized in the fixture before being passed to tests.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

* fix: add MPS support for macOS Apple Silicon

Extend GPU detection to include Apple MPS backend in addition to CUDA.
This ensures macOS users with Apple Silicon use MPS acceleration
instead of being incorrectly routed to the CPU fallback path.
2026-01-20 14:15:51 +01:00
Nicolò Boschi 16b85a4faa chore: drop unused access_count column (#178) 2026-01-20 10:36:02 +01:00
Nicolò Boschi 4c792400c1 feat: new 'worker' service (#176)
* feat: new 'worker' service

* doc

* docs

* tests
2026-01-20 10:17:56 +01:00
Nicolò Boschi 0284595909 fix: pytorch init failures (#175) 2026-01-19 15:02:20 +01:00
Nicolò Boschi fe4ed1db73 feat(clients): mental models api (#172)
* feat(clients): mental models api

* fixes

* more tests

* fixes
2026-01-19 14:49:49 +01:00
Nicolò Boschi bac4b24e30 fix(sec): upgrade vulnerable deps (#174) 2026-01-19 14:27:31 +01:00
Nicolò Boschi 3290f4bfff chore: unify agents.md and claude.md (#173) 2026-01-19 14:18:33 +01:00
Nicolò Boschi 63a65d0723 feat: improve mental model refresh and add directives (#166)
* feat: improve mental model refresh and add directives

* feat: improve mental model refresh and add directives

* tags

* ui

* fix

* fix

* update

* update
2026-01-19 11:38:35 +01:00
Chris Bartholomew 870cfccabb Add structured JSON logging support (#170)
* Add structured JSON logging support

Add HINDSIGHT_API_LOG_FORMAT environment variable to configure log output
format. Options are "text" (default, human-readable) and "json" (structured).

JSON format outputs logs with a "severity" field that cloud logging systems
can parse for proper log level categorization. Also writes to stdout instead
of stderr so log levels are correctly interpreted.

* Rename GCPJsonFormatter to JsonFormatter
2026-01-19 09:01:24 +01:00
456 changed files with 52920 additions and 15992 deletions
+1
View File
@@ -26,6 +26,7 @@ HINDSIGHT_API_LOG_LEVEL=info
# Database (Optional - uses embedded pg0 by default)
# HINDSIGHT_API_DATABASE_URL=postgresql://user:pass@host:5432/db
# HINDSIGHT_API_DATABASE_SCHEMA=public # PostgreSQL schema name (default: public)
# Embeddings Configuration (Optional - uses local by default)
# Provider: "local" (default) or "tei" (HuggingFace Text Embeddings Inference)
+60
View File
@@ -875,6 +875,66 @@ jobs:
echo "=== API Server Logs ==="
cat /tmp/api-server.log || echo "No API server log found"
test-upgrade:
runs-on: ubuntu-latest
env:
HINDSIGHT_API_LLM_PROVIDER: groq
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
with:
fetch-depth: 0 # Full history needed for git clone of tags
- name: Fetch tags
run: git fetch --tags
- 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: Cache HuggingFace models
uses: actions/cache@v4
with:
path: ~/.cache/huggingface
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
restore-keys: |
${{ runner.os }}-huggingface-
- name: Install hindsight-dev dependencies
working-directory: ./hindsight-dev
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
- name: Install current hindsight-api
working-directory: ./hindsight-api
run: uv sync --frozen --index-strategy unsafe-best-match
- name: Pre-download models
working-directory: ./hindsight-api
run: |
uv run python -c "
from sentence_transformers import SentenceTransformer, CrossEncoder
print('Downloading embedding model...')
SentenceTransformer('BAAI/bge-small-en-v1.5')
print('Downloading cross-encoder model...')
CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2')
print('Models downloaded successfully')
"
- name: Run upgrade tests
working-directory: ./hindsight-dev
run: uv run pytest upgrade_tests/ -v --tb=short
verify-generated-files:
runs-on: ubuntu-latest
env:
+5 -2
View File
@@ -29,7 +29,7 @@ nltk_data/
# Monitoring stack (Prometheus/Grafana binaries and data)
.monitoring/
.pgbouncer
.pgbouncer/
# Large benchmark datasets (will be downloaded automatically)
**/longmemeval_s_cleaned.json
@@ -45,9 +45,12 @@ hindsight-docs/static/llms-full.txt
hindsight-dev/benchmarks/locomo/results/
hindsight-dev/benchmarks/longmemeval/results/
hindsight-dev/benchmarks/consolidation/results/
benchmarks/results/
hindsight-cli/target
hindsight-clients/rust/target
.claude
whats-next.md
TASK.md
CHANGELOG.md
# Changelog is now tracked in hindsight-docs/src/pages/changelog.md
# CHANGELOG.md
-39
View File
@@ -1,39 +0,0 @@
[databases]
; Connect to pg0 on port 5433
; The actual pg0 database is called "hindsight"
hindsight = host=127.0.0.1 port=5433 dbname=hindsight user=hindsight password=hindsight
[pgbouncer]
listen_addr = 127.0.0.1
listen_port = 6432
; Use md5 authentication (matches pg0's auth)
auth_type = md5
auth_file = /Users/nicoloboschi/dev/memory-poc/.pgbouncer/userlist.txt
; Transaction pooling mode (recommended for hindsight)
pool_mode = transaction
; Reset connection state after each transaction
server_reset_query = DISCARD ALL
; Pool sizing
default_pool_size = 20
max_client_conn = 200
min_pool_size = 5
; Timeouts
server_idle_timeout = 600
server_lifetime = 3600
query_timeout = 120
; Logging
log_connections = 1
log_disconnections = 1
log_pooler_errors = 1
; Stats
stats_period = 60
; Admin console
admin_users = admin
-2
View File
@@ -1,2 +0,0 @@
"hindsight" "md5d842ccb6249bcd3c53b2f648378092a6"
"admin" ""
+1 -151
View File
@@ -1,153 +1,3 @@
# AGENTS.md
This document captures architectural decisions and coding conventions for the Hindsight project.
## Documentation
- **Main documentation**: [hindsight-docs/docs/developer/](./hindsight-docs/docs/developer/)
- **Use case patterns**: [hindsight-docs/docs/cookbook/](./hindsight-docs/docs/cookbook/)
- **API reference**: Auto-generated from OpenAPI spec
## Project Structure
```
hindsight/ # Python package for embedded usage
hindsight-api/ # FastAPI server (core memory engine)
hindsight-cli/ # Rust CLI client
hindsight-embed/ # Embedded CLI (no server needed)
hindsight-control-plane/ # Next.js admin UI
hindsight-docs/ # Docusaurus documentation site
hindsight-dev/ # Development tools and benchmarks
hindsight-integrations/ # Framework integrations (LangChain, etc.)
hindsight-clients/ # Generated API clients (Python, TypeScript, Rust)
```
## Core Concepts
### Memory Banks
- Each bank is an isolated memory store (like a "brain" for one user/agent)
- Banks contain: memory units (facts), entities, documents, entity links
- Banks have a **disposition** (personality traits) and **background** (context)
- Bank isolation is strict - no cross-bank data leakage
### Memory Types
- **World facts**: General knowledge ("The sky is blue")
- **Experience facts**: Personal experiences ("I visited Paris in 2023")
- **Opinion facts**: Beliefs with confidence scores ("Paris is beautiful" - 0.9 confidence)
### Operations
- **Retain**: Store new memories (extracts facts, entities, relationships)
- **Recall**: Retrieve memories (semantic, BM25, graph, temporal search)
- **Reflect**: Deep analysis to form new insights/opinions
## API Design Decisions
### Single Bank Per Request
- All API endpoints (`recall`, `reflect`, `retain`) operate on a single bank
- Multi-bank queries are the **client/agent's responsibility** to orchestrate
- This keeps the API simple and the isolation model clear
### Disposition Traits (3-trait system)
- **Skepticism** (1-5): How skeptical vs trusting when forming opinions
- **Literalism** (1-5): How literally to interpret information
- **Empathy** (1-5): How much to consider emotional context
- These influence the `reflect` operation, not `recall`
- Background info also only affects `reflect` (opinion formation)
## Multi-Bank Architecture Patterns
See [hindsight-docs/docs/cookbook/](./hindsight-docs/docs/cookbook/) for detailed guides:
- **Per-User Memory**: One bank per user, simplest pattern
- **Support Agent + Shared Knowledge**: User bank + shared docs bank, client orchestrates
## Developer Guide
### Running the API Server
```bash
# From project root
./scripts/dev/start-api.sh
# With options
./scripts/dev/start-api.sh --reload --port 8888 --log-level debug
```
### Running Tests
```bash
# API tests
cd hindsight-api
uv run pytest tests/
# Specific test
uv run pytest tests/test_http_api_integration.py -v
```
### Generating OpenAPI Spec
After changing API endpoints, regenerate the OpenAPI spec and docs:
```bash
./scripts/generate-openapi.sh
```
This will:
1. Generate `openapi.json` at project root
2. Copy to `hindsight-docs/openapi.json`
3. Regenerate API reference documentation
### Generating API Clients
After updating the OpenAPI spec, regenerate all clients:
```bash
./scripts/generate-clients.sh
```
This generates:
- **Rust client**: `hindsight-clients/rust/` (via progenitor in build.rs)
- **Python client**: `hindsight-clients/python/` (via openapi-generator Docker)
- **TypeScript client**: `hindsight-clients/typescript/` (via @hey-api/openapi-ts)
Note: The maintained wrapper `hindsight_client.py` and `README.md` are preserved during regeneration.
### Running the Documentation Site
```bash
./scripts/dev/start-docs.sh
```
### Running the Control Plane
```bash
./scripts/dev/start-control-plane.sh
```
## Code Style
### Python (hindsight-api)
- Use `uv` for package management
- Async throughout (asyncpg, async FastAPI endpoints)
- Pydantic models for request/response validation
- No py files at project root - maintain clean directory structure
### TypeScript (control-plane, clients)
- Next.js with App Router for control plane
- Tailwind CSS with shadcn/ui components
### Rust (CLI)
- Async with tokio
- reqwest for HTTP client
- progenitor for API client generation
## Database
- PostgreSQL with pgvector extension
- Schema managed via Alembic migrations in `hindsight-api/alembic/`, db migrations happen during api startup, no manual commands
- Key tables: `banks`, `memory_units`, `documents`, `entities`, `entity_links`
# Branding
## Colors
- Primary: gradient from #0074d9 to #009296
See [CLAUDE.md](./CLAUDE.md) for project documentation and coding conventions.
+34 -3
View File
@@ -7,8 +7,7 @@ This file provides guidance to Claude Code (claude.ai/code) when working with co
Hindsight is an agent memory system that provides long-term memory for AI agents using biomimetic data structures. Memories are organized as:
- **World facts**: General knowledge ("The sky is blue")
- **Experience facts**: Personal experiences ("I visited Paris in 2023")
- **Opinion facts**: Beliefs with confidence scores ("Paris is beautiful" - 0.9 confidence)
- **Observations**: Complex mental models derived from reflection
- **Mental models**: Consolidated knowledge synthesized from facts ("User prefers functional programming patterns")
## Development Commands
@@ -101,7 +100,7 @@ cd hindsight-control-plane && npm run dev
Main operations:
- **Retain**: Store memories, extracts facts/entities/relationships
- **Recall**: Retrieve memories via 4 parallel strategies (semantic, BM25, graph, temporal) + reranking
- **Reflect**: Deep analysis forming new opinions/observations (disposition-aware)
- **Reflect**: Disposition-aware reasoning using memories and mental models.
### Database
PostgreSQL with pgvector. Schema managed via Alembic migrations in `hindsight-api/hindsight_api/alembic/`. Migrations run automatically on API startup.
@@ -199,6 +198,38 @@ When adding or modifying parameters in the dataplane API (hindsight-api), you mu
- Pydantic models for request/response
- Ruff for linting (line-length 120)
- No Python files at project root - maintain clean directory structure
- **Never use multi-item tuple return values** - prefer dataclass or Pydantic model for structured returns
### Type Safety with Pydantic Models
**NEVER use raw `dict` types for structured data.** Always use Pydantic models:
- Use Pydantic `BaseModel` for all data structures passed between functions
- Add `@field_validator` for type coercion (e.g., ensuring datetimes are timezone-aware)
- Avoid `dict.get()` patterns - use typed model attributes instead
- Parse external data (JSON, API responses) into Pydantic models at the boundary
- This catches type errors at parse time, not deep in business logic
```python
# BAD - error-prone dict access
def process(data: dict) -> str:
return data.get("name", "") # No validation, silent failures
# GOOD - typed and validated
class UserData(BaseModel):
name: str
created_at: datetime
@field_validator("created_at", mode="before")
@classmethod
def ensure_tz_aware(cls, v):
if isinstance(v, str):
v = datetime.fromisoformat(v.replace("Z", "+00:00"))
if v.tzinfo is None:
return v.replace(tzinfo=timezone.utc)
return v
def process(data: UserData) -> str:
return data.name # Type-safe, validated at construction
```
### TypeScript Style
- Next.js App Router for control plane
+53 -40
View File
@@ -1,6 +1,6 @@
<div align="center">
![Hindsight Banner](./hindsight-docs/static/img/banner.svg)
![Hindsight Banner](./hindsight-docs/static/img/hindsight-github-banner.png)
[Documentation](https://hindsight.vectorize.io) • [Paper](https://arxiv.org/abs/2512.12818) • [Cookbook](https://hindsight.vectorize.io/cookbook) • [Hindsight Cloud](https://vectorize.io/hindsight/cloud)
@@ -17,55 +17,31 @@
## What is Hindsight?
Hindsight™ is an agent memory system built to create smarter agents that learn over time. It eliminates the shortcomings of alternative techniques such as RAG and knowledge graph and delivers state-of-the-art performance on long term memory tasks.
Hindsight™ is an agent memory system built to create smarter agents that learn over time. Most agent memory systems focus on recalling conversation history. Hindsight is focused on making agents that learn, not just remember.
Hindsight addresses common challenges that have frustrated AI engineers building agents to automate tasks and assist users with conversational interfaces. Many of these challenges stem directly from a lack of memory.
- **Inconsistency:** Agents complete tasks successfully one time, then fail when asked to complete the same task again. Memory gives the agent a mechanism to remember what worked and what didn't and to use that information to reduce errors and improve consistency.
- **Hallucinations:** Long term memory can be seeded with external knowledge to ground agent behavior in reliable sources to augment training data.
- **Cognitive Overload:** As workflows get complex, retrievals, tool calls, user messages and agent responses can grow to fill the context window leading to context rot. Short term memory optimization allows agents to reduce tokens and focus context by removing irrelevant details.
<video src="https://github.com/user-attachments/assets/923b798d-3581-4897-bb62-9cfa5a931682" controls></video>
## How is Hindsight Different From Other Memory Systems?
![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:
- **World:** Facts about the world ("The stove gets hot")
- **Experiences:** Agent's own experiences ("I touched the stove and it really hurt")
- **Opinion:** Beliefs with confidence scores ("I shouldn't touch the stove again" - .99 confidence)
- **Observation:** Complex mental models derived by reflecting on facts and experiences ("Curling irons, ovens, and fire are also hot. I shouldn't touch those either.")
Memories in Hindsight are stored in banks (i.e. memory banks). When memories are added to Hindsight, they are pushed into either the world facts or experiences memory pathway. They are then represented as a combination of entities, relationships, and time series with sparse/dense vector representations to aid in later recall.
Hindsight provides three simple methods to interact with the system:
- **Retain:** Provide information to Hindsight that you want it to remember
- **Recall:** Retrieve memories from Hindsight
- **Reflect:** Reflect on memories and experiences to generate new observations and insights from existing memories.
### Agent Memory That Learns
A key goal of Hindsight is to build agent memory that enables agents to learn and improve over time. This is the role of the `reflect` operation which provides the agent to form broader opinions and observations over time.
For example, imagine a product support agent that is helping a user troubleshoot a problem. It uses a `search-documentation` tool it found on an MCP server. Later in the conversation, the agent discovers that the documentation returned from the tool wasn't for the product the user was asking about. The agent now has an experience in its memory bank. And just like humans, we want that agent to learn from its experience.
As the agent gains more experiences, `reflect` allows the agent to form observations about what worked, what didn't, and what to do differently the next time it encounters a similar task.
---
It eliminates the shortcomings of alternative techniques such as RAG and knowledge graph and delivers state-of-the-art performance on long term memory tasks.
## Memory Performance & Accuracy
Hindsight has achieved state-of-the-art performance on the LongMemEval benchmark, widely used to assess memory system performance across a variety of conversational
AI scenarios. The current reported performance of Hindsight and other agent memory solutions as of December 2025 is shown here:
Hindsight is the most accurate agent memory system ever tested according to benchmark performance. It has achieved state-of-the-art performance on the LongMemEval benchmark, widely used to assess memory system performance across a variety of conversational AI scenarios. The current reported performance of Hindsight and other agent memory solutions as of January 2026 is shown here:
![Overview](./hindsight-docs/static/img/hindsight-bench.jpg)
The benchmark performance data for Hindsight and GPT-4o (full context) have been reproduced by research collaborators at the Virginia Tech [Sanghani Center for Artificial Intelligence and Data Analytics](https://sanghani.cs.vt.edu/) and The Washington Post. Other scores are self-reported by software vendors.
The benchmark performance data for Hindsight has been independently reproduced by research collaborators at the Virginia Tech [Sanghani Center for Artificial Intelligence and Data Analytics](https://sanghani.cs.vt.edu/) and The Washington Post. Other scores are self-reported by software vendors.
A thorough examination of the techniques implemented in Hindsight and detailed breakdowns of benchmark performance are [available on arXiv](https://arxiv.org/abs/2512.12818). This research is currently being prepared for conference submission and the wider peer review process.
Hindsight is being used in production at Fortune 500 enterprises and by a growing number of AI startups.
## 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.
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.
![Hindsight Banner](./hindsight-docs/static/img/migration-code.png)
The benchmark results from this research can be inspected in our [visual benchmark explorer](https://hindsight-benchmarks.vercel.app). As additional improvements are made to Hindsight, new benchmark data will be available for review using this same tool.
## Quick Start
@@ -148,8 +124,45 @@ await client.recall('my-bank', 'What does Alice like?');
---
## Use Cases
Hindsight is built to support conversational AI agents as well as agents that are intended to perform tasks autonomously. The ideal use case for Hindsight are agents that require a blend of these features such as AI employees that need to handle open-ended tasks, change behavior based on user feedback, and learn to perform complex tasks to automate work at a level that approximates a human work. Hindsight can be used with simple AI workflows like those built with n8n and other similar tools, but may be overkill for such applications.
### Per-User Memories and Chat History
One of the simpler use cases you can use Hindsight for is to personalize AI chatbots and other conversational agents by storing and recalling memories associated with individual users.
The requirements for this use case usually look something like this:
![Per-User Memories](./hindsight-docs/static/img/per-user-memory-requirements.png)
<video src="https://github.com/user-attachments/assets/4805e8e1-e7d1-47c6-a4f8-2344a5ec8906" controls></video>
Satisfying these requirements in Hindsight is straightforward. When new user inputs and tool calls are ingested into Hindsight using the retain operation, custom metadata can be used to enrich the new memories. Metadata provides a convenient way to isolate memories that need to be restricted to a given user. Once these are fed into the retain operation, any raw memories and mental models that get created can be filtered when retrieving relevant memories.
![Per-User Memories](./hindsight-docs/static/img/per-user-memory-howto.png)
---
## Architecture & Operations
![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:
- **World:** Facts about the world ("The stove gets hot")
- **Experiences:** Agent's own experiences ("I touched the stove and it really hurt")
- **Mental Models:** Learned understanding of the agent's world formed by reflecting on raw memories and experiences.
Memories in Hindsight are stored in banks (i.e. memory banks). When memories are added to Hindsight, they are pushed into either the world facts or experiences memory pathway. They are then represented as a combination of entities, relationships, and time series with sparse/dense vector representations to aid in later recall.
Hindsight provides three simple methods to interact with the system:
- **Retain:** Provide information to Hindsight that you want it to remember
- **Recall:** Retrieve memories from Hindsight
- **Reflect:** Reflect on memories and experiences to generate new observations and insights from existing memories.
### Retain
The `retain` operation is used to push new memories into Hindsight. It tells Hindsight to _retain_ the information you pass in as an input.
@@ -208,7 +221,7 @@ The final output is trimmed as needed to fit within the token limit.
### Reflect
The reflect operation is used to perform a more thorough analysis of existing memories. This allows the agent to form new connections between memories which are then persisted as opinions and/or observations. When building agents, the reflect operation is a key capability to enable the agent to learn from its experiences.
The reflect operation is used to perform a more thorough analysis of existing memories. This allows the agent to form new connections between memories and build a more thorough understanding of its world.
For example, the `reflect` operation can be used to support use cases such as:
+2 -2
View File
@@ -2,8 +2,8 @@ apiVersion: v2
name: hindsight
description: Hindsight helm chart
type: application
version: 0.3.0
appVersion: "0.3.0"
version: 0.4.1
appVersion: "0.4.1"
keywords:
- ai
- memory
+16
View File
@@ -80,6 +80,22 @@ Control plane selector labels
app.kubernetes.io/component: control-plane
{{- end }}
{{/*
Worker labels
*/}}
{{- define "hindsight.worker.labels" -}}
{{ include "hindsight.labels" . }}
app.kubernetes.io/component: worker
{{- end }}
{{/*
Worker selector labels
*/}}
{{- define "hindsight.worker.selectorLabels" -}}
{{ include "hindsight.selectorLabels" . }}
app.kubernetes.io/component: worker
{{- end }}
{{/*
Create the name of the service account to use
*/}}
@@ -55,6 +55,11 @@ spec:
{{- end }}
- name: HINDSIGHT_API_DATABASE_URL
value: {{ include "hindsight.databaseUrl" . | quote }}
{{- /* Disable internal worker when dedicated workers are enabled */}}
{{- if .Values.worker.enabled }}
- name: HINDSIGHT_API_WORKER_ENABLED
value: "false"
{{- end }}
{{- range $key, $value := .Values.api.env }}
- name: {{ $key }}
value: {{ $value | quote }}
@@ -0,0 +1,25 @@
{{- if .Values.worker.enabled }}
apiVersion: v1
kind: Service
metadata:
name: {{ include "hindsight.fullname" . }}-worker
labels:
{{- include "hindsight.worker.labels" . | nindent 4 }}
{{- if .Values.podAnnotations }}
annotations:
{{- /* Common Prometheus annotations for metrics scraping */}}
prometheus.io/scrape: "true"
prometheus.io/port: {{ .Values.worker.service.port | quote }}
prometheus.io/path: "/metrics"
{{- end }}
spec:
# Headless service for StatefulSet (enables stable DNS names like worker-0.worker.namespace)
clusterIP: None
ports:
- port: {{ .Values.worker.service.port }}
targetPort: {{ .Values.worker.service.targetPort }}
protocol: TCP
name: http
selector:
{{- include "hindsight.worker.selectorLabels" . | nindent 4 }}
{{- end }}
@@ -0,0 +1,110 @@
{{- if .Values.worker.enabled }}
apiVersion: apps/v1
kind: StatefulSet
metadata:
name: {{ include "hindsight.fullname" . }}-worker
labels:
{{- include "hindsight.worker.labels" . | nindent 4 }}
spec:
serviceName: {{ include "hindsight.fullname" . }}-worker
replicas: {{ .Values.worker.replicaCount }}
selector:
matchLabels:
{{- include "hindsight.worker.selectorLabels" . | nindent 6 }}
template:
metadata:
annotations:
{{- if not .Values.existingSecret }}
checksum/secret: {{ include (print $.Template.BasePath "/secret.yaml") . | sha256sum }}
{{- end }}
{{- with .Values.podAnnotations }}
{{- toYaml . | nindent 8 }}
{{- end }}
labels:
{{- include "hindsight.worker.selectorLabels" . | nindent 8 }}
spec:
{{- if .Values.serviceAccount.create }}
serviceAccountName: {{ include "hindsight.serviceAccountName" . }}
{{- end }}
securityContext:
{{- toYaml .Values.podSecurityContext | nindent 8 }}
containers:
- name: worker
securityContext:
{{- toYaml .Values.securityContext | nindent 10 }}
image: "{{ .Values.worker.image.repository }}:{{ .Values.worker.image.tag | default .Values.version }}"
imagePullPolicy: {{ .Values.worker.image.pullPolicy }}
command: ["hindsight-worker"]
ports:
- name: http
containerPort: {{ .Values.worker.service.targetPort }}
protocol: TCP
{{- if .Values.existingSecret }}
envFrom:
- secretRef:
name: {{ .Values.existingSecret }}
{{- end }}
env:
{{- /* POSTGRES_PASSWORD must be defined before DATABASE_URL for $(VAR) interpolation */}}
{{- if not .Values.postgresql.enabled }}
- name: POSTGRES_PASSWORD
valueFrom:
secretKeyRef:
name: {{ include "hindsight.secretName" . }}
key: postgres-password
{{- end }}
- name: HINDSIGHT_API_DATABASE_URL
value: {{ include "hindsight.databaseUrl" . | quote }}
{{- /* Worker ID uses pod name (StatefulSet provides stable names like worker-0, worker-1) */}}
- name: HINDSIGHT_API_WORKER_ID
valueFrom:
fieldRef:
fieldPath: metadata.name
{{- /* Inherit LLM config from api.env */}}
{{- range $key, $value := .Values.api.env }}
- name: {{ $key }}
value: {{ $value | quote }}
{{- end }}
{{- /* Worker-specific env vars */}}
{{- range $key, $value := .Values.worker.env }}
- name: {{ $key }}
value: {{ $value | quote }}
{{- end }}
{{- /* Only use secrets when not using existingSecret */}}
{{- if not .Values.existingSecret }}
{{- /* Inherit secrets from api.secrets */}}
{{- range $key, $value := .Values.api.secrets }}
- name: {{ $key }}
valueFrom:
secretKeyRef:
name: {{ include "hindsight.secretName" $ }}
key: {{ $key }}
{{- end }}
{{- /* Worker-specific secrets (can override api.secrets) */}}
{{- range $key, $value := .Values.worker.secrets }}
- name: {{ $key }}
valueFrom:
secretKeyRef:
name: {{ include "hindsight.secretName" $ }}
key: {{ $key }}
{{- end }}
{{- end }}
livenessProbe:
{{- toYaml .Values.worker.livenessProbe | nindent 10 }}
readinessProbe:
{{- toYaml .Values.worker.readinessProbe | nindent 10 }}
resources:
{{- toYaml .Values.worker.resources | nindent 10 }}
{{- with .Values.nodeSelector }}
nodeSelector:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.affinity }}
affinity:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.tolerations }}
tolerations:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- end }}
+57
View File
@@ -67,6 +67,63 @@ api:
# HINDSIGHT_API_LLM_API_KEY: "your-api-key"
# HINDSIGHT_API_LLM_BASE_URL: "https://api.groq.com/openai/v1"
# Worker settings (distributed task processing)
# When enabled, dedicated worker pods process tasks and the API's internal worker is disabled
worker:
enabled: false
replicaCount: 2
image:
repository: ghcr.io/vectorize-io/hindsight-api
pullPolicy: IfNotPresent
# tag defaults to .Values.version if not specified
service:
# Service for metrics scraping (headless for StatefulSet)
port: 8889
targetPort: 8889
# Resource limits and requests
resources:
limits:
cpu: 2000m
memory: 4Gi
requests:
cpu: 500m
memory: 1Gi
# Liveness and readiness probes
livenessProbe:
httpGet:
path: /health
port: 8889
initialDelaySeconds: 30
periodSeconds: 10
timeoutSeconds: 5
failureThreshold: 3
readinessProbe:
httpGet:
path: /health
port: 8889
initialDelaySeconds: 10
periodSeconds: 5
timeoutSeconds: 3
failureThreshold: 3
# Worker-specific environment variables
env:
# Poll interval in milliseconds (how often to check for new tasks)
HINDSIGHT_API_WORKER_POLL_INTERVAL_MS: "500"
# Number of tasks to claim per poll cycle
HINDSIGHT_API_WORKER_BATCH_SIZE: "10"
# Max retries before marking a task as failed
HINDSIGHT_API_WORKER_MAX_RETRIES: "3"
# HTTP port for metrics/health (matches service.targetPort)
HINDSIGHT_API_WORKER_HTTP_PORT: "8889"
# Secret environment variables (inherited from api.secrets if not specified)
secrets: {}
# Image settings for control plane
controlPlane:
enabled: true
+1 -1
View File
@@ -46,4 +46,4 @@ __all__ = [
"RemoteTEICrossEncoder",
"LLMConfig",
]
__version__ = "0.1.0"
__version__ = "0.4.1"
+59
View File
@@ -244,6 +244,65 @@ def run_db_migration(
typer.echo("Database migrations completed successfully")
async def _decommission_worker(db_url: str, worker_id: str, schema: str = "public") -> int:
"""Release all tasks owned by a worker, setting them back to pending status."""
is_pg0, instance_name, _ = parse_pg0_url(db_url)
if is_pg0:
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
resolved_url = await resolve_database_url(db_url)
conn = await asyncpg.connect(resolved_url)
try:
table = _fq_table("async_operations", schema)
result = await conn.fetch(
f"""
UPDATE {table}
SET status = 'pending', worker_id = NULL, claimed_at = NULL, updated_at = now()
WHERE worker_id = $1 AND status = 'processing'
RETURNING operation_id
""",
worker_id,
)
return len(result)
finally:
await conn.close()
@app.command(name="decommission-worker")
def decommission_worker(
worker_id: str = typer.Argument(..., help="Worker ID to decommission"),
schema: str = typer.Option("public", "--schema", "-s", help="Database schema"),
yes: bool = typer.Option(False, "--yes", "-y", help="Skip confirmation prompt"),
):
"""Release all tasks owned by a worker (sets status back to pending).
Use this command when a worker has crashed or been removed without graceful shutdown.
All tasks that were being processed by the worker will be released back to the queue
so other workers can pick them up.
"""
config = HindsightConfig.from_env()
if not config.database_url:
typer.echo("Error: Database URL not configured.", err=True)
typer.echo("Set HINDSIGHT_API_DATABASE_URL environment variable.", err=True)
raise typer.Exit(1)
if not yes:
typer.confirm(
f"This will release all tasks owned by worker '{worker_id}' back to pending. Continue?",
abort=True,
)
typer.echo(f"Decommissioning worker '{worker_id}' (schema: {schema})...")
count = asyncio.run(_decommission_worker(config.database_url, worker_id, schema))
if count > 0:
typer.echo(f"Released {count} task(s) from worker '{worker_id}'")
else:
typer.echo(f"No tasks found for worker '{worker_id}'")
def main():
app()
@@ -0,0 +1,95 @@
"""mental_model_versions
Revision ID: j5e6f7g8h9i0
Revises: i4d5e6f7g8h9
Create Date: 2026-01-16 00:00:00.000000
This migration adds versioning support for mental models:
1. Creates mental_model_versions table to store observation snapshots
2. Adds version column to mental_models for tracking current version
This enables changelog/diff functionality for mental model observations.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "j5e6f7g8h9i0"
down_revision: str | Sequence[str] | None = "i4d5e6f7g8h9"
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:
"""Create mental_model_versions table and add version tracking."""
schema = _get_schema_prefix()
# Create mental_model_versions table for storing observation snapshots
op.execute(f"""
CREATE TABLE {schema}mental_model_versions (
id SERIAL PRIMARY KEY,
mental_model_id VARCHAR(64) NOT NULL,
bank_id VARCHAR(64) NOT NULL,
version INT NOT NULL,
observations JSONB NOT NULL DEFAULT '{{"observations": []}}'::jsonb,
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
FOREIGN KEY (mental_model_id, bank_id)
REFERENCES {schema}mental_models(id, bank_id) ON DELETE CASCADE,
UNIQUE (mental_model_id, bank_id, version)
)
""")
# Index for efficient version queries (get latest, list versions)
op.execute(f"""
CREATE INDEX idx_mental_model_versions_lookup
ON {schema}mental_model_versions(mental_model_id, bank_id, version DESC)
""")
# Add version column to mental_models to track current version
op.execute(f"""
ALTER TABLE {schema}mental_models
ADD COLUMN IF NOT EXISTS version INT NOT NULL DEFAULT 0
""")
# Migrate existing mental models: create version 1 for any that have observations
op.execute(f"""
INSERT INTO {schema}mental_model_versions (mental_model_id, bank_id, version, observations, created_at)
SELECT id, bank_id, 1, observations, COALESCE(last_updated, created_at)
FROM {schema}mental_models
WHERE observations IS NOT NULL
AND observations != '{{"observations": []}}'::jsonb
AND (observations->'observations') IS NOT NULL
AND jsonb_array_length(observations->'observations') > 0
""")
# Update version to 1 for migrated mental models
op.execute(f"""
UPDATE {schema}mental_models
SET version = 1
WHERE observations IS NOT NULL
AND observations != '{{"observations": []}}'::jsonb
AND (observations->'observations') IS NOT NULL
AND jsonb_array_length(observations->'observations') > 0
""")
def downgrade() -> None:
"""Remove mental_model_versions table and version column."""
schema = _get_schema_prefix()
# Drop index
op.execute(f"DROP INDEX IF EXISTS {schema}idx_mental_model_versions_lookup")
# Drop versions table
op.execute(f"DROP TABLE IF EXISTS {schema}mental_model_versions")
# Remove version column from mental_models
op.execute(f"ALTER TABLE {schema}mental_models DROP COLUMN IF EXISTS version")
@@ -0,0 +1,58 @@
"""add_directive_subtype
Revision ID: k6f7g8h9i0j1
Revises: j5e6f7g8h9i0
Create Date: 2026-01-16 00:00:00.000000
This migration adds 'directive' to the mental_models subtype constraint.
Directives are hard rules with user-provided observations that the reflect agent must follow.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "k6f7g8h9i0j1"
down_revision: str | Sequence[str] | None = "j5e6f7g8h9i0"
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:
"""Add 'directive' to mental_models subtype constraint."""
schema = _get_schema_prefix()
# Drop existing constraint
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS ck_mental_models_subtype")
# Create new constraint with 'directive' added
op.execute(f"""
ALTER TABLE {schema}mental_models
ADD CONSTRAINT ck_mental_models_subtype
CHECK (subtype IN ('structural', 'emergent', 'pinned', 'learned', 'directive'))
""")
def downgrade() -> None:
"""Remove 'directive' from mental_models subtype constraint."""
schema = _get_schema_prefix()
# First delete any directives (cannot downgrade if they exist)
op.execute(f"DELETE FROM {schema}mental_models WHERE subtype = 'directive'")
# Drop constraint with directive
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS ck_mental_models_subtype")
# Recreate original constraint without directive
op.execute(f"""
ALTER TABLE {schema}mental_models
ADD CONSTRAINT ck_mental_models_subtype
CHECK (subtype IN ('structural', 'emergent', 'pinned', 'learned'))
""")
@@ -0,0 +1,109 @@
"""add_worker_columns
Revision ID: l7g8h9i0j1k2
Revises: k6f7g8h9i0j1
Create Date: 2026-01-19 00:00:00.000000
This migration adds columns to async_operations for distributed worker support:
- worker_id: ID of the worker that claimed the task
- claimed_at: When the task was claimed
- retry_count: Number of retry attempts
- task_payload: The serialized task dictionary
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import context, op
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = "l7g8h9i0j1k2"
down_revision: str | Sequence[str] | None = "k6f7g8h9i0j1"
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:
"""Add worker columns to async_operations."""
schema = _get_schema_prefix()
# Add worker_id column (ID of worker that claimed the task)
op.add_column(
"async_operations",
sa.Column("worker_id", sa.Text(), nullable=True),
schema=context.config.get_main_option("target_schema") or None,
)
# Add claimed_at column (when task was claimed by worker)
op.add_column(
"async_operations",
sa.Column("claimed_at", postgresql.TIMESTAMP(timezone=True), nullable=True),
schema=context.config.get_main_option("target_schema") or None,
)
# Add retry_count column (number of retry attempts)
op.add_column(
"async_operations",
sa.Column("retry_count", sa.Integer(), server_default="0", nullable=False),
schema=context.config.get_main_option("target_schema") or None,
)
# Add task_payload column (serialized task dictionary)
op.add_column(
"async_operations",
sa.Column(
"task_payload",
postgresql.JSONB(astext_type=sa.Text()),
nullable=True,
),
schema=context.config.get_main_option("target_schema") or None,
)
# Add index for efficient worker polling (pending tasks ordered by creation time)
op.execute(
f"CREATE INDEX idx_async_operations_pending_claim ON {schema}async_operations (status, created_at) "
f"WHERE status = 'pending' AND task_payload IS NOT NULL"
)
# Add index for finding tasks by worker_id (for decommissioning)
op.execute(
f"CREATE INDEX idx_async_operations_worker_id ON {schema}async_operations (worker_id) WHERE worker_id IS NOT NULL"
)
def downgrade() -> None:
"""Remove worker columns from async_operations."""
schema = _get_schema_prefix()
# Drop indexes
op.execute(f"DROP INDEX IF EXISTS {schema}idx_async_operations_pending_claim")
op.execute(f"DROP INDEX IF EXISTS {schema}idx_async_operations_worker_id")
# Drop columns
op.drop_column(
"async_operations",
"task_payload",
schema=context.config.get_main_option("target_schema") or None,
)
op.drop_column(
"async_operations",
"retry_count",
schema=context.config.get_main_option("target_schema") or None,
)
op.drop_column(
"async_operations",
"claimed_at",
schema=context.config.get_main_option("target_schema") or None,
)
op.drop_column(
"async_operations",
"worker_id",
schema=context.config.get_main_option("target_schema") or None,
)
@@ -0,0 +1,41 @@
"""mental_model_id_to_text
Revision ID: m8h9i0j1k2l3
Revises: l7g8h9i0j1k2
Create Date: 2026-01-19 00:00:00.000000
This migration changes the mental_models.id column from VARCHAR(64) to TEXT
to support longer model IDs (e.g., entity names that exceed 64 characters).
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "m8h9i0j1k2l3"
down_revision: str | Sequence[str] | None = "l7g8h9i0j1k2"
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:
"""Change mental_models.id from VARCHAR(64) to TEXT."""
schema = _get_schema_prefix()
# Alter the id column type from VARCHAR(64) to TEXT
op.execute(f"ALTER TABLE {schema}mental_models ALTER COLUMN id TYPE TEXT")
def downgrade() -> None:
"""Revert mental_models.id from TEXT to VARCHAR(64)."""
schema = _get_schema_prefix()
# Note: This may fail if any id values exceed 64 characters
op.execute(f"ALTER TABLE {schema}mental_models ALTER COLUMN id TYPE VARCHAR(64)")
@@ -0,0 +1,134 @@
"""learnings_and_pinned_reflections
Revision ID: n9i0j1k2l3m4
Revises: m8h9i0j1k2l3
Create Date: 2026-01-21 00:00:00.000000
This migration:
1. Creates the 'learnings' table for automatic bottom-up consolidation
2. Creates the 'pinned_reflections' table for user-curated living documents
3. Adds consolidation tracking columns to the 'banks' table
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "n9i0j1k2l3m4"
down_revision: str | Sequence[str] | None = "m8h9i0j1k2l3"
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:
"""Create learnings and pinned_reflections tables."""
schema = _get_schema_prefix()
# 1. Create learnings table
op.execute(f"""
CREATE TABLE {schema}learnings (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
bank_id VARCHAR(64) NOT NULL,
text TEXT NOT NULL,
proof_count INT NOT NULL DEFAULT 1,
history JSONB DEFAULT '[]'::jsonb,
mission_context VARCHAR(64),
pre_mission_change BOOLEAN DEFAULT FALSE,
embedding vector(384),
tags VARCHAR[] DEFAULT ARRAY[]::VARCHAR[],
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
)
""")
# Add foreign key constraint
op.execute(f"""
ALTER TABLE {schema}learnings
ADD CONSTRAINT fk_learnings_bank_id
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
""")
# Indexes for learnings
op.execute(f"CREATE INDEX idx_learnings_bank_id ON {schema}learnings(bank_id)")
op.execute(f"""
CREATE INDEX idx_learnings_embedding ON {schema}learnings
USING hnsw (embedding vector_cosine_ops)
""")
op.execute(f"CREATE INDEX idx_learnings_tags ON {schema}learnings USING GIN(tags)")
# Full-text search for learnings
op.execute(f"""
ALTER TABLE {schema}learnings ADD COLUMN search_vector tsvector
GENERATED ALWAYS AS (to_tsvector('english', text)) STORED
""")
op.execute(f"CREATE INDEX idx_learnings_text_search ON {schema}learnings USING gin(search_vector)")
# 2. Create pinned_reflections table
op.execute(f"""
CREATE TABLE {schema}pinned_reflections (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
bank_id VARCHAR(64) NOT NULL,
name VARCHAR(256) NOT NULL,
source_query TEXT NOT NULL,
content TEXT NOT NULL,
embedding vector(384),
tags VARCHAR[] DEFAULT ARRAY[]::VARCHAR[],
last_refreshed_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
)
""")
# Add foreign key constraint
op.execute(f"""
ALTER TABLE {schema}pinned_reflections
ADD CONSTRAINT fk_pinned_reflections_bank_id
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
""")
# Indexes for pinned_reflections
op.execute(f"CREATE INDEX idx_pinned_reflections_bank_id ON {schema}pinned_reflections(bank_id)")
op.execute(f"""
CREATE INDEX idx_pinned_reflections_embedding ON {schema}pinned_reflections
USING hnsw (embedding vector_cosine_ops)
""")
op.execute(f"CREATE INDEX idx_pinned_reflections_tags ON {schema}pinned_reflections USING GIN(tags)")
# Full-text search for pinned_reflections
op.execute(f"""
ALTER TABLE {schema}pinned_reflections ADD COLUMN search_vector tsvector
GENERATED ALWAYS AS (to_tsvector('english', COALESCE(name, '') || ' ' || content)) STORED
""")
op.execute(f"""
CREATE INDEX idx_pinned_reflections_text_search ON {schema}pinned_reflections
USING gin(search_vector)
""")
# 3. Add consolidation tracking columns to banks table
op.execute(f"""
ALTER TABLE {schema}banks
ADD COLUMN IF NOT EXISTS last_consolidated_at TIMESTAMP WITH TIME ZONE
""")
op.execute(f"""
ALTER TABLE {schema}banks
ADD COLUMN IF NOT EXISTS mission_changed_at TIMESTAMP WITH TIME ZONE
""")
def downgrade() -> None:
"""Drop learnings and pinned_reflections tables."""
schema = _get_schema_prefix()
# Drop tables
op.execute(f"DROP TABLE IF EXISTS {schema}learnings CASCADE")
op.execute(f"DROP TABLE IF EXISTS {schema}pinned_reflections CASCADE")
# Remove columns from banks
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS last_consolidated_at")
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS mission_changed_at")
@@ -0,0 +1,113 @@
"""migrate_mental_models_data
Revision ID: o0j1k2l3m4n5
Revises: n9i0j1k2l3m4
Create Date: 2026-01-21 00:00:00.000000
This migration:
1. Migrates existing 'pinned' mental models to the new 'pinned_reflections' table
2. Migrates existing 'learned' mental models to the new 'learnings' table
3. Deletes non-directive mental models (structural, emergent, pinned, learned)
4. Drops the mental_model_versions table (no longer used)
5. Adds a CHECK constraint that only 'directive' subtype is allowed
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "o0j1k2l3m4n5"
down_revision: str | Sequence[str] | None = "n9i0j1k2l3m4"
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:
"""Migrate data and clean up old mental models."""
schema = _get_schema_prefix()
# 1. Migrate 'pinned' mental models to pinned_reflections
# For pinned models, the first observation's content becomes the pinned reflection content
op.execute(f"""
INSERT INTO {schema}pinned_reflections (bank_id, name, source_query, content, tags, created_at)
SELECT
bank_id,
name,
description AS source_query,
COALESCE(
observations->'observations'->0->>'content',
description,
''
) AS content,
tags,
created_at
FROM {schema}mental_models
WHERE subtype = 'pinned'
ON CONFLICT DO NOTHING
""")
# 2. Migrate 'learned' mental models to learnings
# Each observation in a learned model becomes a separate learning
op.execute(f"""
INSERT INTO {schema}learnings (bank_id, text, proof_count, tags, created_at)
SELECT
mm.bank_id,
obs->>'content' AS text,
GREATEST(1, COALESCE(jsonb_array_length(obs->'evidence'), 1)) AS proof_count,
mm.tags,
mm.created_at
FROM {schema}mental_models mm,
LATERAL jsonb_array_elements(mm.observations->'observations') AS obs
WHERE mm.subtype = 'learned'
AND obs->>'content' IS NOT NULL
AND obs->>'content' != ''
ON CONFLICT DO NOTHING
""")
# 3. Delete all non-directive mental models (they've been migrated or are obsolete)
op.execute(f"""
DELETE FROM {schema}mental_models
WHERE subtype != 'directive'
""")
# 4. Drop the mental_model_versions table (no longer used)
op.execute(f"DROP TABLE IF EXISTS {schema}mental_model_versions CASCADE")
# 5. Drop old constraints and add new one that only allows 'directive'
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS ck_mental_models_subtype")
op.execute(f"""
ALTER TABLE {schema}mental_models
ADD CONSTRAINT ck_mental_models_subtype CHECK (subtype = 'directive')
""")
def downgrade() -> None:
"""Reverse the migration (data migration is one-way, so this just removes constraints)."""
schema = _get_schema_prefix()
# Remove the directive-only constraint
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS ck_mental_models_subtype")
# Re-create mental_model_versions table
op.execute(f"""
CREATE TABLE IF NOT EXISTS {schema}mental_model_versions (
id SERIAL PRIMARY KEY,
bank_id VARCHAR(64) NOT NULL,
model_id VARCHAR(128) NOT NULL,
version INT NOT NULL,
observations JSONB NOT NULL,
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
)
""")
op.execute(
f"CREATE INDEX IF NOT EXISTS idx_mm_versions_lookup ON {schema}mental_model_versions(bank_id, model_id, version DESC)"
)
# Note: Data migration cannot be reversed - pinned_reflections and learnings data remains
@@ -0,0 +1,194 @@
"""new_knowledge_architecture
Revision ID: p1k2l3m4n5o6
Revises: o0j1k2l3m4n5
Create Date: 2026-01-21 00:00:00.000000
This migration implements the new knowledge architecture:
1. Drops the 'learnings' table (mental models are now in memory_units)
2. Renames 'pinned_reflections' to 'reflections'
3. Drops the 'mental_models' table completely
4. Creates 'directives' table for hard rules
5. Adds mental model support columns to 'memory_units' (proof_count, source_memory_ids, history)
The new architecture:
- Directives: Hard rules in their own table
- Mental Models: Stored in memory_units with fact_type='mental_model'
- Reflections: User-curated documents (renamed from pinned_reflections)
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "p1k2l3m4n5o6"
down_revision: str | Sequence[str] | None = "o0j1k2l3m4n5"
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:
"""Implement new knowledge architecture."""
schema = _get_schema_prefix()
# 1. Drop the learnings table (mental models will be in memory_units)
op.execute(f"DROP TABLE IF EXISTS {schema}learnings CASCADE")
# 2. Rename pinned_reflections to reflections
op.execute(f"ALTER TABLE IF EXISTS {schema}pinned_reflections RENAME TO reflections")
# Rename indexes for reflections
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_pinned_reflections_bank_id RENAME TO idx_reflections_bank_id")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_pinned_reflections_embedding RENAME TO idx_reflections_embedding")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_pinned_reflections_tags RENAME TO idx_reflections_tags")
op.execute(
f"ALTER INDEX IF EXISTS {schema}idx_pinned_reflections_text_search RENAME TO idx_reflections_text_search"
)
# Rename foreign key constraint
op.execute(f"""
ALTER TABLE {schema}reflections
DROP CONSTRAINT IF EXISTS fk_pinned_reflections_bank_id
""")
op.execute(f"""
ALTER TABLE {schema}reflections
ADD CONSTRAINT fk_reflections_bank_id
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
""")
# 3. Drop the mental_models table completely
op.execute(f"DROP TABLE IF EXISTS {schema}mental_models CASCADE")
# 4. Create directives table
op.execute(f"""
CREATE TABLE {schema}directives (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
bank_id VARCHAR(64) NOT NULL,
name VARCHAR(256) NOT NULL,
content TEXT NOT NULL,
priority INT NOT NULL DEFAULT 0,
is_active BOOLEAN NOT NULL DEFAULT TRUE,
tags VARCHAR[] DEFAULT ARRAY[]::VARCHAR[],
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
)
""")
# Add foreign key and indexes for directives
op.execute(f"""
ALTER TABLE {schema}directives
ADD CONSTRAINT fk_directives_bank_id
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
""")
op.execute(f"CREATE INDEX idx_directives_bank_id ON {schema}directives(bank_id)")
op.execute(f"CREATE INDEX idx_directives_bank_active ON {schema}directives(bank_id, is_active)")
op.execute(f"CREATE INDEX idx_directives_tags ON {schema}directives USING GIN(tags)")
# 5. Add mental model support columns to memory_units
# proof_count: Number of memories that support this mental model
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD COLUMN IF NOT EXISTS proof_count INT DEFAULT 1
""")
# source_memory_ids: Array of memory IDs that consolidated into this mental model
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD COLUMN IF NOT EXISTS source_memory_ids UUID[] DEFAULT ARRAY[]::UUID[]
""")
# history: JSONB array tracking changes to mental models
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD COLUMN IF NOT EXISTS history JSONB DEFAULT '[]'::jsonb
""")
# Add index for finding mental models
op.execute(f"""
CREATE INDEX IF NOT EXISTS idx_memory_units_mental_models
ON {schema}memory_units(bank_id, fact_type)
WHERE fact_type = 'mental_model'
""")
# 6. Update fact_type check constraint to include 'mental_model'
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD CONSTRAINT memory_units_fact_type_check
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation', 'mental_model'))
""")
def downgrade() -> None:
"""Reverse the migration."""
schema = _get_schema_prefix()
# Restore original fact_type check constraint (without 'mental_model')
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD CONSTRAINT memory_units_fact_type_check
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation'))
""")
# Drop mental model columns from memory_units
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS proof_count")
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS source_memory_ids")
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS history")
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_mental_models")
# Drop directives table
op.execute(f"DROP TABLE IF EXISTS {schema}directives CASCADE")
# Rename reflections back to pinned_reflections
op.execute(f"ALTER TABLE IF EXISTS {schema}reflections RENAME TO pinned_reflections")
# Restore indexes
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_bank_id RENAME TO idx_pinned_reflections_bank_id")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_embedding RENAME TO idx_pinned_reflections_embedding")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_tags RENAME TO idx_pinned_reflections_tags")
op.execute(
f"ALTER INDEX IF EXISTS {schema}idx_reflections_text_search RENAME TO idx_pinned_reflections_text_search"
)
# Restore foreign key
op.execute(f"""
ALTER TABLE {schema}pinned_reflections
DROP CONSTRAINT IF EXISTS fk_reflections_bank_id
""")
op.execute(f"""
ALTER TABLE {schema}pinned_reflections
ADD CONSTRAINT fk_pinned_reflections_bank_id
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
""")
# Re-create learnings table
op.execute(f"""
CREATE TABLE IF NOT EXISTS {schema}learnings (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
bank_id VARCHAR(64) NOT NULL,
text TEXT NOT NULL,
proof_count INT NOT NULL DEFAULT 1,
history JSONB DEFAULT '[]'::jsonb,
mission_context VARCHAR(64),
pre_mission_change BOOLEAN DEFAULT FALSE,
embedding vector(384),
tags VARCHAR[] DEFAULT ARRAY[]::VARCHAR[],
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
)
""")
op.execute(f"""
ALTER TABLE {schema}learnings
ADD CONSTRAINT fk_learnings_bank_id
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
""")
# Note: mental_models table recreation is complex and would need separate handling
@@ -0,0 +1,50 @@
"""fix_mental_model_fact_type
Revision ID: q2l3m4n5o6p7
Revises: p1k2l3m4n5o6
Create Date: 2026-01-21 13:30:00.000000
Fix the fact_type check constraint to include 'mental_model'.
This is a fix for p1k2l3m4n5o6 which should have included this change.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "q2l3m4n5o6p7"
down_revision: str | Sequence[str] | None = "p1k2l3m4n5o6"
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:
"""Add 'mental_model' to the fact_type check constraint."""
schema = _get_schema_prefix()
# Drop the old constraint and add the new one with mental_model included
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD CONSTRAINT memory_units_fact_type_check
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation', 'mental_model'))
""")
def downgrade() -> None:
"""Remove 'mental_model' from the fact_type check constraint."""
schema = _get_schema_prefix()
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD CONSTRAINT memory_units_fact_type_check
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation'))
""")
@@ -0,0 +1,47 @@
"""Add reflect_response JSONB column to reflections
Revision ID: r3m4n5o6p7q8
Revises: q2l3m4n5o6p7
Create Date: 2026-01-21
This migration adds a reflect_response JSONB column to store the full
reflect API response payload, including based_on facts and trace data.
Note: Table was renamed from pinned_reflections to reflections in p1k2l3m4n5o6.
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "r3m4n5o6p7q8"
down_revision: str | Sequence[str] | None = "q2l3m4n5o6p7"
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:
"""Add reflect_response JSONB column to reflections."""
schema = _get_schema_prefix()
# Add reflect_response column to store the full reflect API response
op.execute(f"""
ALTER TABLE {schema}reflections
ADD COLUMN IF NOT EXISTS reflect_response JSONB
""")
def downgrade() -> None:
"""Remove reflect_response column from reflections."""
schema = _get_schema_prefix()
op.execute(f"""
ALTER TABLE {schema}reflections
DROP COLUMN IF EXISTS reflect_response
""")
@@ -0,0 +1,53 @@
"""Add consolidated_at column to memory_units for incremental consolidation tracking.
This allows consolidation to track progress at the memory level rather than
using a bank-level watermark. If consolidation crashes, already-processed
memories won't be reprocessed.
Revision ID: s4n5o6p7q8r9
Revises: r3m4n5o6p7q8
Create Date: 2025-01-22
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "s4n5o6p7q8r9"
down_revision: str | Sequence[str] | None = "r3m4n5o6p7q8"
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()
# Add consolidated_at column to memory_units
op.execute(
f"""
ALTER TABLE {schema}memory_units
ADD COLUMN IF NOT EXISTS consolidated_at TIMESTAMPTZ DEFAULT NULL
"""
)
# Create index for efficient querying of unconsolidated memories
op.execute(
f"""
CREATE INDEX IF NOT EXISTS idx_memory_units_unconsolidated
ON {schema}memory_units (bank_id, created_at)
WHERE consolidated_at IS NULL AND fact_type IN ('experience', 'world')
"""
)
def downgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_unconsolidated")
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS consolidated_at")
@@ -0,0 +1,134 @@
"""Rename mental_model fact_type to observation and reflections table to mental_models
Revision ID: t5o6p7q8r9s0
Revises: s4n5o6p7q8r9
Create Date: 2026-01-26
This migration implements the terminology rename:
1. mental_model (fact_type in memory_units) -> observation
2. reflections table -> mental_models table
The new terminology:
- Observations: Consolidated knowledge synthesized from facts (was mental_model)
- Mental Models: Stored reflect responses (was reflections)
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "t5o6p7q8r9s0"
down_revision: str | Sequence[str] | None = "s4n5o6p7q8r9"
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:
"""Rename mental_model -> observation and reflections -> mental_models."""
schema = _get_schema_prefix()
# 1. Update fact_type values: mental_model -> observation
op.execute(f"""
UPDATE {schema}memory_units
SET fact_type = 'observation'
WHERE fact_type = 'mental_model'
""")
# 2. Update the CHECK constraint - remove mental_model, keep observation
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD CONSTRAINT memory_units_fact_type_check
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation'))
""")
# 3. Rename the index for observations (was for mental_models)
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_mental_models")
op.execute(f"""
CREATE INDEX IF NOT EXISTS idx_memory_units_observations
ON {schema}memory_units(bank_id, fact_type)
WHERE fact_type = 'observation'
""")
# 4. Update the unconsolidated index to not filter by fact_type since observations
# are now the consolidated type
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_unconsolidated")
op.execute(f"""
CREATE INDEX IF NOT EXISTS idx_memory_units_unconsolidated
ON {schema}memory_units (bank_id, created_at)
WHERE consolidated_at IS NULL AND fact_type IN ('experience', 'world')
""")
# 5. Rename reflections table to mental_models
op.execute(f"ALTER TABLE IF EXISTS {schema}reflections RENAME TO mental_models")
# 6. Rename indexes for mental_models (was reflections)
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_bank_id RENAME TO idx_mental_models_bank_id")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_embedding RENAME TO idx_mental_models_embedding")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_tags RENAME TO idx_mental_models_tags")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_text_search RENAME TO idx_mental_models_text_search")
# 7. Rename foreign key constraint
op.execute(f"""
ALTER TABLE {schema}mental_models
DROP CONSTRAINT IF EXISTS fk_reflections_bank_id
""")
op.execute(f"""
ALTER TABLE {schema}mental_models
ADD CONSTRAINT fk_mental_models_bank_id
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
""")
def downgrade() -> None:
"""Reverse: observation -> mental_model and mental_models -> reflections."""
schema = _get_schema_prefix()
# 1. Rename mental_models table back to reflections
op.execute(f"ALTER TABLE IF EXISTS {schema}mental_models RENAME TO reflections")
# 2. Rename indexes back
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_mental_models_bank_id RENAME TO idx_reflections_bank_id")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_mental_models_embedding RENAME TO idx_reflections_embedding")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_mental_models_tags RENAME TO idx_reflections_tags")
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_mental_models_text_search RENAME TO idx_reflections_text_search")
# 3. Rename foreign key back
op.execute(f"""
ALTER TABLE {schema}reflections
DROP CONSTRAINT IF EXISTS fk_mental_models_bank_id
""")
op.execute(f"""
ALTER TABLE {schema}reflections
ADD CONSTRAINT fk_reflections_bank_id
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
""")
# 4. Update fact_type values: observation -> mental_model
op.execute(f"""
UPDATE {schema}memory_units
SET fact_type = 'mental_model'
WHERE fact_type = 'observation'
""")
# 5. Update the CHECK constraint back
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
op.execute(f"""
ALTER TABLE {schema}memory_units
ADD CONSTRAINT memory_units_fact_type_check
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation', 'mental_model'))
""")
# 6. Rename index back
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_observations")
op.execute(f"""
CREATE INDEX IF NOT EXISTS idx_memory_units_mental_models
ON {schema}memory_units(bank_id, fact_type)
WHERE fact_type = 'mental_model'
""")
@@ -0,0 +1,41 @@
"""Change mental_models.id from UUID to TEXT
Revision ID: u6p7q8r9s0t1
Revises: t5o6p7q8r9s0
Create Date: 2026-01-27
This migration changes the mental_models.id column from UUID to TEXT
to support user-defined text identifiers like 'team-communication' instead of UUIDs.
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "u6p7q8r9s0t1"
down_revision: str | Sequence[str] | None = "t5o6p7q8r9s0"
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:
"""Change mental_models.id from UUID to TEXT."""
schema = _get_schema_prefix()
# Change the id column type from UUID to TEXT
# Existing UUIDs will be converted to their string representation
op.execute(f"ALTER TABLE {schema}mental_models ALTER COLUMN id TYPE TEXT USING id::TEXT")
def downgrade() -> None:
"""Revert mental_models.id from TEXT to UUID."""
schema = _get_schema_prefix()
# Note: This will fail if any id values are not valid UUIDs
op.execute(f"ALTER TABLE {schema}mental_models ALTER COLUMN id TYPE UUID USING id::UUID")
@@ -0,0 +1,50 @@
"""Add max_tokens and trigger columns to mental_models
Revision ID: v7q8r9s0t1u2
Revises: u6p7q8r9s0t1
Create Date: 2026-01-27
This migration adds:
- max_tokens column: token limit for content generation during refresh
- trigger column: JSONB for trigger settings (e.g., refresh_after_consolidation)
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "v7q8r9s0t1u2"
down_revision: str | Sequence[str] | None = "u6p7q8r9s0t1"
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:
"""Add max_tokens and trigger columns to mental_models."""
schema = _get_schema_prefix()
op.execute(f"""
ALTER TABLE {schema}mental_models
ADD COLUMN IF NOT EXISTS max_tokens INT NOT NULL DEFAULT 2048
""")
# trigger column stores trigger settings as JSONB
# Default: refresh_after_consolidation = false (not "real time")
op.execute(f"""
ALTER TABLE {schema}mental_models
ADD COLUMN IF NOT EXISTS trigger JSONB NOT NULL DEFAULT '{{"refresh_after_consolidation": false}}'::jsonb
""")
def downgrade() -> None:
"""Remove max_tokens and trigger columns from mental_models."""
schema = _get_schema_prefix()
op.execute(f"ALTER TABLE {schema}mental_models DROP COLUMN IF EXISTS max_tokens")
op.execute(f"ALTER TABLE {schema}mental_models DROP COLUMN IF EXISTS trigger")
File diff suppressed because it is too large Load Diff
+10 -190
View File
@@ -1,4 +1,4 @@
"""Hindsight MCP Server implementation using FastMCP."""
"""Hindsight MCP Server implementation using FastMCP (HTTP transport)."""
import json
import logging
@@ -8,8 +8,7 @@ from contextvars import ContextVar
from fastmcp import FastMCP
from hindsight_api import MemoryEngine
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
from hindsight_api.models import RequestContext
from hindsight_api.mcp_tools import MCPToolsConfig, register_mcp_tools
# Configure logging from HINDSIGHT_API_LOG_LEVEL environment variable
_log_level_str = os.environ.get("HINDSIGHT_API_LOG_LEVEL", "info").lower()
@@ -52,194 +51,15 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
# Use stateless_http=True for Claude Code compatibility
mcp = FastMCP("hindsight-mcp-server", stateless_http=True)
@mcp.tool()
async def retain(
content: str,
context: str = "general",
async_processing: bool = True,
bank_id: str | None = None,
) -> str:
"""
Store important information to long-term memory.
# Configure and register tools using shared module
config = MCPToolsConfig(
bank_id_resolver=get_current_bank_id,
include_bank_id_param=True, # HTTP MCP supports multi-bank via parameter
tools=None, # All tools
retain_fire_and_forget=False, # HTTP MCP supports sync/async modes
)
Use this tool PROACTIVELY whenever the user shares:
- Personal facts, preferences, or interests
- Important events or milestones
- User history, experiences, or background
- Decisions, opinions, or stated preferences
- Goals, plans, or future intentions
- Relationships or people mentioned
- Work context, projects, or responsibilities
Args:
content: The fact/memory to store (be specific and include relevant details)
context: Category for the memory (e.g., 'preferences', 'work', 'hobbies', 'family'). Default: 'general'
async_processing: If True, queue for background processing and return immediately. If False, wait for completion. Default: True
bank_id: Optional bank to store in (defaults to session bank). Use for cross-bank operations.
"""
try:
target_bank = bank_id or get_current_bank_id()
if target_bank is None:
return "Error: No bank_id configured"
contents = [{"content": content, "context": context}]
if async_processing:
# Queue for background processing and return immediately
result = await memory.submit_async_retain(
bank_id=target_bank, contents=contents, request_context=RequestContext()
)
return f"Memory queued for background processing (operation_id: {result.get('operation_id', 'N/A')})"
else:
# Wait for completion
await memory.retain_batch_async(
bank_id=target_bank,
contents=contents,
request_context=RequestContext(),
)
return f"Memory stored successfully in bank '{target_bank}'"
except Exception as e:
logger.error(f"Error storing memory: {e}", exc_info=True)
return f"Error: {str(e)}"
@mcp.tool()
async def recall(query: str, max_tokens: int = 4096, bank_id: str | None = None) -> str:
"""
Search memories to provide personalized, context-aware responses.
Use this tool PROACTIVELY to:
- Check user's preferences before making suggestions
- Recall user's history to provide continuity
- Remember user's goals and context
- Personalize responses based on past interactions
Args:
query: Natural language search query (e.g., "user's food preferences", "what projects is user working on")
max_tokens: Maximum tokens in the response (default: 4096)
bank_id: Optional bank to search in (defaults to session bank). Use for cross-bank operations.
"""
try:
target_bank = bank_id or get_current_bank_id()
if target_bank is None:
return "Error: No bank_id configured"
from hindsight_api.engine.memory_engine import Budget
recall_result = await memory.recall_async(
bank_id=target_bank,
query=query,
fact_type=list(VALID_RECALL_FACT_TYPES),
budget=Budget.HIGH,
max_tokens=max_tokens,
request_context=RequestContext(),
)
# Use model's JSON serialization
return recall_result.model_dump_json(indent=2)
except Exception as e:
logger.error(f"Error searching: {e}", exc_info=True)
return f'{{"error": "{e}", "results": []}}'
@mcp.tool()
async def reflect(query: str, context: str | None = None, budget: str = "low", bank_id: str | None = None) -> str:
"""
Generate thoughtful analysis by synthesizing stored memories with the bank's personality.
WHEN TO USE THIS TOOL:
Use reflect when you need reasoned analysis, not just fact retrieval. This tool
thinks through the question using everything the bank knows and its personality traits.
EXAMPLES OF GOOD QUERIES:
- "What patterns have emerged in how I approach debugging?"
- "Based on my past decisions, what architectural style do I prefer?"
- "What might be the best approach for this problem given what you know about me?"
- "How should I prioritize these tasks based on my goals?"
HOW IT DIFFERS FROM RECALL:
- recall: Returns raw facts matching your search (fast lookup)
- reflect: Reasons across memories to form a synthesized answer (deeper analysis)
Use recall for "what did I say about X?" and reflect for "what should I do about X?"
Args:
query: The question or topic to reflect on
context: Optional context about why this reflection is needed
budget: Search budget - 'low', 'mid', or 'high' (default: 'low')
bank_id: Optional bank to reflect in (defaults to session bank). Use for cross-bank operations.
"""
try:
target_bank = bank_id or get_current_bank_id()
if target_bank is None:
return "Error: No bank_id configured"
from hindsight_api.engine.memory_engine import Budget
# Map string budget to enum
budget_map = {"low": Budget.LOW, "mid": Budget.MID, "high": Budget.HIGH}
budget_enum = budget_map.get(budget.lower(), Budget.LOW)
reflect_result = await memory.reflect_async(
bank_id=target_bank,
query=query,
budget=budget_enum,
context=context,
request_context=RequestContext(),
)
return reflect_result.model_dump_json(indent=2)
except Exception as e:
logger.error(f"Error reflecting: {e}", exc_info=True)
return f'{{"error": "{e}", "text": ""}}'
@mcp.tool()
async def list_banks() -> str:
"""
List all available memory banks.
Use this tool to discover what memory banks exist in the system.
Each bank is an isolated memory store (like a separate "brain").
Returns:
JSON list of banks with their IDs, names, dispositions, and missions.
"""
try:
banks = await memory.list_banks(request_context=RequestContext())
return json.dumps({"banks": banks}, indent=2)
except Exception as e:
logger.error(f"Error listing banks: {e}", exc_info=True)
return f'{{"error": "{e}", "banks": []}}'
@mcp.tool()
async def create_bank(bank_id: str, name: str | None = None, mission: str | None = None) -> str:
"""
Create a new memory bank or get an existing one.
Memory banks are isolated stores - each one is like a separate "brain" for a user/agent.
Banks are auto-created with default settings if they don't exist.
Args:
bank_id: Unique identifier for the bank (e.g., 'user-123', 'agent-alpha')
name: Optional human-friendly name for the bank
mission: Optional mission describing who the agent is and what they're trying to accomplish
"""
try:
# get_bank_profile auto-creates bank if it doesn't exist
profile = await memory.get_bank_profile(bank_id, request_context=RequestContext())
# Update name/mission if provided
if name is not None or mission is not None:
await memory.update_bank(
bank_id,
name=name,
mission=mission,
request_context=RequestContext(),
)
# Fetch updated profile
profile = await memory.get_bank_profile(bank_id, request_context=RequestContext())
# Serialize disposition if it's a Pydantic model
if "disposition" in profile and hasattr(profile["disposition"], "model_dump"):
profile["disposition"] = profile["disposition"].model_dump()
return json.dumps(profile, indent=2)
except Exception as e:
logger.error(f"Error creating bank: {e}", exc_info=True)
return f'{{"error": "{e}"}}'
register_mcp_tools(mcp, memory, config)
return mcp
+157 -46
View File
@@ -4,9 +4,12 @@ Centralized configuration for Hindsight API.
All environment variables and their defaults are defined here.
"""
import json
import logging
import os
import sys
from dataclasses import dataclass
from datetime import datetime, timezone
from dotenv import find_dotenv, load_dotenv
@@ -17,6 +20,7 @@ logger = logging.getLogger(__name__)
# Environment variable names
ENV_DATABASE_URL = "HINDSIGHT_API_DATABASE_URL"
ENV_DATABASE_SCHEMA = "HINDSIGHT_API_DATABASE_SCHEMA"
ENV_LLM_PROVIDER = "HINDSIGHT_API_LLM_PROVIDER"
ENV_LLM_API_KEY = "HINDSIGHT_API_LLM_API_KEY"
ENV_LLM_MODEL = "HINDSIGHT_API_LLM_MODEL"
@@ -36,8 +40,14 @@ ENV_REFLECT_LLM_API_KEY = "HINDSIGHT_API_REFLECT_LLM_API_KEY"
ENV_REFLECT_LLM_MODEL = "HINDSIGHT_API_REFLECT_LLM_MODEL"
ENV_REFLECT_LLM_BASE_URL = "HINDSIGHT_API_REFLECT_LLM_BASE_URL"
ENV_CONSOLIDATION_LLM_PROVIDER = "HINDSIGHT_API_CONSOLIDATION_LLM_PROVIDER"
ENV_CONSOLIDATION_LLM_API_KEY = "HINDSIGHT_API_CONSOLIDATION_LLM_API_KEY"
ENV_CONSOLIDATION_LLM_MODEL = "HINDSIGHT_API_CONSOLIDATION_LLM_MODEL"
ENV_CONSOLIDATION_LLM_BASE_URL = "HINDSIGHT_API_CONSOLIDATION_LLM_BASE_URL"
ENV_EMBEDDINGS_PROVIDER = "HINDSIGHT_API_EMBEDDINGS_PROVIDER"
ENV_EMBEDDINGS_LOCAL_MODEL = "HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL"
ENV_EMBEDDINGS_LOCAL_FORCE_CPU = "HINDSIGHT_API_EMBEDDINGS_LOCAL_FORCE_CPU"
ENV_EMBEDDINGS_TEI_URL = "HINDSIGHT_API_EMBEDDINGS_TEI_URL"
ENV_EMBEDDINGS_OPENAI_API_KEY = "HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY"
ENV_EMBEDDINGS_OPENAI_MODEL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL"
@@ -57,6 +67,7 @@ ENV_RERANKER_LITELLM_MODEL = "HINDSIGHT_API_RERANKER_LITELLM_MODEL"
ENV_RERANKER_PROVIDER = "HINDSIGHT_API_RERANKER_PROVIDER"
ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
ENV_RERANKER_LOCAL_FORCE_CPU = "HINDSIGHT_API_RERANKER_LOCAL_FORCE_CPU"
ENV_RERANKER_LOCAL_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT"
ENV_RERANKER_TEI_URL = "HINDSIGHT_API_RERANKER_TEI_URL"
ENV_RERANKER_TEI_BATCH_SIZE = "HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE"
@@ -68,6 +79,7 @@ ENV_RERANKER_FLASHRANK_CACHE_DIR = "HINDSIGHT_API_RERANKER_FLASHRANK_CACHE_DIR"
ENV_HOST = "HINDSIGHT_API_HOST"
ENV_PORT = "HINDSIGHT_API_PORT"
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_GRAPH_RETRIEVER = "HINDSIGHT_API_GRAPH_RETRIEVER"
@@ -78,17 +90,19 @@ ENV_MCP_LOCAL_BANK_ID = "HINDSIGHT_API_MCP_LOCAL_BANK_ID"
ENV_MCP_INSTRUCTIONS = "HINDSIGHT_API_MCP_INSTRUCTIONS"
ENV_MENTAL_MODEL_REFRESH_CONCURRENCY = "HINDSIGHT_API_MENTAL_MODEL_REFRESH_CONCURRENCY"
# Observation thresholds
ENV_OBSERVATION_MIN_FACTS = "HINDSIGHT_API_OBSERVATION_MIN_FACTS"
ENV_OBSERVATION_TOP_ENTITIES = "HINDSIGHT_API_OBSERVATION_TOP_ENTITIES"
# 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_CUSTOM_INSTRUCTIONS = "HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS"
ENV_RETAIN_OBSERVATIONS_ASYNC = "HINDSIGHT_API_RETAIN_OBSERVATIONS_ASYNC"
# Observations settings (consolidated knowledge from facts)
ENV_ENABLE_OBSERVATIONS = "HINDSIGHT_API_ENABLE_OBSERVATIONS"
ENV_CONSOLIDATION_BATCH_SIZE = "HINDSIGHT_API_CONSOLIDATION_BATCH_SIZE"
ENV_CONSOLIDATION_MAX_TOKENS = "HINDSIGHT_API_CONSOLIDATION_MAX_TOKENS"
# Optimization flags
ENV_SKIP_LLM_VERIFICATION = "HINDSIGHT_API_SKIP_LLM_VERIFICATION"
ENV_LAZY_RERANKER = "HINDSIGHT_API_LAZY_RERANKER"
@@ -102,16 +116,20 @@ ENV_DB_POOL_MAX_SIZE = "HINDSIGHT_API_DB_POOL_MAX_SIZE"
ENV_DB_COMMAND_TIMEOUT = "HINDSIGHT_API_DB_COMMAND_TIMEOUT"
ENV_DB_ACQUIRE_TIMEOUT = "HINDSIGHT_API_DB_ACQUIRE_TIMEOUT"
# Background task processing
ENV_TASK_BACKEND = "HINDSIGHT_API_TASK_BACKEND"
ENV_TASK_BACKEND_MEMORY_BATCH_SIZE = "HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_SIZE"
ENV_TASK_BACKEND_MEMORY_BATCH_INTERVAL = "HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_INTERVAL"
# Worker configuration (distributed task processing)
ENV_WORKER_ENABLED = "HINDSIGHT_API_WORKER_ENABLED"
ENV_WORKER_ID = "HINDSIGHT_API_WORKER_ID"
ENV_WORKER_POLL_INTERVAL_MS = "HINDSIGHT_API_WORKER_POLL_INTERVAL_MS"
ENV_WORKER_MAX_RETRIES = "HINDSIGHT_API_WORKER_MAX_RETRIES"
ENV_WORKER_BATCH_SIZE = "HINDSIGHT_API_WORKER_BATCH_SIZE"
ENV_WORKER_HTTP_PORT = "HINDSIGHT_API_WORKER_HTTP_PORT"
# Reflect agent settings
ENV_REFLECT_MAX_ITERATIONS = "HINDSIGHT_API_REFLECT_MAX_ITERATIONS"
# Default values
DEFAULT_DATABASE_URL = "pg0"
DEFAULT_DATABASE_SCHEMA = "public"
DEFAULT_LLM_PROVIDER = "openai"
DEFAULT_LLM_MODEL = "gpt-5-mini"
DEFAULT_LLM_MAX_CONCURRENT = 32
@@ -119,11 +137,13 @@ DEFAULT_LLM_TIMEOUT = 120.0 # seconds
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)
DEFAULT_EMBEDDINGS_OPENAI_MODEL = "text-embedding-3-small"
DEFAULT_EMBEDDING_DIMENSION = 384
DEFAULT_RERANKER_PROVIDER = "local"
DEFAULT_RERANKER_LOCAL_MODEL = "cross-encoder/ms-marco-MiniLM-L-6-v2"
DEFAULT_RERANKER_LOCAL_FORCE_CPU = False # Force CPU mode for local reranker (avoids MPS/XPC issues on macOS)
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT = 4 # Limit concurrent CPU-bound reranking to prevent thrashing
DEFAULT_RERANKER_TEI_BATCH_SIZE = 128
DEFAULT_RERANKER_TEI_MAX_CONCURRENT = 8
@@ -142,6 +162,7 @@ DEFAULT_RERANKER_LITELLM_MODEL = "cohere/rerank-english-v3.0"
DEFAULT_HOST = "0.0.0.0"
DEFAULT_PORT = 8888
DEFAULT_LOG_LEVEL = "info"
DEFAULT_LOG_FORMAT = "text" # Options: "text", "json"
DEFAULT_WORKERS = 1
DEFAULT_MCP_ENABLED = True
DEFAULT_GRAPH_RETRIEVER = "link_expansion" # Options: "link_expansion", "mpfp", "bfs"
@@ -151,18 +172,20 @@ DEFAULT_RECALL_CONNECTION_BUDGET = 4 # Max concurrent DB connections per recall
DEFAULT_MCP_LOCAL_BANK_ID = "mcp"
DEFAULT_MENTAL_MODEL_REFRESH_CONCURRENCY = 8 # Max concurrent mental model refreshes
# Observation thresholds
DEFAULT_OBSERVATION_MIN_FACTS = 5 # Min facts required to generate entity observations
DEFAULT_OBSERVATION_TOP_ENTITIES = 5 # Max entities to process per retain batch
# Retain settings
DEFAULT_RETAIN_MAX_COMPLETION_TOKENS = 64000 # Max tokens for fact extraction LLM call
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" or "verbose"
RETAIN_EXTRACTION_MODES = ("concise", "verbose") # Allowed extraction modes
DEFAULT_RETAIN_EXTRACTION_MODE = "concise" # Extraction mode: "concise", "verbose", or "custom"
RETAIN_EXTRACTION_MODES = ("concise", "verbose", "custom") # Allowed extraction modes
DEFAULT_RETAIN_CUSTOM_INSTRUCTIONS = None # Custom extraction guidelines (only used when mode="custom")
DEFAULT_RETAIN_OBSERVATIONS_ASYNC = False # Run observation generation async (after retain completes)
# 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
# Database migrations
DEFAULT_RUN_MIGRATIONS_ON_STARTUP = True
@@ -172,10 +195,13 @@ DEFAULT_DB_POOL_MAX_SIZE = 100
DEFAULT_DB_COMMAND_TIMEOUT = 60 # seconds
DEFAULT_DB_ACQUIRE_TIMEOUT = 30 # seconds
# Background task processing
DEFAULT_TASK_BACKEND = "memory" # Options: "memory", "noop"
DEFAULT_TASK_BACKEND_MEMORY_BATCH_SIZE = 10
DEFAULT_TASK_BACKEND_MEMORY_BATCH_INTERVAL = 1.0 # seconds
# Worker configuration (distributed task processing)
DEFAULT_WORKER_ENABLED = True # API runs worker by default (standalone mode)
DEFAULT_WORKER_ID = None # Will use hostname if not specified
DEFAULT_WORKER_POLL_INTERVAL_MS = 500 # Poll database every 500ms
DEFAULT_WORKER_MAX_RETRIES = 3 # Max retries before marking task failed
DEFAULT_WORKER_BATCH_SIZE = 10 # Tasks to claim per poll cycle
DEFAULT_WORKER_HTTP_PORT = 8889 # HTTP port for worker metrics/health
# Reflect agent settings
DEFAULT_REFLECT_MAX_ITERATIONS = 10 # Max tool call iterations before forcing response
@@ -204,6 +230,36 @@ Use this tool PROACTIVELY to:
EMBEDDING_DIMENSION = DEFAULT_EMBEDDING_DIMENSION
class JsonFormatter(logging.Formatter):
"""JSON formatter for structured logging.
Outputs logs in JSON format with a 'severity' field that cloud logging
systems (GCP, AWS CloudWatch, etc.) can parse to correctly categorize log levels.
"""
SEVERITY_MAP = {
logging.DEBUG: "DEBUG",
logging.INFO: "INFO",
logging.WARNING: "WARNING",
logging.ERROR: "ERROR",
logging.CRITICAL: "CRITICAL",
}
def format(self, record: logging.LogRecord) -> str:
log_entry = {
"severity": self.SEVERITY_MAP.get(record.levelno, "DEFAULT"),
"message": record.getMessage(),
"timestamp": datetime.now(timezone.utc).isoformat(),
"logger": record.name,
}
# Add exception info if present
if record.exc_info:
log_entry["exception"] = self.formatException(record.exc_info)
return json.dumps(log_entry)
def _validate_extraction_mode(mode: str) -> str:
"""Validate and normalize extraction mode."""
mode_lower = mode.lower()
@@ -222,6 +278,7 @@ class HindsightConfig:
# Database
database_url: str
database_schema: str
# LLM (default, used as fallback for per-operation config)
llm_provider: str
@@ -242,9 +299,15 @@ class HindsightConfig:
reflect_llm_model: str | None
reflect_llm_base_url: str | None
consolidation_llm_provider: str | None
consolidation_llm_api_key: str | None
consolidation_llm_model: str | None
consolidation_llm_base_url: str | None
# Embeddings
embeddings_provider: str
embeddings_local_model: str
embeddings_local_force_cpu: bool
embeddings_tei_url: str | None
embeddings_openai_base_url: str | None
embeddings_cohere_base_url: str | None
@@ -252,6 +315,8 @@ class HindsightConfig:
# Reranker
reranker_provider: str
reranker_local_model: str
reranker_local_force_cpu: bool
reranker_local_max_concurrent: int
reranker_tei_url: str | None
reranker_tei_batch_size: int
reranker_tei_max_concurrent: int
@@ -262,6 +327,7 @@ class HindsightConfig:
host: str
port: int
log_level: str
log_format: str
mcp_enabled: bool
# Recall
@@ -271,17 +337,19 @@ class HindsightConfig:
recall_connection_budget: int
mental_model_refresh_concurrency: int
# Observation thresholds
observation_min_facts: int
observation_top_entities: int
# Retain settings
retain_max_completion_tokens: int
retain_chunk_size: int
retain_extract_causal_links: bool
retain_extraction_mode: str
retain_custom_instructions: str | None
retain_observations_async: bool
# Observations settings (consolidated knowledge from facts)
enable_observations: bool
consolidation_batch_size: int
consolidation_max_tokens: int
# Optimization flags
skip_llm_verification: bool
lazy_reranker: bool
@@ -295,10 +363,13 @@ class HindsightConfig:
db_command_timeout: int
db_acquire_timeout: int
# Background task processing
task_backend: str
task_backend_memory_batch_size: int
task_backend_memory_batch_interval: float
# Worker configuration (distributed task processing)
worker_enabled: bool
worker_id: str | None
worker_poll_interval_ms: int
worker_max_retries: int
worker_batch_size: int
worker_http_port: int
# Reflect agent settings
reflect_max_iterations: int
@@ -309,6 +380,7 @@ class HindsightConfig:
return cls(
# Database
database_url=os.getenv(ENV_DATABASE_URL, DEFAULT_DATABASE_URL),
database_schema=os.getenv(ENV_DATABASE_SCHEMA, DEFAULT_DATABASE_SCHEMA),
# LLM
llm_provider=os.getenv(ENV_LLM_PROVIDER, DEFAULT_LLM_PROVIDER),
llm_api_key=os.getenv(ENV_LLM_API_KEY),
@@ -325,15 +397,30 @@ class HindsightConfig:
reflect_llm_api_key=os.getenv(ENV_REFLECT_LLM_API_KEY) or None,
reflect_llm_model=os.getenv(ENV_REFLECT_LLM_MODEL) or None,
reflect_llm_base_url=os.getenv(ENV_REFLECT_LLM_BASE_URL) or None,
consolidation_llm_provider=os.getenv(ENV_CONSOLIDATION_LLM_PROVIDER) or None,
consolidation_llm_api_key=os.getenv(ENV_CONSOLIDATION_LLM_API_KEY) or None,
consolidation_llm_model=os.getenv(ENV_CONSOLIDATION_LLM_MODEL) or None,
consolidation_llm_base_url=os.getenv(ENV_CONSOLIDATION_LLM_BASE_URL) or None,
# Embeddings
embeddings_provider=os.getenv(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER),
embeddings_local_model=os.getenv(ENV_EMBEDDINGS_LOCAL_MODEL, DEFAULT_EMBEDDINGS_LOCAL_MODEL),
embeddings_local_force_cpu=os.getenv(
ENV_EMBEDDINGS_LOCAL_FORCE_CPU, str(DEFAULT_EMBEDDINGS_LOCAL_FORCE_CPU)
).lower()
in ("true", "1"),
embeddings_tei_url=os.getenv(ENV_EMBEDDINGS_TEI_URL),
embeddings_openai_base_url=os.getenv(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None,
embeddings_cohere_base_url=os.getenv(ENV_EMBEDDINGS_COHERE_BASE_URL) or None,
# Reranker
reranker_provider=os.getenv(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER),
reranker_local_model=os.getenv(ENV_RERANKER_LOCAL_MODEL, DEFAULT_RERANKER_LOCAL_MODEL),
reranker_local_force_cpu=os.getenv(
ENV_RERANKER_LOCAL_FORCE_CPU, str(DEFAULT_RERANKER_LOCAL_FORCE_CPU)
).lower()
in ("true", "1"),
reranker_local_max_concurrent=int(
os.getenv(ENV_RERANKER_LOCAL_MAX_CONCURRENT, str(DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT))
),
reranker_tei_url=os.getenv(ENV_RERANKER_TEI_URL),
reranker_tei_batch_size=int(os.getenv(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE))),
reranker_tei_max_concurrent=int(
@@ -345,6 +432,7 @@ class HindsightConfig:
host=os.getenv(ENV_HOST, DEFAULT_HOST),
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
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",
# Recall
graph_retriever=os.getenv(ENV_GRAPH_RETRIEVER, DEFAULT_GRAPH_RETRIEVER),
@@ -359,11 +447,6 @@ class HindsightConfig:
# Optimization flags
skip_llm_verification=os.getenv(ENV_SKIP_LLM_VERIFICATION, "false").lower() == "true",
lazy_reranker=os.getenv(ENV_LAZY_RERANKER, "false").lower() == "true",
# Observation thresholds
observation_min_facts=int(os.getenv(ENV_OBSERVATION_MIN_FACTS, str(DEFAULT_OBSERVATION_MIN_FACTS))),
observation_top_entities=int(
os.getenv(ENV_OBSERVATION_TOP_ENTITIES, str(DEFAULT_OBSERVATION_TOP_ENTITIES))
),
# Retain settings
retain_max_completion_tokens=int(
os.getenv(ENV_RETAIN_MAX_COMPLETION_TOKENS, str(DEFAULT_RETAIN_MAX_COMPLETION_TOKENS))
@@ -376,10 +459,19 @@ class HindsightConfig:
retain_extraction_mode=_validate_extraction_mode(
os.getenv(ENV_RETAIN_EXTRACTION_MODE, DEFAULT_RETAIN_EXTRACTION_MODE)
),
retain_custom_instructions=os.getenv(ENV_RETAIN_CUSTOM_INSTRUCTIONS) or DEFAULT_RETAIN_CUSTOM_INSTRUCTIONS,
retain_observations_async=os.getenv(
ENV_RETAIN_OBSERVATIONS_ASYNC, str(DEFAULT_RETAIN_OBSERVATIONS_ASYNC)
).lower()
== "true",
# Observations settings (consolidated knowledge from facts)
enable_observations=os.getenv(ENV_ENABLE_OBSERVATIONS, str(DEFAULT_ENABLE_OBSERVATIONS)).lower() == "true",
consolidation_batch_size=int(
os.getenv(ENV_CONSOLIDATION_BATCH_SIZE, str(DEFAULT_CONSOLIDATION_BATCH_SIZE))
),
consolidation_max_tokens=int(
os.getenv(ENV_CONSOLIDATION_MAX_TOKENS, str(DEFAULT_CONSOLIDATION_MAX_TOKENS))
),
# Database migrations
run_migrations_on_startup=os.getenv(ENV_RUN_MIGRATIONS_ON_STARTUP, "true").lower() == "true",
# Database connection pool
@@ -387,14 +479,13 @@ class HindsightConfig:
db_pool_max_size=int(os.getenv(ENV_DB_POOL_MAX_SIZE, str(DEFAULT_DB_POOL_MAX_SIZE))),
db_command_timeout=int(os.getenv(ENV_DB_COMMAND_TIMEOUT, str(DEFAULT_DB_COMMAND_TIMEOUT))),
db_acquire_timeout=int(os.getenv(ENV_DB_ACQUIRE_TIMEOUT, str(DEFAULT_DB_ACQUIRE_TIMEOUT))),
# Background task processing
task_backend=os.getenv(ENV_TASK_BACKEND, DEFAULT_TASK_BACKEND),
task_backend_memory_batch_size=int(
os.getenv(ENV_TASK_BACKEND_MEMORY_BATCH_SIZE, str(DEFAULT_TASK_BACKEND_MEMORY_BATCH_SIZE))
),
task_backend_memory_batch_interval=float(
os.getenv(ENV_TASK_BACKEND_MEMORY_BATCH_INTERVAL, str(DEFAULT_TASK_BACKEND_MEMORY_BATCH_INTERVAL))
),
# Worker configuration
worker_enabled=os.getenv(ENV_WORKER_ENABLED, str(DEFAULT_WORKER_ENABLED)).lower() == "true",
worker_id=os.getenv(ENV_WORKER_ID) or DEFAULT_WORKER_ID,
worker_poll_interval_ms=int(os.getenv(ENV_WORKER_POLL_INTERVAL_MS, str(DEFAULT_WORKER_POLL_INTERVAL_MS))),
worker_max_retries=int(os.getenv(ENV_WORKER_MAX_RETRIES, str(DEFAULT_WORKER_MAX_RETRIES))),
worker_batch_size=int(os.getenv(ENV_WORKER_BATCH_SIZE, str(DEFAULT_WORKER_BATCH_SIZE))),
worker_http_port=int(os.getenv(ENV_WORKER_HTTP_PORT, str(DEFAULT_WORKER_HTTP_PORT))),
# Reflect agent settings
reflect_max_iterations=int(os.getenv(ENV_REFLECT_MAX_ITERATIONS, str(DEFAULT_REFLECT_MAX_ITERATIONS))),
)
@@ -427,16 +518,32 @@ class HindsightConfig:
return log_level_map.get(self.log_level.lower(), logging.INFO)
def configure_logging(self) -> None:
"""Configure Python logging based on the log level."""
logging.basicConfig(
level=self.get_python_log_level(),
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
force=True, # Override any existing configuration
)
"""Configure Python logging based on the log level and format.
When log_format is "json", outputs structured JSON logs with a severity
field that GCP Cloud Logging can parse for proper log level categorization.
"""
root_logger = logging.getLogger()
root_logger.setLevel(self.get_python_log_level())
# Remove existing handlers
for handler in root_logger.handlers[:]:
root_logger.removeHandler(handler)
# Create handler writing to stdout (GCP treats stderr as ERROR)
handler = logging.StreamHandler(sys.stdout)
handler.setLevel(self.get_python_log_level())
if self.log_format == "json":
handler.setFormatter(JsonFormatter())
else:
handler.setFormatter(logging.Formatter("%(asctime)s - %(levelname)s - %(name)s - %(message)s"))
root_logger.addHandler(handler)
def log_config(self) -> None:
"""Log the current configuration (without sensitive values)."""
logger.info(f"Database: {self.database_url}")
logger.info(f"Database: {self.database_url} (schema: {self.database_schema})")
logger.info(f"LLM: provider={self.llm_provider}, model={self.llm_model}")
if self.retain_llm_provider or self.retain_llm_model:
retain_provider = self.retain_llm_provider or self.llm_provider
@@ -446,6 +553,10 @@ class HindsightConfig:
reflect_provider = self.reflect_llm_provider or self.llm_provider
reflect_model = self.reflect_llm_model or self.llm_model
logger.info(f"LLM (reflect): provider={reflect_provider}, model={reflect_model}")
if self.consolidation_llm_provider or self.consolidation_llm_model:
consolidation_provider = self.consolidation_llm_provider or self.llm_provider
consolidation_model = self.consolidation_llm_model or self.llm_model
logger.info(f"LLM (consolidation): provider={consolidation_provider}, model={consolidation_model}")
logger.info(f"Embeddings: provider={self.embeddings_provider}")
logger.info(f"Reranker: provider={self.reranker_provider}")
logger.info(f"Graph retriever: {self.graph_retriever}")
+4 -1
View File
@@ -52,7 +52,10 @@ class IdleTimeoutMiddleware:
logger.info(f"Idle timeout reached ({self.idle_timeout}s), shutting down daemon")
# Give a moment for any in-flight requests
await asyncio.sleep(1)
os._exit(0)
# Send SIGTERM to ourselves to trigger graceful shutdown
import signal
os.kill(os.getpid(), signal.SIGTERM)
class DaemonLock:
@@ -0,0 +1,5 @@
"""Consolidation engine for automatic learning creation from memories."""
from .consolidator import run_consolidation_job
__all__ = ["run_consolidation_job"]
@@ -0,0 +1,926 @@
"""Consolidation engine for automatic observation creation from memories.
The consolidation engine runs as a background job after retain operations complete.
It processes new memories and either:
- Creates new observations from novel facts
- Updates existing observations when new evidence supports/contradicts/refines them
Observations are stored in memory_units with fact_type='observation' and include:
- proof_count: Number of supporting memories
- source_memory_ids: Array of memory UUIDs that contribute to this observation
- history: JSONB tracking changes over time
"""
import json
import logging
import time
import uuid
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any
from ..memory_engine import fq_table
from ..retain import embedding_utils
from .prompts import (
CONSOLIDATION_SYSTEM_PROMPT,
CONSOLIDATION_USER_PROMPT,
)
if TYPE_CHECKING:
from asyncpg import Connection
from ...api.http import RequestContext
from ..memory_engine import MemoryEngine
logger = logging.getLogger(__name__)
class ConsolidationPerfLog:
"""Performance logging for consolidation operations."""
def __init__(self, bank_id: str):
self.bank_id = bank_id
self.start_time = time.time()
self.lines: list[str] = []
self.timings: dict[str, float] = {}
def log(self, message: str) -> None:
"""Add a log line."""
self.lines.append(message)
def record_timing(self, key: str, duration: float) -> None:
"""Record a timing measurement."""
if key in self.timings:
self.timings[key] += duration
else:
self.timings[key] = duration
def flush(self) -> None:
"""Flush all log lines to the logger."""
total_time = time.time() - self.start_time
header = f"\n{'=' * 60}\nCONSOLIDATION for bank {self.bank_id}"
footer = f"{'=' * 60}\nCONSOLIDATION COMPLETE: {total_time:.3f}s total\n{'=' * 60}"
log_output = header + "\n" + "\n".join(self.lines) + "\n" + footer
logger.info(log_output)
async def run_consolidation_job(
memory_engine: "MemoryEngine",
bank_id: str,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Run consolidation job for a bank.
This is called after retain operations to consolidate new memories into mental models.
Args:
memory_engine: MemoryEngine instance
bank_id: Bank identifier
request_context: Request context for authentication
Returns:
Dict with consolidation results
"""
from ...config import get_config
config = get_config()
perf = ConsolidationPerfLog(bank_id)
max_memories_per_batch = config.consolidation_batch_size
# Check if consolidation is enabled
if not config.enable_observations:
logger.debug(f"Consolidation disabled for bank {bank_id}")
return {"status": "disabled", "bank_id": bank_id}
pool = memory_engine._pool
# Get bank profile
async with pool.acquire() as conn:
t0 = time.time()
bank_row = await conn.fetchrow(
f"""
SELECT bank_id, name, mission
FROM {fq_table("banks")}
WHERE bank_id = $1
""",
bank_id,
)
if not bank_row:
logger.warning(f"Bank {bank_id} not found for consolidation")
return {"status": "bank_not_found", "bank_id": bank_id}
mission = bank_row["mission"] or "General memory consolidation"
perf.record_timing("fetch_bank", time.time() - t0)
# Count total unconsolidated memories for progress logging
total_count = await conn.fetchval(
f"""
SELECT COUNT(*)
FROM {fq_table("memory_units")}
WHERE bank_id = $1
AND consolidated_at IS NULL
AND fact_type IN ('experience', 'world')
""",
bank_id,
)
if total_count == 0:
logger.debug(f"No new memories to consolidate for bank {bank_id}")
return {"status": "no_new_memories", "bank_id": bank_id, "memories_processed": 0}
logger.info(f"[CONSOLIDATION] bank={bank_id} total_unconsolidated={total_count}")
perf.log(f"[1] Found {total_count} pending memories to consolidate")
# Process each memory with individual commits for crash recovery
stats = {
"memories_processed": 0,
"observations_created": 0,
"observations_updated": 0,
"observations_merged": 0,
"actions_executed": 0,
"skipped": 0,
}
batch_num = 0
while True:
batch_num += 1
batch_start = time.time()
# Fetch next batch of unconsolidated memories
async with pool.acquire() as conn:
t0 = time.time()
memories = await conn.fetch(
f"""
SELECT id, text, fact_type, occurred_start, occurred_end, event_date, tags, mentioned_at
FROM {fq_table("memory_units")}
WHERE bank_id = $1
AND consolidated_at IS NULL
AND fact_type IN ('experience', 'world')
ORDER BY created_at ASC
LIMIT $2
""",
bank_id,
max_memories_per_batch,
)
perf.record_timing("fetch_memories", time.time() - t0)
if not memories:
break # No more unconsolidated memories
for memory in memories:
mem_start = time.time()
# Process the memory (uses its own connection internally)
async with pool.acquire() as conn:
result = await _process_memory(
conn=conn,
memory_engine=memory_engine,
bank_id=bank_id,
memory=dict(memory),
mission=mission,
request_context=request_context,
perf=perf,
)
# Mark memory as consolidated (committed immediately)
await conn.execute(
f"""
UPDATE {fq_table("memory_units")}
SET consolidated_at = NOW()
WHERE id = $1
""",
memory["id"],
)
mem_time = time.time() - mem_start
perf.record_timing("process_memory_total", mem_time)
stats["memories_processed"] += 1
action = result.get("action")
if action == "created":
stats["observations_created"] += 1
stats["actions_executed"] += 1
elif action == "updated":
stats["observations_updated"] += 1
stats["actions_executed"] += 1
elif action == "merged":
stats["observations_merged"] += 1
stats["actions_executed"] += 1
elif action == "multiple":
stats["observations_created"] += result.get("created", 0)
stats["observations_updated"] += result.get("updated", 0)
stats["observations_merged"] += result.get("merged", 0)
stats["actions_executed"] += result.get("total_actions", 0)
elif action == "skipped":
stats["skipped"] += 1
# Log progress periodically
if stats["memories_processed"] % 10 == 0:
logger.info(
f"[CONSOLIDATION] bank={bank_id} progress: "
f"{stats['memories_processed']}/{total_count} memories processed"
)
batch_time = time.time() - batch_start
perf.log(
f"[2] Batch {batch_num}: {len(memories)} memories in {batch_time:.3f}s "
f"(avg {batch_time / len(memories):.3f}s/memory)"
)
# Build summary
perf.log(
f"[3] Results: {stats['memories_processed']} memories -> "
f"{stats['actions_executed']} actions "
f"({stats['observations_created']} created, "
f"{stats['observations_updated']} updated, "
f"{stats['observations_merged']} merged, "
f"{stats['skipped']} skipped)"
)
# Add timing breakdown
timing_parts = []
if "recall" in perf.timings:
timing_parts.append(f"recall={perf.timings['recall']:.3f}s")
if "llm" in perf.timings:
timing_parts.append(f"llm={perf.timings['llm']:.3f}s")
if "embedding" in perf.timings:
timing_parts.append(f"embedding={perf.timings['embedding']:.3f}s")
if "db_write" in perf.timings:
timing_parts.append(f"db_write={perf.timings['db_write']:.3f}s")
if timing_parts:
perf.log(f"[4] Timing breakdown: {', '.join(timing_parts)}")
# Trigger mental model refreshes for models with refresh_after_consolidation=true
mental_models_refreshed = await _trigger_mental_model_refreshes(
memory_engine=memory_engine,
bank_id=bank_id,
request_context=request_context,
perf=perf,
)
stats["mental_models_refreshed"] = mental_models_refreshed
perf.flush()
return {"status": "completed", "bank_id": bank_id, **stats}
async def _trigger_mental_model_refreshes(
memory_engine: "MemoryEngine",
bank_id: str,
request_context: "RequestContext",
perf: ConsolidationPerfLog | None = None,
) -> int:
"""
Trigger refreshes for mental models with refresh_after_consolidation=true.
Args:
memory_engine: MemoryEngine instance
bank_id: Bank identifier
request_context: Request context for authentication
perf: Performance logging
Returns:
Number of mental models scheduled for refresh
"""
pool = memory_engine._pool
# Find mental models with refresh_after_consolidation=true
async with pool.acquire() as conn:
rows = await conn.fetch(
f"""
SELECT id, name
FROM {fq_table("mental_models")}
WHERE bank_id = $1
AND (trigger->>'refresh_after_consolidation')::boolean = true
""",
bank_id,
)
if not rows:
return 0
if perf:
perf.log(f"[5] Triggering refresh for {len(rows)} mental models with refresh_after_consolidation=true")
# Submit refresh tasks for each mental model
refreshed_count = 0
for row in rows:
mental_model_id = row["id"]
try:
await memory_engine.submit_async_refresh_mental_model(
bank_id=bank_id,
mental_model_id=mental_model_id,
request_context=request_context,
)
refreshed_count += 1
logger.info(
f"[CONSOLIDATION] Triggered refresh for mental model {mental_model_id} "
f"(name: {row['name']}) in bank {bank_id}"
)
except Exception as e:
logger.warning(f"[CONSOLIDATION] Failed to trigger refresh for mental model {mental_model_id}: {e}")
return refreshed_count
async def _process_memory(
conn: "Connection",
memory_engine: "MemoryEngine",
bank_id: str,
memory: dict[str, Any],
mission: str,
request_context: "RequestContext",
perf: ConsolidationPerfLog | None = None,
) -> dict[str, Any]:
"""
Process a single memory for consolidation using a SINGLE LLM call.
This function:
1. Finds related observations (can be empty)
2. Uses ONE LLM call to extract durable knowledge AND decide on actions
3. Executes array of actions (can be multiple creates/updates)
The LLM handles all cases:
- No related observations: returns create action(s) with extracted durable knowledge
- Related observations exist: returns update/create actions based on tag routing
- Purely ephemeral fact: returns empty array (skip)
Returns:
Dict with action summary: created/updated/merged counts
"""
fact_text = memory["text"]
memory_id = memory["id"]
fact_tags = memory.get("tags") or []
# Find related observations using the full recall system (NO tag filtering)
t0 = time.time()
related_observations = await _find_related_observations(
conn=conn,
memory_engine=memory_engine,
bank_id=bank_id,
query=fact_text,
request_context=request_context,
)
if perf:
perf.record_timing("recall", time.time() - t0)
# Single LLM call handles ALL cases (with or without existing observations)
# Note: Tags are NOT passed to LLM - they are handled algorithmically
t0 = time.time()
actions = await _consolidate_with_llm(
memory_engine=memory_engine,
fact_text=fact_text,
observations=related_observations, # Can be empty list
mission=mission,
)
if perf:
perf.record_timing("llm", time.time() - t0)
if not actions:
# LLM returned empty array - fact is purely ephemeral, skip
return {"action": "skipped", "reason": "no_durable_knowledge"}
# Execute all actions and collect results
results = []
for action in actions:
action_type = action.get("action")
if action_type == "update":
result = await _execute_update_action(
conn=conn,
memory_engine=memory_engine,
bank_id=bank_id,
memory_id=memory_id,
action=action,
observations=related_observations,
source_fact_tags=fact_tags, # Pass source fact's tags for security
source_occurred_start=memory.get("occurred_start"),
source_occurred_end=memory.get("occurred_end"),
source_mentioned_at=memory.get("mentioned_at"),
perf=perf,
)
results.append(result)
elif action_type == "create":
result = await _execute_create_action(
conn=conn,
memory_engine=memory_engine,
bank_id=bank_id,
memory_id=memory_id,
action=action,
source_fact_tags=fact_tags, # Pass source fact's tags for security
event_date=memory.get("event_date"),
occurred_start=memory.get("occurred_start"),
occurred_end=memory.get("occurred_end"),
mentioned_at=memory.get("mentioned_at"),
perf=perf,
)
results.append(result)
if not results:
# No valid actions executed
return {"action": "skipped", "reason": "no_valid_actions"}
# Summarize results
created = sum(1 for r in results if r.get("action") == "created")
updated = sum(1 for r in results if r.get("action") == "updated")
merged = sum(1 for r in results if r.get("action") == "merged")
if len(results) == 1:
return results[0]
return {
"action": "multiple",
"created": created,
"updated": updated,
"merged": merged,
"total_actions": len(results),
}
async def _execute_update_action(
conn: "Connection",
memory_engine: "MemoryEngine",
bank_id: str,
memory_id: uuid.UUID,
action: dict[str, Any],
observations: list[dict[str, Any]],
source_fact_tags: list[str] | None = None,
source_occurred_start: datetime | None = None,
source_occurred_end: datetime | None = None,
source_mentioned_at: datetime | None = None,
perf: ConsolidationPerfLog | None = None,
) -> dict[str, Any]:
"""
Execute an update action on an existing observation.
Updates the observation text, adds to history, increments proof_count,
and updates temporal fields:
- occurred_start: uses LEAST to keep the earliest start time
- occurred_end: uses GREATEST to keep the most recent end time
- mentioned_at: uses GREATEST to keep the most recent mention time
SECURITY: Merges source fact's tags into the observation's existing tags.
This ensures all contributors can see the observation they contributed to.
For example, if Lisa's observation (tags=['user_lisa']) is updated with
Mike's fact (tags=['user_mike']), the observation will have both tags.
"""
learning_id = action.get("learning_id")
new_text = action.get("text")
reason = action.get("reason", "Updated with new fact")
if not learning_id or not new_text:
return {"action": "skipped", "reason": "missing_learning_id_or_text"}
# Find the observation
model = next((m for m in observations if str(m["id"]) == learning_id), None)
if not model:
return {"action": "skipped", "reason": "learning_not_found"}
# Build history entry
history = list(model.get("history", []))
history.append(
{
"previous_text": model["text"],
"changed_at": datetime.now(timezone.utc).isoformat(),
"reason": reason,
"source_memory_id": str(memory_id),
}
)
# Update source_memory_ids
source_ids = list(model.get("source_memory_ids", []))
source_ids.append(memory_id)
# SECURITY: Merge source fact's tags into existing observation tags
# This ensures all contributors can see the observation they contributed to
existing_tags = set(model.get("tags", []) or [])
source_tags = set(source_fact_tags or [])
merged_tags = list(existing_tags | source_tags) # Union of both tag sets
if source_tags and source_tags != existing_tags:
logger.debug(
f"Security: Merging tags for observation {learning_id}: "
f"existing={list(existing_tags)}, source={list(source_tags)}, merged={merged_tags}"
)
# Generate new embedding for updated text
t0 = time.time()
embeddings = await embedding_utils.generate_embeddings_batch(memory_engine.embeddings, [new_text])
embedding_str = str(embeddings[0]) if embeddings else None
if perf:
perf.record_timing("embedding", time.time() - t0)
# Update the observation
# - occurred_start: LEAST keeps the earliest start time across all source facts
# - occurred_end: GREATEST keeps the most recent end time across all source facts
# - mentioned_at: GREATEST keeps the most recent mention time
# - tags: merged from existing + source fact (for visibility)
t0 = time.time()
await conn.execute(
f"""
UPDATE {fq_table("memory_units")}
SET text = $1,
embedding = $2::vector,
history = $3,
source_memory_ids = $4,
proof_count = $5,
tags = $10,
updated_at = now(),
occurred_start = LEAST(occurred_start, COALESCE($7, occurred_start)),
occurred_end = GREATEST(occurred_end, COALESCE($8, occurred_end)),
mentioned_at = GREATEST(mentioned_at, COALESCE($9, mentioned_at))
WHERE id = $6
""",
new_text,
embedding_str,
json.dumps(history),
source_ids,
len(source_ids),
uuid.UUID(learning_id),
source_occurred_start,
source_occurred_end,
source_mentioned_at,
merged_tags,
)
# Create links from memory to observation
await _create_memory_links(conn, memory_id, uuid.UUID(learning_id))
if perf:
perf.record_timing("db_write", time.time() - t0)
logger.debug(f"Updated observation {learning_id} with memory {memory_id}")
return {"action": "updated", "observation_id": learning_id}
async def _execute_create_action(
conn: "Connection",
memory_engine: "MemoryEngine",
bank_id: str,
memory_id: uuid.UUID,
action: dict[str, Any],
source_fact_tags: list[str] | None = None,
event_date: datetime | None = None,
occurred_start: datetime | None = None,
occurred_end: datetime | None = None,
mentioned_at: datetime | None = None,
perf: ConsolidationPerfLog | None = None,
) -> dict[str, Any]:
"""
Execute a create action for a new observation.
Creates a new observation with the specified text.
The text comes directly from the classify LLM - no second LLM call needed.
Tags are determined algorithmically (not by LLM):
- Observations always inherit their source fact's tags
- This ensures visibility scope is maintained (security)
"""
text = action.get("text")
# Tags are determined algorithmically - always use source fact's tags
# This ensures private memories create private observations
tags = source_fact_tags or []
if not text:
return {"action": "skipped", "reason": "missing_text"}
# Use text directly from classify - skip the redundant LLM call
result = await _create_observation_directly(
conn=conn,
memory_engine=memory_engine,
bank_id=bank_id,
source_memory_id=memory_id,
observation_text=text, # Text already processed by classify LLM
tags=tags,
event_date=event_date,
occurred_start=occurred_start,
occurred_end=occurred_end,
mentioned_at=mentioned_at,
perf=perf,
)
logger.debug(f"Created observation {result.get('observation_id')} from memory {memory_id} (tags: {tags})")
return result
async def _create_memory_links(
conn: "Connection",
memory_id: uuid.UUID,
observation_id: uuid.UUID,
) -> None:
"""
Placeholder for observation link creation.
Observations do NOT get any memory_links copied from their source facts.
Instead, retrieval uses source_memory_ids to traverse:
- Entity connections: observation → source_memory_ids → unit_entities
- Semantic similarity: observations have their own embeddings
- Temporal proximity: observations have their own temporal fields
This avoids data duplication and ensures observations are always
connected via their source facts' relationships.
The memory_id and observation_id parameters are kept for interface
compatibility but no links are created.
"""
# No links are created - observations rely on source_memory_ids for traversal
pass
async def _find_related_observations(
conn: "Connection",
memory_engine: "MemoryEngine",
bank_id: str,
query: str,
request_context: "RequestContext",
) -> list[dict[str, Any]]:
"""
Find observations related to the given query using optimized recall.
IMPORTANT: We do NOT filter by tags here. Consolidation needs to see ALL
potentially related observations regardless of scope, so the LLM can
decide on tag routing (same scope update vs cross-scope create).
Uses max_tokens to naturally limit observations (no artificial count limit).
Includes source memories with dates for LLM context.
Returns:
List of related observations with their tags, source memories, and dates
"""
# Use recall to find related observations with token budget
# max_tokens naturally limits how many observations are returned
from ...config import get_config
config = get_config()
recall_result = await memory_engine.recall_async(
bank_id=bank_id,
query=query,
max_tokens=config.consolidation_max_tokens, # Token budget for observations (configurable)
fact_type=["observation"], # Only retrieve observations
request_context=request_context,
_quiet=True, # Suppress logging
# NO tags parameter - intentionally get ALL observations
)
# If no observations returned, return empty list
if not recall_result.results:
return []
# Batch fetch all observations in a single query (no artificial limit)
observation_ids = [uuid.UUID(obs.id) for obs in recall_result.results]
rows = await conn.fetch(
f"""
SELECT id, text, proof_count, history, tags, source_memory_ids, created_at, updated_at,
occurred_start, occurred_end, mentioned_at
FROM {fq_table("memory_units")}
WHERE id = ANY($1) AND bank_id = $2 AND fact_type = 'observation'
""",
observation_ids,
bank_id,
)
# Build results list preserving recall order
id_to_row = {row["id"]: row for row in rows}
results = []
for obs in recall_result.results:
obs_id = uuid.UUID(obs.id)
if obs_id not in id_to_row:
continue
row = id_to_row[obs_id]
history = row["history"]
if isinstance(history, str):
history = json.loads(history)
elif history is None:
history = []
# Fetch source memories to include their text and dates
source_memory_ids = row["source_memory_ids"] or []
source_memories = []
if source_memory_ids:
source_rows = await conn.fetch(
f"""
SELECT text, occurred_start, occurred_end, mentioned_at, event_date
FROM {fq_table("memory_units")}
WHERE id = ANY($1) AND bank_id = $2
ORDER BY created_at ASC
LIMIT 5
""",
source_memory_ids[:5], # Limit to first 5 source memories for token efficiency
bank_id,
)
for src_row in source_rows:
source_memories.append(
{
"text": src_row["text"],
"occurred_start": src_row["occurred_start"],
"occurred_end": src_row["occurred_end"],
"mentioned_at": src_row["mentioned_at"],
"event_date": src_row["event_date"],
}
)
results.append(
{
"id": row["id"],
"text": row["text"],
"proof_count": row["proof_count"] or 1,
"tags": row["tags"] or [],
"source_memories": source_memories,
"occurred_start": row["occurred_start"],
"occurred_end": row["occurred_end"],
"mentioned_at": row["mentioned_at"],
"created_at": row["created_at"],
"updated_at": row["updated_at"],
}
)
return results
async def _consolidate_with_llm(
memory_engine: "MemoryEngine",
fact_text: str,
observations: list[dict[str, Any]],
mission: str,
) -> list[dict[str, Any]]:
"""
Single LLM call to extract durable knowledge and decide on consolidation actions.
This handles ALL cases:
- No related observations: extracts durable knowledge, returns create action
- Related observations exist: compares and returns update/create actions
- Purely ephemeral fact: returns empty array
Note: Tags are NOT handled by the LLM. They are determined algorithmically:
- CREATE: observation inherits source fact's tags
- UPDATE: observation merges source fact's tags with existing tags
Returns:
List of actions, each being:
- {"action": "update", "learning_id": "uuid", "text": "...", "reason": "..."}
- {"action": "create", "text": "...", "reason": "..."}
- [] if fact is purely ephemeral (no durable knowledge)
"""
# Format observations as JSON with source memories and dates
if observations:
obs_list = []
for obs in observations:
obs_data = {
"id": str(obs["id"]),
"text": obs["text"],
"proof_count": obs["proof_count"],
"tags": obs["tags"],
"created_at": obs["created_at"].isoformat() if obs.get("created_at") else None,
"updated_at": obs["updated_at"].isoformat() if obs.get("updated_at") else None,
}
# Include temporal info if available
if obs.get("occurred_start"):
obs_data["occurred_start"] = obs["occurred_start"].isoformat()
if obs.get("occurred_end"):
obs_data["occurred_end"] = obs["occurred_end"].isoformat()
if obs.get("mentioned_at"):
obs_data["mentioned_at"] = obs["mentioned_at"].isoformat()
# Include source memories (up to 3 for brevity)
if obs.get("source_memories"):
obs_data["source_memories"] = [
{
"text": sm["text"],
"event_date": sm["event_date"].isoformat() if sm.get("event_date") else None,
"occurred_start": sm["occurred_start"].isoformat() if sm.get("occurred_start") else None,
}
for sm in obs["source_memories"][:3] # Limit to 3 for token efficiency
]
obs_list.append(obs_data)
observations_text = json.dumps(obs_list, indent=2)
else:
observations_text = "[]"
# Only include mission section if mission is set and not the default
mission_section = ""
if mission and mission != "General memory consolidation":
mission_section = f"""
MISSION CONTEXT: {mission}
Focus on DURABLE knowledge that serves this mission, not ephemeral state.
"""
user_prompt = CONSOLIDATION_USER_PROMPT.format(
mission_section=mission_section,
fact_text=fact_text,
observations_text=observations_text,
)
messages = [
{"role": "system", "content": CONSOLIDATION_SYSTEM_PROMPT},
{"role": "user", "content": user_prompt},
]
try:
result = await memory_engine._consolidation_llm_config.call(
messages=messages,
skip_validation=True, # Raw JSON response
scope="consolidation",
)
# Parse JSON response - should be an array
if isinstance(result, str):
result = json.loads(result)
# Ensure result is a list
if isinstance(result, list):
return result
# Handle legacy single-action format for backward compatibility
if isinstance(result, dict):
if result.get("related_ids") and result.get("consolidated_text"):
# Convert old format to new format
return [
{
"action": "update",
"learning_id": result["related_ids"][0],
"text": result["consolidated_text"],
"reason": result.get("reason", ""),
}
]
return []
return []
except Exception as e:
logger.warning(f"Error in consolidation LLM call: {e}")
return []
async def _create_observation_directly(
conn: "Connection",
memory_engine: "MemoryEngine",
bank_id: str,
source_memory_id: uuid.UUID,
observation_text: str,
tags: list[str] | None = None,
event_date: datetime | None = None,
occurred_start: datetime | None = None,
occurred_end: datetime | None = None,
mentioned_at: datetime | None = None,
perf: ConsolidationPerfLog | None = None,
) -> dict[str, Any]:
"""
Create an observation directly with pre-processed text (no LLM call).
Used when the classify LLM has already provided the learning text.
This avoids the redundant second LLM call.
"""
# Generate embedding for the observation (convert to string for pgvector)
t0 = time.time()
embeddings = await embedding_utils.generate_embeddings_batch(memory_engine.embeddings, [observation_text])
embedding_str = str(embeddings[0]) if embeddings else None
if perf:
perf.record_timing("embedding", time.time() - t0)
# Create the observation as a memory_unit
now = datetime.now(timezone.utc)
obs_event_date = event_date or now
obs_occurred_start = occurred_start or now
obs_occurred_end = occurred_end or now
obs_mentioned_at = mentioned_at or now
obs_tags = tags or []
t0 = time.time()
observation_id = uuid.uuid4()
row = await conn.fetchrow(
f"""
INSERT INTO {fq_table("memory_units")} (
id, bank_id, text, fact_type, embedding, proof_count, source_memory_ids, history,
tags, event_date, occurred_start, occurred_end, mentioned_at
)
VALUES ($1, $2, $3, 'observation', $4::vector, 1, $5, '[]'::jsonb, $6, $7, $8, $9, $10)
RETURNING id
""",
observation_id,
bank_id,
observation_text,
embedding_str,
[source_memory_id],
obs_tags,
obs_event_date,
obs_occurred_start,
obs_occurred_end,
obs_mentioned_at,
)
# Create links between memory and observation (includes entity links, memory_links)
await _create_memory_links(conn, source_memory_id, observation_id)
if perf:
perf.record_timing("db_write", time.time() - t0)
logger.debug(f"Created observation {observation_id} from memory {source_memory_id} (tags: {obs_tags})")
return {"action": "created", "observation_id": str(row["id"]), "tags": obs_tags}
@@ -0,0 +1,77 @@
"""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.
You must output ONLY valid JSON with no markdown formatting, no code blocks, and no additional text.
## EXTRACT DURABLE KNOWLEDGE, NOT EPHEMERAL STATE
Facts often describe events or actions. Extract the DURABLE KNOWLEDGE implied by the fact, not the transient state.
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 history (e.g., "used to X, now Y")
3. UPDATE: New state replacing old state → update with history
## CRITICAL RULES:
- NEVER merge facts about DIFFERENT people
- NEVER merge unrelated topics (food preferences vs work vs hobbies)
- When merging contradictions, capture the CHANGE (before → after)
- 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):
{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)
- created_at/updated_at: when observation was created/modified
- 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 []
Output JSON array of actions:
[
{{"action": "update", "learning_id": "uuid-from-observations", "text": "updated knowledge", "reason": "..."}},
{{"action": "create", "text": "new durable knowledge", "reason": "..."}}
]
Return [] if fact contains no durable knowledge."""
@@ -20,6 +20,7 @@ from ..config import (
DEFAULT_RERANKER_FLASHRANK_CACHE_DIR,
DEFAULT_RERANKER_FLASHRANK_MODEL,
DEFAULT_RERANKER_LITELLM_MODEL,
DEFAULT_RERANKER_LOCAL_FORCE_CPU,
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT,
DEFAULT_RERANKER_LOCAL_MODEL,
DEFAULT_RERANKER_PROVIDER,
@@ -33,6 +34,7 @@ from ..config import (
ENV_RERANKER_FLASHRANK_CACHE_DIR,
ENV_RERANKER_FLASHRANK_MODEL,
ENV_RERANKER_LITELLM_MODEL,
ENV_RERANKER_LOCAL_FORCE_CPU,
ENV_RERANKER_LOCAL_MAX_CONCURRENT,
ENV_RERANKER_LOCAL_MODEL,
ENV_RERANKER_PROVIDER,
@@ -99,7 +101,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
_executor: ThreadPoolExecutor | None = None
_max_concurrent: int = 4 # Limit concurrent CPU-bound reranking calls
def __init__(self, model_name: str | None = None, max_concurrent: int = 4):
def __init__(self, model_name: str | None = None, max_concurrent: int = 4, force_cpu: bool = False):
"""
Initialize local SentenceTransformers cross-encoder.
@@ -108,8 +110,11 @@ class LocalSTCrossEncoder(CrossEncoderModel):
Default: cross-encoder/ms-marco-MiniLM-L-6-v2
max_concurrent: Maximum concurrent reranking calls (default: 2).
Higher values may cause CPU thrashing under load.
force_cpu: Force CPU mode (avoids MPS/XPC issues on macOS in daemon mode).
Default: False
"""
self.model_name = model_name or DEFAULT_RERANKER_LOCAL_MODEL
self.force_cpu = force_cpu
self._model = None
LocalSTCrossEncoder._max_concurrent = max_concurrent
@@ -130,13 +135,38 @@ class LocalSTCrossEncoder(CrossEncoderModel):
"Install it with: pip install sentence-transformers"
)
# Note: We use CPU even when GPU/MPS is available because:
# 1. The reranker model (MiniLM) is tiny (~22M params)
# 2. Batch sizes are small (~100-200 pairs)
# 3. Data transfer overhead to GPU outweighs compute benefit
# 4. CPU inference is actually faster for this workload
logger.info(f"Reranker: initializing local provider with model {self.model_name}")
self._model = CrossEncoder(self.model_name)
# Determine device based on hardware availability.
# We always set low_cpu_mem_usage=False to prevent lazy loading (meta tensors)
# which can cause issues when accelerate is installed but no GPU is available.
# Note: We do NOT use device_map because CrossEncoder internally calls .to(device)
# after loading, which conflicts with accelerate's device_map handling.
import torch
# Force CPU mode if configured (used in daemon mode to avoid MPS/XPC issues on macOS)
if self.force_cpu:
device = "cpu"
logger.info("Reranker: forcing CPU mode (HINDSIGHT_API_RERANKER_LOCAL_FORCE_CPU=1)")
else:
# Check for GPU (CUDA) or Apple Silicon (MPS)
# Wrap in try-except to gracefully handle any device detection issues
# (e.g., in CI environments or when PyTorch is built without GPU support)
device = "cpu" # Default to CPU
try:
has_gpu = torch.cuda.is_available() or (
hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
)
if has_gpu:
device = None # Let sentence-transformers auto-detect GPU/MPS
except Exception as e:
logger.warning(f"Failed to detect GPU/MPS, falling back to CPU: {e}")
self._model = CrossEncoder(
self.model_name,
device=device,
model_kwargs={"low_cpu_mem_usage": False},
)
# Initialize shared executor (limited workers naturally limits concurrency)
if LocalSTCrossEncoder._executor is None:
@@ -148,6 +178,11 @@ class LocalSTCrossEncoder(CrossEncoderModel):
else:
logger.info("Reranker: local provider initialized (using existing executor)")
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
"""Synchronous prediction wrapper for thread pool execution."""
scores = self._model.predict(pairs, show_progress_bar=False)
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs for relevance.
@@ -165,11 +200,11 @@ class LocalSTCrossEncoder(CrossEncoderModel):
# Use dedicated executor - limited workers naturally limits concurrency
loop = asyncio.get_event_loop()
scores = await loop.run_in_executor(
return await loop.run_in_executor(
LocalSTCrossEncoder._executor,
lambda: self._model.predict(pairs, show_progress_bar=False),
self._predict_sync,
pairs,
)
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
class RemoteTEICrossEncoder(CrossEncoderModel):
@@ -768,29 +803,33 @@ class LiteLLMCrossEncoder(CrossEncoderModel):
def create_cross_encoder_from_env() -> CrossEncoderModel:
"""
Create a CrossEncoderModel instance based on environment variables.
Create a CrossEncoderModel instance based on configuration.
See hindsight_api.config for environment variable names and defaults.
Reads configuration via get_config() to ensure consistency across the codebase.
Returns:
Configured CrossEncoderModel instance
"""
provider = os.environ.get(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER).lower()
from ..config import get_config
config = get_config()
provider = config.reranker_provider.lower()
if provider == "tei":
url = os.environ.get(ENV_RERANKER_TEI_URL)
url = config.reranker_tei_url
if not url:
raise ValueError(f"{ENV_RERANKER_TEI_URL} is required when {ENV_RERANKER_PROVIDER} is 'tei'")
batch_size = int(os.environ.get(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE)))
max_concurrent = int(os.environ.get(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT)))
return RemoteTEICrossEncoder(base_url=url, batch_size=batch_size, max_concurrent=max_concurrent)
return RemoteTEICrossEncoder(
base_url=url,
batch_size=config.reranker_tei_batch_size,
max_concurrent=config.reranker_tei_max_concurrent,
)
elif provider == "local":
model = os.environ.get(ENV_RERANKER_LOCAL_MODEL)
model_name = model or DEFAULT_RERANKER_LOCAL_MODEL
max_concurrent = int(
os.environ.get(ENV_RERANKER_LOCAL_MAX_CONCURRENT, str(DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT))
return LocalSTCrossEncoder(
model_name=config.reranker_local_model,
max_concurrent=config.reranker_local_max_concurrent,
force_cpu=config.reranker_local_force_cpu,
)
return LocalSTCrossEncoder(model_name=model_name, max_concurrent=max_concurrent)
elif provider == "cohere":
api_key = os.environ.get(ENV_COHERE_API_KEY)
if not api_key:
@@ -0,0 +1,5 @@
"""Directives module for hard rules injected into prompts."""
from .models import Directive
__all__ = ["Directive"]
@@ -0,0 +1,37 @@
"""Pydantic models for directives."""
from datetime import datetime, timezone
from uuid import UUID
from pydantic import BaseModel, Field
class Directive(BaseModel):
"""A directive is a hard rule injected into prompts.
Directives are user-defined rules that guide agent behavior. Unlike mental models
which are automatically consolidated from memories, directives are explicit
instructions that are always included in relevant prompts.
Examples:
- "Always respond in formal English"
- "Never share personal data with third parties"
- "Prefer conservative investment recommendations"
"""
id: UUID = Field(description="Unique identifier")
bank_id: str = Field(description="Bank this directive belongs to")
name: str = Field(description="Human-readable name")
content: str = Field(description="The directive text to inject into prompts")
priority: int = Field(default=0, description="Higher priority directives are injected first")
is_active: bool = Field(default=True, description="Whether this directive is currently active")
tags: list[str] = Field(default_factory=list, description="Tags for filtering")
created_at: datetime = Field(
default_factory=lambda: datetime.now(timezone.utc), description="When this directive was created"
)
updated_at: datetime = Field(
default_factory=lambda: datetime.now(timezone.utc), description="When this directive was last updated"
)
class Config:
from_attributes = True
@@ -18,6 +18,7 @@ import httpx
from ..config import (
DEFAULT_EMBEDDINGS_COHERE_MODEL,
DEFAULT_EMBEDDINGS_LITELLM_MODEL,
DEFAULT_EMBEDDINGS_LOCAL_FORCE_CPU,
DEFAULT_EMBEDDINGS_LOCAL_MODEL,
DEFAULT_EMBEDDINGS_OPENAI_MODEL,
DEFAULT_EMBEDDINGS_PROVIDER,
@@ -26,6 +27,7 @@ from ..config import (
ENV_EMBEDDINGS_COHERE_BASE_URL,
ENV_EMBEDDINGS_COHERE_MODEL,
ENV_EMBEDDINGS_LITELLM_MODEL,
ENV_EMBEDDINGS_LOCAL_FORCE_CPU,
ENV_EMBEDDINGS_LOCAL_MODEL,
ENV_EMBEDDINGS_OPENAI_API_KEY,
ENV_EMBEDDINGS_OPENAI_BASE_URL,
@@ -92,15 +94,18 @@ class LocalSTEmbeddings(Embeddings):
The embedding dimension is auto-detected from the model.
"""
def __init__(self, model_name: str | None = None):
def __init__(self, model_name: str | None = None, force_cpu: bool = False):
"""
Initialize local SentenceTransformers embeddings.
Args:
model_name: Name of the SentenceTransformer model to use.
Default: BAAI/bge-small-en-v1.5
force_cpu: Force CPU mode (avoids MPS/XPC issues on macOS in daemon mode).
Default: False
"""
self.model_name = model_name or DEFAULT_EMBEDDINGS_LOCAL_MODEL
self.force_cpu = force_cpu
self._model = None
self._dimension: int | None = None
@@ -128,11 +133,34 @@ class LocalSTEmbeddings(Embeddings):
)
logger.info(f"Embeddings: initializing local provider with model {self.model_name}")
# Disable lazy loading (meta tensors) which causes issues with newer transformers/accelerate
# Setting low_cpu_mem_usage=False and device_map=None ensures tensors are fully materialized
# Determine device based on hardware availability.
# We always set low_cpu_mem_usage=False to prevent lazy loading (meta tensors)
# which can cause issues when accelerate is installed but no GPU is available.
import torch
# Force CPU mode if configured (used in daemon mode to avoid MPS/XPC issues on macOS)
if self.force_cpu:
device = "cpu"
logger.info("Embeddings: forcing CPU mode")
else:
# Check for GPU (CUDA) or Apple Silicon (MPS)
# Wrap in try-except to gracefully handle any device detection issues
# (e.g., in CI environments or when PyTorch is built without GPU support)
device = "cpu" # Default to CPU
try:
has_gpu = torch.cuda.is_available() or (
hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
)
if has_gpu:
device = None # Let sentence-transformers auto-detect GPU/MPS
except Exception as e:
logger.warning(f"Failed to detect GPU/MPS, falling back to CPU: {e}")
self._model = SentenceTransformer(
self.model_name,
model_kwargs={"low_cpu_mem_usage": False, "device_map": None},
device=device,
model_kwargs={"low_cpu_mem_usage": False},
)
self._dimension = self._model.get_sentence_embedding_dimension()
@@ -150,6 +178,7 @@ class LocalSTEmbeddings(Embeddings):
"""
if self._model is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
embeddings = self._model.encode(texts, convert_to_numpy=True, show_progress_bar=False)
return [emb.tolist() for emb in embeddings]
@@ -673,24 +702,28 @@ class LiteLLMEmbeddings(Embeddings):
def create_embeddings_from_env() -> Embeddings:
"""
Create an Embeddings instance based on environment variables.
Create an Embeddings instance based on configuration.
See hindsight_api.config for environment variable names and defaults.
Reads configuration via get_config() to ensure consistency across the codebase.
Returns:
Configured Embeddings instance
"""
provider = os.environ.get(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER).lower()
from ..config import get_config
config = get_config()
provider = config.embeddings_provider.lower()
if provider == "tei":
url = os.environ.get(ENV_EMBEDDINGS_TEI_URL)
url = config.embeddings_tei_url
if not url:
raise ValueError(f"{ENV_EMBEDDINGS_TEI_URL} is required when {ENV_EMBEDDINGS_PROVIDER} is 'tei'")
return RemoteTEIEmbeddings(base_url=url)
elif provider == "local":
model = os.environ.get(ENV_EMBEDDINGS_LOCAL_MODEL)
model_name = model or DEFAULT_EMBEDDINGS_LOCAL_MODEL
return LocalSTEmbeddings(model_name=model_name)
return LocalSTEmbeddings(
model_name=config.embeddings_local_model,
force_cpu=config.embeddings_local_force_cpu,
)
elif provider == "openai":
# Use dedicated embeddings API key, or fall back to LLM API key
api_key = os.environ.get(ENV_EMBEDDINGS_OPENAI_API_KEY) or os.environ.get(ENV_LLM_API_KEY)
@@ -647,7 +647,13 @@ class LLMProvider:
success=True,
)
return LLMToolCallResult(content=content, tool_calls=tool_calls, finish_reason=finish_reason)
return LLMToolCallResult(
content=content,
tool_calls=tool_calls,
finish_reason=finish_reason,
input_tokens=input_tokens,
output_tokens=output_tokens,
)
except APIConnectionError as e:
last_exception = e
@@ -797,6 +803,10 @@ class LLMProvider:
content = "".join(content_parts) if content_parts else None
finish_reason = "tool_calls" if tool_calls else "stop"
# Extract token usage
input_tokens = response.usage.input_tokens or 0
output_tokens = response.usage.output_tokens or 0
# Record metrics
metrics = get_metrics_collector()
metrics.record_llm_call(
@@ -804,12 +814,18 @@ class LLMProvider:
model=self.model,
scope=scope,
duration=time.time() - start_time,
input_tokens=response.usage.input_tokens or 0,
output_tokens=response.usage.output_tokens or 0,
input_tokens=input_tokens,
output_tokens=output_tokens,
success=True,
)
return LLMToolCallResult(content=content, tool_calls=tool_calls, finish_reason=finish_reason)
return LLMToolCallResult(
content=content,
tool_calls=tool_calls,
finish_reason=finish_reason,
input_tokens=input_tokens,
output_tokens=output_tokens,
)
except (APIConnectionError, APIStatusError) as e:
if isinstance(e, APIStatusError) and e.status_code in (401, 403):
@@ -930,7 +946,13 @@ class LLMProvider:
success=True,
)
return LLMToolCallResult(content=content, tool_calls=tool_calls, finish_reason=finish_reason)
return LLMToolCallResult(
content=content,
tool_calls=tool_calls,
finish_reason=finish_reason,
input_tokens=input_tokens,
output_tokens=output_tokens,
)
except genai_errors.APIError as e:
if e.code in (401, 403):
File diff suppressed because it is too large Load Diff
@@ -1,16 +1,12 @@
"""
Mental models module for Hindsight.
Mental models are synthesized summaries that represent understanding. They come
in different subtypes based on how they were created:
Mental models contain directives - hard rules that are injected into reflect prompts.
Directives are user-defined and their observations are user-provided (not LLM-generated).
- Structural: Derived from the bank's mission (e.g., "Be a PM for engineering team")
These are created upfront based on what any agent with this role would need.
- Emergent: Discovered from data patterns (named entities, temporal clusters, etc.)
These surface organically as facts are retained.
- Pinned: User-defined models that persist across refreshes.
Other types of consolidated knowledge are handled by:
- Learnings: Automatic bottom-up consolidation from facts
- Pinned Reflections: User-curated living documents
"""
from .models import MentalModel, MentalModelSubtype
@@ -1,311 +0,0 @@
"""
Emergent mental model detection and promotion.
Emergent models are discovered from data patterns:
- Named entity extraction (people, projects, systems)
- Temporal clustering (events with multiple references)
- Causal patterns ("Because X, we do Y")
- Behavioral anchors ("After X, we started Y")
- Reference frequency (anything mentioned repeatedly)
When a pattern is detected, it goes through a mission filter to check relevance,
and if relevant, is promoted to a mental model.
"""
import logging
from typing import TYPE_CHECKING
from pydantic import BaseModel, Field
from .models import EmergentCandidate
if TYPE_CHECKING:
from ..llm_wrapper import LLMConfig
logger = logging.getLogger(__name__)
class MissionFilterCandidate(BaseModel):
"""Result of mission filtering for a single candidate."""
name: str
promote: bool = Field(description="True if this is a specific named entity worth tracking")
reason: str = Field(description="Brief explanation for the decision")
class MissionFilterResponse(BaseModel):
"""Response from LLM for mission filtering."""
candidates: list[MissionFilterCandidate] = Field(description="Filtering decision for each candidate")
def build_mission_filter_prompt(mission: str, candidates: list[EmergentCandidate]) -> str:
"""Build the prompt for filtering candidates by mission relevance."""
candidate_list = "\n".join(
[f"- {c.name} (mentions: {c.mention_count}, method: {c.detection_method})" for c in candidates]
)
return f"""Filter these detected entities. For each one, decide: promote=true or promote=false.
MISSION: {mission}
DETECTED ENTITIES:
{candidate_list}
=== DECISION RULES ===
Set promote=true ONLY for specific, named entities:
- Person names: "John", "Maria", "Alice Chen", "Dr. Smith"
- Named organizations: "Google", "Acme Corp", "Frontend Team"
- Named places: "Central Park Zoo", "NYC Office", "Building A"
- Named projects: "Project Phoenix", "Auth Service v2"
Set promote=false for EVERYTHING ELSE, including:
- Common English words: user, support, help, family, kids, parents, friends, people, team, photo, nature, park, office, home, work, school, joy, love, hope, fear, anger, gratitude, kindness, passion, motivation, inspiration, encouragement, positivity, energy, community, connection, commitment, collaboration, growth, impact, difference, success, progress, change, education, volunteering, veterans, homeless, shelter, meeting, project, system, process, event
- Generic categories (even capitalized): Users, Customers, Team, Family, Kids, Veterans, Community
- Abstract concepts: motivation, inspiration, gratitude, commitment, resilience
THE TEST: Is this a specific name you'd find in a contact list or org chart?
- "John" → YES (promote=true)
- "kids" → NO (promote=false)
- "community" → NO (promote=false)
- "Maria" → YES (promote=true)
- "park" → NO (promote=false)
When in doubt, set promote=false."""
def get_mission_filter_system_message() -> str:
"""System message for mission filtering."""
return """You filter entities for promotion. Output JSON with 'candidates' array.
Rules:
- promote=true ONLY for specific names (people, organizations, named places/projects)
- promote=false for common words, generic categories, abstract concepts
Examples:
- "John" → promote=true (person name)
- "kids" → promote=false (generic category)
- "community" → promote=false (abstract concept)
- "Google" → promote=true (organization name)
- "motivation" → promote=false (abstract concept)
When in doubt, promote=false. Most entities should be rejected."""
async def filter_candidates_by_mission(
llm_config: "LLMConfig",
mission: str,
candidates: list[EmergentCandidate],
) -> list[EmergentCandidate]:
"""
Filter emergent candidates to keep only specific, named entities.
Args:
llm_config: LLM configuration
mission: The bank's mission (used for context)
candidates: List of detected candidates
Returns:
Filtered list of candidates that are specific named entities
"""
if not candidates:
return []
if not mission:
# No mission = no filtering, keep all candidates
logger.debug("[EMERGENT] No mission set, skipping filter")
return candidates
prompt = build_mission_filter_prompt(mission, candidates)
try:
result = await llm_config.call(
messages=[
{"role": "system", "content": get_mission_filter_system_message()},
{"role": "user", "content": prompt},
],
response_format=MissionFilterResponse,
scope="mental_model_mission_filter",
)
# Build name -> promote map
promote_map = {c.name: c.promote for c in result.candidates}
# Filter candidates
filtered = []
for candidate in candidates:
if candidate.name in promote_map:
if promote_map[candidate.name]:
filtered.append(candidate)
logger.debug(f"[EMERGENT] Promoting '{candidate.name}'")
else:
logger.debug(f"[EMERGENT] Rejecting '{candidate.name}'")
else:
# Candidate not in response - reject by default
logger.debug(f"[EMERGENT] '{candidate.name}' not in response, rejecting")
logger.info(f"[EMERGENT] Mission filter: {len(filtered)}/{len(candidates)} candidates promoted")
return filtered
except Exception as e:
logger.warning(f"[EMERGENT] Mission filter failed, rejecting all candidates: {e}")
return []
async def evaluate_emergent_models(
llm_config: "LLMConfig",
models: list[dict],
) -> list[str]:
"""
Evaluate existing emergent models to check if they should be kept.
This re-evaluates emergent models using the same filtering criteria
as new candidates. Models that are generic/abstract will be removed.
Args:
llm_config: LLM configuration
models: List of existing emergent model dicts with 'name', 'id'
Returns:
List of model IDs that should be REMOVED (no longer valid)
"""
if not models:
return []
# Convert existing models to candidates for evaluation
candidates = [
EmergentCandidate(
name=m["name"],
detection_method="existing_emergent_model",
mention_count=0,
)
for m in models
]
# Build a simple prompt for re-evaluation
names_list = "\n".join([f"- {m['name']}" for m in models])
prompt = f"""Re-evaluate these existing mental models. For each one, decide: promote=true (keep) or promote=false (remove).
EXISTING MODELS:
{names_list}
=== DECISION RULES ===
Set promote=true ONLY for specific, named entities:
- Person names: "John", "Maria", "Alice Chen", "Dr. Smith"
- Named organizations: "Google", "Acme Corp", "Frontend Team"
- Named places: "Central Park Zoo", "NYC Office", "Building A"
- Named projects: "Project Phoenix", "Auth Service v2"
Set promote=false for EVERYTHING ELSE, including:
- Common English words: user, support, help, family, kids, parents, friends, people, team, photo, nature, park, office, home, work, school, joy, love, hope, fear, anger, gratitude, kindness, passion, motivation, inspiration, encouragement, positivity, energy, community, connection, commitment, collaboration, growth, impact, difference, success, progress, change, education, volunteering, veterans, homeless, shelter, meeting, project, system, process, event
- Generic categories (even capitalized): Users, Customers, Team, Family, Kids, Veterans, Community
- Abstract concepts: motivation, inspiration, gratitude, commitment, resilience
THE TEST: Is this a specific name you'd find in a contact list or org chart?
- "John" → YES (promote=true)
- "kids" → NO (promote=false)
- "community" → NO (promote=false)
When in doubt, set promote=false."""
try:
result = await llm_config.call(
messages=[
{"role": "system", "content": get_mission_filter_system_message()},
{"role": "user", "content": prompt},
],
response_format=MissionFilterResponse,
scope="mental_model_emergent_evaluation",
)
# Build name -> promote map
promote_map = {c.name: c.promote for c in result.candidates}
# Find models to remove
models_to_remove = []
for model in models:
name = model["name"]
if name in promote_map:
if not promote_map[name]:
models_to_remove.append(model["id"])
else:
logger.debug(f"[EMERGENT] Keeping '{name}'")
else:
# Model not in response - remove to be safe
logger.info(f"[EMERGENT] '{name}' not in evaluation response, marking for removal")
models_to_remove.append(model["id"])
logger.info(f"[EMERGENT] Evaluation: {len(models_to_remove)}/{len(models)} emergent models marked for removal")
return models_to_remove
except Exception as e:
logger.warning(f"[EMERGENT] Evaluation failed, keeping all models: {e}")
return []
async def detect_entity_candidates(
pool,
bank_id: str,
min_mentions: int = 5,
top_percent: int = 20,
) -> list[EmergentCandidate]:
"""
Detect entities that are candidates for promotion to mental models.
Args:
pool: Database connection pool
bank_id: Bank identifier
min_mentions: Minimum mention count to consider
top_percent: Only consider top X% by mention count
Returns:
List of entity candidates
"""
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
candidates = []
async with acquire_with_retry(pool) as conn:
# Get entities that meet criteria and don't already have mental models
rows = await conn.fetch(
f"""
WITH ranked AS (
SELECT
e.id,
e.canonical_name,
e.mention_count,
PERCENT_RANK() OVER (ORDER BY e.mention_count DESC) as rank_pct
FROM {fq_table("entities")} e
LEFT JOIN {fq_table("mental_models")} mm
ON mm.entity_id = e.id AND mm.bank_id = e.bank_id
WHERE e.bank_id = $1
AND e.mention_count >= $2
AND mm.id IS NULL -- Not already a mental model
)
SELECT id, canonical_name, mention_count
FROM ranked
WHERE rank_pct <= $3
ORDER BY mention_count DESC
LIMIT 50
""",
bank_id,
min_mentions,
top_percent / 100.0,
)
for row in rows:
candidates.append(
EmergentCandidate(
name=row["canonical_name"],
detection_method="named_entity_extraction",
mention_count=row["mention_count"],
entity_id=str(row["id"]),
relevance_score=0.0,
)
)
logger.debug(f"[EMERGENT] Detected {len(candidates)} entity candidates")
return candidates
@@ -9,12 +9,15 @@ from pydantic import BaseModel, Field
class MentalModelSubtype(str, Enum):
"""Subtype of mental model - how it was created."""
"""Subtype of mental model.
STRUCTURAL = "structural" # Derived from mission, created upfront
EMERGENT = "emergent" # Discovered from data patterns
LEARNED = "learned" # Formed through reflection
PINNED = "pinned" # User-defined, persists across refreshes
Currently only DIRECTIVE is supported. Other types of consolidated knowledge
are handled by:
- Learnings: Automatic bottom-up consolidation from facts
- Pinned Reflections: User-curated living documents
"""
DIRECTIVE = "directive" # User-defined hard rules, observations user-provided
class MentalModel(BaseModel):
@@ -48,50 +51,3 @@ class MentalModel(BaseModel):
created_at: datetime = Field(
default_factory=lambda: datetime.now(timezone.utc), description="When this model was created"
)
class StructuralModelTemplate(BaseModel):
"""
A template for a structural mental model.
Generated by LLM based on the bank's mission. Represents what any agent
with this role would need to track.
"""
id: str = Field(default="", description="Existing model ID to keep, or empty for new models")
name: str = Field(description="Human-readable name")
description: str = Field(description="What this model should track")
initial_probes: list[str] = Field(default_factory=list, description="Initial search queries to populate this model")
class StructuralModelDerivationResponse(BaseModel):
"""Response from LLM for structural model derivation."""
templates: list[StructuralModelTemplate] = Field(description="Structural model templates derived from the mission")
class EmergentCandidate(BaseModel):
"""
A candidate for promotion to emergent mental model.
Detected through pattern analysis of facts.
"""
name: str = Field(description="Name of the detected pattern/entity")
detection_method: str = Field(description="How this candidate was detected")
mention_count: int = Field(default=0, description="How many times referenced")
entity_id: str | None = Field(default=None, description="Entity ID if detected as entity")
relevance_score: float = Field(default=0.0, description="Score from mission filter (0-1)")
class ResearchResult(BaseModel):
"""
Result from the research endpoint.
Contains the answer along with the mental models and facts used.
"""
answer: str = Field(description="The synthesized answer")
mental_models_used: list[str] = Field(default_factory=list, description="IDs of mental models that contributed")
facts_used: list[str] = Field(default_factory=list, description="Fact IDs that contributed")
question_type: str | None = Field(default=None, description="Detected question type (WHO, WHAT, HOW, etc.)")
@@ -1,228 +0,0 @@
"""
Structural mental model derivation from bank mission.
Structural models are derived from the bank's mission - they represent what
any agent with this role would need to track. For example:
Mission: "Be a PM for engineering team"
Structural models:
- Team Structure (who's on the team, roles)
- Project Overview (current projects, status)
- Processes (how releases work, how decisions are made)
- Key Systems (what we own, dependencies)
"""
import logging
from typing import TYPE_CHECKING
from pydantic import BaseModel, Field
from .models import StructuralModelTemplate
if TYPE_CHECKING:
from ..llm_wrapper import LLMConfig
logger = logging.getLogger(__name__)
class StructuralDerivationResponse(BaseModel):
"""Response from LLM for structural model derivation."""
templates: list[StructuralModelTemplate] = Field(description="Structural model templates derived from the mission")
class StructuralRelevanceResult(BaseModel):
"""Result of evaluating a structural model's relevance to the mission."""
name: str
relevant: bool
reason: str
class StructuralRelevanceResponse(BaseModel):
"""Response from LLM for structural model relevance evaluation."""
models: list[StructuralRelevanceResult] = Field(description="Relevance evaluation for each model")
def build_structural_derivation_prompt(mission: str, existing_models: list[dict] | None = None) -> str:
"""Build the prompt for deriving structural models from a mission."""
existing_section = ""
if existing_models:
model_list = "\n".join([f"- id='{m['id']}' name='{m['name']}': {m['description']}" for m in existing_models])
existing_section = f"""
EXISTING STRUCTURAL MODELS:
{model_list}
IMPORTANT: If keeping an existing model, you MUST return its EXACT 'id' value.
Models not included in your output will be REMOVED.
"""
return f"""Given this agent mission, identify the KEY THINGS to track to achieve it.
MISSION: {mission}
{existing_section}
IMPORTANT CONSTRAINTS:
- Return 0-3 structural models MAXIMUM (less is better!)
- Only include models for SPECIFIC, CONCRETE things the agent needs to track
- Each model must be DIRECTLY tied to achieving the mission
- If the mission is simple, return 0 models (empty array is fine)
- If existing models are provided and you want to keep one, use its EXACT id
- Do NOT create near-duplicates (e.g., don't create "topic-map" if "topic-connections" exists)
GOOD examples (specific, actionable):
- Mission: "Be a PM for engineering team""Team Members" (track who's on the team)
- Mission: "Track customer feedback""Customer Issues" (track specific complaints/requests)
- Mission: "Manage project X""Project X Milestones" (track progress)
BAD examples (too generic, don't create these):
- "Processes", "Workflows", "Key Systems", "Important Events"
- "Communication", "Collaboration", "Progress", "Status"
- Generic role-based models not tied to the specific mission
For each model:
1. id: Use EXACT existing id if keeping a model, or leave empty for new models
2. name: Short, specific name (e.g., "Team Members", "Sprint Goals")
3. description: One line describing what to track
4. initial_probes: 2-3 search queries to find relevant information
Return ONLY the models that should exist. Existing models not in your output will be deleted."""
def get_structural_derivation_system_message() -> str:
"""System message for structural model derivation."""
return """You identify the key things to track for a mission. Be VERY selective.
Rules:
- Maximum 3 models (prefer fewer)
- Only SPECIFIC, CONCRETE things - not generic categories
- Each must DIRECTLY help achieve the mission
- Empty array is valid if no models are truly needed
- If existing models are shown and you want to keep one, return its EXACT id
- Never create duplicates - if a similar model exists, keep the existing one
Output JSON with 'templates' array (can be empty)."""
def _normalize_id(text: str) -> str:
"""Normalize a string to a canonical form for comparison.
Removes common suffixes, pluralization, and normalizes separators.
"""
# Lowercase and normalize separators
normalized = text.lower().replace(" ", "-").replace("_", "-")
# Remove common suffixes that indicate the same concept
suffixes_to_remove = ["-map", "-list", "-overview", "-tracker", "-s"]
for suffix in suffixes_to_remove:
if normalized.endswith(suffix) and len(normalized) > len(suffix):
normalized = normalized[: -len(suffix)]
return normalized
def _find_similar_existing_id(new_id: str, existing_models: list[dict]) -> str | None:
"""Find an existing model ID that is similar to the new ID.
Returns the existing ID if a similar one is found, None otherwise.
"""
if not existing_models:
return None
new_normalized = _normalize_id(new_id)
for model in existing_models:
existing_id = model.get("id", "")
existing_normalized = _normalize_id(existing_id)
# Check if one is a prefix of the other (normalized)
if new_normalized.startswith(existing_normalized) or existing_normalized.startswith(new_normalized):
return existing_id
# Check if they're the same when normalized
if new_normalized == existing_normalized:
return existing_id
return None
async def derive_structural_models(
llm_config: "LLMConfig",
mission: str,
existing_models: list[dict] | None = None,
) -> tuple[list[StructuralModelTemplate], list[str]]:
"""
Derive structural model templates from a bank's mission.
This combines derivation and evaluation in one call. The LLM sees existing
models and decides which to keep. Any existing model not in the output
will be marked for removal.
Args:
llm_config: LLM configuration for calling the model
mission: The bank's mission (e.g., "Be a PM for engineering team")
existing_models: Optional list of existing model dicts with 'name', 'description', 'id'
Returns:
Tuple of (templates to create/keep, IDs of existing models to remove)
Raises:
Exception: If LLM call fails
"""
prompt = build_structural_derivation_prompt(mission, existing_models)
result = await llm_config.call(
messages=[
{"role": "system", "content": get_structural_derivation_system_message()},
{"role": "user", "content": prompt},
],
response_format=StructuralDerivationResponse,
scope="mental_model_structural_derivation",
)
templates = result.templates
logger.info(f"[STRUCTURAL] LLM returned {len(templates)} structural models")
# Build set of existing IDs for quick lookup
existing_ids = {m["id"] for m in existing_models} if existing_models else set()
# Process templates: validate IDs, deduplicate, assign stable IDs
processed_templates: list[StructuralModelTemplate] = []
kept_existing_ids: set[str] = set()
for template in templates:
# If LLM returned an ID, check if it's a valid existing ID
if template.id and template.id in existing_ids:
# LLM is keeping an existing model
kept_existing_ids.add(template.id)
processed_templates.append(template)
logger.info(f"[STRUCTURAL] Keeping existing model: {template.id}")
else:
# New model or LLM didn't return a valid ID
# Generate ID from name
generated_id = template.name.lower().replace(" ", "-").replace("_", "-")
# Check for similar existing models to prevent near-duplicates
similar_id = _find_similar_existing_id(generated_id, existing_models)
if similar_id and similar_id not in kept_existing_ids:
# Use the existing similar model instead of creating a new one
logger.info(f"[STRUCTURAL] Detected near-duplicate: '{generated_id}' matches existing '{similar_id}'")
template.id = similar_id
kept_existing_ids.add(similar_id)
else:
template.id = generated_id
processed_templates.append(template)
# Find existing models to remove (not kept in LLM output)
models_to_remove = []
if existing_models:
for model in existing_models:
if model["id"] not in kept_existing_ids:
logger.info(f"[STRUCTURAL] Marking '{model['name']}' (id={model['id']}) for removal")
models_to_remove.append(model["id"])
if models_to_remove:
logger.info(f"[STRUCTURAL] {len(models_to_remove)} existing models will be removed")
return processed_templates, models_to_remove
@@ -4,17 +4,15 @@ Reflect agent module for agentic reflection with tools.
The reflect agent uses an iterative loop with tools to:
1. Lookup mental models (existing knowledge)
2. Recall facts (semantic + temporal search)
3. Learn new insights (create/update mental models)
4. Expand memories (get chunk/document context)
3. Expand memories (get chunk/document context)
"""
from .agent import ReflectAgentResult, run_reflect_agent
from .models import MentalModelInput, ReflectAction, ReflectActionBatch
from .models import ReflectAction, ReflectActionBatch
__all__ = [
"run_reflect_agent",
"ReflectAgentResult",
"ReflectAction",
"ReflectActionBatch",
"MentalModelInput",
]
File diff suppressed because it is too large Load Diff
@@ -7,51 +7,28 @@ from typing import Any, Literal
from pydantic import BaseModel, Field
class MentalModelObservation(BaseModel):
"""An observation within a mental model with its supporting memories."""
class ObservationSection(BaseModel):
"""A section within an observation with its supporting memories."""
title: str = Field(description="Observation header (can be empty for intro)")
text: str = Field(description="Observation content - no headers, use lists/tables/bold")
memory_ids: list[str] = Field(default_factory=list, description="Memory IDs supporting this observation")
class MentalModelInput(BaseModel):
"""Input for the learn tool to create a mental model placeholder.
The agent only specifies name and description - the actual content/observations
are generated during refresh, similar to pinned models.
"""
name: str = Field(description="Human-readable name for the mental model")
description: str = Field(description="What to track - used as prompt for content generation during refresh")
entity_id: str | None = Field(default=None, description="Optional link to existing entity ID")
class AnswerSection(BaseModel):
"""A section of the answer with its supporting evidence (DEPRECATED)."""
title: str = Field(description="Section header/title")
text: str = Field(description="Section content")
title: str = Field(description="Section header (can be empty for intro)")
text: str = Field(description="Section content - no headers, use lists/tables/bold")
memory_ids: list[str] = Field(default_factory=list, description="Memory IDs supporting this section")
model_ids: list[str] = Field(default_factory=list, description="Mental model IDs supporting this section")
class ReflectAction(BaseModel):
"""Single action the reflect agent can take."""
tool: Literal["list_mental_models", "get_mental_model", "recall", "learn", "expand", "done"] = Field(
description="Tool to invoke: list_mental_models, get_mental_model, recall, learn, expand, or done"
tool: Literal["list_observations", "get_observation", "recall", "expand", "done"] = Field(
description="Tool to invoke: list_observations, get_observation, recall, expand, or done"
)
# Tool-specific parameters
model_id: str | None = Field(default=None, description="Mental model ID for get_mental_model")
observation_id: str | None = Field(default=None, description="Observation ID for get_observation")
query: str | None = Field(default=None, description="Search query for recall")
max_tokens: int | None = Field(default=None, description="Max tokens for recall results (default 2048)")
mental_model: MentalModelInput | None = Field(default=None, description="Mental model to create/update for learn")
memory_ids: list[str] | None = Field(default=None, description="Memory unit IDs for expand (batched)")
depth: Literal["chunk", "document"] | None = Field(default=None, description="Expansion depth for expand")
sections: list[AnswerSection] | None = Field(default=None, description="DEPRECATED: Use answer field instead")
observations: list[MentalModelObservation] | None = Field(
default=None, description="Observations for done action (when output_mode=observations)"
observation_sections: list[ObservationSection] | None = Field(
default=None, description="Observation sections for done action (when output_mode=observations)"
)
# Plain text answer fields (for output_mode=answer)
answer: str | None = Field(default=None, description="Plain text answer for done action (no markdown)")
@@ -73,7 +50,8 @@ class ReflectActionBatch(BaseModel):
class ToolCall(BaseModel):
"""A single tool call made during reflect."""
tool: str = Field(description="Tool name: lookup, recall, learn, expand")
tool: str = Field(description="Tool name: lookup, recall, expand")
reason: str | None = Field(default=None, description="Agent's reasoning for making this tool call")
input: dict = Field(description="Tool input parameters")
output: dict = Field(description="Tool output/result")
duration_ms: int = Field(description="Execution time in milliseconds")
@@ -85,30 +63,47 @@ class LLMCall(BaseModel):
scope: str = Field(description="Call scope: agent_1, agent_2, final, etc.")
duration_ms: int = Field(description="Execution time in milliseconds")
input_tokens: int = Field(default=0, description="Input tokens used")
output_tokens: int = Field(default=0, description="Output tokens used")
class Observation(BaseModel):
"""A single observation with supporting memories."""
class DirectiveInfo(BaseModel):
"""Information about a directive that was applied during reflect."""
title: str = Field(description="Observation title/header")
text: str = Field(description="Observation content")
memory_ids: list[str] = Field(default_factory=list, description="Memory IDs supporting this observation")
id: str = Field(description="Directive mental model ID")
name: str = Field(description="Directive name")
content: str = Field(description="Directive content")
class TokenUsageSummary(BaseModel):
"""Total token usage across all LLM calls."""
input_tokens: int = Field(default=0, description="Total input tokens used")
output_tokens: int = Field(default=0, description="Total output tokens used")
total_tokens: int = Field(default=0, description="Total tokens (input + output)")
class ReflectAgentResult(BaseModel):
"""Result from the reflect agent."""
text: str = Field(description="Final answer text")
observations: list[Observation] = Field(
default_factory=list, description="Structured observations (when output_mode=observations)"
)
structured_output: dict[str, Any] | None = Field(
default=None, description="Structured output parsed according to provided response_schema"
)
iterations: int = Field(default=0, description="Number of iterations taken")
tools_called: int = Field(default=0, description="Total number of tool calls made")
mental_models_created: list[str] = Field(default_factory=list, description="IDs of mental models created/updated")
tool_trace: list[ToolCall] = Field(default_factory=list, description="Trace of all tool calls made")
llm_trace: list[LLMCall] = Field(default_factory=list, description="Trace of all LLM calls made")
usage: TokenUsageSummary = Field(
default_factory=TokenUsageSummary, description="Total token usage across all LLM calls"
)
used_memory_ids: list[str] = Field(default_factory=list, description="Validated memory IDs actually used in answer")
used_model_ids: list[str] = Field(default_factory=list, description="Validated model IDs actually used in answer")
used_mental_model_ids: list[str] = Field(
default_factory=list, description="Validated mental model IDs actually used in answer"
)
used_observation_ids: list[str] = Field(
default_factory=list, description="Validated observation IDs actually used in answer"
)
directives_applied: list[DirectiveInfo] = Field(
default_factory=list, description="Directive mental models that affected this reflection"
)
@@ -0,0 +1,186 @@
"""
Models and utilities for evidence-grounded observations with computed trends.
Observations are part of mental models and represent patterns/beliefs derived
from memories. Each observation must be grounded in specific evidence (quotes)
from memories, and trends are computed algorithmically from evidence timestamps.
"""
from datetime import datetime, timedelta, timezone
from enum import Enum
from pydantic import BaseModel, Field, computed_field, field_validator
class Trend(str, Enum):
"""Computed trend for an observation based on evidence timestamps.
Trends indicate how an observation's evidence is distributed over time:
- STABLE: Evidence spread across time, continues to present
- STRENGTHENING: More/denser evidence recently than before
- WEAKENING: Evidence mostly old, sparse recently
- NEW: All evidence within recent window
- STALE: No evidence in recent window (may no longer apply)
"""
STABLE = "stable"
STRENGTHENING = "strengthening"
WEAKENING = "weakening"
NEW = "new"
STALE = "stale"
class ObservationEvidence(BaseModel):
"""A single piece of evidence supporting an observation.
Each evidence item must include an exact quote from the source memory
to ensure observations are grounded and verifiable.
"""
memory_id: str = Field(description="ID of the memory unit this evidence comes from")
quote: str = Field(description="Exact quote from the memory supporting the observation")
relevance: str = Field(default="", description="Brief explanation of how this quote supports the observation")
timestamp: datetime = Field(description="When the source memory was created")
@field_validator("timestamp", mode="before")
@classmethod
def ensure_timezone_aware(cls, v: datetime | str | None) -> datetime:
"""Ensure timestamp is always timezone-aware UTC."""
if v is None:
return datetime.now(timezone.utc)
if isinstance(v, str):
# Parse ISO format string, handling 'Z' suffix
v = datetime.fromisoformat(v.replace("Z", "+00:00"))
if isinstance(v, datetime):
if v.tzinfo is None:
return v.replace(tzinfo=timezone.utc)
return v
raise ValueError(f"Invalid timestamp type: {type(v)}")
class Observation(BaseModel):
"""A single observation within a mental model.
Observations represent patterns, preferences, beliefs, or other insights
derived from memories. Each observation must be grounded in evidence
with exact quotes from source memories.
"""
title: str = Field(description="Short summary title for the observation (5-10 words)")
content: str = Field(description="The observation content - detailed explanation of what we believe to be true")
evidence: list[ObservationEvidence] = Field(default_factory=list, description="Supporting evidence with quotes")
created_at: datetime = Field(
default_factory=lambda: datetime.now(timezone.utc), description="When this observation was first created"
)
@field_validator("created_at", mode="before")
@classmethod
def ensure_created_at_timezone_aware(cls, v: datetime | str | None) -> datetime:
"""Ensure created_at is always timezone-aware UTC."""
if v is None:
return datetime.now(timezone.utc)
if isinstance(v, str):
v = datetime.fromisoformat(v.replace("Z", "+00:00"))
if isinstance(v, datetime):
if v.tzinfo is None:
return v.replace(tzinfo=timezone.utc)
return v
raise ValueError(f"Invalid created_at type: {type(v)}")
@computed_field
@property
def trend(self) -> Trend:
"""Compute trend from evidence timestamps."""
return compute_trend(self.evidence)
@computed_field
@property
def evidence_span(self) -> dict[str, str | None]:
"""Get the time span covered by evidence."""
if not self.evidence:
return {"from": None, "to": None}
timestamps = [e.timestamp for e in self.evidence]
return {
"from": min(timestamps).isoformat(),
"to": max(timestamps).isoformat(),
}
@computed_field
@property
def evidence_count(self) -> int:
"""Number of evidence items supporting this observation."""
return len(self.evidence)
def compute_trend(
evidence: list[ObservationEvidence],
now: datetime | None = None,
recent_days: int = 30,
old_days: int = 90,
) -> Trend:
"""Compute the trend for an observation based on evidence timestamps.
The trend indicates how the evidence is distributed over time:
- STABLE: Evidence spread across time, continues to present
- STRENGTHENING: More evidence recently than historically
- WEAKENING: Evidence mostly old, sparse recently
- NEW: All evidence is recent (within recent_days)
- STALE: No evidence in recent window
Args:
evidence: List of evidence items with timestamps
now: Reference time for calculations (defaults to current UTC time)
recent_days: Number of days to consider "recent" (default 30)
old_days: Number of days to consider "old" (default 90)
Returns:
Computed Trend enum value
"""
if now is None:
now = datetime.now(timezone.utc)
# Ensure now is timezone-aware
if now.tzinfo is None:
now = now.replace(tzinfo=timezone.utc)
if not evidence:
return Trend.STALE
recent_cutoff = now - timedelta(days=recent_days)
old_cutoff = now - timedelta(days=old_days)
# Normalize timestamps to UTC for comparison
def normalize_ts(ts: datetime) -> datetime:
if ts.tzinfo is None:
return ts.replace(tzinfo=timezone.utc)
return ts
recent = [e for e in evidence if normalize_ts(e.timestamp) > recent_cutoff]
old = [e for e in evidence if normalize_ts(e.timestamp) < old_cutoff]
middle = [e for e in evidence if old_cutoff <= normalize_ts(e.timestamp) <= recent_cutoff]
# No recent evidence = stale
if not recent:
return Trend.STALE
# All evidence is recent = new
if not old and not middle:
return Trend.NEW
# Compare density (evidence per day)
recent_density = len(recent) / recent_days if recent_days > 0 else 0
older_period = old_days - recent_days
older_density = (len(old) + len(middle)) / older_period if older_period > 0 else 0
# Avoid division by zero
if older_density == 0:
return Trend.NEW
ratio = recent_density / older_density
if ratio > 1.5:
return Trend.STRENGTHENING
elif ratio < 0.5:
return Trend.WEAKENING
else:
return Trend.STABLE
@@ -1,147 +1,313 @@
"""
System prompts for the reflect agent.
The reflect agent uses hierarchical retrieval:
1. search_mental_models - User-curated summaries (highest quality)
2. search_observations - Consolidated knowledge with freshness awareness
3. recall - Raw facts as ground truth fallback
"""
import json
from typing import Any
def _extract_directive_rules(directives: list[dict[str, Any]]) -> list[str]:
"""
Extract directive rules as a list of strings.
Args:
directives: List of directives with name and content
Returns:
List of directive rule strings
"""
rules = []
for directive in directives:
directive_name = directive.get("name", "")
# New format: directives have direct content field
content = directive.get("content", "")
if content:
if directive_name:
rules.append(f"**{directive_name}**: {content}")
else:
rules.append(content)
else:
# Legacy format: check for observations
observations = directive.get("observations", [])
if observations:
for obs in observations:
# Support both Pydantic Observation objects and dicts
if hasattr(obs, "title"):
title = obs.title
obs_content = obs.content
else:
title = obs.get("title", "")
obs_content = obs.get("content", "")
if title and obs_content:
rules.append(f"**{title}**: {obs_content}")
elif obs_content:
rules.append(obs_content)
elif directive_name:
# Fallback to description
desc = directive.get("description", "")
if desc:
rules.append(f"**{directive_name}**: {desc}")
return rules
def build_directives_section(directives: list[dict[str, Any]]) -> str:
"""
Build the directives section for the system prompt.
Directives are hard rules that MUST be followed in all responses.
Args:
directives: List of directive mental models with observations
"""
if not directives:
return ""
rules = _extract_directive_rules(directives)
if not rules:
return ""
parts = [
"## DIRECTIVES (MANDATORY)",
"These are hard rules you MUST follow in ALL responses:",
"",
]
for rule in rules:
parts.append(f"- {rule}")
parts.extend(
[
"",
"NEVER violate these directives, even if other context suggests otherwise.",
"IMPORTANT: Do NOT explain or justify how you handled directives in your answer. Just follow them silently.",
"",
]
)
return "\n".join(parts)
def build_directives_reminder(directives: list[dict[str, Any]]) -> str:
"""
Build a reminder section for directives to place at the end of the prompt.
Args:
directives: List of directive mental models with observations
"""
if not directives:
return ""
rules = _extract_directive_rules(directives)
if not rules:
return ""
parts = [
"",
"## REMINDER: MANDATORY DIRECTIVES",
"Before responding, ensure your answer complies with ALL of these directives:",
"",
]
for i, rule in enumerate(rules, 1):
parts.append(f"{i}. {rule}")
parts.append("")
parts.append("Your response will be REJECTED if it violates any directive above.")
parts.append("Do NOT include any commentary about how you handled directives - just follow them.")
return "\n".join(parts)
def build_system_prompt_for_tools(
bank_profile: dict[str, Any],
context: str | None = None,
output_mode: str = "answer",
directives: list[dict[str, Any]] | None = None,
has_mental_models: bool = False,
budget: str | None = None,
) -> str:
"""
Build the system prompt for tool-calling reflect agent.
This is a simplified prompt since tools are defined separately via the tools parameter.
The agent uses hierarchical retrieval:
1. search_mental_models - User-curated summaries (try first, if available)
2. search_observations - Consolidated knowledge with freshness
3. recall - Raw facts as ground truth
Args:
bank_profile: Bank profile with name and mission
context: Optional additional context
output_mode: "answer" for plain text response, "observations" for structured observations
directives: Optional list of directive mental models to inject as hard rules
has_mental_models: Whether the bank has any mental models (skip if not)
budget: Search depth budget - "low", "mid", or "high". Controls exploration thoroughness.
"""
name = bank_profile.get("name", "Assistant")
mission = bank_profile.get("mission", "")
# Build critical rules based on mode
if output_mode == "observations":
no_info_rule = "- Only say 'I don't have information' AFTER trying recall with no relevant results"
else:
no_info_rule = (
"- Only say 'I don't have information' AFTER trying list_mental_models AND recall with no relevant results"
)
parts = []
parts = [
"You are a reflection agent that answers questions by reasoning over retrieved memories.",
"",
"## CRITICAL RULES",
"- You must NEVER fabricate information that has no basis in retrieved data",
"- You SHOULD synthesize, infer, and reason from the retrieved memories",
"- You MUST call recall() before saying you don't have information",
no_info_rule,
"",
"## How to Reason",
"- If memories mention someone did an activity, you can infer they likely enjoyed it",
"- Synthesize a coherent narrative from related memories",
"- Be a thoughtful interpreter, not just a literal repeater",
"- When the exact answer isn't stated, use what IS stated to give the best answer",
"",
"## Query Strategy (IMPORTANT)",
"recall() uses semantic search. NEVER just echo the user's question - decompose it into targeted searches:",
"",
"BAD: User asks 'recurring lesson themes between students' → recall('recurring lesson themes between students')",
"GOOD: Break it down into component searches:",
" 1. recall('lessons') - find all lesson-related memories",
" 2. recall('teaching sessions') - alternative phrasing",
" 3. recall('student progress') - find student-related memories",
" 4. recall('topics taught') - find subject matter",
"",
"Think: What ENTITIES and CONCEPTS does this question involve? Search for each separately.",
"- Questions about patterns → search for the individual instances first",
"- Questions comparing things → search for each thing separately",
"- Questions about relationships → search for each party involved",
"",
"## Workflow",
]
# Inject directives at the VERY START for maximum prominence
if directives:
parts.append(build_directives_section(directives))
# Mode-specific workflow and output format
if output_mode == "observations":
# Observations mode: for mental model generation - no mental model lookup tools
parts.extend(
[
"You are a reflection agent that answers questions by reasoning over retrieved memories.",
"",
]
)
parts.extend(
[
"## CRITICAL RULES",
"- You must NEVER fabricate information that has no basis in retrieved data",
"- You SHOULD synthesize, infer, and reason from the retrieved memories",
"- You MUST search before saying you don't have information",
"",
"## How to Reason",
"- If memories mention someone did an activity, you can infer they likely enjoyed it",
"- Synthesize a coherent narrative from related memories",
"- Be a thoughtful interpreter, not just a literal repeater",
"- When the exact answer isn't stated, use what IS stated to give the best answer",
"",
"## HIERARCHICAL RETRIEVAL STRATEGY",
"",
]
)
# Build retrieval levels based on what's available
if has_mental_models:
parts.extend(
[
"1. DECOMPOSE the topic into component searches (see Query Strategy above)",
" - Don't search for the topic name itself - search for related concepts",
" - Example for 'Coffee preferences': search 'coffee', 'drinks', 'morning routine', 'caffeine'",
"2. Run multiple recall() calls with varied, targeted queries",
"3. IMPORTANT: Use expand(memory_ids, 'chunk') to verify memories before using them",
" - Always verify the source chunk to confirm the memory is actually relevant",
" - Don't assume a memory is relevant based on the summary alone",
" - Only include memories you've verified via expand()",
"4. When ready, call done() with MULTIPLE structured observations",
"You have access to THREE levels of knowledge. Use them in this order:",
"",
"## Output Format: MULTIPLE Structured Observations",
"### 1. MENTAL MODELS (search_mental_models) - Try First",
"- User-curated summaries about specific topics",
"- HIGHEST quality - manually created and maintained",
"- If a relevant mental model exists and is FRESH, it may fully answer the question",
"- Check `is_stale` field - if stale, also verify with lower levels",
"",
"CRITICAL: You MUST create MULTIPLE separate observations in the array - one for each theme.",
"Do NOT put all content in a single observation!",
"### 2. OBSERVATIONS (search_observations) - Second Priority",
"- Auto-consolidated knowledge from memories",
"- Check `is_stale` field - if stale, ALSO use recall() to verify",
"- Good for understanding patterns and summaries",
"",
"- Create 3-8 separate observations, each as its OWN item in the observations array",
"- Each observation covers ONE specific theme (preferences, history, relationships, etc.)",
"- Each observation has: title (short header), text (content), memory_ids (full UUIDs)",
"### 3. RAW FACTS (recall) - Ground Truth",
"- Individual memories (world facts and experiences)",
"- Use when: no mental models/observations exist, they're stale, or you need specific details",
"- This is the source of truth that other levels are built from",
"",
"Text format for each observation:",
"- Main insight or finding (no markdown headers)",
"- End with 'Key evidence:' section containing DIRECT QUOTES from memories in *italics*",
"- Quote the actual memory text, don't summarize - use *italics* for citations",
"",
"Example done() call with MULTIPLE observations:",
"```json",
"{",
' "observations": [',
" {",
' "title": "Work Preferences",',
' "text": "Prefers async communication and flexible schedules.\\n\\nKey evidence:\\n- *I prefer Slack over calls for most communication*\\n- *Flexible hours help me do my best work*",',
' "memory_ids": ["abc123-full-uuid", "def456-full-uuid"]',
" },",
" {",
' "title": "Technical Background",',
' "text": "Has extensive ML experience spanning a decade.\\n\\nKey evidence:\\n- *I have 10 years of experience in machine learning*\\n- *Led the ML team at my previous company*",',
' "memory_ids": ["ghi789-full-uuid"]',
" }",
" ]",
"}",
"```",
]
)
else:
# Answer mode: include mental model lookup in workflow
parts.extend(
[
"1. Review the pre-fetched mental models for relevant synthesized knowledge",
"2. If relevant, call get_mental_model(model_id) for full observations",
"3. DECOMPOSE the question into component searches (see Query Strategy above)",
" - Identify entities and concepts in the question",
" - Search for each separately with targeted queries",
"4. Run multiple recall() calls - don't just echo the user's question",
"5. Use expand() if you need more context on specific memories",
"6. If you discover an important recurring topic worth tracking, use learn() to create a mental model",
"7. When ready, call done() with your answer and supporting memory_ids",
"You have access to TWO levels of knowledge. Use them in this order:",
"",
"## When to Use learn()",
"Use learn() to create a new mental model when you discover:",
"- A person, project, or concept that appears frequently in memories",
"- An important topic the user seems to care about but has no mental model for",
"- A pattern or relationship worth synthesizing for future reference",
"Example: learn(name='Project Alpha', description='Track goals, status, and key decisions for Project Alpha')",
"### 1. OBSERVATIONS (search_observations) - Try First",
"- Auto-consolidated knowledge from memories",
"- Check `is_stale` field - if stale, ALSO use recall() to verify",
"- Good for understanding patterns and summaries",
"",
"### 2. RAW FACTS (recall) - Ground Truth",
"- Individual memories (world facts and experiences)",
"- Use when: no observations exist, they're stale, or you need specific details",
"- This is the source of truth that observations are built from",
"",
"## Output Format: Plain Text Answer",
"Call done() with a plain text 'answer' field.",
"- Do NOT use markdown formatting",
"- NEVER include memory IDs, UUIDs, or 'Memory references' in the answer text",
"- Put memory IDs ONLY in the memory_ids array parameter, not in the answer",
]
)
parts.extend(
[
"## Query Strategy",
"recall() uses semantic search. NEVER just echo the user's question - decompose it into targeted searches:",
"",
"BAD: User asks 'recurring lesson themes between students' → recall('recurring lesson themes between students')",
"GOOD: Break it down into component searches:",
" 1. recall('lessons') - find all lesson-related memories",
" 2. recall('teaching sessions') - alternative phrasing",
" 3. recall('student progress') - find student-related memories",
"",
"Think: What ENTITIES and CONCEPTS does this question involve? Search for each separately.",
"",
]
)
# Add budget guidance
if budget:
budget_lower = budget.lower()
if budget_lower == "low":
parts.extend(
[
"## RESEARCH DEPTH: SHALLOW (Quick Response)",
"- Prioritize speed over completeness",
"- If mental models or observations provide a reasonable answer, stop there",
"- Only dig deeper if the initial results are clearly insufficient",
"- Prefer a quick overview rather than exhaustive details",
"- Answer promptly with available information",
"",
]
)
elif budget_lower == "mid":
parts.extend(
[
"## RESEARCH DEPTH: MODERATE (Balanced)",
"- Balance thoroughness with efficiency",
"- Check multiple sources when the question warrants it",
"- Verify stale data if it's central to the answer",
"- Don't over-explore, but ensure reasonable coverage",
"",
]
)
elif budget_lower == "high":
parts.extend(
[
"## RESEARCH DEPTH: DEEP (Thorough Exploration)",
"- Explore comprehensively before answering",
"- Search across all available knowledge levels",
"- Use multiple query variations to ensure coverage",
"- Verify information across different retrieval levels",
"- Use expand() to get full context on important memories",
"- Take time to synthesize a complete, well-researched answer",
"",
]
)
parts.append("## Workflow")
if has_mental_models:
parts.extend(
[
"1. First, try search_mental_models() - check if a curated summary exists",
"2. If no mental model or it's stale, try search_observations() for consolidated knowledge",
"3. If observations are stale OR you need specific details, use recall() for raw facts",
"4. Use expand() if you need more context on specific memories",
"5. When ready, call done() with your answer and supporting IDs",
]
)
else:
parts.extend(
[
"1. First, try search_observations() - check for consolidated knowledge",
"2. If observations are stale OR you need specific details, use recall() for raw facts",
"3. Use expand() if you need more context on specific memories",
"4. When ready, call done() with your answer and supporting IDs",
]
)
parts.extend(
[
"",
"## Output Format: Plain Text Answer",
"Call done() with a plain text 'answer' field.",
"- Do NOT use markdown formatting",
"- 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",
]
)
parts.append("")
parts.append(f"## Memory Bank: {name}")
@@ -164,6 +330,10 @@ def build_system_prompt_for_tools(
if context:
parts.append(f"\n## Additional Context\n{context}")
# Add directive reminder at the END for recency effect
if directives:
parts.append(build_directives_reminder(directives))
return "\n".join(parts)
@@ -228,9 +398,10 @@ def build_agent_prompt(
else:
parts.append(
"\n## Instructions\n"
"Start by calling list_mental_models() to see available mental models - they contain pre-synthesized knowledge. "
"If a relevant model exists, use get_mental_model(model_id) to get its observations. "
"Then use recall(query) for specific details not covered by mental models."
"Start by searching for relevant information using the hierarchical retrieval strategy:\n"
"1. Try search_mental_models() first for curated summaries\n"
"2. Try search_observations() for consolidated knowledge\n"
"3. Use recall() for specific details or to verify stale data"
)
return "\n".join(parts)
@@ -1,14 +1,17 @@
"""
Tool implementations for the reflect agent.
Implements hierarchical retrieval:
1. search_mental_models - User-curated stored reflect responses (highest quality)
2. search_observations - Consolidated knowledge with freshness
3. recall - Raw facts as ground truth
"""
import logging
import re
import uuid
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any
from .models import MentalModelInput
if TYPE_CHECKING:
from asyncpg import Connection
@@ -17,133 +20,215 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def generate_model_id(name: str) -> str:
"""Generate a stable ID from mental model name."""
# Normalize: lowercase, replace spaces/special chars with hyphens
normalized = re.sub(r"[^a-z0-9]+", "-", name.lower()).strip("-")
# Truncate to reasonable length
return normalized[:50]
# Observation is considered stale if not updated in this many days
STALE_THRESHOLD_DAYS = 7
async def tool_lookup(
async def tool_search_mental_models(
conn: "Connection",
bank_id: str,
model_id: str | None = None,
query: str,
query_embedding: list[float],
max_results: int = 5,
tags: list[str] | None = None,
tags_match: str = "any",
exclude_ids: list[str] | None = None,
) -> dict[str, Any]:
"""
List or get mental models.
Search user-curated mental models by semantic similarity.
Mental models are high-quality, manually created summaries about specific topics.
They should be searched FIRST as they represent the most reliable synthesized knowledge.
Args:
conn: Database connection
bank_id: Bank identifier
model_id: Optional specific model ID to get (if None, lists all)
tags: Optional tags to filter models (when listing)
query: Search query (for logging/tracing)
query_embedding: Pre-computed embedding for semantic search
max_results: Maximum number of mental models to return
tags: Optional tags to filter mental models
tags_match: How to match tags - "any" (OR), "all" (AND)
exclude_ids: Optional list of mental model IDs to exclude (e.g., when refreshing a mental model)
Returns:
Dict with either a list of models or a single model's details
Dict with matching mental models including content and freshness info
"""
if model_id:
# Get specific mental model with full details including observations
row = await conn.fetchrow(
"""
SELECT id, subtype, name, description, observations, entity_id, last_updated
FROM mental_models
WHERE id = $1 AND bank_id = $2
""",
model_id,
bank_id,
)
if row:
# Parse observations JSON
obs_data = row["observations"] or {"observations": []}
if isinstance(obs_data, str):
import json
from ..memory_engine import fq_table
obs_data = json.loads(obs_data)
observations_raw = obs_data.get("observations", []) if isinstance(obs_data, dict) else obs_data
# Build filters dynamically
filters = ""
params: list[Any] = [bank_id, str(query_embedding), max_results]
next_param = 4
# Normalize observation format: map memory_ids/fact_ids to based_on
observations = []
for obs in observations_raw:
if isinstance(obs, dict):
based_on = obs.get("memory_ids") or obs.get("fact_ids") or []
observations.append(
{
"title": obs.get("title", ""),
"text": obs.get("text", ""),
"based_on": based_on,
}
)
return {
"found": True,
"model": {
"id": row["id"],
"subtype": row["subtype"],
"name": row["name"],
"description": row["description"],
"observations": observations, # [{title, text, based_on}, ...]
"entity_id": str(row["entity_id"]) if row["entity_id"] else None,
"last_updated": row["last_updated"].isoformat() if row["last_updated"] else None,
},
}
return {"found": False, "model_id": model_id}
else:
# List mental models (compact: id, name, description only)
# Full observations are retrieved via get_mental_model(model_id)
# Filter by tags if provided
if tags:
if tags_match == "all":
# All tags must match
rows = await conn.fetch(
"""
SELECT id, subtype, name, description
FROM mental_models
WHERE bank_id = $1 AND tags @> $2::varchar[]
ORDER BY last_updated DESC NULLS LAST, created_at DESC
""",
bank_id,
tags,
)
else:
# Any tag matches (OR) - default
rows = await conn.fetch(
"""
SELECT id, subtype, name, description
FROM mental_models
WHERE bank_id = $1 AND tags && $2::varchar[]
ORDER BY last_updated DESC NULLS LAST, created_at DESC
""",
bank_id,
tags,
)
if tags:
if tags_match == "all":
filters += f" AND tags @> ${next_param}::varchar[]"
else:
rows = await conn.fetch(
"""
SELECT id, subtype, name, description
FROM mental_models
WHERE bank_id = $1
ORDER BY last_updated DESC NULLS LAST, created_at DESC
filters += f" AND (tags && ${next_param}::varchar[] OR tags IS NULL OR tags = '{{}}')"
params.append(tags)
next_param += 1
if exclude_ids:
filters += f" AND id != ALL(${next_param}::text[])"
params.append(exclude_ids)
next_param += 1
# Search mental models by embedding similarity
rows = await conn.fetch(
f"""
SELECT
id, name, content,
tags, created_at, last_refreshed_at,
1 - (embedding <=> $2::vector) as relevance
FROM {fq_table("mental_models")}
WHERE bank_id = $1 AND embedding IS NOT NULL {filters}
ORDER BY embedding <=> $2::vector
LIMIT $3
""",
*params,
)
now = datetime.now(timezone.utc)
mental_models = []
for row in rows:
last_refreshed_at = row["last_refreshed_at"]
if last_refreshed_at and last_refreshed_at.tzinfo is None:
last_refreshed_at = last_refreshed_at.replace(tzinfo=timezone.utc)
# Calculate freshness
is_stale = False
if last_refreshed_at:
age = now - last_refreshed_at
is_stale = age > timedelta(days=STALE_THRESHOLD_DAYS)
mental_models.append(
{
"id": str(row["id"]),
"name": row["name"],
"content": row["content"],
"tags": row["tags"] or [],
"relevance": round(row["relevance"], 4),
"updated_at": last_refreshed_at.isoformat() if last_refreshed_at else None,
"is_stale": is_stale,
}
)
return {
"query": query,
"count": len(mental_models),
"mental_models": mental_models,
}
async def tool_search_observations(
memory_engine: "MemoryEngine",
bank_id: str,
query: str,
request_context: "RequestContext",
max_tokens: int = 5000,
tags: list[str] | None = None,
tags_match: str = "any",
last_consolidated_at: datetime | None = None,
pending_consolidation: int = 0,
) -> dict[str, Any]:
"""
Search consolidated observations using recall with include_observations.
Observations are auto-generated from memories. Returns freshness info
so the agent knows if it should also verify with recall().
Args:
memory_engine: Memory engine instance
bank_id: Bank identifier
query: Search query
request_context: Request context for authentication
max_tokens: Maximum tokens for results (default 5000)
tags: Optional tags to filter observations
tags_match: How to match tags - "any" (OR), "all" (AND)
last_consolidated_at: When consolidation last ran (for staleness check)
pending_consolidation: Number of memories waiting to be consolidated
Returns:
Dict with matching observations including freshness info
"""
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
enable_trace=False,
request_context=request_context,
tags=tags,
tags_match=tags_match,
_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[])
""",
bank_id,
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 {
"count": len(rows),
"models": [
{
"id": row["id"],
"subtype": row["subtype"],
"name": row["name"],
"description": row["description"],
}
for row in rows
],
}
# Return freshness info (more understandable than raw pending_consolidation count)
if pending_consolidation == 0:
freshness = "up_to_date"
elif pending_consolidation < 10:
freshness = "slightly_stale"
else:
freshness = "stale"
return {
"query": query,
"count": len(observations),
"observations": observations,
"freshness": freshness,
}
async def tool_recall(
@@ -160,6 +245,9 @@ async def tool_recall(
"""
Search memories using TEMPR retrieval.
This is the ground truth - raw facts and experiences.
Use when mental models/observations don't exist, are stale, or need verification.
Args:
memory_engine: Memory engine instance
bank_id: Bank identifier
@@ -177,13 +265,14 @@ async def tool_recall(
result = await memory_engine.recall_async(
bank_id=bank_id,
query=query,
fact_type=["experience", "world"], # Exclude opinions
fact_type=["experience", "world"], # Exclude opinions and observations
max_tokens=max_tokens,
enable_trace=False,
request_context=request_context,
tags=tags,
tags_match=tags_match,
_connection_budget=connection_budget,
_quiet=True, # Suppress logging for internal operations
)
memories = []
@@ -205,85 +294,6 @@ async def tool_recall(
}
async def tool_learn(
conn: "Connection",
bank_id: str,
input: MentalModelInput,
tags: list[str] | None = None,
) -> dict[str, Any]:
"""
Create a mental model placeholder with subtype='learned'.
The agent only specifies name and description - actual observations are generated
in the background via refresh, similar to pinned models.
Args:
conn: Database connection
bank_id: Bank identifier
input: Mental model input data (name, description, optional entity_id)
tags: Tags to apply to new mental models (from reflect context)
Returns:
Dict with created model info including model_id for background generation
"""
model_id = generate_model_id(input.name)
# Parse entity_id if provided
entity_uuid = None
if input.entity_id:
try:
entity_uuid = uuid.UUID(input.entity_id)
except ValueError:
logger.warning(f"Invalid entity_id format: {input.entity_id}")
# Check if model exists
existing = await conn.fetchrow(
"SELECT id FROM mental_models WHERE id = $1 AND bank_id = $2",
model_id,
bank_id,
)
if existing:
# Update description only - observations will be regenerated
await conn.execute(
"""
UPDATE mental_models SET
description = $3,
entity_id = $4
WHERE id = $1 AND bank_id = $2
""",
model_id,
bank_id,
input.description,
entity_uuid,
)
status = "updated"
else:
# Insert new model placeholder - observations will be generated in background
await conn.execute(
"""
INSERT INTO mental_models (id, bank_id, subtype, name, description, observations, entity_id, tags, created_at)
VALUES ($1, $2, 'learned', $3, $4, '{}'::jsonb, $5, $6, NOW())
""",
model_id,
bank_id,
input.name,
input.description,
entity_uuid,
tags or [],
)
status = "created"
logger.info(f"[REFLECT] Mental model '{model_id}' {status} in bank {bank_id} - pending background generation")
return {
"status": status,
"model_id": model_id,
"name": input.name,
"pending_generation": True,
}
async def tool_expand(
conn: "Connection",
bank_id: str,
@@ -302,6 +312,8 @@ async def tool_expand(
Returns:
Dict with results array, each containing memory, chunk, and optionally document data
"""
from ..memory_engine import fq_table
if not memory_ids:
return {"error": "memory_ids is required and must not be empty"}
@@ -319,9 +331,9 @@ async def tool_expand(
# Batch fetch all memory units
memories = await conn.fetch(
"""
f"""
SELECT id, text, chunk_id, document_id, fact_type, context
FROM memory_units
FROM {fq_table("memory_units")}
WHERE id = ANY($1) AND bank_id = $2
""",
valid_uuids,
@@ -338,9 +350,9 @@ async def tool_expand(
chunk_map: dict[str, Any] = {}
if chunk_ids:
chunks = await conn.fetch(
"""
f"""
SELECT chunk_id, chunk_text, chunk_index, document_id
FROM chunks
FROM {fq_table("chunks")}
WHERE chunk_id = ANY($1)
""",
chunk_ids,
@@ -360,9 +372,9 @@ async def tool_expand(
all_doc_ids = list(doc_ids_from_chunks | doc_ids_direct)
if all_doc_ids:
docs = await conn.fetch(
"""
f"""
SELECT id, original_text, metadata, retain_params
FROM documents
FROM {fq_table("documents")}
WHERE id = ANY($1) AND bank_id = $2
""",
all_doc_ids,
@@ -2,38 +2,70 @@
Tool schema definitions for the reflect agent.
These are OpenAI-format tool definitions used with native tool calling.
The reflect agent uses a hierarchical retrieval strategy:
1. search_mental_models - User-curated stored reflect responses (highest quality, if applicable)
2. search_observations - Consolidated knowledge with freshness awareness
3. recall - Raw facts (world/experience) as ground truth fallback
"""
from typing import Literal
# Tool definitions in OpenAI format
TOOL_LIST_MENTAL_MODELS = {
TOOL_SEARCH_MENTAL_MODELS = {
"type": "function",
"function": {
"name": "list_mental_models",
"description": "List all available mental models - your synthesized knowledge about entities, concepts, and events. Returns an array of models with id, name, and description.",
"name": "search_mental_models",
"description": (
"Search user-curated mental models (stored reflect responses). These are high-quality, manually created "
"summaries about specific topics. Use FIRST when the question might be covered by an "
"existing mental model. Returns mental models with their content and last refresh time."
),
"parameters": {
"type": "object",
"properties": {},
"required": [],
"properties": {
"reason": {
"type": "string",
"description": "Brief explanation of why you're making this search (for debugging)",
},
"query": {
"type": "string",
"description": "Search query to find relevant mental models",
},
"max_results": {
"type": "integer",
"description": "Maximum number of mental models to return (default 5)",
},
},
"required": ["reason", "query"],
},
},
}
TOOL_GET_MENTAL_MODEL = {
TOOL_SEARCH_OBSERVATIONS = {
"type": "function",
"function": {
"name": "get_mental_model",
"description": "Get full details of a specific mental model including all observations and memory references.",
"name": "search_observations",
"description": (
"Search consolidated observations (auto-generated knowledge). These are automatically "
"synthesized from memories. Returns observations with freshness info (updated_at, is_stale). "
"If an observation is STALE, you should ALSO use recall() to verify with current facts."
),
"parameters": {
"type": "object",
"properties": {
"model_id": {
"reason": {
"type": "string",
"description": "ID of the mental model (from list_mental_models results)",
"description": "Brief explanation of why you're making this search (for debugging)",
},
"query": {
"type": "string",
"description": "Search query to find relevant observations",
},
"max_tokens": {
"type": "integer",
"description": "Maximum tokens for results (default 5000). Use higher values for broader searches.",
},
},
"required": ["model_id"],
"required": ["reason", "query"],
},
},
}
@@ -42,10 +74,19 @@ TOOL_RECALL = {
"type": "function",
"function": {
"name": "recall",
"description": "Search memories using semantic + temporal retrieval. Returns relevant memories from experience and world knowledge, each with an 'id' you can reference.",
"description": (
"Search raw memories (facts and experiences). This is the ground truth data. "
"Use when: (1) no reflections/mental models exist, (2) mental models are stale, "
"(3) you need specific details not in synthesized knowledge. "
"Returns individual memory facts with their timestamps."
),
"parameters": {
"type": "object",
"properties": {
"reason": {
"type": "string",
"description": "Brief explanation of why you're making this search (for debugging)",
},
"query": {
"type": "string",
"description": "Search query string",
@@ -55,29 +96,7 @@ TOOL_RECALL = {
"description": "Optional limit on result size (default 2048). Use higher values for broader searches.",
},
},
"required": ["query"],
},
},
}
TOOL_LEARN = {
"type": "function",
"function": {
"name": "learn",
"description": "Create a new mental model to track an important recurring topic. Use when you discover a person, project, concept, or pattern that appears frequently and would benefit from synthesized knowledge. The model content will be generated automatically.",
"parameters": {
"type": "object",
"properties": {
"name": {
"type": "string",
"description": "Human-readable name (e.g., 'Project Alpha', 'John Smith', 'Product Strategy')",
},
"description": {
"type": "string",
"description": "What to track and synthesize (e.g., 'Track goals, milestones, blockers, and key decisions for Project Alpha')",
},
},
"required": ["name", "description"],
"required": ["reason", "query"],
},
},
}
@@ -90,6 +109,10 @@ TOOL_EXPAND = {
"parameters": {
"type": "object",
"properties": {
"reason": {
"type": "string",
"description": "Brief explanation of why you need more context (for debugging)",
},
"memory_ids": {
"type": "array",
"items": {"type": "string"},
@@ -101,7 +124,7 @@ TOOL_EXPAND = {
"description": "chunk: surrounding text chunk, document: full source document",
},
},
"required": ["memory_ids", "depth"],
"required": ["reason", "memory_ids", "depth"],
},
},
}
@@ -123,89 +146,104 @@ TOOL_DONE_ANSWER = {
"items": {"type": "string"},
"description": "Array of memory IDs that support your answer (put IDs here, NOT in answer text)",
},
"model_ids": {
"mental_model_ids": {
"type": "array",
"items": {"type": "string"},
"description": "Array of mental model IDs that support your answer",
},
"observation_ids": {
"type": "array",
"items": {"type": "string"},
"description": "Array of observation IDs that support your answer",
},
},
"required": ["answer"],
},
},
}
TOOL_DONE_OBSERVATIONS = {
"type": "function",
"function": {
"name": "done",
"description": "Signal completion with MULTIPLE structured observations. Each observation must be a SEPARATE item in the array covering ONE theme. Do NOT combine all content into a single observation.",
"parameters": {
"type": "object",
"properties": {
"observations": {
"type": "array",
"minItems": 3,
"items": {
"type": "object",
"properties": {
"title": {
"type": "string",
"description": "Short header for this observation's theme (e.g., 'Work Style', 'Technical Skills')",
},
"text": {
"type": "string",
"description": "Observation content about ONE theme. End with 'Key evidence:' containing text citations (summaries of what memories say), NOT memory IDs.",
},
"memory_ids": {
"type": "array",
"items": {"type": "string"},
"description": "Full UUIDs of memories supporting this observation (put IDs here, not in text)",
},
},
"required": ["title", "text", "memory_ids"],
def _build_done_tool_with_directives(directive_rules: list[str]) -> dict:
"""
Build the done tool schema with directive compliance field.
When directives are present, adds a required field that forces the agent
to confirm compliance with each directive before submitting.
Args:
directive_rules: List of directive rule strings
"""
# Build rules list for description
rules_list = "\n".join(f" {i + 1}. {rule}" for i, rule in enumerate(directive_rules))
# Build the tool with directive compliance field
return {
"type": "function",
"function": {
"name": "done",
"description": (
"Signal completion with your final answer. IMPORTANT: You must confirm directive compliance before submitting. "
"Your answer will be REJECTED if it violates any directive."
),
"parameters": {
"type": "object",
"properties": {
"answer": {
"type": "string",
"description": "Your response as plain text. Do NOT use markdown formatting. NEVER include memory IDs, UUIDs, or 'Memory references' in this text - put IDs only in memory_ids array.",
},
"memory_ids": {
"type": "array",
"items": {"type": "string"},
"description": "Array of memory IDs that support your answer (put IDs here, NOT in answer text)",
},
"mental_model_ids": {
"type": "array",
"items": {"type": "string"},
"description": "Array of mental model IDs that support your answer",
},
"observation_ids": {
"type": "array",
"items": {"type": "string"},
"description": "Array of observation IDs that support your answer",
},
"directive_compliance": {
"type": "string",
"description": f"REQUIRED: Confirm your answer complies with ALL directives. List each directive and how your answer follows it:\n{rules_list}\n\nFormat: 'Directive 1: [how answer complies]. Directive 2: [how answer complies]...'",
},
"description": "Array of 3-8 observations, each covering a DIFFERENT aspect/theme. Do NOT put everything in one observation.",
},
"required": ["answer", "directive_compliance"],
},
"required": ["observations"],
},
},
}
}
def get_reflect_tools(
enable_learn: bool = True, output_mode: Literal["answer", "observations"] = "answer"
) -> list[dict]:
def get_reflect_tools(directive_rules: list[str] | None = None) -> list[dict]:
"""
Get the list of tools for the reflect agent.
The tools support a hierarchical retrieval strategy:
1. search_mental_models - User-curated stored reflect responses (try first)
2. search_observations - Consolidated knowledge with freshness
3. recall - Raw facts as ground truth
Args:
enable_learn: Whether to include the learn tool
output_mode: "answer" or "observations" - determines done tool format
In observations mode, mental model tools are excluded to avoid
using potentially outdated models during regeneration.
directive_rules: Optional list of directive rule strings. If provided,
the done() tool will require directive compliance confirmation.
Returns:
List of tool definitions in OpenAI format
"""
tools = []
tools = [
TOOL_SEARCH_MENTAL_MODELS,
TOOL_SEARCH_OBSERVATIONS,
TOOL_RECALL,
TOOL_EXPAND,
]
# In answer mode, include mental model tools for lookup
# In observations mode (mental model generation), exclude them to avoid circular references
if output_mode == "answer":
tools.append(TOOL_LIST_MENTAL_MODELS)
tools.append(TOOL_GET_MENTAL_MODEL)
tools.append(TOOL_RECALL)
if enable_learn:
tools.append(TOOL_LEARN)
tools.append(TOOL_EXPAND)
# Add appropriate done tool based on output mode
if output_mode == "observations":
tools.append(TOOL_DONE_OBSERVATIONS)
# Use directive-aware done tool if directives are present
if directive_rules:
tools.append(_build_done_tool_with_directives(directive_rules))
else:
tools.append(TOOL_DONE_ANSWER)
@@ -10,8 +10,8 @@ from typing import Any
from pydantic import BaseModel, ConfigDict, Field
# Valid fact types for recall operations (excludes 'observation' which is internal, and 'opinion' which is deprecated)
VALID_RECALL_FACT_TYPES = frozenset(["world", "experience"])
# Valid fact types for recall operations (excludes 'opinion' which is deprecated)
VALID_RECALL_FACT_TYPES = frozenset(["world", "experience", "observation"])
class LLMToolCall(BaseModel):
@@ -28,12 +28,15 @@ class LLMToolCallResult(BaseModel):
content: str | None = Field(default=None, description="Text content if any")
tool_calls: list[LLMToolCall] = Field(default_factory=list, description="Tool calls requested by the LLM")
finish_reason: str | None = Field(default=None, description="Reason the LLM stopped: 'stop', 'tool_calls', etc.")
input_tokens: int = Field(default=0, description="Input tokens used in this call")
output_tokens: int = Field(default=0, description="Output tokens used in this call")
class ToolCallTrace(BaseModel):
"""A single tool call made during reflect."""
tool: str = Field(description="Tool name: lookup, recall, learn, expand")
reason: str | None = Field(default=None, description="Agent's reasoning for making this tool call")
input: dict = Field(description="Tool input parameters")
output: dict = Field(description="Tool output/result")
duration_ms: int = Field(description="Execution time in milliseconds")
@@ -47,17 +50,25 @@ class LLMCallTrace(BaseModel):
duration_ms: int = Field(description="Execution time in milliseconds")
class MentalModelRef(BaseModel):
"""Reference to a mental model accessed during reflect."""
class ObservationRef(BaseModel):
"""Reference to an observation accessed during reflect."""
id: str = Field(description="Mental model ID")
name: str = Field(description="Mental model name")
type: str = Field(description="Mental model type: entity, concept, event")
subtype: str = Field(description="Mental model subtype: structural, emergent, learned")
id: str = Field(description="Observation ID")
name: str = Field(description="Observation name")
type: str = Field(description="Observation type: entity, concept, event")
subtype: str = Field(description="Observation subtype: structural, emergent, learned")
description: str = Field(description="Brief description")
summary: str | None = Field(default=None, description="Full summary (when looked up in detail)")
class DirectiveRef(BaseModel):
"""Reference to a directive that was applied during reflect."""
id: str = Field(description="Directive mental model ID")
name: str = Field(description="Directive name")
content: str = Field(description="Directive content")
class TokenUsage(BaseModel):
"""
Token usage metrics for LLM calls.
@@ -158,6 +169,28 @@ class ChunkInfo(BaseModel):
truncated: bool = Field(default=False, description="Whether the chunk was truncated due to token limits")
class ObservationResult(BaseModel):
"""An observation result from recall (consolidated knowledge synthesized from facts)."""
id: str = Field(description="Unique observation ID")
text: str = Field(description="The observation text")
proof_count: int = Field(description="Number of facts supporting this observation")
relevance: float = Field(default=0.0, description="Relevance score to the query")
tags: list[str] | None = Field(default=None, description="Tags for visibility scoping")
source_memory_ids: list[str] = Field(
default_factory=list, description="IDs of facts that contribute to this observation"
)
class MentalModelResult(BaseModel):
"""A mental model result from recall (stored reflect response)."""
id: str = Field(description="Unique mental model ID")
name: str = Field(description="Human-readable name")
content: str = Field(description="The synthesized content")
relevance: float = Field(default=0.0, description="Relevance score to the query")
class RecallResult(BaseModel):
"""
Result from a recall operation.
@@ -221,6 +254,14 @@ class ReflectResult(BaseModel):
],
"experience": [],
"opinion": [],
"mental_models": [],
"directives": [
{
"id": "directive-123",
"name": "Response Style",
"rules": ["Always be concise"],
}
],
},
"new_opinions": ["Machine learning has great potential in healthcare"],
"structured_output": {"summary": "ML in healthcare", "confidence": 0.9},
@@ -230,8 +271,8 @@ class ReflectResult(BaseModel):
)
text: str = Field(description="The formulated answer text")
based_on: dict[str, list[MemoryFact]] = Field(
description="Facts used to formulate the answer, organized by type (world, experience, opinion)"
based_on: dict[str, Any] = Field(
description="Facts used to formulate the answer, organized by type (world, experience, opinion, mental_models, directives)"
)
new_opinions: list[str] = Field(default_factory=list, description="List of newly formed opinions during reflection")
structured_output: dict[str, Any] | None = Field(
@@ -250,9 +291,9 @@ class ReflectResult(BaseModel):
default_factory=list,
description="Trace of LLM calls made during reflection. Only present when include.tool_calls is enabled.",
)
mental_models: list[MentalModelRef] = Field(
directives_applied: list[DirectiveRef] = Field(
default_factory=list,
description="Mental models accessed during reflection. Only present when include.facts is enabled.",
description="Directive mental models that were applied during this reflection.",
)
@@ -114,11 +114,8 @@ class CausalRelation(BaseModel):
"""Causal relationship from this fact to a previous fact (stored format)."""
target_fact_index: int = Field(description="Index of the related fact in the facts array (0-based).")
relation_type: Literal["caused_by", "enabled_by", "prevented_by"] = Field(
description="How this fact relates to the target: "
"'caused_by' = this fact was caused by the target, "
"'enabled_by' = this fact was enabled by the target, "
"'prevented_by' = this fact was prevented by the target"
relation_type: Literal["caused_by"] = Field(
description="How this fact relates to the target: 'caused_by' = this fact was caused by the target"
)
strength: float = Field(
description="Strength of relationship (0.0 to 1.0)",
@@ -141,11 +138,8 @@ class FactCausalRelation(BaseModel):
"MUST be less than this fact's position in the list. "
"Example: if this is fact #5, target_index can only be 0, 1, 2, 3, or 4."
)
relation_type: Literal["caused_by", "enabled_by", "prevented_by"] = Field(
description="How this fact relates to the target fact: "
"'caused_by' = this fact was caused by the target fact, "
"'enabled_by' = this fact was enabled by the target fact, "
"'prevented_by' = this fact was blocked/prevented by the target fact"
relation_type: Literal["caused_by"] = Field(
description="How this fact relates to the target fact: 'caused_by' = this fact was caused by the target fact"
)
strength: float = Field(
description="Strength of relationship (0.0 to 1.0). 1.0 = strong, 0.5 = moderate",
@@ -438,34 +432,15 @@ def _chunk_conversation(turns: list[dict], max_chars: int) -> list[str]:
# FACT EXTRACTION PROMPTS
# =============================================================================
# Concise extraction prompt (default) - selective, high-quality facts
CONCISE_FACT_EXTRACTION_PROMPT = """Extract SIGNIFICANT facts from text. Be SELECTIVE - only extract facts worth remembering long-term.
# Base prompt template (shared by concise and custom modes)
# Uses {extraction_guidelines} placeholder for mode-specific instructions
_BASE_FACT_EXTRACTION_PROMPT = """Extract SIGNIFICANT facts from text. Be SELECTIVE - only extract facts worth remembering long-term.
LANGUAGE RULE (CRITICAL): Output facts in the EXACT SAME language as the input text. If input is Japanese, output Japanese. If input is Chinese, output Chinese. NEVER translate to English. Preserve original language completely.
LANGUAGE REQUIREMENT: Detect the language of the input text. All extracted facts, entity names, descriptions, and other output MUST be in the SAME language as the input. Do not translate to another language.
{fact_types_instruction}
══════════════════════════════════════════════════════════════════════════
SELECTIVITY - CRITICAL (Reduces 90% of unnecessary output)
══════════════════════════════════════════════════════════════════════════
ONLY extract facts that are:
✅ Personal info: names, relationships, roles, background
✅ Preferences: likes, dislikes, habits, interests (e.g., "Alice likes coffee")
✅ Significant events: milestones, decisions, achievements, changes
✅ Plans/goals: future intentions, deadlines, commitments
✅ Expertise: skills, knowledge, certifications, experience
✅ Important context: projects, problems, constraints
✅ Sensory/emotional details: feelings, sensations, perceptions that provide context
✅ Observations: descriptions of people, places, things with specific details
DO NOT extract:
❌ Generic greetings: "how are you", "hello", pleasantries without substance
❌ Pure filler: "thanks", "sounds good", "ok", "got it", "sure"
❌ Process chatter: "let me check", "one moment", "I'll look into it"
❌ Repeated info: if already stated, don't extract again
CONSOLIDATE related statements into ONE fact when possible.
{extraction_guidelines}
══════════════════════════════════════════════════════════════════════════
FACT FORMAT - BE CONCISE
@@ -513,7 +488,33 @@ ENTITIES
══════════════════════════════════════════════════════════════════════════
Include: people names, organizations, places, key objects, abstract concepts (career, friendship, etc.)
Always include "user" when fact is about the user.
Always include "user" when fact is about the user.{examples}"""
# Concise mode guidelines
_CONCISE_GUIDELINES = """══════════════════════════════════════════════════════════════════════════
SELECTIVITY - CRITICAL (Reduces 90% of unnecessary output)
══════════════════════════════════════════════════════════════════════════
ONLY extract facts that are:
✅ Personal info: names, relationships, roles, background
✅ Preferences: likes, dislikes, habits, interests (e.g., "Alice likes coffee")
✅ Significant events: milestones, decisions, achievements, changes
✅ Plans/goals: future intentions, deadlines, commitments
✅ Expertise: skills, knowledge, certifications, experience
✅ Important context: projects, problems, constraints
✅ Sensory/emotional details: feelings, sensations, perceptions that provide context
✅ Observations: descriptions of people, places, things with specific details
DO NOT extract:
❌ Generic greetings: "how are you", "hello", pleasantries without substance
❌ Pure filler: "thanks", "sounds good", "ok", "got it", "sure"
❌ Process chatter: "let me check", "one moment", "I'll look into it"
❌ Repeated info: if already stated, don't extract again
CONSOLIDATE related statements into ONE fact when possible."""
# Concise mode examples
_CONCISE_EXAMPLES = """
══════════════════════════════════════════════════════════════════════════
EXAMPLES
@@ -539,6 +540,20 @@ QUALITY OVER QUANTITY
Ask: "Would this be useful to recall in 6 months?" If no, skip it."""
# Assembled concise prompt (backward compatible - exact same output as before)
CONCISE_FACT_EXTRACTION_PROMPT = _BASE_FACT_EXTRACTION_PROMPT.format(
fact_types_instruction="{fact_types_instruction}",
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}",
extraction_guidelines="{custom_instructions}",
examples="", # No examples for custom mode
)
# Verbose extraction prompt - detailed, comprehensive facts (legacy mode)
VERBOSE_FACT_EXTRACTION_PROMPT = """Extract facts from text into structured format with FIVE required dimensions - BE EXTREMELY DETAILED.
@@ -662,7 +677,7 @@ CAUSAL RELATIONSHIPS
══════════════════════════════════════════════════════════════════════════
Link facts with causal_relations (max 2 per fact). target_index must be < this fact's index.
Types: "caused_by", "enabled_by", "prevented_by"
Type: "caused_by" (this fact was caused by the target fact)
Example: "Lost job → couldn't pay rent → moved apartment"
- Fact 0: Lost job, causal_relations: null
@@ -686,6 +701,12 @@ async def _extract_facts_from_chunk(
Note: event_date parameter is kept for backward compatibility but not used in prompt.
The LLM extracts temporal information from the context string instead.
"""
import logging
from openai import BadRequestError
logger = logging.getLogger(__name__)
memory_bank_context = f"\n- Your name: {agent_name}" if agent_name and extract_opinions else ""
# Determine which fact types to extract based on the flag
@@ -704,13 +725,27 @@ async def _extract_facts_from_chunk(
extract_causal_links = config.retain_extract_causal_links
# Select base prompt based on extraction mode
if extraction_mode == "verbose":
if extraction_mode == "custom":
# Custom mode: inject user-provided guidelines
if not config.retain_custom_instructions:
logger.warning(
"extraction_mode='custom' but HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS not set. "
"Falling back to 'concise' mode."
)
base_prompt = CONCISE_FACT_EXTRACTION_PROMPT
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
else:
base_prompt = CUSTOM_FACT_EXTRACTION_PROMPT
prompt = base_prompt.format(
fact_types_instruction=fact_types_instruction,
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)
else:
base_prompt = CONCISE_FACT_EXTRACTION_PROMPT
# Format the prompt with fact types instruction
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
# Build the full prompt with or without causal relationships section
# Select appropriate response schema based on extraction mode and causal links
@@ -723,12 +758,6 @@ async def _extract_facts_from_chunk(
else:
response_schema = FactExtractionResponseNoCausal
import logging
from openai import BadRequestError
logger = logging.getLogger(__name__)
# Retry logic for JSON validation errors
max_retries = 2
last_error = None
@@ -823,7 +852,8 @@ Text:
# Critical field: fact_type
# LLM uses "assistant" but we convert to "experience" for storage
fact_type = llm_fact.get("fact_type")
original_fact_type = llm_fact.get("fact_type")
fact_type = original_fact_type
# Convert "assistant" → "experience" for storage
if fact_type == "assistant":
@@ -840,7 +870,10 @@ Text:
else:
# Default to 'world' if we can't determine
fact_type = "world"
logger.warning(f"Fact {i}: defaulting to fact_type='world'")
logger.warning(
f"Fact {i}: defaulting to fact_type='world' "
f"(original fact_type={original_fact_type!r}, fact_kind={fact_kind!r})"
)
# Get fact_kind for temporal handling (but don't store it)
fact_kind = llm_fact.get("fact_kind", "conversation")
@@ -41,7 +41,6 @@ async def insert_facts_batch(
contexts = []
fact_types = []
confidence_scores = []
access_counts = []
metadata_jsons = []
chunk_ids = []
document_ids = []
@@ -61,7 +60,6 @@ async def insert_facts_batch(
fact_types.append(fact.fact_type)
# confidence_score is only for opinion facts
confidence_scores.append(1.0 if fact.fact_type == "opinion" else None)
access_counts.append(0) # Initial access count
metadata_jsons.append(json.dumps(fact.metadata))
chunk_ids.append(fact.chunk_id)
# Use per-fact document_id if available, otherwise fallback to batch-level document_id
@@ -76,16 +74,16 @@ async def insert_facts_batch(
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::int[], $12::jsonb[], $13::text[], $14::text[], $15::jsonb[]
$8::text[], $9::text[], $10::float[], $11::jsonb[], $12::text[], $13::text[], $14::jsonb[]
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags_json)
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags_json)
)
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags)
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags)
SELECT
$1,
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id,
context, fact_type, confidence_score, metadata, chunk_id, document_id,
COALESCE(
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
'{{}}'::varchar[]
@@ -103,7 +101,6 @@ async def insert_facts_batch(
contexts,
fact_types,
confidence_scores,
access_counts,
metadata_jsons,
chunk_ids,
document_ids,
@@ -754,17 +754,14 @@ async def create_causal_links_batch(
causal_relations_per_fact: List of causal relations for each fact.
Each element is a list of dicts with:
- target_fact_index: Index into unit_ids for the target fact
- relation_type: "causes", "caused_by", "enables", or "prevents"
- relation_type: "caused_by"
- strength: Float in [0.0, 1.0] representing relationship strength
Returns:
Number of causal links created
Causal link types:
- "causes": This fact directly causes the target fact (forward causation)
- "caused_by": This fact was caused by the target fact (backward causation)
- "enables": This fact enables/allows the target fact (enablement)
- "prevents": This fact prevents/blocks the target fact (prevention)
Causal link type:
- "caused_by": This fact was caused by the target fact
"""
if not unit_ids or not causal_relations_per_fact:
return 0
@@ -787,8 +784,8 @@ async def create_causal_links_batch(
relation_type = relation["relation_type"]
strength = relation.get("strength", 1.0)
# Validate relation_type - must match database constraint
valid_types = {"causes", "caused_by", "enables", "prevents"}
# Validate relation_type - only "caused_by" is supported (DB constraint)
valid_types = {"caused_by"}
if relation_type not in valid_types:
logger.error(
f"Invalid relation_type '{relation_type}' (type: {type(relation_type).__name__}) "
@@ -86,10 +86,10 @@ class CausalRelation:
"""
Causal relationship between facts.
Represents how one fact causes, enables, or prevents another.
Represents how one fact was caused by another.
"""
relation_type: str # "causes", "enables", "prevents", "caused_by"
relation_type: str # "caused_by"
target_fact_index: int # Index of the target fact in the batch
strength: float = 1.0 # Strength of the causal relationship
@@ -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, access_count, embedding, fact_type, document_id, chunk_id, tags,
mentioned_at, embedding, 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.access_count, mu.embedding, mu.fact_type,
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
FROM {fq_table("memory_links")} ml
@@ -45,7 +45,7 @@ async def _find_semantic_seeds(
rows = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
@@ -155,7 +155,6 @@ class LinkExpansionRetriever(GraphRetriever):
all_seeds.extend(temporal_seeds)
if not all_seeds:
logger.debug("[LinkExpansion] No seeds found, returning empty results")
return [], timings
seed_ids = list({s.id for s in all_seeds})
@@ -164,36 +163,108 @@ class LinkExpansionRetriever(GraphRetriever):
# Run entity and causal expansion sequentially on same connection
query_start = time.time()
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.access_count, 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,
)
# 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")
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.access_count, mu.embedding,
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
@@ -211,11 +282,69 @@ class LinkExpansionRetriever(GraphRetriever):
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,
)
timings.edge_load_time = time.time() - query_start
timings.db_queries = 2
timings.edge_count = len(entity_rows) + len(causal_rows)
timings.db_queries = 3
timings.edge_count = len(entity_rows) + len(causal_rows) + len(fallback_rows)
# Merge results, taking max score per fact
# Priority: entity links (unit_entities) > causal links > fallback links
score_map: dict[str, float] = {}
row_map: dict[str, dict] = {}
@@ -230,6 +359,12 @@ class LinkExpansionRetriever(GraphRetriever):
if fact_id not in row_map:
row_map[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)
# 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]
@@ -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, access_count, embedding, fact_type, document_id, chunk_id, tags
mentioned_at, embedding, fact_type, document_id, chunk_id, tags
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
AND fact_type = $2
@@ -116,7 +116,7 @@ async def retrieve_semantic(
results = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
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
FROM {fq_table("memory_units")}
WHERE bank_id = $2
@@ -180,7 +180,7 @@ async def retrieve_bm25(
results = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
ts_rank_cd(search_vector, to_tsquery('english', $1)) AS bm25_score
FROM {fq_table("memory_units")}
WHERE bank_id = $2
@@ -237,7 +237,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, access_count, embedding, fact_type, document_id, chunk_id, tags,
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,
NULL::float AS bm25_score,
'semantic' AS source,
@@ -249,7 +249,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, access_count, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
similarity, bm25_score, source
FROM semantic_ranked
WHERE rn <= $4
@@ -281,7 +281,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, access_count, embedding, fact_type, document_id, chunk_id, tags,
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,
NULL::float AS bm25_score,
'semantic' AS source,
@@ -294,7 +294,7 @@ async def retrieve_semantic_bm25_combined(
{tags_clause}
),
bm25_ranked AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
NULL::float AS similarity,
ts_rank_cd(search_vector, to_tsquery('english', $5)) AS bm25_score,
'bm25' AS source,
@@ -306,12 +306,12 @@ async def retrieve_semantic_bm25_combined(
{tags_clause}
),
semantic AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, 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, access_count, embedding, fact_type, document_id, chunk_id, tags,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
similarity, bm25_score, source
FROM bm25_ranked WHERE rn <= $4
)
@@ -386,7 +386,7 @@ async def retrieve_temporal_combined(
entry_points = await conn.fetch(
f"""
WITH ranked_entries AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
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
FROM {fq_table("memory_units")}
@@ -406,7 +406,7 @@ async def retrieve_temporal_combined(
AND (1 - (embedding <=> $1::vector)) >= $6
{tags_clause}
)
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags, similarity
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
""",
@@ -486,7 +486,7 @@ async def retrieve_temporal_combined(
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.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
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,
1 - (mu.embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_links")} ml
@@ -610,7 +610,7 @@ async def retrieve_temporal(
entry_points = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
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
FROM {fq_table("memory_units")}
WHERE bank_id = $2
@@ -691,7 +691,7 @@ async def retrieve_temporal(
# Batch fetch all neighbors for this batch of nodes
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.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id,
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,
ml.weight, ml.link_type, ml.from_unit_id,
1 - (mu.embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_links")} ml
@@ -1023,7 +1023,7 @@ async def _get_temporal_entry_points(
rows = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at,
access_count, embedding, fact_type, document_id, chunk_id,
embedding, fact_type, document_id, chunk_id,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
@@ -1,159 +0,0 @@
"""
Scoring functions for memory search and retrieval.
Includes recency weighting, frequency weighting, temporal proximity,
and similarity calculations used in memory activation and ranking.
"""
from datetime import datetime
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
"""
Calculate cosine similarity between two vectors.
Args:
vec1: First vector
vec2: Second vector
Returns:
Similarity score between 0 and 1
"""
if len(vec1) != len(vec2):
raise ValueError("Vectors must have same dimension")
dot_product = sum(a * b for a, b in zip(vec1, vec2))
magnitude1 = sum(a * a for a in vec1) ** 0.5
magnitude2 = sum(b * b for b in vec2) ** 0.5
if magnitude1 == 0 or magnitude2 == 0:
return 0.0
return dot_product / (magnitude1 * magnitude2)
def calculate_recency_weight(days_since: float, half_life_days: float = 365.0) -> float:
"""
Calculate recency weight using logarithmic decay.
This provides much better differentiation over long time periods compared to
exponential decay. Uses a log-based decay where the half-life parameter controls
when memories reach 50% weight.
Examples:
- Today (0 days): 1.0
- 1 year (365 days): ~0.5 (with default half_life=365)
- 2 years (730 days): ~0.33
- 5 years (1825 days): ~0.17
- 10 years (3650 days): ~0.09
This ensures that 2-year-old and 5-year-old memories have meaningfully
different weights, unlike exponential decay which makes them both ~0.
Args:
days_since: Number of days since the memory was created
half_life_days: Number of days for weight to reach 0.5 (default: 1 year)
Returns:
Weight between 0 and 1
"""
import math
# Logarithmic decay: 1 / (1 + log(1 + days_since/half_life))
# This decays much slower than exponential, giving better long-term differentiation
normalized_age = days_since / half_life_days
return 1.0 / (1.0 + math.log1p(normalized_age))
def calculate_frequency_weight(access_count: int, max_boost: float = 2.0) -> float:
"""
Calculate frequency weight based on access count.
Frequently accessed memories are weighted higher.
Uses logarithmic scaling to avoid over-weighting.
Args:
access_count: Number of times the memory was accessed
max_boost: Maximum multiplier for frequently accessed memories
Returns:
Weight between 1.0 and max_boost
"""
import math
if access_count <= 0:
return 1.0
# Logarithmic scaling: log(access_count + 1) / log(10)
# This gives: 0 accesses = 1.0, 9 accesses ~= 1.5, 99 accesses ~= 2.0
normalized = math.log(access_count + 1) / math.log(10)
return 1.0 + min(normalized, max_boost - 1.0)
def calculate_temporal_anchor(occurred_start: datetime, occurred_end: datetime) -> datetime:
"""
Calculate a single temporal anchor point from a temporal range.
Used for spreading activation - we need a single representative date
to calculate temporal proximity between facts. This simplifies the
range-to-range distance problem.
Strategy: Use midpoint of the range for balanced representation.
Args:
occurred_start: Start of temporal range
occurred_end: End of temporal range
Returns:
Single datetime representing the temporal anchor (midpoint)
Examples:
- Point event (July 14): start=July 14, end=July 14 anchor=July 14
- Month range (February): start=Feb 1, end=Feb 28 anchor=Feb 14
- Year range (2023): start=Jan 1, end=Dec 31 anchor=July 1
"""
# Calculate midpoint
time_delta = occurred_end - occurred_start
midpoint = occurred_start + (time_delta / 2)
return midpoint
def calculate_temporal_proximity(anchor_a: datetime, anchor_b: datetime, half_life_days: float = 30.0) -> float:
"""
Calculate temporal proximity between two temporal anchors.
Used for spreading activation to determine how "close" two facts are
in time. Uses logarithmic decay so that temporal similarity doesn't
drop off too quickly.
Args:
anchor_a: Temporal anchor of first fact
anchor_b: Temporal anchor of second fact
half_life_days: Number of days for proximity to reach 0.5
(default: 30 days = 1 month)
Returns:
Proximity score in [0, 1] where:
- 1.0 = same day
- 0.5 = ~half_life days apart
- 0.0 = very distant in time
Examples:
- Same day: 1.0
- 1 week apart (half_life=30): ~0.7
- 1 month apart (half_life=30): ~0.5
- 1 year apart (half_life=30): ~0.2
"""
import math
days_apart = abs((anchor_a - anchor_b).days)
if days_apart == 0:
return 1.0
# Logarithmic decay: 1 / (1 + log(1 + days_apart/half_life))
# Similar to calculate_recency_weight but for proximity between events
normalized_distance = days_apart / half_life_days
proximity = 1.0 / (1.0 + math.log1p(normalized_distance))
return proximity
@@ -85,7 +85,6 @@ class NodeVisit(BaseModel):
text: str = Field(description="Memory unit text content")
context: str = Field(description="Memory unit context")
event_date: datetime | None = Field(default=None, description="When the memory occurred")
access_count: int = Field(description="Number of times accessed before this search")
# How this node was reached
is_entry_point: bool = Field(description="Whether this is an entry point")
@@ -136,7 +136,6 @@ class SearchTracer:
text: str,
context: str,
event_date: datetime | None,
access_count: int,
is_entry_point: bool,
parent_node_id: str | None,
link_type: Literal["temporal", "semantic", "entity"] | None,
@@ -155,7 +154,6 @@ class SearchTracer:
text: Memory unit text
context: Memory unit context
event_date: When the memory occurred
access_count: Access count before this search
is_entry_point: Whether this is an entry point
parent_node_id: Node that led here (None for entry points)
link_type: Type of link from parent
@@ -194,7 +192,6 @@ class SearchTracer:
text=text,
context=context,
event_date=event_date,
access_count=access_count,
is_entry_point=is_entry_point,
parent_node_id=parent_node_id,
link_type=link_type,
@@ -333,8 +330,8 @@ class SearchTracer:
RetrievalResult(
rank=rank,
node_id=doc_id,
text=data.get("text", ""),
context=data.get("context", ""),
text=data.get("text") or "",
context=data.get("context") or "",
event_date=data.get("event_date"),
fact_type=data.get("fact_type") or fact_type,
score=score,
@@ -46,7 +46,6 @@ class RetrievalResult:
mentioned_at: datetime | None = None
document_id: str | None = None
chunk_id: str | None = None
access_count: int = 0
embedding: list[float] | None = None
tags: list[str] | None = None # Visibility scope tags
@@ -71,7 +70,6 @@ class RetrievalResult:
mentioned_at=row.get("mentioned_at"),
document_id=row.get("document_id"),
chunk_id=row.get("chunk_id"),
access_count=row.get("access_count", 0),
embedding=row.get("embedding"),
tags=row.get("tags"),
similarity=row.get("similarity"),
@@ -156,7 +154,6 @@ class ScoredResult:
"mentioned_at": self.retrieval.mentioned_at,
"document_id": self.retrieval.document_id,
"chunk_id": self.retrieval.chunk_id,
"access_count": self.retrieval.access_count,
"embedding": self.retrieval.embedding,
"tags": self.retrieval.tags,
"semantic_similarity": self.retrieval.similarity,
+112 -196
View File
@@ -1,31 +1,40 @@
"""
Abstract task backend for running async tasks.
Task backend for distributed task processing.
This provides an abstraction that can be adapted to different execution models:
- AsyncIO queue (default implementation)
- Pub/Sub architectures (future)
- Message brokers (future)
This provides an abstraction for task storage and execution:
- BrokerTaskBackend: Uses PostgreSQL as broker (production)
- SyncTaskBackend: Executes tasks immediately (testing/embedded)
"""
import asyncio
import json
import logging
from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable
from typing import Any
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
import asyncpg
logger = logging.getLogger(__name__)
def fq_table(table: str, schema: str | None = None) -> str:
"""Get fully-qualified table name with optional schema prefix."""
if schema:
return f'"{schema}".{table}'
return table
class TaskBackend(ABC):
"""
Abstract base class for task execution backends.
Implementations must:
1. Store/publish task events (as serializable dicts)
2. Execute tasks through a provided executor callback
2. Execute tasks through a provided executor callback (optional)
The backend treats tasks as pure dictionaries that can be serialized
and sent over the network. The executor (typically MemoryEngine.execute_task)
and stored in the database. The executor (typically MemoryEngine.execute_task)
receives the dict and routes it to the appropriate handler.
"""
@@ -46,7 +55,7 @@ class TaskBackend(ABC):
@abstractmethod
async def initialize(self):
"""
Initialize the backend (e.g., start workers, connect to broker).
Initialize the backend (e.g., connect to database).
"""
pass
@@ -63,7 +72,7 @@ class TaskBackend(ABC):
@abstractmethod
async def shutdown(self):
"""
Shutdown the backend gracefully (e.g., stop workers, close connections).
Shutdown the backend gracefully.
"""
pass
@@ -93,9 +102,8 @@ class SyncTaskBackend(TaskBackend):
"""
Synchronous task backend that executes tasks immediately.
This is useful for embedded/CLI usage where we don't want background
workers that prevent clean exit. Tasks are executed inline rather than
being queued.
This is useful for tests and embedded/CLI usage where we don't want
background workers. Tasks are executed inline rather than being queued.
"""
async def initialize(self):
@@ -121,221 +129,129 @@ class SyncTaskBackend(TaskBackend):
logger.debug("SyncTaskBackend shutdown")
class NoopTaskBackend(TaskBackend):
class BrokerTaskBackend(TaskBackend):
"""
No-op task backend that discards all tasks.
Task backend using PostgreSQL as broker.
This is useful for tests where background task execution is not needed
and would only slow down the test suite.
submit_task() stores task_payload in async_operations table.
Actual polling and execution is handled separately by WorkerPoller.
This backend is used by the API to store tasks. Workers poll
the database separately to claim and execute tasks.
"""
async def initialize(self):
"""No-op."""
self._initialized = True
logger.debug("NoopTaskBackend initialized")
async def submit_task(self, task_dict: dict[str, Any]):
"""Discard the task (do nothing)."""
pass
async def shutdown(self):
"""No-op."""
self._initialized = False
logger.debug("NoopTaskBackend shutdown")
class AsyncIOQueueBackend(TaskBackend):
"""
Task backend implementation using asyncio queues.
This is the default implementation that uses in-process asyncio queues
and a periodic consumer worker.
"""
def __init__(self, batch_size: int = 10, batch_interval: float = 1.0):
def __init__(
self,
pool_getter: Callable[[], "asyncpg.Pool"],
schema: str | None = None,
schema_getter: Callable[[], str | None] | None = None,
):
"""
Initialize AsyncIO queue backend.
Initialize the broker task backend.
Args:
batch_size: Maximum number of tasks to process in one batch
batch_interval: Maximum time (seconds) to wait before processing batch
pool_getter: Callable that returns the asyncpg connection pool
schema: Database schema for multi-tenant support (optional, static)
schema_getter: Callable that returns current schema dynamically (optional).
If set, takes precedence over static schema for submit_task.
"""
super().__init__()
self._queue: asyncio.Queue | None = None
self._worker_task: asyncio.Task | None = None
self._shutdown_event: asyncio.Event | None = None
self._batch_size = batch_size
self._batch_interval = batch_interval
self._in_flight_count = 0
self._in_flight_lock = asyncio.Lock()
self._pool_getter = pool_getter
self._schema = schema
self._schema_getter = schema_getter
async def initialize(self):
"""Initialize the queue and start the worker."""
if self._initialized:
return
self._queue = asyncio.Queue()
self._shutdown_event = asyncio.Event()
self._worker_task = asyncio.create_task(self._worker())
"""Initialize the backend."""
self._initialized = True
logger.info("AsyncIOQueueBackend initialized")
logger.info("BrokerTaskBackend initialized")
async def submit_task(self, task_dict: dict[str, Any]):
"""
Submit a task by putting it in the queue.
Store task payload in async_operations table.
The task_dict should contain an 'operation_id' if updating an existing
operation record, otherwise a new operation will be created.
Args:
task_dict: Task dictionary to execute
task_dict: Task dictionary to store (must be JSON serializable)
"""
if not self._initialized:
await self.initialize()
await self._queue.put(task_dict)
pool = self._pool_getter()
operation_id = task_dict.get("operation_id")
task_type = task_dict.get("type", "unknown")
bank_id = task_dict.get("bank_id")
payload_json = json.dumps(task_dict)
schema = self._schema_getter() if self._schema_getter else self._schema
table = fq_table("async_operations", schema)
if operation_id:
# Update existing operation with task payload
await pool.execute(
f"""
UPDATE {table}
SET task_payload = $1::jsonb, updated_at = now()
WHERE operation_id = $2
""",
payload_json,
operation_id,
)
logger.debug(f"Updated task payload for operation {operation_id}")
else:
# Insert new operation (for tasks without pre-created records)
# e.g., access_count_update tasks
import uuid
new_id = uuid.uuid4()
await pool.execute(
f"""
INSERT INTO {table} (operation_id, bank_id, operation_type, status, task_payload)
VALUES ($1, $2, $3, 'pending', $4::jsonb)
""",
new_id,
bank_id,
task_type,
payload_json,
)
logger.debug(f"Created new operation {new_id} for task type {task_type}")
async def shutdown(self):
"""Shutdown the backend."""
self._initialized = False
logger.info("BrokerTaskBackend shutdown")
async def wait_for_pending_tasks(self, timeout: float = 120.0):
"""
Wait for all pending tasks in the queue and in-flight tasks to complete.
Wait for pending tasks to be processed.
This is useful in tests to ensure background tasks complete before assertions.
In the broker model, this polls the database to check if tasks
for this process have been completed. This is useful in tests
when worker_enabled=True (API processes its own tasks).
Args:
timeout: Maximum time to wait in seconds (default 120s for long-running tasks)
timeout: Maximum time to wait in seconds
"""
if not self._initialized or self._queue is None:
return
import asyncio
pool = self._pool_getter()
schema = self._schema_getter() if self._schema_getter else self._schema
table = fq_table("async_operations", schema)
# Wait for queue to be empty AND no in-flight tasks
start_time = asyncio.get_event_loop().time()
while asyncio.get_event_loop().time() - start_time < timeout:
async with self._in_flight_lock:
in_flight = self._in_flight_count
# Check if there are any pending tasks with payloads
count = await pool.fetchval(
f"""
SELECT COUNT(*) FROM {table}
WHERE status = 'pending' AND task_payload IS NOT NULL
"""
)
if self._queue.empty() and in_flight == 0:
# Queue is empty and no tasks in flight, we're done
if count == 0:
return
# Wait a bit before checking again
await asyncio.sleep(0.5)
async def shutdown(self):
"""Shutdown the worker and drain the queue."""
if not self._initialized:
return
logger.info("Shutting down AsyncIOQueueBackend...")
# Signal shutdown
self._shutdown_event.set()
# Cancel worker
if self._worker_task is not None:
self._worker_task.cancel()
try:
await self._worker_task
except asyncio.CancelledError:
pass # Worker cancelled successfully
self._initialized = False
logger.info("AsyncIOQueueBackend shutdown complete")
async def _execute_task_with_tracking(self, task_dict: dict[str, Any]):
"""Execute a task and track its in-flight status."""
async with self._in_flight_lock:
self._in_flight_count += 1
try:
await self._execute_task(task_dict)
finally:
async with self._in_flight_lock:
self._in_flight_count -= 1
async def _execute_task_no_tracking(self, task_dict: dict[str, Any]):
"""Execute a task without in-flight tracking (tracking done at batch level)."""
await self._execute_task(task_dict)
def _get_queue_stats(self) -> tuple[int, dict[str, int]]:
"""Get current queue size and bank_id distribution."""
queue_size = self._queue.qsize() if self._queue else 0
bank_distribution: dict[str, int] = {}
if queue_size > 0 and self._queue:
# Peek at queue items without removing them
# Note: This is a snapshot and may not be perfectly accurate due to concurrency
try:
# Access internal deque for logging purposes only
items = list(self._queue._queue) # type: ignore[attr-defined]
for item in items:
bank_id = item.get("bank_id", "unknown")
bank_distribution[bank_id] = bank_distribution.get(bank_id, 0) + 1
except Exception:
pass # Queue access failed, return empty distribution
return queue_size, bank_distribution
async def _worker(self):
"""
Background worker that processes tasks in batches.
Collects tasks for up to batch_interval seconds or batch_size items,
then processes them.
"""
while not self._shutdown_event.is_set():
try:
# Collect tasks for batching
tasks = []
deadline = asyncio.get_event_loop().time() + self._batch_interval
while len(tasks) < self._batch_size and asyncio.get_event_loop().time() < deadline:
try:
remaining_time = max(0.1, deadline - asyncio.get_event_loop().time())
task_dict = await asyncio.wait_for(self._queue.get(), timeout=remaining_time)
# Track task as in-flight immediately when picked up from queue
# This prevents wait_for_pending_tasks from returning too early
async with self._in_flight_lock:
self._in_flight_count += 1
tasks.append(task_dict)
except TimeoutError:
break
# Process batch
if tasks:
# Log batch start with queue stats
queue_size, bank_distribution = self._get_queue_stats()
# Summarize batch by task type and bank
batch_summary: dict[str, dict[str, int]] = {}
for task_dict in tasks:
task_type = task_dict.get("type", "unknown")
bank_id = task_dict.get("bank_id", "unknown")
if task_type not in batch_summary:
batch_summary[task_type] = {}
batch_summary[task_type][bank_id] = batch_summary[task_type].get(bank_id, 0) + 1
# Build log message
batch_parts = []
for task_type, banks in sorted(batch_summary.items()):
bank_str = ", ".join(f"{b}:{c}" for b, c in sorted(banks.items()))
batch_parts.append(f"{task_type}[{bank_str}]")
batch_str = ", ".join(batch_parts)
if queue_size > 0:
pending_str = ", ".join(f"{k}:{v}" for k, v in sorted(bank_distribution.items()))
logger.info(
f"Processing {len(tasks)} tasks: {batch_str} (pending={queue_size} [{pending_str}])"
)
else:
logger.info(f"Processing {len(tasks)} tasks: {batch_str}")
# Execute tasks concurrently (in_flight already tracked when picked up)
await asyncio.gather(
*[self._execute_task_no_tracking(task_dict) for task_dict in tasks], return_exceptions=True
)
# Decrement in_flight count after all tasks complete
async with self._in_flight_lock:
self._in_flight_count -= len(tasks)
except asyncio.CancelledError:
break
except Exception as e:
logger.error(f"Worker error: {e}")
await asyncio.sleep(1) # Backoff on error
logger.warning(f"Timeout waiting for pending tasks after {timeout}s")
-151
View File
@@ -65,154 +65,3 @@ async def extract_facts(
return [], chunks
return facts, chunks
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
"""
Calculate cosine similarity between two vectors.
Args:
vec1: First vector
vec2: Second vector
Returns:
Similarity score between 0 and 1
"""
if len(vec1) != len(vec2):
raise ValueError("Vectors must have same dimension")
dot_product = sum(a * b for a, b in zip(vec1, vec2))
magnitude1 = sum(a * a for a in vec1) ** 0.5
magnitude2 = sum(b * b for b in vec2) ** 0.5
if magnitude1 == 0 or magnitude2 == 0:
return 0.0
return dot_product / (magnitude1 * magnitude2)
def calculate_recency_weight(days_since: float, half_life_days: float = 365.0) -> float:
"""
Calculate recency weight using logarithmic decay.
This provides much better differentiation over long time periods compared to
exponential decay. Uses a log-based decay where the half-life parameter controls
when memories reach 50% weight.
Examples:
- Today (0 days): 1.0
- 1 year (365 days): ~0.5 (with default half_life=365)
- 2 years (730 days): ~0.33
- 5 years (1825 days): ~0.17
- 10 years (3650 days): ~0.09
This ensures that 2-year-old and 5-year-old memories have meaningfully
different weights, unlike exponential decay which makes them both ~0.
Args:
days_since: Number of days since the memory was created
half_life_days: Number of days for weight to reach 0.5 (default: 1 year)
Returns:
Weight between 0 and 1
"""
import math
# Logarithmic decay: 1 / (1 + log(1 + days_since/half_life))
# This decays much slower than exponential, giving better long-term differentiation
normalized_age = days_since / half_life_days
return 1.0 / (1.0 + math.log1p(normalized_age))
def calculate_frequency_weight(access_count: int, max_boost: float = 2.0) -> float:
"""
Calculate frequency weight based on access count.
Frequently accessed memories are weighted higher.
Uses logarithmic scaling to avoid over-weighting.
Args:
access_count: Number of times the memory was accessed
max_boost: Maximum multiplier for frequently accessed memories
Returns:
Weight between 1.0 and max_boost
"""
import math
if access_count <= 0:
return 1.0
# Logarithmic scaling: log(access_count + 1) / log(10)
# This gives: 0 accesses = 1.0, 9 accesses ~= 1.5, 99 accesses ~= 2.0
normalized = math.log(access_count + 1) / math.log(10)
return 1.0 + min(normalized, max_boost - 1.0)
def calculate_temporal_anchor(occurred_start: datetime, occurred_end: datetime) -> datetime:
"""
Calculate a single temporal anchor point from a temporal range.
Used for spreading activation - we need a single representative date
to calculate temporal proximity between facts. This simplifies the
range-to-range distance problem.
Strategy: Use midpoint of the range for balanced representation.
Args:
occurred_start: Start of temporal range
occurred_end: End of temporal range
Returns:
Single datetime representing the temporal anchor (midpoint)
Examples:
- Point event (July 14): start=July 14, end=July 14 anchor=July 14
- Month range (February): start=Feb 1, end=Feb 28 anchor=Feb 14
- Year range (2023): start=Jan 1, end=Dec 31 anchor=July 1
"""
# Calculate midpoint
time_delta = occurred_end - occurred_start
midpoint = occurred_start + (time_delta / 2)
return midpoint
def calculate_temporal_proximity(anchor_a: datetime, anchor_b: datetime, half_life_days: float = 30.0) -> float:
"""
Calculate temporal proximity between two temporal anchors.
Used for spreading activation to determine how "close" two facts are
in time. Uses logarithmic decay so that temporal similarity doesn't
drop off too quickly.
Args:
anchor_a: Temporal anchor of first fact
anchor_b: Temporal anchor of second fact
half_life_days: Number of days for proximity to reach 0.5
(default: 30 days = 1 month)
Returns:
Proximity score in [0, 1] where:
- 1.0 = same day
- 0.5 = ~half_life days apart
- 0.0 = very distant in time
Examples:
- Same day: 1.0
- 1 week apart (half_life=30): ~0.7
- 1 month apart (half_life=30): ~0.5
- 1 year apart (half_life=30): ~0.2
"""
import math
days_apart = abs((anchor_a - anchor_b).days)
if days_apart == 0:
return 1.0
# Logarithmic decay: 1 / (1 + log(1 + days_apart/half_life))
# Similar to calculate_recency_weight but for proximity between events
normalized_distance = days_apart / half_life_days
proximity = 1.0 / (1.0 + math.log1p(normalized_distance))
return proximity
@@ -21,6 +21,10 @@ from hindsight_api.extensions.context import DefaultExtensionContext, ExtensionC
from hindsight_api.extensions.http import HttpExtension
from hindsight_api.extensions.loader import load_extension
from hindsight_api.extensions.operation_validator import (
# Consolidation operation
ConsolidateContext,
ConsolidateResult,
# Core operations
OperationValidationError,
OperationValidatorExtension,
RecallContext,
@@ -33,6 +37,7 @@ from hindsight_api.extensions.operation_validator import (
)
from hindsight_api.extensions.tenant import (
AuthenticationError,
Tenant,
TenantContext,
TenantExtension,
)
@@ -47,7 +52,7 @@ __all__ = [
"DefaultExtensionContext",
# HTTP Extension
"HttpExtension",
# Operation Validator
# Operation Validator - Core
"OperationValidationError",
"OperationValidatorExtension",
"RecallContext",
@@ -57,10 +62,14 @@ __all__ = [
"RetainContext",
"RetainResult",
"ValidationResult",
# Operation Validator - Consolidation
"ConsolidateContext",
"ConsolidateResult",
# Tenant/Auth
"ApiKeyTenantExtension",
"AuthenticationError",
"RequestContext",
"Tenant",
"TenantContext",
"TenantExtension",
]
@@ -1,6 +1,7 @@
"""Built-in tenant extension implementations."""
from hindsight_api.extensions.tenant import AuthenticationError, TenantContext, TenantExtension
from hindsight_api.config import get_config
from hindsight_api.extensions.tenant import AuthenticationError, Tenant, TenantContext, TenantExtension
from hindsight_api.models import RequestContext
@@ -10,11 +11,13 @@ class ApiKeyTenantExtension(TenantExtension):
This is a simple implementation that:
1. Validates the API key matches HINDSIGHT_API_TENANT_API_KEY
2. Returns 'public' as the schema for all authenticated requests
2. Returns the configured schema (HINDSIGHT_API_DATABASE_SCHEMA, default 'public')
for all authenticated requests
Configuration:
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.tenant:ApiKeyTenantExtension
HINDSIGHT_API_TENANT_API_KEY=your-secret-key
HINDSIGHT_API_DATABASE_SCHEMA=your-schema (optional, defaults to 'public')
For multi-tenant setups with separate schemas per tenant, implement a custom
TenantExtension that looks up the schema based on the API key or token claims.
@@ -27,7 +30,11 @@ class ApiKeyTenantExtension(TenantExtension):
raise ValueError("HINDSIGHT_API_TENANT_API_KEY is required when using ApiKeyTenantExtension")
async def authenticate(self, context: RequestContext) -> TenantContext:
"""Validate API key and return public schema context."""
"""Validate API key and return configured schema context."""
if context.api_key != self.expected_api_key:
raise AuthenticationError("Invalid API key")
return TenantContext(schema_name="public")
return TenantContext(schema_name=get_config().database_schema)
async def list_tenants(self) -> list[Tenant]:
"""Return configured schema for single-tenant setup."""
return [Tenant(schema=get_config().database_schema)]
@@ -1,4 +1,4 @@
"""Operation Validator Extension for validating retain/recall/reflect operations."""
"""Operation Validator Extension for validating retain/recall/reflect/consolidate operations."""
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
@@ -97,6 +97,19 @@ class ReflectContext:
context: str | None = None
# =============================================================================
# Consolidation Pre-operation Context
# =============================================================================
@dataclass
class ConsolidateContext:
"""Context for a consolidation operation validation (pre-operation)."""
bank_id: str
request_context: "RequestContext"
# =============================================================================
# Post-operation Contexts (includes results)
# =============================================================================
@@ -164,9 +177,28 @@ class ReflectResultContext:
error: str | None = None
# =============================================================================
# Consolidation Post-operation Context
# =============================================================================
@dataclass
class ConsolidateResult:
"""Result context for post-consolidation hook."""
bank_id: str
request_context: "RequestContext"
# Result
processed: int = 0
created: int = 0
updated: int = 0
success: bool = True
error: str | None = None
class OperationValidatorExtension(Extension, ABC):
"""
Validates and hooks into retain/recall/reflect operations.
Validates and hooks into retain/recall/reflect/consolidate operations.
This extension allows implementing custom logic such as:
- Rate limiting (pre-operation)
@@ -185,9 +217,13 @@ class OperationValidatorExtension(Extension, ABC):
-> config = {"max_requests": "100"}
Hook execution order:
1. validate_retain/validate_recall/validate_reflect (pre-operation)
1. validate_* (pre-operation)
2. [operation executes]
3. on_retain_complete/on_recall_complete/on_reflect_complete (post-operation)
3. on_*_complete (post-operation)
Supported operations:
- retain, recall, reflect (core memory operations)
- consolidate (mental models consolidation)
"""
# =========================================================================
@@ -325,3 +361,44 @@ class OperationValidatorExtension(Extension, ABC):
- error: Error message (if failed)
"""
pass
# =========================================================================
# Consolidation - Pre-operation validation hook (optional - override to implement)
# =========================================================================
async def validate_consolidate(self, ctx: ConsolidateContext) -> ValidationResult:
"""
Validate a consolidation operation before execution.
Override to implement custom validation logic for consolidation.
Args:
ctx: Context containing:
- bank_id: Bank identifier
- request_context: Request context with auth info
Returns:
ValidationResult indicating whether the operation is allowed.
"""
return ValidationResult.accept()
# =========================================================================
# Consolidation - Post-operation hook (optional - override to implement)
# =========================================================================
async def on_consolidate_complete(self, result: ConsolidateResult) -> None:
"""
Called after a consolidation operation completes (success or failure).
Override to implement post-operation logic such as usage tracking or audit logging.
Args:
result: Result context containing:
- bank_id: Bank identifier
- processed: Number of memories processed
- created: Number of mental models created
- updated: Number of mental models updated
- success: Whether the operation succeeded
- error: Error message (if failed)
"""
pass
@@ -28,6 +28,18 @@ class TenantContext:
schema_name: str
@dataclass
class Tenant:
"""
Represents a tenant for worker discovery.
Used by list_tenants() to return tenant information including
the PostgreSQL schema name for database operations.
"""
schema: str
class TenantExtension(Extension, ABC):
"""
Extension for multi-tenancy and API key authentication.
@@ -61,3 +73,17 @@ class TenantExtension(Extension, ABC):
AuthenticationError: If authentication fails.
"""
...
@abstractmethod
async def list_tenants(self) -> list[Tenant]:
"""
List all tenants that should be processed by workers.
This method is used by the worker to discover all tenants that need
task polling. Workers will poll for pending tasks in each tenant's schema.
Returns:
List of Tenant objects containing schema information.
For single-tenant setups, return [Tenant(schema="public")].
"""
...
+22 -7
View File
@@ -170,6 +170,7 @@ def main():
if args.log_level != config.log_level:
config = HindsightConfig(
database_url=config.database_url,
database_schema=config.database_schema,
llm_provider=config.llm_provider,
llm_api_key=config.llm_api_key,
llm_model=config.llm_model,
@@ -184,13 +185,20 @@ def main():
reflect_llm_api_key=config.reflect_llm_api_key,
reflect_llm_model=config.reflect_llm_model,
reflect_llm_base_url=config.reflect_llm_base_url,
consolidation_llm_provider=config.consolidation_llm_provider,
consolidation_llm_api_key=config.consolidation_llm_api_key,
consolidation_llm_model=config.consolidation_llm_model,
consolidation_llm_base_url=config.consolidation_llm_base_url,
embeddings_provider=config.embeddings_provider,
embeddings_local_model=config.embeddings_local_model,
embeddings_local_force_cpu=config.embeddings_local_force_cpu,
embeddings_tei_url=config.embeddings_tei_url,
embeddings_openai_base_url=config.embeddings_openai_base_url,
embeddings_cohere_base_url=config.embeddings_cohere_base_url,
reranker_provider=config.reranker_provider,
reranker_local_model=config.reranker_local_model,
reranker_local_force_cpu=config.reranker_local_force_cpu,
reranker_local_max_concurrent=config.reranker_local_max_concurrent,
reranker_tei_url=config.reranker_tei_url,
reranker_tei_batch_size=config.reranker_tei_batch_size,
reranker_tei_max_concurrent=config.reranker_tei_max_concurrent,
@@ -199,18 +207,21 @@ def main():
host=args.host,
port=args.port,
log_level=args.log_level,
log_format=config.log_format,
mcp_enabled=config.mcp_enabled,
graph_retriever=config.graph_retriever,
mpfp_top_k_neighbors=config.mpfp_top_k_neighbors,
recall_max_concurrent=config.recall_max_concurrent,
recall_connection_budget=config.recall_connection_budget,
observation_min_facts=config.observation_min_facts,
observation_top_entities=config.observation_top_entities,
retain_max_completion_tokens=config.retain_max_completion_tokens,
retain_chunk_size=config.retain_chunk_size,
retain_extract_causal_links=config.retain_extract_causal_links,
retain_extraction_mode=config.retain_extraction_mode,
retain_custom_instructions=config.retain_custom_instructions,
retain_observations_async=config.retain_observations_async,
enable_observations=config.enable_observations,
consolidation_batch_size=config.consolidation_batch_size,
consolidation_max_tokens=config.consolidation_max_tokens,
skip_llm_verification=config.skip_llm_verification,
lazy_reranker=config.lazy_reranker,
run_migrations_on_startup=config.run_migrations_on_startup,
@@ -218,9 +229,12 @@ def main():
db_pool_max_size=config.db_pool_max_size,
db_command_timeout=config.db_command_timeout,
db_acquire_timeout=config.db_acquire_timeout,
task_backend=config.task_backend,
task_backend_memory_batch_size=config.task_backend_memory_batch_size,
task_backend_memory_batch_interval=config.task_backend_memory_batch_interval,
worker_enabled=config.worker_enabled,
worker_id=config.worker_id,
worker_poll_interval_ms=config.worker_poll_interval_ms,
worker_max_retries=config.worker_max_retries,
worker_batch_size=config.worker_batch_size,
worker_http_port=config.worker_http_port,
reflect_max_iterations=config.reflect_max_iterations,
mental_model_refresh_concurrency=config.mental_model_refresh_concurrency,
)
@@ -332,6 +346,7 @@ def main():
# Start idle checker in daemon mode
if idle_middleware is not None:
# Start the idle checker in a background thread with its own event loop
import logging
import threading
def run_idle_checker():
@@ -342,8 +357,8 @@ def main():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
loop.run_until_complete(idle_middleware._check_idle())
except Exception:
pass
except Exception as e:
logging.error(f"Idle checker error: {e}", exc_info=True)
threading.Thread(target=run_idle_checker, daemon=True).start()
+11 -52
View File
@@ -44,7 +44,6 @@ import os
import sys
from mcp.server.fastmcp import FastMCP
from mcp.types import Icon
from hindsight_api.config import (
DEFAULT_MCP_LOCAL_BANK_ID,
@@ -53,6 +52,7 @@ from hindsight_api.config import (
ENV_MCP_INSTRUCTIONS,
ENV_MCP_LOCAL_BANK_ID,
)
from hindsight_api.mcp_tools import MCPToolsConfig, register_mcp_tools
# Configure logging - default to warning to avoid polluting stderr during MCP init
# MCP clients interpret stderr output as errors, so we suppress INFO logs by default
@@ -85,9 +85,6 @@ def create_local_mcp_server(bank_id: str, memory=None) -> FastMCP:
"""
# Import here to avoid slow startup if just checking --help
from hindsight_api import MemoryEngine
from hindsight_api.engine.memory_engine import Budget
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
from hindsight_api.models import RequestContext
# Create memory engine with pg0 embedded database if not provided
if memory is None:
@@ -105,55 +102,17 @@ def create_local_mcp_server(bank_id: str, memory=None) -> FastMCP:
mcp = FastMCP("hindsight")
@mcp.tool(description=retain_description)
async def retain(content: str, context: str = "general") -> dict:
"""
Args:
content: The fact/memory to store (be specific and include relevant details)
context: Category for the memory (e.g., 'preferences', 'work', 'hobbies', 'family'). Default: 'general'
"""
import asyncio
# Configure and register tools using shared module
config = MCPToolsConfig(
bank_id_resolver=lambda: bank_id,
include_bank_id_param=False, # Local MCP uses fixed bank_id
tools={"retain", "recall"}, # Local MCP only has retain and recall
retain_description=retain_description,
recall_description=recall_description,
retain_fire_and_forget=True, # Local MCP uses fire-and-forget pattern
)
async def _retain():
try:
await memory.retain_batch_async(
bank_id=bank_id,
contents=[{"content": content, "context": context}],
request_context=RequestContext(),
)
except Exception as e:
logger.error(f"Error storing memory: {e}", exc_info=True)
# Fire and forget - don't block on memory storage
asyncio.create_task(_retain())
return {"status": "accepted", "message": "Memory storage initiated"}
@mcp.tool(description=recall_description)
async def recall(query: str, max_tokens: int = 4096, budget: str = "low") -> dict:
"""
Args:
query: Natural language search query (e.g., "user's food preferences", "what projects is user working on")
max_tokens: Maximum tokens to return in results (default: 4096)
budget: Search budget level - "low", "mid", or "high" (default: "low")
"""
try:
# Map string budget to enum
budget_map = {"low": Budget.LOW, "mid": Budget.MID, "high": Budget.HIGH}
budget_enum = budget_map.get(budget.lower(), Budget.LOW)
search_result = await memory.recall_async(
bank_id=bank_id,
query=query,
fact_type=list(VALID_RECALL_FACT_TYPES),
budget=budget_enum,
max_tokens=max_tokens,
request_context=RequestContext(),
)
return search_result.model_dump()
except Exception as e:
logger.error(f"Error searching: {e}", exc_info=True)
return {"error": str(e), "results": []}
register_mcp_tools(mcp, memory, config)
return mcp
+494
View File
@@ -0,0 +1,494 @@
"""Shared MCP tool implementations for Hindsight.
This module provides the core tool logic used by both:
- mcp_local.py (stdio transport for Claude Code)
- api/mcp.py (HTTP transport for API server)
"""
import json
import logging
from dataclasses import dataclass
from datetime import datetime
from typing import Any, Callable
from fastmcp import FastMCP
from hindsight_api import MemoryEngine
from hindsight_api.config import (
DEFAULT_MCP_RECALL_DESCRIPTION,
DEFAULT_MCP_RETAIN_DESCRIPTION,
)
from hindsight_api.engine.memory_engine import Budget
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
from hindsight_api.models import RequestContext
logger = logging.getLogger(__name__)
@dataclass
class MCPToolsConfig:
"""Configuration for MCP tools registration."""
# How to resolve bank_id for operations
bank_id_resolver: Callable[[], str | None]
# Whether to include bank_id as a parameter on tools (for multi-bank support)
include_bank_id_param: bool = False
# Which tools to register
tools: set[str] | None = None # None means all tools
# Custom descriptions (if None, uses defaults)
retain_description: str | None = None
recall_description: str | None = None
# Retain behavior
retain_fire_and_forget: bool = False # If True, use asyncio.create_task pattern
def parse_timestamp(timestamp: str) -> datetime | None:
"""Parse an ISO format timestamp string.
Args:
timestamp: ISO format timestamp (e.g., '2024-01-15T10:30:00Z')
Returns:
Parsed datetime or None if invalid
Raises:
ValueError: If timestamp format is invalid
"""
try:
return datetime.fromisoformat(timestamp.replace("Z", "+00:00"))
except ValueError as e:
raise ValueError(
f"Invalid timestamp format '{timestamp}'. "
"Expected ISO format like '2024-01-15T10:30:00' or '2024-01-15T10:30:00Z'"
) from e
def build_content_dict(
content: str,
context: str,
timestamp: str | None = None,
) -> tuple[dict[str, Any], str | None]:
"""Build a content dict for retain operations.
Args:
content: The memory content
context: Category for the memory
timestamp: Optional ISO timestamp
Returns:
Tuple of (content_dict, error_message). error_message is None if successful.
"""
content_dict: dict[str, Any] = {"content": content, "context": context}
if timestamp:
try:
parsed_timestamp = parse_timestamp(timestamp)
content_dict["event_date"] = parsed_timestamp
except ValueError as e:
return {}, str(e)
return content_dict, None
def register_mcp_tools(
mcp: FastMCP,
memory: MemoryEngine,
config: MCPToolsConfig,
) -> None:
"""Register MCP tools on a FastMCP server.
Args:
mcp: FastMCP server instance
memory: MemoryEngine instance
config: Tool configuration
"""
tools_to_register = config.tools or {"retain", "recall", "reflect", "list_banks", "create_bank"}
if "retain" in tools_to_register:
_register_retain(mcp, memory, config)
if "recall" in tools_to_register:
_register_recall(mcp, memory, config)
if "reflect" in tools_to_register:
_register_reflect(mcp, memory, config)
if "list_banks" in tools_to_register:
_register_list_banks(mcp, memory, config)
if "create_bank" in tools_to_register:
_register_create_bank(mcp, memory, config)
def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
"""Register the retain tool."""
description = config.retain_description or DEFAULT_MCP_RETAIN_DESCRIPTION
if config.include_bank_id_param:
if config.retain_fire_and_forget:
@mcp.tool(description=description)
async def retain(
content: str,
context: str = "general",
timestamp: str | None = None,
bank_id: str | None = None,
) -> dict:
"""
Args:
content: The fact/memory to store (be specific and include relevant details)
context: Category for the memory (e.g., 'preferences', 'work', 'hobbies', 'family'). Default: 'general'
timestamp: When this event/fact occurred (ISO format, e.g., '2024-01-15T10:30:00Z'). Useful for timeline tracking.
bank_id: Optional bank to store in (defaults to session bank). Use for cross-bank operations.
"""
import asyncio
target_bank = bank_id or config.bank_id_resolver()
if target_bank is None:
return {"status": "error", "message": "No bank_id configured"}
content_dict, error = build_content_dict(content, context, timestamp)
if error:
return {"status": "error", "message": error}
async def _retain():
try:
await memory.retain_batch_async(
bank_id=target_bank,
contents=[content_dict],
request_context=RequestContext(),
)
except Exception as e:
logger.error(f"Error storing memory: {e}", exc_info=True)
asyncio.create_task(_retain())
return {"status": "accepted", "message": "Memory storage initiated"}
else:
@mcp.tool(description=description)
async def retain(
content: str,
context: str = "general",
timestamp: str | None = None,
async_processing: bool = True,
bank_id: str | None = None,
) -> str:
"""
Args:
content: The fact/memory to store (be specific and include relevant details)
context: Category for the memory (e.g., 'preferences', 'work', 'hobbies', 'family'). Default: 'general'
timestamp: When this event/fact occurred (ISO format, e.g., '2024-01-15T10:30:00Z'). Useful for timeline tracking.
async_processing: If True, queue for background processing and return immediately. If False, wait for completion. Default: True
bank_id: Optional bank to store in (defaults to session bank). Use for cross-bank operations.
"""
try:
target_bank = bank_id or config.bank_id_resolver()
if target_bank is None:
return "Error: No bank_id configured"
content_dict, error = build_content_dict(content, context, timestamp)
if error:
return f"Error: {error}"
contents = [content_dict]
if async_processing:
result = await memory.submit_async_retain(
bank_id=target_bank, contents=contents, request_context=RequestContext()
)
return f"Memory queued for background processing (operation_id: {result.get('operation_id', 'N/A')})"
else:
await memory.retain_batch_async(
bank_id=target_bank,
contents=contents,
request_context=RequestContext(),
)
return f"Memory stored successfully in bank '{target_bank}'"
except Exception as e:
logger.error(f"Error storing memory: {e}", exc_info=True)
return f"Error: {str(e)}"
else:
# No bank_id param - use fixed bank from resolver
@mcp.tool(description=description)
async def retain(
content: str,
context: str = "general",
timestamp: str | None = None,
) -> dict:
"""
Args:
content: The fact/memory to store (be specific and include relevant details)
context: Category for the memory (e.g., 'preferences', 'work', 'hobbies', 'family'). Default: 'general'
timestamp: When this event/fact occurred (ISO format, e.g., '2024-01-15T10:30:00Z'). Useful for timeline tracking.
"""
import asyncio
target_bank = config.bank_id_resolver()
if target_bank is None:
return {"status": "error", "message": "No bank_id configured"}
content_dict, error = build_content_dict(content, context, timestamp)
if error:
return {"status": "error", "message": error}
async def _retain():
try:
await memory.retain_batch_async(
bank_id=target_bank,
contents=[content_dict],
request_context=RequestContext(),
)
except Exception as e:
logger.error(f"Error storing memory: {e}", exc_info=True)
asyncio.create_task(_retain())
return {"status": "accepted", "message": "Memory storage initiated"}
def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
"""Register the recall tool."""
description = config.recall_description or DEFAULT_MCP_RECALL_DESCRIPTION
if config.include_bank_id_param:
@mcp.tool(description=description)
async def recall(
query: str,
max_tokens: int = 4096,
bank_id: str | None = None,
) -> str | dict:
"""
Args:
query: Natural language search query (e.g., "user's food preferences", "what projects is user working on")
max_tokens: Maximum tokens to return in results (default: 4096)
bank_id: Optional bank to search in (defaults to session bank). Use for cross-bank operations.
"""
try:
target_bank = bank_id or config.bank_id_resolver()
if target_bank is None:
return "Error: No bank_id configured"
recall_result = await memory.recall_async(
bank_id=target_bank,
query=query,
fact_type=list(VALID_RECALL_FACT_TYPES),
budget=Budget.HIGH,
max_tokens=max_tokens,
request_context=RequestContext(),
)
return recall_result.model_dump_json(indent=2)
except Exception as e:
logger.error(f"Error searching: {e}", exc_info=True)
return f'{{"error": "{e}", "results": []}}'
else:
@mcp.tool(description=description)
async def recall(
query: str,
max_tokens: int = 4096,
) -> dict:
"""
Args:
query: Natural language search query (e.g., "user's food preferences", "what projects is user working on")
max_tokens: Maximum tokens to return in results (default: 4096)
"""
try:
target_bank = config.bank_id_resolver()
if target_bank is None:
return {"error": "No bank_id configured", "results": []}
recall_result = await memory.recall_async(
bank_id=target_bank,
query=query,
fact_type=list(VALID_RECALL_FACT_TYPES),
budget=Budget.HIGH,
max_tokens=max_tokens,
request_context=RequestContext(),
)
return recall_result.model_dump()
except Exception as e:
logger.error(f"Error searching: {e}", exc_info=True)
return {"error": str(e), "results": []}
def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
"""Register the reflect tool."""
if config.include_bank_id_param:
@mcp.tool()
async def reflect(
query: str,
context: str | None = None,
budget: str = "low",
bank_id: str | None = None,
) -> str:
"""
Generate thoughtful analysis by synthesizing stored memories with the bank's personality.
WHEN TO USE THIS TOOL:
Use reflect when you need reasoned analysis, not just fact retrieval. This tool
thinks through the question using everything the bank knows and its personality traits.
EXAMPLES OF GOOD QUERIES:
- "What patterns have emerged in how I approach debugging?"
- "Based on my past decisions, what architectural style do I prefer?"
- "What might be the best approach for this problem given what you know about me?"
- "How should I prioritize these tasks based on my goals?"
HOW IT DIFFERS FROM RECALL:
- recall: Returns raw facts matching your search (fast lookup)
- reflect: Reasons across memories to form a synthesized answer (deeper analysis)
Use recall for "what did I say about X?" and reflect for "what should I do about X?"
Args:
query: The question or topic to reflect on
context: Optional context about why this reflection is needed
budget: Search budget - 'low', 'mid', or 'high' (default: 'low')
bank_id: Optional bank to reflect in (defaults to session bank). Use for cross-bank operations.
"""
try:
target_bank = bank_id or config.bank_id_resolver()
if target_bank is None:
return "Error: No bank_id configured"
budget_map = {"low": Budget.LOW, "mid": Budget.MID, "high": Budget.HIGH}
budget_enum = budget_map.get(budget.lower(), Budget.LOW)
reflect_result = await memory.reflect_async(
bank_id=target_bank,
query=query,
budget=budget_enum,
context=context,
request_context=RequestContext(),
)
return reflect_result.model_dump_json(indent=2)
except Exception as e:
logger.error(f"Error reflecting: {e}", exc_info=True)
return f'{{"error": "{e}", "text": ""}}'
else:
@mcp.tool()
async def reflect(
query: str,
context: str | None = None,
budget: str = "low",
) -> dict:
"""
Generate thoughtful analysis by synthesizing stored memories with the bank's personality.
WHEN TO USE THIS TOOL:
Use reflect when you need reasoned analysis, not just fact retrieval. This tool
thinks through the question using everything the bank knows and its personality traits.
EXAMPLES OF GOOD QUERIES:
- "What patterns have emerged in how I approach debugging?"
- "Based on my past decisions, what architectural style do I prefer?"
- "What might be the best approach for this problem given what you know about me?"
- "How should I prioritize these tasks based on my goals?"
HOW IT DIFFERS FROM RECALL:
- recall: Returns raw facts matching your search (fast lookup)
- reflect: Reasons across memories to form a synthesized answer (deeper analysis)
Use recall for "what did I say about X?" and reflect for "what should I do about X?"
Args:
query: The question or topic to reflect on
context: Optional context about why this reflection is needed
budget: Search budget - 'low', 'mid', or 'high' (default: 'low')
"""
try:
target_bank = config.bank_id_resolver()
if target_bank is None:
return {"error": "No bank_id configured", "text": ""}
budget_map = {"low": Budget.LOW, "mid": Budget.MID, "high": Budget.HIGH}
budget_enum = budget_map.get(budget.lower(), Budget.LOW)
reflect_result = await memory.reflect_async(
bank_id=target_bank,
query=query,
budget=budget_enum,
context=context,
request_context=RequestContext(),
)
return reflect_result.model_dump()
except Exception as e:
logger.error(f"Error reflecting: {e}", exc_info=True)
return {"error": str(e), "text": ""}
def _register_list_banks(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
"""Register the list_banks tool."""
@mcp.tool()
async def list_banks() -> str:
"""
List all available memory banks.
Use this tool to discover what memory banks exist in the system.
Each bank is an isolated memory store (like a separate "brain").
Returns:
JSON list of banks with their IDs, names, dispositions, and missions.
"""
try:
banks = await memory.list_banks(request_context=RequestContext())
return json.dumps({"banks": banks}, indent=2)
except Exception as e:
logger.error(f"Error listing banks: {e}", exc_info=True)
return f'{{"error": "{e}", "banks": []}}'
def _register_create_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
"""Register the create_bank tool."""
@mcp.tool()
async def create_bank(bank_id: str, name: str | None = None, mission: str | None = None) -> str:
"""
Create a new memory bank or get an existing one.
Memory banks are isolated stores - each one is like a separate "brain" for a user/agent.
Banks are auto-created with default settings if they don't exist.
Args:
bank_id: Unique identifier for the bank (e.g., 'user-123', 'agent-alpha')
name: Optional human-friendly name for the bank
mission: Optional mission describing who the agent is and what they're trying to accomplish
"""
try:
# get_bank_profile auto-creates bank if it doesn't exist
profile = await memory.get_bank_profile(bank_id, request_context=RequestContext())
# Update name/mission if provided
if name is not None or mission is not None:
await memory.update_bank(
bank_id,
name=name,
mission=mission,
request_context=RequestContext(),
)
# Fetch updated profile
profile = await memory.get_bank_profile(bank_id, request_context=RequestContext())
# Serialize disposition if it's a Pydantic model
if "disposition" in profile and hasattr(profile["disposition"], "model_dump"):
profile["disposition"] = profile["disposition"].model_dump()
return json.dumps(profile, indent=2)
except Exception as e:
logger.error(f"Error creating bank: {e}", exc_info=True)
return f'{{"error": "{e}"}}'
-2
View File
@@ -95,7 +95,6 @@ class MemoryUnit(Base):
mentioned_at: Mapped[datetime | None] = mapped_column(TIMESTAMP(timezone=True)) # When fact was mentioned
fact_type: Mapped[str] = mapped_column(Text, nullable=False, server_default="world")
confidence_score: Mapped[float | None] = mapped_column(Float)
access_count: Mapped[int] = mapped_column(Integer, server_default="0")
unit_metadata: Mapped[dict] = mapped_column(
"metadata", JSONB, server_default=sql_text("'{}'::jsonb")
) # User-defined metadata (str->str)
@@ -131,7 +130,6 @@ class MemoryUnit(Base):
Index("idx_memory_units_document_id", "document_id"),
Index("idx_memory_units_event_date", "event_date", postgresql_ops={"event_date": "DESC"}),
Index("idx_memory_units_bank_date", "bank_id", "event_date", postgresql_ops={"event_date": "DESC"}),
Index("idx_memory_units_access_count", "access_count", postgresql_ops={"access_count": "DESC"}),
Index("idx_memory_units_fact_type", "fact_type"),
Index("idx_memory_units_bank_fact_type", "bank_id", "fact_type"),
Index(
@@ -0,0 +1,11 @@
"""
Worker package for distributed task processing.
This package provides:
- WorkerPoller: Polls PostgreSQL for pending tasks and executes them
- main: CLI entry point for hindsight-worker
"""
from .poller import WorkerPoller
__all__ = ["WorkerPoller"]
+296
View File
@@ -0,0 +1,296 @@
"""
Command-line interface for Hindsight Worker.
Run the worker with:
hindsight-worker
Stop with Ctrl+C (graceful shutdown).
"""
import argparse
import asyncio
import atexit
import logging
import os
import signal
import socket
import sys
import warnings
from ..config import get_config
from ..engine.task_backend import SyncTaskBackend
from .poller import WorkerPoller
# Filter deprecation warnings from third-party libraries
warnings.filterwarnings("ignore", message="websockets.legacy is deprecated")
warnings.filterwarnings("ignore", message="websockets.server.WebSocketServerProtocol is deprecated")
# Disable tokenizers parallelism to avoid warnings
os.environ["TOKENIZERS_PARALLELISM"] = "false"
logger = logging.getLogger(__name__)
def create_worker_app(poller: WorkerPoller, memory):
"""Create a minimal FastAPI app for worker metrics and health."""
from fastapi import FastAPI
from fastapi.responses import JSONResponse, Response
from prometheus_client import CONTENT_TYPE_LATEST, generate_latest
from ..metrics import create_metrics_collector, get_metrics_collector, initialize_metrics
app = FastAPI(
title="Hindsight Worker",
description="Worker process for distributed task execution",
)
# Initialize OpenTelemetry metrics
try:
prometheus_reader = initialize_metrics(service_name="hindsight-worker", service_version="1.0.0")
create_metrics_collector()
app.state.prometheus_reader = prometheus_reader
logger.info("Metrics initialized - available at /metrics endpoint")
except Exception as e:
logger.warning(f"Failed to initialize metrics: {e}. Metrics will be disabled.")
app.state.prometheus_reader = None
# Set up DB pool metrics if available
metrics_collector = get_metrics_collector()
if memory._pool is not None and hasattr(metrics_collector, "set_db_pool"):
metrics_collector.set_db_pool(memory._pool)
logger.info("DB pool metrics configured")
@app.get(
"/health",
summary="Health check endpoint",
description="Returns worker health status including database connectivity",
tags=["Monitoring"],
)
async def health_endpoint():
"""Health check endpoint."""
health = await memory.health_check()
health["worker_id"] = poller.worker_id
health["is_shutdown"] = poller.is_shutdown
status_code = 200 if health.get("status") == "healthy" else 503
return JSONResponse(content=health, status_code=status_code)
@app.get(
"/metrics",
summary="Prometheus metrics endpoint",
description="Exports metrics in Prometheus format for scraping",
tags=["Monitoring"],
)
async def metrics_endpoint():
"""Return Prometheus metrics."""
metrics_data = generate_latest()
return Response(content=metrics_data, media_type=CONTENT_TYPE_LATEST)
@app.get(
"/",
summary="Worker info",
description="Basic worker information",
tags=["Info"],
)
async def root():
"""Return basic worker info."""
return {
"service": "hindsight-worker",
"worker_id": poller.worker_id,
"is_shutdown": poller.is_shutdown,
}
return app
def main():
"""Main entry point for the hindsight-worker CLI."""
# Load configuration from environment
config = get_config()
parser = argparse.ArgumentParser(
prog="hindsight-worker",
description="Hindsight Worker - distributed task processor",
)
# Worker options
parser.add_argument(
"--worker-id",
default=config.worker_id or socket.gethostname(),
help="Worker identifier (default: hostname, env: HINDSIGHT_API_WORKER_ID)",
)
parser.add_argument(
"--poll-interval",
type=int,
default=config.worker_poll_interval_ms,
help=f"Poll interval in milliseconds (default: {config.worker_poll_interval_ms}, env: HINDSIGHT_API_WORKER_POLL_INTERVAL_MS)",
)
parser.add_argument(
"--batch-size",
type=int,
default=config.worker_batch_size,
help=f"Tasks to claim per poll (default: {config.worker_batch_size}, env: HINDSIGHT_API_WORKER_BATCH_SIZE)",
)
parser.add_argument(
"--max-retries",
type=int,
default=config.worker_max_retries,
help=f"Max retries before marking failed (default: {config.worker_max_retries}, env: HINDSIGHT_API_WORKER_MAX_RETRIES)",
)
# HTTP server options
parser.add_argument(
"--http-port",
type=int,
default=config.worker_http_port,
help=f"HTTP port for metrics/health endpoints (default: {config.worker_http_port}, env: HINDSIGHT_API_WORKER_HTTP_PORT)",
)
parser.add_argument(
"--http-host",
default="0.0.0.0",
help="HTTP host to bind (default: 0.0.0.0)",
)
# Logging options
parser.add_argument(
"--log-level",
default=config.log_level,
choices=["critical", "error", "warning", "info", "debug", "trace"],
help=f"Log level (default: {config.log_level}, env: HINDSIGHT_API_LOG_LEVEL)",
)
args = parser.parse_args()
# Configure logging
config.configure_logging()
# Import MemoryEngine here to avoid circular imports
from .. import MemoryEngine
print(f"Starting Hindsight Worker: {args.worker_id}")
print(f" Poll interval: {args.poll_interval}ms")
print(f" Batch size: {args.batch_size}")
print(f" Max retries: {args.max_retries}")
print(f" HTTP server: {args.http_host}:{args.http_port}")
print()
# Global references for cleanup
memory = None
poller = None
async def run():
nonlocal memory, poller
import uvicorn
from ..extensions import TenantExtension, load_extension
# Initialize MemoryEngine
# Workers use SyncTaskBackend because they execute tasks directly,
# they don't need to store tasks (they poll from DB)
memory = MemoryEngine(
run_migrations=False, # Workers don't run migrations
task_backend=SyncTaskBackend(),
)
await memory.initialize()
print(f"Database connected: {config.database_url}")
# Load tenant extension for dynamic schema discovery
tenant_extension = load_extension("TENANT", TenantExtension)
if tenant_extension:
print("Tenant extension loaded - schemas will be discovered dynamically on each poll")
else:
print("No tenant extension configured, using public schema only")
# Create a single poller that handles all schemas dynamically
poller = WorkerPoller(
pool=memory._pool,
worker_id=args.worker_id,
executor=memory.execute_task,
poll_interval_ms=args.poll_interval,
batch_size=args.batch_size,
max_retries=args.max_retries,
tenant_extension=tenant_extension,
)
# Create the HTTP app for metrics/health
app = create_worker_app(poller, memory)
# Setup signal handlers for graceful shutdown
shutdown_requested = asyncio.Event()
def signal_handler(signum, frame):
print(f"\nReceived signal {signum}, initiating graceful shutdown...")
shutdown_requested.set()
signal.signal(signal.SIGINT, signal_handler)
signal.signal(signal.SIGTERM, signal_handler)
# Create uvicorn config and server
uvicorn_config = uvicorn.Config(
app,
host=args.http_host,
port=args.http_port,
log_level="info", # Reduce uvicorn noise
access_log=False,
)
server = uvicorn.Server(uvicorn_config)
# Run the poller and HTTP server concurrently
poller_task = asyncio.create_task(poller.run())
http_task = asyncio.create_task(server.serve())
print(f"Worker started. Metrics available at http://{args.http_host}:{args.http_port}/metrics")
# Wait for shutdown signal
await shutdown_requested.wait()
# Graceful shutdown
print("Shutting down HTTP server...")
server.should_exit = True
print("Waiting for poller to finish...")
await poller.shutdown_graceful(timeout=30.0)
poller_task.cancel()
try:
await poller_task
except asyncio.CancelledError:
pass
# Wait for HTTP server to finish
try:
await asyncio.wait_for(http_task, timeout=5.0)
except asyncio.TimeoutError:
http_task.cancel()
try:
await http_task
except asyncio.CancelledError:
pass
# Close memory engine
await memory.close()
print("Worker shutdown complete")
def cleanup():
"""Synchronous cleanup for atexit."""
if memory is not None and memory._pg0 is not None:
try:
loop = asyncio.new_event_loop()
loop.run_until_complete(memory._pg0.stop())
loop.close()
print("\npg0 stopped.")
except Exception as e:
print(f"\nError stopping pg0: {e}")
atexit.register(cleanup)
try:
asyncio.run(run())
except KeyboardInterrupt:
print("\nWorker interrupted")
sys.exit(0)
if __name__ == "__main__":
main()
@@ -0,0 +1,486 @@
"""
Worker poller for distributed task execution.
Polls PostgreSQL for pending tasks and executes them using
FOR UPDATE SKIP LOCKED for safe concurrent claiming.
"""
import asyncio
import json
import logging
import time
import traceback
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
import asyncpg
from hindsight_api.extensions.tenant import TenantExtension
logger = logging.getLogger(__name__)
# Progress logging interval in seconds
PROGRESS_LOG_INTERVAL = 30
def fq_table(table: str, schema: str | None = None) -> str:
"""Get fully-qualified table name with optional schema prefix."""
if schema:
return f'"{schema}".{table}'
return table
@dataclass
class ClaimedTask:
"""A task claimed from the database with its schema context."""
operation_id: str
task_dict: dict[str, Any]
schema: str | None
class WorkerPoller:
"""
Polls PostgreSQL for pending tasks and executes them.
Uses FOR UPDATE SKIP LOCKED for safe distributed claiming,
allowing multiple workers to process tasks without conflicts.
Supports dynamic multi-tenant discovery via tenant_extension.
"""
def __init__(
self,
pool: "asyncpg.Pool",
worker_id: str,
executor: Callable[[dict[str, Any]], Awaitable[None]],
poll_interval_ms: int = 500,
batch_size: int = 10,
max_retries: int = 3,
schema: str | None = None,
tenant_extension: "TenantExtension | None" = None,
):
"""
Initialize the worker poller.
Args:
pool: asyncpg connection pool
worker_id: Unique identifier for this worker
executor: Async function to execute tasks (typically MemoryEngine.execute_task)
poll_interval_ms: Interval between polls when no tasks found (milliseconds)
batch_size: Maximum number of tasks to claim per poll cycle
max_retries: Maximum retry attempts before marking task as failed
schema: Database schema for single-tenant support (ignored if tenant_extension is set)
tenant_extension: Extension for dynamic multi-tenant discovery. If set, list_tenants()
is called on each poll cycle to discover schemas dynamically.
"""
self._pool = pool
self._worker_id = worker_id
self._executor = executor
self._poll_interval_ms = poll_interval_ms
self._batch_size = batch_size
self._max_retries = max_retries
self._schema = schema
self._tenant_extension = tenant_extension
self._shutdown = asyncio.Event()
self._current_tasks: set[asyncio.Task] = set()
self._in_flight_count = 0
self._in_flight_lock = asyncio.Lock()
self._last_progress_log = 0.0
self._tasks_completed_since_log = 0
# Track active tasks locally: operation_id -> (op_type, bank_id, schema)
self._active_tasks: dict[str, tuple[str, str, str | None]] = {}
async def _get_schemas(self) -> list[str | None]:
"""Get list of schemas to poll. Returns [None] for public schema."""
if self._tenant_extension is not None:
tenants = await self._tenant_extension.list_tenants()
# Convert "public" to None for SQL compatibility, keep others as-is
return [t.schema if t.schema != "public" else None for t in tenants]
# Single schema mode
return [self._schema]
async def claim_batch(self) -> list[ClaimedTask]:
"""
Claim up to batch_size pending tasks atomically across all tenant schemas.
Uses FOR UPDATE SKIP LOCKED to ensure no conflicts with other workers.
For consolidation tasks specifically, skips pending tasks if there's already
a processing consolidation for the same bank (to avoid duplicate work).
If tenant_extension is configured, dynamically discovers schemas on each call.
Returns:
List of ClaimedTask objects containing operation_id, task_dict, and schema
"""
schemas = await self._get_schemas()
all_tasks: list[ClaimedTask] = []
remaining_batch = self._batch_size
for schema in schemas:
if remaining_batch <= 0:
break
tasks = await self._claim_batch_for_schema(schema, remaining_batch)
all_tasks.extend(tasks)
remaining_batch -= len(tasks)
return all_tasks
async def _claim_batch_for_schema(self, schema: str | None, limit: int) -> list[ClaimedTask]:
"""Claim tasks from a specific schema."""
table = fq_table("async_operations", schema)
async with self._pool.acquire() as conn:
async with conn.transaction():
# Select and lock pending tasks
# For consolidation: skip if same bank already has one processing
rows = await conn.fetch(
f"""
SELECT operation_id, task_payload
FROM {table} AS pending
WHERE status = 'pending' AND task_payload IS NOT NULL
AND (
-- Non-consolidation tasks: always claimable
operation_type != 'consolidation'
OR
-- Consolidation: only if no other consolidation processing for same bank
NOT EXISTS (
SELECT 1 FROM {table} AS processing
WHERE processing.bank_id = pending.bank_id
AND processing.operation_type = 'consolidation'
AND processing.status = 'processing'
)
)
ORDER BY created_at
LIMIT $1
FOR UPDATE SKIP LOCKED
""",
limit,
)
if not rows:
return []
# Claim the tasks by updating status and worker_id
operation_ids = [row["operation_id"] for row in rows]
await conn.execute(
f"""
UPDATE {table}
SET status = 'processing', worker_id = $1, claimed_at = now(), updated_at = now()
WHERE operation_id = ANY($2)
""",
self._worker_id,
operation_ids,
)
# Parse and return task payloads with schema context
return [
ClaimedTask(
operation_id=str(row["operation_id"]),
task_dict=json.loads(row["task_payload"]),
schema=schema,
)
for row in rows
]
async def _mark_completed(self, operation_id: str, schema: str | None):
"""Mark a task as completed."""
table = fq_table("async_operations", schema)
await self._pool.execute(
f"""
UPDATE {table}
SET status = 'completed', completed_at = now(), updated_at = now()
WHERE operation_id = $1
""",
operation_id,
)
async def _mark_failed(self, operation_id: str, error_message: str, schema: str | None):
"""Mark a task as failed with error message."""
table = fq_table("async_operations", schema)
# Truncate error message if too long (max 5000 chars in schema)
error_message = error_message[:5000] if len(error_message) > 5000 else error_message
await self._pool.execute(
f"""
UPDATE {table}
SET status = 'failed', error_message = $2, completed_at = now(), updated_at = now()
WHERE operation_id = $1
""",
operation_id,
error_message,
)
async def _retry_or_fail(self, operation_id: str, error_message: str, schema: str | None):
"""Increment retry count or mark as failed if max retries exceeded."""
table = fq_table("async_operations", schema)
# Get current retry count
row = await self._pool.fetchrow(
f"SELECT retry_count FROM {table} WHERE operation_id = $1",
operation_id,
)
if row is None:
logger.warning(f"Operation {operation_id} not found, cannot retry")
return
retry_count = row["retry_count"]
if retry_count >= self._max_retries:
# Max retries exceeded, mark as failed
await self._mark_failed(
operation_id, f"Max retries ({self._max_retries}) exceeded. Last error: {error_message}", schema
)
logger.error(f"Task {operation_id} failed after {retry_count} retries")
else:
# Increment retry and reset to pending
await self._pool.execute(
f"""
UPDATE {table}
SET status = 'pending', worker_id = NULL, claimed_at = NULL,
retry_count = retry_count + 1, updated_at = now()
WHERE operation_id = $1
""",
operation_id,
)
logger.warning(f"Task {operation_id} failed, will retry (attempt {retry_count + 1}/{self._max_retries})")
async def execute_task(self, task: ClaimedTask):
"""Execute a single task and update its status."""
task_type = task.task_dict.get("type", "unknown")
bank_id = task.task_dict.get("bank_id", "unknown")
# Track this task as active
async with self._in_flight_lock:
self._active_tasks[task.operation_id] = (task_type, bank_id, task.schema)
try:
schema_info = f", schema={task.schema}" if task.schema else ""
logger.debug(f"Executing task {task.operation_id} (type={task_type}, bank={bank_id}{schema_info})")
# Pass schema to executor so it can set the correct context
if task.schema:
task.task_dict["_schema"] = task.schema
await self._executor(task.task_dict)
await self._mark_completed(task.operation_id, task.schema)
logger.debug(f"Task {task.operation_id} completed successfully")
except Exception as e:
error_msg = f"{type(e).__name__}: {e}\n{traceback.format_exc()}"
logger.error(f"Task {task.operation_id} failed: {e}")
await self._retry_or_fail(task.operation_id, error_msg, task.schema)
finally:
# Remove from active tasks
async with self._in_flight_lock:
self._active_tasks.pop(task.operation_id, None)
async def recover_own_tasks(self) -> int:
"""
Recover tasks that were assigned to this worker but not completed.
This handles the case where a worker crashes while processing tasks.
On startup, we reset any tasks stuck in 'processing' for this worker_id
back to 'pending' so they can be picked up again.
If tenant_extension is configured, recovers across all tenant schemas.
Returns:
Number of tasks recovered
"""
schemas = await self._get_schemas()
total_count = 0
for schema in schemas:
table = fq_table("async_operations", schema)
result = await self._pool.execute(
f"""
UPDATE {table}
SET status = 'pending', worker_id = NULL, claimed_at = NULL, updated_at = now()
WHERE status = 'processing' AND worker_id = $1
""",
self._worker_id,
)
# Parse "UPDATE N" to get count
count = int(result.split()[-1]) if result else 0
total_count += count
if total_count > 0:
logger.info(f"Worker {self._worker_id} recovered {total_count} stale tasks from previous run")
return total_count
async def run(self):
"""
Main polling loop.
Continuously polls for pending tasks, claims them, and executes them
until shutdown is signaled.
If tenant_extension is configured, dynamically discovers schemas on each poll.
"""
# Recover any tasks from a previous crash before starting
await self.recover_own_tasks()
logger.info(f"Worker {self._worker_id} starting polling loop")
while not self._shutdown.is_set():
try:
# Claim a batch of tasks (across all tenant schemas if configured)
tasks = await self.claim_batch()
if tasks:
# Log batch info
task_types: dict[str, int] = {}
schemas_seen: set[str | None] = set()
for task in tasks:
t = task.task_dict.get("type", "unknown")
task_types[t] = task_types.get(t, 0) + 1
schemas_seen.add(task.schema)
types_str = ", ".join(f"{k}:{v}" for k, v in task_types.items())
schemas_str = ", ".join(s or "public" for s in schemas_seen)
logger.info(
f"Worker {self._worker_id} claimed {len(tasks)} tasks: {types_str} (schemas: {schemas_str})"
)
# Track in-flight tasks
async with self._in_flight_lock:
self._in_flight_count += len(tasks)
# Execute tasks concurrently
try:
await asyncio.gather(
*[self.execute_task(task) for task in tasks],
return_exceptions=True,
)
finally:
async with self._in_flight_lock:
self._in_flight_count -= len(tasks)
else:
# No tasks found, wait before polling again
try:
await asyncio.wait_for(
self._shutdown.wait(),
timeout=self._poll_interval_ms / 1000,
)
except asyncio.TimeoutError:
pass # Normal timeout, continue polling
# Log progress stats periodically
await self._log_progress_if_due()
except asyncio.CancelledError:
logger.info(f"Worker {self._worker_id} polling loop cancelled")
break
except Exception as e:
logger.error(f"Worker {self._worker_id} error in polling loop: {e}")
traceback.print_exc()
# Backoff on error
await asyncio.sleep(1)
logger.info(f"Worker {self._worker_id} polling loop stopped")
async def shutdown_graceful(self, timeout: float = 30.0):
"""
Signal shutdown and wait for current tasks to complete.
Args:
timeout: Maximum time to wait for in-flight tasks (seconds)
"""
logger.info(f"Worker {self._worker_id} initiating graceful shutdown")
self._shutdown.set()
# Wait for in-flight tasks to complete
start_time = asyncio.get_event_loop().time()
while asyncio.get_event_loop().time() - start_time < timeout:
async with self._in_flight_lock:
in_flight = self._in_flight_count
if in_flight == 0:
logger.info(f"Worker {self._worker_id} graceful shutdown complete")
return
logger.info(f"Worker {self._worker_id} waiting for {in_flight} in-flight tasks")
await asyncio.sleep(0.5)
logger.warning(f"Worker {self._worker_id} shutdown timeout after {timeout}s")
async def _log_progress_if_due(self):
"""Log progress stats every PROGRESS_LOG_INTERVAL seconds."""
now = time.time()
if now - self._last_progress_log < PROGRESS_LOG_INTERVAL:
return
self._last_progress_log = now
try:
# Get local active tasks (this worker only)
async with self._in_flight_lock:
in_flight = self._in_flight_count
active_tasks = dict(self._active_tasks) # Copy to avoid holding lock
# Build local processing breakdown grouped by (op_type, bank_id)
task_groups: dict[tuple[str, str], int] = {}
for op_type, bank_id, _ in active_tasks.values():
key = (op_type, bank_id)
task_groups[key] = task_groups.get(key, 0) + 1
processing_info = [f"{op}:{bank}({cnt})" for (op, bank), cnt in task_groups.items()]
processing_str = ", ".join(processing_info[:10]) if processing_info else "none"
if len(processing_info) > 10:
processing_str += f" +{len(processing_info) - 10} more"
# Get global stats from DB across all schemas
schemas = await self._get_schemas()
global_pending = 0
all_worker_counts: dict[str, int] = {}
async with self._pool.acquire() as conn:
for schema in schemas:
table = fq_table("async_operations", schema)
row = await conn.fetchrow(f"SELECT COUNT(*) as count FROM {table} WHERE status = 'pending'")
global_pending += row["count"] if row else 0
# Get processing breakdown by worker
worker_rows = await conn.fetch(
f"""
SELECT worker_id, COUNT(*) as count
FROM {table}
WHERE status = 'processing'
GROUP BY worker_id
"""
)
for wr in worker_rows:
wid = wr["worker_id"] or "unknown"
all_worker_counts[wid] = all_worker_counts.get(wid, 0) + wr["count"]
# Format other workers' processing counts
other_workers = []
for wid, cnt in all_worker_counts.items():
if wid != self._worker_id:
other_workers.append(f"{wid}:{cnt}")
others_str = ", ".join(other_workers) if other_workers else "none"
schemas_str = ", ".join(s or "public" for s in schemas)
logger.info(
f"[WORKER_STATS] worker={self._worker_id} in_flight={in_flight} | "
f"global: pending={global_pending} (schemas: {schemas_str}) | "
f"others: {others_str} | "
f"my_active: {processing_str}"
)
except Exception as e:
logger.debug(f"Failed to log progress stats: {e}")
@property
def worker_id(self) -> str:
"""Get the worker ID."""
return self._worker_id
@property
def is_shutdown(self) -> bool:
"""Check if shutdown has been signaled."""
return self._shutdown.is_set()
+15 -7
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "hindsight-api"
version = "0.3.0"
version = "0.4.1"
description = "Hindsight: Agent Memory That Works Like Human Memory"
readme = "README.md"
requires-python = ">=3.11"
@@ -25,7 +25,7 @@ dependencies = [
"psycopg2-binary>=2.9.11",
"tiktoken>=0.12.0",
"httpx>=0.27.0",
"fastmcp>=2.3.0",
"fastmcp>=2.14.0", # CVE-2025-66416
"pg0-embedded>=0.11.0",
"python-dateutil>=2.8.0",
"opentelemetry-api>=1.20.0",
@@ -39,10 +39,17 @@ dependencies = [
"cohere>=5.0.0",
"flashrank>=0.2.0",
# Local ML models for embeddings/reranking - can be excluded in Docker with INCLUDE_LOCAL_MODELS=false
"sentence-transformers>=3.0.0,<3.3.0",
"transformers>=4.30.0,<4.46.0",
"torch>=2.0.0",
"sentence-transformers>=3.3.0",
"transformers>=4.53.0", # Security fixes for ReDoS vulnerabilities
"torch>=2.6.0", # CVE fix for remote code execution
"uvloop>=0.22.1",
# Transitive dependency security fixes
"pyasn1>=0.6.2", # DoS vulnerability fix
"urllib3>=2.6.3", # Decompression-bomb safeguards bypass fix
"langchain-core>=1.2.5", # Serialization injection vulnerability fix
"filelock>=3.20.1", # TOCTOU race condition fix
"authlib>=1.6.6", # Account takeover vulnerability fix
"aiohttp>=3.13.3", # Multiple DoS vulnerabilities
]
[project.optional-dependencies]
@@ -51,11 +58,12 @@ test = [
"pytest-asyncio>=0.21.0",
"pytest-timeout>=2.4.0",
"pytest-xdist>=3.0.0",
"filelock>=3.0.0",
"filelock>=3.20.1", # TOCTOU race condition fix
]
[project.scripts]
hindsight-api = "hindsight_api.main:main"
hindsight-worker = "hindsight_api.worker.main:main"
hindsight-local-mcp = "hindsight_api.mcp_local:main"
hindsight-admin = "hindsight_api.admin.cli:main"
@@ -97,7 +105,7 @@ dev = [
"pytest-timeout>=2.4.0",
"pytest-xdist>=3.8.0",
"python-dotenv>=1.2.1",
"filelock>=3.0.0",
"filelock>=3.20.1", # TOCTOU race condition fix
"ruff>=0.8.0",
"ty>=0.0.1",
]
+56 -4
View File
@@ -12,6 +12,7 @@ from hindsight_api import MemoryEngine, LLMConfig, LocalSTEmbeddings, RequestCon
from hindsight_api.engine.cross_encoder import LocalSTCrossEncoder
from hindsight_api.engine.query_analyzer import DateparserQueryAnalyzer
from hindsight_api.engine.task_backend import SyncTaskBackend
from hindsight_api.pg0 import EmbeddedPostgres
# Default pg0 instance configuration for tests
@@ -115,16 +116,65 @@ def llm_config():
@pytest.fixture(scope="session")
def embeddings():
def embeddings(tmp_path_factory, worker_id):
"""
Session-scoped embeddings fixture with filelock to prevent race conditions.
return LocalSTEmbeddings()
When pytest-xdist runs multiple workers in parallel, they all try to load
models from the HuggingFace cache simultaneously, which can cause race
conditions and meta tensor errors. We use a filelock to serialize model
initialization across workers.
"""
# Get shared temp dir for coordination between xdist workers
if worker_id == "master":
root_tmp_dir = tmp_path_factory.getbasetemp()
else:
root_tmp_dir = tmp_path_factory.getbasetemp().parent
lock_file = root_tmp_dir / "embeddings_init.lock"
emb = LocalSTEmbeddings()
# Serialize model initialization across workers
with filelock.FileLock(str(lock_file)):
loop = asyncio.new_event_loop()
try:
loop.run_until_complete(emb.initialize())
finally:
loop.close()
return emb
@pytest.fixture(scope="session")
def cross_encoder():
def cross_encoder(tmp_path_factory, worker_id):
"""
Session-scoped cross-encoder fixture with filelock to prevent race conditions.
return LocalSTCrossEncoder()
When pytest-xdist runs multiple workers in parallel, they all try to load
models from the HuggingFace cache simultaneously, which can cause race
conditions and meta tensor errors. We use a filelock to serialize model
initialization across workers.
"""
# Get shared temp dir for coordination between xdist workers
if worker_id == "master":
root_tmp_dir = tmp_path_factory.getbasetemp()
else:
root_tmp_dir = tmp_path_factory.getbasetemp().parent
lock_file = root_tmp_dir / "cross_encoder_init.lock"
ce = LocalSTCrossEncoder()
# Serialize model initialization across workers
with filelock.FileLock(str(lock_file)):
loop = asyncio.new_event_loop()
try:
loop.run_until_complete(ce.initialize())
finally:
loop.close()
return ce
@pytest.fixture(scope="session")
def query_analyzer():
@@ -147,6 +197,7 @@ async def memory(pg0_db_url, embeddings, cross_encoder, query_analyzer):
Uses pg0_db_url (a postgresql:// URL) directly, so MemoryEngine won't try to
manage pg0 lifecycle - that's handled by the session-scoped pg0_db_url fixture.
Migrations are disabled here since they're run once at session scope in pg0_db_url.
Uses SyncTaskBackend so async tasks execute immediately (no worker needed).
"""
mem = MemoryEngine(
db_url=pg0_db_url, # Direct postgresql:// URL, not pg0://
@@ -160,6 +211,7 @@ async def memory(pg0_db_url, embeddings, cross_encoder, query_analyzer):
pool_min_size=1,
pool_max_size=5,
run_migrations=False, # Migrations already run at session scope
task_backend=SyncTaskBackend(), # Execute tasks immediately in tests
)
await mem.initialize()
yield mem
File diff suppressed because it is too large Load Diff
@@ -9,17 +9,18 @@ Includes tests for:
import asyncio
import os
import pytest
from datetime import datetime
import pytest
from sqlalchemy import create_engine, text
from hindsight_api import MemoryEngine, RequestContext
from hindsight_api.engine.embeddings import LocalSTEmbeddings, OpenAIEmbeddings, CohereEmbeddings
from hindsight_api.engine.cross_encoder import LocalSTCrossEncoder, CohereCrossEncoder
from hindsight_api.engine.cross_encoder import CohereCrossEncoder, LocalSTCrossEncoder
from hindsight_api.engine.embeddings import CohereEmbeddings, LocalSTEmbeddings, OpenAIEmbeddings
from hindsight_api.engine.query_analyzer import DateparserQueryAnalyzer
from hindsight_api.extensions import TenantExtension, TenantContext
from hindsight_api.migrations import run_migrations, ensure_embedding_dimension
from hindsight_api.engine.task_backend import SyncTaskBackend
from hindsight_api.extensions import TenantContext, TenantExtension
from hindsight_api.migrations import ensure_embedding_dimension, run_migrations
# =============================================================================
# Shared Utilities
@@ -35,6 +36,11 @@ class SchemaTenantExtension(TenantExtension):
async def authenticate(self, request_context: RequestContext) -> TenantContext:
return TenantContext(schema_name=self.schema_name)
async def list_tenants(self) -> list:
from hindsight_api.extensions.tenant import Tenant
return [Tenant(schema=self.schema_name)]
def get_test_schema(prefix: str, worker_id: str) -> str:
"""Get unique schema name per xdist worker."""
@@ -323,6 +329,7 @@ class TestOpenAIEmbeddings:
pool_max_size=3,
run_migrations=False,
tenant_extension=SchemaTenantExtension(schema_name),
task_backend=SyncTaskBackend(),
)
try:
@@ -392,6 +399,7 @@ class TestOpenAIEmbeddings:
pool_max_size=3,
run_migrations=False,
tenant_extension=SchemaTenantExtension(schema_name),
task_backend=SyncTaskBackend(),
)
try:
@@ -559,6 +567,7 @@ class TestCohereIntegration:
pool_max_size=3,
run_migrations=False,
tenant_extension=SchemaTenantExtension(schema_name),
task_backend=SyncTaskBackend(),
)
try:
@@ -1,516 +0,0 @@
"""Tests for emergent entity filtering."""
import pytest
from unittest.mock import AsyncMock, MagicMock
from hindsight_api.engine.mental_models.emergent import (
build_mission_filter_prompt,
evaluate_emergent_models,
filter_candidates_by_mission,
MissionFilterResponse,
MissionFilterCandidate,
)
from hindsight_api.engine.mental_models.models import EmergentCandidate
class TestBuildMissionFilterPrompt:
"""Test prompt building for mission filtering."""
def test_prompt_contains_mission(self):
"""Test that prompt includes the mission."""
candidates = [
EmergentCandidate(
name="Alice",
detection_method="named_entity_extraction",
mention_count=10,
)
]
prompt = build_mission_filter_prompt("Be a PM for engineering team", candidates)
assert "Be a PM for engineering team" in prompt
def test_prompt_contains_candidates(self):
"""Test that prompt includes all candidates."""
candidates = [
EmergentCandidate(
name="Alice Chen",
detection_method="named_entity_extraction",
mention_count=10,
),
EmergentCandidate(
name="Project Phoenix",
detection_method="named_entity_extraction",
mention_count=5,
),
]
prompt = build_mission_filter_prompt("Track projects", candidates)
assert "Alice Chen" in prompt
assert "Project Phoenix" in prompt
def test_prompt_contains_rejection_guidance(self):
"""Test that prompt contains guidance to reject generic entities."""
candidates = [
EmergentCandidate(
name="test",
detection_method="named_entity_extraction",
mention_count=1,
)
]
prompt = build_mission_filter_prompt("Test mission", candidates)
# Should contain rejection guidance for generic terms
assert "promote=false" in prompt
assert "kids" in prompt # Example of generic term to reject
assert "community" in prompt # Example of abstract concept to reject
assert "motivation" in prompt # Example of abstract concept to reject
class TestFilterCandidatesByMission:
"""Test the filter_candidates_by_mission function."""
@pytest.fixture
def mock_llm_config(self):
"""Create a mock LLM config."""
config = MagicMock()
config.call = AsyncMock()
return config
async def test_empty_candidates(self, mock_llm_config):
"""Test with empty candidate list."""
result = await filter_candidates_by_mission(
llm_config=mock_llm_config,
mission="Test mission",
candidates=[],
)
assert result == []
mock_llm_config.call.assert_not_called()
async def test_no_mission_keeps_all(self, mock_llm_config):
"""Test that no mission keeps all candidates (skips filtering)."""
candidates = [
EmergentCandidate(
name="Alice",
detection_method="named_entity_extraction",
mention_count=10,
)
]
result = await filter_candidates_by_mission(
llm_config=mock_llm_config,
mission="", # Empty mission
candidates=candidates,
)
assert len(result) == 1
assert result[0].name == "Alice"
mock_llm_config.call.assert_not_called()
async def test_filters_by_promote_flag(self, mock_llm_config):
"""Test that candidates are filtered by promote flag."""
candidates = [
EmergentCandidate(
name="Alice Chen",
detection_method="named_entity_extraction",
mention_count=10,
),
EmergentCandidate(
name="community",
detection_method="named_entity_extraction",
mention_count=5,
),
]
# Mock LLM response - Alice is promoted, community is not
mock_llm_config.call.return_value = MissionFilterResponse(
candidates=[
MissionFilterCandidate(name="Alice Chen", promote=True, reason="Specific person"),
MissionFilterCandidate(name="community", promote=False, reason="Generic abstract concept"),
]
)
result = await filter_candidates_by_mission(
llm_config=mock_llm_config,
mission="Be a PM for engineering team",
candidates=candidates,
)
assert len(result) == 1
assert result[0].name == "Alice Chen"
async def test_rejects_generic_entities(self, mock_llm_config):
"""Test that generic entities are rejected."""
# These are all generic/abstract terms that should be rejected
generic_names = [
"user", "support", "community", "family", "motivation",
"photo", "gratitude", "difference", "volunteering",
"kids", "veterans", "impact", "kindness", "encouragement",
"education", "nature", "joy", "positivity", "inspiration",
"help", "commitment", "passion", "energy", "connection",
]
candidates = [
EmergentCandidate(
name=name,
detection_method="named_entity_extraction",
mention_count=10,
)
for name in generic_names
]
# Add some valid candidates
valid_candidates = [
EmergentCandidate(
name="John",
detection_method="named_entity_extraction",
mention_count=10,
),
EmergentCandidate(
name="Maria",
detection_method="named_entity_extraction",
mention_count=8,
),
EmergentCandidate(
name="Max",
detection_method="named_entity_extraction",
mention_count=6,
),
]
candidates.extend(valid_candidates)
# Mock LLM response - reject all generic, promote only specific names
response_candidates = [
MissionFilterCandidate(name=name, promote=False, reason="Generic/abstract term")
for name in generic_names
]
response_candidates.extend([
MissionFilterCandidate(name=c.name, promote=True, reason="Specific person name")
for c in valid_candidates
])
mock_llm_config.call.return_value = MissionFilterResponse(candidates=response_candidates)
result = await filter_candidates_by_mission(
llm_config=mock_llm_config,
mission="Be a health coach",
candidates=candidates,
)
# Should only have John, Maria, and Max
result_names = {c.name for c in result}
assert result_names == {"John", "Maria", "Max"}
async def test_accepts_specific_named_entities(self, mock_llm_config):
"""Test that specific named entities are accepted."""
# These should all be accepted
valid_names = [
"Alice Chen", # Full name
"Dr. Smith", # Title + name
"John", # First name (when it's clearly a person)
"Google", # Organization
"Frontend Team", # Named team
"Project Phoenix", # Named project
"NYC Office", # Named place
"Q4 Planning", # Named event
"Sprint 23 Review", # Named meeting
]
candidates = [
EmergentCandidate(
name=name,
detection_method="named_entity_extraction",
mention_count=10,
)
for name in valid_names
]
# Mock LLM response - promote all
response_candidates = [
MissionFilterCandidate(name=name, promote=True, reason="Specific named entity")
for name in valid_names
]
mock_llm_config.call.return_value = MissionFilterResponse(candidates=response_candidates)
result = await filter_candidates_by_mission(
llm_config=mock_llm_config,
mission="Be a PM for engineering team",
candidates=candidates,
)
# Should have all valid names
result_names = {c.name for c in result}
assert result_names == set(valid_names)
async def test_llm_error_rejects_all_candidates(self, mock_llm_config):
"""Test that LLM errors result in rejecting all candidates (fail-safe)."""
candidates = [
EmergentCandidate(
name="Alice",
detection_method="named_entity_extraction",
mention_count=10,
)
]
mock_llm_config.call.side_effect = Exception("LLM error")
result = await filter_candidates_by_mission(
llm_config=mock_llm_config,
mission="Test mission",
candidates=candidates,
)
# Should reject all candidates on error (fail-safe)
assert len(result) == 0
async def test_missing_candidate_in_response_is_rejected(self, mock_llm_config):
"""Test that candidates not in LLM response are rejected by default."""
candidates = [
EmergentCandidate(
name="Alice",
detection_method="named_entity_extraction",
mention_count=10,
),
EmergentCandidate(
name="Bob",
detection_method="named_entity_extraction",
mention_count=5,
),
]
# Mock LLM response - only includes Alice, not Bob
mock_llm_config.call.return_value = MissionFilterResponse(
candidates=[
MissionFilterCandidate(name="Alice", promote=True, reason="Specific person"),
]
)
result = await filter_candidates_by_mission(
llm_config=mock_llm_config,
mission="Test mission",
candidates=candidates,
)
# Only Alice should be in result (Bob was missing from response, so rejected)
assert len(result) == 1
assert result[0].name == "Alice"
class TestEvaluateEmergentModels:
"""Test the evaluate_emergent_models function for cleanup of existing models."""
@pytest.fixture
def mock_llm_config(self):
"""Create a mock LLM config."""
config = MagicMock()
config.call = AsyncMock()
return config
async def test_empty_models(self, mock_llm_config):
"""Test with empty model list."""
result = await evaluate_emergent_models(
llm_config=mock_llm_config,
models=[],
)
assert result == []
mock_llm_config.call.assert_not_called()
async def test_removes_generic_models(self, mock_llm_config):
"""Test that generic/abstract models are marked for removal."""
models = [
{"id": "id-kids", "name": "kids"},
{"id": "id-community", "name": "community"},
{"id": "id-motivation", "name": "motivation"},
{"id": "id-john", "name": "John"},
{"id": "id-maria", "name": "Maria"},
]
# Mock LLM response - reject generic, keep specific names
mock_llm_config.call.return_value = MissionFilterResponse(
candidates=[
MissionFilterCandidate(name="kids", promote=False, reason="Generic category"),
MissionFilterCandidate(name="community", promote=False, reason="Abstract concept"),
MissionFilterCandidate(name="motivation", promote=False, reason="Abstract concept"),
MissionFilterCandidate(name="John", promote=True, reason="Person name"),
MissionFilterCandidate(name="Maria", promote=True, reason="Person name"),
]
)
result = await evaluate_emergent_models(
llm_config=mock_llm_config,
models=models,
)
# Should return IDs of generic models to remove
assert set(result) == {"id-kids", "id-community", "id-motivation"}
async def test_keeps_specific_named_models(self, mock_llm_config):
"""Test that specific named models are kept."""
models = [
{"id": "id-john", "name": "John"},
{"id": "id-google", "name": "Google"},
{"id": "id-project", "name": "Project Phoenix"},
]
# Mock LLM response - keep all
mock_llm_config.call.return_value = MissionFilterResponse(
candidates=[
MissionFilterCandidate(name="John", promote=True, reason="Person name"),
MissionFilterCandidate(name="Google", promote=True, reason="Organization"),
MissionFilterCandidate(name="Project Phoenix", promote=True, reason="Named project"),
]
)
result = await evaluate_emergent_models(
llm_config=mock_llm_config,
models=models,
)
# No models should be removed
assert result == []
async def test_llm_error_keeps_all_models(self, mock_llm_config):
"""Test that LLM errors result in keeping all models (safe default)."""
models = [
{"id": "id-kids", "name": "kids"},
{"id": "id-john", "name": "John"},
]
mock_llm_config.call.side_effect = Exception("LLM error")
result = await evaluate_emergent_models(
llm_config=mock_llm_config,
models=models,
)
# Should keep all models on error (return empty removal list)
assert result == []
async def test_missing_model_in_response_is_removed(self, mock_llm_config):
"""Test that models not in LLM response are marked for removal."""
models = [
{"id": "id-alice", "name": "Alice"},
{"id": "id-bob", "name": "Bob"},
]
# Mock LLM response - only includes Alice
mock_llm_config.call.return_value = MissionFilterResponse(
candidates=[
MissionFilterCandidate(name="Alice", promote=True, reason="Person name"),
]
)
result = await evaluate_emergent_models(
llm_config=mock_llm_config,
models=models,
)
# Bob should be marked for removal (missing from response)
assert result == ["id-bob"]
class TestRemovedEntitiesNotRepromoted:
"""Test that entities removed by evaluation are not re-promoted.
This tests the fix for a bug where:
1. evaluate_emergent_models returns model IDs to remove (e.g., 'entity-maya')
2. We delete those models
3. detect_entity_candidates finds the same entities (now eligible since model was deleted)
4. filter_candidates_by_goal approves them (different LLM call)
5. BUG: We were re-promoting the same entities we just removed
The fix tracks removed entity_ids and excludes them from promotion.
"""
async def test_removed_entity_ids_excluded_from_promotion(self):
"""Test that entities whose models were removed are not re-promoted."""
from hindsight_api.engine.mental_models.models import EmergentCandidate
# Simulate the scenario from the bug:
# - existing_emergent has model 'entity-maya' with entity_id='uuid-maya'
# - evaluate_emergent_models says to remove 'entity-maya'
# - detect_entity_candidates returns 'Maya' with entity_id='uuid-maya' (now eligible)
# - filter_candidates_by_goal says to promote 'Maya'
# - But we should NOT promote because we just removed it
existing_emergent = [
{"id": "entity-maya", "name": "Maya", "entity_id": "uuid-maya"},
{"id": "entity-alex", "name": "Alex", "entity_id": "uuid-alex"},
{"id": "entity-john", "name": "John", "entity_id": "uuid-john"}, # This one will be kept
]
# Models to remove (evaluate_emergent_models would return these)
models_to_remove = ["entity-maya", "entity-alex"]
# Build model_id -> entity_id mapping (this is what the fix does)
model_to_entity = {m["id"]: m.get("entity_id") for m in existing_emergent}
# Track removed entity_ids
removed_entity_ids: set[str] = set()
for model_id in models_to_remove:
entity_id = model_to_entity.get(model_id)
if entity_id:
removed_entity_ids.add(str(entity_id))
# Verify we tracked the right entity_ids
assert removed_entity_ids == {"uuid-maya", "uuid-alex"}
# Now simulate candidates that were detected (includes removed entities)
candidates = [
EmergentCandidate(
name="Maya", entity_id="uuid-maya", detection_method="named_entity", mention_count=10
),
EmergentCandidate(
name="Alex", entity_id="uuid-alex", detection_method="named_entity", mention_count=8
),
EmergentCandidate(
name="NewPerson", entity_id="uuid-new", detection_method="named_entity", mention_count=5
),
]
# Filter out candidates whose entity was just removed (the fix)
filtered_candidates = [c for c in candidates if c.entity_id not in removed_entity_ids]
# Only NewPerson should remain - Maya and Alex were removed and should not be re-promoted
assert len(filtered_candidates) == 1
assert filtered_candidates[0].name == "NewPerson"
assert filtered_candidates[0].entity_id == "uuid-new"
async def test_candidates_without_matching_removal_are_kept(self):
"""Test that candidates not in the removed set are still promoted."""
from hindsight_api.engine.mental_models.models import EmergentCandidate
# No models removed
removed_entity_ids: set[str] = set()
candidates = [
EmergentCandidate(
name="Alice", entity_id="uuid-alice", detection_method="named_entity", mention_count=10
),
EmergentCandidate(
name="Bob", entity_id="uuid-bob", detection_method="named_entity", mention_count=8
),
]
# Filter (should keep all since nothing was removed)
filtered_candidates = [c for c in candidates if c.entity_id not in removed_entity_ids]
assert len(filtered_candidates) == 2
assert {c.name for c in filtered_candidates} == {"Alice", "Bob"}
async def test_partial_removal_keeps_other_candidates(self):
"""Test that only removed entities are excluded, others pass through."""
from hindsight_api.engine.mental_models.models import EmergentCandidate
# Only one entity removed
removed_entity_ids = {"uuid-removed"}
candidates = [
EmergentCandidate(
name="Removed", entity_id="uuid-removed", detection_method="named_entity", mention_count=10
),
EmergentCandidate(
name="Kept1", entity_id="uuid-kept1", detection_method="named_entity", mention_count=8
),
EmergentCandidate(
name="Kept2", entity_id="uuid-kept2", detection_method="named_entity", mention_count=5
),
]
filtered_candidates = [c for c in candidates if c.entity_id not in removed_entity_ids]
assert len(filtered_candidates) == 2
assert {c.name for c in filtered_candidates} == {"Kept1", "Kept2"}
+17 -2
View File
@@ -24,6 +24,9 @@ from hindsight_api.extensions import (
TenantExtension,
ValidationResult,
load_extension,
# Consolidation operation
ConsolidateContext,
ConsolidateResult,
)
@@ -128,14 +131,18 @@ class TrackingValidator(OperationValidatorExtension):
def __init__(self, config: dict):
super().__init__(config)
# Pre-hook tracking
# Pre-hook tracking - Core operations
self.pre_retain_calls: list[RetainContext] = []
self.pre_recall_calls: list[RecallContext] = []
self.pre_reflect_calls: list[ReflectContext] = []
# Post-hook tracking
# Post-hook tracking - Core operations
self.post_retain_calls: list[RetainResult] = []
self.post_recall_calls: list[RecallResult] = []
self.post_reflect_calls: list[ReflectResultContext] = []
# Pre-hook tracking - Consolidation
self.pre_consolidate_calls: list[ConsolidateContext] = []
# Post-hook tracking - Consolidation
self.post_consolidate_calls: list[ConsolidateResult] = []
async def validate_retain(self, ctx: RetainContext) -> ValidationResult:
self.pre_retain_calls.append(ctx)
@@ -158,6 +165,14 @@ class TrackingValidator(OperationValidatorExtension):
async def on_reflect_complete(self, result: ReflectResultContext) -> None:
self.post_reflect_calls.append(result)
# Consolidation hooks
async def validate_consolidate(self, ctx: ConsolidateContext) -> ValidationResult:
self.pre_consolidate_calls.append(ctx)
return ValidationResult.accept()
async def on_consolidate_complete(self, result: ConsolidateResult) -> None:
self.post_consolidate_calls.append(result)
class TestMemoryEngineValidation:
"""Tests for validation integration with MemoryEngine.
@@ -969,24 +969,22 @@ async def test_reflect_returns_token_usage(api_client):
assert "text" in result
assert len(result["text"]) > 0
# Verify usage field exists (may be None for agentic reflect which makes multiple LLM calls)
# Verify usage field exists and is populated (agentic reflect aggregates all LLM calls)
assert "usage" in result, "Response should include 'usage' field"
usage = result["usage"]
# Usage is optional - agentic reflect doesn't aggregate multiple LLM call usages
if usage is not None:
assert "input_tokens" in usage, "Usage should have 'input_tokens'"
assert "output_tokens" in usage, "Usage should have 'output_tokens'"
assert "total_tokens" in usage, "Usage should have 'total_tokens'"
# Usage must be present - agentic reflect now aggregates token usage from all LLM calls
assert usage is not None, "Usage should not be None - reflect aggregates all LLM call usages"
assert "input_tokens" in usage, "Usage should have 'input_tokens'"
assert "output_tokens" in usage, "Usage should have 'output_tokens'"
assert "total_tokens" in usage, "Usage should have 'total_tokens'"
# Verify token counts are valid
assert usage["input_tokens"] > 0, f"Expected input_tokens > 0, got {usage['input_tokens']}"
assert usage["output_tokens"] >= 0, f"Expected output_tokens >= 0, got {usage['output_tokens']}"
assert usage["total_tokens"] == usage["input_tokens"] + usage["output_tokens"]
# Verify token counts are valid
assert usage["input_tokens"] > 0, f"Expected input_tokens > 0, got {usage['input_tokens']}"
assert usage["output_tokens"] >= 0, f"Expected output_tokens >= 0, got {usage['output_tokens']}"
assert usage["total_tokens"] == usage["input_tokens"] + usage["output_tokens"]
print(f"Reflect token usage: input={usage['input_tokens']}, output={usage['output_tokens']}, total={usage['total_tokens']}")
else:
print("Reflect usage is None (expected for agentic reflect)")
print(f"Reflect token usage: input={usage['input_tokens']}, output={usage['output_tokens']}, total={usage['total_tokens']}")
@pytest.mark.asyncio
@@ -1065,3 +1063,38 @@ async def test_retain_async_no_usage(api_client):
# Usage should be None for async operations
assert result.get("usage") is None, "Async retain should not include usage"
@pytest.mark.asyncio
async def test_version_endpoint_returns_correct_version(api_client):
"""Test that the /version endpoint returns the correct API version.
The version should match the __version__ defined in hindsight_api.__init__.py
and should not be a hardcoded string.
"""
from hindsight_api import __version__
# Call the /version endpoint
response = await api_client.get("/version")
assert response.status_code == 200
result = response.json()
# Verify response structure
assert "api_version" in result, "Response should include 'api_version' field"
assert "features" in result, "Response should include 'features' field"
# Verify the version matches the package version
assert result["api_version"] == __version__, (
f"API version should be {__version__}, got {result['api_version']}"
)
# Verify features field structure
features = result["features"]
assert "observations" in features
assert "mcp" in features
assert "worker" in features
assert isinstance(features["observations"], bool)
assert isinstance(features["mcp"], bool)
assert isinstance(features["worker"], bool)
print(f"Version endpoint returned: api_version={result['api_version']}, features={features}")
@@ -0,0 +1,278 @@
"""
Tests for LinkExpansion graph retrieval.
Tests cover the entity-based graph traversal for observations.
"""
from datetime import datetime, timezone
import pytest
@pytest.fixture(autouse=True)
def enable_observations():
"""Enable observations for all tests in this module."""
from hindsight_api.config import get_config
config = get_config()
original_value = config.enable_observations
config.enable_observations = True
yield
config.enable_observations = original_value
@pytest.mark.asyncio
async def test_link_expansion_observation_graph_retrieval(memory, request_context):
"""
Test that observations can find other observations via shared entities.
This tests the scenario where:
1. World fact A has entity "Python"
2. World fact B has entity "Python"
3. Observation OA is derived from world fact A
4. Observation OB is derived from world fact B
When searching for observations related to OA, graph retrieval should find OB
because they share the "Python" entity through their source world facts.
Current issue: Graph retrieval returns 0 for observations because:
- Entity links are copied from world facts to observations during consolidation
- But the entity expansion query filters by fact_type
- Observations only share entities with world facts (cross-type), not with other observations
- So filtering to fact_type='observation' returns 0 results
"""
bank_id = f"test_link_expansion_obs_{datetime.now(timezone.utc).timestamp()}"
try:
# Store world facts with shared entities using retain_batch_async
# We need enough facts that semantic search won't return all of them as seeds
# Key: "Alice" query should find Alice's observation but NOT Bob's via semantic search
# Then graph retrieval should find Bob via shared "Python" entity
await memory.retain_batch_async(
bank_id=bank_id,
contents=[
# Python developers - should be connected via "Python" entity
{
"content": "Alice works with Python at TechCorp building REST APIs",
"context": "employee info",
"entities": [{"text": "Python"}, {"text": "Alice"}, {"text": "TechCorp"}],
},
{
"content": "Bob uses Python at DataSoft for machine learning models",
"context": "employee info",
"entities": [{"text": "Python"}, {"text": "Bob"}, {"text": "DataSoft"}],
},
# Many unrelated facts to dilute semantic search and ensure
# "Alice" query only finds Alice-related content as seeds
{
"content": "The weather in San Francisco is often foggy and cool",
"context": "weather info",
"entities": [{"text": "San Francisco"}],
},
{
"content": "Tokyo is the capital city of Japan with many trains",
"context": "geography info",
"entities": [{"text": "Tokyo"}, {"text": "Japan"}],
},
{
"content": "The Great Wall of China is a historic fortification",
"context": "history info",
"entities": [{"text": "Great Wall"}, {"text": "China"}],
},
{
"content": "Coffee beans are grown in tropical regions worldwide",
"context": "food info",
"entities": [{"text": "Coffee"}],
},
{
"content": "Electric vehicles are becoming more popular globally",
"context": "technology info",
"entities": [{"text": "Electric vehicles"}],
},
{
"content": "The Amazon rainforest contains diverse wildlife species",
"context": "nature info",
"entities": [{"text": "Amazon"}, {"text": "Rainforest"}],
},
{
"content": "Basketball is a popular sport in the United States",
"context": "sports info",
"entities": [{"text": "Basketball"}, {"text": "United States"}],
},
{
"content": "Mozart composed many famous classical music pieces",
"context": "music info",
"entities": [{"text": "Mozart"}, {"text": "Classical music"}],
},
],
request_context=request_context,
)
# Consolidation runs automatically after retain - wait for it to complete
# by querying for observations (consolidation creates them)
import asyncio
from hindsight_api.engine.memory_engine import Budget
# Wait for consolidation to complete with retry logic
# Consolidation runs as a background task and may take longer in CI
obs_result = None
for _ in range(30): # Try up to 30 times (30 seconds max)
await asyncio.sleep(1) # Wait 1 second between attempts
obs_result = await memory.recall_async(
bank_id=bank_id,
query="Python developer",
fact_type=["observation"],
budget=Budget.MID,
max_tokens=2048,
request_context=request_context,
)
if obs_result.results and len(obs_result.results) >= 1:
break
assert obs_result is not None and obs_result.results is not None, "Should have observations after consolidation"
# We should have observations from consolidation
assert len(obs_result.results) >= 1, f"Should have at least 1 observation about Python, got {len(obs_result.results)}"
# Now test graph retrieval specifically
# Query for Alice - should find Bob via shared "Python" entity
result = await memory.recall_async(
bank_id=bank_id,
query="Alice",
fact_type=["observation"],
budget=Budget.MID,
max_tokens=2048,
enable_trace=True,
request_context=request_context,
)
# Verify graph retrieval is working by checking the internal debug logs
# The graph retrieval finds observations via entity links, but may not return
# NEW results if semantic search already found all connected observations.
# This is correct behavior - we verify the entity traversal path works.
# Check the trace for graph results
assert result.trace is not None, "Should have trace data"
# The key verification: the entity expansion path works (sources -> entities -> observations)
# We validated this in the debug logs above:
# - Observations have source_memory_ids pointing to world facts ✓
# - World facts have entity links ✓
# - Graph retrieval can traverse this path (seen in logs: potential_obs > 0)
# For a more rigorous test, we need data where semantic search misses something.
# Let's verify the world fact graph retrieval works (it uses direct entity links).
world_result = await memory.recall_async(
bank_id=bank_id,
query="Alice",
fact_type=["world"],
budget=Budget.MID,
max_tokens=2048,
enable_trace=True,
request_context=request_context,
)
assert world_result.trace is not None, "Should have trace data for world facts"
world_retrieval_results = world_result.trace.get("retrieval_results", [])
world_graph_results = [
r for r in world_retrieval_results if r.get("method_name") == "graph"
]
if world_graph_results:
world_graph_result = [r for r in world_graph_results if r.get("fact_type") == "world"][0]
world_graph_results_list = world_graph_result.get("results", [])
# World facts use direct entity links, so graph may find results
if world_graph_results_list:
print(f"\n✓ Graph retrieval found {len(world_graph_results_list)} connected world facts")
graph_texts = [r.get("text", "") for r in world_graph_results_list]
bob_found = any("Bob" in t or "DataSoft" in t for t in graph_texts)
if bob_found:
print(" Found Bob's world fact via shared 'Python' entity!")
print("\n✓ Link expansion observation test passed!")
print(" Entity traversal path verified (observations -> sources -> entities -> connected sources -> observations)")
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_link_expansion_world_fact_graph_retrieval(memory, request_context):
"""
Test that world facts can find other world facts via shared entities.
This verifies the direct entity link traversal for world facts works correctly.
Note: When semantic search finds all world facts as seeds, graph retrieval
won't return NEW results (this is correct - it shouldn't duplicate results).
"""
bank_id = f"test_link_expansion_world_{datetime.now(timezone.utc).timestamp()}"
try:
# Store world facts with shared entities
await memory.retain_batch_async(
bank_id=bank_id,
contents=[
# Python developers - should be connected via "Python" entity
{
"content": "Alice works with Python at TechCorp building REST APIs",
"context": "employee info",
"entities": [{"text": "Python"}, {"text": "Alice"}, {"text": "TechCorp"}],
},
{
"content": "Bob uses Python at DataSoft for machine learning models",
"context": "employee info",
"entities": [{"text": "Python"}, {"text": "Bob"}, {"text": "DataSoft"}],
},
# Unrelated facts
{
"content": "The weather in San Francisco is often foggy",
"context": "weather info",
"entities": [{"text": "San Francisco"}],
},
{
"content": "Coffee beans are grown in tropical regions",
"context": "food info",
"entities": [{"text": "Coffee"}],
},
],
request_context=request_context,
)
from hindsight_api.engine.memory_engine import Budget
# Query for Alice
result = await memory.recall_async(
bank_id=bank_id,
query="Alice",
fact_type=["world"],
budget=Budget.MID,
max_tokens=2048,
enable_trace=True,
request_context=request_context,
)
assert result.trace is not None, "Should have trace data"
# Verify graph retrieval ran (it may or may not find new results depending
# on whether semantic search already found everything)
retrieval_results = result.trace.get("retrieval_results", [])
graph_results = [
r for r in retrieval_results if r.get("method_name") == "graph"
]
assert len(graph_results) > 0, "Should have graph retrieval results in trace"
# The important thing is that recall works and returns relevant results
assert result.results is not None and len(result.results) > 0, (
"Should return results for 'Alice' query"
)
# Alice's result should be at or near the top
result_texts = [r.text for r in result.results]
alice_found = any("Alice" in t for t in result_texts)
assert alice_found, f"Should find Alice in results: {result_texts[:3]}"
print("\n✓ Link expansion world fact test passed!")
print(f" Recall returned {len(result.results)} results for 'Alice' query")
finally:
await memory.delete_bank(bank_id, request_context=request_context)
+10 -18
View File
@@ -241,48 +241,40 @@ class TestReflectToolSchemas:
tools = get_reflect_tools()
tool_names = [t["function"]["name"] for t in tools]
assert "list_mental_models" in tool_names
assert "get_mental_model" in tool_names
assert "search_mental_models" in tool_names
assert "search_observations" in tool_names
assert "recall" in tool_names
assert "learn" in tool_names
assert "expand" in tool_names
assert "done" in tool_names
def test_get_reflect_tools_without_learn(self):
"""Test getting reflect tools without learn."""
def test_get_reflect_tools_with_directives(self):
"""Test getting reflect tools with directive rules."""
from hindsight_api.engine.reflect.tools_schema import get_reflect_tools
tools = get_reflect_tools(enable_learn=False)
tools = get_reflect_tools(directive_rules=["Always respond in French"])
tool_names = [t["function"]["name"] for t in tools]
assert "learn" not in tool_names
assert "recall" in tool_names
assert "done" in tool_names
def test_get_reflect_tools_observations_mode(self):
"""Test getting reflect tools with observations output mode."""
from hindsight_api.engine.reflect.tools_schema import get_reflect_tools
tools = get_reflect_tools(output_mode="observations")
# Done tool should have directive_compliance field when directives are present
done_tool = next(t for t in tools if t["function"]["name"] == "done")
params = done_tool["function"]["parameters"]["properties"]
assert "observations" in params
assert "answer" not in params
assert "directive_compliance" in params
def test_get_reflect_tools_answer_mode(self):
"""Test getting reflect tools with answer output mode."""
from hindsight_api.engine.reflect.tools_schema import get_reflect_tools
tools = get_reflect_tools(output_mode="answer")
tools = get_reflect_tools()
done_tool = next(t for t in tools if t["function"]["name"] == "done")
params = done_tool["function"]["parameters"]["properties"]
assert "answer" in params
assert "memory_ids" in params
assert "model_ids" in params
assert "observation_ids" in params
assert "mental_model_ids" in params
class TestLLMToolCallResult:
@@ -19,6 +19,7 @@ import pytest_asyncio
from hindsight_api import MemoryEngine, LLMConfig, LocalSTEmbeddings, RequestContext
from hindsight_api.engine.cross_encoder import LocalSTCrossEncoder
from hindsight_api.engine.query_analyzer import DateparserQueryAnalyzer
from hindsight_api.engine.task_backend import SyncTaskBackend
from hindsight_api.engine.retain.fact_extraction import FactExtractionResponse, ExtractedFact
from hindsight_api.engine.llm_wrapper import TokenUsage
@@ -106,6 +107,7 @@ class TestLargeBatchRetain:
pool_max_size=10,
run_migrations=False,
skip_llm_verification=True, # Skip LLM verification since we're mocking
task_backend=SyncTaskBackend(), # Execute tasks immediately in tests
)
await mem.initialize()
yield mem
+10 -5
View File
@@ -355,14 +355,14 @@ class TestMainModuleExtensionLoading:
# Mock extensions for testing
from hindsight_api.extensions import (
TenantExtension,
TenantContext,
RequestContext,
OperationValidatorExtension,
ValidationResult,
RetainContext,
RecallContext,
ReflectContext,
RequestContext,
RetainContext,
TenantContext,
TenantExtension,
ValidationResult,
)
@@ -376,6 +376,11 @@ class MockTenantExtension(TenantExtension):
async def authenticate(self, request_context: RequestContext) -> TenantContext:
return TenantContext(schema_name="public")
async def list_tenants(self) -> list:
from hindsight_api.extensions.tenant import Tenant
return [Tenant(schema="public")]
def set_context(self, context) -> None:
self._context_set = True
+55 -5
View File
@@ -62,9 +62,9 @@ async def test_local_mcp_server_recall(mock_memory):
tools = mcp_server._tool_manager._tools
assert "recall" in tools
# Call recall with new params
# Call recall
recall_tool = tools["recall"]
result = await recall_tool.fn(query="test query", max_tokens=2048, budget="mid")
result = await recall_tool.fn(query="test query", max_tokens=2048)
# Result is a dict
assert isinstance(result, dict)
@@ -75,7 +75,7 @@ async def test_local_mcp_server_recall(mock_memory):
assert call_kwargs["bank_id"] == "test-bank"
assert call_kwargs["query"] == "test query"
assert call_kwargs["max_tokens"] == 2048
assert call_kwargs["budget"] == Budget.MID
assert call_kwargs["budget"] == Budget.HIGH
@pytest.mark.asyncio
@@ -141,7 +141,7 @@ async def test_local_mcp_server_recall_error_handling(mock_memory):
@pytest.mark.asyncio
async def test_local_mcp_server_recall_with_defaults(mock_memory):
"""Test that recall uses default max_tokens and budget."""
"""Test that recall uses default max_tokens and HIGH budget."""
from hindsight_api.mcp_local import create_local_mcp_server
from hindsight_api.engine.memory_engine import Budget
@@ -159,4 +159,54 @@ async def test_local_mcp_server_recall_with_defaults(mock_memory):
call_kwargs = mock_memory.recall_async.call_args.kwargs
assert call_kwargs["max_tokens"] == 4096
assert call_kwargs["budget"] == Budget.LOW
assert call_kwargs["budget"] == Budget.HIGH
@pytest.mark.asyncio
async def test_local_mcp_server_retain_with_timestamp(mock_memory):
"""Test that retain passes timestamp as event_date."""
from datetime import datetime, timezone
from hindsight_api.mcp_local import create_local_mcp_server
mcp_server = create_local_mcp_server("test-bank", memory=mock_memory)
tools = mcp_server._tool_manager._tools
retain_tool = tools["retain"]
# Call retain with timestamp
result = await retain_tool.fn(
content="test content", context="test_context", timestamp="2024-01-15T10:30:00Z"
)
assert result["status"] == "accepted"
# Wait for background task
await asyncio.sleep(0.1)
call_kwargs = mock_memory.retain_batch_async.call_args.kwargs
contents = call_kwargs["contents"]
assert len(contents) == 1
assert contents[0]["content"] == "test content"
assert contents[0]["context"] == "test_context"
assert "event_date" in contents[0]
assert contents[0]["event_date"] == datetime(2024, 1, 15, 10, 30, 0, tzinfo=timezone.utc)
@pytest.mark.asyncio
async def test_local_mcp_server_retain_with_invalid_timestamp(mock_memory):
"""Test that retain rejects invalid timestamp format."""
from hindsight_api.mcp_local import create_local_mcp_server
mcp_server = create_local_mcp_server("test-bank", memory=mock_memory)
tools = mcp_server._tool_manager._tools
retain_tool = tools["retain"]
# Call retain with invalid timestamp
result = await retain_tool.fn(content="test content", timestamp="not-a-date")
assert result["status"] == "error"
assert "Invalid timestamp format" in result["message"]
# Verify retain_batch_async was NOT called
mock_memory.retain_batch_async.assert_not_called()
+63
View File
@@ -0,0 +1,63 @@
"""Tests for the shared MCP tools module."""
from datetime import datetime, timezone
import pytest
from hindsight_api.mcp_tools import build_content_dict, parse_timestamp
class TestParseTimestamp:
"""Tests for parse_timestamp function."""
def test_parse_iso_format_with_z(self):
"""Test parsing ISO format with Z suffix."""
result = parse_timestamp("2024-01-15T10:30:00Z")
assert result == datetime(2024, 1, 15, 10, 30, 0, tzinfo=timezone.utc)
def test_parse_iso_format_with_offset(self):
"""Test parsing ISO format with timezone offset."""
result = parse_timestamp("2024-01-15T10:30:00+00:00")
assert result == datetime(2024, 1, 15, 10, 30, 0, tzinfo=timezone.utc)
def test_parse_iso_format_without_tz(self):
"""Test parsing ISO format without timezone."""
result = parse_timestamp("2024-01-15T10:30:00")
assert result == datetime(2024, 1, 15, 10, 30, 0)
def test_parse_invalid_format_raises(self):
"""Test that invalid format raises ValueError."""
with pytest.raises(ValueError) as exc_info:
parse_timestamp("not-a-date")
assert "Invalid timestamp format" in str(exc_info.value)
class TestBuildContentDict:
"""Tests for build_content_dict function."""
def test_basic_content(self):
"""Test building content dict with just content and context."""
result, error = build_content_dict("test content", "test_context")
assert error is None
assert result == {"content": "test content", "context": "test_context"}
def test_with_valid_timestamp(self):
"""Test building content dict with valid timestamp."""
result, error = build_content_dict("test content", "test_context", "2024-01-15T10:30:00Z")
assert error is None
assert result["content"] == "test content"
assert result["context"] == "test_context"
assert result["event_date"] == datetime(2024, 1, 15, 10, 30, 0, tzinfo=timezone.utc)
def test_with_invalid_timestamp(self):
"""Test building content dict with invalid timestamp."""
result, error = build_content_dict("test content", "test_context", "invalid")
assert error is not None
assert "Invalid timestamp format" in error
assert result == {}
def test_with_none_timestamp(self):
"""Test building content dict with None timestamp."""
result, error = build_content_dict("test content", "test_context", None)
assert error is None
assert "event_date" not in result
File diff suppressed because it is too large Load Diff
+159
View File
@@ -275,6 +275,165 @@ async def test_retain_japanese_content(memory, request_context):
pass
@pytest.mark.asyncio
async def test_english_content_stays_english(memory, request_context):
"""
Test that English content is NOT incorrectly translated to Japanese or Chinese.
This test specifically catches the bug where the language instruction in the
CONCISE extraction prompt mentioned Japanese/Chinese explicitly, which primed
the LLM to sometimes output facts in those languages even for English input.
See: https://github.com/vectorize-io/hindsight/issues/181
"""
bank_id = f"test_english_retain_{datetime.now(timezone.utc).timestamp()}"
try:
# English content about a developer
english_content = """
John Smith is a software engineer at TechCorp in Seattle.
He specializes in machine learning and has been working on
recommendation systems for the past three years.
Last month, he launched a new feature that improved click-through rates by 25%.
He prefers working in Python and uses PyTorch for model training.
"""
unit_ids = await memory.retain_async(
bank_id=bank_id,
content=english_content,
context="Team profile",
event_date=datetime(2024, 1, 15, tzinfo=timezone.utc),
request_context=request_context,
)
logger.info(f"Retained {len(unit_ids)} facts from English content")
assert len(unit_ids) > 0, "Should have extracted facts from English content"
# Recall with English query
result = await memory.recall_async(
bank_id=bank_id,
query="Tell me about John Smith",
budget=Budget.MID,
max_tokens=1000,
fact_type=["world"],
request_context=request_context,
)
assert len(result.results) > 0, "Should recall facts about John Smith"
# Verify facts are NOT in Japanese or Chinese
for fact in result.results:
logger.info(f"Fact: {fact.text}")
# Count Japanese characters (hiragana, katakana)
japanese_chars = sum(
1 for char in fact.text
if ("\u3040" <= char <= "\u309f") or ("\u30a0" <= char <= "\u30ff")
)
# Count Chinese/CJK characters (excluding those also used in Japanese)
# Note: Kanji/CJK ideographs overlap between Chinese and Japanese
cjk_chars = sum(1 for char in fact.text if "\u4e00" <= char <= "\u9fff")
# For English input, there should be minimal CJK characters
# Allow for occasional edge cases (e.g., proper nouns) but not full translation
total_chars = len(fact.text)
cjk_ratio = cjk_chars / max(total_chars, 1)
assert cjk_ratio < 0.1, (
f"English content was incorrectly translated to CJK language! "
f"CJK ratio: {cjk_ratio:.1%}, Japanese chars: {japanese_chars}, CJK chars: {cjk_chars}. "
f"Fact: {fact.text}"
)
logger.info("English content test passed - facts stayed in English")
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_italian_content_stays_italian(memory, request_context):
"""
Test that Italian content is NOT incorrectly translated to Japanese or Chinese.
Similar to the English test, this catches the bug where non-CJK languages
could be incorrectly translated due to biased language instruction.
See: https://github.com/vectorize-io/hindsight/issues/181
"""
bank_id = f"test_italian_retain_{datetime.now(timezone.utc).timestamp()}"
try:
# Italian content about a chef
italian_content = """
Marco Rossi è uno chef italiano che lavora in un ristorante a Milano.
È specializzato nella cucina toscana e ha vinto tre premi gastronomici.
Il mese scorso ha aperto un nuovo ristorante nel centro della città.
Preferisce usare ingredienti freschi e locali per i suoi piatti.
"""
unit_ids = await memory.retain_async(
bank_id=bank_id,
content=italian_content,
context="Profilo dello chef",
event_date=datetime(2024, 1, 15, tzinfo=timezone.utc),
request_context=request_context,
)
logger.info(f"Retained {len(unit_ids)} facts from Italian content")
assert len(unit_ids) > 0, "Should have extracted facts from Italian content"
# Recall with Italian query
result = await memory.recall_async(
bank_id=bank_id,
query="Dimmi di Marco Rossi", # "Tell me about Marco Rossi"
budget=Budget.MID,
max_tokens=1000,
fact_type=["world"],
request_context=request_context,
)
assert len(result.results) > 0, "Should recall facts about Marco Rossi"
# Verify facts are NOT in Japanese or Chinese - should stay in Italian
for fact in result.results:
logger.info(f"Fact: {fact.text}")
# Count CJK characters
cjk_chars = sum(1 for char in fact.text if "\u4e00" <= char <= "\u9fff")
japanese_chars = sum(
1 for char in fact.text
if ("\u3040" <= char <= "\u309f") or ("\u30a0" <= char <= "\u30ff")
)
total_chars = len(fact.text)
cjk_ratio = (cjk_chars + japanese_chars) / max(total_chars, 1)
assert cjk_ratio < 0.1, (
f"Italian content was incorrectly translated to CJK language! "
f"CJK ratio: {cjk_ratio:.1%}. Fact: {fact.text}"
)
# Verify facts contain Italian words (basic sanity check)
all_text = " ".join(f.text for f in result.results).lower()
italian_indicators = ["marco", "rossi", "chef", "ristorante", "milano", "cucina", "italiano", "italiana"]
has_italian = any(word in all_text for word in italian_indicators)
# Allow English translation as acceptable (not ideal but not the bug)
english_indicators = ["chef", "restaurant", "milan", "italian", "cooking"]
has_english = any(word in all_text for word in english_indicators)
assert has_italian or has_english, (
f"Expected facts to be in Italian or English, but got neither. Facts: {all_text}"
)
logger.info("Italian content test passed - facts not translated to CJK")
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_mixed_language_entities(memory, request_context):
"""
+15 -4
View File
@@ -8,9 +8,20 @@ populated from the summary for backwards compatibility.
import pytest
from hindsight_api.engine.memory_engine import Budget
from hindsight_api import RequestContext
from hindsight_api.config import get_config
from datetime import datetime, timezone
@pytest.fixture
def disable_observations():
"""Disable observations for a specific test."""
config = get_config()
original_value = config.enable_observations
config.enable_observations = False
yield
config.enable_observations = original_value
@pytest.mark.asyncio
async def test_entity_extraction_on_retain(memory, request_context):
"""
@@ -370,12 +381,12 @@ async def test_get_entity_state(memory, request_context):
@pytest.mark.asyncio
async def test_observation_fact_type_in_database(memory, request_context):
async def test_observation_fact_type_in_database(memory, request_context, disable_observations):
"""
Test that observations are NOT stored as memory_units with fact_type='observation'.
Test that when observations are disabled, no observation records are created.
NOTE: Observations are now handled via mental models, not as memory_units
or entity summaries.
When enable_observations=False, consolidation does not run and no
memory_units with fact_type='observation' should exist.
"""
bank_id = f"test_obs_db_{datetime.now(timezone.utc).timestamp()}"
File diff suppressed because it is too large Load Diff
+448
View File
@@ -0,0 +1,448 @@
"""Tests for mental models (formerly reflections), observations, and learnings functionality."""
import uuid
import pytest
import pytest_asyncio
import httpx
from hindsight_api.api import create_app
from hindsight_api.engine.memory_engine import MemoryEngine
@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"test_mental_models_{uuid.uuid4().hex[:8]}"
class TestMentalModelsCRUD:
"""Test mental models CRUD operations via memory engine."""
@pytest.mark.asyncio
async def test_create_and_get_mental_model(self, memory: MemoryEngine, request_context):
"""Test creating and retrieving a mental model."""
bank_id = f"test-mental-model-{uuid.uuid4().hex[:8]}"
# Create the bank first
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
# Create a mental model
mental_model = await memory.create_mental_model(
bank_id=bank_id,
name="Team Preferences",
source_query="What are the team's communication preferences?",
content="The team prefers async communication via Slack",
tags=["team"],
request_context=request_context,
)
assert mental_model["name"] == "Team Preferences"
assert mental_model["source_query"] == "What are the team's communication preferences?"
assert mental_model["content"] == "The team prefers async communication via Slack"
assert mental_model["tags"] == ["team"]
assert "id" in mental_model
# Get the mental model
fetched = await memory.get_mental_model(
bank_id=bank_id,
mental_model_id=mental_model["id"],
request_context=request_context,
)
assert fetched["id"] == mental_model["id"]
assert fetched["name"] == "Team Preferences"
# Cleanup
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_list_mental_models(self, memory: MemoryEngine, request_context):
"""Test listing mental models with filters."""
bank_id = f"test-mental-model-list-{uuid.uuid4().hex[:8]}"
# Create the bank first
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
# Create multiple mental models
await memory.create_mental_model(
bank_id=bank_id,
name="Mental Model 1",
source_query="Query 1",
content="Content 1",
tags=["tag1"],
request_context=request_context,
)
await memory.create_mental_model(
bank_id=bank_id,
name="Mental Model 2",
source_query="Query 2",
content="Content 2",
tags=["tag2"],
request_context=request_context,
)
# List all
all_mental_models = await memory.list_mental_models(
bank_id=bank_id,
request_context=request_context,
)
assert len(all_mental_models) == 2
# List with tag filter
tag1_mental_models = await memory.list_mental_models(
bank_id=bank_id,
tags=["tag1"],
request_context=request_context,
)
assert len(tag1_mental_models) == 1
# Cleanup
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_update_mental_model(self, memory: MemoryEngine, request_context):
"""Test updating a mental model."""
bank_id = f"test-mental-model-update-{uuid.uuid4().hex[:8]}"
# Create the bank first
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
# Create a mental model
mental_model = await memory.create_mental_model(
bank_id=bank_id,
name="Original Name",
source_query="Original Query",
content="Original Content",
request_context=request_context,
)
# Update the mental model
updated = await memory.update_mental_model(
bank_id=bank_id,
mental_model_id=mental_model["id"],
name="Updated Name",
content="Updated Content",
request_context=request_context,
)
assert updated["name"] == "Updated Name"
assert updated["content"] == "Updated Content"
# Cleanup
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_delete_mental_model(self, memory: MemoryEngine, request_context):
"""Test deleting a mental model."""
bank_id = f"test-mental-model-delete-{uuid.uuid4().hex[:8]}"
# Create the bank first
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
# Create a mental model
mental_model = await memory.create_mental_model(
bank_id=bank_id,
name="To Delete",
source_query="Query",
content="Content",
request_context=request_context,
)
# Delete the mental model
await memory.delete_mental_model(
bank_id=bank_id,
mental_model_id=mental_model["id"],
request_context=request_context,
)
# Verify deletion - should return None
fetched = await memory.get_mental_model(
bank_id=bank_id,
mental_model_id=mental_model["id"],
request_context=request_context,
)
assert fetched is None
# Cleanup
await memory.delete_bank(bank_id, request_context=request_context)
class TestObservationsAPI:
"""Test observations API endpoints.
NOTE: Observations are now stored in memory_units with fact_type='observation'
and accessed via recall with fact_type=["observation"]. The old /observations
endpoint was removed. These tests are skipped.
"""
@pytest.mark.skip(reason="Observations endpoint removed - use recall with fact_type=['observation']")
@pytest.mark.asyncio
async def test_list_observations_empty(self, api_client, test_bank_id):
"""Test listing observations when none exist."""
pass
@pytest.mark.skip(reason="Observations endpoint removed - use recall with fact_type=['observation']")
@pytest.mark.asyncio
async def test_get_observation_not_found(self, api_client, test_bank_id):
"""Test getting a non-existent observation."""
pass
class TestMentalModelsAPI:
"""Test mental models API endpoints."""
@pytest.mark.asyncio
async def test_mental_models_api_crud(self, api_client, test_bank_id):
"""Test full CRUD cycle through API."""
import asyncio
# Create bank first via profile endpoint
await api_client.get(f"/v1/default/banks/{test_bank_id}/profile")
# Create a mental model (async operation)
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/mental-models",
json={
"name": "API Test Mental Model",
"source_query": "What is the API test about?",
"content": "This is an API test mental model",
"tags": ["api-test"],
},
)
assert response.status_code == 200
create_result = response.json()
assert "operation_id" in create_result
operation_id = create_result["operation_id"]
# Wait for the async operation to complete
for _ in range(30): # Wait up to 30 seconds
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/operations/{operation_id}")
if response.status_code == 200:
op_status = response.json()
if op_status.get("status") == "completed":
break
await asyncio.sleep(1)
# List mental models to get the created mental model
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/mental-models")
assert response.status_code == 200
mental_models = response.json()["items"]
assert len(mental_models) >= 1
# Find our mental model
mental_model = next((m for m in mental_models if m["name"] == "API Test Mental Model"), None)
assert mental_model is not None, f"Mental model not found. Items: {mental_models}"
mental_model_id = mental_model["id"]
# Get the mental model
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/mental-models/{mental_model_id}")
assert response.status_code == 200
assert response.json()["name"] == "API Test Mental Model"
# Update the mental model
response = await api_client.patch(
f"/v1/default/banks/{test_bank_id}/mental-models/{mental_model_id}",
json={"name": "Updated API Test Mental Model"},
)
assert response.status_code == 200
assert response.json()["name"] == "Updated API Test Mental Model"
# Delete the mental model
response = await api_client.delete(f"/v1/default/banks/{test_bank_id}/mental-models/{mental_model_id}")
assert response.status_code == 200
# Verify deletion
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/mental-models/{mental_model_id}")
assert response.status_code == 404
# Cleanup
await api_client.delete(f"/v1/default/banks/{test_bank_id}")
class TestRecallWithObservationsAndMentalModels:
"""Test recall integration with observations and mental models."""
@pytest.mark.asyncio
async def test_recall_includes_observations(self, api_client, test_bank_id):
"""Test that recall can include observations in the response."""
# Create bank first via profile endpoint
await api_client.get(f"/v1/default/banks/{test_bank_id}/profile")
# Note: Observations are auto-created via consolidation, not manually
# This test just verifies the include parameter works
# Recall with observations included
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={
"query": "What is machine learning?",
"include": {
"observations": {"max_results": 5},
},
},
)
assert response.status_code == 200
result = response.json()
# Should have observations field in response (may be empty)
assert "observations" in result or result.get("observations") is None
# Cleanup
await api_client.delete(f"/v1/default/banks/{test_bank_id}")
@pytest.mark.asyncio
async def test_recall_includes_mental_models(self, api_client, test_bank_id):
"""Test that recall can include mental models in the response."""
# Create bank first via profile endpoint
await api_client.get(f"/v1/default/banks/{test_bank_id}/profile")
# Create a mental model first
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/mental-models",
json={
"name": "AI Overview",
"source_query": "What is AI?",
"content": "Artificial intelligence is the simulation of human intelligence",
"tags": [],
},
)
assert response.status_code == 200
# Recall with mental models included
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={
"query": "What is artificial intelligence?",
"include": {
"mental_models": {"max_results": 5},
},
},
)
assert response.status_code == 200
result = response.json()
# Should have mental_models in response (may be empty if embedding not generated yet)
assert "mental_models" in result or result.get("mental_models") is None
# Cleanup
await api_client.delete(f"/v1/default/banks/{test_bank_id}")
@pytest.mark.asyncio
async def test_recall_without_observations_by_default(self, api_client, test_bank_id):
"""Test that recall does not include observations by default."""
# Create bank first via profile endpoint
await api_client.get(f"/v1/default/banks/{test_bank_id}/profile")
# Recall without specifying observations
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={
"query": "Test query",
},
)
assert response.status_code == 200
result = response.json()
# Observations should not be in response
assert result.get("observations") is None
# Cleanup
await api_client.delete(f"/v1/default/banks/{test_bank_id}")
class TestReflectUsesMentalModels:
"""Test that reflect searches and uses mental models when available."""
@pytest.mark.asyncio
async def test_reflect_searches_mental_models_when_available(self, memory: MemoryEngine, request_context):
"""Test that reflect uses search_mental_models when the bank has mental models.
Given:
- A bank with a mental model about "team collaboration"
Expected:
- Reflect should call search_mental_models tool
- The mental model content should influence the response
"""
bank_id = f"test-reflect-mm-{uuid.uuid4().hex[:8]}"
# Create the bank
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
# Create a mental model about team collaboration
mental_model = await memory.create_mental_model(
bank_id=bank_id,
mental_model_id=str(uuid.uuid4()),
name="Team Collaboration Practices",
source_query="How does the team collaborate?",
content="The team uses async communication via Slack and holds daily standups at 9am. "
"Code reviews are required before merging. The team values documentation and "
"prefers written communication for complex decisions.",
tags=["team"],
request_context=request_context,
)
# Run reflect with a query about team collaboration
result = await memory.reflect_async(
bank_id=bank_id,
query="How does the team work together?",
request_context=request_context,
)
# Check that mental models were searched
tool_calls = result.tool_trace
search_mm_calls = [tc for tc in tool_calls if tc.tool == "search_mental_models"]
assert len(search_mm_calls) > 0, (
f"Expected search_mental_models to be called when bank has mental models. "
f"Tool calls: {[tc.tool for tc in tool_calls]}"
)
# Check that the reason field is populated for debugging
for tc in search_mm_calls:
assert tc.reason is not None, "Tool call should have a reason for debugging"
# The response should mention concepts from the mental model
response_text = result.text.lower()
has_relevant_content = any(
keyword in response_text
for keyword in ["slack", "async", "standup", "code review", "documentation", "communication"]
)
assert has_relevant_content, (
f"Expected response to reference mental model content. Got: {result.text[:500]}"
)
# Cleanup
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_reflect_tool_trace_includes_reason(self, memory: MemoryEngine, request_context):
"""Test that tool traces include the reason field for debugging."""
bank_id = f"test-reflect-reason-{uuid.uuid4().hex[:8]}"
# Create the bank
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
# Run reflect - it should use observations or recall
result = await memory.reflect_async(
bank_id=bank_id,
query="What is the weather like?",
request_context=request_context,
)
# All tool calls should have a reason
for tc in result.tool_trace:
if tc.tool != "done": # done doesn't need a reason
assert tc.reason is not None, f"Tool {tc.tool} should have a reason for debugging"
# Cleanup
await memory.delete_bank(bank_id, request_context=request_context)
+138
View File
@@ -279,6 +279,7 @@ async def test_event_date_storage(memory, request_context):
@pytest.mark.asyncio
@pytest.mark.xfail(reason="LLM date extraction from content is non-deterministic", strict=False)
async def test_temporal_ordering(memory, request_context):
"""
Test that facts can be stored and retrieved with correct temporal ordering.
@@ -2058,3 +2059,140 @@ async def test_user_provided_entities(memory, request_context):
finally:
await memory.delete_bank(bank_id, request_context=request_context)
def test_recall_result_model_empty_construction():
"""
Test that RecallResultModel can be constructed with empty results.
This is a regression test for the bug where constructing an empty RecallResultModel
would cause an UnboundLocalError because RecallResult was imported as RecallResultModel
but the code mistakenly used the wrong name.
The fix ensures RecallResultModel is used consistently throughout memory_engine.py.
"""
from hindsight_api.engine.response_models import RecallResult
# This should not raise any errors
result = RecallResult(results=[], entities={}, chunks={})
assert result is not None, "Should create a result object"
assert result.results == [], "Should have empty results"
assert result.entities == {}, "Should have empty entities"
assert result.chunks == {}, "Should have empty chunks"
logger.info("✓ RecallResult empty construction works correctly")
@pytest.mark.asyncio
async def test_custom_extraction_mode():
"""
Test that custom extraction mode uses custom guidelines from env variable.
This test verifies that when HINDSIGHT_API_RETAIN_EXTRACTION_MODE=custom and
HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS is set, the fact extraction uses the
custom guidelines while keeping structural parts intact.
"""
import os
from hindsight_api import LLMConfig
from hindsight_api.engine.retain.fact_extraction import extract_facts_from_text
from hindsight_api.config import clear_config_cache
# Save original env vars
original_mode = os.getenv("HINDSIGHT_API_RETAIN_EXTRACTION_MODE")
original_instructions = os.getenv("HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS")
try:
# Set custom extraction mode with challenging language-specific guidelines
os.environ["HINDSIGHT_API_RETAIN_EXTRACTION_MODE"] = "custom"
os.environ["HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS"] = """ONLY extract facts that are in ITALIAN language.
DO NOT extract:
Facts in English
Facts in any other language besides Italian
If the text contains both Italian and English content, extract ONLY the Italian facts."""
# Clear config cache to pick up new env vars
clear_config_cache()
# Test content with BOTH Italian (should extract) and English (should NOT extract) facts
# This is a much harder test than filtering greetings
text = """
The team discussed the new architecture. We will use microservices.
Il database PostgreSQL ha ridotto la latenza delle query del 60%.
Alice ha suggerito di usare il connection pooling per migliorare le prestazioni.
Bob mentioned that the API endpoint is ready for testing.
The deployment pipeline has been updated to use Kubernetes.
Marco ha completato la revisione del codice e ha approvato le modifiche.
Il sistema di autenticazione è stato migrato a OAuth 2.0.
"""
llm_config = LLMConfig.for_memory()
facts, _, _ = await extract_facts_from_text(
text=text,
event_date=datetime(2024, 1, 15, tzinfo=timezone.utc),
context="team meeting notes",
llm_config=llm_config,
agent_name="TestUser"
)
logger.info(f"\nExtracted {len(facts)} facts with custom mode (Italian only):")
for i, fact in enumerate(facts):
logger.info(f" {i+1}. {fact.fact}")
assert len(facts) > 0, "Should extract at least one Italian fact"
# All facts text
all_facts_text = " ".join([f.fact for f in facts])
# Should HAVE Italian content
italian_keywords = ["postgresql", "latenza", "query", "alice", "connection pooling", "prestazioni",
"marco", "revisione", "codice", "autenticazione", "oauth"]
has_italian = any(keyword in all_facts_text.lower() for keyword in italian_keywords)
assert has_italian, f"Should extract Italian facts. Got: {all_facts_text}"
# Should NOT have English-only content
# These are facts that appear ONLY in English sections
english_only_keywords = ["microservices", "bob", "api endpoint", "testing", "deployment pipeline", "kubernetes"]
# Check if facts contain English-only content (this would be wrong)
facts_lower = all_facts_text.lower()
found_english_only = [kw for kw in english_only_keywords if kw in facts_lower]
if found_english_only:
logger.warning(f"⚠ Found English-only keywords in facts: {found_english_only}")
logger.warning(f" Facts: {all_facts_text}")
logger.warning(f" This may indicate the LLM is not strictly following language-specific custom guidelines")
# Log but don't fail - LLM behavior can vary
else:
logger.info("✓ Successfully extracted only Italian facts, ignored English facts")
# At least verify we have some Italian indicators
italian_indicators = ["latenza", "prestazioni", "revisione", "codice", "autenticazione"]
italian_count = sum(1 for ind in italian_indicators if ind in facts_lower)
assert italian_count >= 1, \
f"Should extract facts with Italian words. Found {italian_count} Italian indicators in: {all_facts_text}"
logger.info("✓ Custom extraction mode works with language-specific guidelines")
logger.info(f"✓ Extracted {len(facts)} Italian facts, found {italian_count} Italian indicators")
finally:
# Restore original env vars
if original_mode is not None:
os.environ["HINDSIGHT_API_RETAIN_EXTRACTION_MODE"] = original_mode
else:
os.environ.pop("HINDSIGHT_API_RETAIN_EXTRACTION_MODE", None)
if original_instructions is not None:
os.environ["HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS"] = original_instructions
else:
os.environ.pop("HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS", None)
# Clear cache again to restore original config
clear_config_cache()
+6 -1
View File
@@ -11,8 +11,8 @@ import uuid
import pytest
import pytest_asyncio
from hindsight_api.extensions import RequestContext, TenantContext, TenantExtension
from hindsight_api.engine.memory_engine import _current_schema, fq_table
from hindsight_api.extensions import RequestContext, TenantContext, TenantExtension
from hindsight_api.migrations import run_migrations
@@ -52,6 +52,11 @@ class MultiSchemaTestTenantExtension(TenantExtension):
raise AuthenticationError(f"Unknown API key: {context.api_key}")
async def list_tenants(self) -> list:
from hindsight_api.extensions.tenant import Tenant
return [Tenant(schema=schema) for schema in self.valid_schemas]
async def drop_schema(conn, schema_name: str) -> None:
"""Drop a schema and all its contents."""
+10 -5
View File
@@ -249,14 +249,14 @@ class TestServerModuleExtensionLoading:
# Mock extensions for testing
from hindsight_api.extensions import (
TenantExtension,
TenantContext,
RequestContext,
OperationValidatorExtension,
ValidationResult,
RetainContext,
RecallContext,
ReflectContext,
RequestContext,
RetainContext,
TenantContext,
TenantExtension,
ValidationResult,
)
@@ -270,6 +270,11 @@ class MockTenantExtension(TenantExtension):
async def authenticate(self, request_context: RequestContext) -> TenantContext:
return TenantContext(schema_name="public")
async def list_tenants(self) -> list:
from hindsight_api.extensions.tenant import Tenant
return [Tenant(schema="public")]
def set_context(self, context) -> None:
self._context_set = True

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