Compare commits

...
158 Commits
Author SHA1 Message Date
Nicolò Boschi 45ed02cd6a fix 2026-01-16 10:49:56 +01:00
Nicolò Boschi 2921fa091b fix 2026-01-16 10:21:05 +01:00
Nicolò Boschi aafa7ae17c fix ci 2026-01-16 07:52:38 +01:00
Nicolò Boschi caeab8ac44 more 2026-01-15 18:34:17 +01:00
Nicolò Boschi 42cc098a7c new style 2026-01-14 18:57:16 +01:00
Nicolò Boschi 2871ddc34c reflect agent 2026-01-14 09:29:36 +01:00
Nicolò Boschi 987b47e6a1 agentic 2026-01-14 09:29:36 +01:00
Nicolò Boschi 25fb4b9273 agentic 2026-01-14 09:29:36 +01:00
Nicolò Boschi f72c9f03fb fix db patch 2026-01-14 09:26:01 +01:00
Nicolò Boschi 517eee4eda DRAFT: refactor entity observations 2026-01-14 09:25:27 +01:00
Nicolò Boschi e19c6b9252 mental models 2026-01-14 09:25:27 +01:00
Chris Bartholomew e64d3634a9 feat: add cloud mode to skill installer for team memory sharing (#158)
* doc: update expired Slack invite link

* feat: add cloud mode to skill installer for team memory sharing

Adds support for Hindsight Cloud in the skill installer, enabling teams
to share memories about a codebase. Changes include:

- Add `--mode cloud` option to get-skill installer
- Install hindsight CLI binary for cloud mode (via get-cli)
- Configure ~/.hindsight/config with API URL and key
- Generate cloud-specific SKILL.md with team-aware guidance
- Distinguish between project conventions and individual preferences
- Update skills.md documentation with cloud setup instructions

Cloud mode workflow:
1. Team admin creates a bank in Hindsight Cloud
2. Each developer runs: curl ... | bash -s -- --mode cloud
3. All team members share the same memory bank
4. Knowledge retained by one member benefits everyone
2026-01-14 09:05:40 +01:00
Chris Bartholomew 70ce979fbe doc: update expired Slack invite link (#157) 2026-01-13 16:57:23 -05:00
Nicolò Boschi de132501c6 doc: changelog for 0.3.0 (#156) 2026-01-13 19:09:13 +01:00
Nicolò Boschi a75dcfebf5 Release v0.3.0
- Update version to 0.3.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-13 18:43:33 +01:00
Nicolò Boschi 20c8f8b06a feat: add memory tags (#152)
* feat: add memory tags

* feat: add memory tags

* support tags

* support tags
2026-01-13 18:28:40 +01:00
Chris Bartholomew f5f3fca4ad Fix: Load extensions in server.py for multi-worker deployments (#155)
* Fix: Load extensions in server.py for multi-worker deployments

When running with multiple workers (--workers 2), uvicorn uses
`hindsight_api.server:app` import string instead of passing an app
object. The server.py module was not loading tenant/operation validator
extensions, causing authentication bypass in production.

This fix:
- Adds extension loading to server.py matching main.py behavior
- Sets extension context on tenant extension for schema provisioning
- Adds comprehensive unit tests for server.py extension loading

The tests specifically verify:
- TENANT extension is loaded when HINDSIGHT_API_TENANT_EXTENSION is set
- OPERATION_VALIDATOR is loaded when configured
- Extensions are passed to MemoryEngine constructor
- Extension context is set on tenant extension
- Server works correctly without extensions configured

* Add unit tests for main.py extension loading (single-worker path)
2026-01-13 17:55:33 +01:00
Nicolò Boschi d47c8a28cc feat: support litellm gateway (#154) 2026-01-13 16:55:28 +01:00
Nicolò Boschi 1ffc2a418c feat: add tenant to metrics labels (#151) 2026-01-13 15:31:58 +01:00
Nicolò Boschi fa53917c63 feat: support custom url for openai embeddings & cohere (#150)
* feat: support custom url for openai embeddings & cohere

* feat: support custom url for openai embeddings & cohere
2026-01-13 14:01:44 +01:00
Nicolò Boschi 59913086be fix: batch queries on recall (#149)
* fix: batch queries on recall

* fix: batch queries on recall
2026-01-13 13:20:22 +01:00
Nicolò Boschi 7935b0accd fix: improve mpfp retrieval (#146)
* fix: improve mpfp retrieval

* fix: improve mpfp retrieval

* fix: improve embeddings service performances

* fix: improve embeddings service performances

* fix: improve embeddings service performances

* fix: improve embeddings service performances
2026-01-12 18:58:05 +01:00
Nicolò Boschi 26bf5714cd fix: entities list only show 100 entities (#142)
* fix: entities list only show 100 entities

* fix: update Rust CLI for entities pagination API changes
2026-01-12 18:50:53 +01:00
Nicolò Boschi 6232e690fc fix: improve graph retrieval on large memory banks (#141) 2026-01-09 16:43:31 +01:00
Nicolò Boschi 4135a6cee5 ci: frozen uv sync (#138)
* ci: frozen uv sync

* fix: add missing authorization parameter to get_agent_stats in CLI

The generated Rust client was updated with an authorization header
parameter for get_agent_stats, but the CLI code wasn't updated.
2026-01-09 16:43:00 +01:00
Nicolò Boschi eb2702bcba misc: performance improvements (#140)
* misc: performance improvements

* misc: performance improvements

* misc: performance improvements
2026-01-09 14:47:20 +01:00
Nicolò Boschi 0d0abaaa9f fix(typescript-client): Add error handling to all API methods (#139)
Previously, most methods in HindsightClient would silently return
undefined when API calls failed (e.g., connection refused). Only
the `recall` method had proper error checking.

This change adds a `validateResponse` helper method and applies it
consistently to all API methods:
- retain
- retainBatch
- recall
- reflect
- listMemories
- createBank
- getBankProfile

Now all methods properly throw an error with details when the API
request fails, instead of returning undefined.
2026-01-09 14:25:10 +01:00
Nicolò Boschi a6798f7e2a fix: improve tei client parameters (#137)
* fix: improve tei client parameters

* fix: improve tei client parameters

* fix: improve tei client parameters
2026-01-09 11:31:22 +01:00
Nicolò Boschi fb31a35a86 feat: retain modes (#136)
* feat: retain modes

* fix db patch
2026-01-09 11:30:36 +01:00
Nicolò Boschi ba99b4422a fix: misc perf improvements (#133)
* fix: misc perf improvements

* more tests

* fix test

* fix: update test files for new extract_facts_from_text signature

- Replace test_fact_extraction_token_analysis with test_fact_extraction_basic_analysis
  using inline sample content instead of external file
- Update test_fact_extraction_output_ratio.py to unpack 3 return values
  (facts, chunks, usage) instead of 2

* fix: make temporal tests more flexible for LLM variation

- test_temporal_absolute_conversion: check occurred_start field instead of
  requiring specific text in facts
- test_date_field_calculation_yesterday: make assertions conditional on
  having temporal data, add more content for better extraction
- test_temporal_ordering: reduce minimum required facts from 3 to 2
2026-01-08 22:49:04 +01:00
Chris Bartholomew 6fe93140a7 Fix embedding dimension for tenant schemas (#135)
Call ensure_embedding_dimension after running migrations for tenant
schemas. This ensures the embedding column dimension matches the
model's dimension, which may differ from the default 384 dimensions
used in the initial migration.

Without this fix, using embedding providers with different dimensions
(e.g., Cohere's embed-english-v3.0 with 1024 dims) would fail with
"expected 384 dimensions, not 1024" errors on tenant schemas.
2026-01-08 22:48:25 +01:00
Chris Bartholomew d6ff191198 Fix stats endpoint missing tenant authentication (#134)
The /v1/default/banks/{bank_id}/stats endpoint was missing the
request_context parameter and tenant authentication call, causing
it to query the public schema instead of the tenant's schema.

This resulted in stats always returning zeros for multi-tenant
deployments since the data lives in tenant-specific schemas.

Added request_context dependency and _authenticate_tenant() call
to properly set the tenant schema before querying stats.
2026-01-08 20:38:35 +01:00
Nicolò Boschi 3bb6a38b5c ci: fix flak tests (#131) 2026-01-08 18:44:30 +01:00
Nicolò Boschi b5df8657e8 chore: add flag to not include ml libs in docker image (#130) 2026-01-08 18:22:42 +01:00
Nicolò Boschi 1dacd0e904 feat: add operation_id to retain response (#129) 2026-01-08 17:41:51 +01:00
Derek Bouius 4b82d2d7ec feat: delete memory bank (#127)
* expose the delete API

* add deleteBank

* Add a button and confirmation dialog to delete a memory bank

* commit lint changes

* add CI test for delete bank

* revert alembic lint changes due to version differences

* revert alembic lint changes

* fix the delete bank test

* account for ruff lint third party alembic
2026-01-08 17:41:42 +01:00
Nicolò Boschi 33fac2c5e2 feat: add configs for database connection (#128) 2026-01-08 16:37:08 +01:00
Nicolò Boschi 49e233cdb7 fix: duplicated causal relationships and token optimization (#126)
* fix: duplicated causal relationships and token optimization

* doc

* doc
2026-01-08 14:43:48 +01:00
Nicolò Boschi e6709d541f feat: support different provider/models per operation (#125)
* feat: support different provider/models per operation

* fix tests
2026-01-08 14:02:57 +01:00
Nicolò Boschi 9fd567984c fix(mcp): add back bank list and create_bank tools (#123)
* fix(mcp): add back bank list and create_bank tools

* fix tests

* fix tests
2026-01-08 14:02:28 +01:00
Nicolò Boschi c65c6a9dc0 feat: support for multilingual content (#124)
* feat: support for multilingual content

* feat: support for multilingual content
2026-01-08 12:14:41 +01:00
Nicolò Boschi 4de0730c40 feat: support cohere as embeddings and reranker (#122) 2026-01-08 11:41:15 +01:00
Nicolò Boschi 5e1f13e4f2 feat: add metrics for llm call latency (#120)
* feat: add metrics for llm call latency

* feat: add metrics for llm call latency

* fix
2026-01-08 11:40:34 +01:00
Nicolò Boschi 67c1a4295f fix: ui shows only 1000 memories (#121)
* fix: ui shows only 1000 memories

* fix: ui shows only 1000 memories
2026-01-08 11:22:10 +01:00
37fc7fb8bd feat(mcp): add async_processing parameter to retain tool (#95)
* feat(mcp): add async_processing parameter to retain tool

Add async_processing parameter (default: True) to the MCP retain tool
to allow non-blocking memory storage. When True, memories are queued
for background processing and the tool returns immediately. When False,
the tool waits for completion before returning.

This matches the async behavior available in the HTTP API.

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

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

* feat(mcp): add list_memories and reflect tools

Add two missing MCP tools to achieve feature parity with HTTP API:

- list_memories: browse memories with pagination and full-text search
  (equivalent to GET /memories/list)
- reflect: LLM-based reasoning over memories with disposition awareness
  (equivalent to POST /reflect)

Both tools follow the existing pattern with JSON string responses
and proper error handling.

* docs: improve CLAUDE.md with detailed architecture info

- Add memory types explanation (world, experience, opinion, observation)
- Document retain/ and search/ submodule structure
- Add commands for single test run, ruff format, ty type checking
- Note MCP server implementation in API layer
- Add optional environment variables section
- Clarify conventions (no Python files at root, npm workspaces)

* chore: add .mcp.json and .osgrep to gitignore

These are user-specific development tool configs that should not be committed.

* changes

* refactor(mcp): remove list_memories tool

The list_memories endpoint is for debugging/exploration, not agent use.
Agents should use recall for semantic search instead.

Feedback from maintainer: "this tool is misleading for the agent,
it should use recall, the list method is mostly for debugging and
exploration, not for real usage"

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

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

* refactor(mcp): remove list_banks and create_bank tools

These admin/orchestration tools are not needed for typical agent usage.
Agents work with a single configured bank via X-Bank-Id header.

MCP now exposes only core memory operations:
- retain: store memories
- recall: semantic search
- reflect: LLM reasoning over memories

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

---------

Co-authored-by: Anton Evseev <[email protected]>
Co-authored-by: Claude Opus 4.5 <[email protected]>
2026-01-08 11:17:02 +01:00
Alexander Pinsker 29a542dc23 feat: Add per-request LLM token usage metrics (#117)
* feat: Record LLM token metrics via Prometheus

Wire up the existing token metrics infrastructure to actually record
token usage from LLM calls. The MetricsCollector already had
record_tokens() method and Prometheus counters (hindsight.tokens.input,
hindsight.tokens.output), but they were never being populated.

Changes:
- Import get_metrics_collector in llm_wrapper.py
- Call record_tokens() after successful LLM calls for:
  - OpenAI/Groq (using response.usage.prompt_tokens, completion_tokens)
  - Anthropic (using response.usage.input_tokens, output_tokens)
  - Gemini (using response.usage_metadata.prompt_token_count, candidates_token_count)
- Add test file to verify token metrics are recorded

Note: Ollama's native API doesn't return token usage, so metrics
are not recorded for that provider.

The token metrics will now be available via /metrics endpoint:
- hindsight_tokens_input_total
- hindsight_tokens_output_total

* feat: add per-request token usage tracking to retain and reflect endpoints

- Add TokenUsage model with input_tokens, output_tokens, total_tokens
- Return usage metrics in retain response (sync operations only)
- Return usage metrics in reflect response
- Update Python, TypeScript, and Rust clients
- Add API documentation for usage fields
- Add changelog entry
2026-01-08 10:36:58 +01:00
Anatolii LapytskyiandAnatolii Lapytskyi ecc1f31996 feat(helm): add existingSecret support (#119)
* feat(helm): add existingSecret support

Allow users to reference a pre-existing Kubernetes Secret instead of
having the chart create one. This enables better secret management
through tools like External Secrets Operator or sealed-secrets.

Usage:
```yaml
existingSecret: "my-pre-created-secret"
```

When existingSecret is set:
- The chart skips creating its own Secret resource
- Deployments reference the provided secret name
- Secret checksum annotation is omitted (no auto-rollout on changes)

The existing secret should contain all required keys:
- API secrets (e.g., HINDSIGHT_API_LLM_API_KEY)
- Control plane secrets
- postgres-password (if using external PostgreSQL)

* fix(helm): use envFrom for existingSecret and fix env var ordering

- Add envFrom to inject all keys from existingSecret as env vars automatically
- Fix POSTGRES_PASSWORD ordering (must be before DATABASE_URL for $(VAR) interpolation)
- Only use api.secrets/controlPlane.secrets when existingSecret is not set
- Update values.yaml documentation for existingSecret usage

---------

Co-authored-by: Anatolii Lapytskyi <[email protected]>
2026-01-08 10:36:03 +01:00
Nicolò Boschi 233bd2e5d4 feat: run db migrations offline (optionally) (#114)
* feat: run db migrations offline (optionally)

* fix
2026-01-07 15:49:51 +01:00
Nicolò Boschi b3becb6e9a fix(security): fix qs - CVE-2025-15284 (#113)
* fix(security): fix qs - CVE-2025-15284

* fix
2026-01-07 15:33:07 +01:00
Nicolò Boschi 67b273de69 feat: backup/restore (#110)
* feat: backup/restore

* feat: backup/restore

* fix
2026-01-07 11:29:50 +01:00
Nicolò Boschi 5a3090b5e5 ci: pin rust lock version (#112) 2026-01-07 11:29:41 +01:00
Nicolò Boschi 2a00df0bc0 fix: improve causal links detection (#111)
* fix: improve causal links detection

* fix: improve causal links detection
2026-01-07 11:16:24 +01:00
Nicolò Boschi 7715a5110e fix: make retain max completion tokens configurable (#109)
* fix: make retain max completion tokens configurable

* fix: make retain max completion tokens configurable
2026-01-07 10:26:42 +01:00
Chris Bartholomew c06d9b4e4f Load .env file automatically on startup (#104)
Add automatic .env file loading using python-dotenv. This searches
the current working directory and parent directories for a .env file
and loads environment variables from it.

Uses override=True so .env file values take precedence over existing
shell environment variables, which is the expected behavior when
running from a project directory.
2026-01-07 09:49:13 +01:00
Chris Bartholomew 39e3f7c528 Fix Python SDK not sending Authorization header (#106)
* Fix Python SDK not sending Authorization header

The Python SDK accepts an api_key parameter but never sends it as a
Bearer token in requests. The OpenAPI-generated Configuration class
stores the key in access_token, but auth_settings() returns an empty
dict because the OpenAPI spec doesn't define a security scheme.

This fix manually sets the Authorization header on the ApiClient,
bypassing the broken auth_settings() mechanism.

Tested against api.dev.hindsight.vectorize.io:
- Before: 401 "Authentication failed: API key required"
- After: Success

* chore: update Rust client Cargo.lock for CI verification

Run generate-clients.sh to sync Cargo.lock with current dependencies.
2026-01-07 09:46:50 +01:00
Nicolò Boschi d899d1890d fix: groq llm with free tier doesn't work (#102)
* fix: groq with free tier doens't work

* fix: groq with free tier doens't work
2026-01-05 15:10:35 +01:00
Nicolò Boschi 70de23ed85 feat: configurable embedding dimensions + OpenAI Embeddings (#101)
* feat: configurable embedding dimensions + OpenAI Embeddings

* fix tests
2026-01-05 14:43:05 +01:00
Nicolò Boschi 1984936150 Release v0.2.1
- Update version to 0.2.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-05 12:36:49 +01:00
Nicolò Boschi 4f21886a0e doc: changelog for 0.2.0 (and regenerate clients) (#99)
* doc: changelog for 0.2.0 (and regenerate clients)

* doc: changelog for 0.2.0 (and regenerate clients)

* doc: changelog for 0.2.0 (and regenerate clients)

* doc: changelog for 0.2.0 (and regenerate clients)

* doc: changelog for 0.2.0 (and regenerate clients)

* doc: changelog for 0.2.0 (and regenerate clients)

* doc: changelog for 0.2.0 (and regenerate clients)

* doc: changelog for 0.2.0 (and regenerate clients)
2026-01-05 12:36:29 +01:00
Nicolò Boschi 5e65691743 Release v0.2.0
- Update version to 0.2.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-05 11:34:52 +01:00
Nicolò Boschi 76fd052b3a misc: add mcp integration tests and increase test coverage (#98)
* misc: add mcp integration tests and increase test coverage

* misc: add mcp integration tests and increase test coverage

* misc: add mcp integration tests and increase test coverage
2026-01-05 11:16:55 +01:00
Bjorn SchliebitzandClaude Opus 4.5 6b5f593dca feat(mcp): Add multi-bank access and new MCP tools (#82)
* feat(mcp): Add multi-bank access and new MCP tools

Enables orchestrator agents to access multiple memory banks from a
single MCP connection, with new tools for bank management.

## New MCP Tools
- `reflect` - Thoughtful analysis using bank's personality and memories
- `list_banks` - Discover all available memory banks
- `create_bank` - Create new banks programmatically

## Multi-Bank Access
- Added optional `bank_id` parameter to `retain`, `recall`, `reflect`
- Allows cross-bank operations from a single MCP session
- Defaults to session bank if not specified

## Claude Code Compatibility
- Enabled `stateless_http=True` for proper Claude Code integration
- Responses now include `bank_id` for transparency

## Documentation
- Added docker-compose.example.yml with env var substitution
- Added HINDSIGHT-DOCKER.md setup guide with volume persistence docs
- Updated .gitignore to exclude local docker-compose.yml

## Use Case
Orchestrator agents can now:
- Maintain a private meta-orchestration bank
- Access shared project knowledge banks
- Query across banks for cross-context insights

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

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

* Address PR review feedback: remove docker files, improve reflect description

- Remove HINDSIGHT-DOCKER.md and docker-compose.example.yml per reviewer request
- Improve reflect tool description with clearer guidance for AI agents:
  - Added "WHEN TO USE THIS TOOL" section
  - Added "EXAMPLES OF GOOD QUERIES" with concrete use cases
  - Added "HOW IT DIFFERS FROM RECALL" to clarify when to use each tool

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

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

---------

Co-authored-by: Claude Opus 4.5 <[email protected]>
2026-01-05 10:06:52 +01:00
Phạm Gia Linh dd59bc8ef9 feat: Add user-provided entities support to retain endpoint (#91)
* feat: entities input for retain endpoint

* remove docker-compose.yml
2026-01-05 10:05:17 +01:00
csfet9andClaude Opus 4.5 eea0f27118 feat: Add local LLM improvements for reasoning models and Docker startup (#88)
* feat: Add local LLM improvements for reasoning models and Docker startup

## Reasoning Model Support
- Strip thinking tags from local LLM responses (<think>, <thinking>, <reasoning>, |startthink|/|endthink|)
- Enables Qwen3, DeepSeek, and other reasoning models to work with JSON extraction
- Non-breaking: only affects responses that contain thinking tags

## Docker Retry Start Script
- New retry-start.sh waits for dependencies before starting Hindsight
- Checks LLM Studio availability at /v1/models endpoint
- Checks database connectivity (skipped for embedded pg0)
- Configurable via HINDSIGHT_RETRY_MAX and HINDSIGHT_RETRY_INTERVAL env vars
- Prevents startup failures when LLM Studio isn't ready yet

Tested on Apple Silicon M4 Max with Qwen3 8B via LM Studio.

* refactor: make thinking token stripping opt-in via env var

* refactor: merge retry logic into start-all.sh (opt-in via HINDSIGHT_WAIT_FOR_DEPS)

* fix: resolve pg0 stale instance config in Docker build

- Remove stale pg0 instance data after pre-caching binaries to avoid
  port conflicts (was using hardcoded port 5555 from build time)
- Remove unused cache copy logic from start-all.sh
- Add database backup instructions to CLAUDE.md

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

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

---------

Co-authored-by: Claude Opus 4.5 <[email protected]>
2026-01-05 10:04:58 +01:00
Nicolò Boschi 964537f885 chore: add pre-commit setup instructions 2026-01-05 10:03:15 +01:00
Chris Latimer 1a620697b1 Feature/graph viz (#85)
* Improve graph visualization on the UI

* Fix double animation when loading the graph visualization

* Fix typescript issues

* CI test changes for temporal scenarios

* Fix typescript errors

* Fix animation issue on opinions and experiences
2026-01-02 16:27:29 +01:00
Chris Bartholomew ce45d301ce Add operation validator extension support with proper HTTP error handling (#86)
* Load operation validator extension in main entry point

Enable the operation validator extension to be loaded from environment
configuration and passed to MemoryEngine, allowing pre/post operation
hooks for usage metering, rate limiting, and audit logging.

* Fix reflect background task authentication and add internal flag

- Pass API key to background opinion storage task for proper auth
- Add internal flag to RequestContext for tracking internal operations
- Background opinion storage now authenticates correctly with tenant

* Add api_key_id to RequestContext for usage tracking

- Add api_key_id field to RequestContext to track which API key was used
- Enables per-API-key usage analytics in the metering system

* Fix HTTP error handling for authentication and validation errors

- Add status_code parameter to ValidationResult and OperationValidationError
- Convert OperationValidationError to HTTPException with proper status codes
- Fix authentication errors to return 401 instead of raising internal errors
- Re-raise HTTPException in exception handlers to prevent swallowing errors

* Fix AuthenticationError handling in memory engine

- Raise AuthenticationError from memory_engine._authenticate_tenant instead
  of HTTPException so unit tests pass
- Add AuthenticationError handling in HTTP layer to convert to 401 responses
- Fixes failing TestMemoryEngineTenantAuth tests

* Add global exception handler for AuthenticationError

Returns proper 401 status code for all authentication failures
across all endpoints, not just the ones with explicit handlers.

* Simplify exception handling: use global AuthenticationError handler

- Remove redundant individual exception handlers
- Add 'except AuthenticationError: raise' before generic Exception handlers
  to let global handler process auth errors uniformly

* Refactor background tasks to use tenant_id instead of api_key

This makes the core more generic - it passes tenant_id (which is
extension-agnostic) rather than api_key (which is cloud-specific).

- Add tenant_id field to RequestContext
- Pass tenant_id instead of api_key to background tasks
- Extensions can check internal=True with tenant_id to bypass normal auth

* Fix exception propagation: include HTTPException in re-raise

After cleanup of redundant exception handlers, 404 errors were
returning 500 because HTTPException was caught by the generic
except Exception handler. Fixed by combining AuthenticationError
and HTTPException in the re-raise pattern.
2026-01-01 20:19:52 -05:00
Nicolò Boschi d49e8201b4 feat: add max_tokens and structured output to /reflect (#74)
* feat: add structured output to /reflect

* feat: add structured output to /reflect

* imrpove

* add max_toksn

* fix rust client

* fix rust client

* fix rust client

* try fix

* try fix

* no stricts
2026-01-01 17:09:39 +01:00
Nicolò Boschi c8c7603580 feat(doc): add new config options and supported providers (#84) 2026-01-01 17:09:05 +01:00
csfet9andClaude Opus 4.5 787ed60763 feat: Add Anthropic Claude and LM Studio provider support (#36)
* feat: Add Anthropic Claude and LM Studio provider support

- Add Anthropic as LLM provider with full async support
- Add LM Studio provider for local model inference
- Fix JSON response format compatibility for local models
- Update .env.example with configuration examples
- Update docstrings with all supported providers

Tested with:
- Claude Sonnet 4 (claude-sonnet-4-20250514)
- Claude Haiku 4.5 (claude-haiku-4-5-20251001)
- Qwen 30B via LM Studio

* feat: Add dynamic timeout for local LLM providers

Add configurable timeout support for LLM API calls:
- Environment variable override via HINDSIGHT_API_LLM_TIMEOUT
- Dynamic heuristic for lmstudio/ollama: 20 mins for large models
  (30b, 33b, 34b, 65b, 70b, 72b, 8x7b, 8x22b), 5 mins for others
- Pass timeout to Anthropic, OpenAI, and local model clients

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

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

* fix: Address PR review feedback

- Remove CLAUDE.md from .gitignore (should stay in repository)
- Pass max_completion_tokens to _call_anthropic instead of hardcoding 4096

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

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

* chore: Remove deleted AI assistant files from .gitignore

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

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

* docs: Add CLAUDE.md for Claude Code integration

Provides project context and development commands for AI-assisted coding.

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

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

* chore: Include local dev files and sync changes

- Add docker-compose.yml for local development
- Add test_internal.py for local testing
- Sync uv.lock and llm_wrapper.py changes

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

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

* fix: Address PR review feedback for LLM provider support

- Move LLM config to config.py with HINDSIGHT_API_ prefix
  - Add HINDSIGHT_API_LLM_MAX_CONCURRENT (default: 32)
  - Add HINDSIGHT_API_LLM_TIMEOUT (default: 120s)
- Remove fragile model-size timeout heuristic
- Apply markdown JSON extraction to all providers, not just local
- Fix Anthropic markdown extraction bug (missing split)
- Change LLM request/response logs from info to debug level

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

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

* chore: Remove local dev docker-compose.yml

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

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

* chore: Add local dev docker-compose.yml

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

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

* fix: Update LM Studio port to 2222 in docker-compose

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

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

* chore: Remove obsolete version attribute from docker-compose

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

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

* fix: Remove test file and docker-compose per PR review

- Remove test_internal.py (debug file)
- Remove docker-compose.yml (to be moved to hindsight-cookbook repo)

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

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

---------

Co-authored-by: Claude Opus 4.5 <[email protected]>
2026-01-01 16:34:11 +01:00
Bjorn SchliebitzandClaude Opus 4.5 6b78f7d949 fix(mcp): Chain MCP lifespan with FastAPI app lifespan (#81)
The MCP server's lifespan was not being properly chained with the
FastAPI app's lifespan, causing the MCP server to not start/stop
correctly when mounted as a sub-application.

Changes:
- Create MCP app before FastAPI app to access its lifespan
- Chain MCP lifespan context with FastAPI's lifespan context
- Ensures MCP server lifecycle is properly managed

This fix is required for the MCP server to function correctly when
used with Claude Code and other MCP clients.

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

Co-authored-by: Claude Opus 4.5 <[email protected]>
2026-01-01 16:33:58 +01:00
Bjorn SchliebitzandClaude Opus 4.5 54e2df0baf feat(config): Add configurable observation thresholds (#83)
Allows tuning of entity observation generation via environment variables.

## New Environment Variables
- `HINDSIGHT_API_OBSERVATION_MIN_FACTS` - Minimum facts required to
  generate entity observations (default: 5)
- `HINDSIGHT_API_OBSERVATION_TOP_ENTITIES` - Maximum entities to process
  per retain batch (default: 5)

## Changes
- Added threshold configuration to HindsightConfig
- Updated memory_engine.py to use config values
- Updated observation_regeneration.py to use config values

## Use Case
Lower thresholds generate more observations (better recall, higher cost).
Higher thresholds are more selective (lower cost, may miss patterns).

Example:
```bash
# Generate more observations
docker run -e HINDSIGHT_API_OBSERVATION_MIN_FACTS=3 \
           -e HINDSIGHT_API_OBSERVATION_TOP_ENTITIES=10 ...
```

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

Co-authored-by: Claude Opus 4.5 <[email protected]>
2026-01-01 16:22:25 +01:00
Chris Latimer 967e586e01 Add model providers on README 2025-12-24 10:46:53 -07:00
Chris Bartholomew dfa7cec05b Load operation validator extension in main entry point (#72)
Enable the operation validator extension to be loaded from environment
configuration and passed to MemoryEngine, allowing pre/post operation
hooks for usage metering, rate limiting, and audit logging.
2025-12-23 15:47:26 +01:00
Nicolò Boschi 36e48a7166 doc: add skills documentation (#73)
* doc: add skills documentation

* doc: add skills documentation
2025-12-23 15:42:27 +01:00
Nicolò Boschi 786b1ecbbd Release v0.1.16
- Update version to 0.1.16 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
2025-12-23 14:12:03 +01:00
Nicolò Boschi f14f277692 fix: hindsight-embed release version 2025-12-23 14:11:49 +01:00
Nicolò Boschi c9f3657de6 0.1.15 changelog 2025-12-23 13:54:41 +01:00
Nicolò Boschi 0ae0374dc8 Release v0.1.15
- Update version to 0.1.15 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-23 13:54:14 +01:00
Nicolò Boschi f7ff32d49d feat: delete document from ui (#71)
* feat: delete document from ui

* feat: delete document from ui
2025-12-23 13:54:06 +01:00
Nicolò Boschi e06a6120a3 feat(misc): update clients types, test coverage, improve /health endpoint and add changelog (#70)
* doc: changelog and delete doc info

* others

* others

* fixes

* fixes
2025-12-23 12:49:31 +01:00
Nicolò Boschi e599346e59 Release v0.1.14
- Update version to 0.1.14 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-23 10:37:23 +01:00
Nicolò Boschi 0b352d1bfa fix: embed get-skill installer (#69) 2025-12-23 10:36:36 +01:00
Nicolò Boschi c882511f10 Release v0.1.13
- Update version to 0.1.13 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-22 22:27:49 +01:00
Nicolò Boschi 234d426499 fix(ui): timestamp is not considered in retain (#68) 2025-12-22 22:27:31 +01:00
Nicolò Boschi e6511e7d77 feat: refactor hindsight-embed architecture (#66)
* feat: refactor hindsight-embed architecture

* feat: refactor hindsight-embed architecture

* refactor deamin

* refactor deamin

* refactor deamin

* refactor deamin
2025-12-22 22:02:40 +01:00
Chris Bartholomew 904ea4de24 fix: propagate exceptions from task handlers to enable retry logic (#65)
Task handlers were swallowing exceptions, causing operations to be
marked as completed even when they failed. This prevented the retry
logic in execute_task() from working and led to accumulation of
pending operations that never completed.

Fixed handlers:
- _handle_batch_retain: remove try/except wrapper
- _handle_access_count_update: remove try/except wrapper
- _handle_regenerate_observations: remove outer try/except, keep
  inner one for individual entity failures
2025-12-22 20:42:57 +01:00
Nicolò Boschi 6168a77846 Release v0.1.12
- Update version to 0.1.12 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-22 16:44:05 +01:00
Nicolò Boschi da44a5e839 feat: add hindsight-embed and native agentic skill (#64) 2025-12-22 16:42:11 +01:00
Nicolò Boschi 32bca12c6f fix: ollama structured support (#63)
* fix: ollama structured support

* fix: ollama structured support

* fix: ollama structured support
2025-12-22 16:35:24 +01:00
Nicolò Boschi 26850a0156 doc: add documentation for extensions (#62)
* add doc for extensions

* add doc for extensions
2025-12-22 11:58:05 +01:00
Nicolò Boschi 2a0c490c9e feat: extensions (#54) 2025-12-22 11:05:23 +01:00
cesarandreslopezandCAL a831a7b77b Improve LLM JSON parsing error handling with retry logic and detailed logging (#61)
* Improve LLM JSON parsing error handling with retry logic and detailed logging

* npm changes (packaging)

---------

Co-authored-by: CAL <[email protected]>
2025-12-22 10:44:02 +01:00
DK09876 d405b4feed ci: finalize test for the documentation code (#57)
* Fix main-methods.py: entities is a dict, use .items() and .canonical_name

* Migrate docs to use CodeSnippet components

- Convert quickstart.md, retain.md, recall.md, reflect.md, memory-banks.md to .mdx
- Use CodeSnippet to pull code from validated example scripts
- Add missing 'name' parameter to create_bank calls
- Fix main-methods.py entities iteration (dict not list)
- Remove retain-new.mdx demo file

* Migrate existing docs to match testing pattern with code snippet and add CLI tests to the CI

* Fix doc-id issue + add main-method tests

* CLI fixes

* Update openAPI json

* Fix rust build issues

* increase sleep time for Hindsight to process the document

* Added a polling sleep instead of fixed

* Delete immediately fails, so create the doc a earlier in the test to get the doc ready

* Add debug logs

* Remove debug logs
2025-12-19 12:17:59 -07:00
Nicolò Boschi b94b5cf26e fix: set max_completion_tokens to 100 in llm validation (#59) 2025-12-19 09:32:43 +01:00
Nicolò Boschi 6d820ef91b doc: add openai api compatible note 2025-12-18 16:19:30 +01:00
Nicolò Boschi cf8882a867 changelog for 0.1.11 2025-12-18 14:40:30 +01:00
Nicolò Boschi 490fccdc6f Release v0.1.11
- Update version to 0.1.11 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-18 14:09:44 +01:00
Nicolò Boschi 2948cb62d2 fix: docker image and control plane standalone build 2025-12-18 14:07:41 +01:00
Nicolò Boschi 9053a51a88 update changelog for 0.1.10 2025-12-18 13:36:48 +01:00
Nicolò Boschi f2c28cfd98 Release v0.1.10
- Update version to 0.1.10 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-18 13:10:56 +01:00
Nicolò Boschi 67fc532c43 ci: make release faster and restartable 2025-12-18 13:10:46 +01:00
Nicolò Boschi 9474f950f2 fix release process 2025-12-18 12:10:01 +01:00
Nicolò Boschi 6a0c034f5d Release v0.1.9
- Update version to 0.1.9 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-18 12:00:03 +01:00
Nicolò Boschi b52eb905ad fix: docker image build and startup (#46)
* ci: add docker smoke test to ci

* fix alpine version

* fix

* fix space

* fix space

* comment out

* docker fixes

* docker fixes

* docker fixes

* docker fixes

* docker fixes
2025-12-18 11:59:25 +01:00
Nicolò Boschi 1c6acc3ba0 feat: simplify mcp installation + ui standalone (#41) 2025-12-18 10:24:56 +01:00
DK09876 8ecb5d3a0c Add documentation code validation system (#43)
* Add documentation code validation system

- Create runnable example scripts in examples/api/ (19 files)
- Add CodeSnippet component for extracting marked sections
- Add raw-loader dependency for importing source files
- Create sample retain-new.mdx showing new approach
- Add README documenting coverage and gaps

* Fix wheel glob expansion in test-doc-examples CI job

* Fix CI issue

* Fix wheel path - uv build outputs to repo root dist/

* Fix: use explicit shell expansion for wheel install

* Fix: run cd in subshell so install runs from repo root

* Add documentation code validation CI job

- Use uv sync + uv run pattern (matches existing CI)
- Add requests to test dependencies for cleanup scripts

* Fix async API client usage in documents.py example

* Fix main-methods.py: RecallResult and ReflectFact don't have weight attribute

* Fix opinions.py: use actual API attributes instead of non-existent ones

* Fix example scripts: remove non-existent API attributes

- recall.py: remove .weight, fix entities iteration (dict not list)
- retain.mjs: remove result.async check
2025-12-18 10:21:38 +01:00
Chris Bartholomew ae80876671 fix: add procps to Docker image and smoke test to release workflow (#45)
* fix: add procps to Docker image and smoke test to release workflow

The Docker image was failing to start because pg0 uses `kill -0 <pid>`
to check if PostgreSQL is running, but the python:3.11-slim base image
doesn't include the `kill` command. Adding procps provides it.

This has been broken since release 0.1.6 when the fallback URI code was
removed to support dynamic ports. Without the kill command, pg0 couldn't
detect process status and returned None for the database URI.

Also adds smoke testing to the release workflow:
- Build image locally (single platform) and test before pushing
- Run container and wait for /health endpoint (up to 120s)
- Only push multi-platform release images if smoke test passes
- Each image (api-only, cp-only, standalone) tested independently

This prevents releasing broken Docker images to GHCR.

* refactor: extract smoke test into reusable script

Add scripts/docker-smoke-test.sh that can be run locally or in CI:
- Takes image name and optional target (cp-only vs api)
- Handles LLM credentials for API/standalone images
- Configurable timeout via SMOKE_TEST_TIMEOUT env var
- Colored output and clear error messages
- Proper cleanup on exit

Update release workflow to use the script instead of inline bash.
2025-12-17 22:01:53 +01:00
Chris Bartholomew 476a62da47 Add Hindsight Cloud links to README and docs (#42)
- Add Hindsight Cloud link to README header
- Add Hindsight Cloud navbar item in docs
- Add callout in installation docs for managed alternative
2025-12-17 11:09:36 -05:00
Nicolò Boschi 5aaa769ab9 Release v0.1.8
- Update version to 0.1.8 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-17 13:21:35 +01:00
Nicolò Boschi 04f01ab9ab fix: bank list response with no name banks 2025-12-17 13:21:16 +01:00
Nicolò Boschi 63f51385c4 fix: retain async fails (#40)
* fix: retain async fails

* fix: retain async fails
2025-12-17 13:17:38 +01:00
William Simmonds e468a4e19f fix: bank selector race condition when switching banks (#38) (#39) 2025-12-17 12:56:24 +01:00
Nicolò Boschi c0a0f447b7 Update README.md 2025-12-17 10:20:02 +01:00
Nicolò Boschi 84927ccc99 add run benchmarks instructions 2025-12-16 17:24:45 +01:00
Chris Latimer a6e8944ff0 README updates 2025-12-16 07:09:30 -07:00
Nicolò Boschi f6d890f6ed Release v0.1.7
- Update version to 0.1.7 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-16 14:29:07 +01:00
Nicolò Boschi 1fa8d9150c ci: check compatibility with python 3.11, 3.12 and 3.13 (#35) 2025-12-16 14:28:45 +01:00
Nicolò Boschi 656777c2be 0.1.6 changelog 2025-12-16 14:09:15 +01:00
Nicolò Boschi b36807ad3b Release v0.1.6
- Update version to 0.1.6 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-16 13:49:57 +01:00
Nicolò Boschi 11ac9cd9a5 less verbose git hooks 2025-12-16 13:49:38 +01:00
Nicolò Boschi 9394cf92f2 fix: doc build and lint files (#34)
* fix doc build

* fix doc build
2025-12-16 13:49:09 +01:00
Nicolò Boschi 47be07f97f bump pg0 0.11.x and improve documentation (#33)
* bump pg0 0.11.x and improve documentation

* bump pg0 0.11.x and improve documentation

* bump pg0 0.11.x and improve documentation

* ci: test notebooks on ci

* ci: test notebooks on ci

* rm llms-full from repo

* formatting

* formatting
2025-12-16 13:33:01 +01:00
Nicolò Boschi bb1f9cb221 feat: support for gemini-3-pro and gpt-5.2 (#30)
* feat: support for gemini-3-pro and gpt-5.2

* feat: support for gemini-3-pro and gpt-5.2

* feat: support for gemini-3-pro and gpt-5.2

* feat: support for gemini-3-pro and gpt-5.2

* feat: add local mcp server

* docs

* docs
2025-12-16 11:00:27 +01:00
Nicolò Boschi 7dd68538bb feat: add local mcp server (#32) 2025-12-16 10:50:20 +01:00
Nicolò Boschi 1cef364719 enable model tests on ci (#29) 2025-12-15 15:18:09 +01:00
Nicolò Boschi dff293ca8c fix doc link styling 2025-12-15 14:54:56 +01:00
Nicolò Boschi f4bc8443b3 changelog generator 2025-12-15 14:46:14 +01:00
Nicolò Boschi ae26a8603b models doc 2025-12-15 11:34:34 +01:00
Nicolò Boschi 183b9dacb4 Release v0.1.5
- Update version to 0.1.5 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all, hindsight-litellm
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-15 10:48:53 +01:00
Nicolò Boschi 8a7c6e4e91 litellm release integration 2025-12-15 10:48:27 +01:00
DK09876andClaude Opus 4.5 dfccbf29f1 Added hindsight_liteLLM implementation (#17)
* Added hindsight_liteLLM implementation

* Add instructions for entity vs bank id

* Add another line about entity

* Address PR review comments and enhance litellm integration

- Remove deprecated limit parameter from recall() and arecall() functions
  since Hindsight uses budget/max_tokens for result control
- Remove dead MODEL_MAX_OUTPUT_TOKENS dict and max_output_tokens property
  from LLMProvider (superseded by hardcoded max_completion_tokens)
- Add test-litellm-integration job to CI workflow
- Add reflect API support with use_reflect config option
- Add verbose mode debug info via get_last_injection_debug()
- Add entity_id support for multi-user memory isolation
- Add retain() and reflect() wrapper functions
- Update docstrings and examples

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

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

* Make max_memories optional to allow unlimited memory injection

- Change max_memories default from 10 to None (no limit)
- When max_memories is None, all results from the API are used
- Fix recall result handling to properly detect list vs object return
- Update wrappers (OpenAI, Anthropic) with same optional behavior

This allows users to control memory limits via max_memory_tokens
and recall_budget without an artificial count limit.

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

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

* Remove entity_id from hindsight_litellm; add gpt-4o token cap

Multi-user support now uses separate bank_ids per user instead of
entity_id scoping (e.g., bank_id=f"user-{user_id}"). This simplifies
the API and aligns with the Hindsight architecture.

Also fixes max_completion_tokens error for gpt-4o models by capping
the value at 16384 (gpt-4o's limit) instead of sending the default
65000 which exceeds the model's supported maximum.

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

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

* Fix dark mode styling across Control Plane UI components

Improvements to ensure proper text visibility and contrast in both light
and dark modes:

- Add global CSS rules for datetime-local calendar picker icon visibility
  using filter: invert() for both light (0.5) and dark (1) modes
- Fix text colors in dialog components to use theme-aware foreground colors
- Update memory detail panel, document/chunk modals, and data views to use
  proper dark mode text classes (text-foreground, text-card-foreground)
- Fix form labels, headings, and content text in bank selector dialogs
- Update entities view and documents view table styling for dark mode
- Bump package versions to 0.1.4

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

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

* Remove session_id feature and add How It Works section to README

- Remove session_id and session management (new_session, set_session,
  get_session) from config.py, callbacks.py, and __init__.py
- Session management was a client-only abstraction not backed by core API
- Add "How It Works" section to README with visual flow diagram
- Update README to remove session management documentation

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

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

* Fix readme example

* Add dark mode again

---------

Co-authored-by: Claude Opus 4.5 <[email protected]>
2025-12-15 10:42:19 +01:00
Chris Latimer dfea4dbe15 Trademark to README 2025-12-14 18:58:34 -07:00
Derek Bouius fcea8afa6c Change npm packaging structure and fix contributing info (#16)
* change the package to workspace concept

* add provider name and change default model

* add the node_modules to git ignore

* change the npm runs to use workspace

* fix the start scripts to use the workspace

* update the uv.lock

* updated instructions

* update the docker build to use the npm workspace

* Update package-lock.json after merge to sync workspace dependencies

* fix merge conflict
2025-12-12 14:14:19 -05:00
Nicolò Boschi 94c2b85c81 switch to pg0-embedded (#28)
* switch to pg0-embedded

* switch to pg0-embedded

* stricter mcp lib
2025-12-12 19:13:26 +01:00
Chris Bartholomew 160c5581ec fix: add DOM.Iterable lib to resolve URLSearchParams.entries() type error (#27)
The generated queryKeySerializer.gen.ts uses URLSearchParams.entries() which
requires DOM.Iterable in the TypeScript lib config for proper type definitions.
2025-12-12 17:34:47 +01:00
Nicolò Boschi 70983f5817 fix 400 retries on llm 2025-12-12 17:15:56 +01:00
Chris Latimer 44e9571572 README banner 2025-12-12 09:03:59 -07:00
Nicolò Boschi 7445cef7b7 feat: add optional graph retriever MPFP (#26)
* feat: add optional graph retriever MPFP

* feat: add optional graph retriever MPFP
2025-12-12 16:58:50 +01:00
Derek Bouius f018cc5677 fix: upgrade Next.js to 16.0.10 to patch CVE-2025-55184 and CVE-2025-55183 (#25)
CVE-2025-55184 (High) - Denial of Service via malicious HTTP request
CVE-2025-55183 (Medium) - Source Code Exposure of Server Actions

Reference: https://vercel.com/kb/bulletin/security-bulletin-cve-2025-55184-and-cve-2025-55183
2025-12-12 16:43:26 +01:00
Nicolò Boschi 922164e25c fix recall trace visualization 2025-12-12 14:38:37 +01:00
Derek Bouius d6b7b9b398 Fix base CI issues and the defaults in .env.example (#24)
* Add the LLM_PROVIDER in example

* fix the assert in testing recall

* trial to fix failing client tests

NotImplementedError: Cannot copy out of meta tensor; no data! Please use torch.nn.Module.to_empty() instead of torch.nn.Module.to() when moving module from meta to a different device.

* lock the sentence transformer packages to align with the breaking changes around lazy tensor loading

* Add the LLM_PROVIDER in example

* fix the assert in testing recall

* trial to fix failing client tests

* pre-cache the model so CI doesn't need workarounds

* remove assert that is a race condition

The test was checking that the bank count increased, but with parallel tests (-n 8), other tests can delete their banks while this test is running, causing a race condition. The important assertion is assert test_bank_id in final_banks - which verifies the bank was actually created.

* add debug to figure out why docker build fails sometimes

* use the CPU only version of pytorch to avoid pulling cuda libraries

* add best match strategy to uv

* change the example openai model
2025-12-11 16:48:12 -05:00
Nicolò Boschi 158a6aac9a fix cli installer 2025-12-11 16:26:05 +01:00
Nicolò Boschi 38e73a1414 fix cli installer 2025-12-11 16:22:04 +01:00
Nicolò Boschi 2c1be4cf47 Update Docker run command in README o3 mini 2025-12-11 14:53:52 +01:00
Nicolò Boschi f148d3e338 Release v0.1.4
- Update version to 0.1.4 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-11 14:24:06 +01:00
Nicolò Boschi 99db7b26c3 fix docs on clients 2025-12-11 14:23:56 +01:00
Nicolò Boschi ebc85a5c3d fix docs build 2025-12-11 12:54:36 +01:00
Nicolò Boschi ae30882ec9 Release v0.1.3
- Update version to 0.1.3 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-11 12:48:06 +01:00
Nicolò Boschi fa554b8980 brandind and misc fixes 2025-12-11 12:46:48 +01:00
Chris Latimer f813a807e7 README banner 2025-12-10 23:59:53 -05:00
Chris Latimer f7e8b1097b Fix README images 2025-12-10 10:59:59 -07:00
Nicolò Boschi 522a491fc1 Release v0.1.2
- Update version to 0.1.2 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-10 17:56:57 +01:00
Nicolò Boschi 1056a20e71 fix docker image 2025-12-10 17:56:51 +01:00
Nicolò Boschi 01ba9744e5 Release v0.1.1
- Update version to 0.1.1 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-dev/benchmarks, hindsight-all
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2025-12-10 17:30:24 +01:00
Nicolò Boschi 44e79feb3e helm chart updates v1 2025-12-10 17:30:15 +01:00
Nicolò Boschi 94665b2111 improve docs 2025-12-10 16:47:41 +01:00
Nicolò Boschi f42476bf94 fix: make sure openai provider works + docs updates (#23)
* fix: make sure openai provider works

* fix: make sure openai provider works

* fix
2025-12-10 16:10:10 +01:00
647 changed files with 100489 additions and 63742 deletions
+14 -1
View File
@@ -2,10 +2,23 @@
# Copy this file to .env and fill in your values
# LLM Configuration (Required)
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio
HINDSIGHT_API_LLM_PROVIDER=openai
HINDSIGHT_API_LLM_API_KEY=your-api-key-here
HINDSIGHT_API_LLM_MODEL=gpt-4o-mini
HINDSIGHT_API_LLM_MODEL=o3-mini
HINDSIGHT_API_LLM_BASE_URL=https://api.openai.com/v1
# Example: Anthropic Claude configuration
# HINDSIGHT_API_LLM_PROVIDER=anthropic
# HINDSIGHT_API_LLM_API_KEY=your-anthropic-api-key
# HINDSIGHT_API_LLM_MODEL=claude-sonnet-4-20250514
# Example: LM Studio local configuration (Qwen 2.5 32B recommended)
# HINDSIGHT_API_LLM_PROVIDER=lmstudio
# HINDSIGHT_API_LLM_API_KEY=lmstudio
# HINDSIGHT_API_LLM_BASE_URL=http://localhost:1234/v1
# HINDSIGHT_API_LLM_MODEL=qwen2.5-32b-instruct
# API Configuration (Optional)
HINDSIGHT_API_HOST=0.0.0.0
HINDSIGHT_API_PORT=8888
+27
View File
@@ -0,0 +1,27 @@
#!/bin/bash
# Pre-commit hook - runs all scripts in scripts/hooks/
set -e
REPO_ROOT="$(git rev-parse --show-toplevel)"
HOOKS_DIR="$REPO_ROOT/scripts/hooks"
if [ ! -d "$HOOKS_DIR" ]; then
exit 0
fi
echo ""
echo "=== Running pre-commit hooks ==="
echo ""
# Run all executable scripts in hooks directory
for hook in "$HOOKS_DIR"/*.sh; do
if [ -x "$hook" ]; then
echo "[hook] $(basename "$hook")"
(cd "$REPO_ROOT" && "$hook")
fi
done
echo ""
echo "=== Pre-commit hooks completed ==="
echo ""
+71
View File
@@ -0,0 +1,71 @@
name: Bug Report
description: Report a bug or unexpected behavior
labels: ["bug", "triage"]
body:
- type: markdown
attributes:
value: |
Thanks for taking the time to report a bug! Please fill out the sections below.
- type: textarea
id: description
attributes:
label: Bug Description
description: A clear and concise description of the bug
placeholder: What happened?
validations:
required: true
- type: textarea
id: reproduction
attributes:
label: Steps to Reproduce
description: Steps to reproduce the behavior
placeholder: |
1. Configure '...'
2. Call '...'
3. See error
validations:
required: true
- type: textarea
id: expected
attributes:
label: Expected Behavior
description: What did you expect to happen?
validations:
required: true
- type: textarea
id: actual
attributes:
label: Actual Behavior
description: What actually happened?
validations:
required: true
- type: input
id: version
attributes:
label: Version
description: What version are you using?
placeholder: e.g., 0.1.0 or commit hash
validations:
required: false
- type: dropdown
id: llm-provider
attributes:
label: LLM Provider
description: Which LLM provider are you using?
options:
- OpenAI
- Anthropic
- Gemini
- Groq
- Ollama
- LM Studio
- Other
validations:
required: false
+8
View File
@@ -0,0 +1,8 @@
blank_issues_enabled: false
contact_links:
- name: Questions & Help
url: https://github.com/vectorize-io/hindsight/discussions/categories/q-a
about: Please ask questions and get help in Discussions instead of opening an issue.
- name: Ideas & Feedback
url: https://github.com/vectorize-io/hindsight/discussions/categories/ideas
about: Share ideas or give feedback in Discussions.
@@ -0,0 +1,82 @@
name: Feature Request
description: Suggest a new feature or enhancement
labels: ["enhancement", "triage"]
body:
- type: markdown
attributes:
value: |
Thanks for suggesting a feature! Please describe what you'd like to see added.
- type: textarea
id: use-case
attributes:
label: Use Case
description: Describe your specific use case. What are you building? What's your goal?
placeholder: |
I'm building an AI agent that needs to...
My application handles...
validations:
required: true
- type: textarea
id: problem
attributes:
label: Problem Statement
description: What problem are you facing? What's missing or difficult today?
placeholder: Currently I have to... which causes...
validations:
required: true
- type: textarea
id: benefit
attributes:
label: How This Feature Would Help
description: Explain how this feature would improve your workflow or solve your problem
placeholder: With this feature, I would be able to...
validations:
required: true
- type: textarea
id: solution
attributes:
label: Proposed Solution
description: Describe your ideal solution (optional - we may have ideas too!)
placeholder: It would be great if Hindsight could...
validations:
required: false
- type: textarea
id: alternatives
attributes:
label: Alternatives Considered
description: Have you considered any alternative solutions or workarounds?
validations:
required: false
- type: dropdown
id: priority
attributes:
label: Priority
description: How important is this feature to you?
options:
- Nice to have
- Important - affects my workflow
- Critical - blocking my use case
validations:
required: true
- type: textarea
id: additional
attributes:
label: Additional Context
description: Any other context, mockups, or examples?
validations:
required: false
- type: checkboxes
id: checklist
attributes:
label: Checklist
options:
- label: I would be willing to contribute this feature
required: false
-11
View File
@@ -1,11 +0,0 @@
name: 'Setup pg0'
description: 'Install pg0 embedded PostgreSQL'
runs:
using: 'composite'
steps:
- name: Install pg0
shell: bash
run: |
curl -fsSL https://raw.githubusercontent.com/vectorize-io/pg0/main/install.sh | bash
echo "$HOME/.pg0/bin" >> $GITHUB_PATH
+5 -6
View File
@@ -20,18 +20,17 @@ concurrency:
jobs:
build:
runs-on: ubuntu-latest
defaults:
run:
working-directory: hindsight-docs
steps:
- uses: actions/checkout@v4
- uses: actions/setup-node@v4
with:
node-version: 20
cache: npm
cache-dependency-path: hindsight-docs/package-lock.json
- run: npm ci
- run: npm run build
cache-dependency-path: package-lock.json
- uses: astral-sh/setup-uv@v4
- run: npm ci --workspace=hindsight-docs
- run: uv run generate-llms-full
- run: npm run build --workspace=hindsight-docs
- uses: actions/upload-pages-artifact@v3
with:
path: hindsight-docs/build
+143 -53
View File
@@ -38,6 +38,14 @@ jobs:
working-directory: ./hindsight
run: uv build --out-dir dist
- name: Build hindsight-litellm
working-directory: ./hindsight-integrations/litellm
run: uv build --out-dir dist
- name: Build hindsight-embed
working-directory: ./hindsight-embed
run: uv build --out-dir dist
# Publish in order (client and api first, then hindsight-all which depends on them)
- name: Publish hindsight-client to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
@@ -57,6 +65,18 @@ jobs:
packages-dir: ./hindsight/dist
skip-existing: true
- name: Publish hindsight-litellm to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: ./hindsight-integrations/litellm/dist
skip-existing: true
- name: Publish hindsight-embed to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: ./hindsight-embed/dist
skip-existing: true
# Upload artifacts for GitHub release
- name: Upload artifacts
uses: actions/upload-artifact@v4
@@ -66,6 +86,8 @@ jobs:
hindsight-clients/python/dist/*
hindsight-api/dist/*
hindsight/dist/*
hindsight-integrations/litellm/dist/*
hindsight-embed/dist/*
retention-days: 1
release-typescript-client:
@@ -80,18 +102,29 @@ jobs:
with:
node-version: '20'
registry-url: 'https://registry.npmjs.org'
cache: 'npm'
cache-dependency-path: package-lock.json
- name: Install dependencies
working-directory: ./hindsight-clients/typescript
run: npm ci
run: npm ci --workspace=hindsight-clients/typescript
- name: Build
working-directory: ./hindsight-clients/typescript
run: npm run build
run: npm run build --workspace=hindsight-clients/typescript
- name: Publish to npm
working-directory: ./hindsight-clients/typescript
run: npm publish --access public
run: |
set +e
OUTPUT=$(npm publish --access public 2>&1)
EXIT_CODE=$?
echo "$OUTPUT"
if [ $EXIT_CODE -ne 0 ]; then
if echo "$OUTPUT" | grep -q "cannot publish over"; then
echo "Package version already published, skipping..."
exit 0
fi
exit $EXIT_CODE
fi
env:
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
@@ -106,6 +139,65 @@ jobs:
path: hindsight-clients/typescript/*.tgz
retention-days: 1
release-control-plane:
runs-on: ubuntu-latest
environment: npm
steps:
- uses: actions/checkout@v4
- name: Set up Node.js
uses: actions/setup-node@v4
with:
node-version: '20'
registry-url: 'https://registry.npmjs.org'
cache: 'npm'
cache-dependency-path: package-lock.json
- name: Install dependencies
run: npm ci
- name: Build TypeScript client (dependency)
run: npm run build --workspace=hindsight-clients/typescript
- name: Fix platform-specific native modules
run: |
# npm ci installs from lockfile which may have wrong platform binaries
# Delete hoisted native modules and reinstall for current platform
rm -rf node_modules/lightningcss node_modules/@tailwindcss
npm install lightningcss @tailwindcss/postcss @tailwindcss/node
- name: Build
run: npm run build --workspace=hindsight-control-plane
- name: Publish to npm
working-directory: ./hindsight-control-plane
run: |
set +e
OUTPUT=$(npm publish --access public 2>&1)
EXIT_CODE=$?
echo "$OUTPUT"
if [ $EXIT_CODE -ne 0 ]; then
if echo "$OUTPUT" | grep -q "cannot publish over"; then
echo "Package version already published, skipping..."
exit 0
fi
exit $EXIT_CODE
fi
env:
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
- name: Pack for GitHub release
working-directory: ./hindsight-control-plane
run: npm pack
- name: Upload artifacts
uses: actions/upload-artifact@v4
with:
name: control-plane
path: hindsight-control-plane/*.tgz
retention-days: 1
release-rust-cli:
runs-on: ${{ matrix.os }}
strategy:
@@ -170,7 +262,7 @@ jobs:
- name: Free Disk Space
uses: jlumbroso/free-disk-space@main
with:
tool-cache: false
tool-cache: true
android: true
dotnet: true
haskell: true
@@ -195,7 +287,7 @@ jobs:
id: get_version
run: echo "VERSION=${GITHUB_REF#refs/tags/v}" >> $GITHUB_OUTPUT
- name: Extract metadata
- name: Extract metadata for release tags
id: meta
uses: docker/metadata-action@v5
with:
@@ -206,7 +298,29 @@ jobs:
type=semver,pattern={{major}},value=${{ steps.get_version.outputs.VERSION }}
type=raw,value=latest
- name: Build and push
# TODO: Re-enable smoke test when disk space issue is resolved
# # Step 1: Build for local testing (single platform, no push)
# # This creates an identical image to what will be released, just for one platform
# - name: Build image for testing
# uses: docker/build-push-action@v6
# with:
# context: .
# file: docker/standalone/Dockerfile
# target: ${{ matrix.target }}
# push: false
# load: true
# tags: ${{ matrix.image_name }}:test
# cache-from: type=gha
# cache-to: type=gha,mode=max
# # Step 2: Test the image before pushing anything
# - name: Smoke test - verify container starts
# env:
# GROQ_API_KEY: ${{ secrets.GROQ_API_KEY }}
# run: ./scripts/docker-smoke-test.sh "${{ matrix.image_name }}:test" "${{ matrix.target }}"
# Build multi-platform and push to release tags
- name: Build and push release images
uses: docker/build-push-action@v6
with:
context: .
@@ -219,6 +333,9 @@ jobs:
release-helm-chart:
runs-on: ubuntu-latest
permissions:
contents: read
packages: write
steps:
- uses: actions/checkout@v4
@@ -228,12 +345,18 @@ jobs:
with:
version: 'latest'
- name: Log in to GHCR
run: echo "${{ secrets.GITHUB_TOKEN }}" | helm registry login ghcr.io -u ${{ github.actor }} --password-stdin
- name: Lint Helm chart
run: helm lint helm/hindsight
- name: Package Helm chart
run: helm package helm/hindsight --destination ./helm-packages
- name: Push to GHCR OCI
run: helm push helm-packages/*.tgz oci://ghcr.io/${{ github.repository_owner }}/charts
- name: Upload artifacts
uses: actions/upload-artifact@v4
with:
@@ -243,7 +366,7 @@ jobs:
create-github-release:
runs-on: ubuntu-latest
needs: [release-python-packages, release-typescript-client, release-rust-cli, release-docker-images, release-helm-chart]
needs: [release-python-packages, release-typescript-client, release-control-plane, release-rust-cli, release-docker-images, release-helm-chart]
permissions:
contents: write
@@ -266,6 +389,12 @@ jobs:
name: typescript-client
path: ./artifacts/typescript-client
- name: Download Control Plane
uses: actions/download-artifact@v4
with:
name: control-plane
path: ./artifacts/control-plane
- name: Download Rust CLI (Linux)
uses: actions/download-artifact@v4
with:
@@ -297,8 +426,12 @@ jobs:
cp artifacts/python-packages/hindsight-clients/python/dist/* release-assets/ || true
cp artifacts/python-packages/hindsight-api/dist/* release-assets/ || true
cp artifacts/python-packages/hindsight/dist/* release-assets/ || true
cp artifacts/python-packages/hindsight-integrations/litellm/dist/* release-assets/ || true
cp artifacts/python-packages/hindsight-embed/dist/* release-assets/ || true
# TypeScript client
cp artifacts/typescript-client/*.tgz release-assets/ || true
# Control Plane
cp artifacts/control-plane/*.tgz release-assets/ || true
# Rust CLI binaries
cp artifacts/rust-cli-linux/hindsight-linux-amd64 release-assets/ || true
cp artifacts/rust-cli-darwin-amd64/hindsight-darwin-amd64 release-assets/ || true
@@ -307,54 +440,11 @@ jobs:
cp artifacts/helm-chart/*.tgz release-assets/ || true
ls -la release-assets/
- name: Generate release notes
run: |
cat << 'EOF' > release-notes.md
## Quick Start
```bash
# Install the CLI
curl -fsSL https://raw.githubusercontent.com/vectorize-io/hindsight/refs/heads/main/hindsight-cli/install.sh | bash
# Start the server
docker run -p 8888:8888 -p 9999:9999 \
-e HINDSIGHT_API_LLM_PROVIDER=openai \
-e HINDSIGHT_API_LLM_API_KEY=$OPENAI_API_KEY \
-e HINDSIGHT_API_LLM_MODEL=gpt-4o-mini \
ghcr.io/${{ github.repository_owner }}/hindsight:${{ steps.get_version.outputs.VERSION }}
```
## Docker Images
- `ghcr.io/${{ github.repository_owner }}/hindsight:${{ steps.get_version.outputs.VERSION }}` - Standalone (recommended)
- `ghcr.io/${{ github.repository_owner }}/hindsight-api:${{ steps.get_version.outputs.VERSION }}` - API only
- `ghcr.io/${{ github.repository_owner }}/hindsight-control-plane:${{ steps.get_version.outputs.VERSION }}` - Web UI only
## CLI
```bash
curl -fsSL https://raw.githubusercontent.com/vectorize-io/hindsight/refs/heads/main/hindsight-cli/install.sh | bash
```
## Python
```bash
pip install hindsight-all # or hindsight-api, hindsight-client
```
## TypeScript/JavaScript
```bash
npm install @vectorize-io/hindsight-client
```
## Helm
```bash
helm install hindsight oci://ghcr.io/${{ github.repository_owner }}/charts/hindsight --version ${{ steps.get_version.outputs.VERSION }}
```
EOF
- name: Create GitHub Release
uses: softprops/action-gh-release@v2
with:
files: release-assets/*
body_path: release-notes.md
generate_release_notes: true
draft: false
prerelease: false
env:
+603 -22
View File
@@ -9,6 +9,131 @@ concurrency:
cancel-in-progress: true
jobs:
build-python-packages:
runs-on: ubuntu-latest
strategy:
matrix:
include:
- name: hindsight-all
path: hindsight
- name: hindsight-api
path: hindsight-api
- name: hindsight-client
path: hindsight-clients/python
- name: hindsight-embed
path: hindsight-embed
steps:
- uses: actions/checkout@v4
- name: Install uv
uses: astral-sh/setup-uv@v5
with:
enable-cache: true
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version-file: ".python-version"
- name: Build ${{ matrix.name }}
working-directory: ./${{ matrix.path }}
run: uv build
build-api-python-versions:
runs-on: ubuntu-latest
strategy:
matrix:
python-version: ['3.11', '3.12', '3.13']
steps:
- uses: actions/checkout@v4
- name: Install uv
uses: astral-sh/setup-uv@v5
with:
enable-cache: true
- name: Set up Python ${{ matrix.python-version }}
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
- name: Build hindsight-api
working-directory: ./hindsight-api
run: uv build
build-typescript-client:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Set up Node.js
uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
cache-dependency-path: package-lock.json
- name: Install dependencies
run: npm ci --workspace=hindsight-clients/typescript
- name: Build TypeScript client
run: npm run build --workspace=hindsight-clients/typescript
build-control-plane:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Set up Node.js
uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
cache-dependency-path: package-lock.json
- name: Install SDK dependencies
run: npm ci --workspace=hindsight-clients/typescript
- name: Build SDK
run: npm run build --workspace=hindsight-clients/typescript
# Install control plane deps and fix hoisted lightningcss binary
# lightningcss gets hoisted to root node_modules, so we need to reinstall it there
- name: Install Control Plane dependencies
run: |
npm install --workspace=hindsight-control-plane
rm -rf node_modules/lightningcss node_modules/@tailwindcss
npm install lightningcss @tailwindcss/postcss @tailwindcss/node
- name: Build Control Plane
run: npm run build --workspace=hindsight-control-plane
- name: Verify standalone build
run: |
test -f hindsight-control-plane/standalone/server.js || exit 1
test -d hindsight-control-plane/standalone/node_modules || exit 1
node hindsight-control-plane/bin/cli.js --help
- name: Smoke test - verify server starts
run: |
cd hindsight-control-plane
node bin/cli.js --port 9999 &
SERVER_PID=$!
sleep 5
if curl -sf http://localhost:9999 > /dev/null 2>&1; then
echo "Server started successfully"
kill $SERVER_PID 2>/dev/null || true
exit 0
else
echo "Server failed to respond"
kill $SERVER_PID 2>/dev/null || true
exit 1
fi
build-docs:
runs-on: ubuntu-latest
@@ -19,14 +144,14 @@ jobs:
uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
cache-dependency-path: package-lock.json
- name: Install dependencies
working-directory: ./hindsight-docs
run: npm ci
run: npm ci --workspace=hindsight-docs
- name: Build docs
working-directory: ./hindsight-docs
run: npm run build
run: npm run build --workspace=hindsight-docs
build-rust-cli:
runs-on: ubuntu-latest
@@ -50,6 +175,90 @@ jobs:
working-directory: hindsight-cli
run: cargo build --release
- name: Upload CLI artifact
uses: actions/upload-artifact@v4
with:
name: hindsight-cli
path: hindsight-cli/target/release/hindsight
retention-days: 1
test-rust-cli:
runs-on: ubuntu-latest
needs: build-rust-cli
env:
HINDSIGHT_API_LLM_PROVIDER: groq
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
HINDSIGHT_API_URL: http://localhost:8888
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
- name: Download CLI artifact
uses: actions/download-artifact@v4
with:
name: hindsight-cli
path: /tmp/cli
- name: Make CLI executable
run: chmod +x /tmp/cli/hindsight
- name: Install uv
uses: astral-sh/setup-uv@v5
with:
enable-cache: true
prune-cache: false
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version-file: ".python-version"
- name: Build API
working-directory: ./hindsight-api
run: uv build
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Create .env file
run: |
cat > .env << EOF
HINDSIGHT_API_LLM_PROVIDER=${{ env.HINDSIGHT_API_LLM_PROVIDER }}
HINDSIGHT_API_LLM_API_KEY=${{ env.HINDSIGHT_API_LLM_API_KEY }}
HINDSIGHT_API_LLM_MODEL=${{ env.HINDSIGHT_API_LLM_MODEL }}
EOF
- name: Start API server
run: |
./scripts/dev/start-api.sh > /tmp/api-server.log 2>&1 &
echo "Waiting for API server to be ready..."
for i in {1..60}; do
if curl -sf http://localhost:8888/health > /dev/null 2>&1; then
echo "API server is ready after ${i}s"
break
fi
if [ $i -eq 60 ]; then
echo "API server failed to start after 60s"
cat /tmp/api-server.log
exit 1
fi
sleep 1
done
- name: Run CLI smoke test
run: |
HINDSIGHT_CLI=/tmp/cli/hindsight ./hindsight-cli/smoke-test.sh
- name: Show API server logs
if: always()
run: |
echo "=== API Server Logs ==="
cat /tmp/api-server.log || echo "No API server log found"
lint-helm-chart:
runs-on: ubuntu-latest
@@ -82,7 +291,7 @@ jobs:
- name: Free Disk Space
uses: jlumbroso/free-disk-space@main
with:
tool-cache: false
tool-cache: true
android: true
dotnet: true
haskell: true
@@ -100,14 +309,28 @@ jobs:
file: docker/standalone/Dockerfile
target: ${{ matrix.target }}
push: false
load: false
# TODO: Re-enable smoke test when disk space issue is resolved
# - name: Smoke test - verify container starts
# env:
# GROQ_API_KEY: ${{ secrets.GROQ_API_KEY }}
# run: ./scripts/docker-smoke-test.sh "hindsight-${{ matrix.name }}:test" "${{ matrix.target }}"
test-api:
runs-on: ubuntu-latest
env:
HINDSIGHT_API_LLM_PROVIDER: groq
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
GROQ_API_KEY: ${{ secrets.GROQ_API_KEY }}
GEMINI_API_KEY: ${{ secrets.GEMINI_API_KEY }}
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
COHERE_API_KEY: ${{ secrets.COHERE_API_KEY }}
HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
# Prefer CPU-only PyTorch in CI (but keep PyPI for everything else)
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
@@ -123,16 +346,33 @@ jobs:
with:
python-version-file: ".python-version"
- name: Install pg0
uses: ./.github/actions/setup-pg0
- name: Build API
working-directory: ./hindsight-api
run: uv build
- name: Install dependencies
working-directory: ./hindsight-api
run: uv sync --extra test
run: uv sync --frozen --extra test --no-install-project --index-strategy unsafe-best-match
- 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: 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 tests
working-directory: ./hindsight-api
@@ -146,6 +386,8 @@ jobs:
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
HINDSIGHT_API_URL: http://localhost:8888
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
# Prefer CPU-only PyTorch in CI (but keep PyPI for everything else)
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
@@ -161,9 +403,6 @@ jobs:
with:
python-version-file: ".python-version"
- name: Install pg0
uses: ./.github/actions/setup-pg0
- name: Build API
working-directory: ./hindsight-api
run: uv build
@@ -174,11 +413,11 @@ jobs:
- name: Install client test dependencies
working-directory: ./hindsight-clients/python
run: uv sync --extra test
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Create .env file
run: |
@@ -223,6 +462,8 @@ jobs:
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
HINDSIGHT_API_URL: http://localhost:8888
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
# Prefer CPU-only PyTorch in CI (but keep PyPI for everything else)
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
@@ -243,16 +484,13 @@ jobs:
with:
node-version: '20'
- name: Install pg0
uses: ./.github/actions/setup-pg0
- name: Build API
working-directory: ./hindsight-api
run: uv build
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Install TypeScript client dependencies
working-directory: ./hindsight-clients/typescript
@@ -305,6 +543,8 @@ jobs:
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
HINDSIGHT_API_URL: http://localhost:8888
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
# Prefer CPU-only PyTorch in CI (but keep PyPI for everything else)
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
@@ -332,16 +572,13 @@ jobs:
hindsight-clients/rust/target
key: ${{ runner.os }}-cargo-client-${{ hashFiles('hindsight-clients/rust/Cargo.lock') }}
- name: Install pg0
uses: ./.github/actions/setup-pg0
- name: Build API
working-directory: ./hindsight-api
run: uv build
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Create .env file
run: |
@@ -377,3 +614,347 @@ jobs:
run: |
echo "=== API Server Logs ==="
cat /tmp/api-server.log || echo "No API server log found"
test-integration:
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
HINDSIGHT_API_URL: http://localhost:8888
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
- name: Install uv
uses: astral-sh/setup-uv@v5
with:
enable-cache: true
prune-cache: false
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version-file: ".python-version"
- name: Build API
working-directory: ./hindsight-api
run: uv build
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Install integration test dependencies
working-directory: ./hindsight-integration-tests
run: uv sync --frozen
- 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: 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: Create .env file
run: |
cat > .env << EOF
HINDSIGHT_API_LLM_PROVIDER=${{ env.HINDSIGHT_API_LLM_PROVIDER }}
HINDSIGHT_API_LLM_API_KEY=${{ env.HINDSIGHT_API_LLM_API_KEY }}
HINDSIGHT_API_LLM_MODEL=${{ env.HINDSIGHT_API_LLM_MODEL }}
EOF
- name: Start API server
run: |
./scripts/dev/start-api.sh > /tmp/api-server.log 2>&1 &
echo "Waiting for API server to be ready..."
for i in {1..60}; do
if curl -sf http://localhost:8888/health > /dev/null 2>&1; then
echo "API server is ready after ${i}s"
break
fi
if [ $i -eq 60 ]; then
echo "API server failed to start after 60s"
cat /tmp/api-server.log
exit 1
fi
sleep 1
done
- name: Run integration tests
working-directory: ./hindsight-integration-tests
run: uv run pytest tests/ -v
- name: Show API server logs
if: always()
run: |
echo "=== API Server Logs ==="
cat /tmp/api-server.log || echo "No API server log found"
test-litellm-integration:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Install uv
uses: astral-sh/setup-uv@v5
with:
enable-cache: true
prune-cache: false
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version-file: ".python-version"
- name: Build litellm integration
working-directory: ./hindsight-integrations/litellm
run: uv build
- name: Install dependencies
working-directory: ./hindsight-integrations/litellm
run: uv sync --frozen --extra dev
- name: Run tests
working-directory: ./hindsight-integrations/litellm
run: uv run pytest tests -v
test-embed:
runs-on: ubuntu-latest
env:
HINDSIGHT_EMBED_LLM_PROVIDER: groq
HINDSIGHT_EMBED_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
HINDSIGHT_EMBED_LLM_MODEL: openai/gpt-oss-20b
# Prefer CPU-only PyTorch in CI
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
- name: Install uv
uses: astral-sh/setup-uv@v5
with:
enable-cache: true
prune-cache: false
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version-file: ".python-version"
- name: Install dependencies
working-directory: ./hindsight-embed
run: uv sync --frozen --index-strategy unsafe-best-match
- name: Cache HuggingFace models
uses: actions/cache@v4
with:
path: ~/.cache/huggingface
key: ${{ runner.os }}-huggingface-embed-${{ hashFiles('hindsight-embed/pyproject.toml') }}
restore-keys: |
${{ runner.os }}-huggingface-embed-
${{ runner.os }}-huggingface-
- name: Run smoke test
working-directory: ./hindsight-embed
run: ./test.sh
test-doc-examples:
runs-on: ubuntu-latest
needs: build-rust-cli
env:
HINDSIGHT_API_LLM_PROVIDER: groq
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
HINDSIGHT_API_URL: http://localhost:8888
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
- name: Download CLI artifact
uses: actions/download-artifact@v4
with:
name: hindsight-cli
path: /usr/local/bin
- name: Make CLI executable
run: chmod +x /usr/local/bin/hindsight
- 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: Set up Node.js
uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
cache-dependency-path: package-lock.json
- name: Build and install API
working-directory: ./hindsight-api
run: |
uv build
uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Install Python client dependencies
working-directory: ./hindsight-clients/python
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
- name: Install TypeScript client
run: |
npm ci --workspace=hindsight-clients/typescript
npm run build --workspace=hindsight-clients/typescript
- name: Create .env file
run: |
cat > .env << EOF
HINDSIGHT_API_LLM_PROVIDER=${{ env.HINDSIGHT_API_LLM_PROVIDER }}
HINDSIGHT_API_LLM_API_KEY=${{ env.HINDSIGHT_API_LLM_API_KEY }}
HINDSIGHT_API_LLM_MODEL=${{ env.HINDSIGHT_API_LLM_MODEL }}
EOF
- name: Start API server
run: |
./scripts/dev/start-api.sh > /tmp/api-server.log 2>&1 &
echo "Waiting for API server to be ready..."
for i in {1..60}; do
if curl -sf http://localhost:8888/health > /dev/null 2>&1; then
echo "API server is ready after ${i}s"
break
fi
if [ $i -eq 60 ]; then
echo "API server failed to start after 60s"
cat /tmp/api-server.log
exit 1
fi
sleep 1
done
- name: Run Python doc examples
working-directory: ./hindsight-clients/python
run: |
for f in ../../hindsight-docs/examples/api/*.py; do
echo "Running $f..."
uv run python "$f"
done
- name: Run Node.js doc examples
run: |
for f in hindsight-docs/examples/api/*.mjs; do
echo "Running $f..."
node "$f"
done
- name: Configure CLI
run: hindsight configure --api-url http://localhost:8888
- name: Run CLI doc examples
run: |
for f in hindsight-docs/examples/api/*.sh; do
echo "Running $f..."
bash "$f"
done
- name: Show API server logs
if: always()
run: |
echo "=== API Server Logs ==="
cat /tmp/api-server.log || echo "No API server log found"
verify-generated-files:
runs-on: ubuntu-latest
env:
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
steps:
- uses: actions/checkout@v4
- name: Install uv
uses: astral-sh/setup-uv@v5
with:
enable-cache: true
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version-file: ".python-version"
- name: Set up Node.js
uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
cache-dependency-path: package-lock.json
- name: Install Rust
uses: dtolnay/rust-toolchain@stable
- name: Cache cargo
uses: actions/cache@v4
with:
path: |
~/.cargo/registry
~/.cargo/git
key: ${{ runner.os }}-cargo-gen-${{ hashFiles('**/Cargo.lock') }}
- name: Install Node dependencies
run: npm ci
- name: Install Python dependencies
run: |
cd hindsight-dev && uv sync --frozen --index-strategy unsafe-best-match
cd ../hindsight-api && uv sync --frozen --index-strategy unsafe-best-match
cd ../hindsight-embed && uv sync --frozen --index-strategy unsafe-best-match
- name: Run generate-openapi
run: ./scripts/generate-openapi.sh
- name: Run generate-clients
run: ./scripts/generate-clients.sh
- name: Run lint
run: ./scripts/hooks/lint.sh
- name: Verify no uncommitted changes
run: |
if [ -n "$(git status --porcelain)" ]; then
echo "❌ Error: Generated files are out of sync with committed files."
echo ""
echo "The following files have changed after running generation scripts:"
git status --porcelain
echo ""
echo "Please run the following commands locally and commit the changes:"
echo " ./scripts/generate-openapi.sh"
echo " ./scripts/generate-clients.sh"
echo " ./scripts/hooks/lint.sh"
echo ""
git diff --stat
exit 1
fi
echo "✓ All generated files are up to date"
+20 -3
View File
@@ -5,12 +5,18 @@ build/
dist/
wheels/
*.egg-info
.mcp.json
.osgrep
# Virtual environments
.venv
# Environment variables
# Node
node_modules/
# Environment variables and local config
.env
docker-compose.yml
docker-compose.override.yml
# IDE
.idea/
@@ -21,6 +27,10 @@ wheels/
# NLTK data (will be downloaded automatically)
nltk_data/
# Monitoring stack (Prometheus/Grafana binaries and data)
.monitoring/
.pgbouncer
# Large benchmark datasets (will be downloaded automatically)
**/longmemeval_s_cleaned.json
@@ -29,8 +39,15 @@ logs/
.DS_Store
# Generated docs files
hindsight-docs/static/llms-full.txt
hindsight-dev/benchmarks/locomo/results/
hindsight-dev/benchmarks/longmemeval/results/
hindsight-cli/target
hindsight-clients/rust/target
hindsight-clients/rust/target
.claude
whats-next.md
TASK.md
CHANGELOG.md
+39
View File
@@ -0,0 +1,39 @@
[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
@@ -0,0 +1,2 @@
"hindsight" "md5d842ccb6249bcd3c53b2f648378092a6"
"admin" ""
+6
View File
@@ -14,6 +14,7 @@ This document captures architectural decisions and coding conventions for the Hi
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
@@ -145,3 +146,8 @@ Note: The maintained wrapper `hindsight_client.py` and `README.md` are preserved
- 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
+251
View File
@@ -0,0 +1,251 @@
# CLAUDE.md
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
## Project Overview
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
## Development Commands
### API Server (Python/FastAPI)
```bash
# Start API server (loads .env automatically)
./scripts/dev/start-api.sh
# Run all tests (parallelized with pytest-xdist)
cd hindsight-api && uv run pytest tests/
# Run specific test file
cd hindsight-api && uv run pytest tests/test_http_api_integration.py -v
# Run single test function
cd hindsight-api && uv run pytest tests/test_retain.py::test_retain_simple -v
# Lint and format
cd hindsight-api && uv run ruff check .
cd hindsight-api && uv run ruff format .
# Type checking (uses ty - extremely fast type checker from Astral)
cd hindsight-api && uv run ty check hindsight_api/
```
### Control Plane (Next.js)
```bash
./scripts/dev/start-control-plane.sh
# Or manually:
cd hindsight-control-plane && npm run dev
```
### Documentation Site (Docusaurus)
```bash
./scripts/dev/start-docs.sh
```
### Generating Clients/OpenAPI
```bash
# Regenerate OpenAPI spec after API changes (REQUIRED after changing endpoints)
./scripts/generate-openapi.sh
# Regenerate all client SDKs (Python, TypeScript, Rust)
./scripts/generate-clients.sh
```
### Benchmarks
```bash
./scripts/benchmarks/run-longmemeval.sh
./scripts/benchmarks/run-locomo.sh
./scripts/benchmarks/start-visualizer.sh # View results at localhost:8001
```
## Architecture
### Monorepo Structure
- **hindsight-api/**: Core FastAPI server with memory engine (Python, uv)
- **hindsight/**: Embedded Python bundle (hindsight-all package)
- **hindsight-control-plane/**: Admin UI (Next.js, npm)
- **hindsight-cli/**: CLI tool (Rust, cargo, uses progenitor for API client)
- **hindsight-clients/**: Generated SDK clients (Python, TypeScript, Rust)
- **hindsight-docs/**: Docusaurus documentation site
- **hindsight-integrations/**: Framework integrations (LiteLLM, OpenAI)
- **hindsight-dev/**: Development tools and benchmarks
### Core Engine (hindsight-api/hindsight_api/engine/)
- `memory_engine.py`: Main orchestrator (~170KB) for retain/recall/reflect operations
- `llm_wrapper.py`: LLM abstraction supporting OpenAI, Anthropic, Gemini, Groq, Ollama, LM Studio
- `embeddings.py`: Embedding generation (local sentence-transformers or TEI)
- `cross_encoder.py`: Reranking (local or TEI)
- `entity_resolver.py`: Entity extraction and normalization
- `query_analyzer.py`: Query intent analysis
**retain/**: Memory ingestion pipeline
- `orchestrator.py`: Coordinates the retain flow
- `fact_extraction.py`: LLM-based fact extraction from content
- `link_utils.py`: Entity link creation and management
**search/**: Multi-strategy retrieval
- `retrieval.py`: Main retrieval orchestrator
- `graph_retrieval.py`: Entity/relationship graph traversal
- `mpfp_retrieval.py`: Multi-Path Fact Propagation retrieval
- `fusion.py`: Reciprocal rank fusion for combining results
- `reranking.py`: Cross-encoder reranking
### API Layer (hindsight-api/hindsight_api/api/)
- `http.py`: FastAPI HTTP routers (~80KB) for all REST endpoints
- `mcp.py`: Model Context Protocol server implementation
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)
### Database
PostgreSQL with pgvector. Schema managed via Alembic migrations in `hindsight-api/hindsight_api/alembic/`. Migrations run automatically on API startup.
Key tables: `banks`, `memory_units`, `documents`, `entities`, `entity_links`
### Adding Database Migrations
1. **Create a new migration file** in `hindsight-api/hindsight_api/alembic/versions/`:
- File name format: `<revision_id>_<description>.py` (e.g., `f1a2b3c4d5e6_add_new_index.py`)
- Use a unique hex revision ID (12 chars)
- Set `down_revision` to the previous migration's revision ID
2. **Migration template**:
```python
"""Description of the migration
Revision ID: f1a2b3c4d5e6
Revises: <previous_revision_id>
Create Date: YYYY-MM-DD
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "f1a2b3c4d5e6"
down_revision: str | Sequence[str] | None = "<previous_revision_id>"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
"""Get schema prefix for table names (required for multi-tenant support)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"CREATE INDEX ... ON {schema}table_name(...)")
def downgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"DROP INDEX IF EXISTS {schema}index_name")
```
3. **Run migrations locally**:
```bash
# Set database URL and run migrations
uv run hindsight-admin run-db-migration
# Run on a specific tenant schema
uv run hindsight-admin run-db-migration --schema tenant_xyz
```
## Key Conventions
### Code Quality
**Always run the lint script after making Python or TypeScript/Node changes:**
```bash
./scripts/hooks/lint.sh
```
This runs the same checks as the pre-commit hook (Ruff for Python, ESLint/Prettier for TypeScript).
### Memory Banks
- Each bank is an isolated memory store (like a "brain" for one user/agent)
- Banks have dispositions (skepticism, literalism, empathy traits 1-5) affecting reflect
- Banks can have background context
- Bank isolation is strict - no cross-bank data leakage
### API Design
- All endpoints operate on a single bank per request
- Multi-bank queries are client responsibility to orchestrate
- Disposition traits only affect reflect, not recall
### Control Plane API Routes
When adding or modifying parameters in the dataplane API (hindsight-api), you must also update the control plane routes that proxy to it:
1. **API Routes** (`hindsight-control-plane/src/app/api/`):
- `recall/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/recall`
- `reflect/route.ts` - proxies to `/v1/default/banks/{bank_id}/reflect`
- `memories/retain/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/retain`
- Other routes follow the same pattern
2. **Client types** (`hindsight-control-plane/src/lib/api.ts`):
- Update the TypeScript type definitions for `recall()`, `reflect()`, `retain()` etc.
3. **Checklist when adding new API parameters**:
- Add parameter extraction in the route handler (destructure from `body`)
- Pass the parameter to the SDK call
- Update the client type definition in `lib/api.ts`
- Update any UI components that need to use the new parameter
### Python Style
- Python 3.11+, type hints required
- Async throughout (asyncpg, async FastAPI)
- Pydantic models for request/response
- Ruff for linting (line-length 120)
- No Python files at project root - maintain clean directory structure
### TypeScript Style
- Next.js App Router for control plane
- Tailwind CSS with shadcn/ui components
### Adding New API Configuration Flags
When adding a new environment variable configuration:
1. **config.py** (`hindsight-api/hindsight_api/config.py`):
- Add `ENV_*` constant for the environment variable name
- Add `DEFAULT_*` constant for the default value
- Add field to `HindsightConfig` dataclass
- Add initialization in `from_env()` method
2. **main.py** (`hindsight-api/hindsight_api/main.py`):
- Add field to the manual `HindsightConfig()` constructor call (search for "CLI override")
3. **Use the config** in code:
```python
from ...config import get_config
config = get_config()
value = config.your_new_field
```
4. **Documentation** (`hindsight-docs/docs/developer/configuration.md`):
- Add to appropriate section table with Variable, Description, Default
## Environment Setup
```bash
cp .env.example .env
# Edit .env with LLM API key
# Python deps
uv sync --directory hindsight-api/
# Node deps (uses npm workspaces)
npm install
```
Required env vars:
- `HINDSIGHT_API_LLM_PROVIDER`: openai, anthropic, gemini, groq, ollama, lmstudio
- `HINDSIGHT_API_LLM_API_KEY`: Your API key
- `HINDSIGHT_API_LLM_MODEL`: Model name (e.g., o3-mini, claude-sonnet-4-20250514)
Optional (uses local models by default):
- `HINDSIGHT_API_EMBEDDINGS_PROVIDER`: local (default) or tei
- `HINDSIGHT_API_RERANKER_PROVIDER`: local (default) or tei
- `HINDSIGHT_API_DATABASE_URL`: External PostgreSQL (uses embedded pg0 by default)
+44 -5
View File
@@ -5,13 +5,23 @@ Thanks for your interest in contributing to Hindsight!
## Getting Started
1. Fork and clone the repository
2. Install dependencies:
```bash
cd hindsight-api && uv sync
git clone [email protected]:vectorize-io/hindsight.git
cd hindsight
```
3. Set up your environment:
2. Set up your environment:
```bash
export OPENAI_API_KEY=your-key
cp .env.example .env
```
Edit the .env to add LLM API key and config as required
3. Install dependencies:
```bash
# Python dependencies
uv sync --directory hindsight-api/
# Node dependencies (uses npm workspaces)
npm install
```
## Development
@@ -41,7 +51,36 @@ cd hindsight-api
uv run pytest tests/
```
### Code style
### Code Style
We use [Ruff](https://docs.astral.sh/ruff/) for Python linting and formatting, and ESLint/Prettier for TypeScript.
#### Setting up git hooks (recommended)
Set up git hooks to automatically lint and format code before each commit:
```bash
./scripts/setup-hooks.sh
```
This configures git to use the hooks in `.githooks/`, which run all scripts in `scripts/hooks/` on commit. The lint hook runs in parallel:
- **Python**: `ruff check --fix`, `ruff format`, `ty check`
- **TypeScript**: `eslint --fix`, `prettier`
#### Manual linting and formatting
```bash
# Run all lints (same as pre-commit)
./scripts/hooks/lint.sh
# Or run individually for Python:
cd hindsight-api
uv run ruff check --fix . # Lint and auto-fix
uv run ruff format . # Format code
uv run ty check hindsight_api # Type check
```
#### Style guidelines
- Use Python type hints
- Follow existing code patterns
+69 -58
View File
@@ -1,16 +1,15 @@
<div align="center">
# Hindsight
![Hindsight Banner](./hindsight-docs/static/img/banner.svg)
**Agent Memory that Works Like Human Memory**
[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)
[![CI](https://github.com/vectorize-io/hindsight/actions/workflows/test.yml/badge.svg)](https://github.com/vectorize-io/hindsight/actions/workflows/test.yml)
[![CI](https://github.com/vectorize-io/hindsight/actions/workflows/release.yml/badge.svg)](https://github.com/vectorize-io/hindsight/actions/workflows/release.yml)
[![Slack Community](https://img.shields.io/badge/Slack-Join%20Community-4A154B?logo=slack)](https://join.slack.com/t/hindsight-space/shared_invite/zt-3nhbm4w29-LeSJ5Ixi6j8PdiYOCPlOgg)
[![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
[![PyPI - hindsight-client](https://img.shields.io/pypi/v/hindsight-client?label=hindsight-client)](https://pypi.org/project/hindsight-client/)
[![npm](https://img.shields.io/npm/v/@vectorize-io/hindsight-client)](https://www.npmjs.com/package/@vectorize-io/hindsight-client)
[![Slack Community](https://img.shields.io/badge/Slack-Join%20Community-4A154B?logo=slack)](https://join.slack.com/t/hindsight-space/shared_invite/zt-3klo21kua-VUCC_zHP5rIcXFB1_5yw6A)
![PyPI - Downloads](https://img.shields.io/pypi/dm/hindsight-api?label=PyPI)
![NPM Downloads](https://img.shields.io/npm/dm/%40vectorize-io%2Fhindsight-client?logoColor=orange&label=NPM&color=blue&link=https%3A%2F%2Fwww.npmjs.com%2Fpackage%2F%40vectorize-io%2Fhindsight-client)
[Documentation](https://vectorize-io.github.io/hindsight) • [Paper](./Hindsight.pdf) • [Examples](./examples)
</div>
@@ -18,7 +17,7 @@
## 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.
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 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.
@@ -26,27 +25,48 @@ Hindsight addresses common challenges that have frustrated AI engineers building
- **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.
## How Hindsight Works
## How is Hindsight Different From Other Memory Systems?
![Overview](./hindsight-docs/static/img/hindsight-overview.png)
![Overview](./hindsight-docs/static/img/hindsight-overview.webp)
Hindsight organizes memory into four networks to mimic the way human memory works:
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.
Memories in Hindsight are stored in banks (e.g. memory banks). When memories are retained, they are transformed to construct a series of search indexes, time series data, and entity/relationship graphs.
### 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.
---
## 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:
![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.
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.
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
### Docker (recommended)
@@ -54,21 +74,22 @@ Memories in Hindsight are stored in banks (e.g. memory banks). When memories are
```bash
export OPENAI_API_KEY=your-key
docker run -p 8888:8888 -p 9999:9999 \
-e HINDSIGHT_API_LLM_PROVIDER=openai \
docker run --rm -it --pull always -p 8888:8888 -p 9999:9999 \
-e HINDSIGHT_API_LLM_API_KEY=$OPENAI_API_KEY \
-e HINDSIGHT_API_LLM_MODEL=gpt-4o-mini \
-e HINDSIGHT_API_LLM_MODEL=o3-mini \
-v $HOME/.hindsight-docker:/home/hindsight/.pg0 \
ghcr.io/vectorize-io/hindsight
ghcr.io/vectorize-io/hindsight:latest
```
You can modify the LLM provider by setting `HINDSIGHT_API_LLM_PROVIDER`. Valid options are `openai`, `anthropic`, `gemini`, `groq`, `ollama`, and `lmstudio`. The documentation provides more details on [supported models](https://hindsight.vectorize.io/developer/models).
API: http://localhost:8888
UI: http://localhost:9999
Install client:
```bash
pip install hindsight-client
pip install hindsight-client -U
# or
npm install @vectorize-io/hindsight-client
```
@@ -76,25 +97,24 @@ npm install @vectorize-io/hindsight-client
Python example:
```python
from hindsight import HindsightClient
from hindsight_client import Hindsight
client = HindsightClient(base_url="http://localhost:8888")
client = Hindsight(base_url="http://localhost:8888")
# Store
client.retain(bank_id="my-agent", content="Alice works at Google as a software engineer")
# Retain: Store information
client.retain(bank_id="my-bank", content="Alice works at Google as a software engineer")
# Query
results = client.recall(bank_id="my-agent", query="What does Alice do?")
# Recall: Search memories
client.recall(bank_id="my-bank", query="What does Alice do?")
# Reflect
response = client.reflect(bank_id="my-agent", query="Tell me about Alice")
print(response.text)
# Reflect: Generate disposition-aware response
client.reflect(bank_id="my-bank", query="Tell me about Alice")
```
### Python (embedded, no Docker)
```bash
pip install hindsight-all
pip install hindsight-all -U
```
```python
@@ -103,27 +123,27 @@ from hindsight import HindsightServer, HindsightClient
with HindsightServer(
llm_provider="openai",
llm_model="gpt-4o-mini",
llm_model="gpt-5-mini",
llm_api_key=os.environ["OPENAI_API_KEY"]
) as server:
client = HindsightClient(base_url=server.url)
client.retain(bank_id="my-agent", content="Alice works at Google")
results = client.recall(bank_id="my-agent", query="Where does Alice work?")
client.retain(bank_id="my-bank", content="Alice works at Google")
results = client.recall(bank_id="my-bank", query="Where does Alice work?")
```
### TypeScript
### Node.js / TypeScript
```bash
npm install @vectorize-io/hindsight-client
```
```typescript
import { HindsightClient } from '@vectorize-io/hindsight-client';
```javascript
const { HindsightClient } = require('@vectorize-io/hindsight-client');
const client = new HindsightClient({ baseUrl: 'http://localhost:8888' });
await client.retain('my-agent', 'Alice loves hiking in Yosemite');
const response = await client.recall('my-agent', 'What does Alice like?');
await client.retain('my-bank', 'Alice loves hiking in Yosemite');
await client.recall('my-bank', 'What does Alice like?');
```
---
@@ -156,7 +176,7 @@ client.retain(
Behind the scenes, the retain operation uses an LLM to extract key facts, temporal data, entities, and relationships. It passes these through a normalization process to transform extracted data into canonical entities, time series, and search indexes along with metadata. These representations create the pathways for accurate memory retrieval in the recall and reflect operations.
![Retain Operation](hindsight-docs/static/img/retain-operation.png)
![Retain Operation](hindsight-docs/static/img/retain-operation.webp)
### Recall
@@ -171,9 +191,7 @@ client = Hindsight(base_url="http://localhost:8888")
client.recall(bank_id="my-bank", query="What does Alice do?")
# Temporal
results = client.recall(bank_id="my-bank", query="What happened in June?")
client.recall(bank_id="my-bank", query="What happened in June?")
```
Recall performs 4 retrieval strategies in parallel:
@@ -182,7 +200,7 @@ Recall performs 4 retrieval strategies in parallel:
- Graph: Entity/temporal/causal links
- Temporal: Time range filtering
![Retain Operation](hindsight-docs/static/img/recall-operation.png)
![Retain Operation](hindsight-docs/static/img/recall-operation.webp)
The individual results from the retrievals are merged, then ordered by relevance using reciprocal rank fusion and a cross-encoder reranking model.
@@ -208,36 +226,29 @@ client = Hindsight(base_url="http://localhost:8888")
client.reflect(bank_id="my-bank", query="What should I know about Alice?")
```
![Retain Operation](hindsight-docs/static/img/reflect-operation.png)
## Integrations
### Examples
[Examples Repo]([./examples](https://github.com/vectorize-io/hindsight-cookbook)) includes:
- Basic usage
- Multi-session conversations
- Temporal queries
- Entity reasoning
- Opinion tracking
- Production setup (Docker Compose + monitoring)
![Retain Operation](hindsight-docs/static/img/reflect-operation.webp)
---
## Resources
**Documentation:** [vectorize-io.github.io/hindsight](https://vectorize-io.github.io/hindsight)
**Documentation:**
- [https://hindsight.vectorize.io](https://hindsight.vectorize.io)
**Clients:**
- [Python](http://hindsight.vectorize.io/sdks/python)
- [Node.js](http://hindsight.vectorize.io/sdks/nodejs)
- [REST API](http://hindsight.vectorize.io/api-reference)
- [REST API](https://hindsight.vectorize.io/api-reference)
- [CLI](https://hindsight.vectorize.io/sdks/cli)
**Community:**
- [Slack](https://join.slack.com/t/hindsight-space/shared_invite/zt-3klo21kua-VUCC_zHP5rIcXFB1_5yw6A)
- [Slack](https://join.slack.com/t/hindsight-space/shared_invite/zt-3nhbm4w29-LeSJ5Ixi6j8PdiYOCPlOgg)
- [GitHub Issues](https://github.com/vectorize-io/hindsight/issues)
---
## Star History
[![Star History Chart](https://api.star-history.com/svg?repos=vectorize-io/hindsight&type=date&legend=top-left)](https://www.star-history.com/#vectorize-io/hindsight&type=date&legend=top-left)
---
## Contributing
@@ -250,4 +261,4 @@ MIT — see [LICENSE](./LICENSE)
---
Built by [Vectorize.io](https://vectorize.io)
Built by [Vectorize.io](https://vectorize.io)
+102 -111
View File
@@ -2,16 +2,24 @@
# Supports building API-only, Control Plane-only, or both
#
# Build args:
# INCLUDE_API=true/false - Include API (default: true)
# INCLUDE_CP=true/false - Include Control Plane (default: true)
# INCLUDE_API=true/false - Include API (default: true)
# INCLUDE_CP=true/false - Include Control Plane (default: true)
# INCLUDE_LOCAL_MODELS=true/false - Include local ML models for embeddings/reranking (default: true)
# Set to false when using external providers (TEI, OpenAI, Cohere)
# PRELOAD_ML_MODELS=true/false - Pre-download ML models during build (default: true)
# Only effective when INCLUDE_LOCAL_MODELS=true
#
# Examples:
# docker build -t hindsight . # Both (standalone)
# docker build -t hindsight-api --build-arg INCLUDE_CP=false . # API only
# docker build -t hindsight-cp --build-arg INCLUDE_API=false . # Control Plane only
# docker build -t hindsight . # Both (standalone)
# docker build -t hindsight-api --build-arg INCLUDE_CP=false . # API only
# docker build -t hindsight-cp --build-arg INCLUDE_API=false . # Control Plane only
# docker build -t hindsight --build-arg PRELOAD_ML_MODELS=false . # Skip ML model preload
# docker build -t hindsight --build-arg INCLUDE_LOCAL_MODELS=false . # Skip local ML deps (for external providers)
ARG INCLUDE_API=true
ARG INCLUDE_CP=true
ARG PRELOAD_ML_MODELS=true
ARG INCLUDE_LOCAL_MODELS=true
# =============================================================================
# Stage: API Builder
@@ -19,6 +27,7 @@ ARG INCLUDE_CP=true
FROM python:3.11-slim AS api-builder
ARG INCLUDE_API
ARG INCLUDE_LOCAL_MODELS
RUN if [ "$INCLUDE_API" != "true" ]; then echo "Skipping API build" && exit 0; fi
WORKDIR /app
@@ -37,12 +46,23 @@ COPY hindsight-api/README.md ./api/
WORKDIR /app/api
# Remove local ML model dependencies if INCLUDE_LOCAL_MODELS=false
# This creates a smaller image when using external providers (TEI, OpenAI, Cohere)
RUN if [ "$INCLUDE_LOCAL_MODELS" != "true" ]; then \
echo "Removing local-models dependencies (sentence-transformers, torch, transformers)..." && \
sed -i '/"sentence-transformers/d' pyproject.toml && \
sed -i '/"transformers/d' pyproject.toml && \
sed -i '/"torch/d' pyproject.toml; \
fi
# Sync dependencies (will create lock file if needed)
RUN uv sync
# Copy source code and alembic migrations
# Copy source code (alembic migrations are inside hindsight_api/)
COPY hindsight-api/hindsight_api ./hindsight_api
COPY hindsight-api/alembic ./alembic
# Install the local package (uv sync only installed dependencies, not the package itself)
RUN uv pip install -e .
# =============================================================================
# Stage: SDK Builder (needed for Control Plane)
@@ -52,13 +72,15 @@ FROM node:20-slim AS sdk-builder
ARG INCLUDE_CP
RUN if [ "$INCLUDE_CP" != "true" ]; then echo "Skipping SDK build" && exit 0; fi
WORKDIR /app/sdk
WORKDIR /app
COPY hindsight-clients/typescript/package*.json ./
RUN npm ci
# Copy root package files for npm workspaces
COPY package.json package-lock.json ./
COPY hindsight-clients/typescript/ ./hindsight-clients/typescript/
COPY hindsight-clients/typescript/ ./
RUN npm run build
# Install and build SDK using workspace (--ignore-scripts skips git hooks setup)
RUN npm ci --ignore-scripts -w @vectorize-io/hindsight-client
RUN npm run build -w @vectorize-io/hindsight-client
# =============================================================================
# Stage: Control Plane Builder
@@ -68,30 +90,48 @@ FROM node:20-slim AS cp-builder
ARG INCLUDE_CP
RUN if [ "$INCLUDE_CP" != "true" ]; then echo "Skipping CP build" && exit 0; fi
WORKDIR /app
# Copy built SDK
COPY --from=sdk-builder /app/sdk /app/sdk
# Create directory structure matching the monorepo layout
# This is required because build:standalone script expects .next/standalone/memory-poc/hindsight-control-plane
WORKDIR /app/memory-poc/hindsight-control-plane
# Install Control Plane dependencies
# Only copy package.json (not package-lock.json) to ensure npm installs
# correct platform-specific native bindings for lightningcss/tailwindcss
COPY hindsight-control-plane/package.json ./
# Remove the file: dependency on SDK (we'll copy it directly later)
RUN sed -i '/"@vectorize-io\/hindsight-client":/d' package.json
RUN npm install
# Copy Control Plane source (excluding node_modules via .dockerignore)
COPY hindsight-control-plane/ ./
# Remove package-lock.json to avoid conflicts with installed native bindings
RUN rm -f package-lock.json
# Also remove the file: dependency from package.json (restored by COPY above)
RUN rm -f package-lock.json && sed -i '/"@vectorize-io\/hindsight-client":/d' package.json
# Link SDK (temporary for build)
RUN cd /app/sdk && npm link && cd /app && npm link @vectorize-io/hindsight-client
# Copy built SDK directly into node_modules (more reliable than npm link in Docker)
COPY --from=sdk-builder /app/hindsight-clients/typescript ./node_modules/@vectorize-io/hindsight-client
# Build Control Plane
RUN npm run build
# Build Control Plane - run next build first, then custom standalone copy
# (The build:standalone script expects a specific path structure that differs in Docker)
RUN npm exec -- next build
# Create public directory if it doesn't exist
RUN mkdir -p public
# Create standalone directory structure manually
# Note: Must exclude node_modules from find to avoid wrong server.js from next/dist/experimental/testmode/
# Note: Must explicitly copy .next since glob * doesn't match hidden directories
RUN STANDALONE_ROOT=$(find .next/standalone -path '*/node_modules' -prune -o -name 'server.js' -print | head -1 | xargs dirname) && \
mkdir -p standalone && \
cp -r "$STANDALONE_ROOT"/* standalone/ && \
cp -r "$STANDALONE_ROOT"/.next standalone/.next && \
# Copy node_modules if separate from app dir (monorepo structure)
if [ -d ".next/standalone/node_modules" ] && [ "$STANDALONE_ROOT" != ".next/standalone" ]; then \
cp -r .next/standalone/node_modules standalone/node_modules; \
fi && \
cp -r .next/static standalone/.next/static && \
mkdir -p standalone/public && \
cp -r public/* standalone/public/ 2>/dev/null || true && \
# Verify required files exist
test -f standalone/server.js || (echo "ERROR: server.js missing!" && exit 1) && \
test -f standalone/.next/BUILD_ID || (echo "ERROR: BUILD_ID missing!" && exit 1)
# =============================================================================
# Stage: Final Image - API Only
@@ -100,18 +140,18 @@ FROM python:3.11-slim AS api-only
WORKDIR /app
# Install pg0 dependencies
# Note: libicu version varies by Debian version - try common versions in order
RUN apt-get update && apt-get install -y \
curl \
procps \
libxml2 \
libssl3 \
libgssapi-krb5-2 \
libossp-uuid16 \
&& apt-get install -y libicu72 || apt-get install -y libicu74 || apt-get install -y libicu* \
&& (apt-get install -y libicu72 2>/dev/null || apt-get install -y libicu74 2>/dev/null || apt-get install -y libicu76 2>/dev/null || true) \
&& rm -rf /var/lib/apt/lists/* \
&& pip install --no-cache-dir uv
# Create non-root user (PostgreSQL cannot run as root)
RUN useradd -m -s /bin/bash hindsight
# Copy API with virtual environment from builder
@@ -121,52 +161,26 @@ COPY --from=api-builder /app/api /app/api
COPY docker/standalone/start-all.sh /app/start-all.sh
RUN chmod +x /app/start-all.sh
# Create data directory for pg0 and set ownership
RUN mkdir -p /app/data && chown -R hindsight:hindsight /app
RUN chown -R hindsight:hindsight /app
# Switch to non-root user
USER hindsight
# Set PATH for hindsight user
ENV PATH="/home/hindsight/.hindsight/bin:/app/api/.venv/bin:${PATH}"
ENV PATH="/app/api/.venv/bin:${PATH}"
# Install pg0 binary
RUN mkdir -p /home/hindsight/.hindsight/bin && \
ARCH=$(uname -m) && \
if [ "$ARCH" = "aarch64" ] || [ "$ARCH" = "arm64" ]; then \
PG0_BINARY="pg0-linux-aarch64-gnu"; \
elif [ "$ARCH" = "x86_64" ]; then \
PG0_BINARY="pg0-linux-x86_64-gnu"; \
else \
echo "Unsupported architecture: $ARCH" && exit 1; \
fi && \
echo "Installing pg0 binary: $PG0_BINARY" && \
for i in 1 2 3 4 5; do \
curl -fsSL -o /home/hindsight/.hindsight/bin/pg0 \
"https://github.com/vectorize-io/pg0/releases/latest/download/$PG0_BINARY" && \
chmod +x /home/hindsight/.hindsight/bin/pg0 && \
break || (echo "Retry $i failed, waiting..." && sleep 10); \
done && \
/home/hindsight/.hindsight/bin/pg0 --version
# Pre-download PostgreSQL binaries
ENV PG0_HOME=/home/hindsight/.pg0-cache
RUN pg0 start --help && \
(pg0 start --name hindsight --port 5555 --username hindsight --password hindsight --database hindsight && \
sleep 2 && \
pg0 stop --name hindsight && \
echo "PostgreSQL pre-cached to $PG0_HOME") || echo "Pre-download skipped"
ENV PG0_HOME=/home/hindsight/.pg0
# Pre-download ML models to avoid runtime download
RUN /app/api/.venv/bin/python -c "\
# Pre-download ML models to avoid runtime download (conditional)
# Only runs if both PRELOAD_ML_MODELS=true AND INCLUDE_LOCAL_MODELS=true
ARG PRELOAD_ML_MODELS
ARG INCLUDE_LOCAL_MODELS
RUN if [ "$PRELOAD_ML_MODELS" = "true" ] && [ "$INCLUDE_LOCAL_MODELS" = "true" ]; then \
/app/api/.venv/bin/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 cached successfully')"
print('Models cached successfully')"; \
elif [ "$INCLUDE_LOCAL_MODELS" != "true" ]; then echo "Skipping ML model preload (local-models not included)"; \
else echo "Skipping ML model preload"; fi
EXPOSE 8888
@@ -175,6 +189,7 @@ ENV HINDSIGHT_API_PORT=8888
ENV HINDSIGHT_API_LOG_LEVEL=info
ENV HINDSIGHT_ENABLE_API=true
ENV HINDSIGHT_ENABLE_CP=false
ENV PYTHONUNBUFFERED=1
CMD ["/app/start-all.sh"]
@@ -186,13 +201,13 @@ FROM node:20-alpine AS cp-only
WORKDIR /app
# Copy built SDK
COPY --from=sdk-builder /app/sdk /app/sdk
COPY --from=sdk-builder /app/hindsight-clients/typescript /app/sdk
# Copy Control Plane standalone build
WORKDIR /app/control-plane
COPY --from=cp-builder /app/.next/standalone ./
COPY --from=cp-builder /app/.next/static ./.next/static
COPY --from=cp-builder /app/public ./public
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/standalone ./
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/.next/static ./.next/static
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/public ./public
WORKDIR /app
@@ -219,33 +234,34 @@ FROM python:3.11-slim AS standalone
WORKDIR /app
# Install Node.js, curl, uv, and pg0 dependencies
# Install Node.js, curl, uv, and system dependencies
# Note: libicu version varies by Debian version - try common versions in order
RUN apt-get update && apt-get install -y \
curl \
procps \
libxml2 \
libssl3 \
libgssapi-krb5-2 \
libossp-uuid16 \
&& apt-get install -y libicu72 || apt-get install -y libicu74 || apt-get install -y libicu* \
&& (apt-get install -y libicu72 2>/dev/null || apt-get install -y libicu74 2>/dev/null || apt-get install -y libicu76 2>/dev/null || true) \
&& curl -fsSL https://deb.nodesource.com/setup_20.x | bash - \
&& apt-get install -y nodejs \
&& rm -rf /var/lib/apt/lists/* \
&& pip install --no-cache-dir uv
# Create non-root user (PostgreSQL cannot run as root)
RUN useradd -m -s /bin/bash hindsight
# Copy API with virtual environment from builder
COPY --from=api-builder /app/api /app/api
# Copy built SDK
COPY --from=sdk-builder /app/sdk /app/sdk
COPY --from=sdk-builder /app/hindsight-clients/typescript /app/sdk
# Copy Control Plane standalone build
WORKDIR /app/control-plane
COPY --from=cp-builder /app/.next/standalone ./
COPY --from=cp-builder /app/.next/static ./.next/static
COPY --from=cp-builder /app/public ./public
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/standalone ./
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/.next/static ./.next/static
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/public ./public
WORKDIR /app
@@ -253,52 +269,26 @@ WORKDIR /app
COPY docker/standalone/start-all.sh /app/start-all.sh
RUN chmod +x /app/start-all.sh
# Create data directory for pg0 and set ownership
RUN mkdir -p /app/data && chown -R hindsight:hindsight /app
RUN chown -R hindsight:hindsight /app
# Switch to non-root user
USER hindsight
# Set PATH for hindsight user
ENV PATH="/home/hindsight/.hindsight/bin:/app/api/.venv/bin:${PATH}"
ENV PATH="/app/api/.venv/bin:${PATH}"
# Install pg0 binary
RUN mkdir -p /home/hindsight/.hindsight/bin && \
ARCH=$(uname -m) && \
if [ "$ARCH" = "aarch64" ] || [ "$ARCH" = "arm64" ]; then \
PG0_BINARY="pg0-linux-aarch64-gnu"; \
elif [ "$ARCH" = "x86_64" ]; then \
PG0_BINARY="pg0-linux-x86_64-gnu"; \
else \
echo "Unsupported architecture: $ARCH" && exit 1; \
fi && \
echo "Installing pg0 binary: $PG0_BINARY" && \
for i in 1 2 3 4 5; do \
curl -fsSL -o /home/hindsight/.hindsight/bin/pg0 \
"https://github.com/vectorize-io/pg0/releases/latest/download/$PG0_BINARY" && \
chmod +x /home/hindsight/.hindsight/bin/pg0 && \
break || (echo "Retry $i failed, waiting..." && sleep 10); \
done && \
/home/hindsight/.hindsight/bin/pg0 --version
# Pre-download PostgreSQL binaries
ENV PG0_HOME=/home/hindsight/.pg0-cache
RUN pg0 start --help && \
(pg0 start --name hindsight --port 5555 --username hindsight --password hindsight --database hindsight && \
sleep 2 && \
pg0 stop --name hindsight && \
echo "PostgreSQL pre-cached to $PG0_HOME") || echo "Pre-download skipped"
ENV PG0_HOME=/home/hindsight/.pg0
# Pre-download ML models to avoid runtime download
RUN /app/api/.venv/bin/python -c "\
# Pre-download ML models to avoid runtime download (conditional)
# Only runs if both PRELOAD_ML_MODELS=true AND INCLUDE_LOCAL_MODELS=true
ARG PRELOAD_ML_MODELS
ARG INCLUDE_LOCAL_MODELS
RUN if [ "$PRELOAD_ML_MODELS" = "true" ] && [ "$INCLUDE_LOCAL_MODELS" = "true" ]; then \
/app/api/.venv/bin/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 cached successfully')"
print('Models cached successfully')"; \
elif [ "$INCLUDE_LOCAL_MODELS" != "true" ]; then echo "Skipping ML model preload (local-models not included)"; \
else echo "Skipping ML model preload"; fi
EXPOSE 8888 9999
@@ -309,6 +299,7 @@ ENV NODE_ENV=production
ENV HINDSIGHT_CP_DATAPLANE_API_URL=http://localhost:8888
ENV HINDSIGHT_ENABLE_API=true
ENV HINDSIGHT_ENABLE_CP=true
ENV PYTHONUNBUFFERED=1
CMD ["/app/start-all.sh"]
+68 -18
View File
@@ -1,23 +1,74 @@
#!/bin/bash
set -e
echo "🚀 Starting Hindsight..."
echo ""
# Service flags (default to true if not set)
ENABLE_API="${HINDSIGHT_ENABLE_API:-true}"
ENABLE_CP="${HINDSIGHT_ENABLE_CP:-true}"
# Copy pre-cached PostgreSQL data if runtime directory is empty (first run with volume)
if [ "$ENABLE_API" = "true" ]; then
PG0_CACHE="/home/hindsight/.pg0-cache"
PG0_HOME="/home/hindsight/.pg0"
if [ -d "$PG0_CACHE" ] && [ "$(ls -A $PG0_CACHE 2>/dev/null)" ]; then
if [ ! "$(ls -A $PG0_HOME 2>/dev/null)" ]; then
echo "📦 Copying pre-cached PostgreSQL data..."
cp -r "$PG0_CACHE"/* "$PG0_HOME"/ 2>/dev/null || true
fi
# =============================================================================
# Dependency waiting (opt-in via HINDSIGHT_WAIT_FOR_DEPS=true)
#
# Problem: When running with LM Studio, the LLM may take time to load models.
# If Hindsight starts before LM Studio is ready, it fails on LLM verification.
# This wait loop ensures dependencies are ready before starting.
# =============================================================================
if [ "${HINDSIGHT_WAIT_FOR_DEPS:-false}" = "true" ]; then
LLM_BASE_URL="${HINDSIGHT_API_LLM_BASE_URL:-http://host.docker.internal:1234/v1}"
MAX_RETRIES="${HINDSIGHT_RETRY_MAX:-0}" # 0 = infinite
RETRY_INTERVAL="${HINDSIGHT_RETRY_INTERVAL:-10}"
# Check if external database is configured (skip check for embedded pg0)
SKIP_DB_CHECK=false
if [ -z "${HINDSIGHT_API_DATABASE_URL}" ]; then
SKIP_DB_CHECK=true
else
DB_CHECK_HOST=$(echo "$HINDSIGHT_API_DATABASE_URL" | sed -E 's|.*@([^:/]+):([0-9]+)/.*|\1 \2|')
fi
check_db() {
if $SKIP_DB_CHECK; then
return 0
fi
if command -v pg_isready &> /dev/null; then
pg_isready -h $(echo $DB_CHECK_HOST | cut -d' ' -f1) -p $(echo $DB_CHECK_HOST | cut -d' ' -f2) &>/dev/null
else
python3 -c "import socket; s=socket.socket(); s.settimeout(5); exit(0 if s.connect_ex(('$(echo $DB_CHECK_HOST | cut -d' ' -f1)', $(echo $DB_CHECK_HOST | cut -d' ' -f2))) == 0 else 1)" 2>/dev/null
fi
}
check_llm() {
curl -sf "${LLM_BASE_URL}/models" --connect-timeout 5 &>/dev/null
}
echo "⏳ Waiting for dependencies to be ready..."
attempt=1
while true; do
db_ok=false
llm_ok=false
if check_db; then
db_ok=true
fi
if check_llm; then
llm_ok=true
fi
if $db_ok && $llm_ok; then
echo "✅ Dependencies ready!"
break
fi
if [ "$MAX_RETRIES" -ne 0 ] && [ "$attempt" -ge "$MAX_RETRIES" ]; then
echo "❌ Max retries ($MAX_RETRIES) reached. Dependencies not available."
exit 1
fi
echo " Attempt $attempt: DB=$( $db_ok && echo 'ok' || echo 'waiting' ), LLM=$( $llm_ok && echo 'ok' || echo 'waiting' )"
sleep "$RETRY_INTERVAL"
((attempt++))
done
fi
# Track PIDs for wait
@@ -26,32 +77,31 @@ PIDS=()
# Start API if enabled
if [ "$ENABLE_API" = "true" ]; then
cd /app/api
python -m hindsight_api.web.server 2>&1 | sed -u 's/^/[api] /' &
# Run API directly - Python's PYTHONUNBUFFERED=1 handles output buffering
hindsight-api &
API_PID=$!
PIDS+=($API_PID)
# Wait for API to be ready
echo "⏳ Waiting for API..."
for i in {1..60}; do
if curl -sf http://localhost:8888/health &>/dev/null; then
echo "✅ API is ready"
break
fi
sleep 1
done
else
echo "⏭️ API disabled (HINDSIGHT_ENABLE_API=false)"
echo "API disabled (HINDSIGHT_ENABLE_API=false)"
fi
# Start Control Plane if enabled
if [ "$ENABLE_CP" = "true" ]; then
echo "🎛️ Starting Control Plane..."
cd /app/control-plane
PORT=9999 node server.js 2>&1 | grep -v -E "^[[:space:]]*(▲|✓|-|$)" | sed -u 's/^/[control-plane] /' &
PORT=9999 node server.js &
CP_PID=$!
PIDS+=($CP_PID)
else
echo "⏭️ Control Plane disabled (HINDSIGHT_ENABLE_CP=false)"
echo "Control Plane disabled (HINDSIGHT_ENABLE_CP=false)"
fi
# Print status
-135
View File
@@ -1,135 +0,0 @@
HINDSIGHT HELM CHART INSTALLATION GUIDE
=====================================
PREREQUISITES
-------------
- Kubernetes cluster (1.19+)
- kubectl configured
- Helm 3.x installed
- PostgreSQL database with pgvector extension (if not using bundled PostgreSQL)
BASIC INSTALLATION
------------------
1. Install with default values (requires external PostgreSQL):
helm install hindsight ./hindsight \
--set postgresql.external.host=your-postgres-host \
--set postgresql.external.password=your-password \
--set api.secrets.MEMORY_LLM_API_KEY=your-api-key
2. Install with custom values file:
helm install hindsight ./hindsight -f hindsight/values-production.yaml
3. Install in a specific namespace:
kubectl create namespace hindsight
helm install hindsight ./hindsight -n hindsight
CONFIGURATION OPTIONS
---------------------
Development setup (using values-development.yaml):
helm install hindsight ./hindsight -f hindsight/values-development.yaml
Production setup (using values-production.yaml):
helm install hindsight ./hindsight -f hindsight/values-production.yaml
Custom LLM provider:
helm install hindsight ./hindsight \
--set api.env.MEMORY_LLM_PROVIDER=openai \
--set api.env.MEMORY_LLM_MODEL=gpt-4 \
--set api.secrets.MEMORY_LLM_API_KEY=sk-your-key
Enable ingress:
helm install hindsight ./hindsight \
--set ingress.enabled=true \
--set ingress.hosts[0].host=hindsight.example.com
Enable autoscaling:
helm install hindsight ./hindsight \
--set autoscaling.enabled=true \
--set autoscaling.minReplicas=2 \
--set autoscaling.maxReplicas=10
UPGRADE
-------
Upgrade existing installation:
helm upgrade hindsight ./hindsight
Upgrade with new values:
helm upgrade hindsight ./hindsight -f hindsight/values-production.yaml
UNINSTALL
---------
Remove the Helm release:
helm uninstall hindsight
Remove with namespace:
helm uninstall hindsight -n hindsight
TESTING
-------
Test the installation with dry-run:
helm install hindsight ./hindsight --dry-run --debug
Validate templates:
helm template hindsight ./hindsight
Lint the chart:
helm lint ./hindsight
ACCESSING THE SERVICES
----------------------
Port-forward control plane:
kubectl port-forward svc/hindsight-control-plane 3000:3000
Port-forward API:
kubectl port-forward svc/hindsight-api 8888:8888
Get service URLs:
helm status hindsight
DATABASE INITIALIZATION
-----------------------
NOTE: Database migrations now run automatically when the API service starts.
You typically don't need to run migrations manually.
If you want to pre-initialize the database before deploying (optional):
kubectl run hindsight-init --rm -it --restart=Never \
--image=hindsight/api:latest \
--env="DATABASE_URL=postgresql://user:pass@host:5432/hindsight" \
-- python -c "from hindsight.migrations import run_migrations; run_migrations()"
TROUBLESHOOTING
---------------
Check pod status:
kubectl get pods -l app.kubernetes.io/name=hindsight
View logs for API:
kubectl logs -l app.kubernetes.io/component=api
View logs for control plane:
kubectl logs -l app.kubernetes.io/component=control-plane
Describe a pod:
kubectl describe pod <pod-name>
Check configuration:
kubectl get configmap hindsight-config -o yaml
kubectl get secret hindsight-secret -o yaml
NOTES
-----
- Make sure PostgreSQL has pgvector extension enabled
- Run database migrations before first use
- Configure proper resource limits for production
- Use external secrets management for production
- Enable TLS/SSL for production deployments
+6
View File
@@ -0,0 +1,6 @@
dependencies:
- name: postgresql
repository: https://charts.bitnami.com/bitnami
version: 15.5.38
digest: sha256:f67c7612736803ece8a669f8ca6b0555f3b78557bc0ecb732aa2e43f0df7750d
generated: "2025-12-10T17:20:57.058794+01:00"
+3 -3
View File
@@ -1,9 +1,9 @@
apiVersion: v2
name: hindsight
description: A Helm chart for Hindsight - temporal-semantic-entity memory system for AI agents
description: Hindsight helm chart
type: application
version: 0.1.0
appVersion: "0.1.0"
version: 0.3.0
appVersion: "0.3.0"
keywords:
- ai
- memory
+182
View File
@@ -0,0 +1,182 @@
# Hindsight Helm Chart
Helm chart for deploying Hindsight - a temporal-semantic-entity memory system for AI agents.
## Prerequisites
- Kubernetes 1.19+
- Helm 3.0+
- PostgreSQL database (external or bundled)
## Quick Start
```bash
# Update dependencies first
helm dependency update ./helm/hindsight
# Install (PostgreSQL included by default)
export OPENAI_API_KEY="sk-your-openai-key"
helm upgrade hindsight --install ./helm/hindsight -n hindsight --create-namespace \
--set api.secrets.HINDSIGHT_API_LLM_API_KEY="$OPENAI_API_KEY"
```
To use an external database instead:
```bash
helm install hindsight ./helm/hindsight -n hindsight --create-namespace \
--set api.secrets.HINDSIGHT_API_LLM_API_KEY="sk-your-openai-key" \
--set postgresql.enabled=false \
--set postgresql.external.host=my-postgres.example.com \
--set postgresql.external.password=mypassword
```
## Installation
### Add the repository (if published)
```bash
helm repo add hindsight https://your-helm-repo.com
helm repo update
```
### Install with custom values file
Create a `values-override.yaml`:
```yaml
api:
secrets:
HINDSIGHT_API_LLM_API_KEY: "sk-your-openai-key"
postgresql:
external:
host: "my-postgres.example.com"
password: "mypassword"
```
Then install:
```bash
helm install hindsight ./helm/hindsight -n hindsight --create-namespace -f values-override.yaml
```
## Configuration
### Key Values
| Parameter | Description | Default |
|-----------|-------------|---------|
| `version` | Default image tag for all components | `0.1.0` |
| `api.enabled` | Enable the API component | `true` |
| `api.image.repository` | API image repository | `hindsight/api` |
| `api.image.tag` | API image tag (defaults to `version`) | - |
| `api.service.port` | API service port | `8888` |
| `controlPlane.enabled` | Enable the control plane | `true` |
| `controlPlane.image.repository` | Control plane image repository | `hindsight/control-plane` |
| `controlPlane.image.tag` | Control plane image tag (defaults to `version`) | - |
| `controlPlane.service.port` | Control plane service port | `3000` |
| `postgresql.enabled` | Deploy PostgreSQL as subchart | `true` |
| `postgresql.external.host` | External PostgreSQL host | `postgresql` |
| `postgresql.external.port` | External PostgreSQL port | `5432` |
| `postgresql.external.database` | Database name | `hindsight` |
| `postgresql.external.username` | Database username | `hindsight` |
| `ingress.enabled` | Enable ingress | `false` |
| `autoscaling.enabled` | Enable HPA | `false` |
### Environment Variables
All environment variables in `api.env` and `controlPlane.env` are automatically added to the respective pods. Sensitive values should go in `api.secrets` or `controlPlane.secrets`.
```yaml
api:
env:
HINDSIGHT_API_LLM_PROVIDER: "openai"
HINDSIGHT_API_LLM_MODEL: "gpt-4"
secrets:
HINDSIGHT_API_LLM_API_KEY: "your-api-key"
HINDSIGHT_API_LLM_BASE_URL: "https://api.openai.com/v1"
controlPlane:
env:
NODE_ENV: "production"
secrets: {}
```
### External Database
To connect to an external PostgreSQL database:
```yaml
postgresql:
enabled: false
external:
host: "my-postgres.example.com"
port: 5432
database: "hindsight"
username: "hindsight"
password: "your-password"
```
### Ingress
To expose the services via ingress:
```yaml
ingress:
enabled: true
className: "nginx"
annotations:
cert-manager.io/cluster-issuer: "letsencrypt-prod"
hosts:
- host: hindsight.example.com
paths:
- path: /
pathType: Prefix
service: controlPlane
- path: /api
pathType: Prefix
service: api
tls:
- secretName: hindsight-tls
hosts:
- hindsight.example.com
```
## Upgrading
```bash
helm upgrade hindsight ./helm/hindsight -n hindsight
```
## Uninstalling
```bash
helm uninstall hindsight -n hindsight
```
## Components
The chart deploys:
- **API**: The main Hindsight API server for memory operations
- **Control Plane**: Web UI for managing agents and viewing memories
## Development
### Lint the chart
```bash
helm lint ./helm/hindsight
```
### Template locally
```bash
helm template hindsight ./helm/hindsight --debug
```
### Dry run installation
```bash
helm install hindsight ./helm/hindsight --dry-run --debug
```
+2 -71
View File
@@ -1,71 +1,2 @@
Thank you for installing {{ .Chart.Name }}!
Your release is named {{ .Release.Name }}.
To learn more about the release, try:
$ helm status {{ .Release.Name }}
$ helm get all {{ .Release.Name }}
{{- if .Values.ingress.enabled }}
The application is accessible via the following URL(s):
{{- range .Values.ingress.hosts }}
- http{{ if $.Values.ingress.tls }}s{{ end }}://{{ .host }}
{{- end }}
{{- else }}
1. Get the Control Plane URL by running these commands:
{{- if contains "NodePort" .Values.controlPlane.service.type }}
export NODE_PORT=$(kubectl get --namespace {{ .Release.Namespace }} -o jsonpath="{.spec.ports[0].nodePort}" services {{ include "hindsight.fullname" . }}-control-plane)
export NODE_IP=$(kubectl get nodes --namespace {{ .Release.Namespace }} -o jsonpath="{.items[0].status.addresses[0].address}")
echo "Control Plane URL: http://$NODE_IP:$NODE_PORT"
{{- else if contains "LoadBalancer" .Values.controlPlane.service.type }}
NOTE: It may take a few minutes for the LoadBalancer IP to be available.
You can watch the status by running 'kubectl get --namespace {{ .Release.Namespace }} svc -w {{ include "hindsight.fullname" . }}-control-plane'
export SERVICE_IP=$(kubectl get svc --namespace {{ .Release.Namespace }} {{ include "hindsight.fullname" . }}-control-plane --template "{{"{{ range (index .status.loadBalancer.ingress 0) }}{{.}}{{ end }}"}}")
echo "Control Plane URL: http://$SERVICE_IP:{{ .Values.controlPlane.service.port }}"
{{- else if contains "ClusterIP" .Values.controlPlane.service.type }}
export POD_NAME=$(kubectl get pods --namespace {{ .Release.Namespace }} -l "app.kubernetes.io/component=control-plane,app.kubernetes.io/instance={{ .Release.Name }}" -o jsonpath="{.items[0].metadata.name}")
export CONTAINER_PORT=$(kubectl get pod --namespace {{ .Release.Namespace }} $POD_NAME -o jsonpath="{.spec.containers[0].ports[0].containerPort}")
echo "Control Plane URL: http://127.0.0.1:3000"
kubectl --namespace {{ .Release.Namespace }} port-forward $POD_NAME 3000:$CONTAINER_PORT
{{- end }}
2. Get the API URL by running these commands:
{{- if contains "NodePort" .Values.api.service.type }}
export NODE_PORT=$(kubectl get --namespace {{ .Release.Namespace }} -o jsonpath="{.spec.ports[0].nodePort}" services {{ include "hindsight.fullname" . }}-api)
export NODE_IP=$(kubectl get nodes --namespace {{ .Release.Namespace }} -o jsonpath="{.items[0].status.addresses[0].address}")
echo "API URL: http://$NODE_IP:$NODE_PORT"
{{- else if contains "LoadBalancer" .Values.api.service.type }}
NOTE: It may take a few minutes for the LoadBalancer IP to be available.
You can watch the status by running 'kubectl get --namespace {{ .Release.Namespace }} svc -w {{ include "hindsight.fullname" . }}-api'
export SERVICE_IP=$(kubectl get svc --namespace {{ .Release.Namespace }} {{ include "hindsight.fullname" . }}-api --template "{{"{{ range (index .status.loadBalancer.ingress 0) }}{{.}}{{ end }}"}}")
echo "API URL: http://$SERVICE_IP:{{ .Values.api.service.port }}"
{{- else if contains "ClusterIP" .Values.api.service.type }}
export POD_NAME=$(kubectl get pods --namespace {{ .Release.Namespace }} -l "app.kubernetes.io/component=api,app.kubernetes.io/instance={{ .Release.Name }}" -o jsonpath="{.items[0].metadata.name}")
export CONTAINER_PORT=$(kubectl get pod --namespace {{ .Release.Namespace }} $POD_NAME -o jsonpath="{.spec.containers[0].ports[0].containerPort}")
echo "API URL: http://127.0.0.1:8888"
kubectl --namespace {{ .Release.Namespace }} port-forward $POD_NAME 8888:$CONTAINER_PORT
{{- end }}
{{- end }}
{{- if not .Values.postgresql.enabled }}
NOTE: You are using an external PostgreSQL database.
Please ensure that:
1. The database is accessible from the cluster
2. The pgvector extension is enabled
Database migrations run automatically when the API service starts.
If you want to pre-initialize the database before deploying (optional):
kubectl run --namespace {{ .Release.Namespace }} hindsight-init --rm -it --restart=Never \
--image={{ .Values.api.image.repository }}:{{ .Values.api.image.tag }} \
--env="DATABASE_URL={{ include "hindsight.databaseUrl" . }}" \
-- python -c "from hindsight.migrations import run_migrations; run_migrations()"
{{- end }}
For more information, visit: https://github.com/yourusername/hindsight
Hindsight installed. Access the control plane:
kubectl port-forward -n {{ .Release.Namespace }} svc/{{ include "hindsight.fullname" . }}-control-plane 3000:3000
+12 -1
View File
@@ -98,7 +98,7 @@ Generate database URL
{{- if .Values.databaseUrl }}
{{- .Values.databaseUrl }}
{{- else if .Values.postgresql.enabled }}
{{- printf "postgresql://%s:%s@%s-postgresql:%d/%s" .Values.postgresql.auth.username .Values.postgresql.auth.password (include "hindsight.fullname" .) (.Values.postgresql.primary.service.port | int) .Values.postgresql.auth.database }}
{{- printf "postgresql://%s:%s@%s-postgresql:%d/%s" .Values.postgresql.auth.username .Values.postgresql.auth.password (include "hindsight.fullname" .) (.Values.postgresql.service.port | int) .Values.postgresql.auth.database }}
{{- else }}
{{- printf "postgresql://%s:$(POSTGRES_PASSWORD)@%s:%d/%s" .Values.postgresql.external.username .Values.postgresql.external.host (.Values.postgresql.external.port | int) .Values.postgresql.external.database }}
{{- end }}
@@ -110,3 +110,14 @@ API URL for control plane
{{- define "hindsight.apiUrl" -}}
{{- printf "http://%s-api:%d" (include "hindsight.fullname" .) (.Values.api.service.port | int) }}
{{- end }}
{{/*
Get the name of the secret to use
*/}}
{{- define "hindsight.secretName" -}}
{{- if .Values.existingSecret }}
{{- .Values.existingSecret }}
{{- else }}
{{- printf "%s-secret" (include "hindsight.fullname" .) }}
{{- end }}
{{- end }}
+22 -25
View File
@@ -15,8 +15,9 @@ spec:
template:
metadata:
annotations:
checksum/config: {{ include (print $.Template.BasePath "/configmap.yaml") . | sha256sum }}
{{- if not .Values.existingSecret }}
checksum/secret: {{ include (print $.Template.BasePath "/secret.yaml") . | sha256sum }}
{{- end }}
{{- with .Values.podAnnotations }}
{{- toYaml . | nindent 8 }}
{{- end }}
@@ -32,45 +33,41 @@ spec:
- name: api
securityContext:
{{- toYaml .Values.securityContext | nindent 10 }}
image: "{{ .Values.api.image.repository }}:{{ .Values.api.image.tag }}"
image: "{{ .Values.api.image.repository }}:{{ .Values.api.image.tag | default .Values.version }}"
imagePullPolicy: {{ .Values.api.image.pullPolicy }}
ports:
- name: http
containerPort: {{ .Values.api.service.targetPort }}
protocol: TCP
{{- if .Values.existingSecret }}
envFrom:
- secretRef:
name: {{ .Values.existingSecret }}
{{- end }}
env:
- name: HINDSIGHT_API_DATABASE_URL
value: {{ include "hindsight.databaseUrl" . | quote }}
{{- /* POSTGRES_PASSWORD must be defined before DATABASE_URL for $(VAR) interpolation */}}
{{- if not .Values.postgresql.enabled }}
- name: POSTGRES_PASSWORD
valueFrom:
secretKeyRef:
name: {{ include "hindsight.fullname" . }}-secret
name: {{ include "hindsight.secretName" . }}
key: postgres-password
{{- end }}
- name: HINDSIGHT_API_LLM_PROVIDER
valueFrom:
configMapKeyRef:
name: {{ include "hindsight.fullname" . }}-config
key: llm-provider
- name: HINDSIGHT_API_LLM_MODEL
valueFrom:
configMapKeyRef:
name: {{ include "hindsight.fullname" . }}-config
key: llm-model
{{- if and .Values.api.secrets (hasKey .Values.api.secrets "HINDSIGHT_API_LLM_API_KEY") }}
- name: HINDSIGHT_API_LLM_API_KEY
valueFrom:
secretKeyRef:
name: {{ include "hindsight.fullname" . }}-secret
key: llm-api-key
- name: HINDSIGHT_API_DATABASE_URL
value: {{ include "hindsight.databaseUrl" . | quote }}
{{- range $key, $value := .Values.api.env }}
- name: {{ $key }}
value: {{ $value | quote }}
{{- end }}
{{- if and .Values.api.secrets (hasKey .Values.api.secrets "HINDSIGHT_API_LLM_BASE_URL") }}
- name: HINDSIGHT_API_LLM_BASE_URL
{{- /* Only use api.secrets when not using existingSecret (for chart-managed secrets) */}}
{{- if not .Values.existingSecret }}
{{- range $key, $value := .Values.api.secrets }}
- name: {{ $key }}
valueFrom:
secretKeyRef:
name: {{ include "hindsight.fullname" . }}-secret
key: llm-base-url
name: {{ include "hindsight.secretName" $ }}
key: {{ $key }}
{{- end }}
{{- end }}
livenessProbe:
{{- toYaml .Values.api.livenessProbe | nindent 10 }}
-15
View File
@@ -1,15 +0,0 @@
apiVersion: v1
kind: ConfigMap
metadata:
name: {{ include "hindsight.fullname" . }}-config
labels:
{{- include "hindsight.labels" . | nindent 4 }}
data:
# API configuration
llm-provider: {{ .Values.api.env.HINDSIGHT_API_LLM_PROVIDER | quote }}
llm-model: {{ .Values.api.env.HINDSIGHT_API_LLM_MODEL | quote }}
# Control plane configuration
node-env: {{ .Values.controlPlane.env.NODE_ENV | quote }}
hostname: {{ .Values.controlPlane.env.HINDSIGHT_CP_HOSTNAME | quote }}
control-plane-port: {{ .Values.controlPlane.env.HINDSIGHT_CP_PORT | quote }}
@@ -15,7 +15,9 @@ spec:
template:
metadata:
annotations:
checksum/config: {{ include (print $.Template.BasePath "/configmap.yaml") . | sha256sum }}
{{- if not .Values.existingSecret }}
checksum/secret: {{ include (print $.Template.BasePath "/secret.yaml") . | sha256sum }}
{{- end }}
{{- with .Values.podAnnotations }}
{{- toYaml . | nindent 8 }}
{{- end }}
@@ -31,30 +33,34 @@ spec:
- name: control-plane
securityContext:
{{- toYaml .Values.securityContext | nindent 10 }}
image: "{{ .Values.controlPlane.image.repository }}:{{ .Values.controlPlane.image.tag }}"
image: "{{ .Values.controlPlane.image.repository }}:{{ .Values.controlPlane.image.tag | default .Values.version }}"
imagePullPolicy: {{ .Values.controlPlane.image.pullPolicy }}
ports:
- name: http
containerPort: {{ .Values.controlPlane.service.targetPort }}
protocol: TCP
{{- if .Values.existingSecret }}
envFrom:
- secretRef:
name: {{ .Values.existingSecret }}
{{- end }}
env:
- name: NODE_ENV
valueFrom:
configMapKeyRef:
name: {{ include "hindsight.fullname" . }}-config
key: node-env
- name: HINDSIGHT_CP_HOSTNAME
valueFrom:
configMapKeyRef:
name: {{ include "hindsight.fullname" . }}-config
key: hostname
- name: HINDSIGHT_CP_PORT
valueFrom:
configMapKeyRef:
name: {{ include "hindsight.fullname" . }}-config
key: control-plane-port
- name: HINDSIGHT_CP_DATAPLANE_API_URL
value: {{ include "hindsight.apiUrl" . | quote }}
{{- range $key, $value := .Values.controlPlane.env }}
- name: {{ $key }}
value: {{ $value | quote }}
{{- end }}
{{- /* Only use controlPlane.secrets when not using existingSecret (for chart-managed secrets) */}}
{{- if not .Values.existingSecret }}
{{- range $key, $value := .Values.controlPlane.secrets }}
- name: {{ $key }}
valueFrom:
secretKeyRef:
name: {{ include "hindsight.secretName" $ }}
key: {{ $key }}
{{- end }}
{{- end }}
livenessProbe:
{{- toYaml .Values.controlPlane.livenessProbe | nindent 10 }}
readinessProbe:
@@ -0,0 +1,19 @@
{{- if .Values.postgresql.enabled }}
apiVersion: v1
kind: Service
metadata:
name: {{ include "hindsight.fullname" . }}-postgresql
labels:
{{- include "hindsight.labels" . | nindent 4 }}
app.kubernetes.io/component: postgresql
spec:
type: ClusterIP
ports:
- port: {{ .Values.postgresql.service.port }}
targetPort: postgresql
protocol: TCP
name: postgresql
selector:
{{- include "hindsight.selectorLabels" . | nindent 4 }}
app.kubernetes.io/component: postgresql
{{- end }}
@@ -0,0 +1,85 @@
{{- if .Values.postgresql.enabled }}
apiVersion: apps/v1
kind: StatefulSet
metadata:
name: {{ include "hindsight.fullname" . }}-postgresql
labels:
{{- include "hindsight.labels" . | nindent 4 }}
app.kubernetes.io/component: postgresql
spec:
serviceName: {{ include "hindsight.fullname" . }}-postgresql
replicas: 1
selector:
matchLabels:
{{- include "hindsight.selectorLabels" . | nindent 6 }}
app.kubernetes.io/component: postgresql
template:
metadata:
labels:
{{- include "hindsight.selectorLabels" . | nindent 8 }}
app.kubernetes.io/component: postgresql
spec:
containers:
- name: postgresql
image: "{{ .Values.postgresql.image.repository }}:{{ .Values.postgresql.image.tag }}"
imagePullPolicy: {{ .Values.postgresql.image.pullPolicy }}
ports:
- name: postgresql
containerPort: 5432
protocol: TCP
env:
- name: POSTGRES_USER
value: {{ .Values.postgresql.auth.username | quote }}
- name: POSTGRES_PASSWORD
value: {{ .Values.postgresql.auth.password | quote }}
- name: POSTGRES_DB
value: {{ .Values.postgresql.auth.database | quote }}
- name: PGDATA
value: /var/lib/postgresql/data/pgdata
livenessProbe:
exec:
command:
- pg_isready
- -U
- {{ .Values.postgresql.auth.username }}
- -d
- {{ .Values.postgresql.auth.database }}
initialDelaySeconds: 30
periodSeconds: 10
timeoutSeconds: 5
failureThreshold: 3
readinessProbe:
exec:
command:
- pg_isready
- -U
- {{ .Values.postgresql.auth.username }}
- -d
- {{ .Values.postgresql.auth.database }}
initialDelaySeconds: 5
periodSeconds: 5
timeoutSeconds: 3
failureThreshold: 3
resources:
{{- toYaml .Values.postgresql.resources | nindent 10 }}
volumeMounts:
- name: data
mountPath: /var/lib/postgresql/data
{{- if .Values.postgresql.persistence.enabled }}
volumeClaimTemplates:
- metadata:
name: data
spec:
accessModes: ["ReadWriteOnce"]
{{- if .Values.postgresql.persistence.storageClass }}
storageClassName: {{ .Values.postgresql.persistence.storageClass | quote }}
{{- end }}
resources:
requests:
storage: {{ .Values.postgresql.persistence.size }}
{{- else }}
volumes:
- name: data
emptyDir: {}
{{- end }}
{{- end }}
+8 -8
View File
@@ -1,19 +1,19 @@
{{- if not .Values.existingSecret }}
apiVersion: v1
kind: Secret
metadata:
name: {{ include "hindsight.fullname" . }}-secret
name: {{ include "hindsight.secretName" . }}
labels:
{{- include "hindsight.labels" . | nindent 4 }}
type: Opaque
data:
{{- if and .Values.api.secrets (hasKey .Values.api.secrets "MEMORY_LLM_API_KEY") }}
llm-api-key: {{ .Values.api.secrets.MEMORY_LLM_API_KEY | b64enc | quote }}
{{- range $key, $value := .Values.api.secrets }}
{{ $key }}: {{ $value | b64enc | quote }}
{{- end }}
{{- if and .Values.api.secrets (hasKey .Values.api.secrets "MEMORY_LLM_BASE_URL") }}
llm-base-url: {{ .Values.api.secrets.MEMORY_LLM_BASE_URL | b64enc | quote }}
{{- range $key, $value := .Values.controlPlane.secrets }}
{{ $key }}: {{ $value | b64enc | quote }}
{{- end }}
{{- if not .Values.postgresql.enabled }}
{{- if .Values.postgresql.external.password }}
{{- if and (not .Values.postgresql.enabled) .Values.postgresql.external.password }}
postgres-password: {{ .Values.postgresql.external.password | b64enc | quote }}
{{- end }}
{{- end }}
{{- end }}
+50 -18
View File
@@ -1,5 +1,17 @@
# Default values for hindsight
# Chart version - use this to set a consistent image tag across all components
version: "0.1.1"
# Use an existing secret instead of creating one from values
# When set, all keys from this secret are injected as environment variables via envFrom
# Required keys:
# - postgres-password: PostgreSQL password (when postgresql.enabled=false)
# Optional keys (any key becomes an env var):
# - HINDSIGHT_API_LLM_API_KEY: API key for LLM provider
# - Any other env vars you want to inject
# existingSecret: "my-hindsight-secret"
# Global settings
replicaCount: 1
@@ -8,9 +20,9 @@ api:
enabled: true
replicaCount: 1
image:
repository: hindsight/api
repository: ghcr.io/vectorize-io/hindsight-api
pullPolicy: IfNotPresent
tag: "latest"
# tag defaults to .Values.version if not specified
service:
type: ClusterIP
@@ -29,7 +41,7 @@ api:
# Liveness and readiness probes
livenessProbe:
httpGet:
path: /
path: /health
port: 8888
initialDelaySeconds: 30
periodSeconds: 10
@@ -38,7 +50,7 @@ api:
readinessProbe:
httpGet:
path: /
path: /health
port: 8888
initialDelaySeconds: 10
periodSeconds: 5
@@ -47,7 +59,7 @@ api:
# Environment variables
env:
HINDSIGHT_API_LLM_PROVIDER: "groq"
#HINDSIGHT_API_LLM_PROVIDER: "groq"
HINDSIGHT_API_LLM_MODEL: "openai/gpt-oss-120b"
# Secret environment variables
@@ -60,9 +72,9 @@ controlPlane:
enabled: true
replicaCount: 1
image:
repository: hindsight/hindsight-control-plane
repository: ghcr.io/vectorize-io/hindsight-control-plane
pullPolicy: IfNotPresent
tag: "latest"
# tag defaults to .Values.version if not specified
service:
type: ClusterIP
@@ -78,10 +90,9 @@ controlPlane:
cpu: 250m
memory: 512Mi
# Liveness and readiness probes
# Liveness and readiness probes (TCP check)
livenessProbe:
httpGet:
path: /
tcpSocket:
port: 3000
initialDelaySeconds: 30
periodSeconds: 10
@@ -89,8 +100,7 @@ controlPlane:
failureThreshold: 3
readinessProbe:
httpGet:
path: /
tcpSocket:
port: 3000
initialDelaySeconds: 10
periodSeconds: 5
@@ -106,21 +116,43 @@ controlPlane:
# PostgreSQL configuration
postgresql:
# Set to true to deploy PostgreSQL as part of this chart
enabled: false
enabled: true
image:
repository: ankane/pgvector
tag: latest
pullPolicy: IfNotPresent
auth:
username: "hindsight"
password: "hindsight"
database: "hindsight"
service:
port: 5432
persistence:
enabled: true
size: 8Gi
# storageClass: ""
resources:
limits:
cpu: 1000m
memory: 1Gi
requests:
cpu: 250m
memory: 256Mi
# External PostgreSQL connection details
# If postgresql.enabled is false, provide external database details
# Only used if postgresql.enabled is false
external:
host: "postgresql"
port: 5432
database: "hindsight"
username: "hindsight"
# Password should be provided via secret
# password: ""
# Database URL (auto-generated from postgresql config if not provided)
# databaseUrl: "postgresql://user:pass@host:5432/database"
# Ingress configuration
ingress:
enabled: false
+137 -1
View File
@@ -1 +1,137 @@
# Memory
# Hindsight API
**Memory System for AI Agents** — Temporal + Semantic + Entity Memory Architecture using PostgreSQL with pgvector.
Hindsight gives AI agents persistent memory that works like human memory: it stores facts, tracks entities and relationships, handles temporal reasoning ("what happened last spring?"), and forms opinions based on configurable disposition traits.
## Installation
```bash
pip install hindsight-api
```
## Quick Start
### Run the Server
```bash
# Set your LLM provider
export HINDSIGHT_API_LLM_PROVIDER=openai
export HINDSIGHT_API_LLM_API_KEY=sk-xxxxxxxxxxxx
# Start the server (uses embedded PostgreSQL by default)
hindsight-api
```
The server starts at http://localhost:8888 with:
- REST API for memory operations
- MCP server at `/mcp` for tool-use integration
### Use the Python API
```python
from hindsight_api import MemoryEngine
# Create and initialize the memory engine
memory = MemoryEngine()
await memory.initialize()
# Create a memory bank for your agent
bank = await memory.create_memory_bank(
name="my-assistant",
background="A helpful coding assistant"
)
# Store a memory
await memory.retain(
memory_bank_id=bank.id,
content="The user prefers Python for data science projects"
)
# Recall memories
results = await memory.recall(
memory_bank_id=bank.id,
query="What programming language does the user prefer?"
)
# Reflect with reasoning
response = await memory.reflect(
memory_bank_id=bank.id,
query="Should I recommend Python or R for this ML project?"
)
```
## CLI Options
```bash
hindsight-api --help
# Common options
hindsight-api --port 9000 # Custom port (default: 8888)
hindsight-api --host 127.0.0.1 # Bind to localhost only
hindsight-api --workers 4 # Multiple worker processes
hindsight-api --log-level debug # Verbose logging
```
## Configuration
Configure via environment variables:
| Variable | Description | Default |
|----------|-------------|---------|
| `HINDSIGHT_API_DATABASE_URL` | PostgreSQL connection string | `pg0` (embedded) |
| `HINDSIGHT_API_LLM_PROVIDER` | `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio` | `openai` |
| `HINDSIGHT_API_LLM_API_KEY` | API key for LLM provider | - |
| `HINDSIGHT_API_LLM_MODEL` | Model name | `gpt-4o-mini` |
| `HINDSIGHT_API_HOST` | Server bind address | `0.0.0.0` |
| `HINDSIGHT_API_PORT` | Server port | `8888` |
### Example with External PostgreSQL
```bash
export HINDSIGHT_API_DATABASE_URL=postgresql://user:pass@localhost:5432/hindsight
export HINDSIGHT_API_LLM_PROVIDER=groq
export HINDSIGHT_API_LLM_API_KEY=gsk_xxxxxxxxxxxx
hindsight-api
```
## Docker
```bash
docker run --rm -it -p 8888:8888 \
-e HINDSIGHT_API_LLM_API_KEY=$OPENAI_API_KEY \
-v $HOME/.hindsight-docker:/home/hindsight/.pg0 \
ghcr.io/vectorize-io/hindsight:latest
```
## MCP Server
For local MCP integration without running the full API server:
```bash
hindsight-local-mcp
```
This runs a stdio-based MCP server that can be used directly with MCP-compatible clients.
## Key Features
- **Multi-Strategy Retrieval (TEMPR)** — Semantic, keyword, graph, and temporal search combined with RRF fusion
- **Entity Graph** — Automatic entity extraction and relationship tracking
- **Temporal Reasoning** — Native support for time-based queries
- **Disposition Traits** — Configurable skepticism, literalism, and empathy influence opinion formation
- **Three Memory Types** — World facts, bank actions, and formed opinions with confidence scores
## Documentation
Full documentation: [https://hindsight.vectorize.io](https://hindsight.vectorize.io)
- [Installation Guide](https://hindsight.vectorize.io/developer/installation)
- [Configuration Reference](https://hindsight.vectorize.io/developer/configuration)
- [API Reference](https://hindsight.vectorize.io/api-reference)
- [Python SDK](https://hindsight.vectorize.io/sdks/python)
## License
Apache 2.0
@@ -1,274 +0,0 @@
"""initial_schema
Revision ID: 5a366d414dce
Revises:
Create Date: 2025-11-27 11:54:19.228030
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
from pgvector.sqlalchemy import Vector
# revision identifiers, used by Alembic.
revision: str = '5a366d414dce'
down_revision: Union[str, Sequence[str], None] = None
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
"""Upgrade schema - create all tables from scratch."""
# Enable required extensions
op.execute('CREATE EXTENSION IF NOT EXISTS vector')
# Create banks table
op.create_table(
'banks',
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('name', sa.Text(), nullable=True),
sa.Column('personality', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.Column('background', sa.Text(), nullable=True),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.PrimaryKeyConstraint('bank_id', name=op.f('pk_banks'))
)
# Create documents table
op.create_table(
'documents',
sa.Column('id', sa.Text(), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('original_text', sa.Text(), nullable=True),
sa.Column('content_hash', sa.Text(), nullable=True),
sa.Column('metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.PrimaryKeyConstraint('id', 'bank_id', name=op.f('pk_documents'))
)
op.create_index('idx_documents_bank_id', 'documents', ['bank_id'])
op.create_index('idx_documents_content_hash', 'documents', ['content_hash'])
# Create async_operations table
op.create_table(
'async_operations',
sa.Column('operation_id', postgresql.UUID(as_uuid=True), server_default=sa.text('gen_random_uuid()'), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('operation_type', sa.Text(), nullable=False),
sa.Column('status', sa.Text(), server_default='pending', nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('completed_at', postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column('error_message', sa.Text(), nullable=True),
sa.Column('result_metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.PrimaryKeyConstraint('operation_id', name=op.f('pk_async_operations')),
sa.CheckConstraint("status IN ('pending', 'processing', 'completed', 'failed')", name='async_operations_status_check')
)
op.create_index('idx_async_operations_bank_id', 'async_operations', ['bank_id'])
op.create_index('idx_async_operations_status', 'async_operations', ['status'])
op.create_index('idx_async_operations_bank_status', 'async_operations', ['bank_id', 'status'])
# Create entities table
op.create_table(
'entities',
sa.Column('id', postgresql.UUID(as_uuid=True), server_default=sa.text('gen_random_uuid()'), nullable=False),
sa.Column('canonical_name', sa.Text(), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.Column('first_seen', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('last_seen', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('mention_count', sa.Integer(), server_default='1', nullable=False),
sa.PrimaryKeyConstraint('id', name=op.f('pk_entities'))
)
op.create_index('idx_entities_bank_id', 'entities', ['bank_id'])
op.create_index('idx_entities_canonical_name', 'entities', ['canonical_name'])
op.create_index('idx_entities_bank_name', 'entities', ['bank_id', 'canonical_name'])
# Create unique index on (bank_id, LOWER(canonical_name)) for entity resolution
op.execute('CREATE UNIQUE INDEX idx_entities_bank_lower_name ON entities (bank_id, LOWER(canonical_name))')
# Create memory_units table
op.create_table(
'memory_units',
sa.Column('id', postgresql.UUID(as_uuid=True), server_default=sa.text('gen_random_uuid()'), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('document_id', sa.Text(), nullable=True),
sa.Column('text', sa.Text(), nullable=False),
sa.Column('embedding', Vector(384), nullable=True),
sa.Column('context', sa.Text(), nullable=True),
sa.Column('event_date', postgresql.TIMESTAMP(timezone=True), nullable=False),
sa.Column('occurred_start', postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column('occurred_end', postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column('mentioned_at', postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column('fact_type', sa.Text(), server_default='world', nullable=False),
sa.Column('confidence_score', sa.Float(), nullable=True),
sa.Column('access_count', sa.Integer(), server_default='0', nullable=False),
sa.Column('metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.ForeignKeyConstraint(['document_id', 'bank_id'], ['documents.id', 'documents.bank_id'], name='memory_units_document_fkey', ondelete='CASCADE'),
sa.PrimaryKeyConstraint('id', name=op.f('pk_memory_units')),
sa.CheckConstraint("fact_type IN ('world', 'bank', 'opinion', 'observation')", name='memory_units_fact_type_check'),
sa.CheckConstraint("confidence_score IS NULL OR (confidence_score >= 0.0 AND confidence_score <= 1.0)", name='memory_units_confidence_range_check'),
sa.CheckConstraint(
"(fact_type = 'opinion' AND confidence_score IS NOT NULL) OR "
"(fact_type = 'observation') OR "
"(fact_type NOT IN ('opinion', 'observation') AND confidence_score IS NULL)",
name='confidence_score_fact_type_check'
)
)
# Add search_vector column for full-text search
op.execute("""
ALTER TABLE memory_units
ADD COLUMN search_vector tsvector
GENERATED ALWAYS AS (to_tsvector('english', COALESCE(text, '') || ' ' || COALESCE(context, ''))) STORED
""")
op.create_index('idx_memory_units_bank_id', 'memory_units', ['bank_id'])
op.create_index('idx_memory_units_document_id', 'memory_units', ['document_id'])
op.create_index('idx_memory_units_event_date', 'memory_units', [sa.text('event_date DESC')])
op.create_index('idx_memory_units_bank_date', 'memory_units', ['bank_id', sa.text('event_date DESC')])
op.create_index('idx_memory_units_access_count', 'memory_units', [sa.text('access_count DESC')])
op.create_index('idx_memory_units_fact_type', 'memory_units', ['fact_type'])
op.create_index('idx_memory_units_bank_fact_type', 'memory_units', ['bank_id', 'fact_type'])
op.create_index('idx_memory_units_bank_type_date', 'memory_units', ['bank_id', 'fact_type', sa.text('event_date DESC')])
op.create_index('idx_memory_units_opinion_confidence', 'memory_units', ['bank_id', sa.text('confidence_score DESC')], postgresql_where=sa.text("fact_type = 'opinion'"))
op.create_index('idx_memory_units_opinion_date', 'memory_units', ['bank_id', sa.text('event_date DESC')], postgresql_where=sa.text("fact_type = 'opinion'"))
op.create_index('idx_memory_units_observation_date', 'memory_units', ['bank_id', sa.text('event_date DESC')], postgresql_where=sa.text("fact_type = 'observation'"))
op.create_index('idx_memory_units_embedding', 'memory_units', ['embedding'], postgresql_using='hnsw', postgresql_ops={'embedding': 'vector_cosine_ops'})
# Create BM25 full-text search index on search_vector
op.execute("""
CREATE INDEX idx_memory_units_text_search ON memory_units
USING gin(search_vector)
""")
op.execute("""
CREATE MATERIALIZED VIEW memory_units_bm25 AS
SELECT
id,
bank_id,
text,
to_tsvector('english', text) AS text_vector,
log(1.0 + length(text)::float / (SELECT avg(length(text)) FROM memory_units)) AS doc_length_factor
FROM memory_units
""")
op.create_index('idx_memory_units_bm25_bank', 'memory_units_bm25', ['bank_id'])
op.create_index('idx_memory_units_bm25_text_vector', 'memory_units_bm25', ['text_vector'], postgresql_using='gin')
# Create entity_cooccurrences table
op.create_table(
'entity_cooccurrences',
sa.Column('entity_id_1', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('entity_id_2', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('cooccurrence_count', sa.Integer(), server_default='1', nullable=False),
sa.Column('last_cooccurred', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.ForeignKeyConstraint(['entity_id_1'], ['entities.id'], name=op.f('fk_entity_cooccurrences_entity_id_1_entities'), ondelete='CASCADE'),
sa.ForeignKeyConstraint(['entity_id_2'], ['entities.id'], name=op.f('fk_entity_cooccurrences_entity_id_2_entities'), ondelete='CASCADE'),
sa.PrimaryKeyConstraint('entity_id_1', 'entity_id_2', name=op.f('pk_entity_cooccurrences')),
sa.CheckConstraint('entity_id_1 < entity_id_2', name='entity_cooccurrence_order_check')
)
op.create_index('idx_entity_cooccurrences_entity1', 'entity_cooccurrences', ['entity_id_1'])
op.create_index('idx_entity_cooccurrences_entity2', 'entity_cooccurrences', ['entity_id_2'])
op.create_index('idx_entity_cooccurrences_count', 'entity_cooccurrences', [sa.text('cooccurrence_count DESC')])
# Create memory_links table
op.create_table(
'memory_links',
sa.Column('from_unit_id', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('to_unit_id', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('link_type', sa.Text(), nullable=False),
sa.Column('entity_id', postgresql.UUID(as_uuid=True), nullable=True),
sa.Column('weight', sa.Float(), server_default='1.0', nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.ForeignKeyConstraint(['entity_id'], ['entities.id'], name=op.f('fk_memory_links_entity_id_entities'), ondelete='CASCADE'),
sa.ForeignKeyConstraint(['from_unit_id'], ['memory_units.id'], name=op.f('fk_memory_links_from_unit_id_memory_units'), ondelete='CASCADE'),
sa.ForeignKeyConstraint(['to_unit_id'], ['memory_units.id'], name=op.f('fk_memory_links_to_unit_id_memory_units'), ondelete='CASCADE'),
sa.CheckConstraint("link_type IN ('temporal', 'semantic', 'entity', 'causes', 'caused_by', 'enables', 'prevents')", name='memory_links_link_type_check'),
sa.CheckConstraint('weight >= 0.0 AND weight <= 1.0', name='memory_links_weight_check')
)
# Create unique constraint using COALESCE for nullable entity_id
op.execute("CREATE UNIQUE INDEX idx_memory_links_unique ON memory_links (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid))")
op.create_index('idx_memory_links_from_unit', 'memory_links', ['from_unit_id'])
op.create_index('idx_memory_links_to_unit', 'memory_links', ['to_unit_id'])
op.create_index('idx_memory_links_entity', 'memory_links', ['entity_id'])
op.create_index('idx_memory_links_link_type', 'memory_links', ['link_type'])
# Create unit_entities table
op.create_table(
'unit_entities',
sa.Column('unit_id', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('entity_id', postgresql.UUID(as_uuid=True), nullable=False),
sa.ForeignKeyConstraint(['entity_id'], ['entities.id'], name=op.f('fk_unit_entities_entity_id_entities'), ondelete='CASCADE'),
sa.ForeignKeyConstraint(['unit_id'], ['memory_units.id'], name=op.f('fk_unit_entities_unit_id_memory_units'), ondelete='CASCADE'),
sa.PrimaryKeyConstraint('unit_id', 'entity_id', name=op.f('pk_unit_entities'))
)
op.create_index('idx_unit_entities_unit', 'unit_entities', ['unit_id'])
op.create_index('idx_unit_entities_entity', 'unit_entities', ['entity_id'])
def downgrade() -> None:
"""Downgrade schema - drop all tables."""
# Drop tables in reverse dependency order
op.drop_index('idx_unit_entities_entity', table_name='unit_entities')
op.drop_index('idx_unit_entities_unit', table_name='unit_entities')
op.drop_table('unit_entities')
op.drop_index('idx_memory_links_link_type', table_name='memory_links')
op.drop_index('idx_memory_links_entity', table_name='memory_links')
op.drop_index('idx_memory_links_to_unit', table_name='memory_links')
op.drop_index('idx_memory_links_from_unit', table_name='memory_links')
op.execute('DROP INDEX IF EXISTS idx_memory_links_unique')
op.drop_table('memory_links')
op.drop_index('idx_entity_cooccurrences_count', table_name='entity_cooccurrences')
op.drop_index('idx_entity_cooccurrences_entity2', table_name='entity_cooccurrences')
op.drop_index('idx_entity_cooccurrences_entity1', table_name='entity_cooccurrences')
op.drop_table('entity_cooccurrences')
# Drop BM25 materialized view and index
op.drop_index('idx_memory_units_bm25_text_vector', table_name='memory_units_bm25')
op.drop_index('idx_memory_units_bm25_bank', table_name='memory_units_bm25')
op.execute('DROP MATERIALIZED VIEW IF EXISTS memory_units_bm25')
op.drop_index('idx_memory_units_embedding', table_name='memory_units')
op.drop_index('idx_memory_units_observation_date', table_name='memory_units')
op.drop_index('idx_memory_units_opinion_date', table_name='memory_units')
op.drop_index('idx_memory_units_opinion_confidence', table_name='memory_units')
op.drop_index('idx_memory_units_bank_type_date', table_name='memory_units')
op.drop_index('idx_memory_units_bank_fact_type', table_name='memory_units')
op.drop_index('idx_memory_units_fact_type', table_name='memory_units')
op.drop_index('idx_memory_units_access_count', table_name='memory_units')
op.drop_index('idx_memory_units_bank_date', table_name='memory_units')
op.drop_index('idx_memory_units_event_date', table_name='memory_units')
op.drop_index('idx_memory_units_document_id', table_name='memory_units')
op.drop_index('idx_memory_units_bank_id', table_name='memory_units')
op.execute('DROP INDEX IF EXISTS idx_memory_units_text_search')
op.drop_table('memory_units')
op.execute('DROP INDEX IF EXISTS idx_entities_bank_lower_name')
op.drop_index('idx_entities_bank_name', table_name='entities')
op.drop_index('idx_entities_canonical_name', table_name='entities')
op.drop_index('idx_entities_bank_id', table_name='entities')
op.drop_table('entities')
op.drop_index('idx_async_operations_bank_status', table_name='async_operations')
op.drop_index('idx_async_operations_status', table_name='async_operations')
op.drop_index('idx_async_operations_bank_id', table_name='async_operations')
op.drop_table('async_operations')
op.drop_index('idx_documents_content_hash', table_name='documents')
op.drop_index('idx_documents_bank_id', table_name='documents')
op.drop_table('documents')
op.drop_table('banks')
# Drop extensions (optional - comment out if you want to keep them)
# op.execute('DROP EXTENSION IF EXISTS vector')
# op.execute('DROP EXTENSION IF EXISTS "uuid-ossp"')
@@ -1,70 +0,0 @@
"""add_chunks_table
Revision ID: b7c4d8e9f1a2
Revises: 5a366d414dce
Create Date: 2025-11-28 00:00:00.000000
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = 'b7c4d8e9f1a2'
down_revision: Union[str, Sequence[str], None] = '5a366d414dce'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
"""Add chunks table and link memory_units to chunks."""
# Create chunks table with single text PK (bank_id_document_id_chunk_index)
op.create_table(
'chunks',
sa.Column('chunk_id', sa.Text(), nullable=False),
sa.Column('document_id', sa.Text(), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('chunk_index', sa.Integer(), nullable=False),
sa.Column('chunk_text', sa.Text(), nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.ForeignKeyConstraint(['document_id', 'bank_id'], ['documents.id', 'documents.bank_id'], name='chunks_document_fkey', ondelete='CASCADE'),
sa.PrimaryKeyConstraint('chunk_id', name=op.f('pk_chunks'))
)
# Add indexes for efficient queries
op.create_index('idx_chunks_document_id', 'chunks', ['document_id'])
op.create_index('idx_chunks_bank_id', 'chunks', ['bank_id'])
# Add chunk_id column to memory_units (nullable, as existing records won't have chunks)
op.add_column('memory_units', sa.Column('chunk_id', sa.Text(), nullable=True))
# Add foreign key constraint to chunks table
op.create_foreign_key(
'memory_units_chunk_fkey',
'memory_units',
'chunks',
['chunk_id'],
['chunk_id'],
ondelete='SET NULL'
)
# Add index on chunk_id for efficient lookups
op.create_index('idx_memory_units_chunk_id', 'memory_units', ['chunk_id'])
def downgrade() -> None:
"""Remove chunks table and chunk_id from memory_units."""
# Drop index and foreign key from memory_units
op.drop_index('idx_memory_units_chunk_id', table_name='memory_units')
op.drop_constraint('memory_units_chunk_fkey', 'memory_units', type_='foreignkey')
op.drop_column('memory_units', 'chunk_id')
# Drop chunks table indexes and table
op.drop_index('idx_chunks_bank_id', table_name='chunks')
op.drop_index('idx_chunks_document_id', table_name='chunks')
op.drop_table('chunks')
@@ -1,48 +0,0 @@
"""Rename fact_type 'bank' to 'experience'
Revision ID: d9f6a3b4c5e2
Revises: c8e5f2a3b4d1
Create Date: 2024-12-04 15:00:00.000000
"""
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision = 'd9f6a3b4c5e2'
down_revision = 'c8e5f2a3b4d1'
branch_labels = None
depends_on = None
def upgrade():
# Drop old check constraint FIRST (before updating data)
op.drop_constraint('memory_units_fact_type_check', 'memory_units', type_='check')
# Update existing 'bank' values to 'experience'
op.execute("UPDATE memory_units SET fact_type = 'experience' WHERE fact_type = 'bank'")
# Also update any 'interactions' values (in case of partial migration)
op.execute("UPDATE memory_units SET fact_type = 'experience' WHERE fact_type = 'interactions'")
# Create new check constraint with 'experience' instead of 'bank'
op.create_check_constraint(
'memory_units_fact_type_check',
'memory_units',
"fact_type IN ('world', 'experience', 'opinion', 'observation')"
)
def downgrade():
# Drop new check constraint FIRST
op.drop_constraint('memory_units_fact_type_check', 'memory_units', type_='check')
# Update 'experience' back to 'bank'
op.execute("UPDATE memory_units SET fact_type = 'bank' WHERE fact_type = 'experience'")
# Recreate old check constraint
op.create_check_constraint(
'memory_units_fact_type_check',
'memory_units',
"fact_type IN ('world', 'bank', 'opinion', 'observation')"
)
@@ -1,62 +0,0 @@
"""disposition_to_3_traits
Revision ID: e0a1b2c3d4e5
Revises: rename_personality
Create Date: 2024-12-08
Migrate disposition traits from Big Five (openness, conscientiousness, extraversion,
agreeableness, neuroticism, bias_strength with 0-1 float values) to the new 3-trait
system (skepticism, literalism, empathy with 1-5 integer values).
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'e0a1b2c3d4e5'
down_revision: Union[str, Sequence[str], None] = 'rename_personality'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
"""Convert Big Five disposition to 3-trait disposition."""
conn = op.get_bind()
# Update all existing banks to use the new disposition format
# Convert from old format to new format with reasonable mappings:
# - skepticism: derived from inverse of agreeableness (skeptical people are less agreeable)
# - literalism: derived from conscientiousness (detail-oriented people are more literal)
# - empathy: derived from agreeableness + inverse of neuroticism
# Default all to 3 (neutral) for simplicity
conn.execute(sa.text("""
UPDATE banks
SET disposition = '{"skepticism": 3, "literalism": 3, "empathy": 3}'::jsonb
WHERE disposition IS NOT NULL
"""))
# Update the default for new banks
conn.execute(sa.text("""
ALTER TABLE banks
ALTER COLUMN disposition SET DEFAULT '{"skepticism": 3, "literalism": 3, "empathy": 3}'::jsonb
"""))
def downgrade() -> None:
"""Convert back to Big Five disposition."""
conn = op.get_bind()
# Revert to Big Five format with default values
conn.execute(sa.text("""
UPDATE banks
SET disposition = '{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}'::jsonb
WHERE disposition IS NOT NULL
"""))
# Update the default for new banks
conn.execute(sa.text("""
ALTER TABLE banks
ALTER COLUMN disposition SET DEFAULT '{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}'::jsonb
"""))
@@ -1,65 +0,0 @@
"""rename_personality_to_disposition
Revision ID: rename_personality
Revises: d9f6a3b4c5e2
Create Date: 2024-12-04
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = 'rename_personality'
down_revision: Union[str, Sequence[str], None] = 'd9f6a3b4c5e2'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
"""Rename personality column to disposition in banks table (if it exists)."""
conn = op.get_bind()
# Check if 'personality' column exists (old database)
result = conn.execute(sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_name = 'banks' AND column_name = 'personality'
"""))
has_personality = result.fetchone() is not None
# Check if 'disposition' column exists (new database)
result = conn.execute(sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_name = 'banks' AND column_name = 'disposition'
"""))
has_disposition = result.fetchone() is not None
if has_personality and not has_disposition:
# Old database: rename personality -> disposition
op.alter_column('banks', 'personality', new_column_name='disposition')
elif not has_personality and not has_disposition:
# Neither exists (shouldn't happen, but be safe): add disposition column
op.add_column('banks', sa.Column(
'disposition',
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False
))
# else: disposition already exists, nothing to do
def downgrade() -> None:
"""Revert disposition column back to personality."""
conn = op.get_bind()
result = conn.execute(sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_name = 'banks' AND column_name = 'disposition'
"""))
if result.fetchone():
op.alter_column('banks', 'disposition', new_column_name='personality')
+14 -8
View File
@@ -3,25 +3,31 @@ Memory System for AI Agents.
Temporal + Semantic Memory Architecture using PostgreSQL with pgvector.
"""
from .config import HindsightConfig, get_config
from .engine.cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
from .engine.embeddings import Embeddings, LocalSTEmbeddings, RemoteTEIEmbeddings
from .engine.llm_wrapper import LLMConfig
from .engine.memory_engine import MemoryEngine
from .engine.search.trace import (
SearchTrace,
QueryInfo,
EntryPoint,
NodeVisit,
WeightComponents,
LinkInfo,
NodeVisit,
PruningDecision,
SearchSummary,
QueryInfo,
SearchPhaseMetrics,
SearchSummary,
SearchTrace,
WeightComponents,
)
from .engine.search.tracer import SearchTracer
from .engine.embeddings import Embeddings, LocalSTEmbeddings, RemoteTEIEmbeddings
from .engine.cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
from .engine.llm_wrapper import LLMConfig
from .models import RequestContext
__all__ = [
"MemoryEngine",
"RequestContext",
"HindsightConfig",
"get_config",
"SearchTrace",
"SearchTracer",
"QueryInfo",
@@ -0,0 +1 @@
# Admin CLI for Hindsight
+252
View File
@@ -0,0 +1,252 @@
"""
Hindsight Admin CLI - backup and restore operations.
"""
import asyncio
import io
import json
import logging
import zipfile
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
import asyncpg
import typer
from ..config import HindsightConfig
from ..pg0 import parse_pg0_url, resolve_database_url
def _fq_table(table: str, schema: str) -> str:
"""Get fully-qualified table name with schema prefix."""
return f"{schema}.{table}"
# Setup logging
logging.basicConfig(
level=logging.INFO,
format="%(message)s",
)
logger = logging.getLogger(__name__)
app = typer.Typer(name="hindsight-admin", help="Hindsight administrative commands")
# Tables to backup/restore in dependency order
# Import must happen in this order due to foreign key constraints
BACKUP_TABLES = [
"banks",
"documents",
"entities",
"chunks",
"memory_units",
"unit_entities",
"entity_cooccurrences",
"memory_links",
]
MANIFEST_VERSION = "1"
async def _backup(database_url: str, output_path: Path, schema: str = "public") -> dict[str, Any]:
"""Backup all tables to a zip file using binary COPY protocol."""
conn = await asyncpg.connect(database_url)
try:
tables: dict[str, Any] = {}
manifest: dict[str, Any] = {
"version": MANIFEST_VERSION,
"created_at": datetime.now(timezone.utc).isoformat(),
"schema": schema,
"tables": tables,
}
# Use a transaction with REPEATABLE READ isolation to get a consistent
# snapshot across all tables. This prevents race conditions where
# entity_cooccurrences could reference entities created after the
# entities table was backed up.
async with conn.transaction(isolation="repeatable_read"):
with zipfile.ZipFile(output_path, "w", zipfile.ZIP_DEFLATED) as zf:
for i, table in enumerate(BACKUP_TABLES, 1):
typer.echo(f" [{i}/{len(BACKUP_TABLES)}] Backing up {table}...", nl=False)
buffer = io.BytesIO()
# Use binary COPY for exact type preservation
# asyncpg requires schema_name as separate parameter
await conn.copy_from_table(table, schema_name=schema, output=buffer, format="binary")
data = buffer.getvalue()
zf.writestr(f"{table}.bin", data)
# Get row count for manifest
qualified_table = _fq_table(table, schema)
row_count = await conn.fetchval(f"SELECT COUNT(*) FROM {qualified_table}")
tables[table] = {
"rows": row_count,
"size_bytes": len(data),
}
typer.echo(f" {row_count} rows")
zf.writestr("manifest.json", json.dumps(manifest, indent=2))
return manifest
finally:
await conn.close()
async def _restore(database_url: str, input_path: Path, schema: str = "public") -> dict[str, Any]:
"""Restore all tables from a zip file using binary COPY protocol."""
conn = await asyncpg.connect(database_url)
try:
with zipfile.ZipFile(input_path, "r") as zf:
# Read and validate manifest
manifest: dict[str, Any] = json.loads(zf.read("manifest.json"))
if manifest.get("version") != MANIFEST_VERSION:
raise ValueError(f"Unsupported backup version: {manifest.get('version')}")
# Use a transaction for atomic restore - either all tables are
# restored or none are, preventing partial/inconsistent state.
async with conn.transaction():
typer.echo(" Clearing existing data...")
# Truncate tables in reverse order (respects FK constraints)
for table in reversed(BACKUP_TABLES):
qualified_table = _fq_table(table, schema)
await conn.execute(f"TRUNCATE TABLE {qualified_table} CASCADE")
# Restore tables in forward order
for i, table in enumerate(BACKUP_TABLES, 1):
filename = f"{table}.bin"
if filename not in zf.namelist():
typer.echo(f" [{i}/{len(BACKUP_TABLES)}] {table}: skipped (not in backup)")
continue
expected_rows = manifest["tables"].get(table, {}).get("rows", "?")
typer.echo(f" [{i}/{len(BACKUP_TABLES)}] Restoring {table}... {expected_rows} rows")
data = zf.read(filename)
buffer = io.BytesIO(data)
# asyncpg requires schema_name as separate parameter
await conn.copy_to_table(table, schema_name=schema, source=buffer, format="binary")
# Refresh materialized view
typer.echo(" Refreshing materialized views...")
await conn.execute(f"REFRESH MATERIALIZED VIEW {_fq_table('memory_units_bm25', schema)}")
return manifest
finally:
await conn.close()
async def _run_backup(db_url: str, output: Path, schema: str = "public") -> dict[str, Any]:
"""Resolve database URL and run backup."""
is_pg0, instance_name, _ = parse_pg0_url(db_url)
if is_pg0:
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
resolved_url = await resolve_database_url(db_url)
return await _backup(resolved_url, output, schema)
async def _run_restore(db_url: str, input_file: Path, schema: str = "public") -> dict[str, Any]:
"""Resolve database URL and run restore."""
is_pg0, instance_name, _ = parse_pg0_url(db_url)
if is_pg0:
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
resolved_url = await resolve_database_url(db_url)
return await _restore(resolved_url, input_file, schema)
@app.command()
def backup(
output: Path = typer.Argument(..., help="Output file path (.zip)"),
schema: str = typer.Option("public", "--schema", "-s", help="Database schema to backup"),
):
"""Backup the Hindsight database to a zip file."""
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 output.suffix != ".zip":
output = output.with_suffix(".zip")
typer.echo(f"Backing up database (schema: {schema}) to {output}...")
manifest = asyncio.run(_run_backup(config.database_url, output, schema))
total_rows = sum(t["rows"] for t in manifest["tables"].values())
typer.echo(f"Backed up {total_rows} rows across {len(BACKUP_TABLES)} tables")
typer.echo(f"Backup saved to {output}")
@app.command()
def restore(
input_file: Path = typer.Argument(..., help="Input backup file (.zip)"),
schema: str = typer.Option("public", "--schema", "-s", help="Database schema to restore to"),
yes: bool = typer.Option(False, "--yes", "-y", help="Skip confirmation prompt"),
):
"""Restore the database from a backup file. WARNING: This deletes all existing data."""
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 input_file.exists():
typer.echo(f"Error: File not found: {input_file}", err=True)
raise typer.Exit(1)
if not yes:
typer.confirm(
"This will DELETE all existing data and replace it with the backup. Continue?",
abort=True,
)
typer.echo(f"Restoring database (schema: {schema}) from {input_file}...")
manifest = asyncio.run(_run_restore(config.database_url, input_file, schema))
total_rows = sum(t["rows"] for t in manifest["tables"].values())
typer.echo(f"Restored {total_rows} rows across {len(BACKUP_TABLES)} tables")
typer.echo("Restore complete")
async def _run_migration(db_url: str, schema: str = "public") -> None:
"""Resolve database URL and run migrations."""
from ..migrations import run_migrations
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)
run_migrations(resolved_url, schema=schema)
@app.command(name="run-db-migration")
def run_db_migration(
schema: str = typer.Option("public", "--schema", "-s", help="Database schema to run migrations on"),
):
"""Run database migrations to the latest version."""
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)
typer.echo(f"Running database migrations (schema: {schema})...")
asyncio.run(_run_migration(config.database_url, schema))
typer.echo("Database migrations completed successfully")
def main():
app()
if __name__ == "__main__":
main()
@@ -2,20 +2,19 @@
Alembic environment configuration for SQLAlchemy with pgvector.
Uses synchronous psycopg2 driver for migrations to avoid pgbouncer issues.
"""
import logging
import os
import sys
from pathlib import Path
from sqlalchemy import pool, engine_from_config
from sqlalchemy.engine import Connection
from alembic import context
from dotenv import load_dotenv
from sqlalchemy import engine_from_config, pool
# Import your models here
from hindsight_api.models import Base
# Load environment variables based on HINDSIGHT_API_DATABASE_URL env var or default to local
def load_env():
"""Load environment variables from .env"""
@@ -30,6 +29,7 @@ def load_env():
if env_file.exists():
load_dotenv(env_file)
load_env()
# this is the Alembic Config object, which provides
@@ -109,6 +109,9 @@ def run_migrations_online() -> None:
get_database_url() # Process and set the database URL in config
# Check if we're targeting a specific schema (for multi-tenant isolation)
target_schema = config.get_main_option("target_schema")
connectable = engine_from_config(
config.get_section(config.config_ini_section, {}),
prefix="sqlalchemy.",
@@ -121,17 +124,34 @@ def run_migrations_online() -> None:
def set_read_write_mode(dbapi_connection, connection_record):
cursor = dbapi_connection.cursor()
cursor.execute("SET SESSION CHARACTERISTICS AS TRANSACTION READ WRITE")
# If targeting a specific schema, set search_path
# Include public in search_path for access to shared extensions (pgvector)
if target_schema:
cursor.execute(f'CREATE SCHEMA IF NOT EXISTS "{target_schema}"')
cursor.execute(f'SET search_path TO "{target_schema}", public')
cursor.close()
with connectable.connect() as connection:
# Also explicitly set read-write mode on this connection
connection.execute(text("SET SESSION CHARACTERISTICS AS TRANSACTION READ WRITE"))
# If targeting a specific schema, set search_path
# Include public in search_path for access to shared extensions (pgvector)
if target_schema:
connection.execute(text(f'CREATE SCHEMA IF NOT EXISTS "{target_schema}"'))
connection.execute(text(f'SET search_path TO "{target_schema}", public'))
connection.commit() # Commit the SET command
context.configure(
connection=connection,
target_metadata=target_metadata
)
# Configure context with version_table_schema if using a specific schema
context_opts = {
"connection": connection,
"target_metadata": target_metadata,
}
if target_schema:
context_opts["version_table_schema"] = target_schema
context.configure(**context_opts)
with context.begin_transaction():
context.run_migrations()
@@ -0,0 +1,360 @@
"""initial_schema
Revision ID: 5a366d414dce
Revises:
Create Date: 2025-11-27 11:54:19.228030
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
from pgvector.sqlalchemy import Vector
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = "5a366d414dce"
down_revision: str | Sequence[str] | None = None
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
"""Upgrade schema - create all tables from scratch."""
# Enable required extensions
op.execute("CREATE EXTENSION IF NOT EXISTS vector")
# Create banks table
op.create_table(
"banks",
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("name", sa.Text(), nullable=True),
sa.Column(
"personality",
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False,
),
sa.Column("background", sa.Text(), nullable=True),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.PrimaryKeyConstraint("bank_id", name=op.f("pk_banks")),
)
# Create documents table
op.create_table(
"documents",
sa.Column("id", sa.Text(), nullable=False),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("original_text", sa.Text(), nullable=True),
sa.Column("content_hash", sa.Text(), nullable=True),
sa.Column(
"metadata", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False
),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.PrimaryKeyConstraint("id", "bank_id", name=op.f("pk_documents")),
)
op.create_index("idx_documents_bank_id", "documents", ["bank_id"])
op.create_index("idx_documents_content_hash", "documents", ["content_hash"])
# Create async_operations table
op.create_table(
"async_operations",
sa.Column(
"operation_id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False
),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("operation_type", sa.Text(), nullable=False),
sa.Column("status", sa.Text(), server_default="pending", nullable=False),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("completed_at", postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column("error_message", sa.Text(), nullable=True),
sa.Column(
"result_metadata",
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False,
),
sa.PrimaryKeyConstraint("operation_id", name=op.f("pk_async_operations")),
sa.CheckConstraint(
"status IN ('pending', 'processing', 'completed', 'failed')", name="async_operations_status_check"
),
)
op.create_index("idx_async_operations_bank_id", "async_operations", ["bank_id"])
op.create_index("idx_async_operations_status", "async_operations", ["status"])
op.create_index("idx_async_operations_bank_status", "async_operations", ["bank_id", "status"])
# Create entities table
op.create_table(
"entities",
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
sa.Column("canonical_name", sa.Text(), nullable=False),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column(
"metadata", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False
),
sa.Column("first_seen", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("last_seen", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("mention_count", sa.Integer(), server_default="1", nullable=False),
sa.PrimaryKeyConstraint("id", name=op.f("pk_entities")),
)
op.create_index("idx_entities_bank_id", "entities", ["bank_id"])
op.create_index("idx_entities_canonical_name", "entities", ["canonical_name"])
op.create_index("idx_entities_bank_name", "entities", ["bank_id", "canonical_name"])
# Create unique index on (bank_id, LOWER(canonical_name)) for entity resolution
op.execute("CREATE UNIQUE INDEX idx_entities_bank_lower_name ON entities (bank_id, LOWER(canonical_name))")
# Create memory_units table
op.create_table(
"memory_units",
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("document_id", sa.Text(), nullable=True),
sa.Column("text", sa.Text(), nullable=False),
sa.Column("embedding", Vector(384), nullable=True),
sa.Column("context", sa.Text(), nullable=True),
sa.Column("event_date", postgresql.TIMESTAMP(timezone=True), nullable=False),
sa.Column("occurred_start", postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column("occurred_end", postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column("mentioned_at", postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column("fact_type", sa.Text(), server_default="world", nullable=False),
sa.Column("confidence_score", sa.Float(), nullable=True),
sa.Column("access_count", sa.Integer(), server_default="0", nullable=False),
sa.Column(
"metadata", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False
),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.ForeignKeyConstraint(
["document_id", "bank_id"],
["documents.id", "documents.bank_id"],
name="memory_units_document_fkey",
ondelete="CASCADE",
),
sa.PrimaryKeyConstraint("id", name=op.f("pk_memory_units")),
sa.CheckConstraint(
"fact_type IN ('world', 'bank', 'opinion', 'observation')", name="memory_units_fact_type_check"
),
sa.CheckConstraint(
"confidence_score IS NULL OR (confidence_score >= 0.0 AND confidence_score <= 1.0)",
name="memory_units_confidence_range_check",
),
sa.CheckConstraint(
"(fact_type = 'opinion' AND confidence_score IS NOT NULL) OR "
"(fact_type = 'observation') OR "
"(fact_type NOT IN ('opinion', 'observation') AND confidence_score IS NULL)",
name="confidence_score_fact_type_check",
),
)
# Add search_vector column for full-text search
op.execute("""
ALTER TABLE memory_units
ADD COLUMN search_vector tsvector
GENERATED ALWAYS AS (to_tsvector('english', COALESCE(text, '') || ' ' || COALESCE(context, ''))) STORED
""")
op.create_index("idx_memory_units_bank_id", "memory_units", ["bank_id"])
op.create_index("idx_memory_units_document_id", "memory_units", ["document_id"])
op.create_index("idx_memory_units_event_date", "memory_units", [sa.text("event_date DESC")])
op.create_index("idx_memory_units_bank_date", "memory_units", ["bank_id", sa.text("event_date DESC")])
op.create_index("idx_memory_units_access_count", "memory_units", [sa.text("access_count DESC")])
op.create_index("idx_memory_units_fact_type", "memory_units", ["fact_type"])
op.create_index("idx_memory_units_bank_fact_type", "memory_units", ["bank_id", "fact_type"])
op.create_index(
"idx_memory_units_bank_type_date", "memory_units", ["bank_id", "fact_type", sa.text("event_date DESC")]
)
op.create_index(
"idx_memory_units_opinion_confidence",
"memory_units",
["bank_id", sa.text("confidence_score DESC")],
postgresql_where=sa.text("fact_type = 'opinion'"),
)
op.create_index(
"idx_memory_units_opinion_date",
"memory_units",
["bank_id", sa.text("event_date DESC")],
postgresql_where=sa.text("fact_type = 'opinion'"),
)
op.create_index(
"idx_memory_units_observation_date",
"memory_units",
["bank_id", sa.text("event_date DESC")],
postgresql_where=sa.text("fact_type = 'observation'"),
)
op.create_index(
"idx_memory_units_embedding",
"memory_units",
["embedding"],
postgresql_using="hnsw",
postgresql_ops={"embedding": "vector_cosine_ops"},
)
# Create BM25 full-text search index on search_vector
op.execute("""
CREATE INDEX idx_memory_units_text_search ON memory_units
USING gin(search_vector)
""")
op.execute("""
CREATE MATERIALIZED VIEW memory_units_bm25 AS
SELECT
id,
bank_id,
text,
to_tsvector('english', text) AS text_vector,
log(1.0 + length(text)::float / (SELECT avg(length(text)) FROM memory_units)) AS doc_length_factor
FROM memory_units
""")
op.create_index("idx_memory_units_bm25_bank", "memory_units_bm25", ["bank_id"])
op.create_index("idx_memory_units_bm25_text_vector", "memory_units_bm25", ["text_vector"], postgresql_using="gin")
# Create entity_cooccurrences table
op.create_table(
"entity_cooccurrences",
sa.Column("entity_id_1", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("entity_id_2", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("cooccurrence_count", sa.Integer(), server_default="1", nullable=False),
sa.Column(
"last_cooccurred", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False
),
sa.ForeignKeyConstraint(
["entity_id_1"],
["entities.id"],
name=op.f("fk_entity_cooccurrences_entity_id_1_entities"),
ondelete="CASCADE",
),
sa.ForeignKeyConstraint(
["entity_id_2"],
["entities.id"],
name=op.f("fk_entity_cooccurrences_entity_id_2_entities"),
ondelete="CASCADE",
),
sa.PrimaryKeyConstraint("entity_id_1", "entity_id_2", name=op.f("pk_entity_cooccurrences")),
sa.CheckConstraint("entity_id_1 < entity_id_2", name="entity_cooccurrence_order_check"),
)
op.create_index("idx_entity_cooccurrences_entity1", "entity_cooccurrences", ["entity_id_1"])
op.create_index("idx_entity_cooccurrences_entity2", "entity_cooccurrences", ["entity_id_2"])
op.create_index("idx_entity_cooccurrences_count", "entity_cooccurrences", [sa.text("cooccurrence_count DESC")])
# Create memory_links table
op.create_table(
"memory_links",
sa.Column("from_unit_id", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("to_unit_id", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("link_type", sa.Text(), nullable=False),
sa.Column("entity_id", postgresql.UUID(as_uuid=True), nullable=True),
sa.Column("weight", sa.Float(), server_default="1.0", nullable=False),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.ForeignKeyConstraint(
["entity_id"], ["entities.id"], name=op.f("fk_memory_links_entity_id_entities"), ondelete="CASCADE"
),
sa.ForeignKeyConstraint(
["from_unit_id"],
["memory_units.id"],
name=op.f("fk_memory_links_from_unit_id_memory_units"),
ondelete="CASCADE",
),
sa.ForeignKeyConstraint(
["to_unit_id"],
["memory_units.id"],
name=op.f("fk_memory_links_to_unit_id_memory_units"),
ondelete="CASCADE",
),
sa.CheckConstraint(
"link_type IN ('temporal', 'semantic', 'entity', 'causes', 'caused_by', 'enables', 'prevents')",
name="memory_links_link_type_check",
),
sa.CheckConstraint("weight >= 0.0 AND weight <= 1.0", name="memory_links_weight_check"),
)
# Create unique constraint using COALESCE for nullable entity_id
op.execute(
"CREATE UNIQUE INDEX idx_memory_links_unique ON memory_links (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid))"
)
op.create_index("idx_memory_links_from_unit", "memory_links", ["from_unit_id"])
op.create_index("idx_memory_links_to_unit", "memory_links", ["to_unit_id"])
op.create_index("idx_memory_links_entity", "memory_links", ["entity_id"])
op.create_index("idx_memory_links_link_type", "memory_links", ["link_type"])
# Create unit_entities table
op.create_table(
"unit_entities",
sa.Column("unit_id", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("entity_id", postgresql.UUID(as_uuid=True), nullable=False),
sa.ForeignKeyConstraint(
["entity_id"], ["entities.id"], name=op.f("fk_unit_entities_entity_id_entities"), ondelete="CASCADE"
),
sa.ForeignKeyConstraint(
["unit_id"], ["memory_units.id"], name=op.f("fk_unit_entities_unit_id_memory_units"), ondelete="CASCADE"
),
sa.PrimaryKeyConstraint("unit_id", "entity_id", name=op.f("pk_unit_entities")),
)
op.create_index("idx_unit_entities_unit", "unit_entities", ["unit_id"])
op.create_index("idx_unit_entities_entity", "unit_entities", ["entity_id"])
def downgrade() -> None:
"""Downgrade schema - drop all tables."""
# Drop tables in reverse dependency order
op.drop_index("idx_unit_entities_entity", table_name="unit_entities")
op.drop_index("idx_unit_entities_unit", table_name="unit_entities")
op.drop_table("unit_entities")
op.drop_index("idx_memory_links_link_type", table_name="memory_links")
op.drop_index("idx_memory_links_entity", table_name="memory_links")
op.drop_index("idx_memory_links_to_unit", table_name="memory_links")
op.drop_index("idx_memory_links_from_unit", table_name="memory_links")
op.execute("DROP INDEX IF EXISTS idx_memory_links_unique")
op.drop_table("memory_links")
op.drop_index("idx_entity_cooccurrences_count", table_name="entity_cooccurrences")
op.drop_index("idx_entity_cooccurrences_entity2", table_name="entity_cooccurrences")
op.drop_index("idx_entity_cooccurrences_entity1", table_name="entity_cooccurrences")
op.drop_table("entity_cooccurrences")
# Drop BM25 materialized view and index
op.drop_index("idx_memory_units_bm25_text_vector", table_name="memory_units_bm25")
op.drop_index("idx_memory_units_bm25_bank", table_name="memory_units_bm25")
op.execute("DROP MATERIALIZED VIEW IF EXISTS memory_units_bm25")
op.drop_index("idx_memory_units_embedding", table_name="memory_units")
op.drop_index("idx_memory_units_observation_date", table_name="memory_units")
op.drop_index("idx_memory_units_opinion_date", table_name="memory_units")
op.drop_index("idx_memory_units_opinion_confidence", table_name="memory_units")
op.drop_index("idx_memory_units_bank_type_date", table_name="memory_units")
op.drop_index("idx_memory_units_bank_fact_type", table_name="memory_units")
op.drop_index("idx_memory_units_fact_type", table_name="memory_units")
op.drop_index("idx_memory_units_access_count", table_name="memory_units")
op.drop_index("idx_memory_units_bank_date", table_name="memory_units")
op.drop_index("idx_memory_units_event_date", table_name="memory_units")
op.drop_index("idx_memory_units_document_id", table_name="memory_units")
op.drop_index("idx_memory_units_bank_id", table_name="memory_units")
op.execute("DROP INDEX IF EXISTS idx_memory_units_text_search")
op.drop_table("memory_units")
op.execute("DROP INDEX IF EXISTS idx_entities_bank_lower_name")
op.drop_index("idx_entities_bank_name", table_name="entities")
op.drop_index("idx_entities_canonical_name", table_name="entities")
op.drop_index("idx_entities_bank_id", table_name="entities")
op.drop_table("entities")
op.drop_index("idx_async_operations_bank_status", table_name="async_operations")
op.drop_index("idx_async_operations_status", table_name="async_operations")
op.drop_index("idx_async_operations_bank_id", table_name="async_operations")
op.drop_table("async_operations")
op.drop_index("idx_documents_content_hash", table_name="documents")
op.drop_index("idx_documents_bank_id", table_name="documents")
op.drop_table("documents")
op.drop_table("banks")
# Drop extensions (optional - comment out if you want to keep them)
# op.execute('DROP EXTENSION IF EXISTS vector')
# op.execute('DROP EXTENSION IF EXISTS "uuid-ossp"')
@@ -0,0 +1,70 @@
"""add_chunks_table
Revision ID: b7c4d8e9f1a2
Revises: 5a366d414dce
Create Date: 2025-11-28 00:00:00.000000
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = "b7c4d8e9f1a2"
down_revision: str | Sequence[str] | None = "5a366d414dce"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
"""Add chunks table and link memory_units to chunks."""
# Create chunks table with single text PK (bank_id_document_id_chunk_index)
op.create_table(
"chunks",
sa.Column("chunk_id", sa.Text(), nullable=False),
sa.Column("document_id", sa.Text(), nullable=False),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("chunk_index", sa.Integer(), nullable=False),
sa.Column("chunk_text", sa.Text(), nullable=False),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.ForeignKeyConstraint(
["document_id", "bank_id"],
["documents.id", "documents.bank_id"],
name="chunks_document_fkey",
ondelete="CASCADE",
),
sa.PrimaryKeyConstraint("chunk_id", name=op.f("pk_chunks")),
)
# Add indexes for efficient queries
op.create_index("idx_chunks_document_id", "chunks", ["document_id"])
op.create_index("idx_chunks_bank_id", "chunks", ["bank_id"])
# Add chunk_id column to memory_units (nullable, as existing records won't have chunks)
op.add_column("memory_units", sa.Column("chunk_id", sa.Text(), nullable=True))
# Add foreign key constraint to chunks table
op.create_foreign_key(
"memory_units_chunk_fkey", "memory_units", "chunks", ["chunk_id"], ["chunk_id"], ondelete="SET NULL"
)
# Add index on chunk_id for efficient lookups
op.create_index("idx_memory_units_chunk_id", "memory_units", ["chunk_id"])
def downgrade() -> None:
"""Remove chunks table and chunk_id from memory_units."""
# Drop index and foreign key from memory_units
op.drop_index("idx_memory_units_chunk_id", table_name="memory_units")
op.drop_constraint("memory_units_chunk_fkey", "memory_units", type_="foreignkey")
op.drop_column("memory_units", "chunk_id")
# Drop chunks table indexes and table
op.drop_index("idx_chunks_bank_id", table_name="chunks")
op.drop_index("idx_chunks_document_id", table_name="chunks")
op.drop_table("chunks")
@@ -5,35 +5,35 @@ Revises: b7c4d8e9f1a2
Create Date: 2025-12-02 00:00:00.000000
"""
from typing import Sequence, Union
from alembic import op
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = 'c8e5f2a3b4d1'
down_revision: Union[str, Sequence[str], None] = 'b7c4d8e9f1a2'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
revision: str = "c8e5f2a3b4d1"
down_revision: str | Sequence[str] | None = "b7c4d8e9f1a2"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
"""Add retain_params JSONB column to documents table."""
# Add retain_params column to store parameters passed during retain
op.add_column('documents', sa.Column('retain_params', postgresql.JSONB(), nullable=True))
op.add_column("documents", sa.Column("retain_params", postgresql.JSONB(), nullable=True))
# Add index for efficient queries on retain_params
op.create_index('idx_documents_retain_params', 'documents', ['retain_params'], postgresql_using='gin')
op.create_index("idx_documents_retain_params", "documents", ["retain_params"], postgresql_using="gin")
def downgrade() -> None:
"""Remove retain_params column from documents table."""
# Drop index
op.drop_index('idx_documents_retain_params', table_name='documents')
op.drop_index("idx_documents_retain_params", table_name="documents")
# Drop column
op.drop_column('documents', 'retain_params')
op.drop_column("documents", "retain_params")
@@ -0,0 +1,53 @@
"""Rename fact_type 'bank' to 'experience'
Revision ID: d9f6a3b4c5e2
Revises: c8e5f2a3b4d1
Create Date: 2024-12-04 15:00:00.000000
"""
from alembic import context, op
# revision identifiers, used by Alembic.
revision = "d9f6a3b4c5e2"
down_revision = "c8e5f2a3b4d1"
branch_labels = None
depends_on = None
def _get_schema_prefix() -> str:
"""Get schema prefix for table names (e.g., 'tenant_x.' or '' for public)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade():
schema = _get_schema_prefix()
# Drop old check constraint FIRST (before updating data)
op.drop_constraint("memory_units_fact_type_check", "memory_units", type_="check")
# Update existing 'bank' values to 'experience'
op.execute(f"UPDATE {schema}memory_units SET fact_type = 'experience' WHERE fact_type = 'bank'")
# Also update any 'interactions' values (in case of partial migration)
op.execute(f"UPDATE {schema}memory_units SET fact_type = 'experience' WHERE fact_type = 'interactions'")
# Create new check constraint with 'experience' instead of 'bank'
op.create_check_constraint(
"memory_units_fact_type_check", "memory_units", "fact_type IN ('world', 'experience', 'opinion', 'observation')"
)
def downgrade():
schema = _get_schema_prefix()
# Drop new check constraint FIRST
op.drop_constraint("memory_units_fact_type_check", "memory_units", type_="check")
# Update 'experience' back to 'bank'
op.execute(f"UPDATE {schema}memory_units SET fact_type = 'bank' WHERE fact_type = 'experience'")
# Recreate old check constraint
op.create_check_constraint(
"memory_units_fact_type_check", "memory_units", "fact_type IN ('world', 'bank', 'opinion', 'observation')"
)
@@ -0,0 +1,111 @@
"""disposition_to_3_traits
Revision ID: e0a1b2c3d4e5
Revises: rename_personality
Create Date: 2024-12-08
Migrate disposition traits from Big Five (openness, conscientiousness, extraversion,
agreeableness, neuroticism, bias_strength with 0-1 float values) to the new 3-trait
system (skepticism, literalism, empathy with 1-5 integer values).
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "e0a1b2c3d4e5"
down_revision: str | Sequence[str] | None = "rename_personality"
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 (e.g., 'tenant_x.' or '' for public)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def _get_target_schema() -> str:
"""Get the target schema name (tenant schema or 'public')."""
schema = context.config.get_main_option("target_schema")
return schema if schema else "public"
def upgrade() -> None:
"""Convert Big Five disposition to 3-trait disposition."""
conn = op.get_bind()
schema = _get_schema_prefix()
target_schema = _get_target_schema()
# Check if disposition column exists (should have been created by previous migration)
result = conn.execute(
sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'disposition'
"""),
{"schema": target_schema},
)
if not result.fetchone():
# Column doesn't exist yet (shouldn't happen but be safe)
return
# Update all existing banks to use the new disposition format
# Convert from old format to new format with reasonable mappings:
# - skepticism: derived from inverse of agreeableness (skeptical people are less agreeable)
# - literalism: derived from conscientiousness (detail-oriented people are more literal)
# - empathy: derived from agreeableness + inverse of neuroticism
# Default all to 3 (neutral) for simplicity
conn.execute(
sa.text(f"""
UPDATE {schema}banks
SET disposition = '{{"skepticism": 3, "literalism": 3, "empathy": 3}}'::jsonb
WHERE disposition IS NOT NULL
""")
)
# Update the default for new banks
conn.execute(
sa.text(f"""
ALTER TABLE {schema}banks
ALTER COLUMN disposition SET DEFAULT '{{"skepticism": 3, "literalism": 3, "empathy": 3}}'::jsonb
""")
)
def downgrade() -> None:
"""Convert back to Big Five disposition."""
conn = op.get_bind()
schema = _get_schema_prefix()
target_schema = _get_target_schema()
# Check if disposition column exists
result = conn.execute(
sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'disposition'
"""),
{"schema": target_schema},
)
if not result.fetchone():
return
# Revert to Big Five format with default values
conn.execute(
sa.text(f"""
UPDATE {schema}banks
SET disposition = '{{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}}'::jsonb
WHERE disposition IS NOT NULL
""")
)
# Update the default for new banks
conn.execute(
sa.text(f"""
ALTER TABLE {schema}banks
ALTER COLUMN disposition SET DEFAULT '{{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}}'::jsonb
""")
)
@@ -0,0 +1,44 @@
"""add_memory_links_from_type_weight_index
Revision ID: f1a2b3c4d5e6
Revises: e0a1b2c3d4e5
Create Date: 2025-01-12
Add composite index on memory_links (from_unit_id, link_type, weight DESC)
to optimize MPFP graph traversal queries that need top-k edges per type.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "f1a2b3c4d5e6"
down_revision: str | Sequence[str] | None = "e0a1b2c3d4e5"
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 (e.g., 'tenant_x.' or '' for public)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
"""Add composite index for efficient MPFP edge loading."""
schema = _get_schema_prefix()
# Create composite index for efficient top-k per (from_node, link_type) queries
# This enables LATERAL joins to use index-only scans with early termination
# Note: Not using CONCURRENTLY here as it requires running outside a transaction
# For production with large tables, consider running this manually with CONCURRENTLY
op.execute(
f"CREATE INDEX IF NOT EXISTS idx_memory_links_from_type_weight "
f"ON {schema}memory_links(from_unit_id, link_type, weight DESC)"
)
def downgrade() -> None:
"""Remove the composite index."""
schema = _get_schema_prefix()
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_links_from_type_weight")
@@ -0,0 +1,48 @@
"""add_tags_column
Revision ID: g2a3b4c5d6e7
Revises: f1a2b3c4d5e6
Create Date: 2025-01-13
Add tags column to memory_units and documents tables for visibility scoping.
Tags enable filtering memories by scope (e.g., user IDs, session IDs) during recall/reflect.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "g2a3b4c5d6e7"
down_revision: str | Sequence[str] | None = "f1a2b3c4d5e6"
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 (e.g., 'tenant_x.' or '' for public)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
"""Add tags column to memory_units and documents tables."""
schema = _get_schema_prefix()
# Add tags column to memory_units table
op.execute(f"ALTER TABLE {schema}memory_units ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
# Create GIN index for efficient array containment queries (tags && ARRAY['x'])
op.execute(f"CREATE INDEX IF NOT EXISTS idx_memory_units_tags ON {schema}memory_units USING GIN (tags)")
# Add tags column to documents table for document-level tags
op.execute(f"ALTER TABLE {schema}documents ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
def downgrade() -> None:
"""Remove tags columns and index."""
schema = _get_schema_prefix()
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_tags")
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS tags")
op.execute(f"ALTER TABLE {schema}documents DROP COLUMN IF EXISTS tags")
@@ -0,0 +1,112 @@
"""mental_models_v4
Revision ID: h3c4d5e6f7g8
Revises: g2a3b4c5d6e7
Create Date: 2026-01-08 00:00:00.000000
This migration implements the v4 mental models system:
1. Deletes existing observation memory_units (observations now in mental models)
2. Adds mission column to banks (replacing background)
3. Creates mental_models table with final schema
Mental models can reference entities when an entity is "promoted" to a mental model.
Summary content is stored as JSONB observations with per-observation fact attribution.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "h3c4d5e6f7g8"
down_revision: str | Sequence[str] | None = "g2a3b4c5d6e7"
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:
"""Apply mental models v4 changes."""
schema = _get_schema_prefix()
# Step 1: Delete observation memory_units (cascades to unit_entities links)
# Observations are now handled through mental models, not memory_units
op.execute(f"DELETE FROM {schema}memory_units WHERE fact_type = 'observation'")
# Step 2: Drop observation-specific index (if it exists)
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_observation_date")
# Step 3: Add mission column to banks (replacing background)
op.execute(f"ALTER TABLE {schema}banks ADD COLUMN IF NOT EXISTS mission TEXT")
# Migrate: copy background to mission if background column exists
# Use DO block to check column existence first (idempotent for re-runs)
schema_name = context.config.get_main_option("target_schema") or "public"
op.execute(f"""
DO $$
BEGIN
IF EXISTS (
SELECT 1 FROM information_schema.columns
WHERE table_schema = '{schema_name}' AND table_name = 'banks' AND column_name = 'background'
) THEN
UPDATE {schema}banks
SET mission = background
WHERE mission IS NULL;
END IF;
END $$;
""")
# Remove background column (replaced by mission)
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS background")
# Step 4: Create mental_models table with final v4 schema (if not exists)
op.execute(f"""
CREATE TABLE IF NOT EXISTS {schema}mental_models (
id VARCHAR(64) NOT NULL,
bank_id VARCHAR(64) NOT NULL,
subtype VARCHAR(32) NOT NULL,
name VARCHAR(256) NOT NULL,
description TEXT NOT NULL,
entity_id UUID,
observations JSONB DEFAULT '{{"observations": []}}'::jsonb,
links VARCHAR[],
tags VARCHAR[] DEFAULT '{{}}',
last_updated TIMESTAMP WITH TIME ZONE,
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
PRIMARY KEY (id, bank_id),
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE,
FOREIGN KEY (entity_id) REFERENCES {schema}entities(id) ON DELETE SET NULL,
CONSTRAINT ck_mental_models_subtype CHECK (subtype IN ('structural', 'emergent', 'pinned', 'learned'))
)
""")
# Step 5: Create indexes for efficient queries (if not exist)
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_bank_id ON {schema}mental_models(bank_id)")
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_subtype ON {schema}mental_models(bank_id, subtype)")
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_entity_id ON {schema}mental_models(entity_id)")
# GIN index for efficient tags array filtering
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_tags ON {schema}mental_models USING GIN(tags)")
def downgrade() -> None:
"""Revert mental models v4 changes."""
schema = _get_schema_prefix()
# Drop mental_models table (cascades to indexes)
op.execute(f"DROP TABLE IF EXISTS {schema}mental_models CASCADE")
# Add back background column to banks
op.execute(f"ALTER TABLE {schema}banks ADD COLUMN IF NOT EXISTS background TEXT")
# Migrate mission back to background
op.execute(f"UPDATE {schema}banks SET background = mission WHERE background IS NULL")
# Remove mission column
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS mission")
# Note: Cannot restore deleted observations - they are lost on downgrade
@@ -0,0 +1,41 @@
"""delete_opinions
Revision ID: i4d5e6f7g8h9
Revises: h3c4d5e6f7g8
Create Date: 2026-01-15 00:00:00.000000
This migration removes opinion facts from memory_units.
Opinions are no longer a separate fact type - they are now represented
through mental model observations with confidence scores.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "i4d5e6f7g8h9"
down_revision: str | Sequence[str] | None = "h3c4d5e6f7g8"
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:
"""Delete opinion memory_units."""
schema = _get_schema_prefix()
# Delete opinion memory_units (cascades to unit_entities links)
# Opinions are now handled through mental model observations
op.execute(f"DELETE FROM {schema}memory_units WHERE fact_type = 'opinion'")
def downgrade() -> None:
"""Cannot restore deleted opinions."""
# Note: Cannot restore deleted opinions - they are lost on downgrade
pass
@@ -0,0 +1,85 @@
"""rename_personality_to_disposition
Revision ID: rename_personality
Revises: d9f6a3b4c5e2
Create Date: 2024-12-04
"""
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 = "rename_personality"
down_revision: str | Sequence[str] | None = "d9f6a3b4c5e2"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_target_schema() -> str:
"""Get the target schema name (tenant schema or 'public')."""
schema = context.config.get_main_option("target_schema")
return schema if schema else "public"
def upgrade() -> None:
"""Rename personality column to disposition in banks table (if it exists)."""
conn = op.get_bind()
target_schema = _get_target_schema()
# Check if 'personality' column exists (old database)
result = conn.execute(
sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'personality'
"""),
{"schema": target_schema},
)
has_personality = result.fetchone() is not None
# Check if 'disposition' column exists (new database)
result = conn.execute(
sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'disposition'
"""),
{"schema": target_schema},
)
has_disposition = result.fetchone() is not None
if has_personality and not has_disposition:
# Old database: rename personality -> disposition
op.alter_column("banks", "personality", new_column_name="disposition")
elif not has_personality and not has_disposition:
# Neither exists (shouldn't happen, but be safe): add disposition column
op.add_column(
"banks",
sa.Column(
"disposition",
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False,
),
)
# else: disposition already exists, nothing to do
def downgrade() -> None:
"""Revert disposition column back to personality."""
conn = op.get_bind()
target_schema = _get_target_schema()
result = conn.execute(
sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'disposition'
"""),
{"schema": target_schema},
)
if result.fetchone():
op.alter_column("banks", "disposition", new_column_name="personality")
+49 -25
View File
@@ -3,8 +3,11 @@ Unified API module for Hindsight.
Provides both HTTP REST API and MCP (Model Context Protocol) server.
"""
import logging
from contextlib import asynccontextmanager
from typing import Optional
from fastapi import FastAPI
from hindsight_api import MemoryEngine
@@ -17,7 +20,7 @@ def create_app(
http_api_enabled: bool = True,
mcp_api_enabled: bool = False,
mcp_mount_path: str = "/mcp",
initialize_memory: bool = True
initialize_memory: bool = True,
) -> FastAPI:
"""
Create and configure the unified Hindsight API application.
@@ -43,49 +46,70 @@ def create_app(
# Both HTTP and MCP
app = create_app(memory, mcp_api_enabled=True)
"""
mcp_app = None
# Create MCP app first if enabled (we need its lifespan for chaining)
if mcp_api_enabled:
try:
from .mcp import create_mcp_app
mcp_app = create_mcp_app(memory=memory)
except ImportError as e:
logger.error(f"MCP server requested but dependencies not available: {e}")
logger.error("Install with: pip install hindsight-api[mcp]")
raise
# Import and create HTTP API if enabled
if http_api_enabled:
from .http import create_app as create_http_app
app = create_http_app(
memory=memory,
initialize_memory=initialize_memory
)
app = create_http_app(memory=memory, initialize_memory=initialize_memory)
logger.info("HTTP REST API enabled")
else:
# Create minimal FastAPI app
app = FastAPI(title="Hindsight API", version="0.0.7")
logger.info("HTTP REST API disabled")
# Mount MCP server if enabled
if mcp_api_enabled:
try:
from .mcp import create_mcp_app
# Mount MCP server and chain its lifespan if enabled
if mcp_app is not None:
# Get the MCP app's underlying Starlette app for lifespan access
mcp_starlette_app = mcp_app.mcp_app
# Create MCP app with dynamic bank_id support
# Supports: /mcp/{bank_id}/sse (bank-specific SSE endpoint)
mcp_app = create_mcp_app(memory=memory)
app.mount(mcp_mount_path, mcp_app)
logger.info(f"MCP server enabled at {mcp_mount_path}/{{bank_id}}/sse")
except ImportError as e:
logger.error(f"MCP server requested but dependencies not available: {e}")
logger.error("Install with: pip install hindsight-api[mcp]")
raise
# Store the original lifespan
original_lifespan = app.router.lifespan_context
@asynccontextmanager
async def chained_lifespan(app_instance: FastAPI):
"""Chain the MCP lifespan with the main app lifespan."""
# Start MCP lifespan first
async with mcp_starlette_app.router.lifespan_context(mcp_starlette_app):
logger.info("MCP lifespan started")
# Then start the original app lifespan
async with original_lifespan(app_instance):
yield
logger.info("MCP lifespan stopped")
# Replace the app's lifespan with the chained version
app.router.lifespan_context = chained_lifespan
# Mount the MCP middleware
app.mount(mcp_mount_path, mcp_app)
logger.info(f"MCP server enabled at {mcp_mount_path}/")
return app
# Re-export commonly used items for backwards compatibility
from .http import (
RecallRequest,
RecallResult,
RecallResponse,
MemoryItem,
RetainRequest,
ReflectRequest,
ReflectResponse,
CreateBankRequest,
DispositionTraits,
MemoryItem,
RecallRequest,
RecallResponse,
RecallResult,
ReflectRequest,
ReflectResponse,
RetainRequest,
)
__all__ = [
File diff suppressed because it is too large Load Diff
+230 -75
View File
@@ -4,28 +4,38 @@ import json
import logging
import os
from contextvars import ContextVar
from typing import Optional
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
# Configure logging from HINDSIGHT_API_LOG_LEVEL environment variable
_log_level_str = os.environ.get("HINDSIGHT_API_LOG_LEVEL", "info").lower()
_log_level_map = {"critical": logging.CRITICAL, "error": logging.ERROR, "warning": logging.WARNING,
"info": logging.INFO, "debug": logging.DEBUG, "trace": logging.DEBUG}
_log_level_map = {
"critical": logging.CRITICAL,
"error": logging.ERROR,
"warning": logging.WARNING,
"info": logging.INFO,
"debug": logging.DEBUG,
"trace": logging.DEBUG,
}
logging.basicConfig(
level=_log_level_map.get(_log_level_str, logging.INFO),
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s"
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
)
logger = logging.getLogger(__name__)
# Context variable to hold the current bank_id from the URL path
_current_bank_id: ContextVar[Optional[str]] = ContextVar("current_bank_id", default=None)
# Default bank_id from environment variable
DEFAULT_BANK_ID = os.environ.get("HINDSIGHT_MCP_BANK_ID", "default")
# Context variable to hold the current bank_id
_current_bank_id: ContextVar[str | None] = ContextVar("current_bank_id", default=None)
def get_current_bank_id() -> Optional[str]:
"""Get the current bank_id from context (set from URL path)."""
def get_current_bank_id() -> str | None:
"""Get the current bank_id from context."""
return _current_bank_id.get()
@@ -37,12 +47,18 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
memory: MemoryEngine instance (required)
Returns:
Configured FastMCP server instance
Configured FastMCP server instance with stateless_http enabled
"""
mcp = FastMCP("hindsight-mcp-server")
# 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") -> str:
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.
@@ -58,20 +74,34 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
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:
bank_id = get_current_bank_id()
await memory.put_batch_async(
bank_id=bank_id,
contents=[{"content": content, "context": context}]
)
return "Memory stored successfully"
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_results: int = 10) -> str:
async def recall(query: str, max_tokens: int = 4096, bank_id: str | None = None) -> str:
"""
Search memories to provide personalized, context-aware responses.
@@ -83,49 +113,165 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
Args:
query: Natural language search query (e.g., "user's food preferences", "what projects is user working on")
max_results: Maximum number of results to return (default: 10)
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:
bank_id = get_current_bank_id()
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
search_result = await memory.recall_async(
bank_id=bank_id,
recall_result = await memory.recall_async(
bank_id=target_bank,
query=query,
fact_type=list(VALID_RECALL_FACT_TYPES),
budget=Budget.LOW
budget=Budget.HIGH,
max_tokens=max_tokens,
request_context=RequestContext(),
)
results = [
{
"id": fact.id,
"text": fact.text,
"type": fact.fact_type,
"context": fact.context,
"event_date": fact.event_date,
}
for fact in search_result.results[:max_results]
]
return json.dumps({"results": results}, indent=2)
# 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 json.dumps({"error": str(e), "results": []})
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}"}}'
return mcp
class MCPMiddleware:
"""ASGI middleware that extracts bank_id from path and sets context."""
"""ASGI middleware that extracts bank_id from header or path and sets context.
Bank ID can be provided via:
1. X-Bank-Id header (recommended for Claude Code)
2. URL path: /mcp/{bank_id}/
3. Environment variable HINDSIGHT_MCP_BANK_ID (fallback default)
For Claude Code, configure with:
claude mcp add --transport http hindsight http://localhost:8888/mcp \\
--header "X-Bank-Id: my-bank"
"""
def __init__(self, app, memory: MemoryEngine):
self.app = app
self.memory = memory
self.mcp_server = create_mcp_server(memory)
# Use sse_app - http_app requires lifespan management that's complex with middleware
import warnings
with warnings.catch_warnings():
warnings.simplefilter("ignore", DeprecationWarning)
self.mcp_app = self.mcp_server.sse_app()
self.mcp_app = self.mcp_server.http_app(path="/")
# Expose the lifespan for the parent app to chain
self.lifespan = self.mcp_app.lifespan_handler if hasattr(self.mcp_app, "lifespan_handler") else None
def _get_header(self, scope: dict, name: str) -> str | None:
"""Extract a header value from ASGI scope."""
name_lower = name.lower().encode()
for header_name, header_value in scope.get("headers", []):
if header_name.lower() == name_lower:
return header_value.decode()
return None
async def __call__(self, scope, receive, send):
if scope["type"] != "http":
@@ -137,46 +283,50 @@ class MCPMiddleware:
# Strip any mount prefix (e.g., /mcp) that FastAPI might not have stripped
root_path = scope.get("root_path", "")
if root_path and path.startswith(root_path):
path = path[len(root_path):] or "/"
path = path[len(root_path) :] or "/"
# Also handle case where mount path wasn't stripped (e.g., /mcp/...)
if path.startswith("/mcp/"):
path = path[4:] # Remove /mcp prefix
elif path == "/mcp":
path = "/"
# Extract bank_id from path: /{bank_id}/ or /{bank_id}
# http_app expects requests at /
if not path.startswith("/") or len(path) <= 1:
# No bank_id in path - return error
await self._send_error(send, 400, "bank_id required in path: /mcp/{bank_id}/")
return
# Try to get bank_id from header first (for Claude Code compatibility)
bank_id = self._get_header(scope, "X-Bank-Id")
# Extract bank_id from first path segment
parts = path[1:].split("/", 1)
if not parts[0]:
await self._send_error(send, 400, "bank_id required in path: /mcp/{bank_id}/")
return
# MCP endpoint paths that should not be treated as bank_ids
MCP_ENDPOINTS = {"sse", "messages"}
bank_id = parts[0]
new_path = "/" + parts[1] if len(parts) > 1 else "/"
# If no header, try to extract from path: /{bank_id}/...
new_path = path
if not bank_id and path.startswith("/") and len(path) > 1:
parts = path[1:].split("/", 1)
# Don't treat MCP endpoints as bank_ids
if parts[0] and parts[0] not in MCP_ENDPOINTS:
# First segment looks like a bank_id
bank_id = parts[0]
new_path = "/" + parts[1] if len(parts) > 1 else "/"
# Fall back to default bank_id
if not bank_id:
bank_id = DEFAULT_BANK_ID
logger.debug(f"Using default bank_id: {bank_id}")
# Set bank_id context
token = _current_bank_id.set(bank_id)
try:
new_scope = scope.copy()
new_scope["path"] = new_path
# Clear root_path since we're passing directly to the app
new_scope["root_path"] = ""
# Wrap send to rewrite the SSE endpoint URL to include bank_id
# The SSE app sends "event: endpoint\ndata: /messages\n" but we need
# the client to POST to /{bank_id}/messages instead
# Wrap send to rewrite the SSE endpoint URL to include bank_id if using path-based routing
async def send_wrapper(message):
if message["type"] == "http.response.body":
body = message.get("body", b"")
if body and b"/messages" in body:
# Rewrite /messages to /{bank_id}/messages in SSE endpoint event
body = body.replace(
b"data: /messages",
f"data: /{bank_id}/messages".encode()
)
body = body.replace(b"data: /messages", f"data: /{bank_id}/messages".encode())
message = {**message, "body": body}
await send(message)
@@ -187,24 +337,29 @@ class MCPMiddleware:
async def _send_error(self, send, status: int, message: str):
"""Send an error response."""
body = json.dumps({"error": message}).encode()
await send({
"type": "http.response.start",
"status": status,
"headers": [(b"content-type", b"application/json")],
})
await send({
"type": "http.response.body",
"body": body,
})
await send(
{
"type": "http.response.start",
"status": status,
"headers": [(b"content-type", b"application/json")],
}
)
await send(
{
"type": "http.response.body",
"body": body,
}
)
def create_mcp_app(memory: MemoryEngine):
"""
Create an ASGI app that handles MCP requests.
URL pattern: /mcp/{bank_id}/
The bank_id is extracted from the URL path and made available to tools.
Bank ID can be provided via:
1. X-Bank-Id header: claude mcp add --transport http hindsight http://localhost:8888/mcp --header "X-Bank-Id: my-bank"
2. URL path: /mcp/{bank_id}/
3. Environment variable HINDSIGHT_MCP_BANK_ID (fallback, default: "default")
Args:
memory: MemoryEngine instance
+96
View File
@@ -0,0 +1,96 @@
"""
Banner display for Hindsight API startup.
Shows the logo and tagline with gradient colors.
"""
# Gradient colors: #0074d9 -> #009296
GRADIENT_START = (0, 116, 217) # #0074d9
GRADIENT_END = (0, 146, 150) # #009296
# Pre-generated logo (generated by test-logo.py)
LOGO = """\
\033[38;2;9;127;184m\u2584\033[0m\033[48;2;8;130;178m\033[38;2;5;133;186m\u2584\033[0m \033[48;2;10;143;160m\033[38;2;10;143;165m\u2584\033[0m\033[38;2;7;140;156m\u2584\033[0m
\033[38;2;8;125;192m\u2584\033[0m \033[38;2;3;132;191m\u2580\033[0m\033[38;2;2;133;192m\u2584\033[0m \033[38;2;3;132;180m\u2584\033[0m\033[38;2;1;137;184m\u2584\033[0m\033[38;2;3;133;174m\u2584\033[0m \033[38;2;3;142;176m\u2584\033[0m\033[38;2;4;142;169m\u2580\033[0m \033[38;2;10;144;164m\u2584\033[0m
\033[38;2;6;121;195m\u2580\033[0m\033[38;2;5;128;203m\u2580\033[0m\033[48;2;5;124;195m\033[38;2;3;125;200m\u2584\033[0m\033[38;2;2;126;196m\u2584\033[0m\033[48;2;3;128;188m\033[38;2;1;131;196m\u2584\033[0m\033[48;2;0;152;219m\033[38;2;2;131;191m\u2584\033[0m\033[38;2;1;141;196m\u2580\033[0m\033[38;2;1;135;183m\u2580\033[0m\033[38;2;1;148;198m\u2580\033[0m\033[48;2;1;156;202m\033[38;2;2;135;180m\u2584\033[0m\033[48;2;4;134;169m\033[38;2;1;137;177m\u2584\033[0m\033[38;2;3;138;173m\u2584\033[0m\033[48;2;6;137;165m\033[38;2;2;140;170m\u2584\033[0m\033[38;2;7;144;169m\u2580\033[0m\033[38;2;7;139;158m\u2580\033[0m
\033[48;2;2;128;202m\033[38;2;2;124;201m\u2584\033[0m\033[48;2;1;130;201m\033[38;2;0;135;212m\u2584\033[0m\033[38;2;2;128;196m\u2584\033[0m \033[48;2;2;142;204m\033[38;2;7;138;199m\u2584\033[0m \033[38;2;1;135;186m\u2584\033[0m\033[48;2;1;142;186m\033[38;2;2;144;194m\u2584\033[0m\033[48;2;3;138;176m\033[38;2;2;134;176m\u2584\033[0m
\033[48;2;8;118;200m\033[38;2;8;121;209m\u2584\033[0m\033[38;2;3;121;203m\u2580\033[0m \033[38;2;3;122;192m\u2580\033[0m\033[38;2;1;138;216m\u2580\033[0m\033[48;2;0;138;210m\033[38;2;3;128;198m\u2584\033[0m\033[48;2;0;126;188m\033[38;2;2;131;198m\u2584\033[0m\033[48;2;0;142;205m\033[38;2;3;132;193m\u2584\033[0m\033[38;2;1;140;196m\u2580\033[0m \033[38;2;4;134;175m\u2580\033[0m\033[48;2;13;135;167m\033[38;2;8;136;174m\u2584\033[0m """
def _interpolate_color(start: tuple, end: tuple, t: float) -> tuple:
"""Interpolate between two RGB colors."""
return (
int(start[0] + (end[0] - start[0]) * t),
int(start[1] + (end[1] - start[1]) * t),
int(start[2] + (end[2] - start[2]) * t),
)
def gradient_text(text: str, start: tuple = GRADIENT_START, end: tuple = GRADIENT_END) -> str:
"""Render text with a gradient color effect."""
result = []
length = len(text)
for i, char in enumerate(text):
if char == " ":
result.append(" ")
else:
t = i / max(length - 1, 1)
r, g, b = _interpolate_color(start, end, t)
result.append(f"\033[38;2;{r};{g};{b}m{char}")
result.append("\033[0m")
return "".join(result)
def print_banner():
"""Print the Hindsight startup banner."""
print(LOGO)
tagline = gradient_text("Hindsight: Agent Memory That Works Like Human Memory")
print(f"\n {tagline}\n")
def color(text: str, t: float = 0.0) -> str:
"""Color text using gradient position (0.0 = start, 1.0 = end)."""
r, g, b = _interpolate_color(GRADIENT_START, GRADIENT_END, t)
return f"\033[38;2;{r};{g};{b}m{text}\033[0m"
def color_start(text: str) -> str:
"""Color text with gradient start color (#0074d9)."""
return color(text, 0.0)
def color_end(text: str) -> str:
"""Color text with gradient end color (#009296)."""
return color(text, 1.0)
def color_mid(text: str) -> str:
"""Color text with gradient middle color."""
return color(text, 0.5)
def dim(text: str) -> str:
"""Dim/gray text."""
return f"\033[38;2;128;128;128m{text}\033[0m"
def print_startup_info(
host: str,
port: int,
database_url: str,
llm_provider: str,
llm_model: str,
embeddings_provider: str,
reranker_provider: str,
mcp_enabled: bool = False,
):
"""Print styled startup information."""
print(color_start("Starting Hindsight API..."))
print(f" {dim('URL:')} {color(f'http://{host}:{port}', 0.2)}")
print(f" {dim('Database:')} {color(database_url, 0.4)}")
print(f" {dim('LLM:')} {color(f'{llm_provider} / {llm_model}', 0.6)}")
print(f" {dim('Embeddings:')} {color(embeddings_provider, 0.8)}")
print(f" {dim('Reranker:')} {color(reranker_provider, 1.0)}")
if mcp_enabled:
print(f" {dim('MCP:')} {color_end('enabled at /mcp')}")
print()
-127
View File
@@ -1,127 +0,0 @@
"""
Command-line interface for Hindsight API.
Run the server with:
hindsight-api
Stop with Ctrl+C.
"""
import argparse
import asyncio
import atexit
import os
import signal
import sys
from typing import Optional
import uvicorn
from . import MemoryEngine
from .api import create_app
# Disable tokenizers parallelism to avoid warnings
os.environ["TOKENIZERS_PARALLELISM"] = "false"
# Global reference for cleanup
_memory: Optional[MemoryEngine] = None
def _cleanup():
"""Synchronous cleanup function to stop resources on exit."""
global _memory
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}")
def _signal_handler(signum, frame):
"""Handle SIGINT/SIGTERM to ensure cleanup."""
print(f"\nReceived signal {signum}, shutting down...")
_cleanup()
sys.exit(0)
def main():
"""Main entry point for the CLI."""
global _memory
parser = argparse.ArgumentParser(
prog="hindsight-api",
description="Hindsight API Server",
)
parser.add_argument(
"--host", default="0.0.0.0",
help="Host to bind to (default: 0.0.0.0)"
)
parser.add_argument(
"--port", type=int, default=8888,
help="Port to bind to (default: 8888)"
)
parser.add_argument(
"--log-level", default="info",
choices=["critical", "error", "warning", "info", "debug", "trace"],
help="Log level (default: info)"
)
parser.add_argument(
"--access-log", action="store_true",
help="Enable access log"
)
args = parser.parse_args()
# Register cleanup handlers
atexit.register(_cleanup)
signal.signal(signal.SIGINT, _signal_handler)
signal.signal(signal.SIGTERM, _signal_handler)
# Get configuration from environment variables
db_url = os.getenv("HINDSIGHT_API_DATABASE_URL", "pg0")
llm_provider = os.getenv("HINDSIGHT_API_LLM_PROVIDER", "groq")
llm_api_key = os.getenv("HINDSIGHT_API_LLM_API_KEY", "")
llm_model = os.getenv("HINDSIGHT_API_LLM_MODEL", "openai/gpt-oss-20b")
llm_base_url = os.getenv("HINDSIGHT_API_LLM_BASE_URL") or None
# Create MemoryEngine
_memory = MemoryEngine(
db_url=db_url,
memory_llm_provider=llm_provider,
memory_llm_api_key=llm_api_key,
memory_llm_model=llm_model,
memory_llm_base_url=llm_base_url,
)
# Create FastAPI app
app = create_app(
memory=_memory,
http_api_enabled=True,
mcp_api_enabled=True,
mcp_mount_path="/mcp",
initialize_memory=True,
)
# Prepare uvicorn config
uvicorn_config = {
"app": app,
"host": args.host,
"port": args.port,
"log_level": args.log_level,
"access_log": args.access_log,
}
print(f"\nStarting Hindsight API...")
print(f" URL: http://{args.host}:{args.port}")
print(f" Database: {db_url}")
print(f" LLM Provider: {llm_provider}")
print()
uvicorn.run(**uvicorn_config)
if __name__ == "__main__":
main()
+469
View File
@@ -0,0 +1,469 @@
"""
Centralized configuration for Hindsight API.
All environment variables and their defaults are defined here.
"""
import logging
import os
from dataclasses import dataclass
from dotenv import find_dotenv, load_dotenv
# Load .env file, searching current and parent directories (overrides existing env vars)
load_dotenv(find_dotenv(usecwd=True), override=True)
logger = logging.getLogger(__name__)
# Environment variable names
ENV_DATABASE_URL = "HINDSIGHT_API_DATABASE_URL"
ENV_LLM_PROVIDER = "HINDSIGHT_API_LLM_PROVIDER"
ENV_LLM_API_KEY = "HINDSIGHT_API_LLM_API_KEY"
ENV_LLM_MODEL = "HINDSIGHT_API_LLM_MODEL"
ENV_LLM_BASE_URL = "HINDSIGHT_API_LLM_BASE_URL"
ENV_LLM_MAX_CONCURRENT = "HINDSIGHT_API_LLM_MAX_CONCURRENT"
ENV_LLM_TIMEOUT = "HINDSIGHT_API_LLM_TIMEOUT"
ENV_LLM_GROQ_SERVICE_TIER = "HINDSIGHT_API_LLM_GROQ_SERVICE_TIER"
# Per-operation LLM configuration (optional, falls back to global LLM config)
ENV_RETAIN_LLM_PROVIDER = "HINDSIGHT_API_RETAIN_LLM_PROVIDER"
ENV_RETAIN_LLM_API_KEY = "HINDSIGHT_API_RETAIN_LLM_API_KEY"
ENV_RETAIN_LLM_MODEL = "HINDSIGHT_API_RETAIN_LLM_MODEL"
ENV_RETAIN_LLM_BASE_URL = "HINDSIGHT_API_RETAIN_LLM_BASE_URL"
ENV_REFLECT_LLM_PROVIDER = "HINDSIGHT_API_REFLECT_LLM_PROVIDER"
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_EMBEDDINGS_PROVIDER = "HINDSIGHT_API_EMBEDDINGS_PROVIDER"
ENV_EMBEDDINGS_LOCAL_MODEL = "HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL"
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"
ENV_EMBEDDINGS_OPENAI_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL"
ENV_COHERE_API_KEY = "HINDSIGHT_API_COHERE_API_KEY"
ENV_EMBEDDINGS_COHERE_MODEL = "HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL"
ENV_EMBEDDINGS_COHERE_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL"
ENV_RERANKER_COHERE_MODEL = "HINDSIGHT_API_RERANKER_COHERE_MODEL"
ENV_RERANKER_COHERE_BASE_URL = "HINDSIGHT_API_RERANKER_COHERE_BASE_URL"
# LiteLLM gateway configuration (for embeddings and reranker via LiteLLM proxy)
ENV_LITELLM_API_BASE = "HINDSIGHT_API_LITELLM_API_BASE"
ENV_LITELLM_API_KEY = "HINDSIGHT_API_LITELLM_API_KEY"
ENV_EMBEDDINGS_LITELLM_MODEL = "HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL"
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_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"
ENV_RERANKER_TEI_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT"
ENV_RERANKER_MAX_CANDIDATES = "HINDSIGHT_API_RERANKER_MAX_CANDIDATES"
ENV_RERANKER_FLASHRANK_MODEL = "HINDSIGHT_API_RERANKER_FLASHRANK_MODEL"
ENV_RERANKER_FLASHRANK_CACHE_DIR = "HINDSIGHT_API_RERANKER_FLASHRANK_CACHE_DIR"
ENV_HOST = "HINDSIGHT_API_HOST"
ENV_PORT = "HINDSIGHT_API_PORT"
ENV_LOG_LEVEL = "HINDSIGHT_API_LOG_LEVEL"
ENV_WORKERS = "HINDSIGHT_API_WORKERS"
ENV_MCP_ENABLED = "HINDSIGHT_API_MCP_ENABLED"
ENV_GRAPH_RETRIEVER = "HINDSIGHT_API_GRAPH_RETRIEVER"
ENV_MPFP_TOP_K_NEIGHBORS = "HINDSIGHT_API_MPFP_TOP_K_NEIGHBORS"
ENV_RECALL_MAX_CONCURRENT = "HINDSIGHT_API_RECALL_MAX_CONCURRENT"
ENV_RECALL_CONNECTION_BUDGET = "HINDSIGHT_API_RECALL_CONNECTION_BUDGET"
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_OBSERVATIONS_ASYNC = "HINDSIGHT_API_RETAIN_OBSERVATIONS_ASYNC"
# Optimization flags
ENV_SKIP_LLM_VERIFICATION = "HINDSIGHT_API_SKIP_LLM_VERIFICATION"
ENV_LAZY_RERANKER = "HINDSIGHT_API_LAZY_RERANKER"
# Database migrations
ENV_RUN_MIGRATIONS_ON_STARTUP = "HINDSIGHT_API_RUN_MIGRATIONS_ON_STARTUP"
# Database connection pool
ENV_DB_POOL_MIN_SIZE = "HINDSIGHT_API_DB_POOL_MIN_SIZE"
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"
# Reflect agent settings
ENV_REFLECT_MAX_ITERATIONS = "HINDSIGHT_API_REFLECT_MAX_ITERATIONS"
# Default values
DEFAULT_DATABASE_URL = "pg0"
DEFAULT_LLM_PROVIDER = "openai"
DEFAULT_LLM_MODEL = "gpt-5-mini"
DEFAULT_LLM_MAX_CONCURRENT = 32
DEFAULT_LLM_TIMEOUT = 120.0 # seconds
DEFAULT_EMBEDDINGS_PROVIDER = "local"
DEFAULT_EMBEDDINGS_LOCAL_MODEL = "BAAI/bge-small-en-v1.5"
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_MAX_CONCURRENT = 4 # Limit concurrent CPU-bound reranking to prevent thrashing
DEFAULT_RERANKER_TEI_BATCH_SIZE = 128
DEFAULT_RERANKER_TEI_MAX_CONCURRENT = 8
DEFAULT_RERANKER_MAX_CANDIDATES = 300
DEFAULT_RERANKER_FLASHRANK_MODEL = "ms-marco-MiniLM-L-12-v2" # Best balance of speed and quality
DEFAULT_RERANKER_FLASHRANK_CACHE_DIR = None # Use default cache directory
DEFAULT_EMBEDDINGS_COHERE_MODEL = "embed-english-v3.0"
DEFAULT_RERANKER_COHERE_MODEL = "rerank-english-v3.0"
# LiteLLM defaults
DEFAULT_LITELLM_API_BASE = "http://localhost:4000"
DEFAULT_EMBEDDINGS_LITELLM_MODEL = "text-embedding-3-small"
DEFAULT_RERANKER_LITELLM_MODEL = "cohere/rerank-english-v3.0"
DEFAULT_HOST = "0.0.0.0"
DEFAULT_PORT = 8888
DEFAULT_LOG_LEVEL = "info"
DEFAULT_WORKERS = 1
DEFAULT_MCP_ENABLED = True
DEFAULT_GRAPH_RETRIEVER = "link_expansion" # Options: "link_expansion", "mpfp", "bfs"
DEFAULT_MPFP_TOP_K_NEIGHBORS = 20 # Fan-out limit per node in MPFP graph traversal
DEFAULT_RECALL_MAX_CONCURRENT = 32 # Max concurrent recall operations per worker
DEFAULT_RECALL_CONNECTION_BUDGET = 4 # Max concurrent DB connections per recall operation
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_OBSERVATIONS_ASYNC = False # Run observation generation async (after retain completes)
# Database migrations
DEFAULT_RUN_MIGRATIONS_ON_STARTUP = True
# Database connection pool
DEFAULT_DB_POOL_MIN_SIZE = 5
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
# Reflect agent settings
DEFAULT_REFLECT_MAX_ITERATIONS = 10 # Max tool call iterations before forcing response
# Default MCP tool descriptions (can be customized via env vars)
DEFAULT_MCP_RETAIN_DESCRIPTION = """Store important information to long-term memory.
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"""
DEFAULT_MCP_RECALL_DESCRIPTION = """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"""
# Default embedding dimension (used by initial migration, adjusted at runtime)
EMBEDDING_DIMENSION = DEFAULT_EMBEDDING_DIMENSION
def _validate_extraction_mode(mode: str) -> str:
"""Validate and normalize extraction mode."""
mode_lower = mode.lower()
if mode_lower not in RETAIN_EXTRACTION_MODES:
logger.warning(
f"Invalid extraction mode '{mode}', must be one of {RETAIN_EXTRACTION_MODES}. "
f"Defaulting to '{DEFAULT_RETAIN_EXTRACTION_MODE}'."
)
return DEFAULT_RETAIN_EXTRACTION_MODE
return mode_lower
@dataclass
class HindsightConfig:
"""Configuration container for Hindsight API."""
# Database
database_url: str
# LLM (default, used as fallback for per-operation config)
llm_provider: str
llm_api_key: str | None
llm_model: str
llm_base_url: str | None
llm_max_concurrent: int
llm_timeout: float
# Per-operation LLM configuration (None = use default LLM config)
retain_llm_provider: str | None
retain_llm_api_key: str | None
retain_llm_model: str | None
retain_llm_base_url: str | None
reflect_llm_provider: str | None
reflect_llm_api_key: str | None
reflect_llm_model: str | None
reflect_llm_base_url: str | None
# Embeddings
embeddings_provider: str
embeddings_local_model: str
embeddings_tei_url: str | None
embeddings_openai_base_url: str | None
embeddings_cohere_base_url: str | None
# Reranker
reranker_provider: str
reranker_local_model: str
reranker_tei_url: str | None
reranker_tei_batch_size: int
reranker_tei_max_concurrent: int
reranker_max_candidates: int
reranker_cohere_base_url: str | None
# Server
host: str
port: int
log_level: str
mcp_enabled: bool
# Recall
graph_retriever: str
mpfp_top_k_neighbors: int
recall_max_concurrent: int
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_observations_async: bool
# Optimization flags
skip_llm_verification: bool
lazy_reranker: bool
# Database migrations
run_migrations_on_startup: bool
# Database connection pool
db_pool_min_size: int
db_pool_max_size: int
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
# Reflect agent settings
reflect_max_iterations: int
@classmethod
def from_env(cls) -> "HindsightConfig":
"""Create configuration from environment variables."""
return cls(
# Database
database_url=os.getenv(ENV_DATABASE_URL, DEFAULT_DATABASE_URL),
# LLM
llm_provider=os.getenv(ENV_LLM_PROVIDER, DEFAULT_LLM_PROVIDER),
llm_api_key=os.getenv(ENV_LLM_API_KEY),
llm_model=os.getenv(ENV_LLM_MODEL, DEFAULT_LLM_MODEL),
llm_base_url=os.getenv(ENV_LLM_BASE_URL) or None,
llm_max_concurrent=int(os.getenv(ENV_LLM_MAX_CONCURRENT, str(DEFAULT_LLM_MAX_CONCURRENT))),
llm_timeout=float(os.getenv(ENV_LLM_TIMEOUT, str(DEFAULT_LLM_TIMEOUT))),
# Per-operation LLM config (None = use default)
retain_llm_provider=os.getenv(ENV_RETAIN_LLM_PROVIDER) or None,
retain_llm_api_key=os.getenv(ENV_RETAIN_LLM_API_KEY) or None,
retain_llm_model=os.getenv(ENV_RETAIN_LLM_MODEL) or None,
retain_llm_base_url=os.getenv(ENV_RETAIN_LLM_BASE_URL) or None,
reflect_llm_provider=os.getenv(ENV_REFLECT_LLM_PROVIDER) or None,
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,
# 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_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_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(
os.getenv(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT))
),
reranker_max_candidates=int(os.getenv(ENV_RERANKER_MAX_CANDIDATES, str(DEFAULT_RERANKER_MAX_CANDIDATES))),
reranker_cohere_base_url=os.getenv(ENV_RERANKER_COHERE_BASE_URL) or None,
# Server
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),
mcp_enabled=os.getenv(ENV_MCP_ENABLED, str(DEFAULT_MCP_ENABLED)).lower() == "true",
# Recall
graph_retriever=os.getenv(ENV_GRAPH_RETRIEVER, DEFAULT_GRAPH_RETRIEVER),
mpfp_top_k_neighbors=int(os.getenv(ENV_MPFP_TOP_K_NEIGHBORS, str(DEFAULT_MPFP_TOP_K_NEIGHBORS))),
recall_max_concurrent=int(os.getenv(ENV_RECALL_MAX_CONCURRENT, str(DEFAULT_RECALL_MAX_CONCURRENT))),
recall_connection_budget=int(
os.getenv(ENV_RECALL_CONNECTION_BUDGET, str(DEFAULT_RECALL_CONNECTION_BUDGET))
),
mental_model_refresh_concurrency=int(
os.getenv(ENV_MENTAL_MODEL_REFRESH_CONCURRENCY, str(DEFAULT_MENTAL_MODEL_REFRESH_CONCURRENCY))
),
# 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))
),
retain_chunk_size=int(os.getenv(ENV_RETAIN_CHUNK_SIZE, str(DEFAULT_RETAIN_CHUNK_SIZE))),
retain_extract_causal_links=os.getenv(
ENV_RETAIN_EXTRACT_CAUSAL_LINKS, str(DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS)
).lower()
== "true",
retain_extraction_mode=_validate_extraction_mode(
os.getenv(ENV_RETAIN_EXTRACTION_MODE, DEFAULT_RETAIN_EXTRACTION_MODE)
),
retain_observations_async=os.getenv(
ENV_RETAIN_OBSERVATIONS_ASYNC, str(DEFAULT_RETAIN_OBSERVATIONS_ASYNC)
).lower()
== "true",
# Database migrations
run_migrations_on_startup=os.getenv(ENV_RUN_MIGRATIONS_ON_STARTUP, "true").lower() == "true",
# Database connection pool
db_pool_min_size=int(os.getenv(ENV_DB_POOL_MIN_SIZE, str(DEFAULT_DB_POOL_MIN_SIZE))),
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))
),
# Reflect agent settings
reflect_max_iterations=int(os.getenv(ENV_REFLECT_MAX_ITERATIONS, str(DEFAULT_REFLECT_MAX_ITERATIONS))),
)
def get_llm_base_url(self) -> str:
"""Get the LLM base URL, with provider-specific defaults."""
if self.llm_base_url:
return self.llm_base_url
provider = self.llm_provider.lower()
if provider == "groq":
return "https://api.groq.com/openai/v1"
elif provider == "ollama":
return "http://localhost:11434/v1"
elif provider == "lmstudio":
return "http://localhost:1234/v1"
else:
return ""
def get_python_log_level(self) -> int:
"""Get the Python logging level from the configured log level string."""
log_level_map = {
"critical": logging.CRITICAL,
"error": logging.ERROR,
"warning": logging.WARNING,
"info": logging.INFO,
"debug": logging.DEBUG,
"trace": logging.DEBUG, # Python doesn't have TRACE, use DEBUG
}
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
)
def log_config(self) -> None:
"""Log the current configuration (without sensitive values)."""
logger.info(f"Database: {self.database_url}")
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
retain_model = self.retain_llm_model or self.llm_model
logger.info(f"LLM (retain): provider={retain_provider}, model={retain_model}")
if self.reflect_llm_provider or self.reflect_llm_model:
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}")
logger.info(f"Embeddings: provider={self.embeddings_provider}")
logger.info(f"Reranker: provider={self.reranker_provider}")
logger.info(f"Graph retriever: {self.graph_retriever}")
# Cached config instance
_config_cache: HindsightConfig | None = None
def get_config() -> HindsightConfig:
"""Get the cached configuration, loading from environment on first call."""
global _config_cache
if _config_cache is None:
_config_cache = HindsightConfig.from_env()
return _config_cache
def clear_config_cache() -> None:
"""Clear the config cache. Useful for testing or reloading config."""
global _config_cache
_config_cache = None
+204
View File
@@ -0,0 +1,204 @@
"""
Daemon mode support for Hindsight API.
Provides idle timeout and lockfile management for running as a background daemon.
"""
import asyncio
import fcntl
import logging
import os
import sys
import time
from pathlib import Path
logger = logging.getLogger(__name__)
# Default daemon configuration
DEFAULT_DAEMON_PORT = 8889
DEFAULT_IDLE_TIMEOUT = 0 # 0 = no auto-exit (hindsight-embed passes its own timeout)
LOCKFILE_PATH = Path.home() / ".hindsight" / "daemon.lock"
DAEMON_LOG_PATH = Path.home() / ".hindsight" / "daemon.log"
class IdleTimeoutMiddleware:
"""ASGI middleware that tracks activity and exits after idle timeout."""
def __init__(self, app, idle_timeout: int = DEFAULT_IDLE_TIMEOUT):
self.app = app
self.idle_timeout = idle_timeout
self.last_activity = time.time()
self._checker_task = None
async def __call__(self, scope, receive, send):
# Update activity timestamp on each request
self.last_activity = time.time()
await self.app(scope, receive, send)
def start_idle_checker(self):
"""Start the background task that checks for idle timeout."""
self._checker_task = asyncio.create_task(self._check_idle())
async def _check_idle(self):
"""Background task that exits the process after idle timeout."""
# If idle_timeout is 0, don't auto-exit
if self.idle_timeout <= 0:
return
while True:
await asyncio.sleep(30) # Check every 30 seconds
idle_time = time.time() - self.last_activity
if idle_time > self.idle_timeout:
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)
class DaemonLock:
"""
File-based lock to prevent multiple daemon instances.
Uses fcntl.flock for atomic locking on Unix systems.
"""
def __init__(self, lockfile: Path = LOCKFILE_PATH):
self.lockfile = lockfile
self._fd = None
def acquire(self) -> bool:
"""
Try to acquire the daemon lock.
Returns True if lock acquired, False if another daemon is running.
"""
self.lockfile.parent.mkdir(parents=True, exist_ok=True)
try:
self._fd = open(self.lockfile, "w")
fcntl.flock(self._fd.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
# Write PID for debugging
self._fd.write(str(os.getpid()))
self._fd.flush()
return True
except (IOError, OSError):
# Lock is held by another process
if self._fd:
self._fd.close()
self._fd = None
return False
def release(self):
"""Release the daemon lock."""
if self._fd:
try:
fcntl.flock(self._fd.fileno(), fcntl.LOCK_UN)
self._fd.close()
except Exception:
pass
finally:
self._fd = None
# Remove lockfile
try:
self.lockfile.unlink()
except Exception:
pass
def is_locked(self) -> bool:
"""Check if the lock is held by another process."""
if not self.lockfile.exists():
return False
try:
fd = open(self.lockfile, "r")
fcntl.flock(fd.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
# We got the lock, so no one else has it
fcntl.flock(fd.fileno(), fcntl.LOCK_UN)
fd.close()
return False
except (IOError, OSError):
return True
def get_pid(self) -> int | None:
"""Get the PID of the daemon holding the lock."""
if not self.lockfile.exists():
return None
try:
with open(self.lockfile, "r") as f:
return int(f.read().strip())
except (ValueError, IOError):
return None
def daemonize():
"""
Fork the current process into a background daemon.
Uses double-fork technique to properly detach from terminal.
"""
# First fork
pid = os.fork()
if pid > 0:
# Parent exits
sys.exit(0)
# Create new session
os.setsid()
# Second fork to prevent zombie processes
pid = os.fork()
if pid > 0:
sys.exit(0)
# Redirect standard file descriptors to log file
DAEMON_LOG_PATH.parent.mkdir(parents=True, exist_ok=True)
sys.stdout.flush()
sys.stderr.flush()
# Redirect stdin to /dev/null
with open("/dev/null", "r") as devnull:
os.dup2(devnull.fileno(), sys.stdin.fileno())
# Redirect stdout/stderr to log file
log_fd = open(DAEMON_LOG_PATH, "a")
os.dup2(log_fd.fileno(), sys.stdout.fileno())
os.dup2(log_fd.fileno(), sys.stderr.fileno())
def check_daemon_running(port: int = DEFAULT_DAEMON_PORT) -> bool:
"""Check if a daemon is running and responsive on the given port."""
import socket
try:
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
sock.settimeout(1)
result = sock.connect_ex(("127.0.0.1", port))
sock.close()
return result == 0
except Exception:
return False
def stop_daemon(port: int = DEFAULT_DAEMON_PORT) -> bool:
"""Stop a running daemon by sending SIGTERM to the process."""
lock = DaemonLock()
pid = lock.get_pid()
if pid is None:
return False
try:
import signal
os.kill(pid, signal.SIGTERM)
# Wait for process to exit
for _ in range(50): # Wait up to 5 seconds
time.sleep(0.1)
try:
os.kill(pid, 0) # Check if process exists
except OSError:
return True # Process exited
return False
except OSError:
return False
+20 -9
View File
@@ -7,24 +7,30 @@ This package contains all the implementation details of the memory engine:
- Supporting modules: embeddings, cross_encoder, entity_resolver, etc.
"""
from .memory_engine import MemoryEngine
from .cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
from .db_utils import acquire_with_retry
from .embeddings import Embeddings, LocalSTEmbeddings, RemoteTEIEmbeddings
from .cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
from .llm_wrapper import LLMConfig
from .memory_engine import (
MemoryEngine,
UnqualifiedTableError,
fq_table,
get_current_schema,
validate_sql_schema,
)
from .response_models import MemoryFact, RecallResult, ReflectResult
from .search.trace import (
SearchTrace,
QueryInfo,
EntryPoint,
NodeVisit,
WeightComponents,
LinkInfo,
NodeVisit,
PruningDecision,
SearchSummary,
QueryInfo,
SearchPhaseMetrics,
SearchSummary,
SearchTrace,
WeightComponents,
)
from .search.tracer import SearchTracer
from .llm_wrapper import LLMConfig
from .response_models import RecallResult, ReflectResult, MemoryFact
__all__ = [
"MemoryEngine",
@@ -49,4 +55,9 @@ __all__ = [
"RecallResult",
"ReflectResult",
"MemoryFact",
# Schema safety utilities
"fq_table",
"get_current_schema",
"validate_sql_schema",
"UnqualifiedTableError",
]
@@ -3,26 +3,45 @@ Cross-encoder abstraction for reranking.
Provides an interface for reranking with different backends.
Configuration via environment variables:
- HINDSIGHT_API_RERANKER_PROVIDER: "local" (default) or "tei"
For local provider:
- HINDSIGHT_API_RERANKER_LOCAL_MODEL: Model name (default: cross-encoder/ms-marco-MiniLM-L-6-v2)
For TEI provider:
- HINDSIGHT_API_RERANKER_TEI_URL: TEI server URL (required)
Configuration via environment variables - see hindsight_api.config for all env var names.
"""
from abc import ABC, abstractmethod
from typing import List, Tuple, Optional
import asyncio
import logging
import os
from abc import ABC, abstractmethod
from concurrent.futures import ThreadPoolExecutor
import httpx
logger = logging.getLogger(__name__)
from ..config import (
DEFAULT_LITELLM_API_BASE,
DEFAULT_RERANKER_COHERE_MODEL,
DEFAULT_RERANKER_FLASHRANK_CACHE_DIR,
DEFAULT_RERANKER_FLASHRANK_MODEL,
DEFAULT_RERANKER_LITELLM_MODEL,
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT,
DEFAULT_RERANKER_LOCAL_MODEL,
DEFAULT_RERANKER_PROVIDER,
DEFAULT_RERANKER_TEI_BATCH_SIZE,
DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
ENV_COHERE_API_KEY,
ENV_LITELLM_API_BASE,
ENV_LITELLM_API_KEY,
ENV_RERANKER_COHERE_BASE_URL,
ENV_RERANKER_COHERE_MODEL,
ENV_RERANKER_FLASHRANK_CACHE_DIR,
ENV_RERANKER_FLASHRANK_MODEL,
ENV_RERANKER_LITELLM_MODEL,
ENV_RERANKER_LOCAL_MAX_CONCURRENT,
ENV_RERANKER_LOCAL_MODEL,
ENV_RERANKER_PROVIDER,
ENV_RERANKER_TEI_BATCH_SIZE,
ENV_RERANKER_TEI_MAX_CONCURRENT,
ENV_RERANKER_TEI_URL,
)
# Default model for local cross-encoder
DEFAULT_RERANKER_MODEL = "cross-encoder/ms-marco-MiniLM-L-6-v2"
logger = logging.getLogger(__name__)
class CrossEncoderModel(ABC):
@@ -49,7 +68,7 @@ class CrossEncoderModel(ABC):
pass
@abstractmethod
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs for relevance.
@@ -72,25 +91,34 @@ class LocalSTCrossEncoder(CrossEncoderModel):
- Fast inference (~80ms for 100 pairs on CPU)
- Small model (80MB)
- Trained for passage re-ranking
Uses a dedicated thread pool to limit concurrent CPU-bound work.
"""
def __init__(self, model_name: Optional[str] = None):
# Shared executor across all instances (one model loaded anyway)
_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):
"""
Initialize local SentenceTransformers cross-encoder.
Args:
model_name: Name of the CrossEncoder model to use.
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.
"""
self.model_name = model_name or DEFAULT_RERANKER_MODEL
self.model_name = model_name or DEFAULT_RERANKER_LOCAL_MODEL
self._model = None
LocalSTCrossEncoder._max_concurrent = max_concurrent
@property
def provider_name(self) -> str:
return "local"
async def initialize(self) -> None:
"""Load the cross-encoder model."""
"""Load the cross-encoder model and initialize the executor."""
if self._model is not None:
return
@@ -102,19 +130,30 @@ 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}")
# 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
self._model = CrossEncoder(
self.model_name,
model_kwargs={"low_cpu_mem_usage": False, "device_map": None},
)
logger.info("Reranker: local provider initialized")
self._model = CrossEncoder(self.model_name)
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
# Initialize shared executor (limited workers naturally limits concurrency)
if LocalSTCrossEncoder._executor is None:
LocalSTCrossEncoder._executor = ThreadPoolExecutor(
max_workers=LocalSTCrossEncoder._max_concurrent,
thread_name_prefix="reranker",
)
logger.info(f"Reranker: local provider initialized (max_concurrent={LocalSTCrossEncoder._max_concurrent})")
else:
logger.info("Reranker: local provider initialized (using existing executor)")
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs for relevance.
Uses a dedicated thread pool with limited workers to prevent CPU thrashing.
Args:
pairs: List of (query, document) tuples to score
@@ -123,8 +162,14 @@ class LocalSTCrossEncoder(CrossEncoderModel):
"""
if self._model is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
scores = self._model.predict(pairs, show_progress_bar=False)
return scores.tolist() if hasattr(scores, 'tolist') else list(scores)
# Use dedicated executor - limited workers naturally limits concurrency
loop = asyncio.get_event_loop()
scores = await loop.run_in_executor(
LocalSTCrossEncoder._executor,
lambda: self._model.predict(pairs, show_progress_bar=False),
)
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
class RemoteTEICrossEncoder(CrossEncoderModel):
@@ -135,13 +180,21 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
See: https://github.com/huggingface/text-embeddings-inference
Note: The TEI server must be running a cross-encoder/reranker model.
Requests are made in parallel with configurable batch size and max concurrency (backpressure).
Uses a GLOBAL semaphore to limit concurrent requests across ALL recall operations.
"""
# Global semaphore shared across all instances and calls to prevent thundering herd
_global_semaphore: asyncio.Semaphore | None = None
_global_max_concurrent: int = DEFAULT_RERANKER_TEI_MAX_CONCURRENT
def __init__(
self,
base_url: str,
timeout: float = 30.0,
batch_size: int = 32,
batch_size: int = DEFAULT_RERANKER_TEI_BATCH_SIZE,
max_concurrent: int = DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
max_retries: int = 3,
retry_delay: float = 0.5,
):
@@ -151,75 +204,246 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
Args:
base_url: Base URL of the TEI server (e.g., "http://localhost:8080")
timeout: Request timeout in seconds (default: 30.0)
batch_size: Maximum batch size for rerank requests (default: 32)
batch_size: Maximum batch size for rerank requests (default: 128)
max_concurrent: Maximum concurrent requests for backpressure (default: 8).
This is a GLOBAL limit across all parallel recall operations.
max_retries: Maximum number of retries for failed requests (default: 3)
retry_delay: Initial delay between retries in seconds, doubles each retry (default: 0.5)
"""
self.base_url = base_url.rstrip("/")
self.timeout = timeout
self.batch_size = batch_size
self.max_concurrent = max_concurrent
self.max_retries = max_retries
self.retry_delay = retry_delay
self._client: Optional[httpx.Client] = None
self._model_id: Optional[str] = None
self._async_client: httpx.AsyncClient | None = None
self._model_id: str | None = None
# Update global semaphore if max_concurrent changed
if (
RemoteTEICrossEncoder._global_semaphore is None
or RemoteTEICrossEncoder._global_max_concurrent != max_concurrent
):
RemoteTEICrossEncoder._global_max_concurrent = max_concurrent
RemoteTEICrossEncoder._global_semaphore = asyncio.Semaphore(max_concurrent)
@property
def provider_name(self) -> str:
return "tei"
def _request_with_retry(self, method: str, url: str, **kwargs) -> httpx.Response:
"""Make an HTTP request with automatic retries on transient errors."""
import time
async def _async_request_with_retry(
self,
client: httpx.AsyncClient,
semaphore: asyncio.Semaphore,
method: str,
url: str,
**kwargs,
) -> httpx.Response:
"""Make an async HTTP request with automatic retries on transient errors and semaphore for backpressure."""
last_error = None
delay = self.retry_delay
for attempt in range(self.max_retries + 1):
try:
if method == "GET":
response = self._client.get(url, **kwargs)
else:
response = self._client.post(url, **kwargs)
response.raise_for_status()
return response
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
last_error = e
if attempt < self.max_retries:
logger.warning(f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
time.sleep(delay)
delay *= 2 # Exponential backoff
except httpx.HTTPStatusError as e:
# Retry on 5xx server errors
if e.response.status_code >= 500 and attempt < self.max_retries:
async with semaphore:
for attempt in range(self.max_retries + 1):
try:
if method == "GET":
response = await client.get(url, **kwargs)
else:
response = await client.post(url, **kwargs)
response.raise_for_status()
return response
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
last_error = e
logger.warning(f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
time.sleep(delay)
delay *= 2
else:
raise
if attempt < self.max_retries:
logger.warning(
f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. "
f"Retrying in {delay}s..."
)
await asyncio.sleep(delay)
delay *= 2 # Exponential backoff
except httpx.HTTPStatusError as e:
# Retry on 5xx server errors
if e.response.status_code >= 500 and attempt < self.max_retries:
last_error = e
logger.warning(
f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. "
f"Retrying in {delay}s..."
)
await asyncio.sleep(delay)
delay *= 2
else:
raise
raise last_error
async def initialize(self) -> None:
"""Initialize the HTTP client and verify server connectivity."""
if self._client is not None:
if self._async_client is not None:
return
logger.info(f"Reranker: initializing TEI provider at {self.base_url}")
self._client = httpx.Client(timeout=self.timeout)
logger.info(
f"Reranker: initializing TEI provider at {self.base_url} "
f"(batch_size={self.batch_size}, max_concurrent={self.max_concurrent})"
)
self._async_client = httpx.AsyncClient(timeout=self.timeout)
# Verify server is reachable and get model info
# Use a temporary semaphore for initialization
init_semaphore = asyncio.Semaphore(1)
try:
response = self._request_with_retry("GET", f"{self.base_url}/info")
response = await self._async_request_with_retry(
self._async_client, init_semaphore, "GET", f"{self.base_url}/info"
)
info = response.json()
self._model_id = info.get("model_id", "unknown")
logger.info(f"Reranker: TEI provider initialized (model: {self._model_id})")
except httpx.HTTPError as e:
self._async_client = None
raise RuntimeError(f"Failed to connect to TEI server at {self.base_url}: {e}")
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
async def _rerank_query_group(
self,
client: httpx.AsyncClient,
semaphore: asyncio.Semaphore,
query: str,
texts: list[str],
) -> list[tuple[int, float]]:
"""Rerank a single query group and return list of (original_index, score) tuples."""
try:
response = await self._async_request_with_retry(
client,
semaphore,
"POST",
f"{self.base_url}/rerank",
json={
"query": query,
"texts": texts,
"return_text": False,
},
)
results = response.json()
# TEI returns results sorted by score descending, with original index
return [(result["index"], result["score"]) for result in results]
except httpx.HTTPError as e:
raise RuntimeError(f"TEI rerank request failed: {e}")
async def _predict_async(self, pairs: list[tuple[str, str]]) -> list[float]:
"""Async implementation of predict that runs requests in parallel with backpressure."""
if not pairs:
return []
# Group all pairs by query
query_groups: dict[str, list[tuple[int, str]]] = {}
for idx, (query, text) in enumerate(pairs):
if query not in query_groups:
query_groups[query] = []
query_groups[query].append((idx, text))
# Split each query group into batches
tasks_info: list[tuple[str, list[int], list[str]]] = [] # (query, indices, texts)
for query, indexed_texts in query_groups.items():
indices = [idx for idx, _ in indexed_texts]
texts = [text for _, text in indexed_texts]
# Split into batches
for i in range(0, len(texts), self.batch_size):
batch_indices = indices[i : i + self.batch_size]
batch_texts = texts[i : i + self.batch_size]
tasks_info.append((query, batch_indices, batch_texts))
# Run all requests in parallel with GLOBAL semaphore for backpressure
# This ensures max_concurrent is respected across ALL parallel recall operations
all_scores = [0.0] * len(pairs)
semaphore = RemoteTEICrossEncoder._global_semaphore
tasks = [
self._rerank_query_group(self._async_client, semaphore, query, texts) for query, _, texts in tasks_info
]
results = await asyncio.gather(*tasks)
# Map scores back to original positions
for (_, indices, _), result_scores in zip(tasks_info, results):
for original_idx_in_batch, score in result_scores:
global_idx = indices[original_idx_in_batch]
all_scores[global_idx] = score
return all_scores
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using the remote TEI reranker.
Requests are made in parallel with configurable backpressure.
Args:
pairs: List of (query, document) tuples to score
Returns:
List of relevance scores
"""
if self._async_client is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
return await self._predict_async(pairs)
class CohereCrossEncoder(CrossEncoderModel):
"""
Cohere cross-encoder implementation using the Cohere Rerank API.
Supports rerank-english-v3.0 and rerank-multilingual-v3.0 models.
"""
def __init__(
self,
api_key: str,
model: str = DEFAULT_RERANKER_COHERE_MODEL,
base_url: str | None = None,
timeout: float = 60.0,
):
"""
Initialize Cohere cross-encoder client.
Args:
api_key: Cohere API key
model: Cohere rerank model name (default: rerank-english-v3.0)
base_url: Custom base URL for Cohere-compatible API (e.g., Azure-hosted endpoint)
timeout: Request timeout in seconds (default: 60.0)
"""
self.api_key = api_key
self.model = model
self.base_url = base_url
self.timeout = timeout
self._client = None
@property
def provider_name(self) -> str:
return "cohere"
async def initialize(self) -> None:
"""Initialize the Cohere client."""
if self._client is not None:
return
try:
import cohere
except ImportError:
raise ImportError("cohere is required for CohereCrossEncoder. Install it with: pip install cohere")
base_url_msg = f" at {self.base_url}" if self.base_url else ""
logger.info(f"Reranker: initializing Cohere provider with model {self.model}{base_url_msg}")
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
client_kwargs = {"api_key": self.api_key, "timeout": self.timeout}
if self.base_url:
client_kwargs["base_url"] = self.base_url
self._client = cohere.Client(**client_kwargs)
logger.info("Reranker: Cohere provider initialized")
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using the Cohere Rerank API.
Args:
pairs: List of (query, document) tuples to score
@@ -232,50 +456,312 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
if not pairs:
return []
all_scores = []
# Run sync Cohere API calls in thread pool
loop = asyncio.get_event_loop()
return await loop.run_in_executor(None, self._predict_sync, pairs)
# Process in batches
for i in range(0, len(pairs), self.batch_size):
batch = pairs[i:i + self.batch_size]
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
"""Synchronous predict implementation for Cohere API."""
# Group pairs by query for efficient batching
# Cohere rerank expects one query with multiple documents
query_groups: dict[str, list[tuple[int, str]]] = {}
for idx, (query, text) in enumerate(pairs):
if query not in query_groups:
query_groups[query] = []
query_groups[query].append((idx, text))
# TEI rerank endpoint expects query and texts separately
# All pairs in a batch should have the same query for optimal performance
# but we handle mixed queries by making separate requests per unique query
query_groups: dict[str, list[tuple[int, str]]] = {}
for idx, (query, text) in enumerate(batch):
if query not in query_groups:
query_groups[query] = []
query_groups[query].append((idx, text))
all_scores = [0.0] * len(pairs)
batch_scores = [0.0] * len(batch)
for query, indexed_texts in query_groups.items():
texts = [text for _, text in indexed_texts]
indices = [idx for idx, _ in indexed_texts]
for query, indexed_texts in query_groups.items():
texts = [text for _, text in indexed_texts]
indices = [idx for idx, _ in indexed_texts]
response = self._client.rerank(
query=query,
documents=texts,
model=self.model,
return_documents=False,
)
try:
response = self._request_with_retry(
"POST",
f"{self.base_url}/rerank",
json={
"query": query,
"texts": texts,
"return_text": False,
},
)
results = response.json()
# Map scores back to original positions
for result in response.results:
original_idx = result.index
score = result.relevance_score
all_scores[indices[original_idx]] = score
# TEI returns results sorted by score descending, with original index
for result in results:
original_idx = result["index"]
score = result["score"]
# Map back to batch position
batch_scores[indices[original_idx]] = score
return all_scores
except httpx.HTTPError as e:
raise RuntimeError(f"TEI rerank request failed: {e}")
all_scores.extend(batch_scores)
class RRFPassthroughCrossEncoder(CrossEncoderModel):
"""
Passthrough cross-encoder that preserves RRF scores without neural reranking.
This is useful for:
- Testing retrieval quality without reranking overhead
- Deployments where reranking latency is unacceptable
- Debugging to isolate retrieval vs reranking issues
"""
def __init__(self):
"""Initialize RRF passthrough cross-encoder."""
pass
@property
def provider_name(self) -> str:
return "rrf"
async def initialize(self) -> None:
"""No initialization needed."""
logger.info("Reranker: RRF passthrough provider initialized (neural reranking disabled)")
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Return neutral scores - actual ranking uses RRF scores from retrieval.
Args:
pairs: List of (query, document) tuples (ignored)
Returns:
List of 0.5 scores (neutral, lets RRF scores dominate)
"""
# Return neutral scores so RRF ranking is preserved
return [0.5] * len(pairs)
class FlashRankCrossEncoder(CrossEncoderModel):
"""
FlashRank cross-encoder implementation.
FlashRank is an ultra-lite reranking library that runs on CPU without
requiring PyTorch or Transformers. It's ideal for serverless deployments
with minimal cold-start overhead.
Available models:
- ms-marco-TinyBERT-L-2-v2: Fastest, ~4MB
- ms-marco-MiniLM-L-12-v2: Best quality, ~34MB (default)
- rank-T5-flan: Best zero-shot, ~110MB
- ms-marco-MultiBERT-L-12: Multi-lingual, ~150MB
"""
# Shared executor for CPU-bound reranking
_executor: ThreadPoolExecutor | None = None
_max_concurrent: int = 4
def __init__(
self,
model_name: str | None = None,
cache_dir: str | None = None,
max_length: int = 512,
max_concurrent: int = 4,
):
"""
Initialize FlashRank cross-encoder.
Args:
model_name: FlashRank model name. Default: ms-marco-MiniLM-L-12-v2
cache_dir: Directory to cache downloaded models. Default: system cache
max_length: Maximum sequence length for reranking. Default: 512
max_concurrent: Maximum concurrent reranking calls. Default: 4
"""
self.model_name = model_name or DEFAULT_RERANKER_FLASHRANK_MODEL
self.cache_dir = cache_dir or DEFAULT_RERANKER_FLASHRANK_CACHE_DIR
self.max_length = max_length
self._ranker = None
FlashRankCrossEncoder._max_concurrent = max_concurrent
@property
def provider_name(self) -> str:
return "flashrank"
async def initialize(self) -> None:
"""Load the FlashRank model."""
if self._ranker is not None:
return
try:
from flashrank import Ranker # type: ignore[import-untyped]
except ImportError:
raise ImportError("flashrank is required for FlashRankCrossEncoder. Install it with: pip install flashrank")
logger.info(f"Reranker: initializing FlashRank provider with model {self.model_name}")
# Initialize ranker with optional cache directory
ranker_kwargs = {"model_name": self.model_name, "max_length": self.max_length}
if self.cache_dir:
ranker_kwargs["cache_dir"] = self.cache_dir
self._ranker = Ranker(**ranker_kwargs)
# Initialize shared executor
if FlashRankCrossEncoder._executor is None:
FlashRankCrossEncoder._executor = ThreadPoolExecutor(
max_workers=FlashRankCrossEncoder._max_concurrent,
thread_name_prefix="flashrank",
)
logger.info(
f"Reranker: FlashRank provider initialized (max_concurrent={FlashRankCrossEncoder._max_concurrent})"
)
else:
logger.info("Reranker: FlashRank provider initialized (using existing executor)")
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
"""Synchronous predict - processes each query group."""
from flashrank import RerankRequest # type: ignore[import-untyped]
if not pairs:
return []
# Group pairs by query
query_groups: dict[str, list[tuple[int, str]]] = {}
for idx, (query, text) in enumerate(pairs):
if query not in query_groups:
query_groups[query] = []
query_groups[query].append((idx, text))
all_scores = [0.0] * len(pairs)
for query, indexed_texts in query_groups.items():
# Build passages list for FlashRank
passages = [{"id": i, "text": text} for i, (_, text) in enumerate(indexed_texts)]
global_indices = [idx for idx, _ in indexed_texts]
# Create rerank request
request = RerankRequest(query=query, passages=passages)
results = self._ranker.rerank(request)
# Map scores back to original positions
for result in results:
local_idx = result["id"]
score = result["score"]
global_idx = global_indices[local_idx]
all_scores[global_idx] = score
return all_scores
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using FlashRank.
Args:
pairs: List of (query, document) tuples to score
Returns:
List of relevance scores (higher = more relevant)
"""
if self._ranker is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
# Run in thread pool to avoid blocking event loop
loop = asyncio.get_event_loop()
return await loop.run_in_executor(FlashRankCrossEncoder._executor, self._predict_sync, pairs)
class LiteLLMCrossEncoder(CrossEncoderModel):
"""
LiteLLM cross-encoder implementation using LiteLLM proxy's /rerank endpoint.
LiteLLM provides a unified interface for multiple reranking providers via
the Cohere-compatible /rerank endpoint.
See: https://docs.litellm.ai/docs/rerank
Supported providers via LiteLLM:
- Cohere (rerank-english-v3.0, etc.) - prefix with cohere/
- Together AI - prefix with together_ai/
- Azure AI - prefix with azure_ai/
- Jina AI - prefix with jina_ai/
- AWS Bedrock - prefix with bedrock/
- Voyage AI - prefix with voyage/
"""
def __init__(
self,
api_base: str = DEFAULT_LITELLM_API_BASE,
api_key: str | None = None,
model: str = DEFAULT_RERANKER_LITELLM_MODEL,
timeout: float = 60.0,
):
"""
Initialize LiteLLM cross-encoder client.
Args:
api_base: Base URL of the LiteLLM proxy (default: http://localhost:4000)
api_key: API key for the LiteLLM proxy (optional, depends on proxy config)
model: Reranking model name (default: cohere/rerank-english-v3.0)
Use provider prefix (e.g., cohere/, together_ai/, voyage/)
timeout: Request timeout in seconds (default: 60.0)
"""
self.api_base = api_base.rstrip("/")
self.api_key = api_key
self.model = model
self.timeout = timeout
self._async_client: httpx.AsyncClient | None = None
@property
def provider_name(self) -> str:
return "litellm"
async def initialize(self) -> None:
"""Initialize the async HTTP client."""
if self._async_client is not None:
return
logger.info(f"Reranker: initializing LiteLLM provider at {self.api_base} with model {self.model}")
headers = {"Content-Type": "application/json"}
if self.api_key:
headers["Authorization"] = f"Bearer {self.api_key}"
self._async_client = httpx.AsyncClient(timeout=self.timeout, headers=headers)
logger.info("Reranker: LiteLLM provider initialized")
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using the LiteLLM proxy's /rerank endpoint.
Args:
pairs: List of (query, document) tuples to score
Returns:
List of relevance scores
"""
if self._async_client is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
if not pairs:
return []
# Group pairs by query (LiteLLM rerank expects one query with multiple documents)
query_groups: dict[str, list[tuple[int, str]]] = {}
for idx, (query, text) in enumerate(pairs):
if query not in query_groups:
query_groups[query] = []
query_groups[query].append((idx, text))
all_scores = [0.0] * len(pairs)
for query, indexed_texts in query_groups.items():
texts = [text for _, text in indexed_texts]
indices = [idx for idx, _ in indexed_texts]
# LiteLLM /rerank follows Cohere API format
response = await self._async_client.post(
f"{self.api_base}/rerank",
json={
"model": self.model,
"query": query,
"documents": texts,
"top_n": len(texts), # Return all scores
},
)
response.raise_for_status()
result = response.json()
# Map scores back to original positions
# Response format: {"results": [{"index": 0, "relevance_score": 0.9}, ...]}
for item in result.get("results", []):
original_idx = item["index"]
score = item.get("relevance_score", item.get("score", 0.0))
all_scores[indices[original_idx]] = score
return all_scores
@@ -284,32 +770,46 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
"""
Create a CrossEncoderModel instance based on environment variables.
Environment variables:
- HINDSIGHT_API_RERANKER_PROVIDER: "local" (default) or "tei"
For local provider:
- HINDSIGHT_API_RERANKER_LOCAL_MODEL: Model name (default: cross-encoder/ms-marco-MiniLM-L-6-v2)
For TEI provider:
- HINDSIGHT_API_RERANKER_TEI_URL: TEI server URL (required)
See hindsight_api.config for environment variable names and defaults.
Returns:
Configured CrossEncoderModel instance
"""
provider = os.environ.get("HINDSIGHT_API_RERANKER_PROVIDER", "local").lower()
provider = os.environ.get(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER).lower()
if provider == "tei":
url = os.environ.get("HINDSIGHT_API_RERANKER_TEI_URL")
url = os.environ.get(ENV_RERANKER_TEI_URL)
if not url:
raise ValueError(
"HINDSIGHT_API_RERANKER_TEI_URL is required when HINDSIGHT_API_RERANKER_PROVIDER is 'tei'"
)
return RemoteTEICrossEncoder(base_url=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)
elif provider == "local":
model = os.environ.get("HINDSIGHT_API_RERANKER_LOCAL_MODEL")
model_name = model or DEFAULT_RERANKER_MODEL
return LocalSTCrossEncoder(model_name=model_name)
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=model_name, max_concurrent=max_concurrent)
elif provider == "cohere":
api_key = os.environ.get(ENV_COHERE_API_KEY)
if not api_key:
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_RERANKER_PROVIDER} is 'cohere'")
model = os.environ.get(ENV_RERANKER_COHERE_MODEL, DEFAULT_RERANKER_COHERE_MODEL)
base_url = os.environ.get(ENV_RERANKER_COHERE_BASE_URL) or None
return CohereCrossEncoder(api_key=api_key, model=model, base_url=base_url)
elif provider == "flashrank":
model = os.environ.get(ENV_RERANKER_FLASHRANK_MODEL, DEFAULT_RERANKER_FLASHRANK_MODEL)
cache_dir = os.environ.get(ENV_RERANKER_FLASHRANK_CACHE_DIR, DEFAULT_RERANKER_FLASHRANK_CACHE_DIR)
return FlashRankCrossEncoder(model_name=model, cache_dir=cache_dir)
elif provider == "litellm":
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
api_key = os.environ.get(ENV_LITELLM_API_KEY)
model = os.environ.get(ENV_RERANKER_LITELLM_MODEL, DEFAULT_RERANKER_LITELLM_MODEL)
return LiteLLMCrossEncoder(api_base=api_base, api_key=api_key, model=model)
elif provider == "rrf":
return RRFPassthroughCrossEncoder()
else:
raise ValueError(
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei'"
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'flashrank', 'litellm', 'rrf'"
)
@@ -0,0 +1,284 @@
"""
Database connection budget management.
Limits concurrent database connections per operation to prevent
a single operation (e.g., recall with parallel queries) from
exhausting the connection pool.
"""
import asyncio
import logging
import uuid
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, AsyncIterator
if TYPE_CHECKING:
import asyncpg
logger = logging.getLogger(__name__)
@dataclass
class OperationBudget:
"""
Tracks connection budget for a single operation.
Each operation gets a semaphore limiting its concurrent connections.
"""
operation_id: str
max_connections: int
semaphore: asyncio.Semaphore = field(init=False)
active_count: int = field(default=0, init=False)
def __post_init__(self):
self.semaphore = asyncio.Semaphore(self.max_connections)
class ConnectionBudgetManager:
"""
Manages per-operation connection budgets.
Usage:
manager = ConnectionBudgetManager(default_budget=4)
# Start an operation
async with manager.operation(max_connections=2) as op:
# Acquire connections within the budget
async with op.acquire(pool) as conn:
await conn.fetch(...)
# Multiple connections respect the budget
async with op.acquire(pool) as conn1, op.acquire(pool) as conn2:
# At most 2 concurrent connections for this operation
...
"""
def __init__(self, default_budget: int = 4):
"""
Initialize the budget manager.
Args:
default_budget: Default max connections per operation
"""
self.default_budget = default_budget
self._operations: dict[str, OperationBudget] = {}
self._lock = asyncio.Lock()
@asynccontextmanager
async def operation(
self,
max_connections: int | None = None,
operation_id: str | None = None,
) -> AsyncIterator["BudgetedOperation"]:
"""
Create a budgeted operation context.
Args:
max_connections: Max concurrent connections for this operation.
Defaults to manager's default_budget.
operation_id: Optional custom operation ID. Auto-generated if not provided.
Yields:
BudgetedOperation context for acquiring connections
"""
op_id = operation_id or f"op-{uuid.uuid4().hex[:12]}"
budget = max_connections or self.default_budget
async with self._lock:
if op_id in self._operations:
raise ValueError(f"Operation {op_id} already exists")
self._operations[op_id] = OperationBudget(op_id, budget)
try:
yield BudgetedOperation(self, op_id)
finally:
async with self._lock:
self._operations.pop(op_id, None)
def _get_budget(self, operation_id: str) -> OperationBudget:
"""Get budget for an operation (internal use)."""
budget = self._operations.get(operation_id)
if not budget:
raise ValueError(f"Operation {operation_id} not found")
return budget
class BudgetedOperation:
"""
A single operation with connection budget.
Provides methods to acquire connections within the budget.
"""
def __init__(self, manager: ConnectionBudgetManager, operation_id: str):
self._manager = manager
self.operation_id = operation_id
@property
def budget(self) -> OperationBudget:
"""Get the budget for this operation."""
return self._manager._get_budget(self.operation_id)
@asynccontextmanager
async def acquire(self, pool: "asyncpg.Pool") -> AsyncIterator["asyncpg.Connection"]:
"""
Acquire a connection within the operation's budget.
Blocks if the operation has reached its connection limit.
Args:
pool: asyncpg connection pool
Yields:
Database connection
"""
budget = self.budget
async with budget.semaphore:
budget.active_count += 1
conn = await pool.acquire()
try:
yield conn
finally:
budget.active_count -= 1
await pool.release(conn)
def wrap_pool(self, pool: "asyncpg.Pool") -> "BudgetedPool":
"""
Wrap a pool with this operation's budget.
The returned BudgetedPool can be passed to functions expecting a pool,
and all acquire() calls will be limited by this operation's budget.
Args:
pool: asyncpg connection pool to wrap
Returns:
BudgetedPool that limits connections to this operation's budget
"""
return BudgetedPool(pool, self)
async def acquire_many(
self,
pool: "asyncpg.Pool",
count: int,
) -> AsyncIterator[list["asyncpg.Connection"]]:
"""
Acquire multiple connections within the budget.
Note: This acquires connections sequentially to respect the budget.
For parallel acquisition, use multiple acquire() calls with asyncio.gather().
Args:
pool: asyncpg connection pool
count: Number of connections to acquire
Yields:
List of database connections
"""
connections = []
try:
for _ in range(count):
conn = await pool.acquire()
connections.append(conn)
yield connections
finally:
for conn in connections:
await pool.release(conn)
# Global default manager instance
_default_manager: ConnectionBudgetManager | None = None
def get_budget_manager(default_budget: int = 4) -> ConnectionBudgetManager:
"""
Get or create the global budget manager.
Args:
default_budget: Default max connections per operation
Returns:
Global ConnectionBudgetManager instance
"""
global _default_manager
if _default_manager is None:
_default_manager = ConnectionBudgetManager(default_budget=default_budget)
return _default_manager
@asynccontextmanager
async def budgeted_operation(
max_connections: int | None = None,
operation_id: str | None = None,
default_budget: int = 4,
) -> AsyncIterator[BudgetedOperation]:
"""
Convenience function to create a budgeted operation.
Args:
max_connections: Max concurrent connections for this operation
operation_id: Optional custom operation ID
default_budget: Default budget if manager not yet created
Yields:
BudgetedOperation context
Example:
async with budgeted_operation(max_connections=2) as op:
async with op.acquire(pool) as conn:
await conn.fetch(...)
"""
manager = get_budget_manager(default_budget)
async with manager.operation(max_connections, operation_id) as op:
yield op
class BudgetedPool:
"""
A pool wrapper that limits concurrent connection acquisitions.
This can be passed to functions expecting a pool, and acquire()
calls will be limited by the budget semaphore.
Usage:
async with budgeted_operation(max_connections=4) as op:
budgeted_pool = op.wrap_pool(pool)
# Pass budgeted_pool to functions that expect a pool
await some_function(budgeted_pool, ...)
"""
def __init__(self, pool: "asyncpg.Pool", operation: BudgetedOperation):
self._pool = pool
self._operation = operation
async def acquire(self) -> "asyncpg.Connection":
"""
Acquire a connection within the budget.
Note: Caller must release the connection when done.
Prefer using as context manager via acquire_with_retry or op.acquire().
"""
budget = self._operation.budget
await budget.semaphore.acquire()
budget.active_count += 1
try:
return await self._pool.acquire()
except Exception:
budget.active_count -= 1
budget.semaphore.release()
raise
async def release(self, conn: "asyncpg.Connection") -> None:
"""Release a connection back to the pool."""
budget = self._operation.budget
try:
await self._pool.release(conn)
finally:
budget.active_count -= 1
budget.semaphore.release()
def __getattr__(self, name):
"""Proxy other attributes to the underlying pool."""
return getattr(self._pool, name)
+16 -4
View File
@@ -1,9 +1,11 @@
"""
Database utility functions for connection management with retry logic.
"""
import asyncio
import logging
from contextlib import asynccontextmanager
import asyncpg
logger = logging.getLogger(__name__)
@@ -54,16 +56,14 @@ async def retry_with_backoff(
except retryable_exceptions as e:
last_exception = e
if attempt < max_retries:
delay = min(base_delay * (2 ** attempt), max_delay)
delay = min(base_delay * (2**attempt), max_delay)
logger.warning(
f"Database operation failed (attempt {attempt + 1}/{max_retries + 1}): {e}. "
f"Retrying in {delay:.1f}s..."
)
await asyncio.sleep(delay)
else:
logger.error(
f"Database operation failed after {max_retries + 1} attempts: {e}"
)
logger.error(f"Database operation failed after {max_retries + 1} attempts: {e}")
raise last_exception
@@ -83,10 +83,22 @@ async def acquire_with_retry(pool: asyncpg.Pool, max_retries: int = DEFAULT_MAX_
Yields:
An asyncpg connection
"""
import time
start = time.time()
async def acquire():
return await pool.acquire()
conn = await retry_with_backoff(acquire, max_retries=max_retries)
acquire_time = time.time() - start
# Log slow connection acquisitions (indicates pool contention)
if acquire_time > 0.05: # 50ms threshold
pool_size = pool.get_size()
pool_free = pool.get_idle_size()
logger.warning(f"[DB POOL] Slow acquire: {acquire_time:.3f}s | size={pool_size}, idle={pool_free}")
try:
yield conn
finally:
+483 -66
View File
@@ -3,40 +3,49 @@ Embeddings abstraction for the memory system.
Provides an interface for generating embeddings with different backends.
IMPORTANT: All embeddings must produce 384-dimensional vectors to match
the database schema (pgvector column defined as vector(384)).
The embedding dimension is auto-detected from the model at initialization.
The database schema is automatically adjusted to match the model's dimension.
Configuration via environment variables:
- HINDSIGHT_API_EMBEDDINGS_PROVIDER: "local" (default) or "tei"
For local provider:
- HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL: Model name (default: BAAI/bge-small-en-v1.5)
For TEI provider:
- HINDSIGHT_API_EMBEDDINGS_TEI_URL: TEI server URL (required)
Configuration via environment variables - see hindsight_api.config for all env var names.
"""
from abc import ABC, abstractmethod
from typing import List, Optional
import logging
import os
from abc import ABC, abstractmethod
import httpx
from ..config import (
DEFAULT_EMBEDDINGS_COHERE_MODEL,
DEFAULT_EMBEDDINGS_LITELLM_MODEL,
DEFAULT_EMBEDDINGS_LOCAL_MODEL,
DEFAULT_EMBEDDINGS_OPENAI_MODEL,
DEFAULT_EMBEDDINGS_PROVIDER,
DEFAULT_LITELLM_API_BASE,
ENV_COHERE_API_KEY,
ENV_EMBEDDINGS_COHERE_BASE_URL,
ENV_EMBEDDINGS_COHERE_MODEL,
ENV_EMBEDDINGS_LITELLM_MODEL,
ENV_EMBEDDINGS_LOCAL_MODEL,
ENV_EMBEDDINGS_OPENAI_API_KEY,
ENV_EMBEDDINGS_OPENAI_BASE_URL,
ENV_EMBEDDINGS_OPENAI_MODEL,
ENV_EMBEDDINGS_PROVIDER,
ENV_EMBEDDINGS_TEI_URL,
ENV_LITELLM_API_BASE,
ENV_LITELLM_API_KEY,
ENV_LLM_API_KEY,
)
logger = logging.getLogger(__name__)
# Fixed embedding dimension required by database schema
EMBEDDING_DIMENSION = 384
# Default model for local embeddings
DEFAULT_EMBEDDINGS_MODEL = "BAAI/bge-small-en-v1.5"
class Embeddings(ABC):
"""
Abstract base class for embedding generation.
All implementations MUST generate 384-dimensional embeddings to match
the database schema.
The embedding dimension is determined by the model and detected at initialization.
The database schema is automatically adjusted to match the model's dimension.
"""
@property
@@ -45,6 +54,12 @@ class Embeddings(ABC):
"""Return a human-readable name for this provider (e.g., 'local', 'tei')."""
pass
@property
@abstractmethod
def dimension(self) -> int:
"""Return the embedding dimension produced by this model."""
pass
@abstractmethod
async def initialize(self) -> None:
"""
@@ -56,15 +71,15 @@ class Embeddings(ABC):
pass
@abstractmethod
def encode(self, texts: List[str]) -> List[List[float]]:
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate 384-dimensional embeddings for a list of texts.
Generate embeddings for a list of texts.
Args:
texts: List of text strings to encode
Returns:
List of 384-dimensional embedding vectors (each is a list of floats)
List of embedding vectors (each is a list of floats)
"""
pass
@@ -74,27 +89,31 @@ class LocalSTEmbeddings(Embeddings):
Local embeddings implementation using SentenceTransformers.
Call initialize() during startup to load the model and avoid cold starts.
Default model is BAAI/bge-small-en-v1.5 which produces 384-dimensional
embeddings matching the database schema.
The embedding dimension is auto-detected from the model.
"""
def __init__(self, model_name: Optional[str] = None):
def __init__(self, model_name: str | None = None):
"""
Initialize local SentenceTransformers embeddings.
Args:
model_name: Name of the SentenceTransformer model to use.
Must produce 384-dimensional embeddings.
Default: BAAI/bge-small-en-v1.5
"""
self.model_name = model_name or DEFAULT_EMBEDDINGS_MODEL
self.model_name = model_name or DEFAULT_EMBEDDINGS_LOCAL_MODEL
self._model = None
self._dimension: int | None = None
@property
def provider_name(self) -> str:
return "local"
@property
def dimension(self) -> int:
if self._dimension is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
return self._dimension
async def initialize(self) -> None:
"""Load the embedding model."""
if self._model is not None:
@@ -116,26 +135,18 @@ class LocalSTEmbeddings(Embeddings):
model_kwargs={"low_cpu_mem_usage": False, "device_map": None},
)
# Validate dimension matches database schema
model_dim = self._model.get_sentence_embedding_dimension()
if model_dim != EMBEDDING_DIMENSION:
raise ValueError(
f"Model {self.model_name} produces {model_dim}-dimensional embeddings, "
f"but database schema requires {EMBEDDING_DIMENSION} dimensions. "
f"Use a model that produces {EMBEDDING_DIMENSION}-dimensional embeddings."
)
self._dimension = self._model.get_sentence_embedding_dimension()
logger.info(f"Embeddings: local provider initialized (dim: {self._dimension})")
logger.info(f"Embeddings: local provider initialized (dim: {model_dim})")
def encode(self, texts: List[str]) -> List[List[float]]:
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate 384-dimensional embeddings for a list of texts.
Generate embeddings for a list of texts.
Args:
texts: List of text strings to encode
Returns:
List of 384-dimensional embedding vectors
List of embedding vectors
"""
if self._model is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
@@ -150,7 +161,7 @@ class RemoteTEIEmbeddings(Embeddings):
TEI provides a high-performance inference server for embedding models.
See: https://github.com/huggingface/text-embeddings-inference
The server should be running a model that produces 384-dimensional embeddings.
The embedding dimension is auto-detected from the server at initialization.
"""
def __init__(
@@ -176,16 +187,24 @@ class RemoteTEIEmbeddings(Embeddings):
self.batch_size = batch_size
self.max_retries = max_retries
self.retry_delay = retry_delay
self._client: Optional[httpx.Client] = None
self._model_id: Optional[str] = None
self._client: httpx.Client | None = None
self._model_id: str | None = None
self._dimension: int | None = None
@property
def provider_name(self) -> str:
return "tei"
@property
def dimension(self) -> int:
if self._dimension is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
return self._dimension
def _request_with_retry(self, method: str, url: str, **kwargs) -> httpx.Response:
"""Make an HTTP request with automatic retries on transient errors."""
import time
last_error = None
delay = self.retry_delay
@@ -200,14 +219,18 @@ class RemoteTEIEmbeddings(Embeddings):
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
last_error = e
if attempt < self.max_retries:
logger.warning(f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
logger.warning(
f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
)
time.sleep(delay)
delay *= 2 # Exponential backoff
except httpx.HTTPStatusError as e:
# Retry on 5xx server errors
if e.response.status_code >= 500 and attempt < self.max_retries:
last_error = e
logger.warning(f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
logger.warning(
f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
)
time.sleep(delay)
delay *= 2
else:
@@ -228,11 +251,28 @@ class RemoteTEIEmbeddings(Embeddings):
response = self._request_with_retry("GET", f"{self.base_url}/info")
info = response.json()
self._model_id = info.get("model_id", "unknown")
logger.info(f"Embeddings: TEI provider initialized (model: {self._model_id})")
# Get dimension from server info or by doing a test embedding
if "max_input_length" in info and "model_dtype" in info:
# Try to get dimension from info endpoint (some TEI versions expose it)
# If not available, do a test embedding
pass
# Do a test embedding to detect dimension
test_response = self._request_with_retry(
"POST",
f"{self.base_url}/embed",
json={"inputs": ["test"]},
)
test_embeddings = test_response.json()
if test_embeddings and len(test_embeddings) > 0:
self._dimension = len(test_embeddings[0])
logger.info(f"Embeddings: TEI provider initialized (model: {self._model_id}, dim: {self._dimension})")
except httpx.HTTPError as e:
raise RuntimeError(f"Failed to connect to TEI server at {self.base_url}: {e}")
def encode(self, texts: List[str]) -> List[List[float]]:
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings using the remote TEI server.
@@ -252,7 +292,7 @@ class RemoteTEIEmbeddings(Embeddings):
# Process in batches
for i in range(0, len(texts), self.batch_size):
batch = texts[i:i + self.batch_size]
batch = texts[i : i + self.batch_size]
try:
response = self._request_with_retry(
@@ -268,36 +308,413 @@ class RemoteTEIEmbeddings(Embeddings):
return all_embeddings
class OpenAIEmbeddings(Embeddings):
"""
OpenAI embeddings implementation using the OpenAI API.
Supports text-embedding-3-small (1536 dims), text-embedding-3-large (3072 dims),
and text-embedding-ada-002 (1536 dims, legacy).
The embedding dimension is auto-detected from the model at initialization.
"""
# Known dimensions for OpenAI embedding models
MODEL_DIMENSIONS = {
"text-embedding-3-small": 1536,
"text-embedding-3-large": 3072,
"text-embedding-ada-002": 1536,
}
def __init__(
self,
api_key: str,
model: str = DEFAULT_EMBEDDINGS_OPENAI_MODEL,
base_url: str | None = None,
batch_size: int = 100,
max_retries: int = 3,
):
"""
Initialize OpenAI embeddings client.
Args:
api_key: OpenAI API key
model: OpenAI embedding model name (default: text-embedding-3-small)
base_url: Custom base URL for OpenAI-compatible API (e.g., Azure OpenAI endpoint)
batch_size: Maximum batch size for embedding requests (default: 100)
max_retries: Maximum number of retries for failed requests (default: 3)
"""
self.api_key = api_key
self.model = model
self.base_url = base_url
self.batch_size = batch_size
self.max_retries = max_retries
self._client = None
self._dimension: int | None = None
@property
def provider_name(self) -> str:
return "openai"
@property
def dimension(self) -> int:
if self._dimension is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
return self._dimension
async def initialize(self) -> None:
"""Initialize the OpenAI client and detect dimension."""
if self._client is not None:
return
try:
from openai import OpenAI
except ImportError:
raise ImportError("openai is required for OpenAIEmbeddings. Install it with: pip install openai")
base_url_msg = f" at {self.base_url}" if self.base_url else ""
logger.info(f"Embeddings: initializing OpenAI provider with model {self.model}{base_url_msg}")
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
client_kwargs = {"api_key": self.api_key, "max_retries": self.max_retries}
if self.base_url:
client_kwargs["base_url"] = self.base_url
self._client = OpenAI(**client_kwargs)
# Try to get dimension from known models, otherwise do a test embedding
if self.model in self.MODEL_DIMENSIONS:
self._dimension = self.MODEL_DIMENSIONS[self.model]
else:
# Do a test embedding to detect dimension
response = self._client.embeddings.create(
model=self.model,
input=["test"],
)
if response.data:
self._dimension = len(response.data[0].embedding)
logger.info(f"Embeddings: OpenAI provider initialized (model: {self.model}, dim: {self._dimension})")
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings using the OpenAI API.
Args:
texts: List of text strings to encode
Returns:
List of embedding vectors
"""
if self._client is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
if not texts:
return []
all_embeddings = []
# Process in batches
for i in range(0, len(texts), self.batch_size):
batch = texts[i : i + self.batch_size]
response = self._client.embeddings.create(
model=self.model,
input=batch,
)
# Sort by index to ensure correct order
batch_embeddings = sorted(response.data, key=lambda x: x.index)
all_embeddings.extend([e.embedding for e in batch_embeddings])
return all_embeddings
class CohereEmbeddings(Embeddings):
"""
Cohere embeddings implementation using the Cohere API.
Supports embed-english-v3.0 (1024 dims) and embed-multilingual-v3.0 (1024 dims).
The embedding dimension is auto-detected from the model at initialization.
"""
# Known dimensions for Cohere embedding models
MODEL_DIMENSIONS = {
"embed-english-v3.0": 1024,
"embed-multilingual-v3.0": 1024,
"embed-english-light-v3.0": 384,
"embed-multilingual-light-v3.0": 384,
"embed-english-v2.0": 4096,
"embed-multilingual-v2.0": 768,
}
def __init__(
self,
api_key: str,
model: str = DEFAULT_EMBEDDINGS_COHERE_MODEL,
base_url: str | None = None,
batch_size: int = 96,
timeout: float = 60.0,
input_type: str = "search_document",
):
"""
Initialize Cohere embeddings client.
Args:
api_key: Cohere API key
model: Cohere embedding model name (default: embed-english-v3.0)
base_url: Custom base URL for Cohere-compatible API (e.g., Azure-hosted endpoint)
batch_size: Maximum batch size for embedding requests (default: 96, Cohere's limit)
timeout: Request timeout in seconds (default: 60.0)
input_type: Input type for embeddings (default: search_document).
Options: search_document, search_query, classification, clustering
"""
self.api_key = api_key
self.model = model
self.base_url = base_url
self.batch_size = batch_size
self.timeout = timeout
self.input_type = input_type
self._client = None
self._dimension: int | None = None
@property
def provider_name(self) -> str:
return "cohere"
@property
def dimension(self) -> int:
if self._dimension is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
return self._dimension
async def initialize(self) -> None:
"""Initialize the Cohere client and detect dimension."""
if self._client is not None:
return
try:
import cohere
except ImportError:
raise ImportError("cohere is required for CohereEmbeddings. Install it with: pip install cohere")
base_url_msg = f" at {self.base_url}" if self.base_url else ""
logger.info(f"Embeddings: initializing Cohere provider with model {self.model}{base_url_msg}")
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
client_kwargs = {"api_key": self.api_key, "timeout": self.timeout}
if self.base_url:
client_kwargs["base_url"] = self.base_url
self._client = cohere.Client(**client_kwargs)
# Try to get dimension from known models, otherwise do a test embedding
if self.model in self.MODEL_DIMENSIONS:
self._dimension = self.MODEL_DIMENSIONS[self.model]
else:
# Do a test embedding to detect dimension
response = self._client.embed(
texts=["test"],
model=self.model,
input_type=self.input_type,
)
if response.embeddings:
self._dimension = len(response.embeddings[0])
logger.info(f"Embeddings: Cohere provider initialized (model: {self.model}, dim: {self._dimension})")
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings using the Cohere API.
Args:
texts: List of text strings to encode
Returns:
List of embedding vectors
"""
if self._client is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
if not texts:
return []
all_embeddings = []
# Process in batches
for i in range(0, len(texts), self.batch_size):
batch = texts[i : i + self.batch_size]
response = self._client.embed(
texts=batch,
model=self.model,
input_type=self.input_type,
)
all_embeddings.extend(response.embeddings)
return all_embeddings
class LiteLLMEmbeddings(Embeddings):
"""
LiteLLM embeddings implementation using LiteLLM proxy's /embeddings endpoint.
LiteLLM provides a unified interface for multiple embedding providers.
The proxy exposes an OpenAI-compatible /embeddings endpoint.
See: https://docs.litellm.ai/docs/embedding/supported_embedding
Supported providers via LiteLLM:
- OpenAI (text-embedding-3-small, text-embedding-ada-002, etc.)
- Cohere (embed-english-v3.0, etc.) - prefix with cohere/
- Vertex AI (textembedding-gecko, etc.) - prefix with vertex_ai/
- HuggingFace, Mistral, Voyage AI, etc.
The embedding dimension is auto-detected from the model at initialization.
"""
def __init__(
self,
api_base: str = DEFAULT_LITELLM_API_BASE,
api_key: str | None = None,
model: str = DEFAULT_EMBEDDINGS_LITELLM_MODEL,
batch_size: int = 100,
timeout: float = 60.0,
):
"""
Initialize LiteLLM embeddings client.
Args:
api_base: Base URL of the LiteLLM proxy (default: http://localhost:4000)
api_key: API key for the LiteLLM proxy (optional, depends on proxy config)
model: Embedding model name (default: text-embedding-3-small)
Use provider prefix for non-OpenAI models (e.g., cohere/embed-english-v3.0)
batch_size: Maximum batch size for embedding requests (default: 100)
timeout: Request timeout in seconds (default: 60.0)
"""
self.api_base = api_base.rstrip("/")
self.api_key = api_key
self.model = model
self.batch_size = batch_size
self.timeout = timeout
self._client: httpx.Client | None = None
self._dimension: int | None = None
@property
def provider_name(self) -> str:
return "litellm"
@property
def dimension(self) -> int:
if self._dimension is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
return self._dimension
async def initialize(self) -> None:
"""Initialize the HTTP client and detect embedding dimension."""
if self._client is not None:
return
logger.info(f"Embeddings: initializing LiteLLM provider at {self.api_base} with model {self.model}")
headers = {"Content-Type": "application/json"}
if self.api_key:
headers["Authorization"] = f"Bearer {self.api_key}"
self._client = httpx.Client(timeout=self.timeout, headers=headers)
# Do a test embedding to detect dimension
try:
response = self._client.post(
f"{self.api_base}/embeddings",
json={"model": self.model, "input": ["test"]},
)
response.raise_for_status()
result = response.json()
if result.get("data") and len(result["data"]) > 0:
self._dimension = len(result["data"][0]["embedding"])
logger.info(f"Embeddings: LiteLLM provider initialized (model: {self.model}, dim: {self._dimension})")
except httpx.HTTPError as e:
raise RuntimeError(f"Failed to connect to LiteLLM proxy at {self.api_base}: {e}")
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings using the LiteLLM proxy.
Args:
texts: List of text strings to encode
Returns:
List of embedding vectors
"""
if self._client is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
if not texts:
return []
all_embeddings = []
# Process in batches
for i in range(0, len(texts), self.batch_size):
batch = texts[i : i + self.batch_size]
response = self._client.post(
f"{self.api_base}/embeddings",
json={"model": self.model, "input": batch},
)
response.raise_for_status()
result = response.json()
# Sort by index to ensure correct order
batch_embeddings = sorted(result["data"], key=lambda x: x["index"])
all_embeddings.extend([e["embedding"] for e in batch_embeddings])
return all_embeddings
def create_embeddings_from_env() -> Embeddings:
"""
Create an Embeddings instance based on environment variables.
Environment variables:
- HINDSIGHT_API_EMBEDDINGS_PROVIDER: "local" (default) or "tei"
For local provider:
- HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL: Model name (default: BAAI/bge-small-en-v1.5)
For TEI provider:
- HINDSIGHT_API_EMBEDDINGS_TEI_URL: TEI server URL (required)
See hindsight_api.config for environment variable names and defaults.
Returns:
Configured Embeddings instance
"""
provider = os.environ.get("HINDSIGHT_API_EMBEDDINGS_PROVIDER", "local").lower()
provider = os.environ.get(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER).lower()
if provider == "tei":
url = os.environ.get("HINDSIGHT_API_EMBEDDINGS_TEI_URL")
url = os.environ.get(ENV_EMBEDDINGS_TEI_URL)
if not url:
raise ValueError(
"HINDSIGHT_API_EMBEDDINGS_TEI_URL is required when HINDSIGHT_API_EMBEDDINGS_PROVIDER is 'tei'"
)
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("HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL")
model_name = model or DEFAULT_EMBEDDINGS_MODEL
model = os.environ.get(ENV_EMBEDDINGS_LOCAL_MODEL)
model_name = model or DEFAULT_EMBEDDINGS_LOCAL_MODEL
return LocalSTEmbeddings(model_name=model_name)
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)
if not api_key:
raise ValueError(
f"{ENV_EMBEDDINGS_OPENAI_API_KEY} or {ENV_LLM_API_KEY} is required "
f"when {ENV_EMBEDDINGS_PROVIDER} is 'openai'"
)
model = os.environ.get(ENV_EMBEDDINGS_OPENAI_MODEL, DEFAULT_EMBEDDINGS_OPENAI_MODEL)
base_url = os.environ.get(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None
return OpenAIEmbeddings(api_key=api_key, model=model, base_url=base_url)
elif provider == "cohere":
api_key = os.environ.get(ENV_COHERE_API_KEY)
if not api_key:
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_EMBEDDINGS_PROVIDER} is 'cohere'")
model = os.environ.get(ENV_EMBEDDINGS_COHERE_MODEL, DEFAULT_EMBEDDINGS_COHERE_MODEL)
base_url = os.environ.get(ENV_EMBEDDINGS_COHERE_BASE_URL) or None
return CohereEmbeddings(api_key=api_key, model=model, base_url=base_url)
elif provider == "litellm":
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
api_key = os.environ.get(ENV_LITELLM_API_KEY)
model = os.environ.get(ENV_EMBEDDINGS_LITELLM_MODEL, DEFAULT_EMBEDDINGS_LITELLM_MODEL)
return LiteLLMEmbeddings(api_base=api_base, api_key=api_key, model=model)
else:
raise ValueError(
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei'"
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei', 'openai', 'cohere', 'litellm'"
)
@@ -4,12 +4,14 @@ Entity extraction and resolution for memory system.
Uses spaCy for entity extraction and implements resolution logic
to disambiguate entities across memory units.
"""
import asyncpg
from typing import List, Dict, Optional, Set, Any
from difflib import SequenceMatcher
from datetime import datetime, timezone
from .db_utils import acquire_with_retry
from datetime import UTC, datetime
from difflib import SequenceMatcher
import asyncpg
from .db_utils import acquire_with_retry
from .memory_engine import fq_table
# Load spaCy model (singleton)
_nlp = None
@@ -32,11 +34,11 @@ class EntityResolver:
async def resolve_entities_batch(
self,
bank_id: str,
entities_data: List[Dict],
entities_data: list[dict],
context: str,
unit_event_date,
conn=None,
) -> List[str]:
) -> list[str]:
"""
Resolve multiple entities in batch (MUCH faster than sequential).
@@ -62,36 +64,38 @@ class EntityResolver:
else:
return await self._resolve_entities_batch_impl(conn, bank_id, entities_data, context, unit_event_date)
async def _resolve_entities_batch_impl(self, conn, bank_id: str, entities_data: List[Dict], context: str, unit_event_date) -> List[str]:
async def _resolve_entities_batch_impl(
self, conn, bank_id: str, entities_data: list[dict], context: str, unit_event_date
) -> list[str]:
# Query ALL candidates for this bank
all_entities = await conn.fetch(
"""
f"""
SELECT canonical_name, id, metadata, last_seen, mention_count
FROM entities
FROM {fq_table("entities")}
WHERE bank_id = $1
""",
bank_id
bank_id,
)
# Build entity ID to name mapping for co-occurrence lookups
entity_id_to_name = {row['id']: row['canonical_name'].lower() for row in all_entities}
entity_id_to_name = {row["id"]: row["canonical_name"].lower() for row in all_entities}
# Query ALL co-occurrences for this bank's entities in one query
# This builds a map of entity_id -> set of co-occurring entity names
all_cooccurrences = await conn.fetch(
"""
f"""
SELECT ec.entity_id_1, ec.entity_id_2, ec.cooccurrence_count
FROM entity_cooccurrences ec
WHERE ec.entity_id_1 IN (SELECT id FROM entities WHERE bank_id = $1)
OR ec.entity_id_2 IN (SELECT id FROM entities WHERE bank_id = $1)
FROM {fq_table("entity_cooccurrences")} ec
WHERE ec.entity_id_1 IN (SELECT id FROM {fq_table("entities")} WHERE bank_id = $1)
OR ec.entity_id_2 IN (SELECT id FROM {fq_table("entities")} WHERE bank_id = $1)
""",
bank_id
bank_id,
)
# Build co-occurrence map: entity_id -> set of co-occurring entity names (lowercase)
cooccurrence_map: Dict[str, Set[str]] = {}
cooccurrence_map: dict[str, set[str]] = {}
for row in all_cooccurrences:
eid1, eid2 = row['entity_id_1'], row['entity_id_2']
eid1, eid2 = row["entity_id_1"], row["entity_id_2"]
# Add both directions
if eid1 not in cooccurrence_map:
cooccurrence_map[eid1] = set()
@@ -105,22 +109,24 @@ class EntityResolver:
# Build candidate map for each entity text
all_candidates = {} # Maps entity_text -> list of candidates
entity_texts = list(set(e['text'] for e in entities_data))
entity_texts = list(set(e["text"] for e in entities_data))
for entity_text in entity_texts:
matching = []
entity_text_lower = entity_text.lower()
for row in all_entities:
canonical_name = row['canonical_name']
ent_id = row['id']
metadata = row['metadata']
last_seen = row['last_seen']
mention_count = row['mention_count']
canonical_name = row["canonical_name"]
ent_id = row["id"]
metadata = row["metadata"]
last_seen = row["last_seen"]
mention_count = row["mention_count"]
canonical_lower = canonical_name.lower()
# Match if exact or substring match
if (entity_text_lower == canonical_lower or
entity_text_lower in canonical_lower or
canonical_lower in entity_text_lower):
if (
entity_text_lower == canonical_lower
or entity_text_lower in canonical_lower
or canonical_lower in entity_text_lower
):
matching.append((ent_id, canonical_name, metadata, last_seen, mention_count))
all_candidates[entity_text] = matching
@@ -130,10 +136,10 @@ class EntityResolver:
entities_to_create = [] # (idx, entity_data, event_date)
for idx, entity_data in enumerate(entities_data):
entity_text = entity_data['text']
nearby_entities = entity_data.get('nearby_entities', [])
entity_text = entity_data["text"]
nearby_entities = entity_data.get("nearby_entities", [])
# Use per-entity date if available, otherwise fall back to batch-level date
entity_event_date = entity_data.get('event_date', unit_event_date)
entity_event_date = entity_data.get("event_date", unit_event_date)
candidates = all_candidates.get(entity_text, [])
@@ -146,17 +152,13 @@ class EntityResolver:
best_candidate = None
best_score = 0.0
nearby_entity_set = {e['text'].lower() for e in nearby_entities if e['text'] != entity_text}
nearby_entity_set = {e["text"].lower() for e in nearby_entities if e["text"] != entity_text}
for candidate_id, canonical_name, metadata, last_seen, mention_count in candidates:
score = 0.0
# 1. Name similarity (0-0.5)
name_similarity = SequenceMatcher(
None,
entity_text.lower(),
canonical_name.lower()
).ratio()
name_similarity = SequenceMatcher(None, entity_text.lower(), canonical_name.lower()).ratio()
score += name_similarity * 0.5
# 2. Co-occurring entities (0-0.3)
@@ -169,8 +171,10 @@ class EntityResolver:
# 3. Temporal proximity (0-0.2)
if last_seen and entity_event_date:
# Normalize timezone awareness for comparison
event_date_utc = entity_event_date if entity_event_date.tzinfo else entity_event_date.replace(tzinfo=timezone.utc)
last_seen_utc = last_seen if last_seen.tzinfo else last_seen.replace(tzinfo=timezone.utc)
event_date_utc = (
entity_event_date if entity_event_date.tzinfo else entity_event_date.replace(tzinfo=UTC)
)
last_seen_utc = last_seen if last_seen.tzinfo else last_seen.replace(tzinfo=UTC)
days_diff = abs((event_date_utc - last_seen_utc).total_seconds() / 86400)
if days_diff < 7:
temporal_score = max(0, 1.0 - (days_diff / 7))
@@ -192,23 +196,23 @@ class EntityResolver:
# Batch update existing entities
if entities_to_update:
await conn.executemany(
"""
UPDATE entities SET
f"""
UPDATE {fq_table("entities")} SET
mention_count = mention_count + 1,
last_seen = $2
WHERE id = $1::uuid
""",
entities_to_update
entities_to_update,
)
# Batch create new entities using COPY + INSERT for maximum speed
# This handles duplicates via ON CONFLICT and returns all IDs
if entities_to_create:
# Group entities by canonical name (lowercase) to handle duplicates within batch
# For duplicates, we only insert once and reuse the ID
# For duplicates, we only insert once and reuse the ID, but track the count
unique_entities = {} # lowercase_name -> (entity_data, event_date, [indices])
for idx, entity_data, event_date in entities_to_create:
name_lower = entity_data['text'].lower()
name_lower = entity_data["text"].lower()
if name_lower not in unique_entities:
unique_entities[name_lower] = (entity_data, event_date, [idx])
else:
@@ -219,34 +223,37 @@ class EntityResolver:
# Use a single query with unnest for speed
entity_names = []
entity_dates = []
entity_counts = [] # Track how many times each entity appears in this batch
indices_map = [] # Maps result index -> list of original indices
for name_lower, (entity_data, event_date, indices) in unique_entities.items():
entity_names.append(entity_data['text'])
entity_names.append(entity_data["text"])
entity_dates.append(event_date)
entity_counts.append(len(indices)) # Count of occurrences in this batch
indices_map.append(indices)
# Batch INSERT ... ON CONFLICT with RETURNING
# This is much faster than individual inserts
# Uses the batch count for mention_count instead of always 1
rows = await conn.fetch(
"""
INSERT INTO entities (bank_id, canonical_name, first_seen, last_seen, mention_count)
SELECT $1, name, event_date, event_date, 1
FROM unnest($2::text[], $3::timestamptz[]) AS t(name, event_date)
f"""
INSERT INTO {fq_table("entities")} (bank_id, canonical_name, first_seen, last_seen, mention_count)
SELECT $1, name, event_date, event_date, cnt
FROM unnest($2::text[], $3::timestamptz[], $4::int[]) AS t(name, event_date, cnt)
ON CONFLICT (bank_id, LOWER(canonical_name))
DO UPDATE SET
mention_count = entities.mention_count + 1,
mention_count = {fq_table("entities")}.mention_count + EXCLUDED.mention_count,
last_seen = EXCLUDED.last_seen
RETURNING id
""",
bank_id,
entity_names,
entity_dates
entity_dates,
entity_counts,
)
# Map returned IDs back to original indices
for result_idx, row in enumerate(rows):
entity_id = row['id']
entity_id = row["id"]
for original_idx in indices_map[result_idx]:
entity_ids[original_idx] = entity_id
@@ -257,7 +264,7 @@ class EntityResolver:
bank_id: str,
entity_text: str,
context: str,
nearby_entities: List[Dict],
nearby_entities: list[dict],
unit_event_date,
) -> str:
"""
@@ -276,9 +283,9 @@ class EntityResolver:
async with acquire_with_retry(self.pool) as conn:
# Find candidate entities with similar name
candidates = await conn.fetch(
"""
f"""
SELECT id, canonical_name, metadata, last_seen
FROM entities
FROM {fq_table("entities")}
WHERE bank_id = $1
AND (
canonical_name ILIKE $2
@@ -287,14 +294,14 @@ class EntityResolver:
)
ORDER BY mention_count DESC
""",
bank_id, entity_text, f"%{entity_text}%"
bank_id,
entity_text,
f"%{entity_text}%",
)
if not candidates:
# New entity - create it
return await self._create_entity(
conn, bank_id, entity_text, unit_event_date
)
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
# Score candidates based on:
# 1. Name similarity
@@ -306,31 +313,27 @@ class EntityResolver:
best_score = 0.0
best_name_similarity = 0.0
nearby_entity_set = {e['text'].lower() for e in nearby_entities if e['text'] != entity_text}
nearby_entity_set = {e["text"].lower() for e in nearby_entities if e["text"] != entity_text}
for row in candidates:
candidate_id = row['id']
canonical_name = row['canonical_name']
metadata = row['metadata']
last_seen = row['last_seen']
candidate_id = row["id"]
canonical_name = row["canonical_name"]
metadata = row["metadata"]
last_seen = row["last_seen"]
score = 0.0
# 1. Name similarity (0-1)
name_similarity = SequenceMatcher(
None,
entity_text.lower(),
canonical_name.lower()
).ratio()
name_similarity = SequenceMatcher(None, entity_text.lower(), canonical_name.lower()).ratio()
score += name_similarity * 0.5
# 2. Co-occurring entities (0-0.5)
# Get entities that co-occurred with this candidate before
# Use the materialized co-occurrence cache for fast lookup
co_entity_rows = await conn.fetch(
"""
f"""
SELECT e.canonical_name, ec.cooccurrence_count
FROM entity_cooccurrences ec
JOIN entities e ON (
FROM {fq_table("entity_cooccurrences")} ec
JOIN {fq_table("entities")} e ON (
CASE
WHEN ec.entity_id_1 = $1 THEN ec.entity_id_2
WHEN ec.entity_id_2 = $1 THEN ec.entity_id_1
@@ -338,9 +341,9 @@ class EntityResolver:
)
WHERE ec.entity_id_1 = $1 OR ec.entity_id_2 = $1
""",
candidate_id
candidate_id,
)
co_entities = {r['canonical_name'].lower() for r in co_entity_rows}
co_entities = {r["canonical_name"].lower() for r in co_entity_rows}
# Check overlap with nearby entities
overlap = len(nearby_entity_set & co_entities)
@@ -366,20 +369,19 @@ class EntityResolver:
if best_score > threshold:
# Update entity
await conn.execute(
"""
UPDATE entities
f"""
UPDATE {fq_table("entities")}
SET mention_count = mention_count + 1,
last_seen = $1
WHERE id = $2
""",
unit_event_date, best_candidate
unit_event_date,
best_candidate,
)
return best_candidate
else:
# Not confident - create new entity
return await self._create_entity(
conn, bank_id, entity_text, unit_event_date
)
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
async def _create_entity(
self,
@@ -404,16 +406,19 @@ class EntityResolver:
Entity ID
"""
entity_id = await conn.fetchval(
"""
INSERT INTO entities (bank_id, canonical_name, first_seen, last_seen, mention_count)
f"""
INSERT INTO {fq_table("entities")} (bank_id, canonical_name, first_seen, last_seen, mention_count)
VALUES ($1, $2, $3, $4, 1)
ON CONFLICT (bank_id, LOWER(canonical_name))
DO UPDATE SET
mention_count = entities.mention_count + 1,
mention_count = {fq_table("entities")}.mention_count + 1,
last_seen = EXCLUDED.last_seen
RETURNING id
""",
bank_id, entity_text, event_date, event_date
bank_id,
entity_text,
event_date,
event_date,
)
return entity_id
@@ -429,25 +434,27 @@ class EntityResolver:
async with acquire_with_retry(self.pool) as conn:
# Insert unit-entity link
await conn.execute(
"""
INSERT INTO unit_entities (unit_id, entity_id)
f"""
INSERT INTO {fq_table("unit_entities")} (unit_id, entity_id)
VALUES ($1, $2)
ON CONFLICT DO NOTHING
""",
unit_id, entity_id
unit_id,
entity_id,
)
# Update co-occurrence cache: find other entities in this unit
rows = await conn.fetch(
"""
f"""
SELECT entity_id
FROM unit_entities
FROM {fq_table("unit_entities")}
WHERE unit_id = $1 AND entity_id != $2
""",
unit_id, entity_id
unit_id,
entity_id,
)
other_entities = [row['entity_id'] for row in rows]
other_entities = [row["entity_id"] for row in rows]
# Update co-occurrences for each pair
for other_entity_id in other_entities:
@@ -469,18 +476,19 @@ class EntityResolver:
entity_id_1, entity_id_2 = entity_id_2, entity_id_1
await conn.execute(
"""
INSERT INTO entity_cooccurrences (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
f"""
INSERT INTO {fq_table("entity_cooccurrences")} (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
VALUES ($1, $2, 1, NOW())
ON CONFLICT (entity_id_1, entity_id_2)
DO UPDATE SET
cooccurrence_count = entity_cooccurrences.cooccurrence_count + 1,
cooccurrence_count = {fq_table("entity_cooccurrences")}.cooccurrence_count + 1,
last_cooccurred = NOW()
""",
entity_id_1, entity_id_2
entity_id_1,
entity_id_2,
)
async def link_units_to_entities_batch(self, unit_entity_pairs: List[tuple[str, str]], conn=None):
async def link_units_to_entities_batch(self, unit_entity_pairs: list[tuple[str, str]], conn=None):
"""
Link multiple memory units to entities in batch (MUCH faster than sequential).
@@ -499,15 +507,15 @@ class EntityResolver:
else:
return await self._link_units_to_entities_batch_impl(conn, unit_entity_pairs)
async def _link_units_to_entities_batch_impl(self, conn, unit_entity_pairs: List[tuple[str, str]]):
async def _link_units_to_entities_batch_impl(self, conn, unit_entity_pairs: list[tuple[str, str]]):
# Batch insert all unit-entity links
await conn.executemany(
"""
INSERT INTO unit_entities (unit_id, entity_id)
f"""
INSERT INTO {fq_table("unit_entities")} (unit_id, entity_id)
VALUES ($1, $2)
ON CONFLICT DO NOTHING
""",
unit_entity_pairs
unit_entity_pairs,
)
# Build map of unit -> entities for co-occurrence calculation
@@ -524,7 +532,7 @@ class EntityResolver:
entity_list = list(entity_ids) # Convert set to list for iteration
# For each pair of entities in this unit, create co-occurrence
for i, entity_id_1 in enumerate(entity_list):
for entity_id_2 in entity_list[i+1:]:
for entity_id_2 in entity_list[i + 1 :]:
# Skip if same entity (shouldn't happen with set, but be safe)
if entity_id_1 == entity_id_2:
continue
@@ -535,20 +543,20 @@ class EntityResolver:
# Batch update co-occurrences
if cooccurrence_pairs:
now = datetime.now(timezone.utc)
now = datetime.now(UTC)
await conn.executemany(
"""
INSERT INTO entity_cooccurrences (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
f"""
INSERT INTO {fq_table("entity_cooccurrences")} (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
VALUES ($1, $2, $3, $4)
ON CONFLICT (entity_id_1, entity_id_2)
DO UPDATE SET
cooccurrence_count = entity_cooccurrences.cooccurrence_count + 1,
cooccurrence_count = {fq_table("entity_cooccurrences")}.cooccurrence_count + 1,
last_cooccurred = EXCLUDED.last_cooccurred
""",
[(e1, e2, 1, now) for e1, e2 in cooccurrence_pairs]
[(e1, e2, 1, now) for e1, e2 in cooccurrence_pairs],
)
async def get_units_by_entity(self, entity_id: str, limit: int = 100) -> List[str]:
async def get_units_by_entity(self, entity_id: str, limit: int = 100) -> list[str]:
"""
Get all units that mention an entity.
@@ -561,22 +569,23 @@ class EntityResolver:
"""
async with acquire_with_retry(self.pool) as conn:
rows = await conn.fetch(
"""
f"""
SELECT unit_id
FROM unit_entities
FROM {fq_table("unit_entities")}
WHERE entity_id = $1
ORDER BY unit_id
LIMIT $2
""",
entity_id, limit
entity_id,
limit,
)
return [row['unit_id'] for row in rows]
return [row["unit_id"] for row in rows]
async def get_entity_by_text(
self,
bank_id: str,
entity_text: str,
) -> Optional[str]:
) -> str | None:
"""
Find an entity by text (for query resolution).
@@ -589,14 +598,15 @@ class EntityResolver:
"""
async with acquire_with_retry(self.pool) as conn:
row = await conn.fetchrow(
"""
SELECT id FROM entities
f"""
SELECT id FROM {fq_table("entities")}
WHERE bank_id = $1
AND canonical_name ILIKE $2
ORDER BY mention_count DESC
LIMIT 1
""",
bank_id, entity_text
bank_id,
entity_text,
)
return row['id'] if row else None
return row["id"] if row else None
@@ -0,0 +1,619 @@
"""Abstract interface for MemoryEngine public methods.
This module defines the public API that HTTP endpoints and extensions should use
to interact with the memory system. All methods require a RequestContext for
authentication when a TenantExtension is configured.
"""
from abc import ABC, abstractmethod
from datetime import datetime
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from hindsight_api.engine.memory_engine import Budget
from hindsight_api.engine.response_models import RecallResult, ReflectResult
from hindsight_api.models import RequestContext
class MemoryEngineInterface(ABC):
"""
Abstract interface for the Memory Engine.
This defines the public API that should be used by HTTP endpoints and extensions.
All methods require a RequestContext for authentication.
"""
# =========================================================================
# Health & Status
# =========================================================================
@abstractmethod
async def health_check(self) -> dict:
"""
Check the health of the memory system.
Returns:
Dict with 'status' key ('healthy' or 'unhealthy') and additional info.
"""
...
# =========================================================================
# Core Memory Operations
# =========================================================================
@abstractmethod
async def retain_batch_async(
self,
bank_id: str,
contents: list[dict[str, Any]],
*,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Retain a batch of memory items.
Args:
bank_id: The memory bank ID.
contents: List of content dicts with 'content', optional 'event_date',
'context', 'metadata', 'document_id'.
request_context: Request context for authentication.
Returns:
Dict with processing results.
"""
...
@abstractmethod
async def recall_async(
self,
bank_id: str,
query: str,
*,
budget: "Budget | None" = None,
max_tokens: int = 4096,
enable_trace: bool = False,
fact_type: list[str] | None = None,
question_date: datetime | None = None,
include_entities: bool = False,
max_entity_tokens: int = 500,
include_chunks: bool = False,
max_chunk_tokens: int = 8192,
request_context: "RequestContext",
) -> "RecallResult":
"""
Recall memories relevant to a query.
Args:
bank_id: The memory bank ID.
query: The search query.
budget: Search budget (LOW, MID, HIGH).
max_tokens: Maximum tokens in response.
enable_trace: Include trace information.
fact_type: Filter by fact types.
question_date: Context date for temporal relevance.
include_entities: Include entity observations.
max_entity_tokens: Max tokens for entity observations.
include_chunks: Include raw chunks.
max_chunk_tokens: Max tokens for chunks.
request_context: Request context for authentication.
Returns:
RecallResult with matching memories.
"""
...
@abstractmethod
async def reflect_async(
self,
bank_id: str,
query: str,
*,
budget: "Budget | None" = None,
context: str | None = None,
max_tokens: int = 4096,
response_schema: dict | None = None,
request_context: "RequestContext",
) -> "ReflectResult":
"""
Reflect on a query and generate a thoughtful response.
Args:
bank_id: The memory bank ID.
query: The question to reflect on.
budget: Search budget for retrieving context.
context: Additional context for the reflection.
max_tokens: Maximum tokens for the response.
response_schema: Optional JSON Schema for structured output.
request_context: Request context for authentication.
Returns:
ReflectResult with generated response and supporting facts.
"""
...
# =========================================================================
# Bank Management
# =========================================================================
@abstractmethod
async def list_banks(
self,
*,
request_context: "RequestContext",
) -> list[dict[str, Any]]:
"""
List all memory banks.
Args:
request_context: Request context for authentication.
Returns:
List of bank info dicts.
"""
...
@abstractmethod
async def get_bank_profile(
self,
bank_id: str,
*,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Get bank profile including disposition and mission.
Args:
bank_id: The memory bank ID.
request_context: Request context for authentication.
Returns:
Bank profile dict with bank_id, name, disposition, and mission.
"""
...
@abstractmethod
async def update_bank_disposition(
self,
bank_id: str,
disposition: dict[str, int],
*,
request_context: "RequestContext",
) -> None:
"""
Update bank disposition traits.
Args:
bank_id: The memory bank ID.
disposition: Dict with trait values.
request_context: Request context for authentication.
"""
...
@abstractmethod
async def merge_bank_mission(
self,
bank_id: str,
new_info: str,
*,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Merge new mission information into bank profile.
Args:
bank_id: The memory bank ID.
new_info: New mission information to merge.
request_context: Request context for authentication.
Returns:
Updated mission info.
"""
...
@abstractmethod
async def set_bank_mission(
self,
bank_id: str,
mission: str,
*,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Set the bank's mission (replaces existing).
Args:
bank_id: The memory bank ID.
mission: The mission text.
request_context: Request context for authentication.
Returns:
Dict with bank_id and mission.
"""
...
@abstractmethod
async def delete_bank(
self,
bank_id: str,
*,
fact_type: str | None = None,
request_context: "RequestContext",
) -> dict[str, int]:
"""
Delete a bank or its memories.
Args:
bank_id: The memory bank ID.
fact_type: If specified, only delete memories of this type.
request_context: Request context for authentication.
Returns:
Dict with deletion counts.
"""
...
# =========================================================================
# Memory Units
# =========================================================================
@abstractmethod
async def list_memory_units(
self,
bank_id: str,
*,
fact_type: str | None = None,
search_query: str | None = None,
limit: int = 100,
offset: int = 0,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
List memory units with pagination.
Args:
bank_id: The memory bank ID.
fact_type: Filter by fact type.
search_query: Full-text search query.
limit: Maximum results.
offset: Pagination offset.
request_context: Request context for authentication.
Returns:
Dict with 'items', 'total', 'limit', 'offset'.
"""
...
@abstractmethod
async def delete_memory_unit(
self,
unit_id: str,
*,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Delete a specific memory unit.
Args:
unit_id: The memory unit ID.
request_context: Request context for authentication.
Returns:
Deletion result.
"""
...
@abstractmethod
async def get_graph_data(
self,
bank_id: str,
*,
fact_type: str | None = None,
limit: int = 1000,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Get graph data for visualization.
Args:
bank_id: The memory bank ID.
fact_type: Filter by fact type.
limit: Maximum number of items to return (default: 1000).
request_context: Request context for authentication.
Returns:
Dict with nodes, edges, table_rows, total_units, limit.
"""
...
# =========================================================================
# Documents
# =========================================================================
@abstractmethod
async def list_documents(
self,
bank_id: str,
*,
search_query: str | None = None,
limit: int = 100,
offset: int = 0,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
List documents with pagination.
Args:
bank_id: The memory bank ID.
search_query: Search query.
limit: Maximum results.
offset: Pagination offset.
request_context: Request context for authentication.
Returns:
Dict with 'items', 'total', 'limit', 'offset'.
"""
...
@abstractmethod
async def get_document(
self,
document_id: str,
bank_id: str,
*,
request_context: "RequestContext",
) -> dict[str, Any] | None:
"""
Get a specific document.
Args:
document_id: The document ID.
bank_id: The memory bank ID.
request_context: Request context for authentication.
Returns:
Document dict or None if not found.
"""
...
@abstractmethod
async def delete_document(
self,
document_id: str,
bank_id: str,
*,
request_context: "RequestContext",
) -> dict[str, int]:
"""
Delete a document and its memory units.
Args:
document_id: The document ID.
bank_id: The memory bank ID.
request_context: Request context for authentication.
Returns:
Dict with deletion counts.
"""
...
@abstractmethod
async def get_chunk(
self,
chunk_id: str,
*,
request_context: "RequestContext",
) -> dict[str, Any] | None:
"""
Get a specific chunk.
Args:
chunk_id: The chunk ID.
request_context: Request context for authentication.
Returns:
Chunk dict or None if not found.
"""
...
# =========================================================================
# Entities
# =========================================================================
@abstractmethod
async def list_entities(
self,
bank_id: str,
*,
limit: int = 100,
offset: int = 0,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
List entities for a bank with pagination.
Args:
bank_id: The memory bank ID.
limit: Maximum results.
offset: Offset for pagination.
request_context: Request context for authentication.
Returns:
Dict with items, total, limit, offset.
"""
...
@abstractmethod
async def get_entity_observations(
self,
bank_id: str,
entity_id: str,
*,
limit: int = 10,
request_context: "RequestContext",
) -> list[Any]:
"""
Get observations for an entity.
Args:
bank_id: The memory bank ID.
entity_id: The entity ID.
limit: Maximum observations.
request_context: Request context for authentication.
Returns:
List of EntityObservation objects.
"""
...
@abstractmethod
async def regenerate_entity_observations(
self,
bank_id: str,
entity_id: str,
entity_name: str,
*,
request_context: "RequestContext",
) -> None:
"""
Regenerate observations for an entity.
Args:
bank_id: The memory bank ID.
entity_id: The entity ID.
entity_name: The entity's canonical name.
request_context: Request context for authentication.
"""
...
# =========================================================================
# Statistics & Operations
# =========================================================================
@abstractmethod
async def get_bank_stats(
self,
bank_id: str,
*,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Get statistics about memory nodes and links for a bank.
Args:
bank_id: The memory bank ID.
request_context: Request context for authentication.
Returns:
Dict with node_counts, link_counts, link_counts_by_fact_type,
link_breakdown, and operations stats.
"""
...
@abstractmethod
async def get_entity(
self,
bank_id: str,
entity_id: str,
*,
request_context: "RequestContext",
) -> dict[str, Any] | None:
"""
Get entity details including metadata and observations.
Args:
bank_id: The memory bank ID.
entity_id: The entity ID.
request_context: Request context for authentication.
Returns:
Entity dict with id, canonical_name, mention_count, first_seen,
last_seen, metadata, and observations. None if not found.
"""
...
@abstractmethod
async def list_operations(
self,
bank_id: str,
*,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
List async operations for a bank.
Args:
bank_id: The memory bank ID.
request_context: Request context for authentication.
Returns:
Dict with 'total' (int) and 'operations' (list of operation dicts).
"""
...
@abstractmethod
async def cancel_operation(
self,
bank_id: str,
operation_id: str,
*,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Cancel a pending async operation.
Args:
bank_id: The memory bank ID.
operation_id: The operation ID to cancel.
request_context: Request context for authentication.
Returns:
Dict with success status and message.
Raises:
ValueError: If operation not found.
"""
...
@abstractmethod
async def update_bank(
self,
bank_id: str,
*,
name: str | None = None,
mission: str | None = None,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Update bank name and/or mission.
Args:
bank_id: The memory bank ID.
name: New bank name (optional).
mission: New mission text (optional, replaces existing).
request_context: Request context for authentication.
Returns:
Updated bank profile dict.
"""
...
@abstractmethod
async def submit_async_retain(
self,
bank_id: str,
contents: list[dict[str, Any]],
*,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
Submit a batch retain operation to run asynchronously.
Args:
bank_id: The memory bank ID.
contents: List of content dicts to retain.
request_context: Request context for authentication.
Returns:
Dict with operation_id and items_count.
"""
...
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,18 @@
"""
Mental models module for Hindsight.
Mental models are synthesized summaries that represent understanding. They come
in different subtypes based on how they were created:
- 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.
"""
from .models import MentalModel, MentalModelSubtype
__all__ = ["MentalModel", "MentalModelSubtype"]
@@ -0,0 +1,311 @@
"""
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
@@ -0,0 +1,97 @@
"""
Pydantic models for mental models.
"""
from datetime import datetime, timezone
from enum import Enum
from pydantic import BaseModel, Field
class MentalModelSubtype(str, Enum):
"""Subtype of mental model - how it was created."""
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
class MentalModel(BaseModel):
"""
A mental model representing synthesized understanding.
Mental models are the agent's consolidated knowledge. Unlike raw facts,
mental models provide:
- A one-liner description for quick scanning/retrieval
- A full summary for deep understanding
- Links to related mental models
"""
id: str = Field(description="Unique identifier within the bank")
bank_id: str = Field(description="Bank this mental model belongs to")
subtype: MentalModelSubtype = Field(description="How this model was created")
name: str = Field(description="Human-readable name")
description: str = Field(description="One-liner for quick scanning and retrieval matching")
summary: str | None = Field(default=None, description="Full synthesized understanding")
# References
entity_id: str | None = Field(default=None, description="Reference to entities table when type=entity")
source_facts: list[str] = Field(default_factory=list, description="Fact IDs used to generate summary")
links: list[str] = Field(default_factory=list, description="Related mental model IDs")
# Tags for scoped visibility (similar to document tags)
tags: list[str] = Field(default_factory=list, description="Tags for scoped visibility filtering")
# Timestamps
last_updated: datetime | None = Field(default=None, description="When summary was last regenerated")
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.)")
@@ -0,0 +1,228 @@
"""
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,11 +4,12 @@ Query analysis abstraction for the memory system.
Provides an interface for analyzing natural language queries to extract
structured information like temporal constraints.
"""
from abc import ABC, abstractmethod
from typing import Optional
from datetime import datetime, timedelta
import logging
import re
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from pydantic import BaseModel, Field
logger = logging.getLogger(__name__)
@@ -20,6 +21,7 @@ class TemporalConstraint(BaseModel):
Represents a time range with start and end dates.
"""
start_date: datetime = Field(description="Start of the time range (inclusive)")
end_date: datetime = Field(description="End of the time range (inclusive)")
@@ -33,9 +35,9 @@ class QueryAnalysis(BaseModel):
Contains extracted structured information like temporal constraints.
"""
temporal_constraint: Optional[TemporalConstraint] = Field(
default=None,
description="Extracted temporal constraint, if any"
temporal_constraint: TemporalConstraint | None = Field(
default=None, description="Extracted temporal constraint, if any"
)
@@ -58,9 +60,7 @@ class QueryAnalyzer(ABC):
pass
@abstractmethod
def analyze(
self, query: str, reference_date: Optional[datetime] = None
) -> QueryAnalysis:
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
"""
Analyze a natural language query.
@@ -84,7 +84,7 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
Performance:
- ~10-50ms per query
- No model loading required
- No model loading required (lazy import on first use)
"""
def __init__(self):
@@ -95,11 +95,10 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
"""Load dateparser (lazy import)."""
if self._search_dates is None:
from dateparser.search import search_dates
self._search_dates = search_dates
def analyze(
self, query: str, reference_date: Optional[datetime] = None
) -> QueryAnalysis:
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
"""
Analyze query using dateparser.
@@ -113,8 +112,6 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
Returns:
QueryAnalysis with temporal_constraint if found
"""
self.load()
if reference_date is None:
reference_date = datetime.now()
@@ -124,11 +121,14 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
if period_result is not None:
return QueryAnalysis(temporal_constraint=period_result)
# Lazy load dateparser (only imports on first call, then cached)
self.load()
# Use dateparser's search_dates to find temporal expressions
settings = {
'RELATIVE_BASE': reference_date,
'PREFER_DATES_FROM': 'past',
'RETURN_AS_TIMEZONE_AWARE': False,
"RELATIVE_BASE": reference_date,
"PREFER_DATES_FROM": "past",
"RETURN_AS_TIMEZONE_AWARE": False,
}
results = self._search_dates(query, settings=settings)
@@ -137,11 +137,8 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
return QueryAnalysis(temporal_constraint=None)
# Filter out false positives (common words parsed as dates)
false_positives = {'do', 'may', 'march', 'will', 'can', 'sat', 'sun', 'mon', 'tue', 'wed', 'thu', 'fri'}
valid_results = [
(text, date) for text, date in results
if text.lower() not in false_positives or len(text) > 3
]
false_positives = {"do", "may", "march", "will", "can", "sat", "sun", "mon", "tue", "wed", "thu", "fri"}
valid_results = [(text, date) for text, date in results if text.lower() not in false_positives or len(text) > 3]
if not valid_results:
return QueryAnalysis(temporal_constraint=None)
@@ -153,84 +150,94 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
start_date = parsed_date.replace(hour=0, minute=0, second=0, microsecond=0)
end_date = parsed_date.replace(hour=23, minute=59, second=59, microsecond=999999)
return QueryAnalysis(
temporal_constraint=TemporalConstraint(
start_date=start_date,
end_date=end_date
)
)
return QueryAnalysis(temporal_constraint=TemporalConstraint(start_date=start_date, end_date=end_date))
def _extract_period(
self, query: str, reference_date: datetime
) -> Optional[TemporalConstraint]:
def _extract_period(self, query: str, reference_date: datetime) -> TemporalConstraint | None:
"""
Extract period-based temporal expressions (week, month, year, weekend).
These need special handling as they represent date ranges, not single dates.
Supports multiple languages.
"""
def constraint(start: datetime, end: datetime) -> TemporalConstraint:
return TemporalConstraint(
start_date=start.replace(hour=0, minute=0, second=0, microsecond=0),
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999)
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999),
)
# Yesterday patterns (English, Spanish, Italian, French, German)
if re.search(r'\b(yesterday|ayer|ieri|hier|gestern)\b', query, re.IGNORECASE):
if re.search(r"\b(yesterday|ayer|ieri|hier|gestern)\b", query, re.IGNORECASE):
d = reference_date - timedelta(days=1)
return constraint(d, d)
# Today patterns
if re.search(r'\b(today|hoy|oggi|aujourd\'?hui|heute)\b', query, re.IGNORECASE):
if re.search(r"\b(today|hoy|oggi|aujourd\'?hui|heute)\b", query, re.IGNORECASE):
return constraint(reference_date, reference_date)
# "a couple of days ago" / "a few days ago" patterns
# These are imprecise so we create a range
if re.search(r'\b(a\s+)?couple\s+(of\s+)?days?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?couple\s+(of\s+)?days?\s+ago\b", query, re.IGNORECASE):
# "a couple of days" = approximately 2 days, give range of 1-3 days
return constraint(reference_date - timedelta(days=3), reference_date - timedelta(days=1))
if re.search(r'\b(a\s+)?few\s+days?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?few\s+days?\s+ago\b", query, re.IGNORECASE):
# "a few days" = approximately 3-4 days, give range of 2-5 days
return constraint(reference_date - timedelta(days=5), reference_date - timedelta(days=2))
# "a couple of weeks ago" / "a few weeks ago" patterns
if re.search(r'\b(a\s+)?couple\s+(of\s+)?weeks?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?couple\s+(of\s+)?weeks?\s+ago\b", query, re.IGNORECASE):
# "a couple of weeks" = approximately 2 weeks, give range of 1-3 weeks
return constraint(reference_date - timedelta(weeks=3), reference_date - timedelta(weeks=1))
if re.search(r'\b(a\s+)?few\s+weeks?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?few\s+weeks?\s+ago\b", query, re.IGNORECASE):
# "a few weeks" = approximately 3-4 weeks, give range of 2-5 weeks
return constraint(reference_date - timedelta(weeks=5), reference_date - timedelta(weeks=2))
# "a couple of months ago" / "a few months ago" patterns
if re.search(r'\b(a\s+)?couple\s+(of\s+)?months?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?couple\s+(of\s+)?months?\s+ago\b", query, re.IGNORECASE):
# "a couple of months" = approximately 2 months, give range of 1-3 months
return constraint(reference_date - timedelta(days=90), reference_date - timedelta(days=30))
if re.search(r'\b(a\s+)?few\s+months?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?few\s+months?\s+ago\b", query, re.IGNORECASE):
# "a few months" = approximately 3-4 months, give range of 2-5 months
return constraint(reference_date - timedelta(days=150), reference_date - timedelta(days=60))
# Last week patterns (English, Spanish, Italian, French, German)
if re.search(r'\b(last\s+week|la\s+semana\s+pasada|la\s+settimana\s+scorsa|la\s+semaine\s+derni[eè]re|letzte\s+woche)\b', query, re.IGNORECASE):
if re.search(
r"\b(last\s+week|la\s+semana\s+pasada|la\s+settimana\s+scorsa|la\s+semaine\s+derni[eè]re|letzte\s+woche)\b",
query,
re.IGNORECASE,
):
start = reference_date - timedelta(days=reference_date.weekday() + 7)
return constraint(start, start + timedelta(days=6))
# Last month patterns
if re.search(r'\b(last\s+month|el\s+mes\s+pasado|il\s+mese\s+scorso|le\s+mois\s+dernier|letzten?\s+monat)\b', query, re.IGNORECASE):
if re.search(
r"\b(last\s+month|el\s+mes\s+pasado|il\s+mese\s+scorso|le\s+mois\s+dernier|letzten?\s+monat)\b",
query,
re.IGNORECASE,
):
first = reference_date.replace(day=1)
end = first - timedelta(days=1)
start = end.replace(day=1)
return constraint(start, end)
# Last year patterns
if re.search(r'\b(last\s+year|el\s+a[ñn]o\s+pasado|l\'anno\s+scorso|l\'ann[ée]e\s+derni[eè]re|letztes?\s+jahr)\b', query, re.IGNORECASE):
if re.search(
r"\b(last\s+year|el\s+a[ñn]o\s+pasado|l\'anno\s+scorso|l\'ann[ée]e\s+derni[eè]re|letztes?\s+jahr)\b",
query,
re.IGNORECASE,
):
year = reference_date.year - 1
return constraint(datetime(year, 1, 1), datetime(year, 12, 31))
# Last weekend patterns
if re.search(r'\b(last\s+weekend|el\s+fin\s+de\s+semana\s+pasado|lo\s+scorso\s+fine\s+settimana|le\s+week-?end\s+dernier|letztes?\s+wochenende)\b', query, re.IGNORECASE):
if re.search(
r"\b(last\s+weekend|el\s+fin\s+de\s+semana\s+pasado|lo\s+scorso\s+fine\s+settimana|le\s+week-?end\s+dernier|letztes?\s+wochenende)\b",
query,
re.IGNORECASE,
):
days_since_sat = (reference_date.weekday() + 2) % 7
if days_since_sat == 0:
days_since_sat = 7
@@ -239,22 +246,22 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
# Month + Year patterns (e.g., "June 2024", "junio 2024", "giugno 2024")
month_patterns = {
'january|enero|gennaio|janvier|januar': 1,
'february|febrero|febbraio|f[ée]vrier|februar': 2,
'march|marzo|mars|m[äa]rz': 3,
'april|abril|aprile|avril': 4,
'may|mayo|maggio|mai': 5,
'june|junio|giugno|juin|juni': 6,
'july|julio|luglio|juillet|juli': 7,
'august|agosto|ao[uû]t': 8,
'september|septiembre|settembre|septembre': 9,
'october|octubre|ottobre|octobre|oktober': 10,
'november|noviembre|novembre': 11,
'december|diciembre|dicembre|d[ée]cembre|dezember': 12,
"january|enero|gennaio|janvier|januar": 1,
"february|febrero|febbraio|f[ée]vrier|februar": 2,
"march|marzo|mars|m[äa]rz": 3,
"april|abril|aprile|avril": 4,
"may|mayo|maggio|mai": 5,
"june|junio|giugno|juin|juni": 6,
"july|julio|luglio|juillet|juli": 7,
"august|agosto|ao[uû]t": 8,
"september|septiembre|settembre|septembre": 9,
"october|octubre|ottobre|octobre|oktober": 10,
"november|noviembre|novembre": 11,
"december|diciembre|dicembre|d[ée]cembre|dezember": 12,
}
for pattern, month_num in month_patterns.items():
match = re.search(rf'\b({pattern})\s+(\d{{4}})\b', query, re.IGNORECASE)
match = re.search(rf"\b({pattern})\s+(\d{{4}})\b", query, re.IGNORECASE)
if match:
year = int(match.group(2))
start = datetime(year, month_num, 1)
@@ -279,11 +286,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
- Model size: ~80M params (~300MB download)
"""
def __init__(
self,
model_name: str = "google/flan-t5-small",
device: str = "cpu"
):
def __init__(self, model_name: str = "google/flan-t5-small", device: str = "cpu"):
"""
Initialize T5 query analyzer.
@@ -304,11 +307,10 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
return
try:
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
except ImportError:
raise ImportError(
"transformers is required for TransformerQueryAnalyzer. "
"Install it with: pip install transformers"
"transformers is required for TransformerQueryAnalyzer. Install it with: pip install transformers"
)
logger.info(f"Loading query analyzer model: {self.model_name}...")
@@ -322,9 +324,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
"""Lazy load the T5 model for temporal extraction (calls load())."""
self.load()
def _extract_with_rules(
self, query: str, reference_date: datetime
) -> Optional[TemporalConstraint]:
def _extract_with_rules(self, query: str, reference_date: datetime) -> TemporalConstraint | None:
"""
Extract temporal expressions using rule-based patterns.
@@ -332,6 +332,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
patterns that need model-based extraction.
"""
import re
query_lower = query.lower()
def get_last_weekday(weekday: int) -> datetime:
@@ -343,50 +344,60 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
def constraint(start: datetime, end: datetime) -> TemporalConstraint:
return TemporalConstraint(
start_date=start.replace(hour=0, minute=0, second=0, microsecond=0),
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999)
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999),
)
# Yesterday
if re.search(r'\byesterday\b', query_lower):
if re.search(r"\byesterday\b", query_lower):
d = reference_date - timedelta(days=1)
return constraint(d, d)
# Last week
if re.search(r'\blast\s+week\b', query_lower):
if re.search(r"\blast\s+week\b", query_lower):
start = reference_date - timedelta(days=reference_date.weekday() + 7)
return constraint(start, start + timedelta(days=6))
# Last month
if re.search(r'\blast\s+month\b', query_lower):
if re.search(r"\blast\s+month\b", query_lower):
first = reference_date.replace(day=1)
end = first - timedelta(days=1)
start = end.replace(day=1)
return constraint(start, end)
# Last year
if re.search(r'\blast\s+year\b', query_lower):
if re.search(r"\blast\s+year\b", query_lower):
y = reference_date.year - 1
return constraint(datetime(y, 1, 1), datetime(y, 12, 31))
# Last weekend
if re.search(r'\blast\s+weekend\b', query_lower):
if re.search(r"\blast\s+weekend\b", query_lower):
sat = get_last_weekday(5)
return constraint(sat, sat + timedelta(days=1))
# Last <weekday>
weekdays = {'monday': 0, 'tuesday': 1, 'wednesday': 2, 'thursday': 3,
'friday': 4, 'saturday': 5, 'sunday': 6}
weekdays = {"monday": 0, "tuesday": 1, "wednesday": 2, "thursday": 3, "friday": 4, "saturday": 5, "sunday": 6}
for name, num in weekdays.items():
if re.search(rf'\blast\s+{name}\b', query_lower):
if re.search(rf"\blast\s+{name}\b", query_lower):
d = get_last_weekday(num)
return constraint(d, d)
# Month + Year: "June 2024", "in March 2023"
months = {'january': 1, 'february': 2, 'march': 3, 'april': 4, 'may': 5,
'june': 6, 'july': 7, 'august': 8, 'september': 9, 'october': 10,
'november': 11, 'december': 12}
months = {
"january": 1,
"february": 2,
"march": 3,
"april": 4,
"may": 5,
"june": 6,
"july": 7,
"august": 8,
"september": 9,
"october": 10,
"november": 11,
"december": 12,
}
for name, num in months.items():
match = re.search(rf'\b{name}\s+(\d{{4}})\b', query_lower)
match = re.search(rf"\b{name}\s+(\d{{4}})\b", query_lower)
if match:
year = int(match.group(1))
if num == 12:
@@ -397,9 +408,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
return None
def analyze(
self, query: str, reference_date: Optional[datetime] = None
) -> QueryAnalysis:
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
"""
Analyze query for temporal expressions.
@@ -435,11 +444,11 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
last_saturday = get_last_weekday(5)
# Build prompt for T5
prompt = f"""Today is {reference_date.strftime('%Y-%m-%d')}. Extract date range or "none".
prompt = f"""Today is {reference_date.strftime("%Y-%m-%d")}. Extract date range or "none".
June 2024 = 2024-06-01 to 2024-06-30
yesterday = {yesterday.strftime('%Y-%m-%d')} to {yesterday.strftime('%Y-%m-%d')}
last Saturday = {last_saturday.strftime('%Y-%m-%d')} to {last_saturday.strftime('%Y-%m-%d')}
yesterday = {yesterday.strftime("%Y-%m-%d")} to {yesterday.strftime("%Y-%m-%d")}
last Saturday = {last_saturday.strftime("%Y-%m-%d")} to {last_saturday.strftime("%Y-%m-%d")}
what is the weather = none
{query} ="""
@@ -448,13 +457,7 @@ what is the weather = none
inputs = {k: v.to(self.device) for k, v in inputs.items()}
with self._no_grad():
outputs = self._model.generate(
**inputs,
max_new_tokens=30,
num_beams=3,
do_sample=False,
temperature=1.0
)
outputs = self._model.generate(**inputs, max_new_tokens=30, num_beams=3, do_sample=False, temperature=1.0)
result = self._tokenizer.decode(outputs[0], skip_special_tokens=True).strip()
@@ -466,14 +469,14 @@ what is the weather = none
"""Get torch.no_grad context manager."""
try:
import torch
return torch.no_grad()
except ImportError:
from contextlib import nullcontext
return nullcontext()
def _parse_generated_output(
self, result: str, reference_date: datetime
) -> Optional[TemporalConstraint]:
def _parse_generated_output(self, result: str, reference_date: datetime) -> TemporalConstraint | None:
"""
Parse T5 generated output into TemporalConstraint.
@@ -492,7 +495,8 @@ what is the weather = none
try:
# Parse "YYYY-MM-DD to YYYY-MM-DD"
import re
pattern = r'(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})'
pattern = r"(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})"
match = re.search(pattern, result, re.IGNORECASE)
if match:
@@ -513,7 +517,7 @@ what is the weather = none
return TemporalConstraint(start_date=start_date, end_date=end_date)
except (ValueError, AttributeError) as e:
except (ValueError, AttributeError):
return None
return None
@@ -0,0 +1,20 @@
"""
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)
"""
from .agent import ReflectAgentResult, run_reflect_agent
from .models import MentalModelInput, ReflectAction, ReflectActionBatch
__all__ = [
"run_reflect_agent",
"ReflectAgentResult",
"ReflectAction",
"ReflectActionBatch",
"MentalModelInput",
]
@@ -0,0 +1,731 @@
"""
Reflect agent - agentic loop for reflection with native tool calling.
"""
import asyncio
import json
import logging
import time
from typing import TYPE_CHECKING, Any, Awaitable, Callable, Literal
from .models import LLMCall, MentalModelInput, Observation, ReflectAgentResult, ToolCall
from .prompts import FINAL_SYSTEM_PROMPT, build_final_prompt, build_system_prompt_for_tools
from .tools_schema import get_reflect_tools
if TYPE_CHECKING:
from ..llm_wrapper import LLMProvider
from ..response_models import LLMToolCall
logger = logging.getLogger(__name__)
DEFAULT_MAX_ITERATIONS = 10
async def _generate_structured_output(
answer: str,
response_schema: dict,
llm_config: "LLMProvider",
reflect_id: str,
) -> dict[str, Any] | None:
"""Generate structured output from an answer using the provided JSON schema.
Args:
answer: The text answer to extract structured data from
response_schema: JSON Schema for the expected output structure
llm_config: LLM provider for making the extraction call
reflect_id: Reflect ID for logging
Returns:
Structured output dict if successful, None otherwise
"""
try:
from typing import Any as TypingAny
from pydantic import create_model
def _json_schema_type_to_python(field_schema: dict) -> type:
"""Map JSON schema type to Python type for better LLM guidance."""
json_type = field_schema.get("type", "string")
if json_type == "array":
return list
elif json_type == "object":
return dict
elif json_type == "integer":
return int
elif json_type == "number":
return float
elif json_type == "boolean":
return bool
else:
return str
# Build fields from JSON schema properties
schema_props = response_schema.get("properties", {})
required_fields = set(response_schema.get("required", []))
fields: dict[str, TypingAny] = {}
for field_name, field_schema in schema_props.items():
field_type = _json_schema_type_to_python(field_schema)
default = ... if field_name in required_fields else None
fields[field_name] = (field_type, default)
if not fields:
return None
DynamicModel = create_model("StructuredResponse", **fields)
# Include the full schema in the prompt for better LLM guidance
schema_str = json.dumps(response_schema, indent=2)
# Call LLM with the answer to extract structured data
structured_prompt = f"""Based on this answer, extract the information into the requested structured format.
Answer: {answer}
JSON Schema to follow:
```json
{schema_str}
```
Return ONLY a valid JSON object that matches this exact schema. Pay special attention to field types:
- "type": "array" means the value must be a JSON array/list, NOT a string
- "type": "string" means the value must be a string
- "type": "object" means the value must be a JSON object
Do not include any explanation, only the JSON object."""
structured_result = await llm_config.call(
messages=[
{
"role": "system",
"content": "Extract structured data from the given answer. Return only valid JSON matching the provided schema exactly.",
},
{"role": "user", "content": structured_prompt},
],
response_format=DynamicModel,
scope="reflect_structured",
skip_validation=True, # We'll handle the dict ourselves
)
# Convert to dict
if hasattr(structured_result, "model_dump"):
structured_output = structured_result.model_dump()
elif isinstance(structured_result, dict):
structured_output = structured_result
else:
# Try to parse as JSON
structured_output = json.loads(str(structured_result))
logger.info(f"[REFLECT {reflect_id}] Generated structured output with {len(structured_output)} fields")
return structured_output
except Exception as e:
logger.warning(f"[REFLECT {reflect_id}] Failed to generate structured output: {e}")
return None
async def run_reflect_agent(
llm_config: "LLMProvider",
bank_id: str,
query: str,
bank_profile: dict[str, Any],
lookup_fn: Callable[[str | None], Awaitable[dict[str, Any]]],
recall_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
learn_fn: Callable[[MentalModelInput], Awaitable[dict[str, Any]]] | None = None,
context: str | None = None,
max_iterations: int = DEFAULT_MAX_ITERATIONS,
max_tokens: int | None = None,
response_schema: dict | None = None,
output_mode: Literal["answer", "observations"] = "answer",
) -> ReflectAgentResult:
"""
Execute the reflect agent loop using native tool calling.
The agent iteratively calls tools to gather information and learn,
then provides a final answer via the done() tool.
Args:
llm_config: LLM provider for agent calls
bank_id: Bank identifier
query: Question to answer
bank_profile: Bank profile with name and mission
lookup_fn: Tool callback for lookup (model_id) -> result
recall_fn: Tool callback for recall (query, max_tokens) -> result
expand_fn: Tool callback for expand (memory_id, depth) -> result
learn_fn: Optional tool callback for learn (MentalModelInput) -> result.
If None, learn tool is disabled.
context: Optional additional context
max_iterations: Maximum number of iterations before forcing response
max_tokens: Maximum tokens for the final response
response_schema: Optional JSON Schema for structured output in final response
output_mode: "answer" returns final text, "observations" returns structured observations
Returns:
ReflectAgentResult with final answer and metadata
"""
enable_learn = learn_fn is not None
reflect_id = f"{bank_id[:8]}-{int(time.time() * 1000) % 100000}"
start_time = time.time()
# Get tools for this agent
tools = get_reflect_tools(enable_learn=enable_learn, output_mode=output_mode)
# Build initial messages
system_prompt = build_system_prompt_for_tools(bank_profile, context, output_mode=output_mode)
messages: list[dict[str, Any]] = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": query},
]
# Tracking
mental_models_created: list[str] = []
total_tools_called = 0
tool_trace: list[ToolCall] = []
tool_trace_summary: list[dict[str, Any]] = []
llm_trace: list[dict[str, Any]] = []
context_history: list[dict[str, Any]] = [] # For final prompt fallback
# Track available IDs for validation (prevents hallucinated citations)
available_memory_ids: set[str] = set()
available_model_ids: set[str] = set()
# In answer mode, pre-fetch mental models so the agent always starts with this knowledge
if output_mode == "answer":
prefetch_start = time.time()
models_result = await lookup_fn(None) # List all mental models
prefetch_duration = int((time.time() - prefetch_start) * 1000)
# Track available model IDs
if isinstance(models_result, dict) and "models" in models_result:
for model in models_result["models"]:
if "id" in model:
available_model_ids.add(model["id"])
# Add to context history for the agent
context_history.append({"tool": "list_mental_models", "output": models_result})
# Add to tool trace
tool_trace.append(
ToolCall(
tool="list_mental_models",
input={"tool": "list_mental_models"},
output=models_result,
duration_ms=prefetch_duration,
iteration=0,
)
)
tool_trace_summary.append(
{
"tool": "list_mental_models",
"input_summary": "(prefetch)",
"duration_ms": prefetch_duration,
"output_chars": len(json.dumps(models_result, default=str)),
}
)
total_tools_called += 1
# Include in the user message so the agent sees it
models_info = json.dumps(models_result, indent=2, default=str)
messages[1]["content"] = f"{query}\n\n## Available Mental Models (pre-fetched)\n```json\n{models_info}\n```"
def _get_llm_trace() -> list[LLMCall]:
return [LLMCall(scope=c["scope"], duration_ms=c["duration_ms"]) for c in llm_trace]
def _log_completion(answer: str, iterations: int, forced: bool = False):
elapsed_ms = int((time.time() - start_time) * 1000)
tools_summary = (
", ".join(
f"{t['tool']}({t['input_summary']})={t['duration_ms']}ms/{t.get('output_chars', 0)}c"
for t in tool_trace_summary
)
or "none"
)
llm_summary = ", ".join(f"{c['scope']}={c['duration_ms']}ms" for c in llm_trace) or "none"
total_llm_ms = sum(c["duration_ms"] for c in llm_trace)
total_tools_ms = sum(t["duration_ms"] for t in tool_trace_summary)
answer_preview = answer[:100] + "..." if len(answer) > 100 else answer
mode = "forced" if forced else "done"
logger.info(
f"[REFLECT {reflect_id}] {mode} | "
f"query='{query[:50]}...' | "
f"iterations={iterations} | "
f"llm=[{llm_summary}] ({total_llm_ms}ms) | "
f"tools=[{tools_summary}] ({total_tools_ms}ms) | "
f"answer='{answer_preview}' | "
f"total={elapsed_ms}ms"
)
for iteration in range(max_iterations):
is_last = iteration == max_iterations - 1
if is_last:
# Force text response on last iteration - no tools
prompt = build_final_prompt(query, context_history, bank_profile, context)
llm_start = time.time()
response = await llm_config.call(
messages=[
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
{"role": "user", "content": prompt},
],
scope="reflect_agent_final",
max_completion_tokens=max_tokens,
)
llm_trace.append({"scope": "final", "duration_ms": int((time.time() - llm_start) * 1000)})
answer = response.strip()
# Generate structured output if schema provided
structured_output = None
if response_schema and answer:
structured_output = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
_log_completion(answer, iteration + 1, forced=True)
return ReflectAgentResult(
text=answer,
structured_output=structured_output,
iterations=iteration + 1,
tools_called=total_tools_called,
mental_models_created=mental_models_created,
tool_trace=tool_trace,
llm_trace=_get_llm_trace(),
)
# Call LLM with tools
llm_start = time.time()
try:
result = await llm_config.call_with_tools(
messages=messages,
tools=tools,
scope="reflect_agent",
tool_choice="required" if iteration == 0 else "auto", # Force tool use on first iteration
)
llm_duration = int((time.time() - llm_start) * 1000)
llm_trace.append({"scope": f"agent_{iteration + 1}", "duration_ms": llm_duration})
except Exception:
llm_trace.append(
{"scope": f"agent_{iteration + 1}_err", "duration_ms": int((time.time() - llm_start) * 1000)}
)
# Guardrail: If no evidence gathered yet, retry
has_gathered_evidence = bool(available_memory_ids) or bool(available_model_ids)
if not has_gathered_evidence and iteration < max_iterations - 1:
continue
prompt = build_final_prompt(query, context_history, bank_profile, context)
llm_start = time.time()
response = await llm_config.call(
messages=[
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
{"role": "user", "content": prompt},
],
scope="reflect_agent_final",
max_completion_tokens=max_tokens,
)
llm_trace.append({"scope": "final", "duration_ms": int((time.time() - llm_start) * 1000)})
answer = response.strip()
# Generate structured output if schema provided
structured_output = None
if response_schema and answer:
structured_output = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
_log_completion(answer, iteration + 1, forced=True)
return ReflectAgentResult(
text=answer,
structured_output=structured_output,
iterations=iteration + 1,
tools_called=total_tools_called,
mental_models_created=mental_models_created,
tool_trace=tool_trace,
llm_trace=_get_llm_trace(),
)
# No tool calls - LLM wants to respond with text
if not result.tool_calls:
if result.content:
answer = result.content.strip()
# Generate structured output if schema provided
structured_output = None
if response_schema and answer:
structured_output = await _generate_structured_output(
answer, response_schema, llm_config, reflect_id
)
_log_completion(answer, iteration + 1)
return ReflectAgentResult(
text=answer,
structured_output=structured_output,
iterations=iteration + 1,
tools_called=total_tools_called,
mental_models_created=mental_models_created,
tool_trace=tool_trace,
llm_trace=_get_llm_trace(),
)
# Empty response, force final
prompt = build_final_prompt(query, context_history, bank_profile, context)
llm_start = time.time()
response = await llm_config.call(
messages=[
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
{"role": "user", "content": prompt},
],
scope="reflect_agent_final",
max_completion_tokens=max_tokens,
)
llm_trace.append({"scope": "final", "duration_ms": int((time.time() - llm_start) * 1000)})
answer = response.strip()
# Generate structured output if schema provided
structured_output = None
if response_schema and answer:
structured_output = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
_log_completion(answer, iteration + 1, forced=True)
return ReflectAgentResult(
text=answer,
structured_output=structured_output,
iterations=iteration + 1,
tools_called=total_tools_called,
mental_models_created=mental_models_created,
tool_trace=tool_trace,
llm_trace=_get_llm_trace(),
)
# Check for done tool call
done_call = next((tc for tc in result.tool_calls if tc.name == "done"), None)
if done_call:
# Guardrail: Require evidence before done
has_gathered_evidence = bool(available_memory_ids) or bool(available_model_ids)
if not has_gathered_evidence and iteration < max_iterations - 1:
# Add assistant message and fake tool result asking for evidence
messages.append(
{
"role": "assistant",
"tool_calls": [_tool_call_to_dict(done_call)],
}
)
messages.append(
{
"role": "tool",
"tool_call_id": done_call.id,
"content": json.dumps(
{
"error": "You must call recall() or list_mental_models() to gather evidence before providing your final answer."
}
),
}
)
continue
# Process done tool
return await _process_done_tool(
done_call,
output_mode,
available_memory_ids,
available_model_ids,
iteration + 1,
total_tools_called,
mental_models_created,
tool_trace,
_get_llm_trace(),
_log_completion,
reflect_id,
llm_config=llm_config,
response_schema=response_schema,
)
# Execute other tools in parallel
other_tools = [tc for tc in result.tool_calls if tc.name != "done"]
if other_tools:
# Add assistant message with tool calls
messages.append(
{
"role": "assistant",
"tool_calls": [_tool_call_to_dict(tc) for tc in other_tools],
}
)
# Execute tools in parallel
tool_tasks = [
_execute_tool_with_timing(tc, lookup_fn, recall_fn, expand_fn, learn_fn) for tc in other_tools
]
tool_results = await asyncio.gather(*tool_tasks, return_exceptions=True)
total_tools_called += len(other_tools)
# Process results and add to messages
for tc, result_data in zip(other_tools, tool_results):
if isinstance(result_data, Exception):
# Tool execution failed - log and raise to fail the request
logger.error(f"[REFLECT {reflect_id}] Tool {tc.name} failed with exception: {result_data}")
raise RuntimeError(f"Reflect tool '{tc.name}' failed: {result_data}")
output, duration_ms = result_data
# Check if tool returned an error response
if isinstance(output, dict) and "error" in output:
logger.error(f"[REFLECT {reflect_id}] Tool {tc.name} returned error: {output['error']}")
raise RuntimeError(f"Reflect tool '{tc.name}' error: {output['error']}")
# Track created mental models
if tc.name == "learn" and isinstance(output, dict) and "model_id" in output:
mental_models_created.append(output["model_id"])
# Track available memory IDs from recall
if tc.name == "recall" and isinstance(output, dict) and "memories" in output:
for memory in output["memories"]:
if "id" in memory:
available_memory_ids.add(memory["id"])
# Track available model IDs
if tc.name in ("list_mental_models", "get_mental_model") and isinstance(output, dict):
if output.get("found") and "model" in output:
model_id = output["model"].get("id")
if model_id:
available_model_ids.add(model_id)
elif "models" in output:
for model in output["models"]:
if "id" in model:
available_model_ids.add(model["id"])
# Add tool result message
messages.append(
{
"role": "tool",
"tool_call_id": tc.id,
"content": json.dumps(output, default=str),
}
)
# Track for logging and context history
input_dict = {"tool": tc.name, **tc.arguments}
input_summary = _summarize_input(tc.name, tc.arguments)
tool_trace.append(
ToolCall(
tool=tc.name, input=input_dict, output=output, duration_ms=duration_ms, iteration=iteration + 1
)
)
try:
output_chars = len(json.dumps(output))
except (TypeError, ValueError):
output_chars = len(str(output))
tool_trace_summary.append(
{
"tool": tc.name,
"input_summary": input_summary,
"duration_ms": duration_ms,
"output_chars": output_chars,
}
)
# Keep context history for fallback final prompt
context_history.append({"tool": tc.name, "input": input_dict, "output": output})
# Should not reach here
answer = "I was unable to formulate a complete answer within the iteration limit."
_log_completion(answer, max_iterations, forced=True)
return ReflectAgentResult(
text=answer,
iterations=max_iterations,
tools_called=total_tools_called,
mental_models_created=mental_models_created,
tool_trace=tool_trace,
llm_trace=_get_llm_trace(),
)
def _tool_call_to_dict(tc: "LLMToolCall") -> dict[str, Any]:
"""Convert LLMToolCall to OpenAI message format."""
return {
"id": tc.id,
"type": "function",
"function": {
"name": tc.name,
"arguments": json.dumps(tc.arguments),
},
}
async def _process_done_tool(
done_call: "LLMToolCall",
output_mode: str,
available_memory_ids: set[str],
available_model_ids: set[str],
iterations: int,
total_tools_called: int,
mental_models_created: list[str],
tool_trace: list[ToolCall],
llm_trace: list[LLMCall],
log_completion: Callable,
reflect_id: str,
llm_config: "LLMProvider | None" = None,
response_schema: dict | None = None,
) -> ReflectAgentResult:
"""Process the done tool call and return the result."""
args = done_call.arguments
if output_mode == "observations" and "observations" in args:
# Process observations - handle both list and nested {"observations": [...]} format
observations: list[Observation] = []
used_memory_ids: list[str] = []
obs_list = args["observations"]
# Handle nested format where LLM outputs {"observations": [...]} instead of just [...]
if isinstance(obs_list, dict) and "observations" in obs_list:
obs_list = obs_list["observations"]
for obs_data in obs_list:
validated_mids = []
for mid in obs_data.get("memory_ids", []):
if mid in available_memory_ids:
validated_mids.append(mid)
if mid not in used_memory_ids:
used_memory_ids.append(mid)
observations.append(
Observation(
title=obs_data.get("title", ""),
text=obs_data.get("text", ""),
memory_ids=validated_mids,
)
)
# Build text from observations
text_parts = []
for obs in observations:
if obs.title:
text_parts.append(f"## {obs.title}\n{obs.text}")
else:
text_parts.append(obs.text)
answer = "\n\n".join(text_parts)
log_completion(answer, iterations)
return ReflectAgentResult(
text=answer,
observations=observations,
iterations=iterations,
tools_called=total_tools_called,
mental_models_created=mental_models_created,
tool_trace=tool_trace,
llm_trace=llm_trace,
used_memory_ids=used_memory_ids,
)
# Default: answer mode
answer = args.get("answer", "").strip()
if not answer:
answer = "No answer provided."
# Validate IDs
used_memory_ids = [mid for mid in args.get("memory_ids", []) if mid in available_memory_ids]
used_model_ids = [mid for mid in args.get("model_ids", []) if mid in available_model_ids]
# Generate structured output if schema provided
structured_output = None
if response_schema and llm_config and answer:
structured_output = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
log_completion(answer, iterations)
return ReflectAgentResult(
text=answer,
structured_output=structured_output,
iterations=iterations,
tools_called=total_tools_called,
mental_models_created=mental_models_created,
tool_trace=tool_trace,
llm_trace=llm_trace,
used_memory_ids=used_memory_ids,
used_model_ids=used_model_ids,
)
async def _execute_tool_with_timing(
tc: "LLMToolCall",
lookup_fn: Callable[[str | None], Awaitable[dict[str, Any]]],
recall_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
learn_fn: Callable[[MentalModelInput], Awaitable[dict[str, Any]]] | None = None,
) -> tuple[dict[str, Any], int]:
"""Execute a tool call and return result with timing."""
start = time.time()
result = await _execute_tool(tc.name, tc.arguments, lookup_fn, recall_fn, expand_fn, learn_fn)
duration_ms = int((time.time() - start) * 1000)
return result, duration_ms
async def _execute_tool(
tool_name: str,
args: dict[str, Any],
lookup_fn: Callable[[str | None], Awaitable[dict[str, Any]]],
recall_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
learn_fn: Callable[[MentalModelInput], Awaitable[dict[str, Any]]] | None = None,
) -> dict[str, Any]:
"""Execute a single tool by name."""
if tool_name == "list_mental_models":
return await lookup_fn(None)
elif tool_name == "get_mental_model":
model_id = args.get("model_id")
if not model_id:
return {"error": "get_mental_model requires model_id"}
return await lookup_fn(model_id)
elif tool_name == "recall":
query = args.get("query")
if not query:
return {"error": "recall requires a query parameter"}
max_tokens = max(args.get("max_tokens") or 2048, 1000) # Default 2048, min 1000
return await recall_fn(query, max_tokens)
elif tool_name == "learn":
if learn_fn is None:
return {"error": "learn tool is not available"}
name = args.get("name")
description = args.get("description")
if not name or not description:
return {"error": "learn requires name and description"}
return await learn_fn(MentalModelInput(name=name, description=description))
elif tool_name == "expand":
memory_ids = args.get("memory_ids", [])
if not memory_ids:
return {"error": "expand requires memory_ids"}
depth = args.get("depth", "chunk")
return await expand_fn(memory_ids, depth)
else:
return {"error": f"Unknown tool: {tool_name}"}
def _summarize_input(tool_name: str, args: dict[str, Any]) -> str:
"""Create a summary of tool input for logging, showing all params."""
if tool_name == "list_mental_models":
return "()"
elif tool_name == "get_mental_model":
return f"(model_id={args.get('model_id', '?')})"
elif tool_name == "recall":
query = args.get("query", "")
query_preview = f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
# Show actual value used (default 2048, min 1000)
max_tokens = max(args.get("max_tokens") or 2048, 1000)
return f"(query={query_preview}, max_tokens={max_tokens})"
elif tool_name == "learn":
name = args.get("name", "?")
desc = args.get("description", "")
desc_preview = f"'{desc[:20]}...'" if len(desc) > 20 else f"'{desc}'"
return f"(name='{name}', description={desc_preview})"
elif tool_name == "expand":
memory_ids = args.get("memory_ids", [])
depth = args.get("depth", "chunk")
return f"(memory_ids=[{len(memory_ids)} ids], depth={depth})"
elif tool_name == "done":
answer = args.get("answer", "")
answer_preview = f"'{answer[:30]}...'" if len(answer) > 30 else f"'{answer}'"
memory_ids = args.get("memory_ids", [])
model_ids = args.get("model_ids", [])
return f"(answer={answer_preview}, memory_ids={len(memory_ids)}, model_ids={len(model_ids)})"
return str(args)
@@ -0,0 +1,114 @@
"""
Pydantic models for the reflect agent.
"""
from typing import Any, Literal
from pydantic import BaseModel, Field
class MentalModelObservation(BaseModel):
"""An observation within a mental model 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")
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-specific parameters
model_id: str | None = Field(default=None, description="Mental model ID for get_mental_model")
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)"
)
# Plain text answer fields (for output_mode=answer)
answer: str | None = Field(default=None, description="Plain text answer for done action (no markdown)")
answer_memory_ids: list[str] | None = Field(
default=None, description="Memory IDs supporting the answer", alias="memory_ids"
)
answer_model_ids: list[str] | None = Field(
default=None, description="Mental model IDs supporting the answer", alias="model_ids"
)
reasoning: str | None = Field(default=None, description="Brief reasoning for this action")
class ReflectActionBatch(BaseModel):
"""Batch of actions for parallel execution."""
actions: list[ReflectAction] = Field(description="List of actions to execute in parallel")
class ToolCall(BaseModel):
"""A single tool call made during reflect."""
tool: str = Field(description="Tool name: lookup, recall, learn, expand")
input: dict = Field(description="Tool input parameters")
output: dict = Field(description="Tool output/result")
duration_ms: int = Field(description="Execution time in milliseconds")
iteration: int = Field(default=0, description="Iteration number (1-based) when this tool was called")
class LLMCall(BaseModel):
"""A single LLM call made during reflect."""
scope: str = Field(description="Call scope: agent_1, agent_2, final, etc.")
duration_ms: int = Field(description="Execution time in milliseconds")
class Observation(BaseModel):
"""A single observation with supporting memories."""
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")
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")
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")
@@ -0,0 +1,312 @@
"""
System prompts for the reflect agent.
"""
import json
from typing import Any
def build_system_prompt_for_tools(
bank_profile: dict[str, Any],
context: str | None = None,
output_mode: str = "answer",
) -> str:
"""
Build the system prompt for tool-calling reflect agent.
This is a simplified prompt since tools are defined separately via the tools parameter.
Args:
bank_profile: Bank profile with name and mission
context: Optional additional context
output_mode: "answer" for plain text response, "observations" for structured observations
"""
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 = [
"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",
]
# Mode-specific workflow and output format
if output_mode == "observations":
# Observations mode: for mental model generation - no mental model lookup tools
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",
"",
"## Output Format: MULTIPLE Structured Observations",
"",
"CRITICAL: You MUST create MULTIPLE separate observations in the array - one for each theme.",
"Do NOT put all content in a single observation!",
"",
"- 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)",
"",
"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",
"",
"## 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')",
"",
"## 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.append("")
parts.append(f"## Memory Bank: {name}")
if mission:
parts.append(f"Mission: {mission}")
# Disposition traits
disposition = bank_profile.get("disposition", {})
if disposition:
traits = []
if "skepticism" in disposition:
traits.append(f"skepticism={disposition['skepticism']}")
if "literalism" in disposition:
traits.append(f"literalism={disposition['literalism']}")
if "empathy" in disposition:
traits.append(f"empathy={disposition['empathy']}")
if traits:
parts.append(f"Disposition: {', '.join(traits)}")
if context:
parts.append(f"\n## Additional Context\n{context}")
return "\n".join(parts)
def build_agent_prompt(
query: str,
context_history: list[dict],
bank_profile: dict,
additional_context: str | None = None,
) -> str:
"""Build the user prompt for the reflect agent."""
parts = []
# Bank identity
name = bank_profile.get("name", "Assistant")
mission = bank_profile.get("mission", "")
parts.append(f"## Memory Bank Context\nName: {name}")
if mission:
parts.append(f"Mission: {mission}")
# Disposition traits if present
disposition = bank_profile.get("disposition", {})
if disposition:
traits = []
if "skepticism" in disposition:
traits.append(f"skepticism={disposition['skepticism']}")
if "literalism" in disposition:
traits.append(f"literalism={disposition['literalism']}")
if "empathy" in disposition:
traits.append(f"empathy={disposition['empathy']}")
if traits:
parts.append(f"Disposition: {', '.join(traits)}")
# Additional context from caller
if additional_context:
parts.append(f"\n## Additional Context\n{additional_context}")
# Tool call history
if context_history:
parts.append("\n## Tool Results (synthesize and reason from this data)")
for i, entry in enumerate(context_history, 1):
tool = entry["tool"]
output = entry["output"]
# Format as proper JSON for LLM readability
try:
output_str = json.dumps(output, indent=2, default=str)
except (TypeError, ValueError):
output_str = str(output)
parts.append(f"\n### Call {i}: {tool}\n```json\n{output_str}\n```")
# The question
parts.append(f"\n## Question\n{query}")
# Instructions
if context_history:
parts.append(
"\n## Instructions\n"
"Based on the tool results above, either call more tools or provide your final answer. "
"Synthesize and reason from the data - make reasonable inferences when helpful. "
"If you have related information, use it to give the best possible answer."
)
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."
)
return "\n".join(parts)
def build_final_prompt(
query: str,
context_history: list[dict],
bank_profile: dict,
additional_context: str | None = None,
) -> str:
"""Build the final prompt when forcing a text response (no tools)."""
parts = []
# Bank identity
name = bank_profile.get("name", "Assistant")
mission = bank_profile.get("mission", "")
parts.append(f"## Memory Bank Context\nName: {name}")
if mission:
parts.append(f"Mission: {mission}")
# Disposition traits if present
disposition = bank_profile.get("disposition", {})
if disposition:
traits = []
if "skepticism" in disposition:
traits.append(f"skepticism={disposition['skepticism']}")
if "literalism" in disposition:
traits.append(f"literalism={disposition['literalism']}")
if "empathy" in disposition:
traits.append(f"empathy={disposition['empathy']}")
if traits:
parts.append(f"Disposition: {', '.join(traits)}")
# Additional context from caller
if additional_context:
parts.append(f"\n## Additional Context\n{additional_context}")
# Tool call history
if context_history:
parts.append("\n## Retrieved Data (synthesize and reason from this data)")
for entry in context_history:
tool = entry["tool"]
output = entry["output"]
# Format as proper JSON for LLM readability
try:
output_str = json.dumps(output, indent=2, default=str)
except (TypeError, ValueError):
output_str = str(output)
parts.append(f"\n### From {tool}:\n```json\n{output_str}\n```")
else:
parts.append("\n## Retrieved Data\nNo data was retrieved.")
# The question
parts.append(f"\n## Question\n{query}")
# Final instructions
parts.append(
"\n## Instructions\n"
"Provide a thoughtful answer by synthesizing and reasoning from the retrieved data above. "
"You can make reasonable inferences from the memories, but don't completely fabricate information."
"If the exact answer isn't stated, use what IS stated to give the best possible answer. "
"Only say 'I don't have information' if the retrieved data is truly unrelated to the question."
)
return "\n".join(parts)
FINAL_SYSTEM_PROMPT = """You are a thoughtful assistant that synthesizes answers from retrieved memories.
Your approach:
- Reason over the retrieved memories to answer the question
- Make reasonable inferences when the exact answer isn't explicitly stated
- Connect related memories to form a complete picture
- Be helpful - if you have related information, use it to give the best possible answer
Only say "I don't have information" if the retrieved data is truly unrelated to the question.
Do NOT fabricate information that has no basis in the retrieved data."""
@@ -0,0 +1,425 @@
"""
Tool implementations for the reflect agent.
"""
import logging
import re
import uuid
from typing import TYPE_CHECKING, Any
from .models import MentalModelInput
if TYPE_CHECKING:
from asyncpg import Connection
from ...api.http import RequestContext
from ..memory_engine import MemoryEngine
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]
async def tool_lookup(
conn: "Connection",
bank_id: str,
model_id: str | None = None,
tags: list[str] | None = None,
tags_match: str = "any",
) -> dict[str, Any]:
"""
List or get mental models.
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)
tags_match: How to match tags - "any" (OR), "all" (AND)
Returns:
Dict with either a list of models or a single model's details
"""
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
obs_data = json.loads(obs_data)
observations_raw = obs_data.get("observations", []) if isinstance(obs_data, dict) else obs_data
# 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,
)
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
""",
bank_id,
)
return {
"count": len(rows),
"models": [
{
"id": row["id"],
"subtype": row["subtype"],
"name": row["name"],
"description": row["description"],
}
for row in rows
],
}
async def tool_recall(
memory_engine: "MemoryEngine",
bank_id: str,
query: str,
request_context: "RequestContext",
max_tokens: int = 2048,
max_results: int = 50,
tags: list[str] | None = None,
tags_match: str = "any",
connection_budget: int = 1,
) -> dict[str, Any]:
"""
Search memories using TEMPR retrieval.
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 2048)
max_results: Maximum number of results
tags: Filter by tags (includes untagged memories)
tags_match: How to match tags - "any" (OR), "all" (AND), or "exact"
connection_budget: Max DB connections for this recall (default 1 for internal ops)
Returns:
Dict with list of matching memories
"""
result = await memory_engine.recall_async(
bank_id=bank_id,
query=query,
fact_type=["experience", "world"], # Exclude opinions
max_tokens=max_tokens,
enable_trace=False,
request_context=request_context,
tags=tags,
tags_match=tags_match,
_connection_budget=connection_budget,
)
memories = []
for m in result.results[:max_results]:
memories.append(
{
"id": str(m.id),
"text": m.text,
"type": m.fact_type,
"entities": m.entities or [],
"occurred": m.occurred_start, # Already ISO format string
}
)
return {
"query": query,
"count": len(memories),
"memories": memories,
}
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,
memory_ids: list[str],
depth: str,
) -> dict[str, Any]:
"""
Expand multiple memories to get chunk or document context.
Args:
conn: Database connection
bank_id: Bank identifier
memory_ids: List of memory unit IDs
depth: "chunk" or "document"
Returns:
Dict with results array, each containing memory, chunk, and optionally document data
"""
if not memory_ids:
return {"error": "memory_ids is required and must not be empty"}
# Validate and convert UUIDs
valid_uuids: list[uuid.UUID] = []
errors: dict[str, str] = {}
for mid in memory_ids:
try:
valid_uuids.append(uuid.UUID(mid))
except ValueError:
errors[mid] = f"Invalid memory_id format: {mid}"
if not valid_uuids:
return {"error": "No valid memory IDs provided", "details": errors}
# Batch fetch all memory units
memories = await conn.fetch(
"""
SELECT id, text, chunk_id, document_id, fact_type, context
FROM memory_units
WHERE id = ANY($1) AND bank_id = $2
""",
valid_uuids,
bank_id,
)
memory_map = {row["id"]: row for row in memories}
# Collect chunk_ids and document_ids for batch fetching
chunk_ids = [m["chunk_id"] for m in memories if m["chunk_id"]]
doc_ids_from_chunks: set[str] = set()
doc_ids_direct: set[str] = set()
# Batch fetch all chunks
chunk_map: dict[str, Any] = {}
if chunk_ids:
chunks = await conn.fetch(
"""
SELECT chunk_id, chunk_text, chunk_index, document_id
FROM chunks
WHERE chunk_id = ANY($1)
""",
chunk_ids,
)
chunk_map = {row["chunk_id"]: row for row in chunks}
if depth == "document":
doc_ids_from_chunks = {c["document_id"] for c in chunks if c["document_id"]}
# Collect direct document IDs (memories without chunks)
if depth == "document":
for m in memories:
if not m["chunk_id"] and m["document_id"]:
doc_ids_direct.add(m["document_id"])
# Batch fetch all documents
doc_map: dict[str, Any] = {}
all_doc_ids = list(doc_ids_from_chunks | doc_ids_direct)
if all_doc_ids:
docs = await conn.fetch(
"""
SELECT id, original_text, metadata, retain_params
FROM documents
WHERE id = ANY($1) AND bank_id = $2
""",
all_doc_ids,
bank_id,
)
doc_map = {row["id"]: row for row in docs}
# Build results
results: list[dict[str, Any]] = []
for mid, mem_uuid in zip(memory_ids, valid_uuids):
if mid in errors:
results.append({"memory_id": mid, "error": errors[mid]})
continue
memory = memory_map.get(mem_uuid)
if not memory:
results.append({"memory_id": mid, "error": f"Memory not found: {mid}"})
continue
item: dict[str, Any] = {
"memory_id": mid,
"memory": {
"id": str(memory["id"]),
"text": memory["text"],
"type": memory["fact_type"],
"context": memory["context"],
},
}
# Add chunk if available
if memory["chunk_id"] and memory["chunk_id"] in chunk_map:
chunk = chunk_map[memory["chunk_id"]]
item["chunk"] = {
"id": chunk["chunk_id"],
"text": chunk["chunk_text"],
"index": chunk["chunk_index"],
"document_id": chunk["document_id"],
}
# Add document if depth=document
if depth == "document" and chunk["document_id"] in doc_map:
doc = doc_map[chunk["document_id"]]
item["document"] = {
"id": doc["id"],
"full_text": doc["original_text"],
"metadata": doc["metadata"],
"retain_params": doc["retain_params"],
}
elif memory["document_id"] and depth == "document" and memory["document_id"] in doc_map:
# No chunk, but has document_id
doc = doc_map[memory["document_id"]]
item["document"] = {
"id": doc["id"],
"full_text": doc["original_text"],
"metadata": doc["metadata"],
"retain_params": doc["retain_params"],
}
results.append(item)
return {"results": results, "count": len(results)}
@@ -0,0 +1,212 @@
"""
Tool schema definitions for the reflect agent.
These are OpenAI-format tool definitions used with native tool calling.
"""
from typing import Literal
# Tool definitions in OpenAI format
TOOL_LIST_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.",
"parameters": {
"type": "object",
"properties": {},
"required": [],
},
},
}
TOOL_GET_MENTAL_MODEL = {
"type": "function",
"function": {
"name": "get_mental_model",
"description": "Get full details of a specific mental model including all observations and memory references.",
"parameters": {
"type": "object",
"properties": {
"model_id": {
"type": "string",
"description": "ID of the mental model (from list_mental_models results)",
},
},
"required": ["model_id"],
},
},
}
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.",
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "Search query string",
},
"max_tokens": {
"type": "integer",
"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"],
},
},
}
TOOL_EXPAND = {
"type": "function",
"function": {
"name": "expand",
"description": "Get more context for one or more memories. Memory hierarchy: memory -> chunk -> document.",
"parameters": {
"type": "object",
"properties": {
"memory_ids": {
"type": "array",
"items": {"type": "string"},
"description": "Array of memory IDs from recall results (batch multiple for efficiency)",
},
"depth": {
"type": "string",
"enum": ["chunk", "document"],
"description": "chunk: surrounding text chunk, document: full source document",
},
},
"required": ["memory_ids", "depth"],
},
},
}
TOOL_DONE_ANSWER = {
"type": "function",
"function": {
"name": "done",
"description": "Signal completion with your final answer. Use this when you have gathered enough information to answer the question.",
"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)",
},
"model_ids": {
"type": "array",
"items": {"type": "string"},
"description": "Array of mental model 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"],
},
"description": "Array of 3-8 observations, each covering a DIFFERENT aspect/theme. Do NOT put everything in one observation.",
},
},
"required": ["observations"],
},
},
}
def get_reflect_tools(
enable_learn: bool = True, output_mode: Literal["answer", "observations"] = "answer"
) -> list[dict]:
"""
Get the list of tools for the reflect agent.
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.
Returns:
List of tool definitions in OpenAI format
"""
tools = []
# 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)
else:
tools.append(TOOL_DONE_ANSWER)
return tools
@@ -6,12 +6,87 @@ API response models should be kept separate and convert from these core models t
API stability even if internal models change.
"""
from typing import Optional, List, Dict, Any
from pydantic import BaseModel, Field, ConfigDict
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 'observation' which is internal)
VALID_RECALL_FACT_TYPES = frozenset(["world", "experience", "opinion"])
class LLMToolCall(BaseModel):
"""A tool call requested by the LLM."""
id: str = Field(description="Unique identifier for this tool call")
name: str = Field(description="Name of the tool to call")
arguments: dict[str, Any] = Field(description="Arguments to pass to the tool")
class LLMToolCallResult(BaseModel):
"""Result from an LLM call that may include tool calls."""
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.")
class ToolCallTrace(BaseModel):
"""A single tool call made during reflect."""
tool: str = Field(description="Tool name: lookup, recall, learn, expand")
input: dict = Field(description="Tool input parameters")
output: dict = Field(description="Tool output/result")
duration_ms: int = Field(description="Execution time in milliseconds")
iteration: int = Field(default=0, description="Iteration number (1-based) when this tool was called")
class LLMCallTrace(BaseModel):
"""A single LLM call made during reflect."""
scope: str = Field(description="Call scope: agent_1, agent_2, final, etc.")
duration_ms: int = Field(description="Execution time in milliseconds")
class MentalModelRef(BaseModel):
"""Reference to a mental model 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")
description: str = Field(description="Brief description")
summary: str | None = Field(default=None, description="Full summary (when looked up in detail)")
class TokenUsage(BaseModel):
"""
Token usage metrics for LLM calls.
Tracks input/output tokens for a single request to enable
per-request cost tracking and monitoring.
"""
model_config = ConfigDict(
json_schema_extra={
"example": {
"input_tokens": 1500,
"output_tokens": 500,
"total_tokens": 2000,
}
}
)
input_tokens: int = Field(default=0, description="Number of input/prompt tokens consumed")
output_tokens: int = Field(default=0, description="Number of output/completion tokens generated")
total_tokens: int = Field(default=0, description="Total tokens (input + output)")
def __add__(self, other: "TokenUsage") -> "TokenUsage":
"""Allow aggregating token usage from multiple calls."""
return TokenUsage(
input_tokens=self.input_tokens + other.input_tokens,
output_tokens=self.output_tokens + other.output_tokens,
total_tokens=self.total_tokens + other.total_tokens,
)
class DispositionTraits(BaseModel):
@@ -23,17 +98,12 @@ class DispositionTraits(BaseModel):
- literalism: 1=flexible interpretation, 5=literal interpretation (how strictly to interpret information)
- empathy: 1=detached, 5=empathetic (how much to consider emotional context)
"""
skepticism: int = Field(ge=1, le=5, description="How skeptical vs trusting (1=trusting, 5=skeptical)")
literalism: int = Field(ge=1, le=5, description="How literally to interpret information (1=flexible, 5=literal)")
empathy: int = Field(ge=1, le=5, description="How much to consider emotional context (1=detached, 5=empathetic)")
model_config = ConfigDict(json_schema_extra={
"example": {
"skepticism": 3,
"literalism": 3,
"empathy": 3
}
})
model_config = ConfigDict(json_schema_extra={"example": {"skepticism": 3, "literalism": 3, "empathy": 3}})
class MemoryFact(BaseModel):
@@ -43,38 +113,46 @@ class MemoryFact(BaseModel):
This represents a unit of information stored in the memory system,
including both the content and metadata.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Alice works at Google on the AI team",
"fact_type": "world",
"entities": ["Alice", "Google"],
"context": "work info",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
"mentioned_at": "2024-01-15T10:30:00Z",
"document_id": "session_abc123",
"metadata": {"source": "slack"},
"chunk_id": "bank123_session_abc123_0",
"activation": 0.95
model_config = ConfigDict(
json_schema_extra={
"example": {
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Alice works at Google on the AI team",
"fact_type": "world",
"entities": ["Alice", "Google"],
"context": "work info",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
"mentioned_at": "2024-01-15T10:30:00Z",
"document_id": "session_abc123",
"metadata": {"source": "slack"},
"chunk_id": "bank123_session_abc123_0",
"activation": 0.95,
"tags": ["user_a", "session_123"],
}
}
})
)
id: str = Field(description="Unique identifier for the memory fact")
text: str = Field(description="The actual text content of the memory")
fact_type: str = Field(description="Type of fact: 'world', 'experience', 'opinion', or 'observation'")
entities: Optional[List[str]] = Field(None, description="Entity names mentioned in this fact")
context: Optional[str] = Field(None, description="Additional context for the memory")
occurred_start: Optional[str] = Field(None, description="ISO format date when the event started occurring")
occurred_end: Optional[str] = Field(None, description="ISO format date when the event ended occurring")
mentioned_at: Optional[str] = Field(None, description="ISO format date when the fact was mentioned/learned")
document_id: Optional[str] = Field(None, description="ID of the document this memory belongs to")
metadata: Optional[Dict[str, str]] = Field(None, description="User-defined metadata")
chunk_id: Optional[str] = Field(None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)")
entities: list[str] | None = Field(None, description="Entity names mentioned in this fact")
context: str | None = Field(None, description="Additional context for the memory")
occurred_start: str | None = Field(None, description="ISO format date when the event started occurring")
occurred_end: str | None = Field(None, description="ISO format date when the event ended occurring")
mentioned_at: str | None = Field(None, description="ISO format date when the fact was mentioned/learned")
document_id: str | None = Field(None, description="ID of the document this memory belongs to")
metadata: dict[str, str] | None = Field(None, description="User-defined metadata")
chunk_id: str | None = Field(
None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)"
)
tags: list[str] | None = Field(None, description="Visibility scope tags associated with this fact")
class ChunkInfo(BaseModel):
"""Information about a chunk."""
chunk_text: str = Field(description="The raw chunk text")
chunk_index: int = Field(description="Index of the chunk within the document")
truncated: bool = Field(default=False, description="Whether the chunk was truncated due to token limits")
@@ -87,35 +165,33 @@ class RecallResult(BaseModel):
Contains a list of matching memory facts and optional trace information
for debugging and transparency.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"results": [
{
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Alice works at Google on the AI team",
"fact_type": "world",
"context": "work info",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
"activation": 0.95
}
],
"trace": {
"query": "What did Alice say about machine learning?",
"num_results": 1
model_config = ConfigDict(
json_schema_extra={
"example": {
"results": [
{
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Alice works at Google on the AI team",
"fact_type": "world",
"context": "work info",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
"activation": 0.95,
}
],
"trace": {"query": "What did Alice say about machine learning?", "num_results": 1},
}
}
})
results: List[MemoryFact] = Field(description="List of memory facts matching the query")
trace: Optional[Dict[str, Any]] = Field(None, description="Trace information for debugging")
entities: Optional[Dict[str, "EntityState"]] = Field(
None,
description="Entity states for entities mentioned in results (keyed by canonical name)"
)
chunks: Optional[Dict[str, ChunkInfo]] = Field(
None,
description="Chunks for facts, keyed by '{document_id}_{chunk_index}'"
results: list[MemoryFact] = Field(description="List of memory facts matching the query")
trace: dict[str, Any] | None = Field(None, description="Trace information for debugging")
entities: dict[str, "EntityState"] | None = Field(
None, description="Entity states for entities mentioned in results (keyed by canonical name)"
)
chunks: dict[str, ChunkInfo] | None = Field(
None, description="Chunks for facts, keyed by '{document_id}_{chunk_index}'"
)
@@ -124,38 +200,59 @@ class ReflectResult(BaseModel):
Result from a reflect operation.
Contains the formulated answer, the facts it was based on (organized by type),
and any new opinions that were formed during the reflection process.
any new opinions that were formed during the reflection process, and optionally
structured output if a response schema was provided.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"text": "Based on my knowledge, machine learning is being actively used in healthcare...",
"based_on": {
"world": [
{
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Machine learning is used in medical diagnosis",
"fact_type": "world",
"context": "healthcare",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z"
}
],
"experience": [],
"opinion": []
},
"new_opinions": [
"Machine learning has great potential in healthcare"
]
model_config = ConfigDict(
json_schema_extra={
"example": {
"text": "Based on my knowledge, machine learning is being actively used in healthcare...",
"based_on": {
"world": [
{
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Machine learning is used in medical diagnosis",
"fact_type": "world",
"context": "healthcare",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
}
],
"experience": [],
"opinion": [],
},
"new_opinions": ["Machine learning has great potential in healthcare"],
"structured_output": {"summary": "ML in healthcare", "confidence": 0.9},
"usage": {"input_tokens": 1500, "output_tokens": 500, "total_tokens": 2000},
}
}
})
)
text: str = Field(description="The formulated answer text")
based_on: Dict[str, List[MemoryFact]] = Field(
based_on: dict[str, list[MemoryFact]] = Field(
description="Facts used to formulate the answer, organized by type (world, experience, opinion)"
)
new_opinions: List[str] = Field(
new_opinions: list[str] = Field(default_factory=list, description="List of newly formed opinions during reflection")
structured_output: dict[str, Any] | None = Field(
default=None,
description="Structured output parsed according to the provided response schema. Only present when response_schema was provided.",
)
usage: TokenUsage | None = Field(
default=None,
description="Token usage metrics for the LLM calls made during this reflect operation.",
)
tool_trace: list[ToolCallTrace] = Field(
default_factory=list,
description="List of newly formed opinions during reflection"
description="Trace of tool calls made during reflection. Only present when include.tool_calls is enabled.",
)
llm_trace: list[LLMCallTrace] = Field(
default_factory=list,
description="Trace of LLM calls made during reflection. Only present when include.tool_calls is enabled.",
)
mental_models: list[MentalModelRef] = Field(
default_factory=list,
description="Mental models accessed during reflection. Only present when include.facts is enabled.",
)
@@ -166,12 +263,12 @@ class Opinion(BaseModel):
Opinions represent the bank's formed perspectives on topics,
with a confidence level indicating strength of belief.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"text": "Machine learning has great potential in healthcare",
"confidence": 0.85
model_config = ConfigDict(
json_schema_extra={
"example": {"text": "Machine learning has great potential in healthcare", "confidence": 0.85}
}
})
)
text: str = Field(description="The opinion text")
confidence: float = Field(description="Confidence score between 0.0 and 1.0")
@@ -184,15 +281,15 @@ class EntityObservation(BaseModel):
Observations are objective facts synthesized from multiple memory facts
about an entity, without personality influence.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"text": "John is detail-oriented and works at Google",
"mentioned_at": "2024-01-15T10:30:00Z"
model_config = ConfigDict(
json_schema_extra={
"example": {"text": "John is detail-oriented and works at Google", "mentioned_at": "2024-01-15T10:30:00Z"}
}
})
)
text: str = Field(description="The observation text")
mentioned_at: Optional[str] = Field(None, description="ISO format date when this observation was created")
mentioned_at: str | None = Field(None, description="ISO format date when this observation was created")
class EntityState(BaseModel):
@@ -201,20 +298,51 @@ class EntityState(BaseModel):
Contains observations synthesized from facts about the entity.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"entity_id": "123e4567-e89b-12d3-a456-426614174000",
"canonical_name": "John",
"observations": [
{"text": "John is detail-oriented", "mentioned_at": "2024-01-15T10:30:00Z"},
{"text": "John works at Google on the AI team", "mentioned_at": "2024-01-14T09:00:00Z"}
]
model_config = ConfigDict(
json_schema_extra={
"example": {
"entity_id": "123e4567-e89b-12d3-a456-426614174000",
"canonical_name": "John",
"observations": [
{"text": "John is detail-oriented", "mentioned_at": "2024-01-15T10:30:00Z"},
{"text": "John works at Google on the AI team", "mentioned_at": "2024-01-14T09:00:00Z"},
],
}
}
})
)
entity_id: str = Field(description="Unique identifier for the entity")
canonical_name: str = Field(description="Canonical name of the entity")
observations: List[EntityObservation] = Field(
default_factory=list,
description="List of observations about this entity"
observations: list[EntityObservation] = Field(
default_factory=list, description="List of observations about this entity"
)
class MentalModel(BaseModel):
"""
A manually configured mental model for tracking specific topics/areas.
Mental models are user-defined focus areas that the agent should track
and maintain summaries for, unlike auto-extracted entities.
"""
model_config = ConfigDict(
json_schema_extra={
"example": {
"id": "team-dynamics",
"name": "Team Dynamics",
"description": "Track how the team collaborates, communication patterns, conflicts, and resolutions",
"summary": "The team has strong collaboration...",
"summary_updated_at": "2024-01-15T10:30:00Z",
"created_at": "2024-01-10T08:00:00Z",
}
}
)
id: str = Field(description="Unique identifier (alphanumeric lowercase)")
name: str = Field(description="Display name for the mental model")
description: str = Field(description="Prompt/directions for what to track and summarize")
summary: str | None = Field(None, description="Generated summary based on relevant facts")
summary_updated_at: str | None = Field(None, description="ISO format date when summary was last updated")
created_at: str = Field(description="ISO format date when the mental model was created")
@@ -12,23 +12,16 @@ This package contains modular components for the retain operation:
- fact_storage: Handle fact insertion into database
"""
from .types import (
RetainContent,
ExtractedFact,
ProcessedFact,
ChunkMetadata,
EntityRef,
CausalRelation,
RetainBatch
from . import (
chunk_storage,
deduplication,
embedding_processing,
entity_processing,
fact_extraction,
fact_storage,
link_creation,
)
from . import fact_extraction
from . import embedding_processing
from . import deduplication
from . import entity_processing
from . import link_creation
from . import chunk_storage
from . import fact_storage
from .types import CausalRelation, ChunkMetadata, EntityRef, ExtractedFact, ProcessedFact, RetainBatch, RetainContent
__all__ = [
# Types
@@ -1,13 +1,16 @@
"""
bank profile utilities for disposition and background management.
bank profile utilities for disposition and mission management.
"""
import json
import logging
import re
from typing import Dict, Optional, TypedDict
from typing import TypedDict
from pydantic import BaseModel, Field
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from ..response_models import DispositionTraits
logger = logging.getLogger(__name__)
@@ -21,20 +24,21 @@ DEFAULT_DISPOSITION = {
class BankProfile(TypedDict):
"""Type for bank profile data."""
name: str
disposition: DispositionTraits
background: str
mission: str
class BackgroundMergeResponse(BaseModel):
"""LLM response for background merge with disposition inference."""
background: str = Field(description="Merged background in first person perspective")
disposition: DispositionTraits = Field(description="Inferred disposition traits (skepticism, literalism, empathy)")
class MissionMergeResponse(BaseModel):
"""LLM response for mission merge."""
mission: str = Field(description="Merged mission in first person perspective")
async def get_bank_profile(pool, bank_id: str) -> BankProfile:
"""
Get bank profile (name, disposition + background).
Get bank profile (name, disposition + mission).
Auto-creates bank with default values if not exists.
Args:
@@ -42,16 +46,16 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
bank_id: bank IDentifier
Returns:
BankProfile with name, typed DispositionTraits, and background
BankProfile with name, typed DispositionTraits, and mission
"""
async with acquire_with_retry(pool) as conn:
# Try to get existing bank
row = await conn.fetchrow(
"""
SELECT name, disposition, background
FROM banks WHERE bank_id = $1
f"""
SELECT name, disposition, mission
FROM {fq_table("banks")} WHERE bank_id = $1
""",
bank_id
bank_id,
)
if row:
@@ -63,34 +67,26 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
return BankProfile(
name=row["name"],
disposition=DispositionTraits(**disposition_data),
background=row["background"]
mission=row["mission"] or "",
)
# Bank doesn't exist, create with defaults
await conn.execute(
"""
INSERT INTO banks (bank_id, name, disposition, background)
f"""
INSERT INTO {fq_table("banks")} (bank_id, name, disposition, mission)
VALUES ($1, $2, $3::jsonb, $4)
ON CONFLICT (bank_id) DO NOTHING
""",
bank_id,
bank_id, # Default name is the bank_id
json.dumps(DEFAULT_DISPOSITION),
""
"",
)
return BankProfile(
name=bank_id,
disposition=DispositionTraits(**DEFAULT_DISPOSITION),
background=""
)
return BankProfile(name=bank_id, disposition=DispositionTraits(**DEFAULT_DISPOSITION), mission="")
async def update_bank_disposition(
pool,
bank_id: str,
disposition: Dict[str, int]
) -> None:
async def update_bank_disposition(pool, bank_id: str, disposition: dict[str, int]) -> None:
"""
Update bank disposition traits.
@@ -104,275 +100,132 @@ async def update_bank_disposition(
async with acquire_with_retry(pool) as conn:
await conn.execute(
"""
UPDATE banks
f"""
UPDATE {fq_table("banks")}
SET disposition = $2::jsonb,
updated_at = NOW()
WHERE bank_id = $1
""",
bank_id,
json.dumps(disposition)
json.dumps(disposition),
)
async def merge_bank_background(
pool,
llm_config,
bank_id: str,
new_info: str,
update_disposition: bool = True
) -> dict:
async def set_bank_mission(pool, bank_id: str, mission: str) -> None:
"""
Merge new background information with existing background using LLM.
Normalizes to first person ("I") and resolves conflicts.
Optionally infers disposition traits from the merged background.
Set bank mission (replacing any existing mission).
Args:
pool: Database connection pool
llm_config: LLM configuration for background merging
bank_id: bank IDentifier
new_info: New background information to add/merge
update_disposition: If True, infer Big Five traits from background (default: True)
mission: The mission text
"""
# Ensure bank exists first
await get_bank_profile(pool, bank_id)
async with acquire_with_retry(pool) as conn:
await conn.execute(
f"""
UPDATE {fq_table("banks")}
SET mission = $2,
updated_at = NOW()
WHERE bank_id = $1
""",
bank_id,
mission,
)
async def merge_bank_mission(pool, llm_config, bank_id: str, new_info: str) -> dict:
"""
Merge new mission information with existing mission using LLM.
Normalizes to first person ("I") and resolves conflicts.
Args:
pool: Database connection pool
llm_config: LLM configuration for mission merging
bank_id: bank IDentifier
new_info: New mission information to add/merge
Returns:
Dict with 'background' (str) and optionally 'disposition' (dict) keys
Dict with 'mission' (str) key
"""
# Get current profile
profile = await get_bank_profile(pool, bank_id)
current_background = profile["background"]
current_mission = profile["mission"]
# Use LLM to merge backgrounds and optionally infer disposition
result = await _llm_merge_background(
llm_config,
current_background,
new_info,
infer_disposition=update_disposition
)
# Use LLM to merge missions
result = await _llm_merge_mission(llm_config, current_mission, new_info)
merged_background = result["background"]
inferred_disposition = result.get("disposition")
merged_mission = result["mission"]
# Update in database
async with acquire_with_retry(pool) as conn:
if inferred_disposition:
# Update both background and disposition
await conn.execute(
"""
UPDATE banks
SET background = $2,
disposition = $3::jsonb,
updated_at = NOW()
WHERE bank_id = $1
""",
bank_id,
merged_background,
json.dumps(inferred_disposition)
)
else:
# Update only background
await conn.execute(
"""
UPDATE banks
SET background = $2,
updated_at = NOW()
WHERE bank_id = $1
""",
bank_id,
merged_background
)
await conn.execute(
f"""
UPDATE {fq_table("banks")}
SET mission = $2,
updated_at = NOW()
WHERE bank_id = $1
""",
bank_id,
merged_mission,
)
response = {"background": merged_background}
if inferred_disposition:
response["disposition"] = inferred_disposition
return response
return {"mission": merged_mission}
async def _llm_merge_background(
llm_config,
current: str,
new_info: str,
infer_disposition: bool = False
) -> dict:
async def _llm_merge_mission(llm_config, current: str, new_info: str) -> dict:
"""
Use LLM to intelligently merge background information.
Optionally infer Big Five disposition traits from the merged background.
Use LLM to intelligently merge mission information.
Args:
llm_config: LLM configuration to use
current: Current background text
current: Current mission text
new_info: New information to merge
infer_disposition: If True, also infer disposition traits
Returns:
Dict with 'background' (str) and optionally 'disposition' (dict) keys
Dict with 'mission' (str) key
"""
if infer_disposition:
prompt = f"""You are helping maintain a memory bank's background/profile and infer their disposition. You MUST respond with ONLY valid JSON.
prompt = f"""You are helping maintain an agent's mission statement.
Current background: {current if current else "(empty)"}
Current mission: {current if current else "(empty)"}
New information to add: {new_info}
Instructions:
1. Merge the new information with the current background
2. If there are conflicts (e.g., different birthplaces), the NEW information overwrites the old
3. Keep additions that don't conflict
4. Output in FIRST PERSON ("I") perspective
5. Be concise - keep merged background under 500 characters
6. Infer disposition traits from the merged background (each 1-5 integer):
- Skepticism: 1-5 (1=trusting, takes things at face value; 5=skeptical, questions everything)
- Literalism: 1-5 (1=flexible interpretation, reads between lines; 5=literal, exact interpretation)
- Empathy: 1-5 (1=detached, focuses on facts; 5=empathetic, considers emotional context)
CRITICAL: You MUST respond with ONLY a valid JSON object. No markdown, no code blocks, no explanations. Just the JSON.
Format:
{{
"background": "the merged background text in first person",
"disposition": {{
"skepticism": 3,
"literalism": 3,
"empathy": 3
}}
}}
Trait inference examples:
- "I'm a lawyer" → skepticism: 4, literalism: 5, empathy: 2
- "I'm a therapist" → skepticism: 2, literalism: 2, empathy: 5
- "I'm an engineer" → skepticism: 3, literalism: 4, empathy: 3
- "I've been burned before by trusting people" → skepticism: 5, literalism: 3, empathy: 3
- "I try to understand what people really mean" → skepticism: 3, literalism: 2, empathy: 4
- "I take contracts very seriously" → skepticism: 4, literalism: 5, empathy: 2"""
else:
prompt = f"""You are helping maintain a memory bank's background/profile.
Current background: {current if current else "(empty)"}
New information to add: {new_info}
Instructions:
1. Merge the new information with the current background
2. If there are conflicts (e.g., different birthplaces), the NEW information overwrites the old
1. Merge the new information with the current mission
2. If there are conflicts, the NEW information overwrites the old
3. Keep additions that don't conflict
4. Output in FIRST PERSON ("I") perspective
5. Be concise - keep it under 500 characters
6. Return ONLY the merged background text, no explanations
6. Return ONLY the merged mission text, no explanations
Merged background:"""
Merged mission:"""
try:
# Prepare messages
messages = [{"role": "user", "content": prompt}]
if infer_disposition:
# Use structured output with Pydantic model for disposition inference
try:
parsed = await llm_config.call(
messages=messages,
response_format=BackgroundMergeResponse,
scope="bank_background",
temperature=0.3,
max_tokens=8192
)
logger.info(f"Successfully got structured response: background={parsed.background[:100]}")
# Convert Pydantic model to dict format
return {
"background": parsed.background,
"disposition": parsed.disposition.model_dump()
}
except Exception as e:
logger.warning(f"Structured output failed, falling back to manual parsing: {e}")
# Fall through to manual parsing below
# Manual parsing fallback or non-disposition merge
content = await llm_config.call(
messages=messages,
scope="bank_background",
temperature=0.3,
max_tokens=8192
messages=messages, scope="bank_mission", temperature=0.3, max_completion_tokens=8192
)
logger.info(f"LLM response for background merge (first 500 chars): {content[:500]}")
logger.info(f"LLM response for mission merge (first 500 chars): {content[:500]}")
if infer_disposition:
# Parse JSON response - try multiple extraction methods
result = None
# Method 1: Direct parse
try:
result = json.loads(content)
logger.info("Successfully parsed JSON directly")
except json.JSONDecodeError:
pass
# Method 2: Extract from markdown code blocks
if result is None:
# Remove markdown code blocks
code_block_match = re.search(r'```(?:json)?\s*(\{.*?\})\s*```', content, re.DOTALL)
if code_block_match:
try:
result = json.loads(code_block_match.group(1))
logger.info("Successfully extracted JSON from markdown code block")
except json.JSONDecodeError:
pass
# Method 3: Find nested JSON structure
if result is None:
# Look for JSON object with nested structure
json_match = re.search(r'\{[^{}]*"background"[^{}]*"disposition"[^{}]*\{[^{}]*\}[^{}]*\}', content, re.DOTALL)
if json_match:
try:
result = json.loads(json_match.group())
logger.info("Successfully extracted JSON using nested pattern")
except json.JSONDecodeError:
pass
# All parsing methods failed - use fallback
if result is None:
logger.warning(f"Failed to extract JSON from LLM response. Raw content: {content[:200]}")
# Fallback: use new_info as background with default disposition
return {
"background": new_info if new_info else current if current else "",
"disposition": DEFAULT_DISPOSITION.copy()
}
# Validate disposition values
disposition = result.get("disposition", {})
for key in ["skepticism", "literalism", "empathy"]:
if key not in disposition:
disposition[key] = 3 # Default to neutral
else:
# Clamp to [1, 5] and convert to int
disposition[key] = max(1, min(5, int(disposition[key])))
result["disposition"] = disposition
# Ensure background exists
if "background" not in result or not result["background"]:
result["background"] = new_info if new_info else ""
return result
else:
# Just background merge
merged = content
if not merged or merged.lower() in ["(empty)", "none", "n/a"]:
merged = new_info if new_info else ""
return {"background": merged}
merged = content.strip()
if not merged or merged.lower() in ["(empty)", "none", "n/a"]:
merged = new_info if new_info else ""
return {"mission": merged}
except Exception as e:
logger.error(f"Error merging background with LLM: {e}")
logger.error(f"Error merging mission with LLM: {e}")
# Fallback: just append new info
if current:
merged = f"{current} {new_info}".strip()
else:
merged = new_info
result = {"background": merged}
if infer_disposition:
result["disposition"] = DEFAULT_DISPOSITION.copy()
return result
return {"mission": merged}
async def list_banks(pool) -> list:
@@ -383,13 +236,13 @@ async def list_banks(pool) -> list:
pool: Database connection pool
Returns:
List of dicts with bank_id, name, disposition, background, created_at, updated_at
List of dicts with bank_id, name, disposition, mission, created_at, updated_at
"""
async with acquire_with_retry(pool) as conn:
rows = await conn.fetch(
"""
SELECT bank_id, name, disposition, background, created_at, updated_at
FROM banks
f"""
SELECT bank_id, name, disposition, mission, created_at, updated_at
FROM {fq_table("banks")}
ORDER BY updated_at DESC
"""
)
@@ -401,13 +254,15 @@ async def list_banks(pool) -> list:
if isinstance(disposition_data, str):
disposition_data = json.loads(disposition_data)
result.append({
"bank_id": row["bank_id"],
"name": row["name"],
"disposition": disposition_data,
"background": row["background"],
"created_at": row["created_at"].isoformat() if row["created_at"] else None,
"updated_at": row["updated_at"].isoformat() if row["updated_at"] else None,
})
result.append(
{
"bank_id": row["bank_id"],
"name": row["name"],
"disposition": disposition_data,
"mission": row["mission"] or "",
"created_at": row["created_at"].isoformat() if row["created_at"] else None,
"updated_at": row["updated_at"].isoformat() if row["updated_at"] else None,
}
)
return result
@@ -3,20 +3,16 @@ Chunk storage for retain pipeline.
Handles storage of document chunks in the database.
"""
import logging
from typing import List, Dict, Optional
import logging
from ..memory_engine import fq_table
from .types import ChunkMetadata
logger = logging.getLogger(__name__)
async def store_chunks_batch(
conn,
bank_id: str,
document_id: str,
chunks: List[ChunkMetadata]
) -> Dict[int, str]:
async def store_chunks_batch(conn, bank_id: str, document_id: str, chunks: list[ChunkMetadata]) -> dict[int, str]:
"""
Store document chunks in the database.
@@ -47,24 +43,21 @@ async def store_chunks_batch(
# Batch insert all chunks
await conn.execute(
"""
INSERT INTO chunks (chunk_id, document_id, bank_id, chunk_text, chunk_index)
f"""
INSERT INTO {fq_table("chunks")} (chunk_id, document_id, bank_id, chunk_text, chunk_index)
SELECT * FROM unnest($1::text[], $2::text[], $3::text[], $4::text[], $5::integer[])
""",
chunk_ids,
[document_id] * len(chunk_texts),
[bank_id] * len(chunk_texts),
chunk_texts,
chunk_indices
chunk_indices,
)
return chunk_id_map
def map_facts_to_chunks(
facts_chunk_indices: List[int],
chunk_id_map: Dict[int, str]
) -> List[Optional[str]]:
def map_facts_to_chunks(facts_chunk_indices: list[int], chunk_id_map: dict[int, str]) -> list[str | None]:
"""
Map fact chunk indices to chunk IDs.
@@ -3,22 +3,17 @@ Deduplication logic for retain pipeline.
Checks for duplicate facts using semantic similarity and temporal proximity.
"""
import logging
from datetime import datetime
from typing import List
from collections import defaultdict
from datetime import UTC
from .types import ProcessedFact
logger = logging.getLogger(__name__)
async def check_duplicates_batch(
conn,
bank_id: str,
facts: List[ProcessedFact],
duplicate_checker_fn
) -> List[bool]:
async def check_duplicates_batch(conn, bank_id: str, facts: list[ProcessedFact], duplicate_checker_fn) -> list[bool]:
"""
Check which facts are duplicates using batched time-window queries.
@@ -47,16 +42,12 @@ async def check_duplicates_batch(
# Defensive: if both are None (shouldn't happen), use now()
if fact_date is None:
from datetime import datetime, timezone
fact_date = datetime.now(timezone.utc)
from datetime import datetime
fact_date = datetime.now(UTC)
# Round to 12-hour bucket to group similar times
bucket_key = fact_date.replace(
hour=(fact_date.hour // 12) * 12,
minute=0,
second=0,
microsecond=0
)
bucket_key = fact_date.replace(hour=(fact_date.hour // 12) * 12, minute=0, second=0, microsecond=0)
time_buckets[bucket_key].append((idx, fact))
# Process each bucket in batch
@@ -68,14 +59,7 @@ async def check_duplicates_batch(
embeddings = [item[1].embedding for item in bucket_items]
# Check duplicates for this time bucket
dup_flags = await duplicate_checker_fn(
conn,
bank_id,
texts,
embeddings,
bucket_date,
time_window_hours=24
)
dup_flags = await duplicate_checker_fn(conn, bank_id, texts, embeddings, bucket_date, time_window_hours=24)
# Map results back to original indices
for idx, is_dup in zip(indices, dup_flags):
@@ -84,10 +68,7 @@ async def check_duplicates_batch(
return all_is_duplicate
def filter_duplicates(
facts: List[ProcessedFact],
is_duplicate_flags: List[bool]
) -> List[ProcessedFact]:
def filter_duplicates(facts: list[ProcessedFact], is_duplicate_flags: list[bool]) -> list[ProcessedFact]:
"""
Filter out duplicate facts based on duplicate flags.
@@ -3,9 +3,8 @@ Embedding processing for retain pipeline.
Handles augmenting fact texts with temporal information and generating embeddings.
"""
import logging
from typing import List
from datetime import datetime
from . import embedding_utils
from .types import ExtractedFact
@@ -13,7 +12,7 @@ from .types import ExtractedFact
logger = logging.getLogger(__name__)
def augment_texts_with_dates(facts: List[ExtractedFact], format_date_fn) -> List[str]:
def augment_texts_with_dates(facts: list[ExtractedFact], format_date_fn) -> list[str]:
"""
Augment fact texts with readable dates for better temporal matching.
@@ -37,10 +36,7 @@ def augment_texts_with_dates(facts: List[ExtractedFact], format_date_fn) -> List
return augmented_texts
async def generate_embeddings_batch(
embeddings_model,
texts: List[str]
) -> List[List[float]]:
async def generate_embeddings_batch(embeddings_model, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings for a batch of texts.
@@ -54,9 +50,6 @@ async def generate_embeddings_batch(
if not texts:
return []
embeddings = await embedding_utils.generate_embeddings_batch(
embeddings_model,
texts
)
embeddings = await embedding_utils.generate_embeddings_batch(embeddings_model, texts)
return embeddings
@@ -4,12 +4,11 @@ Embedding generation utilities for memory units.
import asyncio
import logging
from typing import List
logger = logging.getLogger(__name__)
def generate_embedding(embeddings_backend, text: str) -> List[float]:
def generate_embedding(embeddings_backend, text: str) -> list[float]:
"""
Generate embedding for text using the provided embeddings backend.
@@ -27,7 +26,7 @@ def generate_embedding(embeddings_backend, text: str) -> List[float]:
raise Exception(f"Failed to generate embedding: {str(e)}")
async def generate_embeddings_batch(embeddings_backend, texts: List[str]) -> List[List[float]]:
async def generate_embeddings_batch(embeddings_backend, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings for multiple texts using the provided embeddings backend.
@@ -47,7 +46,7 @@ async def generate_embeddings_batch(embeddings_backend, texts: List[str]) -> Lis
embeddings = await loop.run_in_executor(
None, # Use default thread pool
embeddings_backend.encode,
texts
texts,
)
return embeddings
except Exception as e:
@@ -3,12 +3,11 @@ Entity processing for retain pipeline.
Handles entity extraction, resolution, and link creation for stored facts.
"""
import logging
from typing import List, Tuple, Dict, Any
from uuid import UUID
from .types import ProcessedFact, EntityRef, EntityLink
import logging
from . import link_utils
from .types import EntityLink, ProcessedFact
logger = logging.getLogger(__name__)
@@ -17,18 +16,20 @@ async def process_entities_batch(
entity_resolver,
conn,
bank_id: str,
unit_ids: List[str],
facts: List[ProcessedFact],
log_buffer: List[str] = None
) -> List[EntityLink]:
unit_ids: list[str],
facts: list[ProcessedFact],
log_buffer: list[str] = None,
user_entities_per_content: dict[int, list[dict]] = None,
) -> list[EntityLink]:
"""
Process entities for all facts and create entity links.
This function:
1. Extracts entity mentions from fact texts
2. Resolves entity names to canonical entities
3. Creates entity records in the database
4. Returns entity links ready for insertion
2. Merges user-provided entities with LLM-extracted entities
3. Resolves entity names to canonical entities
4. Creates entity records in the database
5. Returns entity links ready for insertion
Args:
entity_resolver: EntityResolver instance for entity resolution
@@ -37,6 +38,7 @@ async def process_entities_batch(
unit_ids: List of unit IDs (same length as facts)
facts: List of ProcessedFact objects
log_buffer: Optional buffer for detailed logging
user_entities_per_content: Dict mapping content_index to list of user-provided entities
Returns:
List of EntityLink objects for batch insertion
@@ -47,15 +49,35 @@ async def process_entities_batch(
if len(unit_ids) != len(facts):
raise ValueError(f"Mismatch between unit_ids ({len(unit_ids)}) and facts ({len(facts)})")
user_entities_per_content = user_entities_per_content or {}
# Extract data for link_utils function
fact_texts = [fact.fact_text for fact in facts]
# Use occurred_start if available, otherwise use mentioned_at for entity timestamps
fact_dates = [fact.occurred_start if fact.occurred_start is not None else fact.mentioned_at for fact in facts]
# Convert EntityRef objects to dict format expected by link_utils
entities_per_fact = [
[{'text': entity.name, 'type': 'CONCEPT'} for entity in (fact.entities or [])]
for fact in facts
]
# Convert EntityRef objects to dict format and merge with user-provided entities
entities_per_fact = []
for fact in facts:
# Start with LLM-extracted entities
llm_entities = [{"text": entity.name, "type": "CONCEPT"} for entity in (fact.entities or [])]
# Get user entities for this content (use content_index from fact)
user_entities = user_entities_per_content.get(fact.content_index, [])
# Merge with case-insensitive deduplication
seen_texts = {e["text"].lower() for e in llm_entities}
for user_entity in user_entities:
if user_entity["text"].lower() not in seen_texts:
llm_entities.append(
{
"text": user_entity["text"],
"type": user_entity.get("type", "CONCEPT"),
}
)
seen_texts.add(user_entity["text"].lower())
entities_per_fact.append(llm_entities)
# Use existing link_utils function for entity processing
entity_links = await link_utils.extract_entities_batch_optimized(
@@ -67,16 +89,13 @@ async def process_entities_batch(
"", # context (not used in current implementation)
fact_dates,
entities_per_fact,
log_buffer # Pass log_buffer for detailed logging
log_buffer, # Pass log_buffer for detailed logging
)
return entity_links
async def insert_entity_links_batch(
conn,
entity_links: List[EntityLink]
) -> None:
async def insert_entity_links_batch(conn, entity_links: list[EntityLink]) -> None:
"""
Insert entity links in batch.
File diff suppressed because it is too large Load Diff
@@ -3,22 +3,19 @@ Fact storage for retain pipeline.
Handles insertion of facts into the database.
"""
import logging
import json
from typing import List, Optional
from uuid import UUID
import json
import logging
from ..memory_engine import fq_table
from .types import ProcessedFact
logger = logging.getLogger(__name__)
async def insert_facts_batch(
conn,
bank_id: str,
facts: List[ProcessedFact],
document_id: Optional[str] = None
) -> List[str]:
conn, bank_id: str, facts: list[ProcessedFact], document_id: str | None = None
) -> list[str]:
"""
Insert facts into the database in batch.
@@ -48,6 +45,7 @@ async def insert_facts_batch(
metadata_jsons = []
chunk_ids = []
document_ids = []
tags_list = []
for fact in facts:
fact_texts.append(fact.fact_text)
@@ -62,22 +60,37 @@ async def insert_facts_batch(
contexts.append(fact.context)
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)
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
document_ids.append(fact.document_id if fact.document_id else document_id)
# Convert tags to JSON string for proper batch insertion (PostgreSQL unnest doesn't handle 2D arrays well)
tags_list.append(json.dumps(fact.tags if fact.tags else []))
# Batch insert all facts
# Note: tags are passed as JSON strings and converted back to varchar[] via jsonb_array_elements_text + array_agg
results = await conn.fetch(
"""
INSERT INTO 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)
SELECT $1, * 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[]
f"""
WITH input_data AS (
SELECT * FROM unnest(
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[], $15::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)
)
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)
SELECT
$1,
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id,
COALESCE(
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
'{{}}'::varchar[]
)
FROM input_data
RETURNING id
""",
bank_id,
@@ -93,10 +106,11 @@ async def insert_facts_batch(
access_counts,
metadata_jsons,
chunk_ids,
document_ids
document_ids,
tags_list,
)
unit_ids = [str(row['id']) for row in results]
unit_ids = [str(row["id"]) for row in results]
return unit_ids
@@ -111,15 +125,15 @@ async def ensure_bank_exists(conn, bank_id: str) -> None:
bank_id: Bank identifier
"""
await conn.execute(
"""
INSERT INTO banks (bank_id, disposition, background)
f"""
INSERT INTO {fq_table("banks")} (bank_id, disposition, mission)
VALUES ($1, $2::jsonb, $3)
ON CONFLICT (bank_id) DO UPDATE
SET updated_at = NOW()
""",
bank_id,
'{"skepticism": 3, "literalism": 3, "empathy": 3}',
""
"",
)
@@ -129,7 +143,8 @@ async def handle_document_tracking(
document_id: str,
combined_content: str,
is_first_batch: bool,
retain_params: Optional[dict] = None
retain_params: dict | None = None,
document_tags: list[str] | None = None,
) -> None:
"""
Handle document tracking in the database.
@@ -141,6 +156,7 @@ async def handle_document_tracking(
combined_content: Combined content text from all content items
is_first_batch: Whether this is the first batch (for chunked operations)
retain_params: Optional parameters passed during retain (context, event_date, etc.)
document_tags: Optional list of tags to associate with the document
"""
import hashlib
@@ -151,20 +167,20 @@ async def handle_document_tracking(
# Only delete on the first batch to avoid deleting data we just inserted
if is_first_batch:
await conn.fetchval(
"DELETE FROM documents WHERE id = $1 AND bank_id = $2 RETURNING id",
document_id, bank_id
f"DELETE FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2 RETURNING id", document_id, bank_id
)
# Insert document (or update if exists from concurrent operations)
await conn.execute(
"""
INSERT INTO documents (id, bank_id, original_text, content_hash, metadata, retain_params)
VALUES ($1, $2, $3, $4, $5, $6)
f"""
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params, tags)
VALUES ($1, $2, $3, $4, $5, $6, $7)
ON CONFLICT (id, bank_id) DO UPDATE
SET original_text = EXCLUDED.original_text,
content_hash = EXCLUDED.content_hash,
metadata = EXCLUDED.metadata,
retain_params = EXCLUDED.retain_params,
tags = EXCLUDED.tags,
updated_at = NOW()
""",
document_id,
@@ -172,5 +188,6 @@ async def handle_document_tracking(
combined_content,
content_hash,
json.dumps({}), # Empty metadata dict
json.dumps(retain_params) if retain_params else None
json.dumps(retain_params) if retain_params else None,
document_tags or [],
)
@@ -3,20 +3,16 @@ Link creation for retain pipeline.
Handles creation of temporal, semantic, and causal links between facts.
"""
import logging
from typing import List
from .types import ProcessedFact, CausalRelation
import logging
from . import link_utils
from .types import ProcessedFact
logger = logging.getLogger(__name__)
async def create_temporal_links_batch(
conn,
bank_id: str,
unit_ids: List[str]
) -> int:
async def create_temporal_links_batch(conn, bank_id: str, unit_ids: list[str]) -> int:
"""
Create temporal links between facts.
@@ -33,20 +29,10 @@ async def create_temporal_links_batch(
if not unit_ids:
return 0
return await link_utils.create_temporal_links_batch_per_fact(
conn,
bank_id,
unit_ids,
log_buffer=[]
)
return await link_utils.create_temporal_links_batch_per_fact(conn, bank_id, unit_ids, log_buffer=[])
async def create_semantic_links_batch(
conn,
bank_id: str,
unit_ids: List[str],
embeddings: List[List[float]]
) -> int:
async def create_semantic_links_batch(conn, bank_id: str, unit_ids: list[str], embeddings: list[list[float]]) -> int:
"""
Create semantic links between facts.
@@ -67,20 +53,10 @@ async def create_semantic_links_batch(
if len(unit_ids) != len(embeddings):
raise ValueError(f"Mismatch between unit_ids ({len(unit_ids)}) and embeddings ({len(embeddings)})")
return await link_utils.create_semantic_links_batch(
conn,
bank_id,
unit_ids,
embeddings,
log_buffer=[]
)
return await link_utils.create_semantic_links_batch(conn, bank_id, unit_ids, embeddings, log_buffer=[])
async def create_causal_links_batch(
conn,
unit_ids: List[str],
facts: List[ProcessedFact]
) -> int:
async def create_causal_links_batch(conn, unit_ids: list[str], facts: list[ProcessedFact]) -> int:
"""
Create causal links between facts.
@@ -108,9 +84,9 @@ async def create_causal_links_batch(
# Convert CausalRelation objects to dicts
relations_dicts = [
{
'relation_type': rel.relation_type,
'target_fact_index': rel.target_fact_index,
'strength': rel.strength
"relation_type": rel.relation_type,
"target_fact_index": rel.target_fact_index,
"strength": rel.strength,
}
for rel in fact.causal_relations
]
@@ -118,10 +94,6 @@ async def create_causal_links_batch(
else:
causal_relations_per_fact.append([])
link_count = await link_utils.create_causal_links_batch(
conn,
unit_ids,
causal_relations_per_fact
)
link_count = await link_utils.create_causal_links_batch(conn, unit_ids, causal_relations_per_fact)
return link_count
@@ -2,12 +2,12 @@
Link creation utilities for temporal, semantic, and entity links.
"""
import time
import logging
from typing import List
from datetime import timedelta, datetime, timezone
import time
from datetime import UTC, datetime, timedelta
from uuid import UUID
from ..memory_engine import fq_table
from .types import EntityLink
logger = logging.getLogger(__name__)
@@ -19,7 +19,7 @@ def _normalize_datetime(dt):
return None
if dt.tzinfo is None:
# Naive datetime - assume UTC
return dt.replace(tzinfo=timezone.utc)
return dt.replace(tzinfo=UTC)
return dt
@@ -54,24 +54,26 @@ def compute_temporal_links(
try:
time_lower = unit_event_date_norm - timedelta(hours=time_window_hours)
except OverflowError:
time_lower = datetime.min.replace(tzinfo=timezone.utc)
time_lower = datetime.min.replace(tzinfo=UTC)
try:
time_upper = unit_event_date_norm + timedelta(hours=time_window_hours)
except OverflowError:
time_upper = datetime.max.replace(tzinfo=timezone.utc)
time_upper = datetime.max.replace(tzinfo=UTC)
# Filter candidates within this unit's time window
matching_neighbors = [
(row['id'], row['event_date'])
(row["id"], row["event_date"])
for row in candidates
if time_lower <= _normalize_datetime(row['event_date']) <= time_upper
if time_lower <= _normalize_datetime(row["event_date"]) <= time_upper
][:10] # Limit to top 10
for recent_id, recent_event_date in matching_neighbors:
# Calculate temporal proximity weight
time_diff_hours = abs((unit_event_date_norm - _normalize_datetime(recent_event_date)).total_seconds() / 3600)
time_diff_hours = abs(
(unit_event_date_norm - _normalize_datetime(recent_event_date)).total_seconds() / 3600
)
weight = max(0.3, 1.0 - (time_diff_hours / time_window_hours))
links.append((unit_id, str(recent_id), 'temporal', weight, None))
links.append((unit_id, str(recent_id), "temporal", weight, None))
return links
@@ -99,17 +101,17 @@ def compute_temporal_query_bounds(
try:
min_date = min(all_dates) - timedelta(hours=time_window_hours)
except OverflowError:
min_date = datetime.min.replace(tzinfo=timezone.utc)
min_date = datetime.min.replace(tzinfo=UTC)
try:
max_date = max(all_dates) + timedelta(hours=time_window_hours)
except OverflowError:
max_date = datetime.max.replace(tzinfo=timezone.utc)
max_date = datetime.max.replace(tzinfo=UTC)
return min_date, max_date
def _log(log_buffer, message, level='info'):
def _log(log_buffer, message, level="info"):
"""Helper to log to buffer if available, otherwise use logger.
Args:
@@ -117,7 +119,7 @@ def _log(log_buffer, message, level='info'):
message: The log message
level: 'info', 'debug', 'warning', or 'error'. Debug messages are not added to buffer.
"""
if level == 'debug':
if level == "debug":
# Debug messages only go to logger, not to buffer
logger.debug(message)
return
@@ -125,23 +127,23 @@ def _log(log_buffer, message, level='info'):
if log_buffer is not None:
log_buffer.append(message)
else:
if level == 'info':
if level == "info":
logger.info(message)
else:
logger.log(logging.WARNING if level == 'warning' else logging.ERROR, message)
logger.log(logging.WARNING if level == "warning" else logging.ERROR, message)
async def extract_entities_batch_optimized(
entity_resolver,
conn,
bank_id: str,
unit_ids: List[str],
sentences: List[str],
unit_ids: list[str],
sentences: list[str],
context: str,
fact_dates: List,
llm_entities: List[List[dict]],
log_buffer: List[str] = None,
) -> List[tuple]:
fact_dates: list,
llm_entities: list[list[dict]],
log_buffer: list[str] = None,
) -> list[tuple]:
"""
Process LLM-extracted entities for ALL facts in batch.
@@ -171,15 +173,19 @@ async def extract_entities_batch_optimized(
formatted_entities = []
for ent in entity_list:
# Handle both Entity objects and dicts
if hasattr(ent, 'text'):
if hasattr(ent, "text"):
# Entity objects only have 'text', default type to 'CONCEPT'
formatted_entities.append({'text': ent.text, 'type': 'CONCEPT'})
formatted_entities.append({"text": ent.text, "type": "CONCEPT"})
elif isinstance(ent, dict):
formatted_entities.append({'text': ent.get('text', ''), 'type': ent.get('type', 'CONCEPT')})
formatted_entities.append({"text": ent.get("text", ""), "type": ent.get("type", "CONCEPT")})
all_entities.append(formatted_entities)
total_entities = sum(len(ents) for ents in all_entities)
_log(log_buffer, f" [6.1] Process LLM entities: {total_entities} entities from {len(sentences)} facts in {time.time() - substep_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.1] Process LLM entities: {total_entities} entities from {len(sentences)} facts in {time.time() - substep_start:.3f}s",
level="debug",
)
# Step 2: Resolve entities in BATCH (much faster!)
substep_start = time.time()
@@ -195,13 +201,19 @@ async def extract_entities_batch_optimized(
continue
for local_idx, entity in enumerate(entities):
all_entities_flat.append({
'text': entity['text'],
'type': entity['type'],
'nearby_entities': entities,
})
all_entities_flat.append(
{
"text": entity["text"],
"type": entity["type"],
"nearby_entities": entities,
}
)
entity_to_unit.append((unit_id, local_idx, fact_date))
_log(log_buffer, f" [6.2.1] Prepare entities: {len(all_entities_flat)} entities in {time.time() - substep_6_2_1_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2.1] Prepare entities: {len(all_entities_flat)} entities in {time.time() - substep_6_2_1_start:.3f}s",
level="debug",
)
# Resolve ALL entities in one batch call
if all_entities_flat:
@@ -210,7 +222,7 @@ async def extract_entities_batch_optimized(
# Add per-entity dates to entity data for batch resolution
for idx, (unit_id, local_idx, fact_date) in enumerate(entity_to_unit):
all_entities_flat[idx]['event_date'] = fact_date
all_entities_flat[idx]["event_date"] = fact_date
# Resolve ALL entities in ONE batch call (much faster than sequential buckets)
# INSERT ... ON CONFLICT handles any race conditions at the DB level
@@ -219,10 +231,14 @@ async def extract_entities_batch_optimized(
entities_data=all_entities_flat,
context=context,
unit_event_date=None, # Not used when per-entity dates provided
conn=conn # Use main transaction connection
conn=conn, # Use main transaction connection
)
_log(log_buffer, f" [6.2.2] Resolve entities: {len(all_entities_flat)} entities in single batch in {time.time() - substep_6_2_2_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2.2] Resolve entities: {len(all_entities_flat)} entities in single batch in {time.time() - substep_6_2_2_start:.3f}s",
level="debug",
)
# [6.2.3] Create unit-entity links in BATCH
substep_6_2_3_start = time.time()
@@ -239,12 +255,24 @@ async def extract_entities_batch_optimized(
# Batch insert all unit-entity links (MUCH faster!)
await entity_resolver.link_units_to_entities_batch(unit_entity_pairs, conn=conn)
_log(log_buffer, f" [6.2.3] Create unit-entity links (batched): {len(unit_entity_pairs)} links in {time.time() - substep_6_2_3_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2.3] Create unit-entity links (batched): {len(unit_entity_pairs)} links in {time.time() - substep_6_2_3_start:.3f}s",
level="debug",
)
_log(log_buffer, f" [6.2] Entity resolution (batched): {len(all_entities_flat)} entities resolved in {time.time() - step_6_2_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2] Entity resolution (batched): {len(all_entities_flat)} entities resolved in {time.time() - step_6_2_start:.3f}s",
level="debug",
)
else:
unit_to_entity_ids = {}
_log(log_buffer, f" [6.2] Entity resolution (batched): 0 entities in {time.time() - step_6_2_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2] Entity resolution (batched): 0 entities in {time.time() - step_6_2_start:.3f}s",
level="debug",
)
# Step 3: Create entity links between units that share entities
substep_start = time.time()
@@ -253,39 +281,44 @@ async def extract_entities_batch_optimized(
for entity_ids in unit_to_entity_ids.values():
all_entity_ids.update(entity_ids)
_log(log_buffer, f" [6.3] Creating entity links for {len(all_entity_ids)} unique entities...", level='debug')
_log(log_buffer, f" [6.3] Creating entity links for {len(all_entity_ids)} unique entities...", level="debug")
# Find all units that reference these entities (ONE batched query)
entity_to_units = {}
if all_entity_ids:
query_start = time.time()
import uuid
entity_id_list = [uuid.UUID(eid) if isinstance(eid, str) else eid for eid in all_entity_ids]
rows = await conn.fetch(
"""
f"""
SELECT entity_id, unit_id
FROM unit_entities
FROM {fq_table("unit_entities")}
WHERE entity_id = ANY($1::uuid[])
""",
entity_id_list
entity_id_list,
)
_log(
log_buffer,
f" [6.3.1] Query unit_entities: {len(rows)} rows in {time.time() - query_start:.3f}s",
level="debug",
)
_log(log_buffer, f" [6.3.1] Query unit_entities: {len(rows)} rows in {time.time() - query_start:.3f}s", level='debug')
# Group by entity_id
group_start = time.time()
for row in rows:
entity_id = row['entity_id']
entity_id = row["entity_id"]
if entity_id not in entity_to_units:
entity_to_units[entity_id] = []
entity_to_units[entity_id].append(row['unit_id'])
_log(log_buffer, f" [6.3.2] Group by entity_id: {time.time() - group_start:.3f}s", level='debug')
entity_to_units[entity_id].append(row["unit_id"])
_log(log_buffer, f" [6.3.2] Group by entity_id: {time.time() - group_start:.3f}s", level="debug")
# Create bidirectional links between units that share entities
# OPTIMIZATION: Limit links per entity to avoid N² explosion
# Only link each new unit to the most recent MAX_LINKS_PER_ENTITY units
MAX_LINKS_PER_ENTITY = 50 # Limit to prevent explosion when entity appears in many facts
link_gen_start = time.time()
links: List[EntityLink] = []
links: list[EntityLink] = []
new_unit_set = set(unit_ids) # Units from this batch
def to_uuid(val) -> UUID:
@@ -299,27 +332,52 @@ async def extract_entities_batch_optimized(
# Link new units to each other (within batch) - also limited
# For very common entities, limit within-batch links too
new_units_to_link = new_units[-MAX_LINKS_PER_ENTITY:] if len(new_units) > MAX_LINKS_PER_ENTITY else new_units
new_units_to_link = (
new_units[-MAX_LINKS_PER_ENTITY:] if len(new_units) > MAX_LINKS_PER_ENTITY else new_units
)
for i, unit_id_1 in enumerate(new_units_to_link):
for unit_id_2 in new_units_to_link[i+1:]:
links.append(EntityLink(from_unit_id=to_uuid(unit_id_1), to_unit_id=to_uuid(unit_id_2), entity_id=entity_uuid))
links.append(EntityLink(from_unit_id=to_uuid(unit_id_2), to_unit_id=to_uuid(unit_id_1), entity_id=entity_uuid))
for unit_id_2 in new_units_to_link[i + 1 :]:
links.append(
EntityLink(
from_unit_id=to_uuid(unit_id_1), to_unit_id=to_uuid(unit_id_2), entity_id=entity_uuid
)
)
links.append(
EntityLink(
from_unit_id=to_uuid(unit_id_2), to_unit_id=to_uuid(unit_id_1), entity_id=entity_uuid
)
)
# Link new units to LIMITED existing units (most recent)
existing_to_link = existing_units[-MAX_LINKS_PER_ENTITY:] # Take most recent
for new_unit in new_units:
for existing_unit in existing_to_link:
links.append(EntityLink(from_unit_id=to_uuid(new_unit), to_unit_id=to_uuid(existing_unit), entity_id=entity_uuid))
links.append(EntityLink(from_unit_id=to_uuid(existing_unit), to_unit_id=to_uuid(new_unit), entity_id=entity_uuid))
links.append(
EntityLink(
from_unit_id=to_uuid(new_unit), to_unit_id=to_uuid(existing_unit), entity_id=entity_uuid
)
)
links.append(
EntityLink(
from_unit_id=to_uuid(existing_unit), to_unit_id=to_uuid(new_unit), entity_id=entity_uuid
)
)
_log(log_buffer, f" [6.3.3] Generate {len(links)} links: {time.time() - link_gen_start:.3f}s", level='debug')
_log(log_buffer, f" [6.3] Entity link creation: {len(links)} links for {len(all_entity_ids)} unique entities in {time.time() - substep_start:.3f}s", level='debug')
_log(
log_buffer, f" [6.3.3] Generate {len(links)} links: {time.time() - link_gen_start:.3f}s", level="debug"
)
_log(
log_buffer,
f" [6.3] Entity link creation: {len(links)} links for {len(all_entity_ids)} unique entities in {time.time() - substep_start:.3f}s",
level="debug",
)
return links
except Exception as e:
logger.error(f"Failed to extract entities in batch: {str(e)}")
import traceback
traceback.print_exc()
raise
@@ -327,9 +385,9 @@ async def extract_entities_batch_optimized(
async def create_temporal_links_batch_per_fact(
conn,
bank_id: str,
unit_ids: List[str],
unit_ids: list[str],
time_window_hours: int = 24,
log_buffer: List[str] = None,
log_buffer: list[str] = None,
) -> int:
"""
Create temporal links for multiple units, each with their own event_date.
@@ -356,15 +414,18 @@ async def create_temporal_links_batch_per_fact(
# Get the event_date for each new unit
fetch_dates_start = time_mod.time()
rows = await conn.fetch(
"""
f"""
SELECT id, event_date
FROM memory_units
FROM {fq_table("memory_units")}
WHERE id::text = ANY($1)
""",
unit_ids
unit_ids,
)
new_units = {str(row["id"]): row["event_date"] for row in rows}
_log(
log_buffer,
f" [7.1] Fetch event_dates for {len(unit_ids)} units: {time_mod.time() - fetch_dates_start:.3f}s",
)
new_units = {str(row['id']): row['event_date'] for row in rows}
_log(log_buffer, f" [7.1] Fetch event_dates for {len(unit_ids)} units: {time_mod.time() - fetch_dates_start:.3f}s")
# Fetch ALL potential temporal neighbors in ONE query (much faster!)
# Get time range across all units with overflow protection
@@ -372,9 +433,9 @@ async def create_temporal_links_batch_per_fact(
fetch_neighbors_start = time_mod.time()
all_candidates = await conn.fetch(
"""
f"""
SELECT id, event_date
FROM memory_units
FROM {fq_table("memory_units")}
WHERE bank_id = $1
AND event_date BETWEEN $2 AND $3
AND id::text != ALL($4)
@@ -383,25 +444,53 @@ async def create_temporal_links_batch_per_fact(
bank_id,
min_date,
max_date,
unit_ids
unit_ids,
)
_log(
log_buffer,
f" [7.2] Fetch {len(all_candidates)} candidate neighbors (1 query): {time_mod.time() - fetch_neighbors_start:.3f}s",
)
_log(log_buffer, f" [7.2] Fetch {len(all_candidates)} candidate neighbors (1 query): {time_mod.time() - fetch_neighbors_start:.3f}s")
# Filter and create links in memory (much faster than N queries)
link_gen_start = time_mod.time()
links = compute_temporal_links(new_units, all_candidates, time_window_hours)
# Also compute temporal links WITHIN the new batch (new units to each other)
if len(new_units) > 1:
# Convert new_units dict to candidate format for within-batch linking
new_unit_items = list(new_units.items())
for i, (unit_id, event_date) in enumerate(new_unit_items):
unit_event_date_norm = _normalize_datetime(event_date)
# Compare with other new units (only those after this one to avoid duplicates)
for j in range(i + 1, len(new_unit_items)):
other_id, other_event_date = new_unit_items[j]
other_event_date_norm = _normalize_datetime(other_event_date)
# Check if within time window
time_diff_hours = abs((unit_event_date_norm - other_event_date_norm).total_seconds() / 3600)
if time_diff_hours <= time_window_hours:
weight = max(0.3, 1.0 - (time_diff_hours / time_window_hours))
# Create bidirectional links
links.append((unit_id, other_id, "temporal", weight, None))
links.append((other_id, unit_id, "temporal", weight, None))
_log(log_buffer, f" [7.3] Generate {len(links)} temporal links: {time_mod.time() - link_gen_start:.3f}s")
if links:
insert_start = time_mod.time()
await conn.executemany(
"""
INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
links
)
# Batch inserts to avoid timeout on large batches
BATCH_SIZE = 1000
for batch_start in range(0, len(links), BATCH_SIZE):
batch = links[batch_start : batch_start + BATCH_SIZE]
await conn.executemany(
f"""
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
batch,
)
_log(log_buffer, f" [7.4] Insert {len(links)} temporal links: {time_mod.time() - insert_start:.3f}s")
return len(links)
@@ -409,6 +498,7 @@ async def create_temporal_links_batch_per_fact(
except Exception as e:
logger.error(f"Failed to create temporal links: {str(e)}")
import traceback
traceback.print_exc()
raise
@@ -416,11 +506,11 @@ async def create_temporal_links_batch_per_fact(
async def create_semantic_links_batch(
conn,
bank_id: str,
unit_ids: List[str],
embeddings: List[List[float]],
unit_ids: list[str],
embeddings: list[list[float]],
top_k: int = 5,
threshold: float = 0.7,
log_buffer: List[str] = None,
log_buffer: list[str] = None,
) -> int:
"""
Create semantic links for multiple units efficiently.
@@ -444,22 +534,26 @@ async def create_semantic_links_batch(
try:
import time as time_mod
import numpy as np
# Fetch ALL existing units with embeddings in ONE query
fetch_start = time_mod.time()
all_existing = await conn.fetch(
"""
f"""
SELECT id, embedding
FROM memory_units
FROM {fq_table("memory_units")}
WHERE bank_id = $1
AND embedding IS NOT NULL
AND id::text != ALL($2)
""",
bank_id,
unit_ids
unit_ids,
)
_log(
log_buffer,
f" [8.1] Fetch {len(all_existing)} existing embeddings (1 query): {time_mod.time() - fetch_start:.3f}s",
)
_log(log_buffer, f" [8.1] Fetch {len(all_existing)} existing embeddings (1 query): {time_mod.time() - fetch_start:.3f}s")
# Convert to numpy for vectorized similarity computation
compute_start = time_mod.time()
@@ -467,15 +561,16 @@ async def create_semantic_links_batch(
if all_existing:
# Convert existing embeddings to numpy array
existing_ids = [str(row['id']) for row in all_existing]
existing_ids = [str(row["id"]) for row in all_existing]
# Stack embeddings as 2D array: (num_embeddings, embedding_dim)
embedding_arrays = []
for row in all_existing:
raw_emb = row['embedding']
raw_emb = row["embedding"]
# Handle different pgvector formats
if isinstance(raw_emb, str):
# Parse string format: "[1.0, 2.0, ...]"
import json
emb = np.array(json.loads(raw_emb), dtype=np.float32)
elif isinstance(raw_emb, (list, tuple)):
emb = np.array(raw_emb, dtype=np.float32)
@@ -514,33 +609,72 @@ async def create_semantic_links_batch(
for idx in sorted_indices:
similar_id = existing_ids[idx]
similarity = float(similarities[idx])
all_links.append((unit_id, similar_id, 'semantic', similarity, None))
# Clamp to [0, 1] to handle floating point precision issues
similarity = float(min(1.0, max(0.0, similarities[idx])))
all_links.append((unit_id, similar_id, "semantic", similarity, None))
_log(log_buffer, f" [8.2] Compute similarities & generate {len(all_links)} semantic links: {time_mod.time() - compute_start:.3f}s")
# Also compute similarities WITHIN the new batch (new units to each other)
# Apply the same top_k limit per unit as we do for existing units
if len(unit_ids) > 1:
new_embeddings_matrix = np.array(embeddings)
for i, unit_id in enumerate(unit_ids):
# Compute similarities with all OTHER new units
other_indices = [j for j in range(len(unit_ids)) if j != i]
if not other_indices:
continue
other_embeddings = new_embeddings_matrix[other_indices]
similarities = np.dot(other_embeddings, new_embeddings_matrix[i])
# Find top-k above threshold (same logic as existing units)
above_threshold = np.where(similarities >= threshold)[0]
if len(above_threshold) > 0:
# Sort by similarity (descending) and take top-k
sorted_local_indices = above_threshold[np.argsort(-similarities[above_threshold])][:top_k]
for local_idx in sorted_local_indices:
other_idx = other_indices[local_idx]
other_id = unit_ids[other_idx]
# Clamp to [0, 1] to handle floating point precision issues
similarity = float(min(1.0, max(0.0, similarities[local_idx])))
all_links.append((unit_id, other_id, "semantic", similarity, None))
_log(
log_buffer,
f" [8.2] Compute similarities & generate {len(all_links)} semantic links: {time_mod.time() - compute_start:.3f}s",
)
if all_links:
insert_start = time_mod.time()
await conn.executemany(
"""
INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
all_links
# Batch inserts to avoid timeout on large batches
BATCH_SIZE = 1000
for batch_start in range(0, len(all_links), BATCH_SIZE):
batch = all_links[batch_start : batch_start + BATCH_SIZE]
await conn.executemany(
f"""
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
batch,
)
_log(
log_buffer, f" [8.3] Insert {len(all_links)} semantic links: {time_mod.time() - insert_start:.3f}s"
)
_log(log_buffer, f" [8.3] Insert {len(all_links)} semantic links: {time_mod.time() - insert_start:.3f}s")
return len(all_links)
except Exception as e:
logger.error(f"Failed to create semantic links: {str(e)}")
import traceback
traceback.print_exc()
raise
async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: int = 50000):
async def insert_entity_links_batch(conn, links: list[EntityLink], chunk_size: int = 50000):
"""
Insert all entity links using COPY to temp table + INSERT for maximum speed.
@@ -556,7 +690,6 @@ async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: i
if not links:
return
import uuid as uuid_mod
import time as time_mod
total_start = time_mod.time()
@@ -583,28 +716,22 @@ async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: i
convert_start = time_mod.time()
records = []
for link in links:
records.append((
link.from_unit_id,
link.to_unit_id,
link.link_type,
link.weight,
link.entity_id
))
records.append((link.from_unit_id, link.to_unit_id, link.link_type, link.weight, link.entity_id))
logger.debug(f" [9.3] Convert {len(records)} records: {time_mod.time() - convert_start:.3f}s")
# Bulk load using COPY (fastest method)
copy_start = time_mod.time()
await conn.copy_records_to_table(
'_temp_entity_links',
"_temp_entity_links",
records=records,
columns=['from_unit_id', 'to_unit_id', 'link_type', 'weight', 'entity_id']
columns=["from_unit_id", "to_unit_id", "link_type", "weight", "entity_id"],
)
logger.debug(f" [9.4] COPY {len(records)} records to temp table: {time_mod.time() - copy_start:.3f}s")
# Insert from temp table with ON CONFLICT (single query for all rows)
insert_start = time_mod.time()
await conn.execute("""
INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id)
await conn.execute(f"""
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
SELECT from_unit_id, to_unit_id, link_type, weight, entity_id
FROM _temp_entity_links
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
@@ -615,8 +742,8 @@ async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: i
async def create_causal_links_batch(
conn,
unit_ids: List[str],
causal_relations_per_fact: List[List[dict]],
unit_ids: list[str],
causal_relations_per_fact: list[list[dict]],
) -> int:
"""
Create causal links between facts based on LLM-extracted causal relationships.
@@ -644,6 +771,7 @@ async def create_causal_links_batch(
try:
import time as time_mod
create_start = time_mod.time()
# Build links list
@@ -655,12 +783,12 @@ async def create_causal_links_batch(
from_unit_id = unit_ids[fact_idx]
for relation in causal_relations:
target_idx = relation['target_fact_index']
relation_type = relation['relation_type']
strength = relation.get('strength', 1.0)
target_idx = relation["target_fact_index"]
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'}
valid_types = {"causes", "caused_by", "enables", "prevents"}
if relation_type not in valid_types:
logger.error(
f"Invalid relation_type '{relation_type}' (type: {type(relation_type).__name__}) "
@@ -685,24 +813,25 @@ async def create_causal_links_batch(
# weight is the strength of the relationship
links.append((from_unit_id, to_unit_id, relation_type, strength, None))
if links:
insert_start = time_mod.time()
try:
await conn.executemany(
"""
INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id)
f"""
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
links
links,
)
except Exception as db_error:
# Log the actual data being inserted for debugging
logger.error(f"Database insert failed for causal links. Error: {db_error}")
logger.error(f"Attempted to insert {len(links)} links. First few:")
for i, link in enumerate(links[:3]):
logger.error(f" Link {i}: from={link[0]}, to={link[1]}, type='{link[2]}' (repr={repr(link[2])}), weight={link[3]}, entity={link[4]}")
logger.error(
f" Link {i}: from={link[0]}, to={link[1]}, type='{link[2]}' (repr={repr(link[2])}), weight={link[3]}, entity={link[4]}"
)
raise
return len(links)
@@ -710,5 +839,6 @@ async def create_causal_links_batch(
except Exception as e:
logger.error(f"Failed to create causal links: {str(e)}")
import traceback
traceback.print_exc()
raise
@@ -1,264 +0,0 @@
"""
Observation regeneration for retain pipeline.
Regenerates entity observations as part of the retain transaction.
"""
import logging
import time
import uuid
from datetime import datetime, timezone
from typing import List, Dict, Optional
from ..search import observation_utils
from . import embedding_utils
from ..db_utils import acquire_with_retry
from .types import EntityLink
logger = logging.getLogger(__name__)
def utcnow():
"""Get current UTC time."""
return datetime.now(timezone.utc)
# Simple dataclass-like container for facts (avoid importing from memory_engine)
class MemoryFactForObservation:
def __init__(self, id: str, text: str, fact_type: str, context: str, occurred_start: Optional[str]):
self.id = id
self.text = text
self.fact_type = fact_type
self.context = context
self.occurred_start = occurred_start
async def regenerate_observations_batch(
conn,
embeddings_model,
llm_config,
bank_id: str,
entity_links: List[EntityLink],
log_buffer: List[str] = None
) -> None:
"""
Regenerate observations for top entities in this batch.
Called INSIDE the retain transaction for atomicity - if observations
fail, the entire retain batch is rolled back.
Args:
conn: Database connection (from the retain transaction)
embeddings_model: Embeddings model for generating observation embeddings
llm_config: LLM configuration for observation extraction
bank_id: Bank identifier
entity_links: Entity links from this batch
log_buffer: Optional log buffer for timing
"""
TOP_N_ENTITIES = 5
MIN_FACTS_THRESHOLD = 5
if not entity_links:
return
# Count mentions per entity in this batch
entity_mention_counts: Dict[str, int] = {}
for link in entity_links:
if link.entity_id:
entity_id = str(link.entity_id)
entity_mention_counts[entity_id] = entity_mention_counts.get(entity_id, 0) + 1
if not entity_mention_counts:
return
# Sort by mention count descending and take top N
sorted_entities = sorted(
entity_mention_counts.items(),
key=lambda x: x[1],
reverse=True
)
entities_to_process = [e[0] for e in sorted_entities[:TOP_N_ENTITIES]]
obs_start = time.time()
# Convert to UUIDs
entity_uuids = [uuid.UUID(eid) if isinstance(eid, str) else eid for eid in entities_to_process]
# Batch query for entity names
entity_rows = await conn.fetch(
"""
SELECT id, canonical_name FROM entities
WHERE id = ANY($1) AND bank_id = $2
""",
entity_uuids, bank_id
)
entity_names = {row['id']: row['canonical_name'] for row in entity_rows}
# Batch query for fact counts
fact_counts = await conn.fetch(
"""
SELECT ue.entity_id, COUNT(*) as cnt
FROM unit_entities ue
JOIN memory_units mu ON ue.unit_id = mu.id
WHERE ue.entity_id = ANY($1) AND mu.bank_id = $2
GROUP BY ue.entity_id
""",
entity_uuids, bank_id
)
entity_fact_counts = {row['entity_id']: row['cnt'] for row in fact_counts}
# Filter entities that meet the threshold
entities_with_names = []
for entity_id in entities_to_process:
entity_uuid = uuid.UUID(entity_id) if isinstance(entity_id, str) else entity_id
if entity_uuid not in entity_names:
continue
fact_count = entity_fact_counts.get(entity_uuid, 0)
if fact_count >= MIN_FACTS_THRESHOLD:
entities_with_names.append((entity_id, entity_names[entity_uuid]))
if not entities_with_names:
return
# Process entities SEQUENTIALLY (asyncpg doesn't allow concurrent queries on same connection)
# We must use the same connection to stay in the retain transaction
total_observations = 0
for entity_id, entity_name in entities_with_names:
try:
obs_ids = await _regenerate_entity_observations(
conn, embeddings_model, llm_config,
bank_id, entity_id, entity_name
)
total_observations += len(obs_ids)
except Exception as e:
logger.error(f"[OBSERVATIONS] Error processing entity {entity_id}: {e}")
obs_time = time.time() - obs_start
if log_buffer is not None:
log_buffer.append(f"[11] Observations: {total_observations} observations for {len(entities_with_names)} entities in {obs_time:.3f}s")
async def _regenerate_entity_observations(
conn,
embeddings_model,
llm_config,
bank_id: str,
entity_id: str,
entity_name: str
) -> List[str]:
"""
Regenerate observations for a single entity.
Uses the provided connection (part of retain transaction).
Args:
conn: Database connection (from the retain transaction)
embeddings_model: Embeddings model
llm_config: LLM configuration
bank_id: Bank identifier
entity_id: Entity UUID
entity_name: Canonical name of the entity
Returns:
List of created observation IDs
"""
entity_uuid = uuid.UUID(entity_id) if isinstance(entity_id, str) else entity_id
# Get all facts mentioning this entity (exclude observations themselves)
rows = await conn.fetch(
"""
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.fact_type
FROM memory_units mu
JOIN unit_entities ue ON mu.id = ue.unit_id
WHERE mu.bank_id = $1
AND ue.entity_id = $2
AND mu.fact_type IN ('world', 'experience')
ORDER BY mu.occurred_start DESC
LIMIT 50
""",
bank_id, entity_uuid
)
if not rows:
return []
# Convert to fact objects for observation extraction
facts = []
for row in rows:
occurred_start = row['occurred_start'].isoformat() if row['occurred_start'] else None
facts.append(MemoryFactForObservation(
id=str(row['id']),
text=row['text'],
fact_type=row['fact_type'],
context=row['context'],
occurred_start=occurred_start
))
# Extract observations using LLM
observations = await observation_utils.extract_observations_from_facts(
llm_config,
entity_name,
facts
)
if not observations:
return []
# Delete old observations for this entity
await conn.execute(
"""
DELETE FROM memory_units
WHERE id IN (
SELECT mu.id
FROM memory_units mu
JOIN unit_entities ue ON mu.id = ue.unit_id
WHERE mu.bank_id = $1
AND mu.fact_type = 'observation'
AND ue.entity_id = $2
)
""",
bank_id, entity_uuid
)
# Generate embeddings for new observations
embeddings = await embedding_utils.generate_embeddings_batch(
embeddings_model, observations
)
# Insert new observations
current_time = utcnow()
created_ids = []
for obs_text, embedding in zip(observations, embeddings):
result = await conn.fetchrow(
"""
INSERT INTO memory_units (
bank_id, text, embedding, context, event_date,
occurred_start, occurred_end, mentioned_at,
fact_type, access_count
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, 'observation', 0)
RETURNING id
""",
bank_id,
obs_text,
str(embedding),
f"observation about {entity_name}",
current_time,
current_time,
current_time,
current_time
)
obs_id = str(result['id'])
created_ids.append(obs_id)
# Link observation to entity
await conn.execute(
"""
INSERT INTO unit_entities (unit_id, entity_id)
VALUES ($1, $2)
""",
uuid.UUID(obs_id), entity_uuid
)
return created_ids
@@ -3,31 +3,32 @@ Main orchestrator for the retain pipeline.
Coordinates all retain pipeline modules to store memories efficiently.
"""
import logging
import time
import uuid
from datetime import datetime, timezone
from typing import List, Dict, Any, Optional
from datetime import UTC, datetime
from . import bank_utils
from ..db_utils import acquire_with_retry
from . import bank_utils
def utcnow():
"""Get current UTC time."""
return datetime.now(timezone.utc)
return datetime.now(UTC)
from .types import RetainContent, ExtractedFact, ProcessedFact, EntityLink
from ..response_models import TokenUsage
from . import (
fact_extraction,
embedding_processing,
deduplication,
chunk_storage,
fact_storage,
deduplication,
embedding_processing,
entity_processing,
fact_extraction,
fact_storage,
link_creation,
observation_regeneration
)
from .types import EntityLink, ExtractedFact, ProcessedFact, RetainContent, RetainContentDict
logger = logging.getLogger(__name__)
@@ -37,16 +38,16 @@ async def retain_batch(
embeddings_model,
llm_config,
entity_resolver,
task_backend,
format_date_fn,
duplicate_checker_fn,
bank_id: str,
contents_dicts: List[Dict[str, Any]],
document_id: Optional[str] = None,
contents_dicts: list[RetainContentDict],
document_id: str | None = None,
is_first_batch: bool = True,
fact_type_override: Optional[str] = None,
confidence_score: Optional[float] = None,
) -> List[List[str]]:
fact_type_override: str | None = None,
confidence_score: float | None = None,
document_tags: list[str] | None = None,
) -> tuple[list[list[str]], TokenUsage]:
"""
Process a batch of content through the retain pipeline.
@@ -55,7 +56,6 @@ async def retain_batch(
embeddings_model: Embeddings model for generating embeddings
llm_config: LLM configuration for fact extraction
entity_resolver: Entity resolver for entity processing
task_backend: Task backend for background jobs
format_date_fn: Function to format datetime to readable string
duplicate_checker_fn: Function to check for duplicate facts
bank_id: Bank identifier
@@ -64,19 +64,20 @@ async def retain_batch(
is_first_batch: Whether this is the first batch
fact_type_override: Override fact type for all facts
confidence_score: Confidence score for opinions
document_tags: Tags applied to all items in this batch
Returns:
List of unit ID lists (one list per content item)
Tuple of (unit ID lists, token usage for fact extraction)
"""
start_time = time.time()
total_chars = sum(len(item.get("content", "")) for item in contents_dicts)
# Buffer all logs
log_buffer = []
log_buffer.append(f"{'='*60}")
log_buffer.append(f"{'=' * 60}")
log_buffer.append(f"RETAIN_BATCH START: {bank_id}")
log_buffer.append(f"Batch size: {len(contents_dicts)} content items, {total_chars:,} chars")
log_buffer.append(f"{'='*60}")
log_buffer.append(f"{'=' * 60}")
# Get bank profile
profile = await bank_utils.get_bank_profile(pool, bank_id)
@@ -85,28 +86,89 @@ async def retain_batch(
# Convert dicts to RetainContent objects
contents = []
for item in contents_dicts:
# Merge item-level tags with document-level tags
item_tags = item.get("tags", []) or []
merged_tags = list(set(item_tags + (document_tags or [])))
content = RetainContent(
content=item["content"],
context=item.get("context", ""),
event_date=item.get("event_date") or utcnow(),
metadata=item.get("metadata", {})
metadata=item.get("metadata", {}),
entities=item.get("entities", []),
tags=merged_tags,
)
contents.append(content)
# Step 1: Extract facts from all contents
step_start = time.time()
extract_opinions = (fact_type_override == 'opinion')
extract_opinions = fact_type_override == "opinion"
extracted_facts, chunks = await fact_extraction.extract_facts_from_contents(
contents,
llm_config,
agent_name,
extract_opinions
extracted_facts, chunks, usage = await fact_extraction.extract_facts_from_contents(
contents, llm_config, agent_name, extract_opinions
)
log_buffer.append(
f"[1] Extract facts: {len(extracted_facts)} facts, {len(chunks)} chunks from {len(contents)} contents in {time.time() - step_start:.3f}s"
)
log_buffer.append(f"[1] Extract facts: {len(extracted_facts)} facts, {len(chunks)} chunks from {len(contents)} contents in {time.time() - step_start:.3f}s")
if not extracted_facts:
return [[] for _ in contents]
# Still need to create document if document_id was provided
async with acquire_with_retry(pool) as conn:
async with conn.transaction():
await fact_storage.ensure_bank_exists(conn, bank_id)
# Handle document tracking even with no facts
if document_id:
combined_content = "\n".join([c.get("content", "") for c in contents_dicts])
retain_params = {}
if contents_dicts:
first_item = contents_dicts[0]
if first_item.get("context"):
retain_params["context"] = first_item["context"]
if first_item.get("event_date"):
retain_params["event_date"] = (
first_item["event_date"].isoformat()
if hasattr(first_item["event_date"], "isoformat")
else str(first_item["event_date"])
)
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
)
else:
# Check for per-item document_ids
from collections import defaultdict
contents_by_doc = defaultdict(list)
for idx, content_dict in enumerate(contents_dicts):
doc_id = content_dict.get("document_id")
if doc_id:
contents_by_doc[doc_id].append((idx, content_dict))
for doc_id, doc_contents in contents_by_doc.items():
combined_content = "\n".join([c.get("content", "") for _, c in doc_contents])
retain_params = {}
if doc_contents:
first_item = doc_contents[0][1]
if first_item.get("context"):
retain_params["context"] = first_item["context"]
if first_item.get("event_date"):
retain_params["event_date"] = (
first_item["event_date"].isoformat()
if hasattr(first_item["event_date"], "isoformat")
else str(first_item["event_date"])
)
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params, document_tags
)
total_time = time.time() - start_time
logger.info(
f"RETAIN_BATCH COMPLETE: 0 facts extracted from {len(contents)} contents in {total_time:.3f}s (document tracked, no facts)"
)
return [[] for _ in contents], usage
# Apply fact_type_override if provided
if fact_type_override:
@@ -130,6 +192,7 @@ async def retain_batch(
# Group contents by document_id for document tracking and chunk storage
from collections import defaultdict
contents_by_doc = defaultdict(list)
for idx, content_dict in enumerate(contents_dicts):
doc_id = content_dict.get("document_id")
@@ -155,12 +218,16 @@ async def retain_batch(
if first_item.get("context"):
retain_params["context"] = first_item["context"]
if first_item.get("event_date"):
retain_params["event_date"] = first_item["event_date"].isoformat() if hasattr(first_item["event_date"], "isoformat") else str(first_item["event_date"])
retain_params["event_date"] = (
first_item["event_date"].isoformat()
if hasattr(first_item["event_date"], "isoformat")
else str(first_item["event_date"])
)
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
)
document_ids_added.append(document_id)
doc_id_mapping[None] = document_id # For backwards compatibility
@@ -195,17 +262,29 @@ async def retain_batch(
if first_item.get("context"):
retain_params["context"] = first_item["context"]
if first_item.get("event_date"):
retain_params["event_date"] = first_item["event_date"].isoformat() if hasattr(first_item["event_date"], "isoformat") else str(first_item["event_date"])
retain_params["event_date"] = (
first_item["event_date"].isoformat()
if hasattr(first_item["event_date"], "isoformat")
else str(first_item["event_date"])
)
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, actual_doc_id, combined_content, is_first_batch, retain_params
conn,
bank_id,
actual_doc_id,
combined_content,
is_first_batch,
retain_params,
document_tags,
)
document_ids_added.append(actual_doc_id)
if document_ids_added:
log_buffer.append(f"[2.5] Document tracking: {len(document_ids_added)} documents in {time.time() - step_start:.3f}s")
log_buffer.append(
f"[2.5] Document tracking: {len(document_ids_added)} documents in {time.time() - step_start:.3f}s"
)
# Store chunks and map to facts for all documents
step_start = time.time()
@@ -230,7 +309,9 @@ async def retain_batch(
for chunk_idx, chunk_id in chunk_id_map.items():
chunk_id_map_by_doc[(doc_id, chunk_idx)] = chunk_id
log_buffer.append(f"[3] Store chunks: {len(chunks)} chunks for {len(chunks_by_doc)} documents in {time.time() - step_start:.3f}s")
log_buffer.append(
f"[3] Store chunks: {len(chunks)} chunks for {len(chunks_by_doc)} documents in {time.time() - step_start:.3f}s"
)
# Map chunk_ids and document_ids to facts
for fact, processed_fact in zip(extracted_facts, processed_facts):
@@ -265,13 +346,15 @@ async def retain_batch(
is_duplicate_flags = await deduplication.check_duplicates_batch(
conn, bank_id, processed_facts, duplicate_checker_fn
)
log_buffer.append(f"[4] Deduplication: {sum(is_duplicate_flags)} duplicates in {time.time() - step_start:.3f}s")
log_buffer.append(
f"[4] Deduplication: {sum(is_duplicate_flags)} duplicates in {time.time() - step_start:.3f}s"
)
# Filter out duplicates
non_duplicate_facts = deduplication.filter_duplicates(processed_facts, is_duplicate_flags)
if not non_duplicate_facts:
return [[] for _ in contents]
return [[] for _ in contents], usage
# Insert facts (document_id is now stored per-fact)
step_start = time.time()
@@ -280,8 +363,18 @@ async def retain_batch(
# Process entities
step_start = time.time()
# Build map of content_index -> user entities for merging
user_entities_per_content = {
idx: content.entities for idx, content in enumerate(contents) if content.entities
}
entity_links = await entity_processing.process_entities_batch(
entity_resolver, conn, bank_id, unit_ids, non_duplicate_facts, log_buffer
entity_resolver,
conn,
bank_id,
unit_ids,
non_duplicate_facts,
log_buffer,
user_entities_per_content=user_entities_per_content,
)
log_buffer.append(f"[6] Process entities: {len(entity_links)} links in {time.time() - step_start:.3f}s")
@@ -293,62 +386,46 @@ async def retain_batch(
# Create semantic links
step_start = time.time()
embeddings_for_links = [fact.embedding for fact in non_duplicate_facts]
semantic_link_count = await link_creation.create_semantic_links_batch(conn, bank_id, unit_ids, embeddings_for_links)
semantic_link_count = await link_creation.create_semantic_links_batch(
conn, bank_id, unit_ids, embeddings_for_links
)
log_buffer.append(f"[8] Semantic links: {semantic_link_count} links in {time.time() - step_start:.3f}s")
# Insert entity links
step_start = time.time()
if entity_links:
await entity_processing.insert_entity_links_batch(conn, entity_links)
log_buffer.append(f"[9] Entity links: {len(entity_links) if entity_links else 0} links in {time.time() - step_start:.3f}s")
log_buffer.append(
f"[9] Entity links: {len(entity_links) if entity_links else 0} links in {time.time() - step_start:.3f}s"
)
# Create causal links
step_start = time.time()
causal_link_count = await link_creation.create_causal_links_batch(conn, unit_ids, non_duplicate_facts)
log_buffer.append(f"[10] Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s")
# Regenerate observations INSIDE transaction for atomicity
await observation_regeneration.regenerate_observations_batch(
conn,
embeddings_model,
llm_config,
bank_id,
entity_links,
log_buffer
)
# Map results back to original content items
result_unit_ids = _map_results_to_contents(
contents, extracted_facts, is_duplicate_flags, unit_ids
)
# Trigger background tasks AFTER transaction commits (opinion reinforcement only)
await _trigger_background_tasks(
task_backend,
bank_id,
unit_ids,
non_duplicate_facts
)
result_unit_ids = _map_results_to_contents(contents, extracted_facts, is_duplicate_flags, unit_ids)
# Log final summary
total_time = time.time() - start_time
log_buffer.append(f"{'='*60}")
log_buffer.append(f"{'=' * 60}")
log_buffer.append(f"RETAIN_BATCH COMPLETE: {len(unit_ids)} units in {total_time:.3f}s")
if document_ids_added:
log_buffer.append(f"Documents: {', '.join(document_ids_added)}")
log_buffer.append(f"{'='*60}")
log_buffer.append(f"{'=' * 60}")
logger.info("\n" + "\n".join(log_buffer) + "\n")
return result_unit_ids
return result_unit_ids, usage
def _map_results_to_contents(
contents: List[RetainContent],
extracted_facts: List[ExtractedFact],
is_duplicate_flags: List[bool],
unit_ids: List[str]
) -> List[List[str]]:
contents: list[RetainContent],
extracted_facts: list[ExtractedFact],
is_duplicate_flags: list[bool],
unit_ids: list[str],
) -> list[list[str]]:
"""
Map created unit IDs back to original content items.
@@ -371,22 +448,3 @@ def _map_results_to_contents(
result_unit_ids.append(content_unit_ids)
return result_unit_ids
async def _trigger_background_tasks(
task_backend,
bank_id: str,
unit_ids: List[str],
facts: List[ProcessedFact],
) -> None:
"""Trigger opinion reinforcement as background task (after transaction commits)."""
# Trigger opinion reinforcement if there are entities
fact_entities = [[e.name for e in fact.entities] for fact in facts]
if any(fact_entities):
await task_backend.submit_task({
'type': 'reinforce_opinion',
'bank_id': bank_id,
'created_unit_ids': unit_ids,
'unit_texts': [fact.fact_text for fact in facts],
'unit_entities': fact_entities
})
@@ -6,11 +6,38 @@ from content input to fact storage.
"""
from dataclasses import dataclass, field
from typing import List, Optional, Dict, Any
from datetime import datetime
from datetime import UTC, datetime
from typing import TypedDict
from uuid import UUID
class RetainContentDict(TypedDict, total=False):
"""Type definition for content items in retain_batch_async.
Fields:
content: Text content to store (required)
context: Context about the content (optional)
event_date: When the content occurred (optional, defaults to now)
metadata: Custom key-value metadata (optional)
document_id: Document ID for this content item (optional)
entities: User-provided entities to merge with extracted entities (optional)
tags: Visibility scope tags for this content item (optional)
"""
content: str # Required
context: str
event_date: datetime
metadata: dict[str, str]
document_id: str
entities: list[dict[str, str]] # [{"text": "...", "type": "..."}]
tags: list[str] # Visibility scope tags
def _now_utc() -> datetime:
"""Factory function for default event_date."""
return datetime.now(UTC)
@dataclass
class RetainContent:
"""
@@ -18,16 +45,13 @@ class RetainContent:
Represents a single piece of content to extract facts from.
"""
content: str
context: str = ""
event_date: Optional[datetime] = None
metadata: Dict[str, str] = field(default_factory=dict)
def __post_init__(self):
"""Ensure event_date is set."""
if self.event_date is None:
from datetime import datetime, timezone
self.event_date = datetime.now(timezone.utc)
event_date: datetime = field(default_factory=_now_utc)
metadata: dict[str, str] = field(default_factory=dict)
entities: list[dict[str, str]] = field(default_factory=list) # User-provided entities
tags: list[str] = field(default_factory=list) # Visibility scope tags
@dataclass
@@ -37,6 +61,7 @@ class ChunkMetadata:
Used to track which facts were extracted from which chunks.
"""
chunk_text: str
fact_count: int
content_index: int # Index of the source content
@@ -50,9 +75,10 @@ class EntityRef:
Entities are extracted by the LLM during fact extraction.
"""
name: str
canonical_name: Optional[str] = None # Resolved canonical name
entity_id: Optional[UUID] = None # Resolved entity ID
canonical_name: str | None = None # Resolved canonical name
entity_id: UUID | None = None # Resolved entity ID
@dataclass
@@ -62,6 +88,7 @@ class CausalRelation:
Represents how one fact causes, enables, or prevents another.
"""
relation_type: str # "causes", "enables", "prevents", "caused_by"
target_fact_index: int # Index of the target fact in the batch
strength: float = 1.0 # Strength of the causal relationship
@@ -74,20 +101,22 @@ class ExtractedFact:
This is the raw output from fact extraction before processing.
"""
fact_text: str
fact_type: str # "world", "experience", "opinion", "observation"
entities: List[str] = field(default_factory=list)
occurred_start: Optional[datetime] = None
occurred_end: Optional[datetime] = None
where: Optional[str] = None # WHERE the fact occurred or is about
causal_relations: List[CausalRelation] = field(default_factory=list)
entities: list[str] = field(default_factory=list)
occurred_start: datetime | None = None
occurred_end: datetime | None = None
where: str | None = None # WHERE the fact occurred or is about
causal_relations: list[CausalRelation] = field(default_factory=list)
# Context from the content item
content_index: int = 0 # Which content this fact came from
chunk_index: int = 0 # Which chunk this fact came from
context: str = ""
mentioned_at: Optional[datetime] = None
metadata: Dict[str, str] = field(default_factory=dict)
mentioned_at: datetime | None = None
metadata: dict[str, str] = field(default_factory=dict)
tags: list[str] = field(default_factory=list) # Visibility scope tags
@dataclass
@@ -97,37 +126,44 @@ class ProcessedFact:
Includes resolved entities, embeddings, and all necessary fields.
"""
# Core fact data
fact_text: str
fact_type: str
embedding: List[float]
embedding: list[float]
# Temporal data
occurred_start: Optional[datetime]
occurred_end: Optional[datetime]
occurred_start: datetime | None
occurred_end: datetime | None
mentioned_at: datetime
# Context and metadata
context: str
metadata: Dict[str, str]
metadata: dict[str, str]
# Location data
where: Optional[str] = None
where: str | None = None
# Entities
entities: List[EntityRef] = field(default_factory=list)
entities: list[EntityRef] = field(default_factory=list)
# Causal relations
causal_relations: List[CausalRelation] = field(default_factory=list)
causal_relations: list[CausalRelation] = field(default_factory=list)
# Chunk reference
chunk_id: Optional[str] = None
chunk_id: str | None = None
# Document reference (denormalized for query performance)
document_id: Optional[str] = None
document_id: str | None = None
# DB fields (set after insertion)
unit_id: Optional[UUID] = None
unit_id: UUID | None = None
# Track which content this fact came from (for user entity merging)
content_index: int = 0
# Visibility scope tags
tags: list[str] = field(default_factory=list)
@property
def is_duplicate(self) -> bool:
@@ -136,10 +172,8 @@ class ProcessedFact:
@staticmethod
def from_extracted_fact(
extracted_fact: 'ExtractedFact',
embedding: List[float],
chunk_id: Optional[str] = None
) -> 'ProcessedFact':
extracted_fact: "ExtractedFact", embedding: list[float], chunk_id: str | None = None
) -> "ProcessedFact":
"""
Create ProcessedFact from ExtractedFact.
@@ -151,12 +185,12 @@ class ProcessedFact:
Returns:
ProcessedFact ready for storage
"""
from datetime import datetime, timezone
from datetime import datetime
# Use occurred dates only if explicitly provided by LLM
occurred_start = extracted_fact.occurred_start
occurred_end = extracted_fact.occurred_end
mentioned_at = extracted_fact.mentioned_at or datetime.now(timezone.utc)
mentioned_at = extracted_fact.mentioned_at or datetime.now(UTC)
# Convert entity strings to EntityRef objects
entities = [EntityRef(name=name) for name in extracted_fact.entities]
@@ -172,7 +206,9 @@ class ProcessedFact:
metadata=extracted_fact.metadata,
entities=entities,
causal_relations=extracted_fact.causal_relations,
chunk_id=chunk_id
chunk_id=chunk_id,
content_index=extracted_fact.content_index,
tags=extracted_fact.tags,
)
@@ -183,10 +219,11 @@ class EntityLink:
Used for entity-based graph connections in the memory graph.
"""
from_unit_id: UUID
to_unit_id: UUID
entity_id: UUID
link_type: str = 'entity'
link_type: str = "entity"
weight: float = 1.0
@@ -197,24 +234,26 @@ class RetainBatch:
Tracks all facts, chunks, and metadata for a batch operation.
"""
bank_id: str
contents: List[RetainContent]
document_id: Optional[str] = None
fact_type_override: Optional[str] = None
confidence_score: Optional[float] = None
contents: list[RetainContent]
document_id: str | None = None
fact_type_override: str | None = None
confidence_score: float | None = None
document_tags: list[str] = field(default_factory=list) # Tags applied to all items
# Extracted data (populated during processing)
extracted_facts: List[ExtractedFact] = field(default_factory=list)
processed_facts: List[ProcessedFact] = field(default_factory=list)
chunks: List[ChunkMetadata] = field(default_factory=list)
extracted_facts: list[ExtractedFact] = field(default_factory=list)
processed_facts: list[ProcessedFact] = field(default_factory=list)
chunks: list[ChunkMetadata] = field(default_factory=list)
# Results (populated after storage)
unit_ids_by_content: List[List[str]] = field(default_factory=list)
unit_ids_by_content: list[list[str]] = field(default_factory=list)
def get_facts_for_content(self, content_index: int) -> List[ExtractedFact]:
def get_facts_for_content(self, content_index: int) -> list[ExtractedFact]:
"""Get all extracted facts for a specific content item."""
return [f for f in self.extracted_facts if f.content_index == content_index]
def get_chunks_for_content(self, content_index: int) -> List[ChunkMetadata]:
def get_chunks_for_content(self, content_index: int) -> list[ChunkMetadata]:
"""Get all chunks for a specific content item."""
return [c for c in self.chunks if c.content_index == content_index]
@@ -3,13 +3,27 @@ Search module for memory retrieval.
Provides modular search architecture:
- Retrieval: 4-way parallel (semantic + BM25 + graph + temporal)
- Graph retrieval: Pluggable strategies (BFS, PPR)
- Reranking: Pluggable strategies (heuristic, cross-encoder)
"""
from .retrieval import retrieve_parallel
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
from .mpfp_retrieval import MPFPGraphRetriever
from .reranking import CrossEncoderReranker
from .retrieval import (
ParallelRetrievalResult,
get_default_graph_retriever,
retrieve_parallel,
set_default_graph_retriever,
)
__all__ = [
"retrieve_parallel",
"get_default_graph_retriever",
"set_default_graph_retriever",
"ParallelRetrievalResult",
"GraphRetriever",
"BFSGraphRetriever",
"MPFPGraphRetriever",
"CrossEncoderReranker",
]
@@ -2,15 +2,12 @@
Helper functions for hybrid search (semantic + BM25 + graph).
"""
from typing import List, Dict, Any, Tuple
import asyncio
from .types import RetrievalResult, MergedCandidate
from typing import Any
from .types import MergedCandidate, RetrievalResult
def reciprocal_rank_fusion(
result_lists: List[List[RetrievalResult]],
k: int = 60
) -> List[MergedCandidate]:
def reciprocal_rank_fusion(result_lists: list[list[RetrievalResult]], k: int = 60) -> list[MergedCandidate]:
"""
Merge multiple ranked result lists using Reciprocal Rank Fusion.
@@ -73,20 +70,14 @@ def reciprocal_rank_fusion(
sorted(rrf_scores.items(), key=lambda x: x[1], reverse=True), start=1
):
merged_candidate = MergedCandidate(
retrieval=all_retrievals[doc_id],
rrf_score=rrf_score,
rrf_rank=rrf_rank,
source_ranks=source_ranks[doc_id]
retrieval=all_retrievals[doc_id], rrf_score=rrf_score, rrf_rank=rrf_rank, source_ranks=source_ranks[doc_id]
)
merged_results.append(merged_candidate)
return merged_results
def normalize_scores_on_deltas(
results: List[Dict[str, Any]],
score_keys: List[str]
) -> List[Dict[str, Any]]:
def normalize_scores_on_deltas(results: list[dict[str, Any]], score_keys: list[str]) -> list[dict[str, Any]]:
"""
Normalize scores based on deltas (min-max normalization within result set).
@@ -0,0 +1,264 @@
"""
Graph retrieval strategies for memory recall.
This module provides an abstraction for graph-based memory retrieval,
allowing different algorithms (BFS spreading activation, PPR, etc.) to be
swapped without changing the rest of the recall pipeline.
"""
import logging
from abc import ABC, abstractmethod
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .tags import TagsMatch, filter_results_by_tags
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
class GraphRetriever(ABC):
"""
Abstract base class for graph-based memory retrieval.
Implementations traverse the memory graph (entity links, temporal links,
causal links) to find relevant facts that might not be found by
semantic or keyword search alone.
"""
@property
@abstractmethod
def name(self) -> str:
"""Return identifier for this retrieval strategy (e.g., 'bfs', 'mpfp')."""
pass
@abstractmethod
async def retrieve(
self,
pool,
query_embedding_str: str,
bank_id: str,
fact_type: str,
budget: int,
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None, # TypedAdjacency, optional pre-loaded graph
tags: list[str] | None = None, # Visibility scope tags for filtering
tags_match: TagsMatch = "any", # How to match tags: 'any' (OR) or 'all' (AND)
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve relevant facts via graph traversal.
Args:
pool: Database connection pool
query_embedding_str: Query embedding as string (for finding entry points)
bank_id: Memory bank identifier
fact_type: Fact type to filter ('world', 'experience', 'opinion', 'observation')
budget: Maximum number of nodes to explore/return
query_text: Original query text (optional, for some strategies)
semantic_seeds: Pre-computed semantic entry points (from semantic retrieval)
temporal_seeds: Pre-computed temporal entry points (from temporal retrieval)
adjacency: Pre-loaded typed adjacency graph (optional, for MPFP)
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
Tuple of (List of RetrievalResult with activation scores, optional timing info)
"""
pass
class BFSGraphRetriever(GraphRetriever):
"""
Graph retrieval using BFS-style spreading activation.
Starting from semantic entry points, spreads activation through
the memory graph (entity, temporal, causal links) using breadth-first
traversal with decaying activation.
This is the original Hindsight graph retrieval algorithm.
"""
def __init__(
self,
entry_point_limit: int = 5,
entry_point_threshold: float = 0.5,
activation_decay: float = 0.8,
min_activation: float = 0.1,
batch_size: int = 20,
):
"""
Initialize BFS graph retriever.
Args:
entry_point_limit: Maximum number of entry points to start from
entry_point_threshold: Minimum semantic similarity for entry points
activation_decay: Decay factor per hop (activation *= decay)
min_activation: Minimum activation to continue spreading
batch_size: Number of nodes to process per batch (for neighbor fetching)
"""
self.entry_point_limit = entry_point_limit
self.entry_point_threshold = entry_point_threshold
self.activation_decay = activation_decay
self.min_activation = min_activation
self.batch_size = batch_size
@property
def name(self) -> str:
return "bfs"
async def retrieve(
self,
pool,
query_embedding_str: str,
bank_id: str,
fact_type: str,
budget: int,
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None, # Not used by BFS
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve facts using BFS spreading activation.
Algorithm:
1. Find entry points (top semantic matches above threshold)
2. BFS traversal: visit neighbors, propagate decaying activation
3. Boost causal links (causes, enables, prevents)
4. Return visited nodes up to budget
Note: BFS finds its own entry points via embedding search.
The semantic_seeds, temporal_seeds, and adjacency parameters are accepted
for interface compatibility but not used.
"""
async with acquire_with_retry(pool) as conn:
results = await self._retrieve_with_conn(
conn, query_embedding_str, bank_id, fact_type, budget, tags=tags, tags_match=tags_match
)
return results, None
async def _retrieve_with_conn(
self,
conn,
query_embedding_str: str,
bank_id: str,
fact_type: str,
budget: int,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> list[RetrievalResult]:
"""Internal implementation with connection."""
from .tags import build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_embedding_str, bank_id, fact_type, self.entry_point_threshold, self.entry_point_limit]
if tags:
params.append(tags)
# Step 1: Find entry points
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,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= $4
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
*params,
)
if not entry_points:
logger.debug(
f"[BFS] No entry points found for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
)
return []
logger.debug(
f"[BFS] Found {len(entry_points)} entry points for fact_type={fact_type} "
f"(tags={tags}, tags_match={tags_match})"
)
# Step 2: BFS spreading activation
visited = set()
results = []
queue = [(RetrievalResult.from_db_row(dict(r)), r["similarity"]) for r in entry_points]
budget_remaining = budget
while queue and budget_remaining > 0:
# Collect a batch of nodes to process
batch_nodes = []
batch_activations = {}
while queue and len(batch_nodes) < self.batch_size and budget_remaining > 0:
current, activation = queue.pop(0)
unit_id = current.id
if unit_id not in visited:
visited.add(unit_id)
budget_remaining -= 1
current.activation = activation
results.append(current)
batch_nodes.append(current.id)
batch_activations[unit_id] = activation
# Batch fetch neighbors
if batch_nodes and budget_remaining > 0:
max_neighbors = len(batch_nodes) * 20
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.document_id, mu.chunk_id, mu.tags,
ml.weight, ml.link_type, ml.from_unit_id
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.weight >= $2
AND mu.fact_type = $3
ORDER BY ml.weight DESC
LIMIT $4
""",
batch_nodes,
self.min_activation,
fact_type,
max_neighbors,
)
for n in neighbors:
neighbor_id = str(n["id"])
if neighbor_id not in visited:
parent_id = str(n["from_unit_id"])
parent_activation = batch_activations.get(parent_id, 0.5)
# Boost causal links
link_type = n["link_type"]
base_weight = n["weight"]
if link_type in ("causes", "caused_by"):
causal_boost = 2.0
elif link_type in ("enables", "prevents"):
causal_boost = 1.5
else:
causal_boost = 1.0
effective_weight = base_weight * causal_boost
new_activation = parent_activation * effective_weight * self.activation_decay
if new_activation > self.min_activation:
neighbor_result = RetrievalResult.from_db_row(dict(n))
queue.append((neighbor_result, new_activation))
# Apply tags filtering (BFS may traverse into memories that don't match tags criteria)
if tags:
results = filter_results_by_tags(results, tags, match=tags_match)
return results
@@ -0,0 +1,256 @@
"""
Link Expansion graph retrieval.
A simple, fast graph retrieval that expands from seeds via:
1. Entity links: Find facts sharing entities with seeds (filtered by entity frequency)
2. Causal links: Find facts causally linked to seeds (top-k by weight)
Characteristics:
- 2-3 DB queries (seed finding + parallel entity/causal expansion)
- Sublinear: only touches connected facts via indexes
- No iteration, no propagation, no normalization
- Target: <100ms
"""
import logging
import time
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .graph_retrieval import GraphRetriever
from .tags import TagsMatch, filter_results_by_tags
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
async def _find_semantic_seeds(
conn,
query_embedding_str: str,
bank_id: str,
fact_type: str,
limit: int = 20,
threshold: float = 0.3,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> list[RetrievalResult]:
"""Find semantic seeds via embedding search."""
from .tags import build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
if tags:
params.append(tags)
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,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= $4
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
*params,
)
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
class LinkExpansionRetriever(GraphRetriever):
"""
Graph retrieval via direct link expansion from seeds.
Expands through entity co-occurrence and causal links in a single query.
Fast and simple alternative to MPFP.
"""
def __init__(
self,
max_entity_frequency: int = 500,
causal_weight_threshold: float = 0.3,
causal_limit_per_seed: int = 10,
):
"""
Initialize link expansion retriever.
Args:
max_entity_frequency: Skip entities appearing in more than this many facts
causal_weight_threshold: Minimum weight for causal links
causal_limit_per_seed: Max causal links to follow per seed
"""
self.max_entity_frequency = max_entity_frequency
self.causal_weight_threshold = causal_weight_threshold
self.causal_limit_per_seed = causal_limit_per_seed
@property
def name(self) -> str:
return "link_expansion"
async def retrieve(
self,
pool,
query_embedding_str: str,
bank_id: str,
fact_type: str,
budget: int,
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve facts by expanding links from seeds.
Args:
pool: Database connection pool
query_embedding_str: Query embedding (unused, kept for interface)
bank_id: Memory bank ID
fact_type: Fact type to filter
budget: Maximum results to return
query_text: Original query text (unused)
semantic_seeds: Pre-computed semantic entry points
temporal_seeds: Pre-computed temporal entry points
adjacency: Unused, kept for interface compatibility
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
Tuple of (results, timings)
"""
start_time = time.time()
timings = MPFPTimings(fact_type=fact_type)
# Use single connection for all queries to reduce pool pressure
# (queries are fast ~50ms each, connection acquisition is the bottleneck)
async with acquire_with_retry(pool) as conn:
# Find seeds if not provided
if semantic_seeds:
all_seeds = list(semantic_seeds)
else:
seeds_start = time.time()
all_seeds = await _find_semantic_seeds(
conn,
query_embedding_str,
bank_id,
fact_type,
limit=20,
threshold=0.3,
tags=tags,
tags_match=tags_match,
)
timings.seeds_time = time.time() - seeds_start
logger.debug(
f"[LinkExpansion] Found {len(all_seeds)} semantic seeds for fact_type={fact_type} "
f"(tags={tags}, tags_match={tags_match})"
)
# Add temporal seeds if provided
if temporal_seeds:
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})
timings.pattern_count = len(seed_ids)
# 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,
)
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.fact_type, mu.document_id, mu.chunk_id, mu.tags,
ml.weight + 1.0 AS score
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.link_type IN ('causes', 'caused_by', 'enables', 'prevents')
AND ml.weight >= $2
AND mu.fact_type = $3
ORDER BY mu.id, ml.weight DESC
LIMIT $4
""",
seed_ids,
self.causal_weight_threshold,
fact_type,
budget,
)
timings.edge_load_time = time.time() - query_start
timings.db_queries = 2
timings.edge_count = len(entity_rows) + len(causal_rows)
# Merge results, taking max score per fact
score_map: dict[str, float] = {}
row_map: dict[str, dict] = {}
for row in entity_rows:
fact_id = str(row["id"])
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
row_map[fact_id] = dict(row)
for row in causal_rows:
fact_id = str(row["id"])
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
if fact_id not in row_map:
row_map[fact_id] = dict(row)
# Sort by score and limit
sorted_ids = sorted(score_map.keys(), key=lambda x: score_map[x], reverse=True)[:budget]
rows = [row_map[fact_id] for fact_id in sorted_ids]
# Convert to results
results = []
for row in rows:
result = RetrievalResult.from_db_row(dict(row))
result.activation = row["score"]
results.append(result)
# Apply tags filtering (graph expansion may reach untagged memories)
if tags:
results = filter_results_by_tags(results, tags, match=tags_match)
timings.result_count = len(results)
timings.traverse = time.time() - start_time
logger.debug(
f"LinkExpansion: {len(results)} results from {len(seed_ids)} seeds "
f"in {timings.traverse * 1000:.1f}ms (query: {timings.edge_load_time * 1000:.1f}ms)"
)
return results, timings
@@ -0,0 +1,684 @@
"""
Meta-Path Forward Push (MPFP) graph retrieval.
A sublinear graph traversal algorithm for memory retrieval over heterogeneous
graphs with multiple edge types (semantic, temporal, causal, entity).
Combines meta-path patterns from HIN literature with Forward Push local
propagation from Approximate PPR.
Key properties:
- Sublinear in graph size (threshold pruning bounds active nodes)
- Lazy edge loading: only loads edges for frontier nodes, not entire graph
- Predefined patterns capture different retrieval intents
- All patterns run in parallel, results fused via RRF
- No LLM in the loop during traversal
"""
import asyncio
import logging
from collections import defaultdict
from dataclasses import dataclass, field
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .graph_retrieval import GraphRetriever
from .tags import TagsMatch
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
# -----------------------------------------------------------------------------
# Data Classes
# -----------------------------------------------------------------------------
@dataclass
class EdgeTarget:
"""A neighbor node with its edge weight."""
node_id: str
weight: float
@dataclass
class EdgeCache:
"""
Cache for lazily-loaded edges.
Grows per-hop as edges are loaded for frontier nodes.
Shared across patterns to avoid redundant loads.
Loads ALL edge types at once to minimize DB queries.
Thread-safe via asyncio lock to prevent redundant concurrent loads.
"""
# edge_type -> from_node_id -> list of EdgeTarget
graphs: dict[str, dict[str, list[EdgeTarget]]] = field(default_factory=dict)
# Track which nodes have been fully loaded (all edge types)
_fully_loaded: set[str] = field(default_factory=set)
# Timing stats
db_queries: int = 0
edge_load_time: float = 0.0
# Detailed hop timing for debugging
hop_details: list[dict] = field(default_factory=list)
# Lock to prevent redundant concurrent loads
_lock: asyncio.Lock = field(default_factory=asyncio.Lock)
def get_neighbors(self, edge_type: str, node_id: str) -> list[EdgeTarget]:
"""Get neighbors for a node via a specific edge type."""
return self.graphs.get(edge_type, {}).get(node_id, [])
def get_normalized_neighbors(self, edge_type: str, node_id: str, top_k: int) -> list[EdgeTarget]:
"""Get top-k neighbors with weights normalized to sum to 1."""
neighbors = self.get_neighbors(edge_type, node_id)[:top_k]
if not neighbors:
return []
total = sum(n.weight for n in neighbors)
if total == 0:
return []
return [EdgeTarget(node_id=n.node_id, weight=n.weight / total) for n in neighbors]
def is_fully_loaded(self, node_id: str) -> bool:
"""Check if all edges for this node have been loaded."""
return node_id in self._fully_loaded
def get_uncached(self, node_ids: list[str]) -> list[str]:
"""Get node IDs that haven't been fully loaded yet."""
return [n for n in node_ids if not self.is_fully_loaded(n)]
def add_all_edges(self, edges_by_type: dict[str, dict[str, list[EdgeTarget]]], all_queried: list[str]):
"""
Add loaded edges to the cache (all edge types at once).
Args:
edges_by_type: Dict mapping edge_type -> from_node_id -> list of EdgeTarget
all_queried: All node IDs that were queried (marks them as fully loaded)
"""
for edge_type, edges in edges_by_type.items():
if edge_type not in self.graphs:
self.graphs[edge_type] = {}
for node_id, neighbors in edges.items():
self.graphs[edge_type][node_id] = neighbors
# Mark all queried nodes as fully loaded (even if they have no edges)
self._fully_loaded.update(all_queried)
@dataclass
class PatternResult:
"""Result from a single pattern traversal."""
pattern: list[str]
scores: dict[str, float] # node_id -> accumulated mass
@dataclass
class MPFPConfig:
"""Configuration for MPFP algorithm."""
alpha: float = 0.15 # teleport/keep probability
threshold: float = 1e-6 # mass pruning threshold (lower = explore more)
top_k_neighbors: int = 20 # fan-out limit per node
# Patterns from semantic seeds
patterns_semantic: list[list[str]] = field(
default_factory=lambda: [
["semantic", "semantic"], # topic expansion
["entity", "temporal"], # entity timeline
["semantic", "causes"], # reasoning chains (forward)
["semantic", "caused_by"], # reasoning chains (backward)
["entity", "semantic"], # entity context
]
)
# Patterns from temporal seeds
patterns_temporal: list[list[str]] = field(
default_factory=lambda: [
["temporal", "semantic"], # what was happening then
["temporal", "entity"], # who was involved then
]
)
@dataclass
class SeedNode:
"""An entry point node with its initial score."""
node_id: str
score: float # initial mass (e.g., similarity score)
# -----------------------------------------------------------------------------
# Lazy Edge Loading
# -----------------------------------------------------------------------------
async def load_all_edges_for_frontier(
pool,
node_ids: list[str],
top_k_per_type: int = 20,
) -> dict[str, dict[str, list[EdgeTarget]]]:
"""
Load top-k edges per (node, edge_type) for frontier nodes.
Uses a LATERAL join to efficiently fetch only the top-k edges per type,
avoiding loading hundreds of entity edges when only 20 are needed.
Requires composite index: (from_unit_id, link_type, weight DESC)
Args:
pool: Database connection pool
node_ids: Frontier node IDs to load edges for
top_k_per_type: Max edges to load per (node, link_type) pair
Returns:
Dict mapping edge_type -> from_node_id -> list of EdgeTarget
"""
if not node_ids:
return {}
async with acquire_with_retry(pool) as conn:
# Use LATERAL join to get top-k per (from_node, link_type)
# This leverages the composite index for efficient early termination
rows = await conn.fetch(
f"""
WITH frontier(node_id) AS (SELECT unnest($1::uuid[]))
SELECT f.node_id as from_unit_id, lt.link_type, edges.to_unit_id, edges.weight
FROM frontier f
CROSS JOIN (VALUES ('semantic'), ('temporal'), ('entity'), ('causes'), ('caused_by')) AS lt(link_type)
CROSS JOIN LATERAL (
SELECT ml.to_unit_id, ml.weight
FROM {fq_table("memory_links")} ml
WHERE ml.from_unit_id = f.node_id
AND ml.link_type = lt.link_type
AND ml.weight >= 0.1
ORDER BY ml.weight DESC
LIMIT $2
) edges
""",
node_ids,
top_k_per_type,
)
# Group by edge_type -> from_node -> neighbors
result: dict[str, dict[str, list[EdgeTarget]]] = defaultdict(lambda: defaultdict(list))
for row in rows:
edge_type = row["link_type"]
from_id = str(row["from_unit_id"])
to_id = str(row["to_unit_id"])
weight = row["weight"]
result[edge_type][from_id].append(EdgeTarget(node_id=to_id, weight=weight))
# Convert nested defaultdicts to regular dicts
return {edge_type: dict(edges) for edge_type, edges in result.items()}
# -----------------------------------------------------------------------------
# Core Algorithm (Async with Lazy Loading)
# -----------------------------------------------------------------------------
@dataclass
class PatternState:
"""State for a pattern traversal between hops."""
pattern: list[str]
hop_index: int
scores: dict[str, float]
frontier: dict[str, float]
def _init_pattern_state(seeds: list[SeedNode], pattern: list[str]) -> PatternState:
"""Initialize pattern state from seeds."""
if not seeds:
return PatternState(pattern=pattern, hop_index=0, scores={}, frontier={})
total_seed_score = sum(s.score for s in seeds)
if total_seed_score == 0:
total_seed_score = len(seeds)
frontier = {s.node_id: s.score / total_seed_score for s in seeds}
return PatternState(pattern=pattern, hop_index=0, scores={}, frontier=frontier)
def _execute_hop(state: PatternState, cache: EdgeCache, config: MPFPConfig) -> set[str]:
"""
Execute ONE hop of traversal, return frontier nodes for next hop.
This is a pure function that uses cached edges (no DB access).
Returns set of uncached nodes needed for next hop.
"""
if state.hop_index >= len(state.pattern):
return set()
edge_type = state.pattern[state.hop_index]
# Collect active nodes above threshold
active_nodes = [node_id for node_id, mass in state.frontier.items() if mass >= config.threshold]
if not active_nodes:
state.frontier = {}
return set()
# Propagate mass using cached edges
next_frontier: dict[str, float] = {}
uncached_for_next: set[str] = set()
for node_id, mass in state.frontier.items():
if mass < config.threshold:
continue
# Keep α portion for this node
state.scores[node_id] = state.scores.get(node_id, 0) + config.alpha * mass
# Push (1-α) to neighbors
push_mass = (1 - config.alpha) * mass
neighbors = cache.get_normalized_neighbors(edge_type, node_id, config.top_k_neighbors)
for neighbor in neighbors:
next_frontier[neighbor.node_id] = next_frontier.get(neighbor.node_id, 0) + push_mass * neighbor.weight
# Track if we'll need edges for this node in the next hop
if not cache.is_fully_loaded(neighbor.node_id):
uncached_for_next.add(neighbor.node_id)
state.frontier = next_frontier
state.hop_index += 1
return uncached_for_next
def _finalize_pattern(state: PatternState, config: MPFPConfig) -> PatternResult:
"""Finalize pattern by adding remaining frontier mass to scores."""
for node_id, mass in state.frontier.items():
if mass >= config.threshold:
state.scores[node_id] = state.scores.get(node_id, 0) + mass
return PatternResult(pattern=state.pattern, scores=state.scores)
async def mpfp_traverse_hop_synchronized(
pool,
pattern_jobs: list[tuple[list[SeedNode], list[str]]],
config: MPFPConfig,
cache: EdgeCache,
) -> list[PatternResult]:
"""
Execute ALL patterns with hop-synchronized edge loading.
Instead of running each pattern independently (causing multiple DB queries),
this function:
1. Runs hop 1 for ALL patterns (using pre-warmed seed edges)
2. Collects ALL unique hop-2 frontier nodes across patterns
3. Pre-warms hop-2 edges in ONE query
4. Runs hop 2 for ALL patterns
This reduces DB queries from O(patterns * hops) to O(hops).
Args:
pool: Database connection pool
pattern_jobs: List of (seeds, pattern) tuples
config: Algorithm parameters
cache: Shared edge cache (should be pre-warmed with seed edges)
Returns:
List of PatternResult for each pattern
"""
import time
# Initialize all pattern states
states = [_init_pattern_state(seeds, pattern) for seeds, pattern in pattern_jobs]
# Determine max hops (all patterns should be same length, but be safe)
max_hops = max((len(p) for _, p in pattern_jobs), default=0)
# Detailed timing for debugging
hop_times: list[dict] = []
# Execute hop-by-hop across ALL patterns
for hop in range(max_hops):
hop_start = time.time()
hop_timing = {"hop": hop, "patterns_executed": 0, "uncached_count": 0, "load_time": 0.0}
# Execute this hop for all patterns, collect uncached nodes for next hop
all_uncached: set[str] = set()
exec_start = time.time()
for state in states:
if state.hop_index < len(state.pattern):
uncached = _execute_hop(state, cache, config)
all_uncached.update(uncached)
hop_timing["patterns_executed"] += 1
hop_timing["exec_time"] = time.time() - exec_start
# Pre-warm edges for ALL uncached nodes before next hop
hop_timing["uncached_count"] = len(all_uncached)
if all_uncached:
uncached_list = list(all_uncached - cache._fully_loaded)
hop_timing["uncached_after_filter"] = len(uncached_list)
if uncached_list:
load_start = time.time()
edges_by_type = await load_all_edges_for_frontier(pool, uncached_list, config.top_k_neighbors)
hop_timing["load_time"] = time.time() - load_start
cache.edge_load_time += hop_timing["load_time"]
cache.db_queries += 1
cache.add_all_edges(edges_by_type, uncached_list)
hop_timing["edges_loaded"] = sum(
len(neighbors) for edges in edges_by_type.values() for neighbors in edges.values()
)
hop_timing["total_time"] = time.time() - hop_start
hop_times.append(hop_timing)
# Store hop timing details in cache for logging
cache.hop_details = hop_times
# Finalize all patterns
return [_finalize_pattern(state, config) for state in states]
async def mpfp_traverse_async(
pool,
seeds: list[SeedNode],
pattern: list[str],
config: MPFPConfig,
cache: EdgeCache,
) -> PatternResult:
"""
Async Forward Push traversal with lazy edge loading.
NOTE: For better performance with multiple patterns, use mpfp_traverse_hop_synchronized().
This function is kept for single-pattern use cases.
"""
if not seeds:
return PatternResult(pattern=pattern, scores={})
results = await mpfp_traverse_hop_synchronized(pool, [(seeds, pattern)], config, cache)
return results[0] if results else PatternResult(pattern=pattern, scores={})
def rrf_fusion(
results: list[PatternResult],
k: int = 60,
top_k: int = 50,
) -> list[tuple[str, float]]:
"""
Reciprocal Rank Fusion to combine pattern results.
Args:
results: List of pattern results
k: RRF constant (higher = more uniform weighting)
top_k: Number of results to return
Returns:
List of (node_id, fused_score) tuples, sorted by score descending
"""
fused: dict[str, float] = {}
for result in results:
if not result.scores:
continue
# Rank nodes by their score in this pattern
ranked = sorted(result.scores.keys(), key=lambda n: result.scores[n], reverse=True)
for rank, node_id in enumerate(ranked):
fused[node_id] = fused.get(node_id, 0) + 1.0 / (k + rank + 1)
# Sort by fused score and return top-k
sorted_results = sorted(fused.items(), key=lambda x: x[1], reverse=True)
return sorted_results[:top_k]
# -----------------------------------------------------------------------------
# Database Loading
# -----------------------------------------------------------------------------
async def fetch_memory_units_by_ids(
pool,
node_ids: list[str],
fact_type: str,
) -> list[RetrievalResult]:
"""Fetch full memory unit details for a list of node IDs."""
if not node_ids:
return []
async with acquire_with_retry(pool) as conn:
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
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
AND fact_type = $2
""",
node_ids,
fact_type,
)
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
# -----------------------------------------------------------------------------
# Graph Retriever Implementation
# -----------------------------------------------------------------------------
class MPFPGraphRetriever(GraphRetriever):
"""
Graph retrieval using Meta-Path Forward Push with lazy edge loading.
Runs predefined patterns in parallel from semantic and temporal seeds,
loading edges on-demand per hop instead of loading entire graph upfront.
"""
def __init__(self, config: MPFPConfig | None = None):
"""
Initialize MPFP retriever.
Args:
config: Algorithm configuration (uses defaults if None)
"""
if config is None:
# Read top_k_neighbors from global config
from ...config import get_config
global_config = get_config()
config = MPFPConfig(top_k_neighbors=global_config.mpfp_top_k_neighbors)
self.config = config
@property
def name(self) -> str:
return "mpfp"
async def retrieve(
self,
pool,
query_embedding_str: str,
bank_id: str,
fact_type: str,
budget: int,
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None, # Ignored - kept for interface compatibility
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve facts using MPFP algorithm with lazy edge loading.
Args:
pool: Database connection pool
query_embedding_str: Query embedding (used for fallback seed finding)
bank_id: Memory bank ID
fact_type: Fact type to filter
budget: Maximum results to return
query_text: Original query text (optional)
semantic_seeds: Pre-computed semantic entry points
temporal_seeds: Pre-computed temporal entry points
adjacency: Ignored (kept for interface compatibility)
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
Tuple of (List of RetrievalResult with activation scores, MPFPTimings)
"""
import time
timings = MPFPTimings(fact_type=fact_type)
# Convert seeds to SeedNode format
semantic_seed_nodes = self._convert_seeds(semantic_seeds, "similarity")
temporal_seed_nodes = self._convert_seeds(temporal_seeds, "temporal_score")
# If no semantic seeds provided, fall back to finding our own
if not semantic_seed_nodes:
seeds_start = time.time()
semantic_seed_nodes = await self._find_semantic_seeds(
pool, query_embedding_str, bank_id, fact_type, tags=tags, tags_match=tags_match
)
timings.seeds_time = time.time() - seeds_start
logger.debug(
f"[MPFP] Found {len(semantic_seed_nodes)} semantic seeds for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
)
# Collect all pattern jobs
pattern_jobs = []
# Patterns from semantic seeds
for pattern in self.config.patterns_semantic:
if semantic_seed_nodes:
pattern_jobs.append((semantic_seed_nodes, pattern))
# Patterns from temporal seeds
for pattern in self.config.patterns_temporal:
if temporal_seed_nodes:
pattern_jobs.append((temporal_seed_nodes, pattern))
if not pattern_jobs:
logger.debug(
f"[MPFP] No pattern jobs (semantic_seeds={len(semantic_seed_nodes)}, temporal_seeds={len(temporal_seed_nodes)})"
)
return [], timings
timings.pattern_count = len(pattern_jobs)
# Shared edge cache across all patterns
cache = EdgeCache()
# Pre-warm cache with ALL seed node edges BEFORE running patterns
# This prevents redundant DB queries at hop 1
all_seed_ids = list({s.node_id for seeds, _ in pattern_jobs for s in seeds})
if all_seed_ids:
import time as time_module
prewarm_start = time_module.time()
edges_by_type = await load_all_edges_for_frontier(pool, all_seed_ids, self.config.top_k_neighbors)
cache.edge_load_time += time_module.time() - prewarm_start
cache.db_queries += 1
cache.add_all_edges(edges_by_type, all_seed_ids)
# Run all patterns with HOP-SYNCHRONIZED edge loading
# This batches hop-2 edge loads across ALL patterns into ONE query
# Reduces DB queries from O(patterns * hops) to O(hops)
step_start = time.time()
pattern_results = await mpfp_traverse_hop_synchronized(pool, pattern_jobs, self.config, cache)
timings.traverse = time.time() - step_start
# Record edge loading stats from cache
timings.edge_count = sum(len(neighbors) for g in cache.graphs.values() for neighbors in g.values())
timings.db_queries = cache.db_queries
timings.edge_load_time = cache.edge_load_time
timings.hop_details = cache.hop_details
# Fuse results
step_start = time.time()
fused = rrf_fusion(pattern_results, top_k=budget)
timings.fusion = time.time() - step_start
if not fused:
logger.debug(f"[MPFP] No fused results after RRF fusion (pattern_count={len(pattern_results)})")
return [], timings
# Get top result IDs
result_ids = [node_id for node_id, score in fused][:budget]
# Fetch full details
step_start = time.time()
results = await fetch_memory_units_by_ids(pool, result_ids, fact_type)
timings.fetch = time.time() - step_start
# Filter results by tags (graph traversal may have picked up unfiltered memories)
if tags:
from .tags import filter_results_by_tags
results = filter_results_by_tags(results, tags, match=tags_match)
timings.result_count = len(results)
# Add activation scores from fusion
score_map = {node_id: score for node_id, score in fused}
for result in results:
result.activation = score_map.get(result.id, 0.0)
# Sort by activation
results.sort(key=lambda r: r.activation or 0, reverse=True)
return results, timings
def _convert_seeds(
self,
seeds: list[RetrievalResult] | None,
score_attr: str,
) -> list[SeedNode]:
"""Convert RetrievalResult seeds to SeedNode format."""
if not seeds:
return []
result = []
for seed in seeds:
score = getattr(seed, score_attr, None)
if score is None:
score = seed.activation or seed.similarity or 1.0
result.append(SeedNode(node_id=seed.id, score=score))
return result
async def _find_semantic_seeds(
self,
pool,
query_embedding_str: str,
bank_id: str,
fact_type: str,
limit: int = 20,
threshold: float = 0.3,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> list[SeedNode]:
"""Fallback: find semantic seeds via embedding search."""
from .tags import build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
if tags:
params.append(tags)
async with acquire_with_retry(pool) as conn:
rows = await conn.fetch(
f"""
SELECT id, 1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= $4
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
*params,
)
return [SeedNode(node_id=str(r["id"]), score=r["similarity"]) for r in rows]

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