Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ff1886eeff |
@@ -1,7 +1,6 @@
|
||||
{
|
||||
"$schema": "https://anthropic.com/claude-code/marketplace.schema.json",
|
||||
"name": "hindsight",
|
||||
"version": "0.7.2",
|
||||
"description": "Official Hindsight integrations for Claude Code",
|
||||
"owner": {
|
||||
"name": "vectorize-io"
|
||||
|
||||
+1
-54
@@ -2,7 +2,7 @@
|
||||
# Copy this file to .env and fill in your values
|
||||
|
||||
# LLM Configuration (Required)
|
||||
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio, vertexai, minimax, deepseek, zai, atlas, volcano
|
||||
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio, vertexai, minimax, deepseek, zai, volcano
|
||||
HINDSIGHT_API_LLM_PROVIDER=openai
|
||||
HINDSIGHT_API_LLM_API_KEY=your-api-key-here
|
||||
HINDSIGHT_API_LLM_MODEL=gpt-4o-mini
|
||||
@@ -10,17 +10,6 @@ HINDSIGHT_API_LLM_BASE_URL=https://api.openai.com/v1
|
||||
# Reasoning effort for providers/models that support it. Examples: low, medium, high, xhigh.
|
||||
# HINDSIGHT_API_LLM_REASONING_EFFORT=low
|
||||
|
||||
# Sampling temperature for internal LLM calls. Set a number in [0.0, 2.0], or `none`
|
||||
# to omit the temperature parameter entirely (required for models that reject explicit
|
||||
# temperatures, e.g. Azure gpt-5.5). The global override below applies to every operation;
|
||||
# per-operation overrides (defaults: verification=0.0, retain=0.1, reflect=0.9,
|
||||
# consolidation=0.0) take precedence.
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE=none
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE_VERIFICATION=0.0
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE_RETAIN=0.1
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE_REFLECT=0.9
|
||||
# HINDSIGHT_API_LLM_TEMPERATURE_CONSOLIDATION=0.0
|
||||
|
||||
# Example: Anthropic Claude configuration
|
||||
# HINDSIGHT_API_LLM_PROVIDER=anthropic
|
||||
# HINDSIGHT_API_LLM_API_KEY=your-anthropic-api-key
|
||||
@@ -48,41 +37,16 @@ HINDSIGHT_API_LLM_BASE_URL=https://api.openai.com/v1
|
||||
# HINDSIGHT_API_LLM_API_KEY=your-zai-api-key
|
||||
# HINDSIGHT_API_LLM_MODEL=glm-4.5-flash # or glm-4.5-air for the paid tier
|
||||
|
||||
# Example: Atlas Cloud configuration (OpenAI-compatible, https://www.atlascloud.ai)
|
||||
# HINDSIGHT_API_LLM_PROVIDER=atlas
|
||||
# HINDSIGHT_API_LLM_API_KEY=your-atlascloud-api-key
|
||||
# HINDSIGHT_API_LLM_MODEL=deepseek-ai/deepseek-v4-pro # reasoning model; also Qwen / GLM / Kimi / MiniMax, etc.
|
||||
|
||||
# Example: LM Studio local configuration (Qwen 2.5 32B recommended)
|
||||
# HINDSIGHT_API_LLM_PROVIDER=lmstudio
|
||||
# HINDSIGHT_API_LLM_API_KEY=lmstudio
|
||||
# HINDSIGHT_API_LLM_BASE_URL=http://localhost:1234/v1
|
||||
# HINDSIGHT_API_LLM_MODEL=qwen2.5-32b-instruct
|
||||
|
||||
# Multi-LLM strategies: configure extra LLMs by index alongside the primary above,
|
||||
# then pick a routing strategy. Unset = single primary LLM (default). Members are
|
||||
# numbered from 1; indices must be contiguous. Each operation can override with a
|
||||
# RETAIN_/REFLECT_/CONSOLIDATION_ prefix (e.g. HINDSIGHT_API_RETAIN_LLM_1_PROVIDER).
|
||||
# HINDSIGHT_API_LLM_1_PROVIDER=groq
|
||||
# HINDSIGHT_API_LLM_1_API_KEY=your-groq-api-key
|
||||
# HINDSIGHT_API_LLM_1_MODEL=openai/gpt-oss-120b
|
||||
# HINDSIGHT_API_LLM_2_PROVIDER=anthropic
|
||||
# HINDSIGHT_API_LLM_2_API_KEY=your-anthropic-api-key
|
||||
# Strategy JSON: {"mode": "failover"} or {"mode": "round-robin"}.
|
||||
# Round-robin accepts optional positive-int "weights" (one per member, primary first).
|
||||
# HINDSIGHT_API_LLM_STRATEGY={"mode": "failover"}
|
||||
|
||||
# API Configuration (Optional)
|
||||
HINDSIGHT_API_HOST=0.0.0.0
|
||||
HINDSIGHT_API_PORT=8888
|
||||
HINDSIGHT_API_LOG_LEVEL=info
|
||||
# Optional retain chunking override for structured logs/transcripts.
|
||||
# Unset uses HINDSIGHT_API_RETAIN_CHUNK_SIZE as the structured-chunk limit.
|
||||
# HINDSIGHT_API_RETAIN_STRUCTURED_CHUNK_SIZE=
|
||||
|
||||
# Dry-run extraction preview endpoint (POST /memories/dry-run-extract). Enabled by default; it makes
|
||||
# a real LLM call but stores nothing. Set to false to remove the endpoint (returns 404).
|
||||
# HINDSIGHT_API_ENABLE_DRY_RUN_EXTRACT=true
|
||||
|
||||
# Base Path / Reverse Proxy Support (Optional)
|
||||
# Set these when deploying behind a reverse proxy with path-based routing
|
||||
@@ -95,7 +59,6 @@ HINDSIGHT_API_LOG_LEVEL=info
|
||||
# HINDSIGHT_API_READ_DATABASE_URL= # Optional read-replica URL. When set, recall queries (semantic, BM25, graph, temporal) flow through a separate pool against this URL, offloading the primary. Typically points to a read-only endpoint (CNPG's <cluster>-ro service or Aurora reader endpoint).
|
||||
# HINDSIGHT_API_MIGRATION_DATABASE_URL= # Direct PostgreSQL URL for migrations (bypasses PgBouncer). Falls back to DATABASE_URL.
|
||||
# HINDSIGHT_API_DATABASE_SCHEMA=public # PostgreSQL schema name (default: public)
|
||||
# HINDSIGHT_API_MIGRATION_CONCURRENCY=1 # Tenant schemas to migrate concurrently (PG only, each in its own process; per-schema work stays sequential). Each worker has ~1-2s startup cost + uses ~3 DB connections, so it only pays off with many schemas (tens+) or slow migrations; keep concurrency*3 <= spare max_connections. Default: 1 (sequential).
|
||||
|
||||
# Vector Extension (Optional - uses pgvector by default)
|
||||
# Options: "pgvector" (default), "vchord", "pgvectorscale" (DiskANN)
|
||||
@@ -116,18 +79,6 @@ HINDSIGHT_API_LOG_LEVEL=info
|
||||
# korean_lindera/lindera(korean), ngram(min,max), edge_ngram(min,max)
|
||||
# HINDSIGHT_API_TEXT_SEARCH_EXTENSION_PG_SEARCH_TOKENIZER=
|
||||
|
||||
# File Parser (Optional - uses markitdown by default)
|
||||
# HINDSIGHT_API_FILE_PARSER=markitdown
|
||||
# Enable image OCR for MarkItDown using an OpenAI-compatible OCR/vision endpoint.
|
||||
# These OCR settings are independent from HINDSIGHT_API_LLM_* because MarkItDown
|
||||
# uses the OpenAI SDK directly and requires Chat Completions image input support.
|
||||
# When OCR is enabled, API_KEY, BASE_URL, and MODEL are required.
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_ENABLED=false
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_API_KEY=
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_BASE_URL=
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_MODEL=
|
||||
# HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_PROMPT=
|
||||
|
||||
# Embeddings Configuration (Optional - uses local by default)
|
||||
# Provider: "local" (default), "onnx", "tei", "openai", "cohere", "google", "openrouter", "zeroentropy", "litellm", or "litellm-sdk"
|
||||
# HINDSIGHT_API_EMBEDDINGS_PROVIDER=local
|
||||
@@ -191,10 +142,6 @@ HINDSIGHT_API_LOG_LEVEL=info
|
||||
# Custom service name and environment (optional, defaults: hindsight-api, development)
|
||||
# HINDSIGHT_API_OTEL_SERVICE_NAME=hindsight-production
|
||||
# HINDSIGHT_API_OTEL_DEPLOYMENT_ENVIRONMENT=production
|
||||
#
|
||||
# Expose async-operation queue + consolidation-backlog gauges on /metrics.
|
||||
# Runs periodic per-schema COUNT queries on a background task (disabled by default).
|
||||
# HINDSIGHT_API_METRICS_BACKLOG_ENABLED=true
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Control Plane (Optional)
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
version: 2
|
||||
updates:
|
||||
- package-ecosystem: "github-actions"
|
||||
directory: "/"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
@@ -75,7 +75,7 @@ jobs:
|
||||
if: steps.type.outputs.type == 'plugin'
|
||||
run: |
|
||||
echo "Plugin integration ${{ steps.info.outputs.integration }} v${{ steps.info.outputs.version }} — no package to publish."
|
||||
echo "Users install via: claude plugin marketplace add vectorize-io/hindsight"
|
||||
echo "Users install via: claude plugin marketplace add vectorize-io/hindsight --sparse hindsight-integrations"
|
||||
|
||||
# ── TypeScript integrations (ai-sdk, chat, openclaw) ────────────────────
|
||||
|
||||
|
||||
@@ -266,7 +266,7 @@ jobs:
|
||||
strategy:
|
||||
matrix:
|
||||
include:
|
||||
- os: ubuntu-22.04
|
||||
- os: ubuntu-latest
|
||||
target: x86_64-unknown-linux-gnu
|
||||
artifact_name: hindsight
|
||||
asset_name: hindsight-linux-amd64
|
||||
@@ -278,7 +278,7 @@ jobs:
|
||||
target: aarch64-apple-darwin
|
||||
artifact_name: hindsight
|
||||
asset_name: hindsight-darwin-arm64
|
||||
- os: ubuntu-22.04-arm
|
||||
- os: ubuntu-24.04-arm
|
||||
target: aarch64-unknown-linux-gnu
|
||||
artifact_name: hindsight
|
||||
asset_name: hindsight-linux-arm64
|
||||
|
||||
@@ -33,35 +33,26 @@ jobs:
|
||||
integrations-openclaw: ${{ steps.filter.outputs.integrations-openclaw }}
|
||||
integrations-ai-sdk: ${{ steps.filter.outputs.integrations-ai-sdk }}
|
||||
integrations-agent-framework: ${{ steps.filter.outputs.integrations-agent-framework }}
|
||||
integrations-composio: ${{ steps.filter.outputs.integrations-composio }}
|
||||
integrations-chat: ${{ steps.filter.outputs.integrations-chat }}
|
||||
integrations-claude-code: ${{ steps.filter.outputs.integrations-claude-code }}
|
||||
integrations-cline: ${{ steps.filter.outputs.integrations-cline }}
|
||||
integrations-codex: ${{ steps.filter.outputs.integrations-codex }}
|
||||
integrations-github-copilot: ${{ steps.filter.outputs.integrations-github-copilot }}
|
||||
integrations-continue: ${{ steps.filter.outputs.integrations-continue }}
|
||||
integrations-cursor-cli: ${{ steps.filter.outputs.integrations-cursor-cli }}
|
||||
integrations-crewai: ${{ steps.filter.outputs.integrations-crewai }}
|
||||
integrations-litellm: ${{ steps.filter.outputs.integrations-litellm }}
|
||||
integrations-pydantic-ai: ${{ steps.filter.outputs.integrations-pydantic-ai }}
|
||||
integrations-ag2: ${{ steps.filter.outputs.integrations-ag2 }}
|
||||
integrations-autogen: ${{ steps.filter.outputs.integrations-autogen }}
|
||||
integrations-aider: ${{ steps.filter.outputs.integrations-aider }}
|
||||
integrations-langgraph: ${{ steps.filter.outputs.integrations-langgraph }}
|
||||
integrations-llamaindex: ${{ steps.filter.outputs.integrations-llamaindex }}
|
||||
integrations-paperclip: ${{ steps.filter.outputs.integrations-paperclip }}
|
||||
integrations-opencode: ${{ steps.filter.outputs.integrations-opencode }}
|
||||
integrations-eve: ${{ steps.filter.outputs.integrations-eve }}
|
||||
integrations-cursor: ${{ steps.filter.outputs.integrations-cursor }}
|
||||
integrations-zed: ${{ steps.filter.outputs.integrations-zed }}
|
||||
integrations-n8n: ${{ steps.filter.outputs.integrations-n8n }}
|
||||
integrations-zapier: ${{ steps.filter.outputs.integrations-zapier }}
|
||||
integrations-cloudflare-oauth-proxy: ${{ steps.filter.outputs.integrations-cloudflare-oauth-proxy }}
|
||||
integrations-superagent: ${{ steps.filter.outputs.integrations-superagent }}
|
||||
integrations-lockfiles: ${{ steps.filter.outputs.integrations-lockfiles }}
|
||||
integrations-openai-agents: ${{ steps.filter.outputs.integrations-openai-agents }}
|
||||
integrations-openhands: ${{ steps.filter.outputs.integrations-openhands }}
|
||||
integrations-devin-desktop: ${{ steps.filter.outputs.integrations-devin-desktop }}
|
||||
integrations-pipecat: ${{ steps.filter.outputs.integrations-pipecat }}
|
||||
integrations-agentcore: ${{ steps.filter.outputs.integrations-agentcore }}
|
||||
integrations-smolagents: ${{ steps.filter.outputs.integrations-smolagents }}
|
||||
@@ -136,8 +127,6 @@ jobs:
|
||||
- 'hindsight-integrations/ai-sdk/**'
|
||||
integrations-agent-framework:
|
||||
- 'hindsight-integrations/agent-framework/**'
|
||||
integrations-composio:
|
||||
- 'hindsight-integrations/composio/**'
|
||||
integrations-chat:
|
||||
- 'hindsight-integrations/chat/**'
|
||||
integrations-claude-code:
|
||||
@@ -146,10 +135,6 @@ jobs:
|
||||
- 'hindsight-integrations/cline/**'
|
||||
integrations-codex:
|
||||
- 'hindsight-integrations/codex/**'
|
||||
integrations-github-copilot:
|
||||
- 'hindsight-integrations/github-copilot/**'
|
||||
integrations-continue:
|
||||
- 'hindsight-integrations/continue/**'
|
||||
integrations-cursor-cli:
|
||||
- 'hindsight-integrations/cursor-cli/**'
|
||||
integrations-crewai:
|
||||
@@ -162,8 +147,6 @@ jobs:
|
||||
- 'hindsight-integrations/ag2/**'
|
||||
integrations-autogen:
|
||||
- 'hindsight-integrations/autogen/**'
|
||||
integrations-aider:
|
||||
- 'hindsight-integrations/aider/**'
|
||||
integrations-langgraph:
|
||||
- 'hindsight-integrations/langgraph/**'
|
||||
integrations-llamaindex:
|
||||
@@ -174,16 +157,10 @@ jobs:
|
||||
- 'hindsight-integrations/paperclip/**'
|
||||
integrations-opencode:
|
||||
- 'hindsight-integrations/opencode/**'
|
||||
integrations-eve:
|
||||
- 'hindsight-integrations/eve/**'
|
||||
integrations-cursor:
|
||||
- 'hindsight-integrations/cursor/**'
|
||||
integrations-zed:
|
||||
- 'hindsight-integrations/zed/**'
|
||||
integrations-n8n:
|
||||
- 'hindsight-integrations/n8n/**'
|
||||
integrations-zapier:
|
||||
- 'hindsight-integrations/zapier/**'
|
||||
integrations-cloudflare-oauth-proxy:
|
||||
- 'hindsight-integrations/cloudflare-oauth-proxy/**'
|
||||
integrations-superagent:
|
||||
@@ -194,10 +171,6 @@ jobs:
|
||||
- 'scripts/check-integration-lockfiles.sh'
|
||||
integrations-openai-agents:
|
||||
- 'hindsight-integrations/openai-agents/**'
|
||||
integrations-openhands:
|
||||
- 'hindsight-integrations/openhands/**'
|
||||
integrations-devin-desktop:
|
||||
- 'hindsight-integrations/devin-desktop/**'
|
||||
integrations-pipecat:
|
||||
- 'hindsight-integrations/pipecat/**'
|
||||
integrations-agentcore:
|
||||
@@ -506,37 +479,6 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/cursor
|
||||
run: python -m pytest tests/ -v
|
||||
|
||||
test-zed-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
github.event_name != 'pull_request_review' &&
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-zed == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: '3.11'
|
||||
|
||||
- name: Install package and pytest
|
||||
working-directory: ./hindsight-integrations/zed
|
||||
# Installs the package (incl. the zstandard runtime dep) so the threads.db
|
||||
# reader tests can decompress Zed's zstd blobs.
|
||||
run: pip install -e . pytest
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/zed
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: python -m pytest tests/ -v -m "not requires_real_llm"
|
||||
|
||||
test-omo-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -600,45 +542,6 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/cline
|
||||
run: uv run pytest tests -v
|
||||
|
||||
test-github-copilot-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-github-copilot == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build github-copilot integration
|
||||
working-directory: ./hindsight-integrations/github-copilot
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/github-copilot
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/github-copilot
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-codex-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -796,37 +699,6 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/opencode
|
||||
run: npm run build
|
||||
|
||||
test-eve-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-eve == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '24'
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/eve
|
||||
run: npm ci
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/eve
|
||||
run: npm test
|
||||
|
||||
- name: Build
|
||||
working-directory: ./hindsight-integrations/eve
|
||||
run: npm run build
|
||||
|
||||
test-n8n-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -858,37 +730,6 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/n8n
|
||||
run: npm run build
|
||||
|
||||
test-zapier-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-zapier == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '22'
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/zapier
|
||||
run: npm install --no-fund --no-audit
|
||||
|
||||
- name: Validate app definition
|
||||
working-directory: ./hindsight-integrations/zapier
|
||||
run: npm run validate
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/zapier
|
||||
run: npm test
|
||||
|
||||
test-hindsight-agent-sdk:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -3140,45 +2981,6 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/ag2
|
||||
run: uv run pytest tests -v
|
||||
|
||||
test-aider-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-aider == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build aider integration
|
||||
working-directory: ./hindsight-integrations/aider
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/aider
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/aider
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-autogen-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -3218,88 +3020,6 @@ jobs:
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-composio-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-composio == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build composio integration
|
||||
working-directory: ./hindsight-integrations/composio
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/composio
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Lint
|
||||
working-directory: ./hindsight-integrations/composio
|
||||
run: uv run ruff check .
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/composio
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-continue-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-continue == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build continue integration
|
||||
working-directory: ./hindsight-integrations/continue
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/continue
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/continue
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-smolagents-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -3827,84 +3547,6 @@ jobs:
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-openhands-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-openhands == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build openhands integration
|
||||
working-directory: ./hindsight-integrations/openhands
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/openhands
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/openhands
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-devin-desktop-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.integrations-devin-desktop == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build devin-desktop integration
|
||||
working-directory: ./hindsight-integrations/devin-desktop
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/devin-desktop
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/devin-desktop
|
||||
# PR CI runs only the deterministic bucket; the real-LLM E2E bucket
|
||||
# (requires_real_llm) needs a live Hindsight server and runs separately.
|
||||
run: uv run pytest tests -v -m "not requires_real_llm"
|
||||
|
||||
test-claude-agent-sdk-integration:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
@@ -4669,60 +4311,6 @@ jobs:
|
||||
fi
|
||||
done
|
||||
|
||||
# Dead-code detection beyond what ruff's F401/F841 catch (those are already
|
||||
# BLOCKING via the ruff config + the verify-generated-files job).
|
||||
#
|
||||
# - knip (control plane): BLOCKING on unused files / dependencies / unlisted
|
||||
# dependencies. These are unambiguous — an orphaned file or a dead
|
||||
# package.json entry — so they fail the build.
|
||||
# - vulture (Python) + knip unused *exports*: ADVISORY only. vulture's
|
||||
# function/argument heuristics false-positive on FastAPI/SQLAlchemy/Pydantic
|
||||
# patterns, and the control plane intentionally keeps an unused shadcn/ui
|
||||
# component surface, so these are surfaced in the step summary, not gated.
|
||||
check-unused-code:
|
||||
needs: [detect-changes]
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' ||
|
||||
needs.detect-changes.outputs.core == 'true' ||
|
||||
needs.detect-changes.outputs.control-plane == 'true' ||
|
||||
needs.detect-changes.outputs.ci == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || '' }}
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
with:
|
||||
node-version: '20'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: package-lock.json
|
||||
|
||||
- name: Install Control Plane dependencies
|
||||
run: npm install --workspace=hindsight-control-plane
|
||||
|
||||
- name: knip — unused files / dependencies (blocking)
|
||||
working-directory: hindsight-control-plane
|
||||
run: npx --yes knip@5 --no-progress --include files,dependencies,unlisted
|
||||
|
||||
- name: Advisory scan — vulture + knip exports
|
||||
continue-on-error: true
|
||||
run: |
|
||||
{
|
||||
echo '## Dead-code scan (advisory)'
|
||||
echo ''
|
||||
echo '```'
|
||||
./scripts/hooks/check-unused.sh 2>&1 | sed 's/\x1b\[[0-9;]*m//g'
|
||||
echo '```'
|
||||
} | tee -a "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
verify-generated-files:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
@@ -4906,13 +4494,11 @@ jobs:
|
||||
- test-claude-code-integration
|
||||
- test-cursor-integration
|
||||
- test-cline-integration
|
||||
- test-github-copilot-integration
|
||||
- test-codex-integration
|
||||
- test-cursor-cli-integration
|
||||
- build-ai-sdk-integration
|
||||
- test-ai-sdk-integration-deno
|
||||
- test-opencode-integration
|
||||
- test-eve-integration
|
||||
- test-omo-integration
|
||||
- test-cloudflare-oauth-proxy-integration
|
||||
- build-chat-integration
|
||||
@@ -4941,9 +4527,7 @@ jobs:
|
||||
- test-openclaw-integration
|
||||
- test-integration
|
||||
- test-ag2-integration
|
||||
- test-aider-integration
|
||||
- test-autogen-integration
|
||||
- test-continue-integration
|
||||
- test-smolagents-integration
|
||||
- test-dify-integration
|
||||
- test-flowise-integration
|
||||
@@ -4956,8 +4540,6 @@ jobs:
|
||||
- test-pydantic-ai-integration
|
||||
- test-llamaindex-integration
|
||||
- test-openai-agents-integration
|
||||
- test-openhands-integration
|
||||
- test-devin-desktop-integration
|
||||
- test-agentcore-integration
|
||||
- test-haystack-integration
|
||||
- test-pip-slim
|
||||
|
||||
@@ -6,7 +6,6 @@ dist/
|
||||
wheels/
|
||||
*.egg-info
|
||||
.mcp.json
|
||||
.playwright-mcp/
|
||||
.osgrep
|
||||
# Virtual environments
|
||||
.venv
|
||||
|
||||
@@ -216,18 +216,6 @@ migration file dispatches through `run_for_dialect`, which calls either
|
||||
./scripts/hooks/lint.sh
|
||||
```
|
||||
|
||||
Dead-code detection runs in CI (the `check-unused-code` job) at two levels:
|
||||
- **Blocking:** unused imports (ruff `F401`) and variables (`F841`) — `lint.sh` auto-removes
|
||||
them and `verify-generated-files` fails on any leftover diff; and **knip** for orphaned
|
||||
control-plane files / unused (or unlisted) `package.json` dependencies.
|
||||
- **Advisory:** whole unused Python functions (vulture) and unused control-plane *exports*
|
||||
(the shadcn/ui surface is kept on purpose) — surfaced, not gated.
|
||||
|
||||
Run both locally with:
|
||||
```bash
|
||||
./scripts/hooks/check-unused.sh
|
||||
```
|
||||
|
||||
**After completing any implementation work, run `/code-review`** to verify your changes against project standards (missing tests, dead code, type safety, etc.). Fix any "must fix" issues before considering the task done.
|
||||
|
||||
**MANDATORY: Run `/code-review` before pushing code or creating a pull request.** Do not push or create a PR until all "must fix" issues are resolved.
|
||||
@@ -327,10 +315,7 @@ Fields must be categorized as either **hierarchical** (can be overridden per-ten
|
||||
```
|
||||
|
||||
2. **main.py** (`hindsight-api-slim/hindsight_api/main.py`):
|
||||
- No change is needed for ordinary environment-backed config fields. The CLI starts from `_get_raw_config()`,
|
||||
so new `HindsightConfig` fields are carried through automatically.
|
||||
- If the new field should be overridable by a CLI flag, add the argparse option in `_parse_cli_args()` and include
|
||||
that field in the `dataclasses.replace(config, ...)` call near the "CLI override" comment.
|
||||
- Add field to the manual `HindsightConfig()` constructor call (search for "CLI override")
|
||||
|
||||
3. **Use hierarchical config in MemoryEngine**:
|
||||
```python
|
||||
@@ -376,7 +361,7 @@ Fields must be categorized as either **hierarchical** (can be overridden per-ten
|
||||
|
||||
```bash
|
||||
cp .env.example .env
|
||||
# Edit .env with the LLM provider/model and credentials for your setup
|
||||
# Edit .env with LLM API key
|
||||
|
||||
# Python deps
|
||||
uv sync --directory hindsight-api-slim/
|
||||
@@ -385,10 +370,10 @@ uv sync --directory hindsight-api-slim/
|
||||
npm install
|
||||
```
|
||||
|
||||
Common LLM settings:
|
||||
Required env vars:
|
||||
- `HINDSIGHT_API_LLM_PROVIDER`: openai, anthropic, gemini, groq, minimax, ollama, lmstudio
|
||||
- `HINDSIGHT_API_LLM_API_KEY`: API key for providers that require one
|
||||
- `HINDSIGHT_API_LLM_MODEL`: Model name (defaults are provider-specific)
|
||||
- `HINDSIGHT_API_LLM_API_KEY`: Your API key
|
||||
- `HINDSIGHT_API_LLM_MODEL`: Model name (e.g., gpt-4o-mini, claude-sonnet-4-20250514)
|
||||
|
||||
Optional (uses local models by default):
|
||||
- `HINDSIGHT_API_EMBEDDINGS_PROVIDER`: local (default) or tei
|
||||
|
||||
@@ -70,7 +70,7 @@ docker run -it --pull always --name hindsight --restart unless-stopped -p 8888:8
|
||||
>API: http://localhost:8888
|
||||
>UI: http://localhost:9999
|
||||
|
||||
You can modify the LLM provider by setting `HINDSIGHT_API_LLM_PROVIDER`. Valid options are `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio`, `minimax`, and `atlas` ([Atlas Cloud](https://www.atlascloud.ai/?utm_source=github&utm_medium=link&utm_campaign=hindsight)). The documentation provides more details on [supported models](https://hindsight.vectorize.io/developer/models).
|
||||
You can modify the LLM provider by setting `HINDSIGHT_API_LLM_PROVIDER`. Valid options are `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio`, and `minimax`. The documentation provides more details on [supported models](https://hindsight.vectorize.io/developer/models).
|
||||
|
||||
|
||||
|
||||
@@ -250,7 +250,7 @@ Recall performs 4 retrieval strategies in parallel:
|
||||
- Graph: Entity/temporal/causal links
|
||||
- Temporal: Time range filtering
|
||||
|
||||

|
||||

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

|
||||

|
||||
|
||||
---
|
||||
|
||||
@@ -310,7 +310,7 @@ client.reflect(bank_id="my-bank", query="What should I know about Alice?")
|
||||
| **macOS** (Intel / x86_64) | ✅ | ⚠️ | ✅ |
|
||||
| **Windows** (x86_64) | ✅ | ✅ | ✅ |
|
||||
|
||||
⚠️ Intel Macs: use `hindsight-all-slim` — see the [installation guide](https://hindsight.vectorize.io/developer/installation#supported-platforms) for details.
|
||||
⚠️ Intel Macs: use `hindsight-all-slim` — see the [installation guide](https://docs.hindsight.vectorize.io/docs/developer/installation#supported-platforms) for details.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -2,8 +2,8 @@ apiVersion: v2
|
||||
name: hindsight
|
||||
description: Hindsight helm chart
|
||||
type: application
|
||||
version: 0.8.4
|
||||
appVersion: "0.8.4"
|
||||
version: 0.8.1
|
||||
appVersion: "0.8.1"
|
||||
keywords:
|
||||
- ai
|
||||
- memory
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@vectorize-io/hindsight-all",
|
||||
"version": "0.8.4",
|
||||
"version": "0.8.1",
|
||||
"description": "Node.js programmatic lifecycle manager for Hindsight — embeds a local hindsight daemon in a Node application. Pair with @vectorize-io/hindsight-client for memory operations.",
|
||||
"main": "dist/index.js",
|
||||
"types": "dist/index.d.ts",
|
||||
|
||||
@@ -4,12 +4,12 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "hindsight-all-slim"
|
||||
version = "0.8.4"
|
||||
version = "0.8.1"
|
||||
description = "Hindsight: Agent Memory That Works Like Human Memory - Slim All-in-One Bundle"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
"hindsight-api-slim==0.8.4",
|
||||
"hindsight-api-slim==0.8.1",
|
||||
"hindsight-client>=0.0.7",
|
||||
"hindsight-embed>=0.1.0",
|
||||
]
|
||||
|
||||
@@ -4,12 +4,12 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "hindsight-all"
|
||||
version = "0.8.4"
|
||||
version = "0.8.1"
|
||||
description = "Hindsight: Agent Memory That Works Like Human Memory - All-in-One Bundle"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
"hindsight-api-slim[all]==0.8.4",
|
||||
"hindsight-api-slim[all]==0.8.1",
|
||||
"hindsight-client>=0.0.7",
|
||||
"hindsight-embed>=0.1.0",
|
||||
]
|
||||
@@ -21,7 +21,7 @@ hindsight-embed = { workspace = true }
|
||||
|
||||
[project.optional-dependencies]
|
||||
local-llm = [
|
||||
"hindsight-api-slim[local-llm]==0.8.4",
|
||||
"hindsight-api-slim[local-llm]==0.8.1",
|
||||
]
|
||||
test = [
|
||||
"pytest>=7.0.0",
|
||||
|
||||
@@ -121,7 +121,7 @@ This runs a stdio-based MCP server that can be used directly with MCP-compatible
|
||||
- **Entity Graph** — Automatic entity extraction and relationship tracking
|
||||
- **Temporal Reasoning** — Native support for time-based queries
|
||||
- **Disposition Traits** — Configurable skepticism, literalism, and empathy influence opinion formation
|
||||
- **Three Memory Types** — World facts, experience facts (the bank's own actions), and observations
|
||||
- **Three Memory Types** — World facts, bank actions, and formed opinions with confidence scores
|
||||
|
||||
## Documentation
|
||||
|
||||
|
||||
@@ -53,4 +53,4 @@ __all__ = [
|
||||
"RemoteTEICrossEncoder",
|
||||
"LLMConfig",
|
||||
]
|
||||
__version__ = "0.8.4"
|
||||
__version__ = "0.8.1"
|
||||
|
||||
@@ -56,7 +56,6 @@ BACKUP_TABLES = [
|
||||
"observation_history",
|
||||
"mental_models",
|
||||
"mental_model_history",
|
||||
"knowledge_pages",
|
||||
"directives",
|
||||
"async_operations",
|
||||
"webhooks",
|
||||
@@ -257,10 +256,14 @@ async def _run_migration(
|
||||
schema: str | None = None,
|
||||
base_schema: str = DEFAULT_DATABASE_SCHEMA,
|
||||
embedding_dimension: int | None = None,
|
||||
ensure_extensions: bool = True,
|
||||
) -> list[str]:
|
||||
"""Resolve database URL and run migrations for one schema or all discovered schemas."""
|
||||
from ..migrations import run_migrations_for_schemas
|
||||
from ..migrations import (
|
||||
ensure_embedding_dimension,
|
||||
ensure_text_search_extension,
|
||||
ensure_vector_extension,
|
||||
run_migrations,
|
||||
)
|
||||
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
if is_pg0:
|
||||
@@ -281,21 +284,32 @@ async def _run_migration(
|
||||
# Preserve order while removing duplicates.
|
||||
schemas = list(dict.fromkeys(schemas))
|
||||
|
||||
# Migrate up to `migration_concurrency` schemas at once (each in its own
|
||||
# process); within a schema the work stays sequential. Run off the event
|
||||
# loop so the process pool's blocking joins don't stall it.
|
||||
await asyncio.to_thread(
|
||||
run_migrations_for_schemas,
|
||||
resolved_url,
|
||||
schemas,
|
||||
concurrency=config.migration_concurrency,
|
||||
migration_database_url=config.migration_database_url,
|
||||
embedding_dimension=embedding_dimension,
|
||||
vector_extension=config.vector_extension,
|
||||
text_search_extension=config.text_search_extension,
|
||||
pg_search_tokenizer=config.text_search_extension_pg_search_tokenizer,
|
||||
ensure_extensions=ensure_extensions,
|
||||
)
|
||||
for schema in schemas:
|
||||
run_migrations(resolved_url, schema=schema, migration_database_url=config.migration_database_url)
|
||||
|
||||
if embedding_dimension is not None:
|
||||
for schema in schemas:
|
||||
ensure_embedding_dimension(
|
||||
resolved_url,
|
||||
embedding_dimension,
|
||||
schema=schema,
|
||||
vector_extension=config.vector_extension,
|
||||
)
|
||||
|
||||
for schema in schemas:
|
||||
ensure_vector_extension(
|
||||
resolved_url,
|
||||
vector_extension=config.vector_extension,
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
for schema in schemas:
|
||||
ensure_text_search_extension(
|
||||
resolved_url,
|
||||
text_search_extension=config.text_search_extension,
|
||||
pg_search_tokenizer=config.text_search_extension_pg_search_tokenizer,
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
return schemas
|
||||
|
||||
@@ -313,18 +327,6 @@ def run_db_migration(
|
||||
"--embedding-dimension",
|
||||
help="Expected embedding dimension to enforce after migrations. Omit to skip dimension sync.",
|
||||
),
|
||||
skip_extension_reconcile: bool = typer.Option(
|
||||
False,
|
||||
"--skip-extension-reconcile",
|
||||
help=(
|
||||
"Skip the post-migration vector / text-search index reconcile. This step only does "
|
||||
"work when the configured backend (HINDSIGHT_API_VECTOR_EXTENSION / "
|
||||
"HINDSIGHT_API_TEXT_SEARCH_EXTENSION) differs from a schema's existing indexes — a "
|
||||
"rare, operator-driven change. Skipping it makes a no-change re-migration over many "
|
||||
"tenant schemas much faster. Only use when you have NOT changed the backend; a "
|
||||
"backend change still needs a normal run to reshape the indexes."
|
||||
),
|
||||
),
|
||||
):
|
||||
"""Run database migrations to the latest version."""
|
||||
config = HindsightConfig.from_env()
|
||||
@@ -338,8 +340,6 @@ def run_db_migration(
|
||||
typer.echo(f"Running database migrations for schema: {schema}...")
|
||||
else:
|
||||
typer.echo("Running database migrations for base schema and all discovered tenant schemas...")
|
||||
if skip_extension_reconcile:
|
||||
typer.echo("Skipping post-migration extension reconcile (--skip-extension-reconcile).")
|
||||
|
||||
schemas = asyncio.run(
|
||||
_run_migration(
|
||||
@@ -347,7 +347,6 @@ def run_db_migration(
|
||||
schema=schema,
|
||||
base_schema=config.database_schema,
|
||||
embedding_dimension=embedding_dimension,
|
||||
ensure_extensions=not skip_extension_reconcile,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
-105
@@ -1,105 +0,0 @@
|
||||
"""Add a composite index on memory_links(bank_id, link_type) (PostgreSQL).
|
||||
|
||||
``bank_id`` was added to ``memory_links`` in ``c5d6e7f8a9b0`` precisely so that
|
||||
bank-scoped reads (e.g. the stats endpoint) could filter on the link table
|
||||
directly instead of joining ``memory_units`` — that JOIN took 18+ seconds on
|
||||
banks with millions of links. The column landed without an index, so every
|
||||
``bank_id = $1`` predicate still falls back to a sequential scan over the whole
|
||||
table.
|
||||
|
||||
This adds the missing btree. It is composite on ``(bank_id, link_type)`` rather
|
||||
than ``bank_id`` alone because the hot query is the stats endpoint's
|
||||
``SELECT link_type, COUNT(*) ... WHERE bank_id = $1 GROUP BY link_type``: a
|
||||
``(bank_id, link_type)`` index serves that filter, grouping and count as an
|
||||
index-only scan, never touching the heap, whereas a ``bank_id``-only index would
|
||||
still have to read every matching row to recover ``link_type``. ``link_type`` is
|
||||
low-cardinality (only ``temporal``/``semantic``/``caused_by`` are written —
|
||||
entity edges were dropped in ``e9b2c7d1f3a4``), so the trailing column adds
|
||||
little to the index size while removing the heap fetch.
|
||||
|
||||
The Oracle baseline (``o1a2b3c4d5e6``) already creates ``idx_ml_bank_id`` on
|
||||
``memory_links(bank_id)``; that single-column index already covers Oracle's
|
||||
bank-scoped filter, so the Oracle slot here is intentionally absent and only the
|
||||
PostgreSQL dialect gets the composite index.
|
||||
|
||||
``memory_links`` can hold tens of millions of rows, so the index is built
|
||||
CONCURRENTLY to avoid taking a write lock on the table. CONCURRENTLY cannot run
|
||||
inside a transaction block, so the statement runs in an ``autocommit_block()``;
|
||||
``IF NOT EXISTS`` keeps it idempotent across retries and re-migrated tenant
|
||||
schemas. A CONCURRENTLY build interrupted partway (lock conflict, disk
|
||||
pressure, signal) leaves the index behind as *invalid*; ``IF NOT EXISTS`` would
|
||||
then skip over it forever, so the upgrade first drops any invalid leftover of
|
||||
this name before (re)creating it.
|
||||
|
||||
Revision ID: 2071c7518f88
|
||||
Revises: a1d3f5b7c9e2
|
||||
Create Date: 2026-06-16
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
from sqlalchemy import text
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "2071c7518f88"
|
||||
down_revision: str | Sequence[str] | None = "a1d3f5b7c9e2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
_INDEX_NAME = "idx_memory_links_bank_id_link_type"
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Schema-qualifier for raw SQL on PG (multi-tenant search_path)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
# `or None` collapses an unset option and an explicit empty string into NULL
|
||||
# so the COALESCE below falls back to current_schema() in both cases.
|
||||
target_schema = context.config.get_main_option("target_schema") or None
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# CREATE INDEX CONCURRENTLY cannot run inside a transaction block; the
|
||||
# autocommit_block runs each statement outside Alembic's migration
|
||||
# transaction.
|
||||
with op.get_context().autocommit_block():
|
||||
# A CONCURRENTLY build that errored on a previous run leaves an INVALID
|
||||
# index of this name behind. `CREATE INDEX ... IF NOT EXISTS` would see
|
||||
# that relation and skip, so bank_id queries would keep seq-scanning.
|
||||
# Drop only the invalid leftover — never a healthy index — so the retry
|
||||
# actually rebuilds a usable one.
|
||||
leftover_invalid = bind.execute(
|
||||
text(
|
||||
"SELECT NOT i.indisvalid "
|
||||
"FROM pg_class c "
|
||||
"JOIN pg_index i ON c.oid = i.indexrelid "
|
||||
"JOIN pg_namespace n ON c.relnamespace = n.oid "
|
||||
"WHERE c.relname = :index_name "
|
||||
" AND n.nspname = COALESCE(:target_schema, current_schema())"
|
||||
),
|
||||
{"index_name": _INDEX_NAME, "target_schema": target_schema},
|
||||
).scalar()
|
||||
if leftover_invalid:
|
||||
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}{_INDEX_NAME}")
|
||||
|
||||
# IF NOT EXISTS keeps the create idempotent across retries and schemas.
|
||||
op.execute(f"CREATE INDEX CONCURRENTLY IF NOT EXISTS {_INDEX_NAME} ON {schema}memory_links(bank_id, link_type)")
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
with op.get_context().autocommit_block():
|
||||
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}{_INDEX_NAME}")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
-85
@@ -1,85 +0,0 @@
|
||||
"""Repair: widen the remaining live ``bank_id`` columns from VARCHAR(64) to TEXT on PostgreSQL.
|
||||
|
||||
Follow-up to ``c3e5a7b9d1f4`` (issue #2106), which widened the two *history*
|
||||
tables (``observation_history``, ``mental_model_history``) to ``TEXT`` after the
|
||||
narrow ``VARCHAR(64)`` declaration bricked startup. The same VARCHAR(64) / TEXT
|
||||
inconsistency still affects the live tables that store a user-supplied
|
||||
``bank_id``:
|
||||
|
||||
* ``directives`` -- created VARCHAR(64) in ``p1k2l3m4n5o6``
|
||||
* ``mental_models`` -- VARCHAR(64) (origin ``pinned_reflections`` in
|
||||
``n9i0j1k2l3m4``; recreated in ``h3c4d5e6f7g8``)
|
||||
|
||||
``mental_model_versions`` is intentionally *not* widened here: it is created in
|
||||
``j5e6f7g8h9i0`` but dropped (``DROP TABLE ... CASCADE``) in ``o0j1k2l3m4n5`` and
|
||||
never recreated on the upgrade path, so it does not exist at head. Issuing
|
||||
``ALTER TABLE mental_model_versions ...`` would raise ``UndefinedTable`` and --
|
||||
because migrations run inside the lifespan-startup transaction -- roll the whole
|
||||
migration back, bricking the API. (It is unrelated to the live
|
||||
``mental_model_history`` table widened by ``c3e5a7b9d1f4``.)
|
||||
|
||||
``banks.bank_id`` is ``TEXT`` (unbounded), so a deployment can create a bank
|
||||
whose id exceeds 64 chars -- the 78-char hierarchical org-unit shape reported in
|
||||
issue #2106 -- and the bank insert succeeds. The next write that propagates that
|
||||
id (``create_directive``, ``create_mental_model`` / consolidation, or
|
||||
mental-model versioning) then aborts with::
|
||||
|
||||
psycopg2.errors.StringDataRightTruncation: value too long for type
|
||||
character varying(64)
|
||||
|
||||
i.e. a 500 on core write endpoints, instead of the startup brick that
|
||||
``c3e5a7b9d1f4`` already repaired.
|
||||
|
||||
``ALTER COLUMN ... TYPE TEXT`` is a no-op on a column that is already ``TEXT``,
|
||||
so every upgrade path converges on ``TEXT``. These tables are per-tenant (they
|
||||
live in each tenant schema, not ``public``), so this runs for every migrated
|
||||
schema via the search-path-aware prefix -- the same mechanism as
|
||||
``c3e5a7b9d1f4``.
|
||||
|
||||
PostgreSQL only: these tables are created by PostgreSQL-only migrations
|
||||
(``run_for_dialect(pg=...)``); on Oracle they are absent or already
|
||||
``VARCHAR2(256)`` (consistent, never truncates), so the Oracle slot is
|
||||
intentionally absent -- mirroring ``c3e5a7b9d1f4``.
|
||||
|
||||
Revision ID: a1d3f5b7c9e2
|
||||
Revises: c3e5a7b9d1f4
|
||||
Create Date: 2026-06-13
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "a1d3f5b7c9e2"
|
||||
down_revision: str | Sequence[str] | None = "c3e5a7b9d1f4"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Schema-qualifier for raw SQL on PG (multi-tenant search_path)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
op.execute(f"ALTER TABLE {schema}directives ALTER COLUMN bank_id TYPE TEXT")
|
||||
op.execute(f"ALTER TABLE {schema}mental_models ALTER COLUMN bank_id TYPE TEXT")
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
# No-op: narrowing back to VARCHAR(64) could truncate real data and would
|
||||
# re-introduce the bug this migration repairs. The column types are owned by
|
||||
# the migrations that created the tables.
|
||||
pass
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
-52
@@ -1,52 +0,0 @@
|
||||
"""Add managed flag to knowledge_pages.
|
||||
|
||||
The knowledge base is managed by clients (CRUD over folders/pages). ``managed``
|
||||
lets a client tag a node as system-owned vs. hand-authored; it carries no
|
||||
server-side behaviour.
|
||||
|
||||
Revision ID: a5b6c7d8e9f0
|
||||
Revises: a9b8c7d6e5f4
|
||||
Create Date: 2026-06-26
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "a5b6c7d8e9f0"
|
||||
down_revision: str | Sequence[str] | None = "a9b8c7d6e5f4"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _pg_schema_prefix() -> str:
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
op.execute(f"ALTER TABLE {schema}knowledge_pages ADD COLUMN IF NOT EXISTS managed BOOLEAN NOT NULL DEFAULT false")
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
op.execute(f"ALTER TABLE {schema}knowledge_pages DROP COLUMN IF EXISTS managed")
|
||||
|
||||
|
||||
def _oracle_upgrade() -> None:
|
||||
op.execute("ALTER TABLE knowledge_pages ADD (managed NUMBER(1) DEFAULT 0 NOT NULL)")
|
||||
|
||||
|
||||
def _oracle_downgrade() -> None:
|
||||
op.execute("ALTER TABLE knowledge_pages DROP COLUMN managed")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade, oracle=_oracle_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade, oracle=_oracle_downgrade)
|
||||
-110
@@ -1,110 +0,0 @@
|
||||
"""Add knowledge_pages table (knowledge-base hierarchy).
|
||||
|
||||
The knowledge base organizes synthesized mental models into a navigable tree of
|
||||
**folders** and **pages**. A page references the mental model that holds its
|
||||
content (``mental_model_id``); a folder is a pure container (``mental_model_id``
|
||||
NULL). Hierarchy is a single self-referential ``parent_id`` so folders can nest
|
||||
arbitrarily. Content stays in ``mental_models`` — this table is metadata + tree
|
||||
structure only.
|
||||
|
||||
Revision ID: a9b8c7d6e5f4
|
||||
Revises: b57a7c9e0d13
|
||||
Create Date: 2026-06-25
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "a9b8c7d6e5f4"
|
||||
down_revision: str | Sequence[str] | None = "b57a7c9e0d13"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _pg_schema_prefix() -> str:
|
||||
"""Schema-qualifier for raw SQL on PG (multi-tenant search_path)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
# parent_id self-FK cascades so deleting a folder row removes its whole
|
||||
# subtree of rows in one shot. The mental_model FK is composite (matches the
|
||||
# mental_models (id, bank_id) PK) and cascades too, so deleting a page's
|
||||
# mental model removes the page row — folders skip the FK because a NULL
|
||||
# column in a composite FK is not enforced (MATCH SIMPLE).
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {schema}knowledge_pages (
|
||||
id VARCHAR(64) NOT NULL,
|
||||
bank_id TEXT NOT NULL,
|
||||
parent_id VARCHAR(64),
|
||||
kind VARCHAR(16) NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
mental_model_id VARCHAR(64),
|
||||
sort_order INTEGER NOT NULL DEFAULT 0,
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
|
||||
updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
|
||||
CONSTRAINT pk_knowledge_pages PRIMARY KEY (id),
|
||||
CONSTRAINT ck_knowledge_pages_kind CHECK (kind IN ('folder', 'page')),
|
||||
CONSTRAINT fk_kp_bank FOREIGN KEY (bank_id)
|
||||
REFERENCES {schema}banks(bank_id) ON DELETE CASCADE,
|
||||
CONSTRAINT fk_kp_parent FOREIGN KEY (parent_id)
|
||||
REFERENCES {schema}knowledge_pages(id) ON DELETE CASCADE,
|
||||
CONSTRAINT fk_kp_mm FOREIGN KEY (mental_model_id, bank_id)
|
||||
REFERENCES {schema}mental_models(id, bank_id) ON DELETE CASCADE
|
||||
)
|
||||
"""
|
||||
)
|
||||
op.execute(
|
||||
f"CREATE INDEX IF NOT EXISTS idx_kp_bank_parent ON {schema}knowledge_pages (bank_id, parent_id, sort_order)"
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_kp_bank_parent")
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}knowledge_pages")
|
||||
|
||||
|
||||
def _oracle_upgrade() -> None:
|
||||
op.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS knowledge_pages (
|
||||
id VARCHAR2(64) NOT NULL,
|
||||
bank_id VARCHAR2(256) NOT NULL,
|
||||
parent_id VARCHAR2(64),
|
||||
kind VARCHAR2(16) NOT NULL,
|
||||
name CLOB NOT NULL,
|
||||
mental_model_id VARCHAR2(64),
|
||||
sort_order NUMBER DEFAULT 0 NOT NULL,
|
||||
created_at TIMESTAMP WITH TIME ZONE DEFAULT SYSTIMESTAMP NOT NULL,
|
||||
updated_at TIMESTAMP WITH TIME ZONE DEFAULT SYSTIMESTAMP NOT NULL,
|
||||
CONSTRAINT pk_knowledge_pages PRIMARY KEY (id),
|
||||
CONSTRAINT ck_knowledge_pages_kind CHECK (kind IN ('folder', 'page')),
|
||||
CONSTRAINT fk_kp_bank FOREIGN KEY (bank_id)
|
||||
REFERENCES banks(bank_id) ON DELETE CASCADE,
|
||||
CONSTRAINT fk_kp_parent FOREIGN KEY (parent_id)
|
||||
REFERENCES knowledge_pages(id) ON DELETE CASCADE,
|
||||
CONSTRAINT fk_kp_mm FOREIGN KEY (mental_model_id, bank_id)
|
||||
REFERENCES mental_models(id, bank_id) ON DELETE CASCADE
|
||||
)
|
||||
"""
|
||||
)
|
||||
op.execute("CREATE INDEX idx_kp_bank_parent ON knowledge_pages (bank_id, parent_id, sort_order)")
|
||||
|
||||
|
||||
def _oracle_downgrade() -> None:
|
||||
op.execute("DROP TABLE knowledge_pages CASCADE CONSTRAINTS")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade, oracle=_oracle_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade, oracle=_oracle_downgrade)
|
||||
-61
@@ -1,61 +0,0 @@
|
||||
"""Add bank_stats_cache table for distributed get_bank_stats caching
|
||||
|
||||
Revision ID: b57a7c9e0d13
|
||||
Revises: c3f7a1b9d2e4
|
||||
Create Date: 2026-07-01
|
||||
|
||||
get_bank_stats aggregates over memory_links / unit_entities — a multi-second scan
|
||||
on banks with millions of rows. The result was cached per-process (in-memory), so
|
||||
every API worker recomputed it once per TTL and the first caller after expiry
|
||||
stalled. This table backs a shared, cross-process TTL cache: one worker's compute
|
||||
is written here and served to all the others.
|
||||
|
||||
PostgreSQL only. Oracle keeps the in-process cache (the runtime picks the backing
|
||||
store by dialect), so the Oracle upgrade slot is intentionally absent.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "b57a7c9e0d13"
|
||||
down_revision: str | Sequence[str] | None = "c3f7a1b9d2e4"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Schema-qualifier for raw SQL on PG (multi-tenant search_path)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
# One row per bank: payload is the full get_bank_stats result, computed_at
|
||||
# drives logical TTL expiry. Rows are overwritten in place (ON CONFLICT), so
|
||||
# the table never grows beyond the number of banks and needs no purge job.
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {schema}bank_stats_cache (
|
||||
bank_id TEXT PRIMARY KEY,
|
||||
payload JSONB NOT NULL,
|
||||
computed_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}bank_stats_cache")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade) # oracle slot intentionally absent → no-op
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
-71
@@ -1,71 +0,0 @@
|
||||
"""Unique page name per folder in knowledge_pages.
|
||||
|
||||
The folder curator can fire concurrently (folder-create trigger + the
|
||||
post-consolidation sweep), and an in-process lock can't serialize runs that
|
||||
execute in different threads/loops. A partial unique index on
|
||||
(bank_id, parent, lower(name)) for pages makes duplicate-named pages in the same
|
||||
folder impossible at the DB level — the second concurrent insert fails and the
|
||||
curator treats it as "already exists".
|
||||
|
||||
PostgreSQL only: the Oracle ``name`` column is a CLOB and cannot back a
|
||||
functional unique index; Oracle relies on the in-process serialization instead.
|
||||
|
||||
Revision ID: c3d4e5f6a7b8
|
||||
Revises: a5b6c7d8e9f0
|
||||
Create Date: 2026-06-26
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "c3d4e5f6a7b8"
|
||||
down_revision: str | Sequence[str] | None = "a5b6c7d8e9f0"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _pg_schema_prefix() -> str:
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
# First drop any pre-existing duplicate pages (created by the racy curator
|
||||
# before this guard existed), keeping the earliest row of each duplicate set,
|
||||
# so the unique index can be built. Their backing mental models are left in
|
||||
# place (harmless orphans).
|
||||
op.execute(
|
||||
f"""
|
||||
DELETE FROM {schema}knowledge_pages a
|
||||
USING {schema}knowledge_pages b
|
||||
WHERE a.kind = 'page' AND b.kind = 'page'
|
||||
AND a.bank_id = b.bank_id
|
||||
AND COALESCE(a.parent_id, '') = COALESCE(b.parent_id, '')
|
||||
AND lower(a.name) = lower(b.name)
|
||||
AND a.ctid > b.ctid
|
||||
"""
|
||||
)
|
||||
# COALESCE(parent_id, '') so root-level pages (NULL parent) are also unique by
|
||||
# name — NULLs would otherwise compare distinct and allow duplicates.
|
||||
op.execute(
|
||||
"CREATE UNIQUE INDEX IF NOT EXISTS uq_kp_folder_pagename "
|
||||
f"ON {schema}knowledge_pages (bank_id, COALESCE(parent_id, ''), lower(name)) "
|
||||
"WHERE kind = 'page'"
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}uq_kp_folder_pagename")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade) # oracle slot intentionally absent (CLOB name)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
-144
@@ -1,144 +0,0 @@
|
||||
"""Backfill search_vector for native-backend observations.
|
||||
|
||||
Observations created or updated by the consolidator landed with a NULL
|
||||
``search_vector`` under the ``native`` text-search backend: the
|
||||
single-row INSERT/UPDATE paths in ``consolidator.py`` never populated the
|
||||
tsvector (only the batch raw-fact path in ``ops_postgresql.insert_facts_batch``
|
||||
did). Those observations were therefore invisible to the BM25 retrieval arm
|
||||
until they were re-written by a later consolidation pass. The writer is fixed
|
||||
in the same change set (all four consolidator sites now call
|
||||
``to_tsvector($lang, COALESCE(text, ''))``); this migration repairs the
|
||||
historical residue so existing observations become BM25-searchable without a
|
||||
re-ingest.
|
||||
|
||||
Scope mirrors the writer fix exactly:
|
||||
* Only the ``native`` backend is touched. The gate is the column *type*:
|
||||
under ``native`` ``search_vector`` is a regular (non-generated) tsvector
|
||||
column; under ``vchord`` it is a ``bm25vector`` and under
|
||||
``pg_textsearch`` / ``pgroonga`` / ``pg_search`` it is a dummy ``text``
|
||||
column. ``_is_regular_tsvector`` is true only for ``native``, so every
|
||||
other backend is a no-op.
|
||||
* The tsvector is built from the observation's own ``text`` only — matching
|
||||
the consolidator INSERT/UPDATE paths (entity / source / temporal signals
|
||||
are intentionally excluded; the other retrieval arms cover those).
|
||||
* Only ``fact_type = 'observation'`` rows with a NULL ``search_vector`` are
|
||||
rewritten. Raw facts already carry a populated tsvector, and the
|
||||
``IS NULL`` predicate makes the migration idempotent and re-runnable.
|
||||
|
||||
The configured ``HINDSIGHT_API_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE`` is used
|
||||
so backfilled rows are lexically identical to newly-created observations. The
|
||||
value is validated as a PG identifier (mirroring
|
||||
``HindsightConfig.validate``) before being embedded as a SQL literal.
|
||||
|
||||
This is a single UPDATE per schema: it locks the targeted observation rows for
|
||||
its duration. It is one-time and only touches unpopulated rows, so subsequent
|
||||
online writes (which now carry the tsvector via the writer fix) are unaffected.
|
||||
|
||||
Oracle slot is intentionally absent: the consolidator INSERT/UPDATE paths that
|
||||
this repairs are PostgreSQL-specific (``ops_postgresql``), and the native
|
||||
tsvector ``search_vector`` column only exists on PostgreSQL. There is no Oracle
|
||||
residue to repair.
|
||||
|
||||
Revision ID: c3f7a1b9d2e4
|
||||
Revises: f4d1c2b3a5e6
|
||||
Create Date: 2026-06-29
|
||||
"""
|
||||
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
from sqlalchemy import Connection, text
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
from hindsight_api.config import (
|
||||
DEFAULT_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE,
|
||||
ENV_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE,
|
||||
)
|
||||
|
||||
revision: str = "c3f7a1b9d2e4"
|
||||
down_revision: str | Sequence[str] | None = "f4d1c2b3a5e6"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
# Matches HindsightConfig.validate(): a tsvector regconfig name embedded as a
|
||||
# SQL literal must be a bare PG identifier.
|
||||
_PG_IDENTIFIER = re.compile(r"[a-zA-Z_][a-zA-Z0-9_]*")
|
||||
|
||||
|
||||
def _schema_prefix() -> str:
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _schema_name() -> str:
|
||||
return (context.config.get_main_option("target_schema") or "public").strip('"')
|
||||
|
||||
|
||||
def _native_language() -> str:
|
||||
"""Configured native tsvector language, validated as a PG identifier."""
|
||||
lang = os.getenv(
|
||||
ENV_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE,
|
||||
DEFAULT_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE,
|
||||
)
|
||||
if not _PG_IDENTIFIER.fullmatch(lang):
|
||||
return DEFAULT_TEXT_SEARCH_EXTENSION_NATIVE_LANGUAGE
|
||||
return lang
|
||||
|
||||
|
||||
def _is_regular_tsvector(conn: Connection, schema: str, table: str) -> bool:
|
||||
"""True iff ``schema.table.search_vector`` is a non-generated tsvector column.
|
||||
|
||||
This is the ``native`` backend signature. ``vchord`` (bm25vector) and
|
||||
``pg_textsearch`` / ``pgroonga`` / ``pg_search`` (dummy text column) all
|
||||
fail this check, so the backfill is a no-op for them.
|
||||
"""
|
||||
row = conn.execute(
|
||||
text(
|
||||
"""
|
||||
SELECT is_generated, udt_name
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = :schema
|
||||
AND table_name = :table
|
||||
AND column_name = 'search_vector'
|
||||
"""
|
||||
),
|
||||
{"schema": schema, "table": table},
|
||||
).fetchone()
|
||||
if not row:
|
||||
return False
|
||||
is_generated, udt_name = row[0], row[1]
|
||||
return udt_name == "tsvector" and is_generated != "ALWAYS"
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
schema_name = _schema_name()
|
||||
if not _is_regular_tsvector(conn, schema_name, "memory_units"):
|
||||
# Non-native backend (or column absent) — nothing to backfill.
|
||||
return
|
||||
schema_prefix = _schema_prefix()
|
||||
lang = _native_language()
|
||||
op.execute(
|
||||
f"""
|
||||
UPDATE {schema_prefix}memory_units
|
||||
SET search_vector = to_tsvector('{lang}'::regconfig, COALESCE(text, ''))
|
||||
WHERE fact_type = 'observation' AND search_vector IS NULL
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
# No-op: backfilled rows are indistinguishable from observations that were
|
||||
# populated by the post-fix writer, and reverting either to NULL would
|
||||
# re-break BM25 retrieval. The column simply stays populated.
|
||||
pass
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
-158
@@ -1,158 +0,0 @@
|
||||
"""Make maintenance routines resilient to schemas that vanish mid-scan.
|
||||
|
||||
``public.banks_needing_consolidation()`` and
|
||||
``public.schemas_with_expired_rows(...)`` snapshot the set of schemas owning a
|
||||
target table from ``pg_class`` and then run a dynamic query against each schema
|
||||
in turn. That is a time-of-check/time-of-use race: a schema (or its tables) can
|
||||
be dropped — a tenant being deleted, or a tenant migration that recreates
|
||||
tables — between the snapshot and the per-schema query, which then aborts the
|
||||
whole routine with::
|
||||
|
||||
relation "<schema>.memory_units" does not exist
|
||||
relation "<schema>.audit_log" does not exist
|
||||
|
||||
In the test suite this surfaces as cross-worker contamination: the multi-tenant
|
||||
maintenance test creates and drops ~100 ``mt<hash>_NNN`` schemas while
|
||||
``test_maintenance_routines`` (on another xdist worker, same DB) calls the
|
||||
routines. In production the background maintenance loop hits the same race when
|
||||
a tenant is removed or mid-migration.
|
||||
|
||||
Wrap each per-schema query in its own ``BEGIN ... EXCEPTION`` block so a schema
|
||||
that disappears (``undefined_table`` / ``invalid_schema_name`` /
|
||||
``undefined_column``) is skipped instead of aborting the scan. The routines stay
|
||||
``CREATE OR REPLACE`` and PostgreSQL-only, and are (re)installed only on the run
|
||||
that targets the shared ``public`` schema — same gating as the original
|
||||
install (``e5f6a7b8c9d0``) and its repair (``b2d4f6a8c1e3``).
|
||||
|
||||
Revision ID: c7e9f1a3b5d2
|
||||
Revises: e1f2a3b4c5d6
|
||||
Create Date: 2026-06-19
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "c7e9f1a3b5d2"
|
||||
down_revision: str | Sequence[str] | None = "e1f2a3b4c5d6"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _should_install_public_routines(target_schema: str | None) -> bool:
|
||||
"""True for the run that must (re)create the shared ``public.*`` routines.
|
||||
|
||||
The routines physically live in ``public``, so they are installed exactly
|
||||
once — on the base run (no ``target_schema``) or the run that explicitly
|
||||
targets ``public``. Mirrors ``b2d4f6a8c1e3``.
|
||||
"""
|
||||
return not target_schema or target_schema == "public"
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
if not _should_install_public_routines(context.config.get_main_option("target_schema")):
|
||||
return
|
||||
|
||||
# Same body as b2d4f6a8c1e3, but each per-schema query runs in its own
|
||||
# subtransaction so a schema dropped mid-scan is skipped, not fatal.
|
||||
op.execute(
|
||||
"""
|
||||
CREATE OR REPLACE FUNCTION public.banks_needing_consolidation()
|
||||
RETURNS TABLE(schema_name text, bank_id text)
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
BEGIN
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = 'memory_units' AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
RETURN QUERY EXECUTE format($q$
|
||||
SELECT %1$L::text, m.bank_id
|
||||
FROM %1$I.memory_units m
|
||||
JOIN %1$I.banks b ON b.bank_id = m.bank_id
|
||||
WHERE m.consolidated_at IS NULL
|
||||
AND m.consolidation_failed_at IS NULL
|
||||
AND m.fact_type IN ('experience', 'world')
|
||||
AND COALESCE(b.config -> 'enable_auto_consolidation', 'true'::jsonb) <> 'false'::jsonb
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM %1$I.async_operations o
|
||||
WHERE o.bank_id = m.bank_id
|
||||
AND o.operation_type = 'consolidation'
|
||||
AND o.status IN ('pending', 'processing')
|
||||
)
|
||||
GROUP BY m.bank_id
|
||||
$q$, sch);
|
||||
EXCEPTION
|
||||
-- Schema or its tables vanished between the pg_class
|
||||
-- snapshot and this query (tenant dropped or migrating).
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
op.execute(
|
||||
"""
|
||||
CREATE OR REPLACE FUNCTION public.schemas_with_expired_rows(
|
||||
p_table text, p_ts_col text, p_days int
|
||||
)
|
||||
RETURNS SETOF text
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
has_expired boolean;
|
||||
BEGIN
|
||||
IF p_days IS NULL OR p_days <= 0 THEN
|
||||
RETURN;
|
||||
END IF;
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = p_table AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
EXECUTE format(
|
||||
'SELECT EXISTS (SELECT 1 FROM %I.%I WHERE %I < NOW() - make_interval(days => $1))',
|
||||
sch, p_table, p_ts_col
|
||||
) INTO has_expired USING p_days;
|
||||
EXCEPTION
|
||||
-- Schema or its table vanished mid-scan; skip it.
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
IF has_expired THEN
|
||||
RETURN NEXT sch;
|
||||
END IF;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
# No-op: e5f6a7b8c9d0 owns these functions' lifecycle and drops them on its
|
||||
# own downgrade. This migration only re-installs them (the resilient body is
|
||||
# a strict superset of the previous behaviour), so there is nothing to undo
|
||||
# without racing that migration's DROP.
|
||||
pass
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
+6
-14
@@ -6,11 +6,8 @@ place. If a row is in ``memory_units`` it is live; if it is in
|
||||
``invalidated_memory_units`` it has been retired. Recall/consolidation/graph
|
||||
queries never need a state predicate — the rows simply aren't there.
|
||||
|
||||
The archive mirrors ``memory_units`` column-for-column — except ``embedding``,
|
||||
which it never keeps: the archive is cold storage, never a recall surface, and
|
||||
revert recomputes the embedding from the unit's text/dates/entities. Keeping no
|
||||
archive vector also means a later embedding-model switch (which re-dimensions
|
||||
``memory_units``) can't trip a dimension mismatch on the move (#2209). Plus:
|
||||
The archive mirrors ``memory_units`` column-for-column (so a row round-trips
|
||||
losslessly on revert) plus:
|
||||
- ``invalidation_reason`` optional free text recorded on invalidate
|
||||
- ``invalidated_at`` when it was retired
|
||||
- ``entity_ids`` snapshot of the unit's entity associations, so revert
|
||||
@@ -52,18 +49,13 @@ def _pg_upgrade() -> None:
|
||||
# Add edited_at to the live table FIRST so the archive's LIKE clone below
|
||||
# inherits it (keeps the two tables column-for-column identical for round-trip).
|
||||
op.execute(f"ALTER TABLE {schema}memory_units ADD COLUMN IF NOT EXISTS edited_at TIMESTAMPTZ")
|
||||
# LIKE ... INCLUDING DEFAULTS clones every memory_units column (incl.
|
||||
# edited_at) so an invalidated row can move back verbatim. We deliberately
|
||||
# omit indexes/constraints — the archive is cold storage, not a recall
|
||||
# surface; only the lookups below need indexing.
|
||||
# LIKE ... INCLUDING DEFAULTS clones every memory_units column (incl. the
|
||||
# embedding vector and edited_at) so an invalidated row can move back verbatim.
|
||||
# We deliberately omit indexes/constraints — the archive is cold storage, not a
|
||||
# recall surface; only the lookups below need indexing.
|
||||
op.execute(
|
||||
f"CREATE TABLE IF NOT EXISTS {schema}invalidated_memory_units (LIKE {schema}memory_units INCLUDING DEFAULTS)"
|
||||
)
|
||||
# ...then drop the inherited embedding: the archive never stores one (revert
|
||||
# recomputes it), so it isn't created here only to be dropped again later by
|
||||
# d4f6a8c2e1b3. That migration still runs as a no-op (DROP ... IF EXISTS) on
|
||||
# fresh DBs and does the real drop on DBs created before this column was removed.
|
||||
op.execute(f"ALTER TABLE {schema}invalidated_memory_units DROP COLUMN IF EXISTS embedding")
|
||||
op.execute(
|
||||
f"ALTER TABLE {schema}invalidated_memory_units "
|
||||
f"ADD COLUMN IF NOT EXISTS invalidation_reason TEXT, "
|
||||
|
||||
-93
@@ -1,93 +0,0 @@
|
||||
"""Drop the embedding column from the curation archive (invalidated_memory_units).
|
||||
|
||||
The archive is cold storage, never a recall surface, so it has no business
|
||||
keeping an embedding. Earlier curation code copied the live row's embedding into
|
||||
``invalidated_memory_units`` on invalidate; the engine now leaves it out on
|
||||
invalidate and recomputes it on revert, so the column is dead weight.
|
||||
|
||||
Dropping it makes "the archive holds no embedding" a schema-enforced invariant
|
||||
rather than a convention the move queries have to honour, and removes a latent
|
||||
failure mode (#2209): after an embedding-model switch the live tables are
|
||||
re-dimensioned but the archive was not, so a stale old-dimension embedding in
|
||||
the archive tripped a vector-dimension mismatch on the INSERT … SELECT
|
||||
round-trip. With no column at all, there is nothing to mismatch.
|
||||
|
||||
The creation sites no longer add the column (the PG ``LIKE`` clone in
|
||||
c9a1b2d3e4f5 drops it; the Oracle baseline omits it), so on a fresh database
|
||||
this migration is a no-op (DROP ... IF EXISTS / Oracle ORA-00904 swallow). It
|
||||
does the real work on databases created before the column was removed there.
|
||||
|
||||
DROP COLUMN is a metadata-only operation on both PostgreSQL and Oracle 23ai (no
|
||||
table rewrite), so it is cheap even across many tenant schemas. The downgrade
|
||||
re-adds an unconstrained vector column (any dimension) — empty, since the
|
||||
embeddings are intentionally discarded.
|
||||
|
||||
Revision ID: d4f6a8c2e1b3
|
||||
Revises: a1d3f5b7c9e2
|
||||
Create Date: 2026-06-15
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "d4f6a8c2e1b3"
|
||||
down_revision: str | Sequence[str] | None = "a1d3f5b7c9e2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _pg_schema_prefix() -> str:
|
||||
"""Schema-qualifier for raw SQL on PG (multi-tenant search_path)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
op.execute(f"ALTER TABLE {schema}invalidated_memory_units DROP COLUMN IF EXISTS embedding")
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
schema = _pg_schema_prefix()
|
||||
# Unconstrained `vector` (no dimension) so the re-added column accepts any
|
||||
# model's embeddings; it comes back empty regardless.
|
||||
op.execute(f"ALTER TABLE {schema}invalidated_memory_units ADD COLUMN IF NOT EXISTS embedding vector")
|
||||
|
||||
|
||||
def _oracle_upgrade() -> None:
|
||||
# Oracle has no `DROP COLUMN IF EXISTS`; swallow ORA-00904 (column does not
|
||||
# exist) so the migration is idempotent and safe on a fresh schema whose
|
||||
# baseline already omits the column.
|
||||
op.execute(
|
||||
"""
|
||||
BEGIN
|
||||
EXECUTE IMMEDIATE 'ALTER TABLE invalidated_memory_units DROP COLUMN embedding';
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
IF SQLCODE != -904 THEN RAISE; END IF;
|
||||
END;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _oracle_downgrade() -> None:
|
||||
# Swallow ORA-01430 (column already exists) for idempotency.
|
||||
op.execute(
|
||||
"""
|
||||
BEGIN
|
||||
EXECUTE IMMEDIATE 'ALTER TABLE invalidated_memory_units ADD (embedding VECTOR)';
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
IF SQLCODE != -1430 THEN RAISE; END IF;
|
||||
END;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade, oracle=_oracle_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade, oracle=_oracle_downgrade)
|
||||
-39
@@ -1,39 +0,0 @@
|
||||
"""Merge two divergent migration heads.
|
||||
|
||||
``d4f6a8c2e1b3`` (drop the curation-archive embedding column) and
|
||||
``2071c7518f88`` (add the memory_links(bank_id, link_type) index) were authored
|
||||
in parallel off the same parent (``a1d3f5b7c9e2``) and merged independently,
|
||||
leaving the DAG with two heads. This is a no-op merge that re-unifies them so
|
||||
``alembic upgrade head`` is unambiguous again (enforced by
|
||||
``tests/test_alembic_dag.py::test_single_head``).
|
||||
|
||||
Revision ID: e1f2a3b4c5d6
|
||||
Revises: d4f6a8c2e1b3, 2071c7518f88
|
||||
Create Date: 2026-06-16
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "e1f2a3b4c5d6"
|
||||
down_revision: str | Sequence[str] | None = ("d4f6a8c2e1b3", "2071c7518f88")
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
# Pure DAG merge — both parents already applied their schema changes.
|
||||
pass
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
pass
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
-110
@@ -1,110 +0,0 @@
|
||||
"""Add server-side routine for cron-scheduled mental model refresh.
|
||||
|
||||
Installs ``public.mental_models_with_cron()`` — a discovery routine that returns
|
||||
every mental model carrying a non-empty ``trigger->>'refresh_cron'`` across all
|
||||
tenant schemas in one round-trip (the same per-schema scan as the other
|
||||
maintenance routines from ``e5f6a7b8c9d0``). The maintenance loop evaluates each
|
||||
candidate's cron expression in Python (``croniter``) against ``last_refreshed_at``
|
||||
to decide whether a scheduled refresh is due — cron arithmetic isn't expressible
|
||||
in plain SQL — and only the cron *candidate set* is discovered here.
|
||||
|
||||
Models that already have a ``refresh_mental_model`` operation pending/processing
|
||||
are excluded so a slow refresh isn't double-queued (mirrors the in-flight guard
|
||||
in ``banks_needing_consolidation``). Each per-schema query runs in its own
|
||||
``BEGIN ... EXCEPTION`` subtransaction so a schema dropped mid-scan (tenant
|
||||
deletion / migration) is skipped, not fatal — same resilience as
|
||||
``c7e9f1a3b5d2``.
|
||||
|
||||
Read-only (STABLE) discovery routine — the caller performs the refresh enqueue —
|
||||
so installing it never mutates data. PostgreSQL only: the worker poller and the
|
||||
maintenance loop are PG-only (Oracle slot intentionally absent, mirroring
|
||||
``e5f6a7b8c9d0``). The routine lives in ``public`` and is CREATE OR REPLACE, so
|
||||
it is installed exactly once (base / ``public`` run) to avoid the
|
||||
``tuple concurrently updated`` race on concurrent per-tenant runs.
|
||||
|
||||
Revision ID: f4d1c2b3a5e6
|
||||
Revises: c7e9f1a3b5d2
|
||||
Create Date: 2026-06-23
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
from hindsight_api.alembic._dialect import run_for_dialect
|
||||
|
||||
revision: str = "f4d1c2b3a5e6"
|
||||
down_revision: str | Sequence[str] | None = "c7e9f1a3b5d2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _should_install_public_routines(target_schema: str | None) -> bool:
|
||||
"""True for the run that must (re)create the shared ``public.*`` routine.
|
||||
|
||||
The routine physically lives in ``public``, so it is installed exactly once —
|
||||
on the base run (no ``target_schema``) or the run that explicitly targets
|
||||
``public``. Mirrors ``c7e9f1a3b5d2``.
|
||||
"""
|
||||
return not target_schema or target_schema == "public"
|
||||
|
||||
|
||||
def _pg_upgrade() -> None:
|
||||
if not _should_install_public_routines(context.config.get_main_option("target_schema")):
|
||||
return
|
||||
|
||||
op.execute(
|
||||
"""
|
||||
CREATE OR REPLACE FUNCTION public.mental_models_with_cron()
|
||||
RETURNS TABLE(schema_name text, bank_id text, mental_model_id text,
|
||||
refresh_cron text, last_refreshed_at timestamptz)
|
||||
LANGUAGE plpgsql STABLE
|
||||
AS $fn$
|
||||
DECLARE
|
||||
sch text;
|
||||
BEGIN
|
||||
FOR sch IN
|
||||
SELECT n.nspname
|
||||
FROM pg_class c
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE c.relname = 'mental_models' AND c.relkind = 'r'
|
||||
LOOP
|
||||
BEGIN
|
||||
RETURN QUERY EXECUTE format($q$
|
||||
SELECT %1$L::text, mm.bank_id::text, mm.id::text,
|
||||
mm.trigger->>'refresh_cron', mm.last_refreshed_at
|
||||
FROM %1$I.mental_models mm
|
||||
WHERE COALESCE(mm.trigger->>'refresh_cron', '') <> ''
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM %1$I.async_operations o
|
||||
WHERE o.bank_id = mm.bank_id
|
||||
AND o.operation_type = 'refresh_mental_model'
|
||||
AND o.status IN ('pending', 'processing')
|
||||
AND o.task_payload->>'mental_model_id' = mm.id::text
|
||||
)
|
||||
$q$, sch);
|
||||
EXCEPTION
|
||||
-- Schema or its tables vanished between the pg_class
|
||||
-- snapshot and this query (tenant dropped or migrating).
|
||||
WHEN undefined_table OR invalid_schema_name OR undefined_column THEN
|
||||
CONTINUE;
|
||||
END;
|
||||
END LOOP;
|
||||
END;
|
||||
$fn$;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _pg_downgrade() -> None:
|
||||
if not _should_install_public_routines(context.config.get_main_option("target_schema")):
|
||||
return
|
||||
op.execute("DROP FUNCTION IF EXISTS public.mental_models_with_cron()")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
run_for_dialect(pg=_pg_upgrade)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
run_for_dialect(pg=_pg_downgrade)
|
||||
@@ -142,9 +142,6 @@ _TABLES: tuple[str, ...] = (
|
||||
# Cold archive for curation: invalidated facts are MOVED here out of
|
||||
# memory_units so the recall hot-path never sees them. Mirrors memory_units
|
||||
# plus invalidation bookkeeping and an entity-id snapshot for lossless revert.
|
||||
# No `embedding` column: the archive is cold storage and revert recomputes the
|
||||
# embedding, so there is no archive vector to fall out of sync with the live
|
||||
# model's dimension on a model switch (#2209).
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS invalidated_memory_units (
|
||||
id RAW(16) NOT NULL,
|
||||
@@ -152,6 +149,7 @@ _TABLES: tuple[str, ...] = (
|
||||
document_id VARCHAR2(512),
|
||||
chunk_id VARCHAR2(512),
|
||||
text CLOB NOT NULL,
|
||||
embedding VECTOR(384, FLOAT32),
|
||||
context CLOB,
|
||||
event_date TIMESTAMP WITH TIME ZONE NOT NULL,
|
||||
occurred_start TIMESTAMP WITH TIME ZONE,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -9,7 +9,7 @@ from fastmcp import FastMCP
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
from hindsight_api import __version__ as HINDSIGHT_VERSION
|
||||
from hindsight_api.config import DEFAULT_MCP_RECALL_DESCRIPTION, DEFAULT_MCP_RETAIN_DESCRIPTION, _get_raw_config
|
||||
from hindsight_api.config import _get_raw_config
|
||||
from hindsight_api.engine.memory_engine import _current_schema
|
||||
from hindsight_api.extensions import MCPExtension, load_extension
|
||||
from hindsight_api.extensions.tenant import AuthenticationError
|
||||
@@ -78,19 +78,6 @@ def get_current_mcp_authenticated() -> bool:
|
||||
return _current_mcp_authenticated.get()
|
||||
|
||||
|
||||
def _build_mcp_tool_descriptions(extra_instructions: str | None) -> tuple[str | None, str | None]:
|
||||
"""Return custom retain/recall descriptions when server-level MCP instructions are set."""
|
||||
if not isinstance(extra_instructions, str):
|
||||
return None, None
|
||||
|
||||
extra_instructions = extra_instructions.strip()
|
||||
if not extra_instructions:
|
||||
return None, None
|
||||
|
||||
suffix = f"\n\nAdditional instructions: {extra_instructions}"
|
||||
return DEFAULT_MCP_RETAIN_DESCRIPTION + suffix, DEFAULT_MCP_RECALL_DESCRIPTION + suffix
|
||||
|
||||
|
||||
def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||
"""
|
||||
Create and configure the Hindsight MCP server.
|
||||
@@ -148,10 +135,6 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||
allowed = frozenset(global_config.mcp_enabled_tools)
|
||||
base_tools = (base_tools if base_tools is not None else _ALL_TOOLS) & allowed
|
||||
|
||||
retain_description, recall_description = _build_mcp_tool_descriptions(
|
||||
getattr(global_config, "mcp_instructions", None)
|
||||
)
|
||||
|
||||
# Configure and register tools using shared module
|
||||
config = MCPToolsConfig(
|
||||
bank_id_resolver=get_current_bank_id,
|
||||
@@ -161,8 +144,6 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||
mcp_authenticated_resolver=get_current_mcp_authenticated, # Propagate MCP pre-auth flag
|
||||
include_bank_id_param=multi_bank,
|
||||
tools=base_tools,
|
||||
retain_description=retain_description,
|
||||
recall_description=recall_description,
|
||||
)
|
||||
|
||||
register_mcp_tools(mcp, memory, config)
|
||||
|
||||
@@ -1,263 +0,0 @@
|
||||
"""Open Knowledge Format (OKF) projection for knowledge pages.
|
||||
|
||||
Knowledge pages are a *read-only* OKF view over the existing mental models: each
|
||||
mental model is projected into an OKF document — a markdown body with YAML
|
||||
frontmatter (``type`` required; ``title``/``description``/``tags``/``timestamp``
|
||||
optional) — and pages are linked into a constellation graph via shared tags.
|
||||
|
||||
See the Open Knowledge Format spec:
|
||||
https://github.com/GoogleCloudPlatform/knowledge-catalog/tree/main/okf
|
||||
|
||||
This module is intentionally pure: every function transforms the mental-model
|
||||
dicts returned by ``MemoryEngine.list_mental_models`` / ``get_mental_model`` and
|
||||
never touches the database. That keeps the OKF contract unit-testable without a
|
||||
DB or LLM and lets the HTTP layer stay a thin wrapper.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
# OKF requires exactly one frontmatter field — ``type``. We default to this when
|
||||
# a page does not declare one via a ``type:<x>`` tag.
|
||||
DEFAULT_PAGE_TYPE = "knowledge-page"
|
||||
|
||||
# A page declares its OKF ``type`` through a tag of the form ``type:runbook``.
|
||||
# This keeps the projection schema-free (no new mental_models column): the type
|
||||
# is lifted from the existing tags array.
|
||||
TYPE_TAG_PREFIX = "type:"
|
||||
|
||||
INDEX_FILENAME = "index.md"
|
||||
|
||||
# Deterministic, colour-blind-friendly palette. Type → colour is stable across
|
||||
# requests so the constellation keeps the same colours between reloads.
|
||||
_PALETTE = (
|
||||
"#0074d9", # blue
|
||||
"#2ecc40", # green
|
||||
"#b10dc9", # purple
|
||||
"#ff851b", # orange
|
||||
"#39cccc", # teal
|
||||
"#f012be", # magenta
|
||||
"#3d9970", # olive
|
||||
"#ff4136", # red
|
||||
)
|
||||
|
||||
_EDGE_COLOR = "#9aa5b1"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PageType:
|
||||
"""A page's OKF ``type`` and the tags that remain after the type tag is split off."""
|
||||
|
||||
type: str
|
||||
display_tags: list[str]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class KnowledgeGraph:
|
||||
"""Cytoscape-style node/edge graph of knowledge pages linked by shared tags."""
|
||||
|
||||
nodes: list[dict[str, Any]] = field(default_factory=list)
|
||||
edges: list[dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
|
||||
def _color_for(key: str) -> str:
|
||||
"""Stable colour for a string key (FNV-ish hash into the fixed palette)."""
|
||||
h = 0
|
||||
for ch in key:
|
||||
h = (h * 31 + ord(ch)) & 0xFFFFFFFF
|
||||
return _PALETTE[h % len(_PALETTE)]
|
||||
|
||||
|
||||
def _scalar(value: Any) -> str:
|
||||
"""Emit a YAML-safe double-quoted scalar.
|
||||
|
||||
We always double-quote so arbitrary page names / source queries can't be
|
||||
misread as YAML special forms (``true``, ``2026-01-01``, ``- x``, etc.).
|
||||
"""
|
||||
text = str(value)
|
||||
escaped = text.replace("\\", "\\\\").replace('"', '\\"').replace("\n", "\\n").replace("\r", "")
|
||||
return f'"{escaped}"'
|
||||
|
||||
|
||||
def page_type(tags: list[str] | None) -> PageType:
|
||||
"""Split an OKF ``type`` out of the tag list.
|
||||
|
||||
The first ``type:<x>`` tag wins; all ``type:`` tags are removed from the
|
||||
returned ``display_tags`` so they don't pollute the constellation's
|
||||
shared-tag edges. Falls back to :data:`DEFAULT_PAGE_TYPE`.
|
||||
"""
|
||||
resolved = DEFAULT_PAGE_TYPE
|
||||
display: list[str] = []
|
||||
for tag in tags or []:
|
||||
if tag.startswith(TYPE_TAG_PREFIX):
|
||||
suffix = tag[len(TYPE_TAG_PREFIX) :].strip()
|
||||
if suffix and resolved == DEFAULT_PAGE_TYPE:
|
||||
resolved = suffix
|
||||
continue
|
||||
display.append(tag)
|
||||
return PageType(type=resolved, display_tags=display)
|
||||
|
||||
|
||||
def _timestamp(mm: dict[str, Any]) -> str | None:
|
||||
return mm.get("last_refreshed_at") or mm.get("created_at")
|
||||
|
||||
|
||||
def frontmatter(mm: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Build the ordered OKF frontmatter mapping for a mental model.
|
||||
|
||||
``None``/empty values are dropped by :func:`render_frontmatter`.
|
||||
"""
|
||||
pt = page_type(mm.get("tags"))
|
||||
return {
|
||||
"id": mm.get("id"),
|
||||
"type": pt.type,
|
||||
"title": mm.get("name"),
|
||||
"description": mm.get("source_query"),
|
||||
"tags": pt.display_tags,
|
||||
"timestamp": _timestamp(mm),
|
||||
}
|
||||
|
||||
|
||||
def render_frontmatter(fm: dict[str, Any]) -> str:
|
||||
"""Render a frontmatter mapping into a ``---`` fenced YAML block."""
|
||||
lines = ["---"]
|
||||
for key, value in fm.items():
|
||||
if value is None:
|
||||
continue
|
||||
if isinstance(value, list):
|
||||
if not value:
|
||||
continue
|
||||
lines.append(f"{key}:")
|
||||
lines.extend(f" - {_scalar(item)}" for item in value)
|
||||
else:
|
||||
lines.append(f"{key}: {_scalar(value)}")
|
||||
lines.append("---")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def render_document(mm: dict[str, Any]) -> str:
|
||||
"""Render a full OKF document: frontmatter block + markdown body."""
|
||||
body = (mm.get("content") or "").strip()
|
||||
return f"{render_frontmatter(frontmatter(mm))}\n\n{body}\n" if body else f"{render_frontmatter(frontmatter(mm))}\n"
|
||||
|
||||
|
||||
def page_filename(page_id: str) -> str:
|
||||
"""OKF bundle filename for a page id."""
|
||||
return f"{page_id}.md"
|
||||
|
||||
|
||||
def log_filename(page_id: str) -> str:
|
||||
"""OKF reserved per-page history filename."""
|
||||
return f"{page_id}.log.md"
|
||||
|
||||
|
||||
def render_index(nodes: list[dict[str, Any]]) -> str:
|
||||
"""Render the reserved ``index.md`` — nested OKF navigation over the tree.
|
||||
|
||||
``nodes`` is the flat folder/page list (each with ``id``, ``kind``, ``name``,
|
||||
``parent_id``); folders nest their children, pages link to their ``.md``.
|
||||
"""
|
||||
fm = render_frontmatter({"type": "index", "title": "Knowledge base"})
|
||||
lines = [fm, "", "# Knowledge base", ""]
|
||||
|
||||
children: dict[Any, list[dict[str, Any]]] = {}
|
||||
for node in nodes:
|
||||
children.setdefault(node.get("parent_id"), []).append(node)
|
||||
|
||||
def walk(parent: Any, depth: int) -> None:
|
||||
ordered = sorted(children.get(parent, []), key=lambda n: (n.get("sort_order", 0), n.get("name") or ""))
|
||||
for node in ordered:
|
||||
indent = " " * depth
|
||||
if node.get("kind") == "folder":
|
||||
lines.append(f"{indent}- **{node['name']}/**")
|
||||
walk(node["id"], depth + 1)
|
||||
else:
|
||||
description = node.get("source_query") or node.get("description")
|
||||
link = f"{indent}- [{node['name']}](./{page_filename(node['id'])})"
|
||||
lines.append(f"{link} — {description}" if description else link)
|
||||
|
||||
walk(None, 0)
|
||||
if len(lines) == 4:
|
||||
lines.append("_No knowledge pages yet._")
|
||||
return "\n".join(lines) + "\n"
|
||||
|
||||
|
||||
def render_log(mm: dict[str, Any], history: list[dict[str, Any]]) -> str:
|
||||
"""Render the reserved per-page ``log.md`` from refresh history.
|
||||
|
||||
Each history entry is ``{previous_content, previous_reflect_response,
|
||||
changed_at}`` (newest first), capturing the content *before* a refresh.
|
||||
"""
|
||||
name = mm.get("name") or mm.get("id")
|
||||
fm = render_frontmatter({"type": "log", "title": f"{name} — history"})
|
||||
lines = [fm, "", f"# {name} — history", ""]
|
||||
if not history:
|
||||
lines.append("_No refresh history._")
|
||||
return "\n".join(lines) + "\n"
|
||||
for entry in history:
|
||||
changed_at = entry.get("changed_at") or "unknown"
|
||||
previous = (entry.get("previous_content") or "").strip()
|
||||
lines.append(f"## {changed_at}")
|
||||
lines.append("")
|
||||
lines.append(previous if previous else "_(empty)_")
|
||||
lines.append("")
|
||||
return "\n".join(lines).rstrip() + "\n"
|
||||
|
||||
|
||||
def knowledge_graph(
|
||||
pages: list[dict[str, Any]],
|
||||
cluster_for: "Callable[[dict[str, Any]], str] | None" = None,
|
||||
) -> KnowledgeGraph:
|
||||
"""Derive the constellation graph: pages as nodes, shared tags as edges.
|
||||
|
||||
Two pages are linked when they share at least one (non-``type:``) tag; the
|
||||
edge weight is the number of shared tags. Each node's cluster (``type`` field
|
||||
+ colour) comes from ``cluster_for(page)`` — the knowledge base groups by
|
||||
parent folder; the default groups by OKF ``type``.
|
||||
"""
|
||||
nodes: list[dict[str, Any]] = []
|
||||
tag_sets: list[tuple[str, frozenset[str]]] = []
|
||||
for mm in pages:
|
||||
page_id = mm["id"]
|
||||
pt = page_type(mm.get("tags"))
|
||||
cluster = cluster_for(mm) if cluster_for else pt.type
|
||||
tag_sets.append((page_id, frozenset(pt.display_tags)))
|
||||
nodes.append(
|
||||
{
|
||||
"data": {
|
||||
"id": page_id,
|
||||
"label": mm.get("name") or page_id,
|
||||
"type": cluster,
|
||||
"tagCount": len(pt.display_tags),
|
||||
"color": _color_for(cluster),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
edges: list[dict[str, Any]] = []
|
||||
for i in range(len(tag_sets)):
|
||||
source_id, source_tags = tag_sets[i]
|
||||
if not source_tags:
|
||||
continue
|
||||
for j in range(i + 1, len(tag_sets)):
|
||||
target_id, target_tags = tag_sets[j]
|
||||
shared = source_tags & target_tags
|
||||
if not shared:
|
||||
continue
|
||||
edges.append(
|
||||
{
|
||||
"data": {
|
||||
"id": f"{source_id}--{target_id}",
|
||||
"source": source_id,
|
||||
"target": target_id,
|
||||
"sharedTags": sorted(shared),
|
||||
"weight": len(shared),
|
||||
"color": _EDGE_COLOR,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
return KnowledgeGraph(nodes=nodes, edges=edges)
|
||||
@@ -142,35 +142,11 @@ ENV_LLM_REASONING_EFFORT = "HINDSIGHT_API_LLM_REASONING_EFFORT"
|
||||
ENV_LLM_GROQ_SERVICE_TIER = "HINDSIGHT_API_LLM_GROQ_SERVICE_TIER"
|
||||
ENV_LLM_OPENAI_SERVICE_TIER = "HINDSIGHT_API_LLM_OPENAI_SERVICE_TIER"
|
||||
ENV_LLM_BEDROCK_SERVICE_TIER = "HINDSIGHT_API_LLM_BEDROCK_SERVICE_TIER"
|
||||
ENV_LLM_GEMINI_SERVICE_TIER = "HINDSIGHT_API_LLM_GEMINI_SERVICE_TIER"
|
||||
ENV_LLM_EXTRA_BODY = "HINDSIGHT_API_LLM_EXTRA_BODY"
|
||||
ENV_LLM_DEFAULT_HEADERS = "HINDSIGHT_API_LLM_DEFAULT_HEADERS"
|
||||
ENV_LLM_STRICT_SCHEMA = "HINDSIGHT_API_LLM_STRICT_SCHEMA"
|
||||
ENV_LLM_SEND_BANK_AS_USER = "HINDSIGHT_API_LLM_SEND_BANK_AS_USER"
|
||||
|
||||
# Per-operation sampling temperature. Each internal LLM call uses a temperature
|
||||
# tuned for its task (deterministic extraction vs. creative reflection). These
|
||||
# expose those as overridable knobs. Resolution per operation:
|
||||
# per-operation env -> global env (ENV_LLM_TEMPERATURE) -> built-in default.
|
||||
# A value of "none"/"default"/"" (or "off") omits the temperature parameter
|
||||
# entirely, for models that reject explicit temperatures (e.g. Azure GPT-5.5,
|
||||
# which only accepts the default value) -- see issue #2459.
|
||||
ENV_LLM_TEMPERATURE = "HINDSIGHT_API_LLM_TEMPERATURE"
|
||||
ENV_LLM_TEMPERATURE_VERIFICATION = "HINDSIGHT_API_LLM_TEMPERATURE_VERIFICATION"
|
||||
ENV_LLM_TEMPERATURE_RETAIN = "HINDSIGHT_API_LLM_TEMPERATURE_RETAIN"
|
||||
ENV_LLM_TEMPERATURE_REFLECT = "HINDSIGHT_API_LLM_TEMPERATURE_REFLECT"
|
||||
ENV_LLM_TEMPERATURE_CONSOLIDATION = "HINDSIGHT_API_LLM_TEMPERATURE_CONSOLIDATION"
|
||||
|
||||
# Multi-LLM strategy. Extra LLMs are configured by index alongside the unindexed
|
||||
# primary (e.g. HINDSIGHT_API_LLM_1_PROVIDER, HINDSIGHT_API_LLM_2_PROVIDER, ...),
|
||||
# and HINDSIGHT_API_LLM_STRATEGY (JSON) selects how to route across them — see
|
||||
# _parse_llm_members / _parse_llm_strategy below. Each operation can override the
|
||||
# global chain with its own HINDSIGHT_API_<OP>_LLM_<n>_* members + _STRATEGY.
|
||||
ENV_LLM_STRATEGY = "HINDSIGHT_API_LLM_STRATEGY"
|
||||
ENV_RETAIN_LLM_STRATEGY = "HINDSIGHT_API_RETAIN_LLM_STRATEGY"
|
||||
ENV_REFLECT_LLM_STRATEGY = "HINDSIGHT_API_REFLECT_LLM_STRATEGY"
|
||||
ENV_CONSOLIDATION_LLM_STRATEGY = "HINDSIGHT_API_CONSOLIDATION_LLM_STRATEGY"
|
||||
|
||||
# LiteLLM Router chain — provider-specific config consumed by the "litellmrouter"
|
||||
# provider. Each entry is a deployment; the Router tries them in declared order and
|
||||
# falls back to the next on transient errors (5xx, rate-limit, timeout).
|
||||
@@ -179,75 +155,15 @@ ENV_CONSOLIDATION_LLM_STRATEGY = "HINDSIGHT_API_CONSOLIDATION_LLM_STRATEGY"
|
||||
# disambiguates from the embeddings/reranker LITELLM_* settings.
|
||||
ENV_LLM_LITELLMROUTER_CONFIG = "HINDSIGHT_API_LLM_LITELLMROUTER_CONFIG"
|
||||
|
||||
# Per-operation temperature defaults (preserve historical hardcoded values).
|
||||
DEFAULT_LLM_TEMPERATURE_VERIFICATION = 0.0 # connection check
|
||||
DEFAULT_LLM_TEMPERATURE_RETAIN = 0.1 # fact extraction
|
||||
DEFAULT_LLM_TEMPERATURE_REFLECT = 0.9 # reflect "thinking"
|
||||
DEFAULT_LLM_TEMPERATURE_CONSOLIDATION = 0.0 # mental-model delta / dedup
|
||||
|
||||
# Defaults for service tiers
|
||||
DEFAULT_LLM_GROQ_SERVICE_TIER = "auto" # "on_demand", "flex", or "auto"
|
||||
DEFAULT_LLM_OPENAI_SERVICE_TIER = None # None (default) or "flex" (50% cheaper)
|
||||
DEFAULT_LLM_BEDROCK_SERVICE_TIER = None # None (default), "flex", "priority", or "reserved"
|
||||
DEFAULT_LLM_GEMINI_SERVICE_TIER = None # None (default) or "flex" (50% cheaper best-effort tier)
|
||||
DEFAULT_LLM_EXTRA_BODY = None # None = no extra body params; JSON dict merged into OpenAI extra_body
|
||||
DEFAULT_LLM_DEFAULT_HEADERS = (
|
||||
None # None = no extra headers; JSON dict passed as default_headers to provider SDK clients
|
||||
)
|
||||
|
||||
|
||||
def parse_gemini_service_tier(value: str | None) -> str | None:
|
||||
"""Normalize and validate the Gemini service tier."""
|
||||
tier = value or None
|
||||
valid_tiers = (None, "flex")
|
||||
if tier not in valid_tiers:
|
||||
raise ValueError(
|
||||
f"Invalid HINDSIGHT_API_LLM_GEMINI_SERVICE_TIER: "
|
||||
f"{tier!r}. Must be one of: {', '.join(t for t in valid_tiers if t is not None)}."
|
||||
)
|
||||
return tier
|
||||
|
||||
|
||||
# Sentinel strings that, as a temperature value, mean "omit the temperature
|
||||
# parameter entirely" rather than a numeric setting.
|
||||
_TEMPERATURE_OMIT_VALUES = frozenset({"", "none", "default", "off", "unset"})
|
||||
|
||||
|
||||
def _parse_temperature(raw: str) -> float | None:
|
||||
"""Parse a raw temperature env value into a float, or None to omit it.
|
||||
|
||||
Returns None for the omit sentinels (so the temperature parameter is dropped
|
||||
from the LLM call); otherwise parses a float and validates the 0.0-2.0 range.
|
||||
"""
|
||||
if raw.strip().lower() in _TEMPERATURE_OMIT_VALUES:
|
||||
return None
|
||||
try:
|
||||
value = float(raw)
|
||||
except ValueError as e:
|
||||
raise ValueError(
|
||||
f"Invalid LLM temperature {raw!r}: must be a number in [0.0, 2.0] "
|
||||
f"or one of {sorted(_TEMPERATURE_OMIT_VALUES)} to omit it."
|
||||
) from e
|
||||
if not 0.0 <= value <= 2.0:
|
||||
raise ValueError(f"Invalid LLM temperature {value}: must be in [0.0, 2.0].")
|
||||
return value
|
||||
|
||||
|
||||
def _resolve_operation_temperature(operation_env: str, default: float) -> float | None:
|
||||
"""Resolve a per-operation temperature: per-op env -> global env -> default.
|
||||
|
||||
The omit sentinels resolve to None at any layer, so a single
|
||||
``HINDSIGHT_API_LLM_TEMPERATURE=none`` drops temperature from every operation
|
||||
that has no explicit per-operation override.
|
||||
"""
|
||||
raw = os.getenv(operation_env)
|
||||
if raw is None:
|
||||
raw = os.getenv(ENV_LLM_TEMPERATURE)
|
||||
if raw is None:
|
||||
return default
|
||||
return _parse_temperature(raw)
|
||||
|
||||
|
||||
# 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"
|
||||
@@ -341,11 +257,6 @@ ENV_RERANKER_OPENROUTER_API_KEY = "HINDSIGHT_API_RERANKER_OPENROUTER_API_KEY"
|
||||
ENV_RERANKER_OPENROUTER_MODEL = "HINDSIGHT_API_RERANKER_OPENROUTER_MODEL"
|
||||
ENV_RERANKER_OPENROUTER_BASE_URL = "HINDSIGHT_API_RERANKER_OPENROUTER_BASE_URL"
|
||||
|
||||
# Requesty configuration (OpenAI-compatible gateway; embeddings)
|
||||
ENV_REQUESTY_API_KEY = "HINDSIGHT_API_REQUESTY_API_KEY"
|
||||
ENV_EMBEDDINGS_REQUESTY_API_KEY = "HINDSIGHT_API_EMBEDDINGS_REQUESTY_API_KEY"
|
||||
ENV_EMBEDDINGS_REQUESTY_MODEL = "HINDSIGHT_API_EMBEDDINGS_REQUESTY_MODEL"
|
||||
|
||||
# ZeroEntropy configuration (embeddings)
|
||||
ENV_EMBEDDINGS_ZEROENTROPY_API_KEY = "HINDSIGHT_API_EMBEDDINGS_ZEROENTROPY_API_KEY"
|
||||
ENV_EMBEDDINGS_ZEROENTROPY_MODEL = "HINDSIGHT_API_EMBEDDINGS_ZEROENTROPY_MODEL"
|
||||
@@ -443,10 +354,8 @@ ENV_ACCESS_LOG = "HINDSIGHT_API_ACCESS_LOG"
|
||||
ENV_MCP_ENABLED = "HINDSIGHT_API_MCP_ENABLED"
|
||||
ENV_MCP_ENABLED_TOOLS = "HINDSIGHT_API_MCP_ENABLED_TOOLS"
|
||||
ENV_MCP_STATELESS = "HINDSIGHT_API_MCP_STATELESS"
|
||||
ENV_MCP_INSTRUCTIONS = "HINDSIGHT_API_MCP_INSTRUCTIONS"
|
||||
ENV_ENABLE_BANK_CONFIG_API = "HINDSIGHT_API_ENABLE_BANK_CONFIG_API"
|
||||
ENV_ENABLE_BANK_LLM_HEALTH = "HINDSIGHT_API_ENABLE_BANK_LLM_HEALTH"
|
||||
ENV_ENABLE_DRY_RUN_EXTRACT = "HINDSIGHT_API_ENABLE_DRY_RUN_EXTRACT"
|
||||
ENV_DEFAULT_BANK_TEMPLATE = "HINDSIGHT_API_DEFAULT_BANK_TEMPLATE"
|
||||
ENV_GRAPH_RETRIEVER = "HINDSIGHT_API_GRAPH_RETRIEVER"
|
||||
ENV_RECALL_MAX_CONCURRENT = "HINDSIGHT_API_RECALL_MAX_CONCURRENT"
|
||||
@@ -465,7 +374,6 @@ ENV_OTEL_EXPORTER_OTLP_HEADERS = "HINDSIGHT_API_OTEL_EXPORTER_OTLP_HEADERS"
|
||||
ENV_OTEL_SERVICE_NAME = "HINDSIGHT_API_OTEL_SERVICE_NAME"
|
||||
ENV_OTEL_DEPLOYMENT_ENVIRONMENT = "HINDSIGHT_API_OTEL_DEPLOYMENT_ENVIRONMENT"
|
||||
ENV_METRICS_INCLUDE_BANK_ID = "HINDSIGHT_API_METRICS_INCLUDE_BANK_ID"
|
||||
ENV_METRICS_BACKLOG_ENABLED = "HINDSIGHT_API_METRICS_BACKLOG_ENABLED"
|
||||
|
||||
# Vertex AI configuration
|
||||
ENV_LLM_VERTEXAI_PROJECT_ID = "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID"
|
||||
@@ -488,7 +396,6 @@ ENV_LLM_PROMPT_CACHE_ENABLED = "HINDSIGHT_API_LLM_PROMPT_CACHE_ENABLED"
|
||||
# Retain settings
|
||||
ENV_RETAIN_MAX_COMPLETION_TOKENS = "HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS"
|
||||
ENV_RETAIN_CHUNK_SIZE = "HINDSIGHT_API_RETAIN_CHUNK_SIZE"
|
||||
ENV_RETAIN_STRUCTURED_CHUNK_SIZE = "HINDSIGHT_API_RETAIN_STRUCTURED_CHUNK_SIZE"
|
||||
ENV_RETAIN_EXTRACT_CAUSAL_LINKS = "HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS"
|
||||
ENV_RETAIN_EXTRACTION_MODE = "HINDSIGHT_API_RETAIN_EXTRACTION_MODE"
|
||||
ENV_RETAIN_MISSION = "HINDSIGHT_API_RETAIN_MISSION"
|
||||
@@ -515,11 +422,6 @@ ENV_FILE_STORAGE_AZURE_ACCOUNT_NAME = "HINDSIGHT_API_FILE_STORAGE_AZURE_ACCOUNT_
|
||||
ENV_FILE_STORAGE_AZURE_ACCOUNT_KEY = "HINDSIGHT_API_FILE_STORAGE_AZURE_ACCOUNT_KEY"
|
||||
ENV_FILE_PARSER = "HINDSIGHT_API_FILE_PARSER"
|
||||
ENV_FILE_PARSER_ALLOWLIST = "HINDSIGHT_API_FILE_PARSER_ALLOWLIST"
|
||||
ENV_FILE_PARSER_MARKITDOWN_OCR_ENABLED = "HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_ENABLED"
|
||||
ENV_FILE_PARSER_MARKITDOWN_OCR_API_KEY = "HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_API_KEY"
|
||||
ENV_FILE_PARSER_MARKITDOWN_OCR_BASE_URL = "HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_BASE_URL"
|
||||
ENV_FILE_PARSER_MARKITDOWN_OCR_MODEL = "HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_MODEL"
|
||||
ENV_FILE_PARSER_MARKITDOWN_OCR_PROMPT = "HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_PROMPT"
|
||||
ENV_FILE_PARSER_IRIS_TOKEN = "HINDSIGHT_API_FILE_PARSER_IRIS_TOKEN"
|
||||
ENV_FILE_PARSER_IRIS_ORG_ID = "HINDSIGHT_API_FILE_PARSER_IRIS_ORG_ID"
|
||||
ENV_FILE_PARSER_LLAMA_PARSE_API_KEY = "HINDSIGHT_API_FILE_PARSER_LLAMA_PARSE_API_KEY"
|
||||
@@ -573,10 +475,10 @@ ENV_LLAMACPP_EXTRA_ARGS = "HINDSIGHT_API_LLAMACPP_EXTRA_ARGS"
|
||||
|
||||
# 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"
|
||||
ENV_MIGRATION_CONCURRENCY = "HINDSIGHT_API_MIGRATION_CONCURRENCY"
|
||||
|
||||
# Database connection pool
|
||||
ENV_DB_POOL_MIN_SIZE = "HINDSIGHT_API_DB_POOL_MIN_SIZE"
|
||||
@@ -648,14 +550,6 @@ ENV_RECALL_MAX_CANDIDATES_PER_SOURCE = "HINDSIGHT_API_RECALL_MAX_CANDIDATES_PER_
|
||||
# Empty disables the feature.
|
||||
ENV_RECALL_STRATEGY_BOOSTS = "HINDSIGHT_API_RECALL_STRATEGY_BOOSTS"
|
||||
|
||||
# Recency decay used by recall reranking (engine/search/reranking.py). The decay
|
||||
# function maps a memory's age onto a freshness signal that nudges its final
|
||||
# ranking via a small multiplicative boost. "linear" (default) preserves the
|
||||
# historical behaviour; "exponential" decays by half-life; "none" disables it.
|
||||
ENV_RECENCY_DECAY_FUNCTION = "HINDSIGHT_API_RECENCY_DECAY_FUNCTION"
|
||||
ENV_RECENCY_DECAY_LINEAR_WINDOW_DAYS = "HINDSIGHT_API_RECENCY_DECAY_LINEAR_WINDOW_DAYS"
|
||||
ENV_RECENCY_DECAY_HALFLIFE_DAYS = "HINDSIGHT_API_RECENCY_DECAY_HALFLIFE_DAYS"
|
||||
|
||||
# Audit log settings
|
||||
ENV_AUDIT_LOG_ENABLED = "HINDSIGHT_API_AUDIT_LOG_ENABLED"
|
||||
ENV_AUDIT_LOG_ACTIONS = "HINDSIGHT_API_AUDIT_LOG_ACTIONS"
|
||||
@@ -669,7 +563,6 @@ ENV_LLM_TRACE_MAX_CHARS = "HINDSIGHT_API_LLM_TRACE_MAX_CHARS"
|
||||
|
||||
# Background maintenance settings
|
||||
ENV_CONSOLIDATION_RECONCILE_INTERVAL_SECONDS = "HINDSIGHT_API_CONSOLIDATION_RECONCILE_INTERVAL_SECONDS"
|
||||
ENV_MENTAL_MODEL_REFRESH_TICK_SECONDS = "HINDSIGHT_API_MENTAL_MODEL_REFRESH_TICK_SECONDS"
|
||||
|
||||
# Disposition settings
|
||||
ENV_DISPOSITION_SKEPTICISM = "HINDSIGHT_API_DISPOSITION_SKEPTICISM"
|
||||
@@ -692,7 +585,6 @@ PROVIDER_DEFAULT_MODELS = {
|
||||
"deepseek": "deepseek-v4-flash",
|
||||
"zai": "glm-4.5-flash",
|
||||
"opencode-go": "deepseek-v4-flash",
|
||||
"atlas": "deepseek-ai/deepseek-v4-pro",
|
||||
"ollama": "gemma3:12b",
|
||||
"ollama-cloud": "gemma3:12b",
|
||||
"llamacpp": "gemma-4-e2b-it",
|
||||
@@ -706,7 +598,6 @@ PROVIDER_DEFAULT_MODELS = {
|
||||
"bedrock": "us.amazon.nova-2-lite-v1:0",
|
||||
"volcano": "doubao-pro-32k",
|
||||
"openrouter": "qwen/qwen3.5-9b",
|
||||
"requesty": "openai/gpt-4o-mini",
|
||||
"fireworks": "accounts/fireworks/models/llama-v3p1-8b-instruct",
|
||||
"nous": "deepseek/deepseek-v4-flash",
|
||||
}
|
||||
@@ -797,14 +688,6 @@ DEFAULT_RECALL_MAX_CANDIDATES_PER_SOURCE = 0
|
||||
# "graph:high,semantic:low"). Empty disables the feature. See
|
||||
# ENV_RECALL_STRATEGY_BOOSTS for the full rationale.
|
||||
DEFAULT_RECALL_STRATEGY_BOOSTS = ""
|
||||
# Recency decay shape used by recall reranking. "linear" reproduces the
|
||||
# historical straight-line decay; defaults below keep behaviour unchanged.
|
||||
RECENCY_DECAY_FUNCTIONS = ("linear", "exponential", "none")
|
||||
DEFAULT_RECENCY_DECAY_FUNCTION = "linear"
|
||||
# Linear: days over which freshness decays from 1.0 to its 0.1 floor.
|
||||
DEFAULT_RECENCY_DECAY_LINEAR_WINDOW_DAYS = 365.0
|
||||
# Exponential: age (days) at which the recency signal is neutral (0.5).
|
||||
DEFAULT_RECENCY_DECAY_HALFLIFE_DAYS = 90.0
|
||||
# Retrieval arms that can be boosted; mirrors fusion.py source_names.
|
||||
RECALL_STRATEGY_NAMES = ("semantic", "bm25", "graph", "temporal")
|
||||
# User-facing priority levels. Kept in sync with recall_boost.BOOST_LEVELS by a
|
||||
@@ -863,9 +746,6 @@ DEFAULT_EMBEDDINGS_OPENROUTER_MODEL = "perplexity/pplx-embed-v1-0.6b"
|
||||
DEFAULT_RERANKER_OPENROUTER_MODEL = "cohere/rerank-v3.5"
|
||||
DEFAULT_RERANKER_OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1/rerank"
|
||||
|
||||
# Requesty defaults
|
||||
DEFAULT_EMBEDDINGS_REQUESTY_MODEL = "openai/text-embedding-3-small"
|
||||
|
||||
# ZeroEntropy defaults
|
||||
DEFAULT_EMBEDDINGS_ZEROENTROPY_MODEL = "zembed-1"
|
||||
# Shared between embeddings (zembed-1) and reranker (zerank-*) — the host is the same.
|
||||
@@ -921,12 +801,7 @@ DEFAULT_ACCESS_LOG = False
|
||||
DEFAULT_MCP_ENABLED = True
|
||||
DEFAULT_MCP_ENABLED_TOOLS: list[str] | None = None # None = all tools enabled
|
||||
DEFAULT_MCP_STATELESS = False # False = stateful (supports SSE/GET); True = stateless (POST-only)
|
||||
DEFAULT_MCP_INSTRUCTIONS = None
|
||||
DEFAULT_ENABLE_BANK_CONFIG_API = True
|
||||
# Dry-run extraction is a preview tool that makes a real LLM call but stores nothing. Enabled by
|
||||
# default; set HINDSIGHT_API_ENABLE_DRY_RUN_EXTRACT=false to remove the endpoint (e.g. to cap
|
||||
# provider cost/abuse on untrusted deployments).
|
||||
DEFAULT_ENABLE_DRY_RUN_EXTRACT = True
|
||||
# The per-bank LLM connectivity probe makes a real provider call, so it's OFF by
|
||||
# default (cost/abuse concerns) and must be explicitly enabled to expose the endpoint.
|
||||
DEFAULT_ENABLE_BANK_LLM_HEALTH = False
|
||||
@@ -965,10 +840,6 @@ DEFAULT_RETAIN_BATCH_POLL_INTERVAL_SECONDS = 60 # Batch API polling interval in
|
||||
DEFAULT_FILE_STORAGE_TYPE = "native" # PostgreSQL BYTEA storage
|
||||
DEFAULT_FILE_PARSER = "markitdown" # Default parser fallback chain (comma-separated, e.g. "iris,markitdown")
|
||||
DEFAULT_FILE_PARSER_ALLOWLIST = None # Allowlist of parsers clients may request (None = all registered parsers)
|
||||
DEFAULT_FILE_PARSER_MARKITDOWN_OCR_ENABLED = False
|
||||
DEFAULT_FILE_PARSER_MARKITDOWN_OCR_PROMPT = """You are a precise OCR transcription engine.
|
||||
|
||||
Transcribe only the visible text in the image. Do not describe the image, summarize it, translate it, infer missing content, or add commentary. Preserve the original language, wording, numbers, punctuation, capitalization, and reading order. Reconstruct headings, lists, key-value fields, stamps, and tables as clean Markdown when the layout is clear. If text is unreadable or uncertain, write [unclear] for that span. Return only the extracted Markdown."""
|
||||
DEFAULT_FILE_CONVERSION_MAX_BATCH_SIZE_MB = 100 # Max total batch size in MB (all files combined)
|
||||
DEFAULT_FILE_CONVERSION_MAX_BATCH_SIZE = 10 # Max files per batch upload
|
||||
DEFAULT_ENABLE_FILE_UPLOAD_API = True # Enable file upload endpoint
|
||||
@@ -1029,10 +900,6 @@ DEFAULT_OBSERVATION_SCOPE_LIMITS: list | None = None
|
||||
|
||||
# Database migrations
|
||||
DEFAULT_RUN_MIGRATIONS_ON_STARTUP = True
|
||||
# Number of tenant schemas to migrate concurrently. Each schema runs in its own
|
||||
# process (Alembic's command.upgrade() is not thread-safe); within a schema the
|
||||
# work is always sequential. 1 = fully sequential (the safe default).
|
||||
DEFAULT_MIGRATION_CONCURRENCY = 1
|
||||
|
||||
# Database connection pool
|
||||
DEFAULT_DB_POOL_MIN_SIZE = 5
|
||||
@@ -1087,7 +954,6 @@ DEFAULT_OTEL_TRACES_ENABLED = False # Disabled by default for backward compatib
|
||||
DEFAULT_OTEL_SERVICE_NAME = "hindsight-api"
|
||||
DEFAULT_OTEL_DEPLOYMENT_ENVIRONMENT = "development"
|
||||
DEFAULT_METRICS_INCLUDE_BANK_ID = False # Disabled by default to avoid high-cardinality OTel metric growth
|
||||
DEFAULT_METRICS_BACKLOG_ENABLED = False # Disabled by default: runs periodic per-schema COUNT queries
|
||||
|
||||
# Audit log defaults
|
||||
DEFAULT_AUDIT_LOG_ENABLED = False # Disabled by default
|
||||
@@ -1106,11 +972,6 @@ DEFAULT_LLM_TRACE_MAX_CHARS = 50000 # Truncate stored input/output beyond this
|
||||
# 0 disables the reconcile sweep.
|
||||
DEFAULT_CONSOLIDATION_RECONCILE_INTERVAL_SECONDS = 300
|
||||
|
||||
# How often the maintenance loop checks for cron-scheduled mental models that are
|
||||
# due for a refresh. This is the *check* cadence; the actual schedule is the
|
||||
# per-model cron expression in the mental model's trigger. 0 disables the sweep.
|
||||
DEFAULT_MENTAL_MODEL_REFRESH_TICK_SECONDS = 60
|
||||
|
||||
# Default MCP tool descriptions (can be customized via env vars)
|
||||
DEFAULT_MCP_RETAIN_DESCRIPTION = """Store important information to long-term memory.
|
||||
|
||||
@@ -1216,63 +1077,6 @@ def _parse_optional_positive_int(name: str, raw: str | None) -> int | None:
|
||||
return _parse_positive_int(name, raw, 1)
|
||||
|
||||
|
||||
def _validate_retain_chunking_int(name: str, value: Any) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError(f"{name} must be an integer, got {value!r}")
|
||||
if value < 1:
|
||||
raise ValueError(f"{name} must be >= 1, got {value}")
|
||||
return value
|
||||
|
||||
|
||||
def validate_retain_chunking_config(
|
||||
retain_chunk_size: Any,
|
||||
retain_structured_chunk_size: Any,
|
||||
*,
|
||||
retain_chunk_size_name: str = "retain_chunk_size",
|
||||
retain_structured_chunk_size_name: str = "retain_structured_chunk_size",
|
||||
) -> None:
|
||||
"""Validate retain chunking size fields.
|
||||
|
||||
Defaults emit field-style names ("retain_chunk_size") so API/PATCH callers
|
||||
don't have to override them. The startup validator (HindsightConfig.validate)
|
||||
overrides to env-style names ("HINDSIGHT_API_RETAIN_CHUNK_SIZE") for env
|
||||
misconfig errors.
|
||||
"""
|
||||
_validate_retain_chunking_int(retain_chunk_size_name, retain_chunk_size)
|
||||
if retain_structured_chunk_size is None:
|
||||
return
|
||||
_validate_retain_chunking_int(
|
||||
retain_structured_chunk_size_name,
|
||||
retain_structured_chunk_size,
|
||||
)
|
||||
|
||||
|
||||
def validate_retain_completion_token_budget(
|
||||
*,
|
||||
llm_provider: str,
|
||||
retain_max_completion_tokens: int,
|
||||
retain_chunk_size: int,
|
||||
retain_llm_model: str | None = None,
|
||||
llm_model: str | None = None,
|
||||
retain_llm_provider: str | None = None,
|
||||
retain_max_completion_tokens_name: str = "retain_max_completion_tokens",
|
||||
retain_chunk_size_name: str = "retain_chunk_size",
|
||||
) -> None:
|
||||
"""Validate that retain LLM output capacity exceeds the configured chunk size."""
|
||||
if llm_provider == "none" or retain_max_completion_tokens > retain_chunk_size:
|
||||
return
|
||||
raise ValueError(
|
||||
f"Invalid configuration: {retain_max_completion_tokens_name} "
|
||||
f"({retain_max_completion_tokens}) must be greater than "
|
||||
f"{retain_chunk_size_name} ({retain_chunk_size}). "
|
||||
f"\n\nYou have two options to fix this:"
|
||||
f"\n 1. Increase {retain_max_completion_tokens_name} to a value > {retain_chunk_size}"
|
||||
f"\n 2. Use a model that supports at least {retain_max_completion_tokens} output tokens"
|
||||
f"\n (current model: {retain_llm_model or llm_model}, "
|
||||
f"provider: {retain_llm_provider or llm_provider})"
|
||||
)
|
||||
|
||||
|
||||
def _parse_optional_choice(name: str, raw: str | None, allowed: frozenset[str]) -> str | None:
|
||||
"""Parse an optional string env var constrained to a small allowlist."""
|
||||
if raw is None or raw == "":
|
||||
@@ -1308,18 +1112,6 @@ def _validate_recall_budget_function(function: str) -> str:
|
||||
return function_lower
|
||||
|
||||
|
||||
def _validate_recency_decay_function(function: str) -> str:
|
||||
"""Validate and normalize the recency decay function."""
|
||||
function_lower = function.lower()
|
||||
if function_lower not in RECENCY_DECAY_FUNCTIONS:
|
||||
logger.warning(
|
||||
f"Invalid recency decay function '{function}', must be one of {RECENCY_DECAY_FUNCTIONS}. "
|
||||
f"Defaulting to '{DEFAULT_RECENCY_DECAY_FUNCTION}'."
|
||||
)
|
||||
return DEFAULT_RECENCY_DECAY_FUNCTION
|
||||
return function_lower
|
||||
|
||||
|
||||
def _parse_bank_priority(raw: str) -> dict[str, int]:
|
||||
"""Parse ``bank-pattern:priority,...`` into ``{pattern: priority}``.
|
||||
|
||||
@@ -1375,132 +1167,6 @@ def _parse_llm_router_config(env_var: str) -> dict | None:
|
||||
raise ValueError(f"Invalid {env_var}: invalid JSON: {e}") from e
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMMemberConfig:
|
||||
"""One extra LLM in a multi-LLM chain, configured via indexed env vars.
|
||||
|
||||
Mirrors the subset of LLM settings an indexed member supports
|
||||
(``HINDSIGHT_API_<OP>LLM_<n>_*``). The unindexed config remains the primary
|
||||
member (index 0); these describe members 1..N.
|
||||
"""
|
||||
|
||||
provider: str
|
||||
api_key: str | None
|
||||
model: str
|
||||
base_url: str | None
|
||||
reasoning_effort: str | None
|
||||
extra_body: dict | None
|
||||
default_headers: dict | None
|
||||
bedrock_service_tier: str | None
|
||||
gemini_service_tier: str | None
|
||||
vertexai_project_id: str | None = None
|
||||
vertexai_region: str | None = None
|
||||
vertexai_service_account_key: str | None = None
|
||||
litellmrouter_config: dict | None = None
|
||||
|
||||
|
||||
# Valid multi-LLM strategy modes.
|
||||
LLM_STRATEGY_FAILOVER = "failover"
|
||||
LLM_STRATEGY_ROUND_ROBIN = "round-robin"
|
||||
_VALID_LLM_STRATEGY_MODES = (LLM_STRATEGY_FAILOVER, LLM_STRATEGY_ROUND_ROBIN)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMStrategyConfig:
|
||||
"""How to route a request across the members of a multi-LLM chain.
|
||||
|
||||
``mode`` is "failover" (try members in order) or "round-robin" (rotate the
|
||||
starting member per request, then fall through the rest on error). ``weights``
|
||||
is round-robin only: positive integers, one per member (primary first), giving
|
||||
an unbalanced rotation; ``None`` means uniform.
|
||||
"""
|
||||
|
||||
mode: str
|
||||
weights: list[int] | None = None
|
||||
|
||||
|
||||
def _parse_llm_strategy(raw: str | None) -> LLMStrategyConfig | None:
|
||||
"""Parse a multi-LLM strategy from a JSON env var.
|
||||
|
||||
Returns ``None`` when unset. The value must be a JSON object with a ``mode``
|
||||
of "failover" or "round-robin"; ``weights`` (round-robin only) must be a list
|
||||
of positive ints. Raises ``ValueError`` on any malformed input so
|
||||
misconfiguration fails fast at startup rather than silently degrading.
|
||||
"""
|
||||
text = (raw or "").strip()
|
||||
if not text:
|
||||
return None
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid {ENV_LLM_STRATEGY}: invalid JSON: {e}") from e
|
||||
if not isinstance(parsed, dict):
|
||||
raise ValueError(f"Invalid LLM strategy: expected a JSON object, got {type(parsed).__name__}")
|
||||
|
||||
mode = parsed.get("mode")
|
||||
if mode not in _VALID_LLM_STRATEGY_MODES:
|
||||
raise ValueError(f"Invalid LLM strategy mode {mode!r}. Must be one of: {', '.join(_VALID_LLM_STRATEGY_MODES)}.")
|
||||
|
||||
weights = parsed.get("weights")
|
||||
if weights is not None:
|
||||
if mode != LLM_STRATEGY_ROUND_ROBIN:
|
||||
raise ValueError(f"LLM strategy 'weights' is only valid with mode '{LLM_STRATEGY_ROUND_ROBIN}'.")
|
||||
if not isinstance(weights, list) or not weights or not all(isinstance(w, int) and w > 0 for w in weights):
|
||||
raise ValueError("LLM strategy 'weights' must be a non-empty list of positive integers.")
|
||||
|
||||
return LLMStrategyConfig(mode=mode, weights=weights)
|
||||
|
||||
|
||||
def _parse_llm_members(prefix: str) -> list[LLMMemberConfig]:
|
||||
"""Parse indexed extra-LLM members for an operation env prefix.
|
||||
|
||||
``prefix`` is the operation segment in the env name: ``""`` (global),
|
||||
``"RETAIN_"``, ``"REFLECT_"`` or ``"CONSOLIDATION_"``. Members are read from
|
||||
``HINDSIGHT_API_{prefix}LLM_{n}_PROVIDER`` for n = 1, 2, ... and scanning
|
||||
stops at the first index whose ``_PROVIDER`` is unset (so indices must be
|
||||
contiguous from 1). ``MODEL`` defaults to the provider's default model.
|
||||
"""
|
||||
from .engine.llm_wrapper import requires_api_key
|
||||
|
||||
members: list[LLMMemberConfig] = []
|
||||
index = 1
|
||||
while True:
|
||||
base = f"HINDSIGHT_API_{prefix}LLM_{index}_"
|
||||
provider = os.getenv(base + "PROVIDER")
|
||||
if not provider:
|
||||
break
|
||||
|
||||
api_key = os.getenv(base + "API_KEY") or None
|
||||
if not api_key and requires_api_key(provider):
|
||||
raise ValueError(
|
||||
f"{base}API_KEY is required for provider '{provider}' (member {index} of the multi-LLM chain)."
|
||||
)
|
||||
|
||||
gemini_service_tier = os.getenv(base + "GEMINI_SERVICE_TIER")
|
||||
members.append(
|
||||
LLMMemberConfig(
|
||||
provider=provider,
|
||||
api_key=api_key,
|
||||
model=os.getenv(base + "MODEL") or _get_default_model_for_provider(provider),
|
||||
base_url=os.getenv(base + "BASE_URL") or None,
|
||||
reasoning_effort=os.getenv(base + "REASONING_EFFORT") or None,
|
||||
extra_body=json.loads(os.getenv(base + "EXTRA_BODY", "null")),
|
||||
default_headers=json.loads(os.getenv(base + "DEFAULT_HEADERS", "null")),
|
||||
bedrock_service_tier=os.getenv(base + "BEDROCK_SERVICE_TIER") or None,
|
||||
gemini_service_tier=(
|
||||
parse_gemini_service_tier(gemini_service_tier) if provider.lower() == "gemini" else None
|
||||
),
|
||||
vertexai_project_id=os.getenv(base + "VERTEXAI_PROJECT_ID") or None,
|
||||
vertexai_region=os.getenv(base + "VERTEXAI_REGION") or None,
|
||||
vertexai_service_account_key=os.getenv(base + "VERTEXAI_SERVICE_ACCOUNT_KEY") or None,
|
||||
litellmrouter_config=_parse_llm_router_config(base + "LITELLMROUTER_CONFIG"),
|
||||
)
|
||||
)
|
||||
index += 1
|
||||
|
||||
return members
|
||||
|
||||
|
||||
def _parse_default_bank_template(raw: str | None) -> dict | None:
|
||||
"""
|
||||
Parse HINDSIGHT_API_DEFAULT_BANK_TEMPLATE as JSON.
|
||||
@@ -1564,7 +1230,6 @@ class HindsightConfig:
|
||||
llm_groq_service_tier: str # Groq: "on_demand", "flex", or "auto"
|
||||
llm_openai_service_tier: str | None # OpenAI: None (default) or "flex" (50% cheaper)
|
||||
llm_bedrock_service_tier: str | None # Bedrock: None (default), "flex", "priority", or "reserved"
|
||||
llm_gemini_service_tier: str | None # Gemini: None (default) or "flex" (50% cheaper)
|
||||
llm_extra_body: (
|
||||
dict | None
|
||||
) # Extra body params merged into OpenAI-compatible API calls (e.g. {"chat_template_kwargs": {"enable_thinking": true}})
|
||||
@@ -1578,14 +1243,6 @@ class HindsightConfig:
|
||||
# overrides a `user` the caller already set.
|
||||
llm_send_bank_as_user: bool
|
||||
|
||||
# Per-operation sampling temperature. None means the temperature parameter is
|
||||
# omitted from the call (for models that reject explicit temperatures). See
|
||||
# ENV_LLM_TEMPERATURE and _resolve_operation_temperature.
|
||||
llm_temperature_verification: float | None
|
||||
llm_temperature_retain: float | None
|
||||
llm_temperature_reflect: float | None
|
||||
llm_temperature_consolidation: float | None
|
||||
|
||||
# LiteLLM Router chain (provider-specific; consumed by the "litellmrouter" provider).
|
||||
# List of deployment dicts evaluated in order with fallback on transient errors.
|
||||
# Each entry: {"provider": str, "model": str, "api_key": str | None, "base_url": str | None}.
|
||||
@@ -1675,8 +1332,6 @@ class HindsightConfig:
|
||||
embeddings_cohere_output_dimensions: int | None
|
||||
embeddings_openrouter_api_key: str | None
|
||||
embeddings_openrouter_model: str
|
||||
embeddings_requesty_api_key: str | None
|
||||
embeddings_requesty_model: str
|
||||
embeddings_litellm_api_base: str
|
||||
embeddings_litellm_api_key: str | None
|
||||
embeddings_litellm_model: str
|
||||
@@ -1712,9 +1367,6 @@ class HindsightConfig:
|
||||
bm25_min_score: float
|
||||
recall_max_candidates_per_source: int
|
||||
recall_strategy_boosts: dict[str, str]
|
||||
recency_decay_function: str
|
||||
recency_decay_linear_window_days: float
|
||||
recency_decay_halflife_days: float
|
||||
reranker_cohere_api_key: str | None
|
||||
reranker_cohere_model: str
|
||||
reranker_cohere_base_url: str | None
|
||||
@@ -1758,10 +1410,8 @@ class HindsightConfig:
|
||||
mcp_enabled: bool
|
||||
mcp_enabled_tools: list[str] | None # None = all tools; explicit list = allowlist
|
||||
mcp_stateless: bool # True = stateless HTTP (POST-only); False = stateful (supports GET/SSE)
|
||||
mcp_instructions: str | None # Additional instructions appended to retain/recall MCP tool descriptions
|
||||
enable_bank_config_api: bool
|
||||
enable_bank_llm_health: bool
|
||||
enable_dry_run_extract: bool
|
||||
# Default bank template (static, server-level only). When set, the manifest is applied
|
||||
# to every newly-created bank, overriding the env/config defaults for any fields it sets.
|
||||
default_bank_template: dict | None
|
||||
@@ -1780,7 +1430,6 @@ class HindsightConfig:
|
||||
# Retain settings
|
||||
retain_max_completion_tokens: int
|
||||
retain_chunk_size: int
|
||||
retain_structured_chunk_size: int | None
|
||||
retain_extract_causal_links: bool
|
||||
retain_extraction_mode: str
|
||||
retain_mission: str | None
|
||||
@@ -1885,10 +1534,10 @@ class HindsightConfig:
|
||||
|
||||
# Optimization flags
|
||||
skip_llm_verification: bool
|
||||
lazy_reranker: bool
|
||||
|
||||
# Database migrations
|
||||
run_migrations_on_startup: bool
|
||||
migration_concurrency: int
|
||||
|
||||
# Database connection pool
|
||||
db_pool_min_size: int
|
||||
@@ -1922,7 +1571,6 @@ class HindsightConfig:
|
||||
otel_service_name: str
|
||||
otel_deployment_environment: str
|
||||
metrics_include_bank_id: bool
|
||||
metrics_backlog_enabled: bool
|
||||
|
||||
# Audit log configuration (static - server-level only)
|
||||
audit_log_enabled: bool # Master switch for audit logging
|
||||
@@ -1939,9 +1587,6 @@ class HindsightConfig:
|
||||
# Interval for the periodic sweep that re-schedules consolidation for banks with
|
||||
# eligible-but-unscheduled facts. 0 = disabled.
|
||||
consolidation_reconcile_interval_seconds: int
|
||||
# How often the maintenance loop checks for cron-scheduled mental models due for
|
||||
# refresh (the per-model schedule lives in the mental model trigger). 0 = disabled.
|
||||
mental_model_refresh_tick_seconds: int
|
||||
|
||||
# Webhook configuration (static - server-level only, not per-bank)
|
||||
webhook_url: str | None # Global webhook URL (None = disabled)
|
||||
@@ -1960,25 +1605,6 @@ class HindsightConfig:
|
||||
embeddings_zeroentropy_encoding_format: str = DEFAULT_EMBEDDINGS_ZEROENTROPY_ENCODING_FORMAT
|
||||
embeddings_zeroentropy_batch_size: int = DEFAULT_EMBEDDINGS_ZEROENTROPY_BATCH_SIZE
|
||||
embeddings_zeroentropy_latency: str | None = DEFAULT_EMBEDDINGS_ZEROENTROPY_LATENCY
|
||||
file_parser_markitdown_ocr_enabled: bool = DEFAULT_FILE_PARSER_MARKITDOWN_OCR_ENABLED
|
||||
file_parser_markitdown_ocr_api_key: str | None = None
|
||||
file_parser_markitdown_ocr_base_url: str | None = None
|
||||
file_parser_markitdown_ocr_model: str | None = None
|
||||
file_parser_markitdown_ocr_prompt: str = DEFAULT_FILE_PARSER_MARKITDOWN_OCR_PROMPT
|
||||
|
||||
# Multi-LLM chains (static, server-level). Index 0 of each chain is the
|
||||
# corresponding unindexed/base LLM config above; these hold the extra indexed
|
||||
# members and the routing strategy. Per-op members fall back to the global
|
||||
# members when unset (see MemoryEngine._build_llm). Credential fields (members
|
||||
# embed api_keys/base_urls).
|
||||
llm_members: list[LLMMemberConfig] = field(default_factory=list)
|
||||
llm_strategy: LLMStrategyConfig | None = None
|
||||
retain_llm_members: list[LLMMemberConfig] = field(default_factory=list)
|
||||
retain_llm_strategy: LLMStrategyConfig | None = None
|
||||
reflect_llm_members: list[LLMMemberConfig] = field(default_factory=list)
|
||||
reflect_llm_strategy: LLMStrategyConfig | None = None
|
||||
consolidation_llm_members: list[LLMMemberConfig] = field(default_factory=list)
|
||||
consolidation_llm_strategy: LLMStrategyConfig | None = None
|
||||
|
||||
# Class-level sets for configuration categorization
|
||||
|
||||
@@ -1994,11 +1620,6 @@ class HindsightConfig:
|
||||
"retain_llm_litellmrouter_config",
|
||||
"reflect_llm_litellmrouter_config",
|
||||
"consolidation_llm_litellmrouter_config",
|
||||
# Multi-LLM chains — members embed api_keys and base_urls
|
||||
"llm_members",
|
||||
"retain_llm_members",
|
||||
"reflect_llm_members",
|
||||
"consolidation_llm_members",
|
||||
# Base URLs (could expose infrastructure)
|
||||
"llm_base_url",
|
||||
"retain_llm_base_url",
|
||||
@@ -2024,8 +1645,6 @@ class HindsightConfig:
|
||||
"file_storage_gcs_service_account_key",
|
||||
"file_storage_azure_account_key",
|
||||
# File parser credentials
|
||||
"file_parser_markitdown_ocr_api_key",
|
||||
"file_parser_markitdown_ocr_base_url",
|
||||
"file_parser_iris_token",
|
||||
"file_parser_llama_parse_api_key",
|
||||
}
|
||||
@@ -2038,7 +1657,6 @@ class HindsightConfig:
|
||||
"mcp_enabled_tools",
|
||||
# Retention settings (behavioral)
|
||||
"retain_chunk_size",
|
||||
"retain_structured_chunk_size",
|
||||
"retain_extraction_mode",
|
||||
"retain_mission",
|
||||
"retain_custom_instructions",
|
||||
@@ -2189,9 +1807,6 @@ class HindsightConfig:
|
||||
f"Note: 'standard' is not a valid Bedrock service tier -- use unset for default tier."
|
||||
)
|
||||
|
||||
# Validate gemini_service_tier
|
||||
self.llm_gemini_service_tier = parse_gemini_service_tier(self.llm_gemini_service_tier)
|
||||
|
||||
# When LLM provider is "none", force chunks-only mode and disable LLM-dependent features
|
||||
if self.llm_provider == "none":
|
||||
self.retain_extraction_mode = "chunks"
|
||||
@@ -2201,23 +1816,20 @@ class HindsightConfig:
|
||||
"disabling observations/consolidation. Reflect will return HTTP 400."
|
||||
)
|
||||
|
||||
validate_retain_chunking_config(
|
||||
self.retain_chunk_size,
|
||||
self.retain_structured_chunk_size,
|
||||
retain_chunk_size_name="HINDSIGHT_API_RETAIN_CHUNK_SIZE",
|
||||
retain_structured_chunk_size_name="HINDSIGHT_API_RETAIN_STRUCTURED_CHUNK_SIZE",
|
||||
)
|
||||
|
||||
validate_retain_completion_token_budget(
|
||||
llm_provider=self.llm_provider,
|
||||
retain_max_completion_tokens=self.retain_max_completion_tokens,
|
||||
retain_chunk_size=self.retain_chunk_size,
|
||||
retain_llm_model=self.retain_llm_model,
|
||||
llm_model=self.llm_model,
|
||||
retain_llm_provider=self.retain_llm_provider,
|
||||
retain_max_completion_tokens_name="HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS",
|
||||
retain_chunk_size_name="HINDSIGHT_API_RETAIN_CHUNK_SIZE",
|
||||
)
|
||||
# RETAIN_MAX_COMPLETION_TOKENS must be greater than RETAIN_CHUNK_SIZE
|
||||
# to ensure the LLM has enough output capacity to extract facts from chunks
|
||||
# (not applicable when provider is "none" since no LLM calls are made)
|
||||
if self.llm_provider != "none" and self.retain_max_completion_tokens <= self.retain_chunk_size:
|
||||
raise ValueError(
|
||||
f"Invalid configuration: HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS "
|
||||
f"({self.retain_max_completion_tokens}) must be greater than "
|
||||
f"HINDSIGHT_API_RETAIN_CHUNK_SIZE ({self.retain_chunk_size}). "
|
||||
f"\n\nYou have two options to fix this:"
|
||||
f"\n 1. Increase HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS to a value > {self.retain_chunk_size}"
|
||||
f"\n 2. Use a model that supports at least {self.retain_max_completion_tokens} output tokens"
|
||||
f"\n (current model: {self.retain_llm_model or self.llm_model}, "
|
||||
f"provider: {self.retain_llm_provider or self.llm_provider})"
|
||||
)
|
||||
|
||||
# Warn if local ML dependencies are missing when configured.
|
||||
# Don't hard-fail here — the actual ImportError fires at model init time
|
||||
@@ -2309,28 +1921,11 @@ class HindsightConfig:
|
||||
llm_groq_service_tier=os.getenv(ENV_LLM_GROQ_SERVICE_TIER, DEFAULT_LLM_GROQ_SERVICE_TIER),
|
||||
llm_openai_service_tier=os.getenv(ENV_LLM_OPENAI_SERVICE_TIER, DEFAULT_LLM_OPENAI_SERVICE_TIER),
|
||||
llm_bedrock_service_tier=os.getenv(ENV_LLM_BEDROCK_SERVICE_TIER) or None,
|
||||
llm_gemini_service_tier=(
|
||||
parse_gemini_service_tier(os.getenv(ENV_LLM_GEMINI_SERVICE_TIER) or DEFAULT_LLM_GEMINI_SERVICE_TIER)
|
||||
if llm_provider.lower() == "gemini"
|
||||
else None
|
||||
),
|
||||
llm_extra_body=json.loads(os.getenv(ENV_LLM_EXTRA_BODY, "null")),
|
||||
llm_default_headers=json.loads(os.getenv(ENV_LLM_DEFAULT_HEADERS, "null")),
|
||||
llm_strict_schema=os.getenv(ENV_LLM_STRICT_SCHEMA, str(DEFAULT_LLM_STRICT_SCHEMA)).lower() in ("true", "1"),
|
||||
llm_send_bank_as_user=os.getenv(ENV_LLM_SEND_BANK_AS_USER, str(DEFAULT_LLM_SEND_BANK_AS_USER)).lower()
|
||||
in ("true", "1"),
|
||||
llm_temperature_verification=_resolve_operation_temperature(
|
||||
ENV_LLM_TEMPERATURE_VERIFICATION, DEFAULT_LLM_TEMPERATURE_VERIFICATION
|
||||
),
|
||||
llm_temperature_retain=_resolve_operation_temperature(
|
||||
ENV_LLM_TEMPERATURE_RETAIN, DEFAULT_LLM_TEMPERATURE_RETAIN
|
||||
),
|
||||
llm_temperature_reflect=_resolve_operation_temperature(
|
||||
ENV_LLM_TEMPERATURE_REFLECT, DEFAULT_LLM_TEMPERATURE_REFLECT
|
||||
),
|
||||
llm_temperature_consolidation=_resolve_operation_temperature(
|
||||
ENV_LLM_TEMPERATURE_CONSOLIDATION, DEFAULT_LLM_TEMPERATURE_CONSOLIDATION
|
||||
),
|
||||
llm_litellmrouter_config=_parse_llm_router_config(ENV_LLM_LITELLMROUTER_CONFIG),
|
||||
# Vertex AI
|
||||
llm_vertexai_project_id=os.getenv(ENV_LLM_VERTEXAI_PROJECT_ID) or DEFAULT_LLM_VERTEXAI_PROJECT_ID,
|
||||
@@ -2430,15 +2025,6 @@ class HindsightConfig:
|
||||
if os.getenv(ENV_CONSOLIDATION_LLM_TIMEOUT)
|
||||
else None,
|
||||
consolidation_llm_litellmrouter_config=_parse_llm_router_config(ENV_CONSOLIDATION_LLM_LITELLMROUTER_CONFIG),
|
||||
# Multi-LLM chains (indexed members + routing strategy)
|
||||
llm_members=_parse_llm_members(""),
|
||||
llm_strategy=_parse_llm_strategy(os.getenv(ENV_LLM_STRATEGY)),
|
||||
retain_llm_members=_parse_llm_members("RETAIN_"),
|
||||
retain_llm_strategy=_parse_llm_strategy(os.getenv(ENV_RETAIN_LLM_STRATEGY)),
|
||||
reflect_llm_members=_parse_llm_members("REFLECT_"),
|
||||
reflect_llm_strategy=_parse_llm_strategy(os.getenv(ENV_REFLECT_LLM_STRATEGY)),
|
||||
consolidation_llm_members=_parse_llm_members("CONSOLIDATION_"),
|
||||
consolidation_llm_strategy=_parse_llm_strategy(os.getenv(ENV_CONSOLIDATION_LLM_STRATEGY)),
|
||||
# Embeddings
|
||||
embeddings_provider=os.getenv(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER),
|
||||
embeddings_local_model=os.getenv(ENV_EMBEDDINGS_LOCAL_MODEL, DEFAULT_EMBEDDINGS_LOCAL_MODEL),
|
||||
@@ -2503,11 +2089,6 @@ class HindsightConfig:
|
||||
or os.getenv(ENV_OPENROUTER_API_KEY)
|
||||
or os.getenv(ENV_LLM_API_KEY),
|
||||
embeddings_openrouter_model=os.getenv(ENV_EMBEDDINGS_OPENROUTER_MODEL, DEFAULT_EMBEDDINGS_OPENROUTER_MODEL),
|
||||
# Requesty embeddings (with fallback to shared Requesty key, then LLM key)
|
||||
embeddings_requesty_api_key=os.getenv(ENV_EMBEDDINGS_REQUESTY_API_KEY)
|
||||
or os.getenv(ENV_REQUESTY_API_KEY)
|
||||
or os.getenv(ENV_LLM_API_KEY),
|
||||
embeddings_requesty_model=os.getenv(ENV_EMBEDDINGS_REQUESTY_MODEL, DEFAULT_EMBEDDINGS_REQUESTY_MODEL),
|
||||
# ZeroEntropy embeddings
|
||||
embeddings_zeroentropy_api_key=os.getenv(ENV_EMBEDDINGS_ZEROENTROPY_API_KEY)
|
||||
or os.getenv("ZEROENTROPY_API_KEY"),
|
||||
@@ -2614,15 +2195,6 @@ class HindsightConfig:
|
||||
recall_strategy_boosts=_parse_strategy_boosts(
|
||||
os.getenv(ENV_RECALL_STRATEGY_BOOSTS, DEFAULT_RECALL_STRATEGY_BOOSTS)
|
||||
),
|
||||
recency_decay_function=_validate_recency_decay_function(
|
||||
os.getenv(ENV_RECENCY_DECAY_FUNCTION, DEFAULT_RECENCY_DECAY_FUNCTION)
|
||||
),
|
||||
recency_decay_linear_window_days=float(
|
||||
os.getenv(ENV_RECENCY_DECAY_LINEAR_WINDOW_DAYS, str(DEFAULT_RECENCY_DECAY_LINEAR_WINDOW_DAYS))
|
||||
),
|
||||
recency_decay_halflife_days=float(
|
||||
os.getenv(ENV_RECENCY_DECAY_HALFLIFE_DAYS, str(DEFAULT_RECENCY_DECAY_HALFLIFE_DAYS))
|
||||
),
|
||||
# Cohere reranker (with backward-compatible fallback to shared API key)
|
||||
reranker_cohere_api_key=os.getenv(ENV_RERANKER_COHERE_API_KEY) or os.getenv(ENV_COHERE_API_KEY),
|
||||
reranker_cohere_model=os.getenv(ENV_RERANKER_COHERE_MODEL, DEFAULT_RERANKER_COHERE_MODEL),
|
||||
@@ -2698,13 +2270,10 @@ class HindsightConfig:
|
||||
if os.getenv(ENV_MCP_ENABLED_TOOLS)
|
||||
else DEFAULT_MCP_ENABLED_TOOLS,
|
||||
mcp_stateless=os.getenv(ENV_MCP_STATELESS, str(DEFAULT_MCP_STATELESS)).lower() == "true",
|
||||
mcp_instructions=os.getenv(ENV_MCP_INSTRUCTIONS) or DEFAULT_MCP_INSTRUCTIONS,
|
||||
enable_bank_llm_health=os.getenv(ENV_ENABLE_BANK_LLM_HEALTH, str(DEFAULT_ENABLE_BANK_LLM_HEALTH)).lower()
|
||||
== "true",
|
||||
enable_bank_config_api=os.getenv(ENV_ENABLE_BANK_CONFIG_API, str(DEFAULT_ENABLE_BANK_CONFIG_API)).lower()
|
||||
== "true",
|
||||
enable_dry_run_extract=os.getenv(ENV_ENABLE_DRY_RUN_EXTRACT, str(DEFAULT_ENABLE_DRY_RUN_EXTRACT)).lower()
|
||||
== "true",
|
||||
default_bank_template=_parse_default_bank_template(os.getenv(ENV_DEFAULT_BANK_TEMPLATE)),
|
||||
# Recall
|
||||
graph_retriever=os.getenv(ENV_GRAPH_RETRIEVER, DEFAULT_GRAPH_RETRIEVER),
|
||||
@@ -2728,15 +2297,12 @@ class HindsightConfig:
|
||||
),
|
||||
# Optimization flags
|
||||
skip_llm_verification=os.getenv(ENV_SKIP_LLM_VERIFICATION, "false").lower() == "true",
|
||||
lazy_reranker=os.getenv(ENV_LAZY_RERANKER, "false").lower() == "true",
|
||||
# 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_structured_chunk_size=_parse_optional_positive_int(
|
||||
ENV_RETAIN_STRUCTURED_CHUNK_SIZE,
|
||||
os.getenv(ENV_RETAIN_STRUCTURED_CHUNK_SIZE),
|
||||
),
|
||||
retain_extract_causal_links=os.getenv(
|
||||
ENV_RETAIN_EXTRACT_CAUSAL_LINKS, str(DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS)
|
||||
).lower()
|
||||
@@ -2777,18 +2343,6 @@ class HindsightConfig:
|
||||
file_parser_allowlist=_parse_str_list(os.getenv(ENV_FILE_PARSER_ALLOWLIST))
|
||||
if os.getenv(ENV_FILE_PARSER_ALLOWLIST)
|
||||
else None,
|
||||
file_parser_markitdown_ocr_enabled=os.getenv(
|
||||
ENV_FILE_PARSER_MARKITDOWN_OCR_ENABLED,
|
||||
str(DEFAULT_FILE_PARSER_MARKITDOWN_OCR_ENABLED),
|
||||
).lower()
|
||||
in ("1", "true", "yes", "on"),
|
||||
file_parser_markitdown_ocr_api_key=os.getenv(ENV_FILE_PARSER_MARKITDOWN_OCR_API_KEY) or None,
|
||||
file_parser_markitdown_ocr_base_url=os.getenv(ENV_FILE_PARSER_MARKITDOWN_OCR_BASE_URL) or None,
|
||||
file_parser_markitdown_ocr_model=os.getenv(ENV_FILE_PARSER_MARKITDOWN_OCR_MODEL) or None,
|
||||
file_parser_markitdown_ocr_prompt=os.getenv(
|
||||
ENV_FILE_PARSER_MARKITDOWN_OCR_PROMPT,
|
||||
DEFAULT_FILE_PARSER_MARKITDOWN_OCR_PROMPT,
|
||||
),
|
||||
file_parser_iris_token=os.getenv(ENV_FILE_PARSER_IRIS_TOKEN) or None,
|
||||
file_parser_iris_org_id=os.getenv(ENV_FILE_PARSER_IRIS_ORG_ID) or None,
|
||||
file_parser_llama_parse_api_key=os.getenv(ENV_FILE_PARSER_LLAMA_PARSE_API_KEY) or None,
|
||||
@@ -2895,7 +2449,6 @@ class HindsightConfig:
|
||||
memory_defense=None,
|
||||
# Database migrations
|
||||
run_migrations_on_startup=os.getenv(ENV_RUN_MIGRATIONS_ON_STARTUP, "true").lower() == "true",
|
||||
migration_concurrency=int(os.getenv(ENV_MIGRATION_CONCURRENCY, str(DEFAULT_MIGRATION_CONCURRENCY))),
|
||||
# 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))),
|
||||
@@ -2979,8 +2532,6 @@ class HindsightConfig:
|
||||
otel_deployment_environment=os.getenv(ENV_OTEL_DEPLOYMENT_ENVIRONMENT, DEFAULT_OTEL_DEPLOYMENT_ENVIRONMENT),
|
||||
metrics_include_bank_id=os.getenv(ENV_METRICS_INCLUDE_BANK_ID, str(DEFAULT_METRICS_INCLUDE_BANK_ID)).lower()
|
||||
in ("true", "1", "yes"),
|
||||
metrics_backlog_enabled=os.getenv(ENV_METRICS_BACKLOG_ENABLED, str(DEFAULT_METRICS_BACKLOG_ENABLED)).lower()
|
||||
in ("true", "1", "yes"),
|
||||
# Audit log configuration (static, server-level only)
|
||||
audit_log_enabled=os.getenv(ENV_AUDIT_LOG_ENABLED, str(DEFAULT_AUDIT_LOG_ENABLED)).lower() == "true",
|
||||
audit_log_actions=[
|
||||
@@ -3005,12 +2556,6 @@ class HindsightConfig:
|
||||
str(DEFAULT_CONSOLIDATION_RECONCILE_INTERVAL_SECONDS),
|
||||
)
|
||||
),
|
||||
mental_model_refresh_tick_seconds=int(
|
||||
os.getenv(
|
||||
ENV_MENTAL_MODEL_REFRESH_TICK_SECONDS,
|
||||
str(DEFAULT_MENTAL_MODEL_REFRESH_TICK_SECONDS),
|
||||
)
|
||||
),
|
||||
# Webhook configuration (static, server-level only)
|
||||
webhook_url=os.getenv(ENV_WEBHOOK_URL) or DEFAULT_WEBHOOK_URL,
|
||||
webhook_secret=os.getenv(ENV_WEBHOOK_SECRET) or DEFAULT_WEBHOOK_SECRET,
|
||||
|
||||
@@ -8,7 +8,6 @@ Config values are resolved on every request to ensure consistency across
|
||||
multiple API servers.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import asdict, replace
|
||||
@@ -19,8 +18,6 @@ from hindsight_api.config import (
|
||||
HindsightConfig,
|
||||
_get_raw_config,
|
||||
normalize_config_dict,
|
||||
validate_retain_chunking_config,
|
||||
validate_retain_completion_token_budget,
|
||||
)
|
||||
from hindsight_api.engine.memory_engine import fq_table
|
||||
from hindsight_api.extensions.tenant import TenantExtension
|
||||
@@ -32,35 +29,6 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _validate_retain_strategy_chunking(base_config: HindsightConfig, strategies: Any) -> None:
|
||||
"""Validate retain strategy chunking with the same semantics as apply_strategy()."""
|
||||
if not isinstance(strategies, dict):
|
||||
return
|
||||
configurable = HindsightConfig.get_configurable_fields()
|
||||
for strategy_name, overrides in strategies.items():
|
||||
if not isinstance(overrides, dict):
|
||||
raise ValueError(f"Invalid retain strategy {strategy_name!r}: must be an object")
|
||||
filtered = {k: v for k, v in overrides.items() if k in configurable}
|
||||
if not filtered:
|
||||
continue
|
||||
try:
|
||||
resolved = replace(base_config, **filtered)
|
||||
validate_retain_chunking_config(
|
||||
resolved.retain_chunk_size,
|
||||
resolved.retain_structured_chunk_size,
|
||||
)
|
||||
validate_retain_completion_token_budget(
|
||||
llm_provider=resolved.llm_provider,
|
||||
retain_max_completion_tokens=resolved.retain_max_completion_tokens,
|
||||
retain_chunk_size=resolved.retain_chunk_size,
|
||||
retain_llm_model=resolved.retain_llm_model,
|
||||
llm_model=resolved.llm_model,
|
||||
retain_llm_provider=resolved.retain_llm_provider,
|
||||
)
|
||||
except ValueError as e:
|
||||
raise ValueError(f"Invalid retain strategy {strategy_name!r}: {e}") from e
|
||||
|
||||
|
||||
class ConfigResolver:
|
||||
"""Resolves hierarchical configuration with tenant/bank overrides."""
|
||||
|
||||
@@ -78,26 +46,6 @@ class ConfigResolver:
|
||||
self._configurable_fields = HindsightConfig.get_configurable_fields()
|
||||
self._credential_fields = HindsightConfig.get_credential_fields()
|
||||
|
||||
async def _resolve_parent_config_dict(self, bank_id: str, context: RequestContext | None = None) -> dict[str, Any]:
|
||||
"""Resolve global + tenant config before bank-level overrides."""
|
||||
config_dict = asdict(self._global_config)
|
||||
|
||||
if self.tenant_extension and context:
|
||||
try:
|
||||
tenant_overrides = await self.tenant_extension.get_tenant_config(context)
|
||||
if tenant_overrides:
|
||||
# Normalize keys and filter to configurable fields only
|
||||
normalized_tenant = normalize_config_dict(tenant_overrides)
|
||||
configurable_tenant = {k: v for k, v in normalized_tenant.items() if k in self._configurable_fields}
|
||||
config_dict.update(configurable_tenant)
|
||||
logger.debug(
|
||||
f"Applied tenant config overrides for bank {bank_id}: {list(configurable_tenant.keys())}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load tenant config for bank {bank_id}: {e}")
|
||||
|
||||
return config_dict
|
||||
|
||||
async def resolve_full_config(self, bank_id: str, context: RequestContext | None = None) -> HindsightConfig:
|
||||
"""
|
||||
Resolve full HindsightConfig for a bank with hierarchical overrides applied.
|
||||
@@ -117,7 +65,23 @@ class ConfigResolver:
|
||||
Returns:
|
||||
Complete HindsightConfig with hierarchical overrides applied
|
||||
"""
|
||||
config_dict = await self._resolve_parent_config_dict(bank_id, context)
|
||||
# Start with global config (all fields)
|
||||
config_dict = asdict(self._global_config)
|
||||
|
||||
# Load tenant config overrides (if tenant extension available)
|
||||
if self.tenant_extension and context:
|
||||
try:
|
||||
tenant_overrides = await self.tenant_extension.get_tenant_config(context)
|
||||
if tenant_overrides:
|
||||
# Normalize keys and filter to configurable fields only
|
||||
normalized_tenant = normalize_config_dict(tenant_overrides)
|
||||
configurable_tenant = {k: v for k, v in normalized_tenant.items() if k in self._configurable_fields}
|
||||
config_dict.update(configurable_tenant)
|
||||
logger.debug(
|
||||
f"Applied tenant config overrides for bank {bank_id}: {list(configurable_tenant.keys())}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load tenant config for bank {bank_id}: {e}")
|
||||
|
||||
# Load bank config overrides
|
||||
bank_overrides = await self._load_bank_config(bank_id)
|
||||
@@ -128,25 +92,6 @@ class ConfigResolver:
|
||||
# Return full config object (dataclass doesn't have __init__ that accepts kwargs, so we update the object)
|
||||
# Create a new config instance by copying the global config and updating fields
|
||||
resolved_config = HindsightConfig(**config_dict)
|
||||
# Multi-LLM chains are static credential fields (never tenant/bank-overridable),
|
||||
# but asdict() above flattened their member dataclasses into plain dicts. Restore
|
||||
# the original typed objects from the global config so the resolved object stays
|
||||
# well-typed for any consumer that reads them.
|
||||
resolved_config = replace(
|
||||
resolved_config,
|
||||
llm_members=self._global_config.llm_members,
|
||||
llm_strategy=self._global_config.llm_strategy,
|
||||
retain_llm_members=self._global_config.retain_llm_members,
|
||||
retain_llm_strategy=self._global_config.retain_llm_strategy,
|
||||
reflect_llm_members=self._global_config.reflect_llm_members,
|
||||
reflect_llm_strategy=self._global_config.reflect_llm_strategy,
|
||||
consolidation_llm_members=self._global_config.consolidation_llm_members,
|
||||
consolidation_llm_strategy=self._global_config.consolidation_llm_strategy,
|
||||
)
|
||||
validate_retain_chunking_config(
|
||||
resolved_config.retain_chunk_size,
|
||||
resolved_config.retain_structured_chunk_size,
|
||||
)
|
||||
return resolved_config
|
||||
|
||||
async def get_bank_config(self, bank_id: str, context: RequestContext | None = None) -> dict[str, Any]:
|
||||
@@ -177,83 +122,26 @@ class ConfigResolver:
|
||||
resolved_config = await self.resolve_full_config(bank_id, context)
|
||||
config_dict = asdict(resolved_config)
|
||||
|
||||
# SECURITY: drop static/infrastructure + credential fields, then permission-filter.
|
||||
filtered = self._strip_static_and_credential_fields(config_dict)
|
||||
return await self._apply_permission_filter(filtered, bank_id, context)
|
||||
# SECURITY: Filter to only configurable fields (exclude static/infrastructure)
|
||||
filtered = {k: v for k, v in config_dict.items() if k in self._configurable_fields}
|
||||
|
||||
def _strip_static_and_credential_fields(self, config_dict: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Keep only configurable, non-credential fields.
|
||||
# SECURITY: Remove ALL credential fields (API keys, base URLs, etc.)
|
||||
filtered = {k: v for k, v in filtered.items() if k not in self._credential_fields}
|
||||
|
||||
SECURITY: excludes static/infrastructure fields and ALL credential fields
|
||||
(API keys, base URLs, etc.) so a resolved config is safe to return over the API.
|
||||
"""
|
||||
return {
|
||||
k: v for k, v in config_dict.items() if k in self._configurable_fields and k not in self._credential_fields
|
||||
}
|
||||
|
||||
async def _apply_permission_filter(
|
||||
self, filtered: dict[str, Any], bank_id: str, context: RequestContext | None
|
||||
) -> dict[str, Any]:
|
||||
"""Further restrict already-stripped config to the tenant/bank permission allow-list.
|
||||
|
||||
On extension error, leaves ``filtered`` unchanged (parity with the historical
|
||||
single-bank path: a permissions lookup failure must not leak or drop fields).
|
||||
"""
|
||||
if not (self.tenant_extension and context):
|
||||
return filtered
|
||||
try:
|
||||
allowed_fields = await self.tenant_extension.get_allowed_config_fields(context, bank_id)
|
||||
if allowed_fields is not None: # None means "allow all"
|
||||
filtered = {k: v for k, v in filtered.items() if k in allowed_fields}
|
||||
logger.debug(
|
||||
f"Applied permission filter for bank {bank_id}: allowed={len(allowed_fields)} fields, "
|
||||
f"returned={len(filtered)} fields"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load permissions for bank {bank_id}: {e}")
|
||||
return filtered
|
||||
|
||||
async def get_bank_configs(
|
||||
self, bank_ids: list[str], context: RequestContext | None = None
|
||||
) -> dict[str, dict[str, Any]]:
|
||||
"""Batch variant of :meth:`get_bank_config` for many banks.
|
||||
|
||||
Equivalent to calling ``get_bank_config`` per bank, but resolves the
|
||||
global + tenant base once and loads every bank's ``banks.config`` JSONB
|
||||
in a single query, instead of one config round-trip per bank. Used by
|
||||
``list_banks`` to overlay disposition + mission without an N+1.
|
||||
|
||||
Returns a mapping of bank_id -> filtered configurable-field dict. A bank
|
||||
with no config row still appears, mapped to the global+tenant base.
|
||||
"""
|
||||
if not bank_ids:
|
||||
return {}
|
||||
|
||||
# Global + tenant base, resolved once (tenant override is per-request, not per-bank).
|
||||
base_dict = asdict(self._global_config)
|
||||
# PERMISSIONS: Further filter based on tenant/bank permissions
|
||||
if self.tenant_extension and context:
|
||||
try:
|
||||
tenant_overrides = await self.tenant_extension.get_tenant_config(context)
|
||||
if tenant_overrides:
|
||||
normalized_tenant = normalize_config_dict(tenant_overrides)
|
||||
base_dict.update({k: v for k, v in normalized_tenant.items() if k in self._configurable_fields})
|
||||
allowed_fields = await self.tenant_extension.get_allowed_config_fields(context, bank_id)
|
||||
if allowed_fields is not None: # None means "allow all"
|
||||
filtered = {k: v for k, v in filtered.items() if k in allowed_fields}
|
||||
logger.debug(
|
||||
f"Applied permission filter for bank {bank_id}: allowed={len(allowed_fields)} fields, "
|
||||
f"returned={len(filtered)} fields"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load tenant config for bulk resolve: {e}")
|
||||
logger.warning(f"Failed to load permissions for bank {bank_id}: {e}")
|
||||
|
||||
# All bank overrides in one query, then merge + strip per bank.
|
||||
bank_overrides = await self._load_bank_configs(bank_ids)
|
||||
stripped = {
|
||||
bank_id: self._strip_static_and_credential_fields({**base_dict, **bank_overrides.get(bank_id, {})})
|
||||
for bank_id in bank_ids
|
||||
}
|
||||
|
||||
# Permission filter is per-bank; resolve concurrently when an extension is present.
|
||||
if not (self.tenant_extension and context):
|
||||
return stripped
|
||||
permission_filtered = await asyncio.gather(
|
||||
*(self._apply_permission_filter(stripped[bank_id], bank_id, context) for bank_id in bank_ids)
|
||||
)
|
||||
return dict(zip(bank_ids, permission_filtered, strict=True))
|
||||
return filtered
|
||||
|
||||
async def _load_bank_config(self, bank_id: str) -> dict[str, Any]:
|
||||
"""
|
||||
@@ -292,45 +180,6 @@ class ConfigResolver:
|
||||
|
||||
return {}
|
||||
|
||||
async def _load_bank_configs(self, bank_ids: list[str]) -> dict[str, dict[str, Any]]:
|
||||
"""Bulk variant of :meth:`_load_bank_config`: load many banks' overrides in one query.
|
||||
|
||||
Returns a mapping of bank_id -> normalized active overrides. Banks with no row
|
||||
(or an empty/all-tombstone config) are simply absent from the mapping.
|
||||
"""
|
||||
result: dict[str, dict[str, Any]] = {}
|
||||
if not bank_ids:
|
||||
return result
|
||||
try:
|
||||
async with self._backend.acquire() as conn:
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT bank_id, config FROM {fq_table("banks")} WHERE bank_id = ANY($1)
|
||||
""",
|
||||
bank_ids,
|
||||
)
|
||||
for row in rows:
|
||||
config_data = row["config"]
|
||||
if not config_data:
|
||||
continue
|
||||
# Handle case where JSONB is returned as JSON string
|
||||
if isinstance(config_data, str):
|
||||
config_data = json.loads(config_data)
|
||||
|
||||
# Normalize keys (handle both env var format and Python field format)
|
||||
normalized = normalize_config_dict(config_data)
|
||||
|
||||
# Only active overrides for configurable fields. JSON null is a tombstone
|
||||
# for "Server Default" in the bank-config UI and must not override defaults.
|
||||
overrides = {
|
||||
k: v for k, v in normalized.items() if k in self._configurable_fields and v is not None
|
||||
}
|
||||
if overrides:
|
||||
result[row["bank_id"]] = overrides
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to bulk-load bank configs: {e}")
|
||||
return result
|
||||
|
||||
async def update_bank_config(
|
||||
self, bank_id: str, updates: dict[str, Any], context: RequestContext | None = None
|
||||
) -> None:
|
||||
@@ -417,32 +266,6 @@ class ConfigResolver:
|
||||
# Validate recall budget fields
|
||||
_validate_recall_budget_updates(normalized_updates)
|
||||
|
||||
# Validate disposition trait fields (1-5 integer scale)
|
||||
_validate_disposition_updates(normalized_updates)
|
||||
|
||||
chunking_fields_updated = (
|
||||
"retain_chunk_size" in normalized_updates
|
||||
or "retain_structured_chunk_size" in normalized_updates
|
||||
or "retain_strategies" in normalized_updates
|
||||
)
|
||||
if chunking_fields_updated:
|
||||
config_dict = await self._resolve_parent_config_dict(bank_id, context)
|
||||
active_bank_overrides = await self._load_bank_config(bank_id)
|
||||
for key, value in normalized_updates.items():
|
||||
if key not in self._configurable_fields:
|
||||
continue
|
||||
if value is None:
|
||||
active_bank_overrides.pop(key, None)
|
||||
else:
|
||||
active_bank_overrides[key] = value
|
||||
config_dict.update(active_bank_overrides)
|
||||
base_config = HindsightConfig(**config_dict)
|
||||
validate_retain_chunking_config(
|
||||
base_config.retain_chunk_size,
|
||||
base_config.retain_structured_chunk_size,
|
||||
)
|
||||
_validate_retain_strategy_chunking(base_config, base_config.retain_strategies)
|
||||
|
||||
# Persist the override. Banks are created lazily (on first retain), so a
|
||||
# PATCH that precedes any ingestion would otherwise UPDATE zero rows and
|
||||
# silently no-op while returning 200. Ensure the bank row exists first
|
||||
@@ -534,31 +357,6 @@ def _validate_recall_budget_updates(updates: dict[str, Any]) -> None:
|
||||
)
|
||||
|
||||
|
||||
_DISPOSITION_KEYS = (
|
||||
"disposition_skepticism",
|
||||
"disposition_literalism",
|
||||
"disposition_empathy",
|
||||
)
|
||||
|
||||
|
||||
def _validate_disposition_updates(updates: dict[str, Any]) -> None:
|
||||
"""Validate disposition trait config updates. Raises ValueError on invalid input.
|
||||
|
||||
Each trait is an integer on a 1-5 scale (or None to clear the per-bank
|
||||
override). The read overlay injects the stored value verbatim into a strict
|
||||
``DispositionTraits(int, ge=1, le=5)``; an out-of-contract value (a float, a
|
||||
0-1 scale, or an int outside 1-5) accepted here would later 500 the whole
|
||||
bank list when any bank profile is serialized (issue #2348).
|
||||
"""
|
||||
for key in _DISPOSITION_KEYS:
|
||||
if key in updates:
|
||||
value = updates[key]
|
||||
if value is None:
|
||||
continue
|
||||
if not isinstance(value, int) or isinstance(value, bool) or not (1 <= value <= 5):
|
||||
raise ValueError(f"{key} must be an integer between 1 and 5, got {value!r}")
|
||||
|
||||
|
||||
def apply_strategy(config: HindsightConfig, strategy_name: str) -> HindsightConfig:
|
||||
"""
|
||||
Apply a named retain strategy's overrides on top of a resolved config.
|
||||
@@ -566,8 +364,7 @@ def apply_strategy(config: HindsightConfig, strategy_name: str) -> HindsightConf
|
||||
A strategy is a named set of hierarchical field overrides stored in
|
||||
config.retain_strategies. Any field in _HIERARCHICAL_FIELDS can be
|
||||
overridden, including retain_extraction_mode, retain_chunk_size,
|
||||
retain_structured_chunk_size, entity_labels,
|
||||
entities_allow_free_form, etc.
|
||||
entity_labels, entities_allow_free_form, etc.
|
||||
|
||||
Unknown strategy names log a warning and return config unchanged.
|
||||
Unknown or non-hierarchical fields in the strategy are silently ignored.
|
||||
@@ -589,17 +386,4 @@ def apply_strategy(config: HindsightConfig, strategy_name: str) -> HindsightConf
|
||||
return config
|
||||
|
||||
logger.debug(f"Applying retain strategy '{strategy_name}': {list(filtered.keys())}")
|
||||
resolved = replace(config, **filtered)
|
||||
validate_retain_chunking_config(
|
||||
resolved.retain_chunk_size,
|
||||
resolved.retain_structured_chunk_size,
|
||||
)
|
||||
validate_retain_completion_token_budget(
|
||||
llm_provider=resolved.llm_provider,
|
||||
retain_max_completion_tokens=resolved.retain_max_completion_tokens,
|
||||
retain_chunk_size=resolved.retain_chunk_size,
|
||||
retain_llm_model=resolved.retain_llm_model,
|
||||
llm_model=resolved.llm_model,
|
||||
retain_llm_provider=resolved.retain_llm_provider,
|
||||
)
|
||||
return resolved
|
||||
return replace(config, **filtered)
|
||||
|
||||
@@ -13,18 +13,9 @@ in-flight task so that N concurrent callers produce one query rather than N.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable
|
||||
|
||||
from .db_utils import acquire_with_retry
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .db.base import DatabaseBackend
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
from typing import Any, Awaitable, Callable
|
||||
|
||||
|
||||
class BankStatsCache:
|
||||
@@ -75,28 +66,17 @@ class BankStatsCache:
|
||||
schema: str,
|
||||
bank_id: str,
|
||||
loader: Callable[[], Awaitable[dict[str, Any]]],
|
||||
*,
|
||||
force_refresh: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""Return cached stats for `(schema, bank_id)` or call `loader()`.
|
||||
|
||||
Concurrent misses on the same key are coalesced onto a single
|
||||
in-flight loader. When ``force_refresh`` is set the cached value is
|
||||
ignored: the loader runs and its result replaces the cached entry.
|
||||
in-flight loader.
|
||||
"""
|
||||
if not self.enabled:
|
||||
return await loader()
|
||||
|
||||
key = (schema, bank_id)
|
||||
|
||||
if force_refresh:
|
||||
value = await loader()
|
||||
async with self._lock:
|
||||
self._store_unlocked(key, value)
|
||||
# Supersede any loader that was in flight for this key.
|
||||
self._in_flight.pop(key, None)
|
||||
return value
|
||||
|
||||
async with self._lock:
|
||||
cached = self._get_fresh_unlocked(key)
|
||||
if cached is not None:
|
||||
@@ -116,10 +96,7 @@ class BankStatsCache:
|
||||
value = await loader()
|
||||
except BaseException as exc:
|
||||
async with self._lock:
|
||||
# Invalidation may have detached this loader and allowed a new
|
||||
# one to claim the key. Never remove that newer loader's slot.
|
||||
if self._in_flight.get(key) is in_flight:
|
||||
self._in_flight.pop(key, None)
|
||||
self._in_flight.pop(key, None)
|
||||
if not in_flight.done():
|
||||
in_flight.set_exception(exc)
|
||||
# Suppress "Future exception was never retrieved" when no other
|
||||
@@ -129,12 +106,8 @@ class BankStatsCache:
|
||||
raise
|
||||
|
||||
async with self._lock:
|
||||
# Only the loader that still owns the key may populate the cache.
|
||||
# An invalidated loader can finish for its original callers, but its
|
||||
# pre-invalidation result must not overwrite a newer load.
|
||||
if self._in_flight.get(key) is in_flight:
|
||||
self._store_unlocked(key, value)
|
||||
self._in_flight.pop(key, None)
|
||||
self._store_unlocked(key, value)
|
||||
self._in_flight.pop(key, None)
|
||||
if not in_flight.done():
|
||||
in_flight.set_result(value)
|
||||
return value
|
||||
@@ -142,113 +115,8 @@ class BankStatsCache:
|
||||
async def invalidate(self, schema: str, bank_id: str) -> None:
|
||||
"""Drop any cached stats for `(schema, bank_id)`."""
|
||||
async with self._lock:
|
||||
key = (schema, bank_id)
|
||||
self._entries.pop(key, None)
|
||||
# Detach rather than cancel: existing callers may finish with the
|
||||
# snapshot they requested, while post-invalidation callers reload.
|
||||
self._in_flight.pop(key, None)
|
||||
self._entries.pop((schema, bank_id), None)
|
||||
|
||||
async def clear(self) -> None:
|
||||
async with self._lock:
|
||||
self._entries.clear()
|
||||
self._in_flight.clear()
|
||||
|
||||
|
||||
class DistributedBankStatsCache:
|
||||
"""Table-backed (cross-process) TTL cache for `get_bank_stats`.
|
||||
|
||||
Same ``get_or_load`` / ``invalidate`` / ``clear`` contract as
|
||||
:class:`BankStatsCache`, but the store is the per-schema ``bank_stats_cache``
|
||||
table instead of a per-process dict — so one worker's computation is shared
|
||||
with every other worker, and no caller recomputes while a fresh row exists.
|
||||
|
||||
On a hit, a call is a single primary-key ``SELECT`` (sub-millisecond); only a
|
||||
miss runs the (expensive) ``loader`` and writes the row back. Concurrent
|
||||
misses are *not* coalesced across processes (that would need a lock): they
|
||||
each compute and ``UPSERT``, last write wins — all results are correct, at the
|
||||
cost of a brief redundant compute at expiry.
|
||||
|
||||
Every DB touch is best-effort: if the cache table is unreachable or missing
|
||||
(e.g. a schema mid-migration), the call degrades to computing without caching
|
||||
rather than failing ``get_bank_stats``. PostgreSQL only — the engine keeps the
|
||||
in-process :class:`BankStatsCache` for Oracle.
|
||||
"""
|
||||
|
||||
def __init__(self, *, backend: "DatabaseBackend", ttl_seconds: float) -> None:
|
||||
self._backend = backend
|
||||
self._ttl = float(ttl_seconds)
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return self._ttl > 0
|
||||
|
||||
@staticmethod
|
||||
def _qualified(schema: str) -> str:
|
||||
return f'"{schema}".bank_stats_cache' if schema else "bank_stats_cache"
|
||||
|
||||
async def get_or_load(
|
||||
self,
|
||||
schema: str,
|
||||
bank_id: str,
|
||||
loader: Callable[[], Awaitable[dict[str, Any]]],
|
||||
*,
|
||||
force_refresh: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
if not self.enabled:
|
||||
return await loader()
|
||||
|
||||
table = self._qualified(schema)
|
||||
|
||||
# 1. Fresh row? Single PK lookup; ``payload::text`` sidesteps any
|
||||
# jsonb->object codec so we always decode the same way. Skipped when
|
||||
# the caller forces a refresh — then we recompute and overwrite below.
|
||||
if not force_refresh:
|
||||
try:
|
||||
async with acquire_with_retry(self._backend) as conn:
|
||||
row = await conn.fetchrow(
|
||||
f"SELECT payload::text AS payload FROM {table} "
|
||||
f"WHERE bank_id = $1 AND computed_at > now() - make_interval(secs => $2::double precision)",
|
||||
bank_id,
|
||||
self._ttl,
|
||||
)
|
||||
if row is not None:
|
||||
return json.loads(row["payload"])
|
||||
except Exception as exc: # noqa: BLE001 — cache read must never break the endpoint
|
||||
logger.debug("bank_stats_cache read failed for %s.%s (%s); computing uncached", schema, bank_id, exc)
|
||||
return await loader()
|
||||
|
||||
# 2. Miss — compute, then write the row back (best-effort).
|
||||
value = await loader()
|
||||
try:
|
||||
async with acquire_with_retry(self._backend) as conn:
|
||||
await conn.execute(
|
||||
f"INSERT INTO {table} (bank_id, payload, computed_at) VALUES ($1, $2::jsonb, now()) "
|
||||
f"ON CONFLICT (bank_id) DO UPDATE SET payload = EXCLUDED.payload, computed_at = now()",
|
||||
bank_id,
|
||||
json.dumps(value),
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 — a failed write just means no caching this round
|
||||
logger.warning("bank_stats_cache write failed for %s.%s (%s)", schema, bank_id, exc)
|
||||
return value
|
||||
|
||||
async def invalidate(self, schema: str, bank_id: str) -> None:
|
||||
"""Drop the cached row so the next read recomputes."""
|
||||
if not self.enabled:
|
||||
return
|
||||
try:
|
||||
async with acquire_with_retry(self._backend) as conn:
|
||||
await conn.execute(f"DELETE FROM {self._qualified(schema)} WHERE bank_id = $1", bank_id)
|
||||
except Exception as exc: # noqa: BLE001 — invalidation must never break the write path
|
||||
logger.debug("bank_stats_cache invalidate failed for %s.%s (%s)", schema, bank_id, exc)
|
||||
|
||||
async def clear(self) -> None:
|
||||
"""Drop all cached rows in the current schema (best-effort)."""
|
||||
if not self.enabled:
|
||||
return
|
||||
from .memory_engine import get_current_schema
|
||||
|
||||
try:
|
||||
async with acquire_with_retry(self._backend) as conn:
|
||||
await conn.execute(f"DELETE FROM {self._qualified(get_current_schema())}")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("bank_stats_cache clear failed (%s)", exc)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -98,7 +98,7 @@ _DEDUP_TOP_K = 5
|
||||
class _DedupDecision(BaseModel):
|
||||
"""Focused 1-by-1 verdict for whether a new observation duplicates an existing one."""
|
||||
|
||||
action: Literal["merge", "keep"] = "keep"
|
||||
action: Literal["merge", "keep"]
|
||||
text: str = "" # the synthesized merged observation (when action == "merge")
|
||||
reason: str = ""
|
||||
|
||||
@@ -224,18 +224,13 @@ async def _dedup_reconcile_create(
|
||||
# Fold the new source facts into the twin and persist the merged text. We keep the twin's
|
||||
# existing embedding: the merged text is >= threshold similar, so the stored vector stays
|
||||
# representative and we avoid a re-embed + a dialect-specific vector UPDATE.
|
||||
search_vector_clause = (
|
||||
f",\n search_vector = to_tsvector('{config.text_search_extension_native_language}'::regconfig, COALESCE($1, ''))"
|
||||
if config.text_search_extension == "native"
|
||||
else ""
|
||||
)
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("memory_units")}
|
||||
SET text = $1,
|
||||
source_memory_ids = (SELECT array_agg(DISTINCT e) FROM unnest(source_memory_ids || $2::uuid[]) e),
|
||||
proof_count = (SELECT count(DISTINCT e) FROM unnest(source_memory_ids || $2::uuid[]) e),
|
||||
updated_at = now(){search_vector_clause}
|
||||
updated_at = now()
|
||||
WHERE id = $3::uuid
|
||||
""",
|
||||
outcome.merged_text,
|
||||
@@ -284,11 +279,6 @@ async def _dedup_reconcile_update(
|
||||
# the create path) then delete the now-redundant updated row. The all_strict/any tag match
|
||||
# guarantees twin and updated share scope, so dropping the updated row's tags loses no
|
||||
# visibility. Temporal fields follow the surviving twin (minimal scope; matches create).
|
||||
search_vector_clause = (
|
||||
f",\n search_vector = to_tsvector('{config.text_search_extension_native_language}'::regconfig, COALESCE($1, ''))"
|
||||
if config.text_search_extension == "native"
|
||||
else ""
|
||||
)
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("memory_units")} t
|
||||
@@ -299,7 +289,7 @@ async def _dedup_reconcile_update(
|
||||
proof_count = (
|
||||
SELECT count(DISTINCT e) FROM unnest(t.source_memory_ids || u.source_memory_ids) e
|
||||
),
|
||||
updated_at = now(){search_vector_clause}
|
||||
updated_at = now()
|
||||
FROM {fq_table("memory_units")} u
|
||||
WHERE t.id = $2::uuid AND u.id = $3::uuid
|
||||
""",
|
||||
@@ -345,15 +335,7 @@ def _resolve_obs_tags_list(memory: dict[str, Any]) -> list[list[str]] | None:
|
||||
|
||||
Returns ``None`` for the default ``combined``-mode single pass (caller uses
|
||||
the memory's own tags). Returns a list[list[str]] when the memory requested
|
||||
multi-pass scoping (``per_tag``, ``all_combinations``, ``shared``, or an
|
||||
explicit list).
|
||||
|
||||
``shared`` resolves to ``[[]]`` — a single pass over the empty (untagged)
|
||||
scope. The created observation carries no tags and recall/dedup match it with
|
||||
``tags_match="any"``, so every memory consolidates into one shared observation
|
||||
regardless of its own tags. Use it to deduplicate across volatile per-call
|
||||
provenance tags (e.g. per-session ids) without dropping those tags from the
|
||||
source facts.
|
||||
multi-pass scoping (``per_tag``, ``all_combinations``, or an explicit list).
|
||||
"""
|
||||
parsed = _parse_observation_scopes(memory)
|
||||
tags = list(memory.get("tags") or [])
|
||||
@@ -364,8 +346,6 @@ def _resolve_obs_tags_list(memory: dict[str, Any]) -> list[list[str]] | None:
|
||||
if not tags:
|
||||
return None
|
||||
return [list(c) for r in range(1, len(tags) + 1) for c in combinations(tags, r)]
|
||||
if parsed == "shared":
|
||||
return [[]]
|
||||
if parsed == "combined" or parsed is None:
|
||||
return None
|
||||
return parsed # explicit list[list[str]]
|
||||
@@ -382,7 +362,6 @@ def _resolve_write_scopes(memory: dict[str, Any]) -> list[frozenset[str]]:
|
||||
- ``combined`` / ``None`` -> ``[frozenset(memory.tags)]``
|
||||
- ``per_tag`` -> ``[frozenset({t}) for t in memory.tags]``
|
||||
- ``all_combinations`` -> one frozenset per nonempty subset of tags
|
||||
- ``shared`` -> ``[frozenset()]`` (the single untagged scope)
|
||||
- explicit ``list[list[str]]`` -> one frozenset per declared scope
|
||||
|
||||
Empty-tag memories collapse to a single ``frozenset()`` in all modes so they
|
||||
@@ -397,8 +376,6 @@ def _resolve_write_scopes(memory: dict[str, Any]) -> list[frozenset[str]]:
|
||||
if not tags:
|
||||
return [frozenset()]
|
||||
return [frozenset(c) for r in range(1, len(tags) + 1) for c in combinations(tags, r)]
|
||||
if parsed == "shared":
|
||||
return [frozenset()]
|
||||
if parsed == "combined" or parsed is None:
|
||||
return [frozenset(tags)]
|
||||
return [frozenset(s) for s in parsed] # explicit list[list[str]]
|
||||
@@ -459,13 +436,6 @@ class _CreateAction(BaseModel):
|
||||
def sanitize_text(cls, v: str) -> str:
|
||||
return sanitize_llm_output(v) or ""
|
||||
|
||||
@field_validator("source_fact_ids", mode="before")
|
||||
@classmethod
|
||||
def ensure_list(cls, v: str | list[str]) -> list[str]:
|
||||
if isinstance(v, str):
|
||||
return [v]
|
||||
return v
|
||||
|
||||
|
||||
class _UpdateAction(BaseModel):
|
||||
text: str
|
||||
@@ -478,13 +448,6 @@ class _UpdateAction(BaseModel):
|
||||
def sanitize_text(cls, v: str) -> str:
|
||||
return sanitize_llm_output(v) or ""
|
||||
|
||||
@field_validator("source_fact_ids", mode="before")
|
||||
@classmethod
|
||||
def ensure_list(cls, v: str | list[str]) -> list[str]:
|
||||
if isinstance(v, str):
|
||||
return [v]
|
||||
return v
|
||||
|
||||
|
||||
class _DeleteAction(BaseModel):
|
||||
observation_id: str # UUID of the observation to remove
|
||||
@@ -664,7 +627,6 @@ class ConsolidationPerfLog:
|
||||
self.start_time = time.time()
|
||||
self.lines: list[str] = []
|
||||
self.timings: dict[str, float] = {}
|
||||
self.timing_counts: dict[str, int] = {}
|
||||
self.llm_calls: int = 0
|
||||
self.total_obs_in_context: int = 0
|
||||
self.total_prompt_chars: int = 0
|
||||
@@ -674,13 +636,11 @@ class ConsolidationPerfLog:
|
||||
self.lines.append(message)
|
||||
|
||||
def record_timing(self, key: str, duration: float) -> None:
|
||||
"""Record a timing measurement.
|
||||
|
||||
Tracks both total seconds and call count so the summary can
|
||||
distinguish one slow call from many fast calls in aggregate.
|
||||
"""
|
||||
self.timings[key] = self.timings.get(key, 0.0) + duration
|
||||
self.timing_counts[key] = self.timing_counts.get(key, 0) + 1
|
||||
"""Record a timing measurement."""
|
||||
if key in self.timings:
|
||||
self.timings[key] += duration
|
||||
else:
|
||||
self.timings[key] = duration
|
||||
|
||||
def record_llm_call(self, obs_count: int, prompt_chars: int) -> None:
|
||||
"""Record stats for a single LLM call."""
|
||||
@@ -703,8 +663,6 @@ class ConsolidationPerfLog:
|
||||
"""
|
||||
for key, value in other.timings.items():
|
||||
self.timings[key] = self.timings.get(key, 0.0) + value
|
||||
for key, count in other.timing_counts.items():
|
||||
self.timing_counts[key] = self.timing_counts.get(key, 0) + count
|
||||
self.llm_calls += other.llm_calls
|
||||
self.total_obs_in_context += other.total_obs_in_context
|
||||
self.total_prompt_chars += other.total_prompt_chars
|
||||
@@ -1305,22 +1263,16 @@ async def _run_consolidation_job(
|
||||
f"{stats['skipped']} skipped)"
|
||||
)
|
||||
|
||||
# Add timing breakdown. Each phase is recorded once per call, so the count
|
||||
# disambiguates a single slow call from many fast calls — important for
|
||||
# operators triaging "the recall phase took 15s" log lines, where the
|
||||
# total is the sum of many serial sub-calls rather than one slow query.
|
||||
def _fmt(key: str) -> str:
|
||||
total = perf.timings[key]
|
||||
count = perf.timing_counts.get(key, 0)
|
||||
if count > 1:
|
||||
avg_ms = total * 1000.0 / count
|
||||
return f"{key}={total:.3f}s ({count} calls, avg={avg_ms:.0f}ms)"
|
||||
return f"{key}={total:.3f}s"
|
||||
|
||||
# Add timing breakdown
|
||||
timing_parts = []
|
||||
for key in ("recall", "llm", "embedding", "db_write"):
|
||||
if key in perf.timings:
|
||||
timing_parts.append(_fmt(key))
|
||||
if "recall" in perf.timings:
|
||||
timing_parts.append(f"recall={perf.timings['recall']:.3f}s")
|
||||
if "llm" in perf.timings:
|
||||
timing_parts.append(f"llm={perf.timings['llm']:.3f}s")
|
||||
if "embedding" in perf.timings:
|
||||
timing_parts.append(f"embedding={perf.timings['embedding']:.3f}s")
|
||||
if "db_write" in perf.timings:
|
||||
timing_parts.append(f"db_write={perf.timings['db_write']:.3f}s")
|
||||
|
||||
if perf.llm_calls > 0:
|
||||
timing_parts.append(f"avg_obs={perf.total_obs_in_context / perf.llm_calls:.1f}")
|
||||
@@ -1537,10 +1489,8 @@ async def _process_memory_batch(
|
||||
# the bank-wide max_observations_per_scope for scopes matching its tag pattern.
|
||||
max_obs = _effective_scope_limit(config, fact_tags)
|
||||
remaining_observation_slots: int | None = None
|
||||
if max_obs >= 0 and fact_tags:
|
||||
# max_obs == 0 means "no new observations": there are no slots regardless
|
||||
# of the current count, so skip the count query for that case.
|
||||
current_count = await _count_observations_for_scope(conn, bank_id, fact_tags) if max_obs > 0 else 0
|
||||
if max_obs > 0 and fact_tags:
|
||||
current_count = await _count_observations_for_scope(conn, bank_id, fact_tags)
|
||||
remaining_observation_slots = max(max_obs - current_count, 0)
|
||||
if remaining_observation_slots == 0:
|
||||
logger.info(
|
||||
@@ -1855,12 +1805,6 @@ async def _execute_update_action(
|
||||
|
||||
config = get_config()
|
||||
|
||||
search_vector_clause = (
|
||||
f",\n search_vector = to_tsvector('{config.text_search_extension_native_language}'::regconfig, COALESCE($1, ''))"
|
||||
if config.text_search_extension == "native"
|
||||
else ""
|
||||
)
|
||||
|
||||
t0 = time.time()
|
||||
await conn.execute(
|
||||
f"""
|
||||
@@ -1873,7 +1817,7 @@ async def _execute_update_action(
|
||||
updated_at = now(),
|
||||
occurred_start = LEAST(occurred_start, COALESCE($6, occurred_start)),
|
||||
occurred_end = GREATEST(occurred_end, COALESCE($7, occurred_end)),
|
||||
mentioned_at = GREATEST(mentioned_at, COALESCE($8, mentioned_at)){search_vector_clause}
|
||||
mentioned_at = GREATEST(mentioned_at, COALESCE($8, mentioned_at))
|
||||
WHERE id = $5
|
||||
""",
|
||||
new_text,
|
||||
@@ -2184,7 +2128,7 @@ async def _consolidate_batch_with_llm(
|
||||
|
||||
# Build capacity note for the prompt when observation limit is configured
|
||||
observation_capacity_note: str | None = None
|
||||
if remaining_observation_slots is not None and max_observations_per_scope >= 0:
|
||||
if remaining_observation_slots is not None and max_observations_per_scope > 0:
|
||||
if remaining_observation_slots == 0:
|
||||
observation_capacity_note = (
|
||||
f"OBSERVATION LIMIT REACHED ({max_observations_per_scope}/{max_observations_per_scope}). "
|
||||
@@ -2349,20 +2293,16 @@ async def _create_observation_directly(
|
||||
tokenize($3, 'llmlingua2')::bm25_catalog.bm25vector)
|
||||
RETURNING id
|
||||
"""
|
||||
elif config.text_search_extension == "native":
|
||||
# Native: search_vector is populated with to_tsvector() using the
|
||||
# configured native language dictionary, matching the batch insert
|
||||
# path in ops_postgresql.insert_facts_batch.
|
||||
query = f"""
|
||||
INSERT INTO {fq_table("memory_units")} (
|
||||
id, bank_id, text, fact_type, embedding, proof_count, source_memory_ids,
|
||||
tags, event_date, occurred_start, occurred_end, mentioned_at, search_vector
|
||||
)
|
||||
VALUES ($1, $2, $3, 'observation', $4::vector, 1, $5, $6, $7, $8, $9, $10,
|
||||
to_tsvector('{config.text_search_extension_native_language}'::regconfig, COALESCE($3, '')))
|
||||
RETURNING id
|
||||
"""
|
||||
else: # pg_textsearch, pgroonga, pg_search: indexes operate on base text columns directly
|
||||
else: # native, pg_textsearch, pgroonga, or pg_search
|
||||
# pg_textsearch / pgroonga / pg_search: indexes operate on base text
|
||||
# columns directly, so the dummy search_vector column is left NULL.
|
||||
# Native: the migration p4q5r6s7t8u9 dropped the GENERATED expression on
|
||||
# search_vector to allow per-deployment language configuration; the
|
||||
# batch insert path in ops_postgresql.insert_facts_batch now populates
|
||||
# it via to_tsvector($lang, ...). This single-observation INSERT does
|
||||
# not, so observations under the native backend currently land with
|
||||
# NULL search_vector and are not BM25-searchable until reflected/
|
||||
# re-ingested. Tracking a separate fix for that gap.
|
||||
query = f"""
|
||||
INSERT INTO {fq_table("memory_units")} (
|
||||
id, bank_id, text, fact_type, embedding, proof_count, source_memory_ids,
|
||||
|
||||
@@ -212,7 +212,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
device = "cpu"
|
||||
logger.info("Reranker: forcing CPU mode (HINDSIGHT_API_RERANKER_LOCAL_FORCE_CPU=1)")
|
||||
else:
|
||||
# Check for GPU (CUDA), Apple Silicon (MPS), or Intel XPU
|
||||
# Check for GPU (CUDA) or Apple Silicon (MPS)
|
||||
# Wrap in try-except to gracefully handle any device detection issues
|
||||
# (e.g., in CI environments or when PyTorch is built without GPU support)
|
||||
device = "cpu" # Default to CPU
|
||||
@@ -220,13 +220,10 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
has_gpu = torch.cuda.is_available() or (
|
||||
hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
|
||||
)
|
||||
# Intel Arc XPU support — torch.xpu is available when the XPU build is loaded
|
||||
if not has_gpu and hasattr(torch, "xpu"):
|
||||
has_gpu = torch.xpu.is_available()
|
||||
if has_gpu:
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS/XPU
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to detect GPU/MPS/XPU, falling back to CPU: {e}")
|
||||
logger.warning(f"Failed to detect GPU/MPS, falling back to CPU: {e}")
|
||||
|
||||
# Patch transformers 5.x compatibility for models using XLM-RoBERTa
|
||||
# (e.g., jina-reranker-v2-base-multilingual). transformers 5.x removed
|
||||
|
||||
@@ -256,21 +256,13 @@ class OracleOps(DataAccessOps):
|
||||
# Oracle doesn't support ON CONFLICT; rely on the PK and the
|
||||
# IGNORE_ROW_ON_DUPKEY_INDEX hint to skip duplicates server-side.
|
||||
# The hint name must match the PK constraint exactly.
|
||||
#
|
||||
# Sort to enforce a global lock-acquisition order on the
|
||||
# (bank_id, unit_id) PK. Without this, two concurrent
|
||||
# transactions inserting overlapping unit_id sets in different
|
||||
# orders can deadlock on the unique-check row locks. Sorting
|
||||
# gives every concurrent caller the same lock order, so
|
||||
# conflicting inserts queue cleanly instead of cycling.
|
||||
sorted_unit_ids = sorted(unit_ids)
|
||||
await conn.executemany(
|
||||
f"""
|
||||
INSERT /*+ IGNORE_ROW_ON_DUPKEY_INDEX({table}, pk_graph_maintenance_queue) */
|
||||
INTO {table} (bank_id, unit_id)
|
||||
VALUES ($1, $2)
|
||||
""",
|
||||
[(bank_id, uid) for uid in sorted_unit_ids],
|
||||
[(bank_id, uid) for uid in unit_ids],
|
||||
)
|
||||
|
||||
async def claim_graph_maintenance_batch(
|
||||
|
||||
@@ -348,15 +348,6 @@ class PostgreSQLOps(DataAccessOps):
|
||||
) -> None:
|
||||
if not unit_ids:
|
||||
return
|
||||
# Sort to enforce a global lock-acquisition order on the
|
||||
# (bank_id, unit_id) unique-key. Without this, two concurrent
|
||||
# transactions inserting overlapping unit_id sets in different
|
||||
# orders can deadlock on the ON CONFLICT row locks — Postgres
|
||||
# acquires a short-lived lock per row being checked, and cycle
|
||||
# detection then aborts one transaction. Sorting gives every
|
||||
# concurrent caller the same lock order, so conflicting inserts
|
||||
# queue cleanly instead of cycling.
|
||||
sorted_unit_ids = sorted(unit_ids)
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {table} (bank_id, unit_id)
|
||||
@@ -364,7 +355,7 @@ class PostgreSQLOps(DataAccessOps):
|
||||
ON CONFLICT (bank_id, unit_id) DO NOTHING
|
||||
""",
|
||||
bank_id,
|
||||
sorted_unit_ids,
|
||||
unit_ids,
|
||||
)
|
||||
|
||||
async def claim_graph_maintenance_batch(
|
||||
|
||||
@@ -106,15 +106,6 @@ SCHEMAS_WITH_PENDING_WORK = OptionalRoutine(
|
||||
deployment.
|
||||
* Should be cheap and idempotent — called every poll cycle (~30s).
|
||||
|
||||
The poller trusts the result wholesale: any schema the routine does
|
||||
not return is treated as having no work this cycle. It does NOT
|
||||
second-guess omissions with a per-schema scan — that would re-run the
|
||||
exact queries this routine exists to avoid. Consequently the routine
|
||||
is *only* appropriate for multi-tenant deployments. Single-schema
|
||||
(default/public only) installs should NOT create it: the per-schema
|
||||
fallback below is a single cheap EXISTS check that covers ``public``
|
||||
correctly and cannot starve.
|
||||
|
||||
Fallback when the routine is absent: per-schema ``EXISTS`` queries
|
||||
from Python (~4ms per schema). The server-side path is a single-
|
||||
round-trip optimisation worth ~200ms in deployments with thousands
|
||||
|
||||
@@ -190,7 +190,7 @@ class LocalSTEmbeddings(Embeddings):
|
||||
device = "cpu"
|
||||
logger.info("Embeddings: forcing CPU mode")
|
||||
else:
|
||||
# Check for GPU (CUDA), Apple Silicon (MPS), or Intel XPU
|
||||
# Check for GPU (CUDA) or Apple Silicon (MPS)
|
||||
# Wrap in try-except to gracefully handle any device detection issues
|
||||
# (e.g., in CI environments or when PyTorch is built without GPU support)
|
||||
device = "cpu" # Default to CPU
|
||||
@@ -198,13 +198,10 @@ class LocalSTEmbeddings(Embeddings):
|
||||
has_gpu = torch.cuda.is_available() or (
|
||||
hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
|
||||
)
|
||||
# Intel Arc XPU support — torch.xpu is available when the XPU build is loaded
|
||||
if not has_gpu and hasattr(torch, "xpu"):
|
||||
has_gpu = torch.xpu.is_available()
|
||||
if has_gpu:
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS/XPU
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to detect GPU/MPS/XPU, falling back to CPU: {e}")
|
||||
logger.warning(f"Failed to detect GPU/MPS, falling back to CPU: {e}")
|
||||
|
||||
# Suppress verbose transformers warnings during model loading
|
||||
# This suppresses the "UNEXPECTED" warnings from BertModel which are harmless
|
||||
@@ -712,8 +709,7 @@ class OpenAIEmbeddings(Embeddings):
|
||||
|
||||
class CodexOAuthEmbeddings(OpenAIEmbeddings):
|
||||
"""
|
||||
OpenAI embeddings using the Codex/ChatGPT OAuth token from the Codex
|
||||
``auth.json`` (``$CODEX_HOME/auth.json``, or ``~/.codex/auth.json`` when unset).
|
||||
OpenAI embeddings using the Codex/ChatGPT OAuth token from ``~/.codex/auth.json``.
|
||||
|
||||
Codex OAuth is an LLM-provider auth path in Hindsight, but the same bearer token
|
||||
can also authenticate against the standard OpenAI embeddings endpoint. This keeps
|
||||
@@ -1638,20 +1634,6 @@ def create_embeddings_from_env() -> Embeddings:
|
||||
batch_size=config.embeddings_openai_batch_size,
|
||||
dimensions=config.embeddings_openai_dimensions,
|
||||
)
|
||||
elif provider == "requesty":
|
||||
api_key = config.embeddings_requesty_api_key
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"HINDSIGHT_API_EMBEDDINGS_REQUESTY_API_KEY, HINDSIGHT_API_REQUESTY_API_KEY, "
|
||||
f"or {ENV_LLM_API_KEY} is required when {ENV_EMBEDDINGS_PROVIDER} is 'requesty'"
|
||||
)
|
||||
return OpenAIEmbeddings(
|
||||
api_key=api_key,
|
||||
model=config.embeddings_requesty_model,
|
||||
base_url="https://router.requesty.ai/v1",
|
||||
batch_size=config.embeddings_openai_batch_size,
|
||||
dimensions=config.embeddings_openai_dimensions,
|
||||
)
|
||||
elif provider == "zeroentropy":
|
||||
api_key = config.embeddings_zeroentropy_api_key
|
||||
if not api_key:
|
||||
@@ -1715,6 +1697,6 @@ def create_embeddings_from_env() -> Embeddings:
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown embeddings provider: {provider}. "
|
||||
f"Supported: 'local', 'onnx', 'tei', 'openai', 'openai-codex', 'openrouter', 'requesty', 'cohere', 'google', "
|
||||
f"Supported: 'local', 'onnx', 'tei', 'openai', 'openai-codex', 'openrouter', 'cohere', 'google', "
|
||||
f"'zeroentropy', 'litellm', 'litellm-sdk'"
|
||||
)
|
||||
|
||||
@@ -782,6 +782,236 @@ class EntityResolver:
|
||||
|
||||
return entity_ids
|
||||
|
||||
async def resolve_entity(
|
||||
self,
|
||||
bank_id: str,
|
||||
entity_text: str,
|
||||
context: str,
|
||||
nearby_entities: list[dict],
|
||||
unit_event_date,
|
||||
) -> str:
|
||||
"""
|
||||
Resolve an entity to a canonical entity ID.
|
||||
|
||||
Args:
|
||||
bank_id: bank ID (entities are scoped to agents)
|
||||
entity_text: Entity text ("Alice", "Google", etc.)
|
||||
context: Context where entity appears
|
||||
nearby_entities: Other entities in the same unit
|
||||
unit_event_date: When this unit was created
|
||||
|
||||
Returns:
|
||||
Entity ID (creates new entity if needed)
|
||||
"""
|
||||
async with acquire_with_retry(self.pool) as conn:
|
||||
# Find candidate entities with similar name
|
||||
candidates = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, canonical_name, metadata, last_seen
|
||||
FROM {fq_table("entities")}
|
||||
WHERE bank_id = $1
|
||||
AND (
|
||||
canonical_name ILIKE $2
|
||||
OR canonical_name ILIKE $3
|
||||
OR $2 ILIKE canonical_name || '%%'
|
||||
)
|
||||
ORDER BY mention_count DESC
|
||||
""",
|
||||
bank_id,
|
||||
entity_text,
|
||||
f"%{entity_text}%",
|
||||
)
|
||||
|
||||
if not candidates:
|
||||
# New entity - create it
|
||||
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
|
||||
|
||||
# Score candidates based on:
|
||||
# 1. Name similarity
|
||||
# 2. Context overlap (TODO: could use embeddings)
|
||||
# 3. Co-occurring entities
|
||||
# 4. Temporal proximity
|
||||
|
||||
best_candidate = None
|
||||
best_score = 0.0
|
||||
|
||||
nearby_entity_set = {e["text"].lower() for e in nearby_entities if e["text"] != entity_text}
|
||||
|
||||
for row in candidates:
|
||||
candidate_id = row["id"]
|
||||
canonical_name = row["canonical_name"]
|
||||
last_seen = row["last_seen"]
|
||||
score = 0.0
|
||||
|
||||
# 1. Name similarity (0-1)
|
||||
name_similarity = SequenceMatcher(None, entity_text.lower(), canonical_name.lower()).ratio()
|
||||
score += name_similarity * 0.5
|
||||
|
||||
# 2. Co-occurring entities (0-0.5)
|
||||
# Get entities that co-occurred with this candidate before
|
||||
# Use the materialized co-occurrence cache for fast lookup
|
||||
co_entity_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT e.canonical_name, ec.cooccurrence_count
|
||||
FROM {fq_table("entity_cooccurrences")} ec
|
||||
JOIN {fq_table("entities")} e ON (
|
||||
CASE
|
||||
WHEN ec.entity_id_1 = $1 THEN ec.entity_id_2
|
||||
WHEN ec.entity_id_2 = $1 THEN ec.entity_id_1
|
||||
END = e.id
|
||||
)
|
||||
WHERE ec.entity_id_1 = $1 OR ec.entity_id_2 = $1
|
||||
""",
|
||||
candidate_id,
|
||||
)
|
||||
co_entities = {r["canonical_name"].lower() for r in co_entity_rows}
|
||||
|
||||
# Check overlap with nearby entities
|
||||
overlap = len(nearby_entity_set & co_entities)
|
||||
if nearby_entity_set:
|
||||
co_entity_score = overlap / len(nearby_entity_set)
|
||||
score += co_entity_score * 0.3
|
||||
|
||||
# 3. Temporal proximity (0-0.2)
|
||||
if last_seen:
|
||||
# Normalize both to UTC-aware to avoid naive/aware mismatch
|
||||
# (Oracle returns naive datetimes from fromisoformat)
|
||||
_evt = unit_event_date if unit_event_date.tzinfo else unit_event_date.replace(tzinfo=UTC)
|
||||
_seen = last_seen if last_seen.tzinfo else last_seen.replace(tzinfo=UTC)
|
||||
days_diff = abs((_evt - _seen).total_seconds() / 86400)
|
||||
if days_diff < 7: # Within a week
|
||||
temporal_score = max(0, 1.0 - (days_diff / 7))
|
||||
score += temporal_score * 0.2
|
||||
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_candidate = candidate_id
|
||||
|
||||
# Threshold for considering it the same entity
|
||||
threshold = 0.6
|
||||
|
||||
if best_score > threshold:
|
||||
# Update entity
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("entities")}
|
||||
SET mention_count = mention_count + 1,
|
||||
last_seen = $1
|
||||
WHERE id = $2
|
||||
""",
|
||||
unit_event_date,
|
||||
best_candidate,
|
||||
)
|
||||
return best_candidate
|
||||
else:
|
||||
# Not confident - create new entity
|
||||
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
|
||||
|
||||
async def _create_entity(
|
||||
self,
|
||||
conn,
|
||||
bank_id: str,
|
||||
entity_text: str,
|
||||
event_date,
|
||||
) -> str:
|
||||
"""
|
||||
Create a new entity or get existing one if it already exists.
|
||||
|
||||
Uses INSERT ... ON CONFLICT to handle race conditions where
|
||||
two concurrent transactions try to create the same entity.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
bank_id: bank ID
|
||||
entity_text: Entity text
|
||||
event_date: When first seen
|
||||
|
||||
Returns:
|
||||
Entity ID
|
||||
"""
|
||||
entity_id = await conn.fetchval(
|
||||
f"""
|
||||
INSERT INTO {fq_table("entities")} (bank_id, canonical_name, first_seen, last_seen, mention_count)
|
||||
VALUES ($1, $2, COALESCE($3, now()), COALESCE($4, now()), 1)
|
||||
ON CONFLICT (bank_id, LOWER(canonical_name))
|
||||
DO UPDATE SET
|
||||
mention_count = {fq_table("entities")}.mention_count + 1,
|
||||
last_seen = EXCLUDED.last_seen
|
||||
RETURNING id
|
||||
""",
|
||||
bank_id,
|
||||
entity_text,
|
||||
event_date,
|
||||
event_date,
|
||||
)
|
||||
return entity_id
|
||||
|
||||
async def link_unit_to_entity(self, unit_id: str, entity_id: str):
|
||||
"""
|
||||
Link a memory unit to an entity.
|
||||
Also updates co-occurrence cache with other entities in the same unit.
|
||||
|
||||
Args:
|
||||
unit_id: Memory unit ID
|
||||
entity_id: Entity ID
|
||||
"""
|
||||
async with acquire_with_retry(self.pool) as conn:
|
||||
# Insert unit-entity link
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("unit_entities")} (unit_id, entity_id)
|
||||
VALUES ($1, $2)
|
||||
ON CONFLICT DO NOTHING
|
||||
""",
|
||||
unit_id,
|
||||
entity_id,
|
||||
)
|
||||
|
||||
# Update co-occurrence cache: find other entities in this unit
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT entity_id
|
||||
FROM {fq_table("unit_entities")}
|
||||
WHERE unit_id = $1 AND entity_id != $2
|
||||
""",
|
||||
unit_id,
|
||||
entity_id,
|
||||
)
|
||||
|
||||
other_entities = [row["entity_id"] for row in rows]
|
||||
|
||||
# Update co-occurrences for each pair
|
||||
for other_entity_id in other_entities:
|
||||
await self._update_cooccurrence(conn, entity_id, other_entity_id)
|
||||
|
||||
async def _update_cooccurrence(self, conn, entity_id_1: str, entity_id_2: str):
|
||||
"""
|
||||
Update the co-occurrence cache for two entities.
|
||||
|
||||
Uses CHECK constraint ordering (entity_id_1 < entity_id_2) to avoid duplicates.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
entity_id_1: First entity ID
|
||||
entity_id_2: Second entity ID
|
||||
"""
|
||||
# Ensure consistent ordering (smaller UUID first)
|
||||
if entity_id_1 > entity_id_2:
|
||||
entity_id_1, entity_id_2 = entity_id_2, entity_id_1
|
||||
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("entity_cooccurrences")} (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
|
||||
VALUES ($1, $2, 1, NOW())
|
||||
ON CONFLICT (entity_id_1, entity_id_2)
|
||||
DO UPDATE SET
|
||||
cooccurrence_count = {fq_table("entity_cooccurrences")}.cooccurrence_count + 1,
|
||||
last_cooccurred = NOW()
|
||||
""",
|
||||
entity_id_1,
|
||||
entity_id_2,
|
||||
)
|
||||
|
||||
async def link_units_to_entities_batch(
|
||||
self,
|
||||
unit_entity_pairs: list[tuple[str, str]] | list[tuple[str, str, datetime | None]],
|
||||
|
||||
@@ -449,7 +449,6 @@ class MemoryEngineInterface(ABC):
|
||||
bank_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
force_refresh: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Get statistics about memory nodes and links for a bank.
|
||||
@@ -457,8 +456,6 @@ class MemoryEngineInterface(ABC):
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
request_context: Request context for authentication.
|
||||
force_refresh: Bypass the cached value and recompute (also refreshes
|
||||
the cache for subsequent callers).
|
||||
|
||||
Returns:
|
||||
Dict with node_counts, link_counts, link_counts_by_fact_type
|
||||
|
||||
@@ -6,7 +6,6 @@ enabling support for multiple LLM backends (OpenAI, Anthropic, Gemini, Codex, et
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from .response_models import LLMToolCallResult
|
||||
@@ -253,11 +252,3 @@ class OutputTooLongError(Exception):
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class ProviderRateLimitResetError(Exception):
|
||||
"""Raised when an upstream provider says quota will reopen at a known time."""
|
||||
|
||||
def __init__(self, retry_at: datetime, message: str = "") -> None:
|
||||
self.retry_at = retry_at
|
||||
super().__init__(message)
|
||||
|
||||
@@ -76,51 +76,6 @@ _request_ctx: ContextVar[dict[str, Any] | None] = ContextVar("hindsight_llm_requ
|
||||
_call_metadata_ctx: ContextVar[dict[str, Any] | None] = ContextVar("hindsight_llm_call_metadata_ctx", default=None)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMResponseUsage:
|
||||
"""Provider-reported token usage for the in-flight LLM call.
|
||||
|
||||
Stashed by provider implementations as soon as a response is received —
|
||||
*before* local JSON parsing / schema validation, which may still fail. The
|
||||
wrapper reads it to attach real token counts to an error trace when the
|
||||
provider call itself succeeded but the structured output couldn't be parsed
|
||||
or validated (providers charge for those tokens regardless). See #2387.
|
||||
"""
|
||||
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cached_tokens: int = 0
|
||||
|
||||
|
||||
# Per-call provider usage, set by providers right after a response is received.
|
||||
_response_usage_ctx: ContextVar[LLMResponseUsage | None] = ContextVar("hindsight_llm_response_usage_ctx", default=None)
|
||||
|
||||
|
||||
def set_response_usage(usage: LLMResponseUsage | None) -> Token:
|
||||
"""Bind provider-reported usage for the current call. Returns a reset token."""
|
||||
return _response_usage_ctx.set(usage)
|
||||
|
||||
|
||||
def stash_response_usage(usage: LLMResponseUsage | None) -> None:
|
||||
"""Record provider-reported usage so an error trace can attach it later.
|
||||
|
||||
Called by provider implementations once a response (with usage) is in hand,
|
||||
before parsing/validation that may raise. Overwrites any prior value from an
|
||||
earlier retry attempt so the last attempt's usage wins.
|
||||
"""
|
||||
_response_usage_ctx.set(usage)
|
||||
|
||||
|
||||
def reset_response_usage(token: Token) -> None:
|
||||
"""Unwind a binding made by :func:`set_response_usage`."""
|
||||
_response_usage_ctx.reset(token)
|
||||
|
||||
|
||||
def current_response_usage() -> LLMResponseUsage | None:
|
||||
"""Return the active call's provider-reported usage, or None."""
|
||||
return _response_usage_ctx.get()
|
||||
|
||||
|
||||
def set_trace_context(ctx: LLMTraceContext | None) -> Token:
|
||||
"""Bind trace attribution to the current context. Returns a reset token."""
|
||||
return _trace_ctx.set(ctx)
|
||||
|
||||
@@ -10,6 +10,7 @@ import re
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import AsyncExitStack
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
# Vertex AI imports (conditional - for LLMProvider to pass credentials to GeminiLLM)
|
||||
@@ -252,8 +253,6 @@ def create_llm_provider(
|
||||
gemini_safety_settings: list | None = None,
|
||||
prompt_cache_enabled: bool = False,
|
||||
litellmrouter_config: dict[str, Any] | None = None,
|
||||
gemini_service_tier: str | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> Any: # Returns LLMInterface
|
||||
"""
|
||||
Factory function to create the appropriate LLM provider implementation.
|
||||
@@ -267,26 +266,17 @@ def create_llm_provider(
|
||||
groq_service_tier: Groq service tier (for Groq provider) - "on_demand", "flex", or "auto".
|
||||
openai_service_tier: OpenAI service tier (for OpenAI provider) - None (default) or "flex" (50% cheaper).
|
||||
bedrock_service_tier: Bedrock service tier (for Bedrock provider) - None (default), "flex", "priority", or "reserved".
|
||||
gemini_service_tier: Gemini service tier (for Gemini provider) - None (default) or "flex" (50% cheaper).
|
||||
extra_body: Extra request-body params merged into the provider's native
|
||||
call. Threaded into OpenAI-compatible, Fireworks, Anthropic, Gemini/
|
||||
VertexAI and LiteLLM providers (each merges them in its own parameter
|
||||
space). Keys must use each provider's native names (e.g. ``max_tokens``
|
||||
for OpenAI/Anthropic vs ``max_output_tokens`` for Gemini).
|
||||
default_headers: Custom headers passed to provider SDK clients (used by operators
|
||||
routing through proxies / request-tracing middleware). Wired into the Anthropic
|
||||
provider (SDK ``default_headers``) and the LiteLLM-backed providers — ``litellm``,
|
||||
``litellmrouter`` and ``bedrock`` — as the LiteLLM ``extra_headers`` completion
|
||||
kwarg; other providers may opt in as needed.
|
||||
default_headers: Custom headers passed as ``default_headers`` to provider SDK clients
|
||||
(used by operators routing through proxies / request-tracing middleware). Currently
|
||||
wired into the Anthropic provider; other providers may opt in as needed.
|
||||
vertexai_project_id: Vertex AI project ID (for VertexAI provider).
|
||||
vertexai_region: Vertex AI region (for VertexAI provider).
|
||||
vertexai_credentials: Vertex AI credentials object (for VertexAI provider).
|
||||
timeout: Per-request LLM timeout in seconds (resolved by the caller from the
|
||||
per-operation/global config). Threaded into the providers that honour a
|
||||
configurable request timeout (LiteLLM, LiteLLM Router, OpenAI-compatible,
|
||||
Nous). ``None`` lets each provider fall back to its own default
|
||||
(``HINDSIGHT_API_LLM_TIMEOUT`` / ``DEFAULT_LLM_TIMEOUT`` for those four;
|
||||
Anthropic and Gemini keep their provider-specific defaults).
|
||||
|
||||
Returns:
|
||||
LLMInterface implementation for the specified provider.
|
||||
@@ -306,12 +296,6 @@ def create_llm_provider(
|
||||
)
|
||||
|
||||
provider_lower = provider.lower()
|
||||
if provider_lower == "gemini":
|
||||
from ..config import parse_gemini_service_tier
|
||||
|
||||
gemini_service_tier = parse_gemini_service_tier(gemini_service_tier)
|
||||
else:
|
||||
gemini_service_tier = None
|
||||
|
||||
if provider_lower == "openai-codex":
|
||||
return CodexLLM(
|
||||
@@ -360,7 +344,6 @@ def create_llm_provider(
|
||||
vertexai_region=vertexai_region,
|
||||
vertexai_credentials=vertexai_credentials,
|
||||
gemini_safety_settings=gemini_safety_settings,
|
||||
gemini_service_tier=gemini_service_tier,
|
||||
prompt_cache_enabled=prompt_cache_enabled,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
@@ -384,8 +367,6 @@ def create_llm_provider(
|
||||
model=model,
|
||||
reasoning_effort=reasoning_effort,
|
||||
extra_body=extra_body,
|
||||
default_headers=default_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
elif provider_lower == "litellmrouter":
|
||||
@@ -404,8 +385,6 @@ def create_llm_provider(
|
||||
config=litellmrouter_config,
|
||||
reasoning_effort=reasoning_effort,
|
||||
extra_body=extra_body,
|
||||
default_headers=default_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
elif provider_lower == "bedrock":
|
||||
@@ -418,9 +397,7 @@ def create_llm_provider(
|
||||
model=bedrock_model,
|
||||
reasoning_effort=reasoning_effort,
|
||||
extra_body=extra_body,
|
||||
default_headers=default_headers,
|
||||
bedrock_service_tier=bedrock_service_tier,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
elif provider_lower == "llamacpp":
|
||||
@@ -467,7 +444,6 @@ def create_llm_provider(
|
||||
model=model,
|
||||
reasoning_effort=reasoning_effort,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
elif provider_lower in (
|
||||
@@ -480,10 +456,8 @@ def create_llm_provider(
|
||||
"deepseek",
|
||||
"volcano",
|
||||
"openrouter",
|
||||
"requesty",
|
||||
"zai",
|
||||
"opencode-go",
|
||||
"atlas",
|
||||
):
|
||||
return OpenAICompatibleLLM(
|
||||
provider=provider,
|
||||
@@ -494,7 +468,6 @@ def create_llm_provider(
|
||||
groq_service_tier=groq_service_tier,
|
||||
openai_service_tier=openai_service_tier,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
else:
|
||||
@@ -523,14 +496,6 @@ class LLMProvider:
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
default_headers: dict[str, str] | None = None,
|
||||
litellmrouter_config: dict[str, Any] | None = None,
|
||||
gemini_service_tier: str | None = None,
|
||||
vertexai_project_id: str | None = None,
|
||||
vertexai_region: str | None = None,
|
||||
vertexai_service_account_key: str | None = None,
|
||||
timeout: float | None = None,
|
||||
max_retries: int | None = None,
|
||||
initial_backoff: float | None = None,
|
||||
max_backoff: float | None = None,
|
||||
):
|
||||
"""
|
||||
Initialize LLM provider.
|
||||
@@ -544,60 +509,29 @@ class LLMProvider:
|
||||
groq_service_tier: Groq service tier ("on_demand", "flex", "auto") - from config.
|
||||
openai_service_tier: OpenAI service tier (None or "flex") - from config.
|
||||
bedrock_service_tier: Bedrock service tier (None, "flex", "priority", "reserved") - from config.
|
||||
gemini_service_tier: Gemini service tier (None or "flex") - from config.
|
||||
gemini_safety_settings: Safety settings for Gemini/VertexAI providers.
|
||||
extra_body: Extra request-body params merged into the provider's native call
|
||||
(OpenAI-compatible, Fireworks, Anthropic, Gemini/VertexAI, LiteLLM).
|
||||
default_headers: Custom headers passed as ``default_headers`` to provider SDK clients.
|
||||
Used by operators routing through proxies / request-tracing middleware.
|
||||
Used by operators routing through proxies / request-tracing middleware. Falls
|
||||
back to ``HindsightConfig.llm_default_headers`` (env: ``HINDSIGHT_API_LLM_DEFAULT_HEADERS``)
|
||||
when ``None``.
|
||||
litellmrouter_config: Provider-specific config for ``provider="litellmrouter"``.
|
||||
JSON object passed verbatim to ``litellm.Router(**config)`` — see
|
||||
https://docs.litellm.ai/docs/routing. Ignored unless ``provider == "litellmrouter"``.
|
||||
vertexai_project_id: Vertex AI project ID for ``provider="vertexai"`` (required for
|
||||
that provider).
|
||||
vertexai_region: Vertex AI region for ``provider="vertexai"`` (defaults to
|
||||
``"us-central1"`` when ``None``).
|
||||
vertexai_service_account_key: Path to a Vertex AI service-account key file for
|
||||
``provider="vertexai"`` (uses ADC when ``None``).
|
||||
timeout: Per-request LLM timeout in seconds. Resolved by the caller from the
|
||||
per-operation/global config (``retain_llm_timeout`` falling back to
|
||||
``llm_timeout``, etc.). ``None`` lets each provider apply its own default.
|
||||
max_retries: Default retry-attempt budget for ``call`` / ``call_with_tools``
|
||||
when the per-call argument is omitted. Resolved by the caller from the
|
||||
per-operation/global config (``reflect_llm_max_retries`` falling back to
|
||||
``llm_max_retries``, etc.). ``None`` keeps each method's own fallback.
|
||||
initial_backoff: Default initial retry backoff (seconds), same resolution as
|
||||
``max_retries``. ``None`` keeps each method's own fallback.
|
||||
max_backoff: Default maximum retry backoff (seconds), same resolution as
|
||||
``max_retries``. ``None`` keeps each method's own fallback.
|
||||
|
||||
This constructor uses every argument as passed and does not read global
|
||||
``HindsightConfig``: resolving the server-level default for a ``None`` argument is the
|
||||
caller's responsibility (see ``MemoryEngine``'s per-op builds, ``_member_to_llm``, and
|
||||
``LLMProvider.from_env``). Keeping it config-free makes a provider's effective settings a
|
||||
pure function of its arguments — which is what lets each member of a multi-LLM chain be
|
||||
configured independently.
|
||||
When None and the provider is ``litellmrouter``, falls back to
|
||||
``HindsightConfig.llm_litellmrouter_config``.
|
||||
"""
|
||||
self.provider = provider.lower()
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.reasoning_effort = reasoning_effort
|
||||
# Per-request timeout (seconds). Used verbatim — the caller resolves the
|
||||
# per-operation/global fallback. ``None`` defers to the provider default.
|
||||
self.timeout = timeout
|
||||
# Default retry policy for call()/call_with_tools(). The caller resolves the
|
||||
# per-operation/global fallback; ``None`` keeps each method's own fallback so
|
||||
# providers built without a resolved config (from_env, tests) are unchanged.
|
||||
self.max_retries = max_retries
|
||||
self.initial_backoff = initial_backoff
|
||||
self.max_backoff = max_backoff
|
||||
self.litellmrouter_config = litellmrouter_config
|
||||
# Service tiers from hierarchical config (not env vars)
|
||||
self.groq_service_tier = groq_service_tier
|
||||
self.openai_service_tier = openai_service_tier
|
||||
self.bedrock_service_tier = bedrock_service_tier
|
||||
self.gemini_service_tier = gemini_service_tier
|
||||
# Gemini safety settings (instance default; can be overridden per-request via context var)
|
||||
self.gemini_safety_settings = gemini_safety_settings
|
||||
# Gemini prompt caching: when True, retain extraction (and any future
|
||||
@@ -608,9 +542,16 @@ class LLMProvider:
|
||||
# Extra body params for OpenAI-compatible providers (e.g. chat_template_kwargs)
|
||||
self.extra_body = extra_body
|
||||
# Default headers passed to provider SDK clients (e.g. proxy auth, request tracing).
|
||||
# Used verbatim — callers resolve the global fallback (see _member_to_llm /
|
||||
# the per-op builds in MemoryEngine, and LLMProvider.from_env).
|
||||
# Same pattern as ``gemini_safety_settings``: explicit override wins; otherwise read
|
||||
# the static server-level default from ``HindsightConfig`` via ``_get_raw_config()``.
|
||||
self.default_headers = default_headers
|
||||
if self.default_headers is None:
|
||||
from ..config import _get_raw_config
|
||||
|
||||
try:
|
||||
self.default_headers = _get_raw_config().llm_default_headers
|
||||
except Exception:
|
||||
pass # Config may not be initialized in test environments
|
||||
|
||||
# Validate provider
|
||||
valid_providers = [
|
||||
@@ -634,10 +575,8 @@ class LLMProvider:
|
||||
"bedrock",
|
||||
"volcano",
|
||||
"openrouter",
|
||||
"requesty",
|
||||
"zai",
|
||||
"opencode-go",
|
||||
"atlas",
|
||||
"fireworks",
|
||||
"nous",
|
||||
]
|
||||
@@ -660,31 +599,32 @@ class LLMProvider:
|
||||
self.base_url = "https://api.deepseek.com"
|
||||
elif self.provider == "openrouter":
|
||||
self.base_url = "https://openrouter.ai/api/v1"
|
||||
elif self.provider == "requesty":
|
||||
self.base_url = "https://router.requesty.ai/v1"
|
||||
elif self.provider == "zai":
|
||||
self.base_url = "https://api.z.ai/api/coding/paas/v4"
|
||||
elif self.provider == "opencode-go":
|
||||
self.base_url = "https://opencode.ai/zen/go/v1"
|
||||
elif self.provider == "atlas":
|
||||
self.base_url = "https://api.atlascloud.ai/v1"
|
||||
elif self.provider == "nous":
|
||||
self.base_url = "https://inference-api.nousresearch.com/v1"
|
||||
|
||||
# Prepare Vertex AI config (if applicable). Values are used as passed; the
|
||||
# caller resolves the global-config fallback (MemoryEngine builds /
|
||||
# _member_to_llm / from_env). The region keeps a constant default here.
|
||||
# Prepare Vertex AI config (if applicable)
|
||||
vertexai_project_id = None
|
||||
vertexai_region = None
|
||||
vertexai_credentials = None
|
||||
|
||||
if self.provider == "vertexai":
|
||||
from ..config import get_config
|
||||
|
||||
config = get_config()
|
||||
|
||||
vertexai_project_id = config.llm_vertexai_project_id
|
||||
if not vertexai_project_id:
|
||||
raise ValueError(
|
||||
"HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID is required for Vertex AI provider. "
|
||||
"Set it to your GCP project ID."
|
||||
)
|
||||
|
||||
vertexai_region = vertexai_region or "us-central1"
|
||||
service_account_key = vertexai_service_account_key
|
||||
vertexai_region = config.llm_vertexai_region or "us-central1"
|
||||
service_account_key = config.llm_vertexai_service_account_key
|
||||
|
||||
# Load explicit service account credentials if provided
|
||||
if service_account_key:
|
||||
@@ -708,20 +648,45 @@ class LLMProvider:
|
||||
f"model={self.model}, auth={'service_account' if service_account_key else 'ADC'}"
|
||||
)
|
||||
|
||||
# Normalize the Gemini service tier (pure: maps/validates the passed value,
|
||||
# no global config read). Non-Gemini providers never carry a tier. The
|
||||
# server-level default is resolved by the caller, like the other fields.
|
||||
if self.provider == "gemini":
|
||||
from ..config import parse_gemini_service_tier
|
||||
# For Gemini/VertexAI providers: read safety settings from global config if not explicitly provided
|
||||
# Use _get_raw_config() to bypass StaticConfigProxy (which blocks configurable fields),
|
||||
# since LLMProvider initialization legitimately needs the server-level default.
|
||||
if self.provider in ("gemini", "vertexai") and self.gemini_safety_settings is None:
|
||||
from ..config import _get_raw_config
|
||||
|
||||
self.gemini_service_tier = parse_gemini_service_tier(self.gemini_service_tier)
|
||||
else:
|
||||
self.gemini_service_tier = None
|
||||
try:
|
||||
raw_config = _get_raw_config()
|
||||
self.gemini_safety_settings = raw_config.llm_gemini_safety_settings
|
||||
except Exception:
|
||||
pass # Config may not be initialized in test environments
|
||||
|
||||
# gemini_safety_settings / prompt_cache_enabled / litellmrouter_config are
|
||||
# used as passed — the caller resolves the global-config fallback. Providers
|
||||
# that don't support prompt caching ignore the flag.
|
||||
# Prompt-prefix caching is a provider-agnostic toggle (default on): resolve
|
||||
# it from the static server config for every provider when the caller didn't
|
||||
# pass an explicit override. Providers that don't support caching ignore the
|
||||
# value; only those that implement get_or_create_cached_prefix act on it.
|
||||
if not self.prompt_cache_enabled:
|
||||
from ..config import DEFAULT_LLM_PROMPT_CACHE_ENABLED, _get_raw_config
|
||||
|
||||
try:
|
||||
raw_config = _get_raw_config()
|
||||
self.prompt_cache_enabled = bool(
|
||||
getattr(raw_config, "llm_prompt_cache_enabled", DEFAULT_LLM_PROMPT_CACHE_ENABLED)
|
||||
)
|
||||
except Exception:
|
||||
pass # Config may not be initialized in test environments
|
||||
|
||||
# For litellmrouter: prefer an explicit chain from the caller (per-op
|
||||
# construction in MemoryEngine threads the right chain through). If the caller
|
||||
# didn't supply one, fall back to the global ``llm_litellmrouter_config`` so
|
||||
# ad-hoc constructions (e.g. ``LLMProvider.from_env()``) keep working.
|
||||
router_config: dict[str, Any] | None = self.litellmrouter_config
|
||||
if self.provider == "litellmrouter" and router_config is None:
|
||||
from ..config import _get_raw_config
|
||||
|
||||
try:
|
||||
router_config = _get_raw_config().llm_litellmrouter_config
|
||||
except Exception:
|
||||
router_config = None
|
||||
|
||||
# Create provider implementation using factory
|
||||
self._provider_impl = create_llm_provider(
|
||||
@@ -733,7 +698,6 @@ class LLMProvider:
|
||||
groq_service_tier=self.groq_service_tier,
|
||||
openai_service_tier=self.openai_service_tier,
|
||||
bedrock_service_tier=self.bedrock_service_tier,
|
||||
gemini_service_tier=self.gemini_service_tier,
|
||||
extra_body=self.extra_body,
|
||||
default_headers=self.default_headers,
|
||||
vertexai_project_id=vertexai_project_id,
|
||||
@@ -742,7 +706,6 @@ class LLMProvider:
|
||||
gemini_safety_settings=self.gemini_safety_settings,
|
||||
prompt_cache_enabled=self.prompt_cache_enabled,
|
||||
litellmrouter_config=router_config,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
|
||||
# Backward compatibility: Keep mock provider properties
|
||||
@@ -799,9 +762,9 @@ class LLMProvider:
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int | None = None,
|
||||
initial_backoff: float | None = None,
|
||||
max_backoff: float | None = None,
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 60.0,
|
||||
skip_validation: bool = False,
|
||||
strict_schema: bool = False,
|
||||
return_usage: bool = False,
|
||||
@@ -816,12 +779,9 @@ class LLMProvider:
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts. ``None`` uses the provider's configured
|
||||
default (per-operation/global ``llm_max_retries``), else 10.
|
||||
initial_backoff: Initial backoff time in seconds. ``None`` uses the provider's
|
||||
configured default (``llm_initial_backoff``), else 1.0.
|
||||
max_backoff: Maximum backoff time in seconds. ``None`` uses the provider's
|
||||
configured default (``llm_max_backoff``), else 60.0.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
skip_validation: Return raw JSON without Pydantic validation.
|
||||
strict_schema: Per-call override requesting grammar-enforced (json_schema strict)
|
||||
structured output instead of the soft json_object path. The server-level
|
||||
@@ -846,20 +806,6 @@ class LLMProvider:
|
||||
structured = "+structured" if response_format is not None else ""
|
||||
set_stage(f"llm.{self.provider}.{scope}{structured}")
|
||||
|
||||
# Resolve the retry policy: explicit per-call arg wins, else the provider's
|
||||
# configured per-operation/global default, else this method's own fallback.
|
||||
max_retries = (
|
||||
max_retries if max_retries is not None else (self.max_retries if self.max_retries is not None else 10)
|
||||
)
|
||||
initial_backoff = (
|
||||
initial_backoff
|
||||
if initial_backoff is not None
|
||||
else (self.initial_backoff if self.initial_backoff is not None else 1.0)
|
||||
)
|
||||
max_backoff = (
|
||||
max_backoff if max_backoff is not None else (self.max_backoff if self.max_backoff is not None else 60.0)
|
||||
)
|
||||
|
||||
# Resolve strict-schema once, here, rather than in each provider: the
|
||||
# per-call argument OR the server-level HINDSIGHT_API_LLM_STRICT_SCHEMA
|
||||
# flag. Providers with a json_schema response_format (OpenAI-compatible,
|
||||
@@ -876,13 +822,7 @@ class LLMProvider:
|
||||
# The requested params are stashed in a contextvar (only what the caller
|
||||
# actually set) so the recorder can attach them to either path.
|
||||
from ..tracing import get_span_recorder
|
||||
from .llm_trace import (
|
||||
current_response_usage,
|
||||
reset_request_context,
|
||||
reset_response_usage,
|
||||
set_request_context,
|
||||
set_response_usage,
|
||||
)
|
||||
from .llm_trace import reset_request_context, set_request_context
|
||||
|
||||
call_start = time.monotonic()
|
||||
request_token = set_request_context(
|
||||
@@ -893,9 +833,6 @@ class LLMProvider:
|
||||
response_format=response_format,
|
||||
)
|
||||
)
|
||||
# Cleared per call; the provider stashes real usage once a response is in
|
||||
# hand so the error path below can attach it if parsing/validation fails.
|
||||
usage_token = set_response_usage(None)
|
||||
try:
|
||||
async with AsyncExitStack() as stack:
|
||||
for sem in _semaphores_for_scope(scope):
|
||||
@@ -923,19 +860,14 @@ class LLMProvider:
|
||||
**cache_kwarg,
|
||||
)
|
||||
except Exception as e:
|
||||
# The provider call may have succeeded (and incurred token
|
||||
# cost) before local parsing/validation raised; attach the
|
||||
# provider-reported usage to the error trace when available.
|
||||
usage = current_response_usage()
|
||||
get_span_recorder().record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
messages=messages,
|
||||
response_content=None,
|
||||
input_tokens=usage.input_tokens if usage else 0,
|
||||
output_tokens=usage.output_tokens if usage else 0,
|
||||
cached_tokens=usage.cached_tokens if usage else 0,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
duration=time.monotonic() - call_start,
|
||||
error=e,
|
||||
)
|
||||
@@ -951,7 +883,6 @@ class LLMProvider:
|
||||
self._mock_calls = self._provider_impl.get_mock_calls()
|
||||
finally:
|
||||
reset_request_context(request_token)
|
||||
reset_response_usage(usage_token)
|
||||
|
||||
return result
|
||||
|
||||
@@ -962,9 +893,9 @@ class LLMProvider:
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "tools",
|
||||
max_retries: int | None = None,
|
||||
initial_backoff: float | None = None,
|
||||
max_backoff: float | None = None,
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
cached_prefix: str | None = None,
|
||||
) -> "LLMToolCallResult":
|
||||
@@ -977,12 +908,9 @@ class LLMProvider:
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts. ``None`` uses the provider's configured
|
||||
default (per-operation/global ``llm_max_retries``), else 5.
|
||||
initial_backoff: Initial backoff time in seconds. ``None`` uses the provider's
|
||||
configured default (``llm_initial_backoff``), else 1.0.
|
||||
max_backoff: Maximum backoff time in seconds. ``None`` uses the provider's
|
||||
configured default (``llm_max_backoff``), else 30.0.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or {"type": "function", "function": {"name": "..."}}
|
||||
|
||||
Returns:
|
||||
@@ -992,29 +920,9 @@ class LLMProvider:
|
||||
|
||||
set_stage(f"llm.{self.provider}.{scope}+tools")
|
||||
|
||||
# Resolve the retry policy: explicit per-call arg wins, else the provider's
|
||||
# configured per-operation/global default, else this method's own fallback.
|
||||
max_retries = (
|
||||
max_retries if max_retries is not None else (self.max_retries if self.max_retries is not None else 5)
|
||||
)
|
||||
initial_backoff = (
|
||||
initial_backoff
|
||||
if initial_backoff is not None
|
||||
else (self.initial_backoff if self.initial_backoff is not None else 1.0)
|
||||
)
|
||||
max_backoff = (
|
||||
max_backoff if max_backoff is not None else (self.max_backoff if self.max_backoff is not None else 30.0)
|
||||
)
|
||||
|
||||
# Failures forwarded to the GenAI recorder; successes recorded by providers.
|
||||
from ..tracing import get_span_recorder
|
||||
from .llm_trace import (
|
||||
current_response_usage,
|
||||
reset_request_context,
|
||||
reset_response_usage,
|
||||
set_request_context,
|
||||
set_response_usage,
|
||||
)
|
||||
from .llm_trace import reset_request_context, set_request_context
|
||||
|
||||
call_start = time.monotonic()
|
||||
request_token = set_request_context(
|
||||
@@ -1025,9 +933,6 @@ class LLMProvider:
|
||||
tool_choice=tool_choice,
|
||||
)
|
||||
)
|
||||
# Cleared per call; the provider stashes real usage once a response is in
|
||||
# hand so the error path below can attach it if parsing/validation fails.
|
||||
usage_token = set_response_usage(None)
|
||||
try:
|
||||
async with AsyncExitStack() as stack:
|
||||
for sem in _semaphores_for_scope(scope):
|
||||
@@ -1052,19 +957,14 @@ class LLMProvider:
|
||||
**cache_kwarg,
|
||||
)
|
||||
except Exception as e:
|
||||
# The provider call may have succeeded (and incurred token
|
||||
# cost) before local parsing/validation raised; attach the
|
||||
# provider-reported usage to the error trace when available.
|
||||
usage = current_response_usage()
|
||||
get_span_recorder().record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
messages=messages,
|
||||
response_content=None,
|
||||
input_tokens=usage.input_tokens if usage else 0,
|
||||
output_tokens=usage.output_tokens if usage else 0,
|
||||
cached_tokens=usage.cached_tokens if usage else 0,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
duration=time.monotonic() - call_start,
|
||||
error=e,
|
||||
)
|
||||
@@ -1080,7 +980,6 @@ class LLMProvider:
|
||||
self._mock_calls = self._provider_impl.get_mock_calls()
|
||||
finally:
|
||||
reset_request_context(request_token)
|
||||
reset_response_usage(usage_token)
|
||||
|
||||
return result
|
||||
|
||||
@@ -1124,9 +1023,7 @@ class LLMProvider:
|
||||
|
||||
def _load_codex_auth(self) -> tuple[str, str]:
|
||||
"""
|
||||
Load OAuth credentials from the Codex ``auth.json``.
|
||||
|
||||
Honors ``CODEX_HOME`` (falling back to ``~/.codex``).
|
||||
Load OAuth credentials from ~/.codex/auth.json.
|
||||
|
||||
Returns:
|
||||
Tuple of (access_token, account_id).
|
||||
@@ -1135,9 +1032,7 @@ class LLMProvider:
|
||||
FileNotFoundError: If auth file doesn't exist.
|
||||
ValueError: If auth file is invalid.
|
||||
"""
|
||||
from .providers.codex_auth import default_codex_auth_file
|
||||
|
||||
auth_file = default_codex_auth_file()
|
||||
auth_file = Path.home() / ".codex" / "auth.json"
|
||||
|
||||
if not auth_file.exists():
|
||||
raise FileNotFoundError(
|
||||
@@ -1239,38 +1134,18 @@ class LLMProvider:
|
||||
@classmethod
|
||||
def from_env(cls) -> "LLMProvider":
|
||||
"""Create provider from environment variables using config.py constants."""
|
||||
# Read every field straight from the environment. The constructor no longer
|
||||
# resolves global-config fallbacks, so this factory must supply them — and it
|
||||
# does so without building the full HindsightConfig, keeping from_env() a
|
||||
# lightweight env-only loader (see test_llm_provider_from_env_keeps_lightweight_loader).
|
||||
from ..config import (
|
||||
DEFAULT_LLM_GROQ_SERVICE_TIER,
|
||||
DEFAULT_LLM_OPENAI_SERVICE_TIER,
|
||||
DEFAULT_LLM_PROMPT_CACHE_ENABLED,
|
||||
DEFAULT_LLM_PROVIDER,
|
||||
DEFAULT_LLM_REASONING_EFFORT,
|
||||
DEFAULT_LLM_TIMEOUT,
|
||||
ENV_LLM_API_KEY,
|
||||
ENV_LLM_BASE_URL,
|
||||
ENV_LLM_BEDROCK_SERVICE_TIER,
|
||||
ENV_LLM_DEFAULT_HEADERS,
|
||||
ENV_LLM_EXTRA_BODY,
|
||||
ENV_LLM_GEMINI_SAFETY_SETTINGS,
|
||||
ENV_LLM_GEMINI_SERVICE_TIER,
|
||||
ENV_LLM_GROQ_SERVICE_TIER,
|
||||
ENV_LLM_LITELLMROUTER_CONFIG,
|
||||
ENV_LLM_MODEL,
|
||||
ENV_LLM_OPENAI_SERVICE_TIER,
|
||||
ENV_LLM_PROMPT_CACHE_ENABLED,
|
||||
ENV_LLM_PROVIDER,
|
||||
ENV_LLM_REASONING_EFFORT,
|
||||
ENV_LLM_TIMEOUT,
|
||||
ENV_LLM_VERTEXAI_PROJECT_ID,
|
||||
ENV_LLM_VERTEXAI_REGION,
|
||||
ENV_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY,
|
||||
_get_default_model_for_provider,
|
||||
_parse_llm_router_config,
|
||||
parse_gemini_service_tier,
|
||||
)
|
||||
|
||||
provider = os.getenv(ENV_LLM_PROVIDER, DEFAULT_LLM_PROVIDER)
|
||||
@@ -1287,14 +1162,6 @@ class LLMProvider:
|
||||
model = os.getenv(ENV_LLM_MODEL) or _get_default_model_for_provider(provider)
|
||||
extra_body = json.loads(os.getenv(ENV_LLM_EXTRA_BODY, "null"))
|
||||
default_headers = json.loads(os.getenv(ENV_LLM_DEFAULT_HEADERS, "null"))
|
||||
prompt_cache_enabled = os.getenv(
|
||||
ENV_LLM_PROMPT_CACHE_ENABLED, str(DEFAULT_LLM_PROMPT_CACHE_ENABLED)
|
||||
).lower() in (
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
"on",
|
||||
)
|
||||
|
||||
return cls(
|
||||
provider=provider,
|
||||
@@ -1304,21 +1171,7 @@ class LLMProvider:
|
||||
reasoning_effort=os.getenv(ENV_LLM_REASONING_EFFORT, DEFAULT_LLM_REASONING_EFFORT),
|
||||
extra_body=extra_body,
|
||||
default_headers=default_headers,
|
||||
groq_service_tier=os.getenv(ENV_LLM_GROQ_SERVICE_TIER, DEFAULT_LLM_GROQ_SERVICE_TIER),
|
||||
openai_service_tier=os.getenv(ENV_LLM_OPENAI_SERVICE_TIER, DEFAULT_LLM_OPENAI_SERVICE_TIER),
|
||||
bedrock_service_tier=os.getenv(ENV_LLM_BEDROCK_SERVICE_TIER) or None,
|
||||
gemini_service_tier=(
|
||||
parse_gemini_service_tier(os.getenv(ENV_LLM_GEMINI_SERVICE_TIER))
|
||||
if provider.lower() == "gemini"
|
||||
else None
|
||||
),
|
||||
gemini_safety_settings=json.loads(os.getenv(ENV_LLM_GEMINI_SAFETY_SETTINGS, "null")),
|
||||
prompt_cache_enabled=prompt_cache_enabled,
|
||||
litellmrouter_config=_parse_llm_router_config(ENV_LLM_LITELLMROUTER_CONFIG),
|
||||
vertexai_project_id=os.getenv(ENV_LLM_VERTEXAI_PROJECT_ID) or None,
|
||||
vertexai_region=os.getenv(ENV_LLM_VERTEXAI_REGION) or None,
|
||||
vertexai_service_account_key=os.getenv(ENV_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY) or None,
|
||||
timeout=float(os.getenv(ENV_LLM_TIMEOUT, str(DEFAULT_LLM_TIMEOUT))),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -11,12 +11,6 @@ from one place, so we don't spawn a separate ``asyncio`` task per concern:
|
||||
consolidation operation failed terminally and left them with
|
||||
``consolidated_at IS NULL AND consolidation_failed_at IS NULL`` and nothing to
|
||||
re-trigger them.
|
||||
- **Scheduled mental model refresh** (configurable check cadence, default 60s):
|
||||
refresh mental models whose ``trigger.refresh_cron`` schedule is due, but only
|
||||
when the model is stale (new memories in its scope since its last refresh), so
|
||||
a scheduled tick never burns an LLM call to regenerate identical content. The
|
||||
per-model schedule lives in the cron expression; this loop only decides when to
|
||||
*check*.
|
||||
|
||||
The loop wakes on a short fixed tick and runs each job when its own
|
||||
``last_run + interval`` is due (run-at-start, then on interval), so adding jobs
|
||||
@@ -31,14 +25,12 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Coroutine
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..config import HindsightConfig, get_config
|
||||
from ..models import RequestContext
|
||||
from .db_utils import acquire_with_retry
|
||||
from .schema import _is_oracle, fq_table
|
||||
from .schema import _is_oracle
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .memory_engine import MemoryEngine
|
||||
@@ -99,8 +91,7 @@ class MaintenanceLoop:
|
||||
reconcile_on = cfg.consolidation_reconcile_interval_seconds > 0
|
||||
audit_on = cfg.audit_log_enabled and cfg.audit_log_retention_days > 0
|
||||
llm_on = cfg.llm_trace_enabled and cfg.llm_trace_retention_days > 0
|
||||
mm_refresh_on = cfg.mental_model_refresh_tick_seconds > 0
|
||||
return reconcile_on or audit_on or llm_on or mm_refresh_on
|
||||
return reconcile_on or audit_on or llm_on
|
||||
|
||||
# ── loop ───────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -127,25 +118,10 @@ class MaintenanceLoop:
|
||||
async def _tick(self) -> None:
|
||||
cfg = get_config()
|
||||
if self._is_due("retention", _RETENTION_INTERVAL_SECONDS):
|
||||
await self._run_timed("retention", self._run_retention(cfg))
|
||||
await self._run_retention(cfg)
|
||||
interval = cfg.consolidation_reconcile_interval_seconds
|
||||
if interval > 0 and self._is_due("reconcile", interval):
|
||||
await self._run_timed("consolidation reconcile", self._run_reconcile())
|
||||
mm_interval = cfg.mental_model_refresh_tick_seconds
|
||||
if mm_interval > 0 and self._is_due("mm_refresh", mm_interval):
|
||||
await self._run_timed("scheduled mental model refresh", self._run_scheduled_mm_refresh())
|
||||
|
||||
async def _run_timed(self, name: str, coro: Coroutine[Any, Any, None]) -> None:
|
||||
"""Run a maintenance job and emit one timing line for it.
|
||||
|
||||
Each job keeps its own summary log (counts of work done); this adds a
|
||||
single, uniform line per run so the cost of every sweep is observable.
|
||||
"""
|
||||
start = time.monotonic()
|
||||
try:
|
||||
await coro
|
||||
finally:
|
||||
logger.info(f"Maintenance: {name} took {time.monotonic() - start:.3f}s")
|
||||
await self._run_reconcile()
|
||||
|
||||
# ── retention ──────────────────────────────────────────────────────────
|
||||
|
||||
@@ -236,112 +212,3 @@ class MaintenanceLoop:
|
||||
f"Consolidation reconcile: scheduled {submitted} bank(s)"
|
||||
+ (f", skipped {skipped_unknown} in unrecognized schema(s)" if skipped_unknown else "")
|
||||
)
|
||||
|
||||
# ── scheduled mental model refresh ───────────────────────────────────────
|
||||
|
||||
async def _run_scheduled_mm_refresh(self) -> None:
|
||||
"""Refresh mental models whose ``trigger.refresh_cron`` is due.
|
||||
|
||||
Discovery (the set of cron-scheduled models, minus any with an in-flight
|
||||
refresh) is one cross-tenant round-trip via
|
||||
``public.mental_models_with_cron()``. Cron *due-ness* is evaluated here in
|
||||
Python — a scheduled fire has elapsed when the most recent cron boundary at
|
||||
or before now is later than ``last_refreshed_at`` — because cron arithmetic
|
||||
isn't expressible in plain SQL. Each due model is refreshed only when it is
|
||||
actually stale, so a schedule that fires while nothing changed costs a
|
||||
cheap staleness query, not an LLM call.
|
||||
"""
|
||||
engine = self._engine
|
||||
try:
|
||||
async with acquire_with_retry(engine._backend, max_retries=1) as conn:
|
||||
rows = await conn.fetch(
|
||||
"SELECT schema_name, bank_id, mental_model_id, refresh_cron, last_refreshed_at "
|
||||
"FROM public.mental_models_with_cron()"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Scheduled mental model refresh discovery failed: {e}")
|
||||
return
|
||||
if not rows:
|
||||
return
|
||||
|
||||
from croniter import croniter
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
due = []
|
||||
for row in rows:
|
||||
cron = row["refresh_cron"]
|
||||
last = row["last_refreshed_at"]
|
||||
try:
|
||||
prev_fire = croniter(cron, now).get_prev(datetime)
|
||||
except (ValueError, KeyError) as e:
|
||||
logger.warning(
|
||||
f"Scheduled mental model refresh: skipping invalid cron {cron!r} for "
|
||||
f"{row['schema_name']}/{row['mental_model_id']}: {e}"
|
||||
)
|
||||
continue
|
||||
if last is None or prev_fire > last:
|
||||
due.append(row)
|
||||
if not due:
|
||||
return
|
||||
|
||||
# Only enqueue into schemas the worker actually polls (tenant discovery),
|
||||
# otherwise the op would never be claimed. The tenant_id (when provided)
|
||||
# lets config resolution honor tenant-level overrides.
|
||||
try:
|
||||
tenants = await engine._tenant_extension.list_tenants()
|
||||
except Exception as e:
|
||||
logger.warning(f"Scheduled mental model refresh tenant discovery failed: {e}")
|
||||
return
|
||||
tenant_by_schema = {t.schema: t for t in tenants}
|
||||
default_schema = get_config().database_schema
|
||||
|
||||
from .memory_engine import _current_schema
|
||||
|
||||
submitted = 0
|
||||
skipped_unknown = 0
|
||||
skipped_fresh = 0
|
||||
for row in due:
|
||||
schema = row["schema_name"]
|
||||
bank_id = row["bank_id"]
|
||||
mm_id = row["mental_model_id"]
|
||||
tenant = tenant_by_schema.get(schema)
|
||||
if tenant is None and schema != default_schema:
|
||||
skipped_unknown += 1
|
||||
continue
|
||||
tenant_id = tenant.tenant_id if tenant else None
|
||||
token = _current_schema.set(schema)
|
||||
try:
|
||||
context = RequestContext(internal=True, tenant_id=tenant_id)
|
||||
# Skip if nothing in the model's scope changed since its last
|
||||
# refresh — a scheduled refresh must not regenerate identical
|
||||
# content. compute_mental_model_is_stale needs the model's tags +
|
||||
# trigger, which the discovery routine doesn't return, so re-read
|
||||
# the row under the bank's schema context.
|
||||
async with acquire_with_retry(engine._backend, max_retries=1) as conn:
|
||||
mm_row = await conn.fetchrow(
|
||||
f"SELECT id, tags, trigger, last_refreshed_at FROM {fq_table('mental_models')} "
|
||||
"WHERE bank_id = $1 AND id = $2",
|
||||
bank_id,
|
||||
mm_id,
|
||||
)
|
||||
if mm_row is None:
|
||||
continue
|
||||
is_stale = await engine.compute_mental_model_is_stale(conn, bank_id, mm_row)
|
||||
if not is_stale:
|
||||
skipped_fresh += 1
|
||||
continue
|
||||
await engine.submit_async_refresh_mental_model(
|
||||
bank_id=bank_id, mental_model_id=mm_id, request_context=context
|
||||
)
|
||||
submitted += 1
|
||||
except Exception as e:
|
||||
logger.warning(f"Scheduled mental model refresh failed for {mm_id} in {schema}: {e}")
|
||||
finally:
|
||||
_current_schema.reset(token)
|
||||
|
||||
if submitted or skipped_unknown or skipped_fresh:
|
||||
logger.info(
|
||||
f"Scheduled mental model refresh: scheduled {submitted} model(s)"
|
||||
+ (f", {skipped_fresh} up-to-date" if skipped_fresh else "")
|
||||
+ (f", skipped {skipped_unknown} in unrecognized schema(s)" if skipped_unknown else "")
|
||||
)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,204 +0,0 @@
|
||||
"""Multi-LLM routing: failover and (weighted) round-robin across N providers.
|
||||
|
||||
``MultiLLMProvider`` wraps an ordered list of :class:`LLMProvider` members and a
|
||||
:class:`~hindsight_api.config.LLMStrategyConfig`, exposing the same public surface
|
||||
as a single ``LLMProvider`` so it drops into every existing call path (including
|
||||
``with_config()`` / ``ConfiguredLLMProvider``).
|
||||
|
||||
Member 0 is the **primary** (the operation's unindexed/base LLM); members 1..N are
|
||||
the indexed extras (``HINDSIGHT_API_<OP>LLM_<n>_*``). Each member keeps its own
|
||||
internal retry budget, so we only advance to the next member after a member has
|
||||
exhausted its retries and raised.
|
||||
|
||||
Strategies:
|
||||
- ``failover``: try members in declared order ``[0..N]``.
|
||||
- ``round-robin``: rotate the starting member per request (optionally weighted),
|
||||
then fall through the remaining members on error.
|
||||
|
||||
Batch retain and any direct ``_provider_impl`` access operate on the **primary
|
||||
member only** (via attribute passthrough) — failover/round-robin apply to the
|
||||
interactive ``call`` / ``call_with_tools`` paths.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from ..config import LLM_STRATEGY_FAILOVER, LLMStrategyConfig
|
||||
from .llm_wrapper import LLMProvider, OutputTooLongError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .llm_wrapper import ConfiguredLLMProvider, LLMToolCallResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _should_failover(exc: BaseException) -> bool:
|
||||
"""Whether ``exc`` from one member should trigger a try on the next member.
|
||||
|
||||
Generic ``Exception`` instances (network errors, provider 5xx, timeouts after
|
||||
a member's own retries) fail over. ``OutputTooLongError`` is propagated — a
|
||||
different provider won't fit an over-length output either. ``CancelledError``,
|
||||
``KeyboardInterrupt`` and ``SystemExit`` are ``BaseException`` (not
|
||||
``Exception``) and therefore propagate unchanged.
|
||||
"""
|
||||
if isinstance(exc, OutputTooLongError):
|
||||
return False
|
||||
return isinstance(exc, Exception)
|
||||
|
||||
|
||||
class _WeightedRoundRobin:
|
||||
"""Smooth weighted round-robin scheduler (nginx SWRR).
|
||||
|
||||
Produces a starting member index per request such that, over time, member
|
||||
``i`` is chosen in proportion to ``weights[i]`` while keeping selections
|
||||
interleaved rather than bursty. Uniform weights degrade to plain round-robin.
|
||||
The tiny selection critical section is mutex-guarded so concurrent callers
|
||||
don't corrupt the running totals (they may still interleave, which only
|
||||
affects distribution, never correctness).
|
||||
"""
|
||||
|
||||
def __init__(self, weights: list[int]) -> None:
|
||||
self._weights = list(weights)
|
||||
self._current = [0] * len(weights)
|
||||
self._total = sum(weights)
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def next(self) -> int:
|
||||
with self._lock:
|
||||
best = 0
|
||||
for i, w in enumerate(self._weights):
|
||||
self._current[i] += w
|
||||
if self._current[i] > self._current[best]:
|
||||
best = i
|
||||
self._current[best] -= self._total
|
||||
return best
|
||||
|
||||
|
||||
class MultiLLMProvider:
|
||||
"""Route LLM calls across multiple members per a failover / round-robin strategy."""
|
||||
|
||||
def __init__(self, members: list[LLMProvider], strategy: LLMStrategyConfig) -> None:
|
||||
if not members:
|
||||
raise ValueError("MultiLLMProvider requires at least one member")
|
||||
self._members = members
|
||||
self._strategy = strategy
|
||||
|
||||
weights = strategy.weights or [1] * len(members)
|
||||
if len(weights) != len(members):
|
||||
raise ValueError(
|
||||
f"LLM strategy 'weights' has {len(weights)} entries but the chain has "
|
||||
f"{len(members)} members (primary + indexed); they must match."
|
||||
)
|
||||
self._scheduler = _WeightedRoundRobin(weights)
|
||||
|
||||
# ── routing ────────────────────────────────────────────────────────────────
|
||||
|
||||
def _member_order(self) -> list[int]:
|
||||
"""Indices to try, in order, for one request."""
|
||||
n = len(self._members)
|
||||
if self._strategy.mode == LLM_STRATEGY_FAILOVER:
|
||||
return list(range(n))
|
||||
start = self._scheduler.next()
|
||||
return [(start + i) % n for i in range(n)]
|
||||
|
||||
async def _dispatch(self, method_name: str, **kwargs: Any) -> Any:
|
||||
last_exc: BaseException | None = None
|
||||
order = self._member_order()
|
||||
for position, idx in enumerate(order):
|
||||
member = self._members[idx]
|
||||
try:
|
||||
return await getattr(member, method_name)(**kwargs)
|
||||
except BaseException as e: # noqa: BLE001 - re-raised unless it should fail over
|
||||
if not _should_failover(e):
|
||||
raise
|
||||
last_exc = e
|
||||
remaining = len(order) - position - 1
|
||||
logger.warning(
|
||||
"LLM member %d (%s/%s) failed on %s: %s%s",
|
||||
idx,
|
||||
member.provider,
|
||||
member.model,
|
||||
method_name,
|
||||
e,
|
||||
f"; trying next member ({remaining} left)" if remaining else "; no members left",
|
||||
)
|
||||
# All members failed; surface the last error (loop ran at least once).
|
||||
assert last_exc is not None
|
||||
raise last_exc
|
||||
|
||||
async def call(self, messages: list[dict[str, Any]], **kwargs: Any) -> Any:
|
||||
return await self._dispatch("call", messages=messages, **kwargs)
|
||||
|
||||
async def call_with_tools(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
**kwargs: Any,
|
||||
) -> "LLMToolCallResult":
|
||||
return await self._dispatch("call_with_tools", messages=messages, tools=tools, **kwargs)
|
||||
|
||||
# ── lifecycle ────────────────────────────────────────────────────────────────
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""Strictly verify the primary; soft-verify the rest (warn, don't fail).
|
||||
|
||||
A failover member being unreachable at startup must not block the server —
|
||||
it may come back before it's needed. The primary is the steady-state path,
|
||||
so its failure is still surfaced (the caller already wraps this in a
|
||||
warn-only try/except at startup).
|
||||
"""
|
||||
await self._members[0].verify_connection()
|
||||
for member in self._members[1:]:
|
||||
try:
|
||||
await member.verify_connection()
|
||||
except Exception as e: # noqa: BLE001 - soft verification
|
||||
logger.warning(
|
||||
"Failover LLM member %s/%s failed connection verification: %s. "
|
||||
"It will be tried at request time if the primary fails.",
|
||||
member.provider,
|
||||
member.model,
|
||||
e,
|
||||
)
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
for member in self._members:
|
||||
await member.cleanup()
|
||||
|
||||
def with_config(
|
||||
self,
|
||||
config: Any,
|
||||
*,
|
||||
bank_id: str | None = None,
|
||||
operation: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> "ConfiguredLLMProvider":
|
||||
"""Mirror ``LLMProvider.with_config`` so the strategy runs inside the
|
||||
per-operation configured wrapper (gemini-safety + trace contextvars wrap
|
||||
every member call)."""
|
||||
from .llm_trace import LLMTraceContext
|
||||
from .llm_wrapper import ConfiguredLLMProvider
|
||||
|
||||
trace_ctx = None
|
||||
if bank_id is not None or operation is not None or metadata:
|
||||
trace_ctx = LLMTraceContext(
|
||||
bank_id=bank_id,
|
||||
operation=operation,
|
||||
metadata=dict(metadata or {}),
|
||||
trace_id=str(uuid.uuid4()),
|
||||
operation_span_id=str(uuid.uuid4()),
|
||||
)
|
||||
return ConfiguredLLMProvider(self, config.llm_gemini_safety_settings, trace_ctx)
|
||||
|
||||
# ── attribute passthrough ────────────────────────────────────────────────────
|
||||
|
||||
@property
|
||||
def members(self) -> list[LLMProvider]:
|
||||
return self._members
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
# Anything not defined here (provider, model, api_key, base_url,
|
||||
# _provider_impl, mock helpers, batch helpers, ...) delegates to the
|
||||
# primary member so existing call sites keep working unchanged.
|
||||
return getattr(object.__getattribute__(self, "_members")[0], name)
|
||||
@@ -3,138 +3,43 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from hindsight_api.config import DEFAULT_FILE_PARSER_MARKITDOWN_OCR_PROMPT
|
||||
|
||||
from .base import FileParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from markitdown import StreamInfo
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Extensions whose markitdown converters decode the raw bytes as text. markitdown
|
||||
# samples only the first chunk for charset detection, so a UTF-8 file with a long
|
||||
# ASCII-only prefix is mis-detected as ASCII; the JSON/ipynb converter then crashes
|
||||
# decoding the first multibyte byte. Passing an explicit UTF-8 hint when the bytes
|
||||
# are valid UTF-8 sidesteps the faulty detection without affecting other encodings.
|
||||
_TEXT_EXTENSIONS = {
|
||||
".json",
|
||||
".jsonl",
|
||||
".ipynb",
|
||||
".txt",
|
||||
".text",
|
||||
".md",
|
||||
".markdown",
|
||||
".csv",
|
||||
".html",
|
||||
".htm",
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MarkitdownOcrOptions:
|
||||
"""OpenAI-compatible OCR options passed through to MarkItDown."""
|
||||
|
||||
# Keep this typed as object so the OpenAI SDK import stays lazy for non-OCR users.
|
||||
llm_client: object
|
||||
llm_model: str
|
||||
llm_prompt: str
|
||||
|
||||
|
||||
class MarkitdownParser(FileParser):
|
||||
"""
|
||||
Markitdown file parser.
|
||||
|
||||
Uses Microsoft's markitdown library to convert various file formats
|
||||
to markdown including PDF, Office docs, images with optional OCR,
|
||||
audio, HTML.
|
||||
to markdown including PDF, Office docs, images (via OCR), audio, HTML.
|
||||
|
||||
Supported formats:
|
||||
- PDF (.pdf)
|
||||
- Word (.docx, .doc)
|
||||
- PowerPoint (.pptx, .ppt)
|
||||
- Excel (.xlsx, .xls)
|
||||
- Images (.jpg, .jpeg, .png) - optional OCR
|
||||
- Images (.jpg, .jpeg, .png) - with OCR
|
||||
- HTML (.html, .htm)
|
||||
- Text (.txt, .md)
|
||||
- Audio (.mp3, .wav) - with transcription
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
ocr_enabled: bool = False,
|
||||
ocr_api_key: str | None = None,
|
||||
ocr_base_url: str | None = None,
|
||||
ocr_model: str | None = None,
|
||||
ocr_prompt: str | None = None,
|
||||
):
|
||||
def __init__(self):
|
||||
"""Initialize markitdown parser."""
|
||||
# Lazy import to avoid requiring markitdown for all users
|
||||
try:
|
||||
from markitdown import MarkItDown
|
||||
|
||||
self._markitdown = MarkItDown()
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"markitdown package is required for file parsing. Install with: pip install markitdown"
|
||||
) from e
|
||||
|
||||
self._ocr_enabled = ocr_enabled
|
||||
if ocr_enabled:
|
||||
ocr_options = self._build_ocr_options(
|
||||
api_key=ocr_api_key,
|
||||
base_url=ocr_base_url,
|
||||
model=ocr_model,
|
||||
prompt=ocr_prompt,
|
||||
)
|
||||
self._markitdown = MarkItDown(
|
||||
llm_client=ocr_options.llm_client,
|
||||
llm_model=ocr_options.llm_model,
|
||||
llm_prompt=ocr_options.llm_prompt,
|
||||
)
|
||||
else:
|
||||
self._markitdown = MarkItDown()
|
||||
|
||||
def _build_ocr_options(
|
||||
self,
|
||||
*,
|
||||
api_key: str | None,
|
||||
base_url: str | None,
|
||||
model: str | None,
|
||||
prompt: str | None,
|
||||
) -> MarkitdownOcrOptions:
|
||||
"""Build MarkItDown options for OpenAI-compatible image OCR."""
|
||||
if not model or not model.strip():
|
||||
raise ValueError(
|
||||
"Markitdown OCR is enabled but no model is configured. "
|
||||
"Set HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_MODEL to an OpenAI-compatible OCR/vision model "
|
||||
"with image-input support."
|
||||
)
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"Markitdown OCR is enabled but no API key is configured. "
|
||||
"Set HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_API_KEY."
|
||||
)
|
||||
if not base_url or not base_url.strip():
|
||||
raise ValueError(
|
||||
"Markitdown OCR is enabled but no base URL is configured. "
|
||||
"Set HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_BASE_URL to an OpenAI-compatible OCR/vision endpoint."
|
||||
)
|
||||
|
||||
try:
|
||||
from openai import OpenAI
|
||||
except ImportError as e:
|
||||
raise RuntimeError("openai package is required when Markitdown OCR is enabled.") from e
|
||||
|
||||
return MarkitdownOcrOptions(
|
||||
llm_client=OpenAI(api_key=api_key, base_url=base_url.strip()),
|
||||
llm_model=model.strip(),
|
||||
llm_prompt=prompt or DEFAULT_FILE_PARSER_MARKITDOWN_OCR_PROMPT,
|
||||
)
|
||||
|
||||
async def convert(self, file_data: bytes, filename: str) -> str:
|
||||
"""Parse file to markdown using markitdown."""
|
||||
# markitdown is synchronous, so we run it in executor to avoid blocking
|
||||
@@ -143,22 +48,14 @@ class MarkitdownParser(FileParser):
|
||||
|
||||
def _convert_sync(self, file_data: bytes, filename: str) -> str:
|
||||
"""Synchronous parsing (runs in thread pool)."""
|
||||
if self._is_image_file(filename) and not self._ocr_enabled:
|
||||
raise RuntimeError(
|
||||
"Image OCR is not enabled for the markitdown parser. "
|
||||
"Set HINDSIGHT_API_FILE_PARSER_MARKITDOWN_OCR_ENABLED=true and configure an OpenAI-compatible "
|
||||
"OCR/vision endpoint with image-input support, or choose an OCR-capable parser."
|
||||
)
|
||||
|
||||
# Write to temp file (markitdown requires file path)
|
||||
with tempfile.NamedTemporaryFile(suffix=Path(filename).suffix, delete=False) as tmp:
|
||||
tmp.write(file_data)
|
||||
tmp_path = tmp.name
|
||||
|
||||
try:
|
||||
# Parse using markitdown, passing an explicit charset hint for text
|
||||
# files to avoid markitdown's sample-based (and crash-prone) detection.
|
||||
result = self._markitdown.convert(tmp_path, stream_info=self._utf8_stream_info(file_data, filename))
|
||||
# Parse using markitdown
|
||||
result = self._markitdown.convert(tmp_path)
|
||||
|
||||
if not result or not result.text_content:
|
||||
raise RuntimeError(f"No content extracted from '{filename}'")
|
||||
@@ -176,28 +73,6 @@ class MarkitdownParser(FileParser):
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _utf8_stream_info(file_data: bytes, filename: str) -> "StreamInfo | None":
|
||||
"""Return a UTF-8 charset hint for text files that decode cleanly as UTF-8.
|
||||
|
||||
Returns None for binary files or non-UTF-8 text so markitdown falls back
|
||||
to its own detection.
|
||||
"""
|
||||
if Path(filename).suffix.lower() not in _TEXT_EXTENSIONS:
|
||||
return None
|
||||
try:
|
||||
file_data.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
return None
|
||||
from markitdown import StreamInfo
|
||||
|
||||
return StreamInfo(charset="utf-8")
|
||||
|
||||
@staticmethod
|
||||
def _is_image_file(filename: str) -> bool:
|
||||
"""Return whether the file type needs OCR to extract useful text."""
|
||||
return Path(filename).suffix.lower() in {".jpg", ".jpeg", ".png"}
|
||||
|
||||
def supports(self, filename: str, content_type: str | None = None) -> bool:
|
||||
"""Check if markitdown supports this file type."""
|
||||
# Supported extensions (from markitdown docs)
|
||||
@@ -210,7 +85,7 @@ class MarkitdownParser(FileParser):
|
||||
".ppt",
|
||||
".xlsx",
|
||||
".xls",
|
||||
# Images (optional OCR)
|
||||
# Images (with OCR)
|
||||
".jpg",
|
||||
".jpeg",
|
||||
".png",
|
||||
|
||||
@@ -15,25 +15,12 @@ import time
|
||||
from typing import Any
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _usage_from_anthropic_response(response: Any) -> LLMResponseUsage:
|
||||
"""Extract input/output/cached token counts from an Anthropic usage block."""
|
||||
usage = getattr(response, "usage", None)
|
||||
if not usage:
|
||||
return LLMResponseUsage()
|
||||
return LLMResponseUsage(
|
||||
input_tokens=usage.input_tokens or 0,
|
||||
output_tokens=usage.output_tokens or 0,
|
||||
cached_tokens=getattr(usage, "cache_read_input_tokens", 0) or 0,
|
||||
)
|
||||
|
||||
|
||||
class AnthropicLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using Anthropic's Claude models.
|
||||
@@ -149,9 +136,7 @@ class AnthropicLLM(LLMInterface):
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
skip_validation: Return raw JSON without Pydantic validation.
|
||||
strict_schema: Route structured output through a forced tool_use tool for
|
||||
native constrained decoding (issue #1002). When False, falls back to
|
||||
schema-in-prompt + JSON parse.
|
||||
strict_schema: Use strict JSON schema enforcement (not supported by Anthropic).
|
||||
return_usage: If True, return tuple (result, TokenUsage) instead of just result.
|
||||
|
||||
Returns:
|
||||
@@ -182,21 +167,14 @@ class AnthropicLLM(LLMInterface):
|
||||
else:
|
||||
anthropic_messages.append({"role": role, "content": content})
|
||||
|
||||
# Structured output: prefer Anthropic-native constrained decoding via a single
|
||||
# forced tool_use tool (strict_schema) over text-injecting the schema and
|
||||
# parsing the reply. Native constrained decoding guarantees schema-valid JSON,
|
||||
# eliminating the invalid-JSON retry storm (issue #1002). When strict_schema is
|
||||
# off we keep the text-inject + json.loads fallback for backward compatibility.
|
||||
schema = None
|
||||
use_forced_tool = False
|
||||
_tool_name = "structured_response"
|
||||
# Add JSON schema instruction if response_format is provided
|
||||
if response_format is not None and hasattr(response_format, "model_json_schema"):
|
||||
schema = response_format.model_json_schema()
|
||||
if strict_schema:
|
||||
use_forced_tool = True
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
if system_prompt:
|
||||
system_prompt += schema_msg
|
||||
else:
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
system_prompt = (system_prompt + schema_msg) if system_prompt else schema_msg
|
||||
system_prompt = schema_msg
|
||||
|
||||
# Prepare parameters
|
||||
call_params: dict[str, Any] = {
|
||||
@@ -208,14 +186,6 @@ class AnthropicLLM(LLMInterface):
|
||||
if system_prompt:
|
||||
call_params["system"] = system_prompt
|
||||
|
||||
if use_forced_tool:
|
||||
# Single tool whose input_schema IS the response schema; force the model to
|
||||
# emit it via tool_choice so the SDK does constrained decoding for us.
|
||||
call_params["tools"] = [
|
||||
{"name": _tool_name, "description": "Return the structured response.", "input_schema": schema}
|
||||
]
|
||||
call_params["tool_choice"] = {"type": "tool", "name": _tool_name}
|
||||
|
||||
if self._extra_body:
|
||||
call_params["extra_body"] = self._extra_body
|
||||
|
||||
@@ -224,61 +194,40 @@ class AnthropicLLM(LLMInterface):
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.messages.create(**call_params)
|
||||
# Stash usage before parse/validate, which may raise locally
|
||||
# even though the provider charged for these tokens (#2387).
|
||||
stash_response_usage(_usage_from_anthropic_response(response))
|
||||
|
||||
if use_forced_tool:
|
||||
# Forced tool_use → the validated args are already a dict; no parsing,
|
||||
# no markdown-strip, no JSON-decode retry possible.
|
||||
tool_input = None
|
||||
for block in response.content:
|
||||
if block.type == "tool_use" and block.name == _tool_name:
|
||||
tool_input = block.input or {}
|
||||
break
|
||||
if tool_input is None:
|
||||
# Model ignored the forced tool (rare, e.g. a gateway that drops
|
||||
# tool_choice). Fall back to text parse so we don't hard-fail; the
|
||||
# existing retry loop still covers genuine errors.
|
||||
content = "".join(b.text for b in response.content if b.type == "text")
|
||||
tool_input = json.loads(content)
|
||||
content = json.dumps(tool_input)
|
||||
result = tool_input if skip_validation else response_format.model_validate(tool_input)
|
||||
else:
|
||||
# Anthropic response content is a list of blocks
|
||||
content = ""
|
||||
for block in response.content:
|
||||
if block.type == "text":
|
||||
content += block.text
|
||||
# Anthropic response content is a list of blocks
|
||||
content = ""
|
||||
for block in response.content:
|
||||
if block.type == "text":
|
||||
content += block.text
|
||||
|
||||
if response_format is not None:
|
||||
# Models may wrap JSON in markdown code blocks
|
||||
clean_content = content
|
||||
if "```json" in content:
|
||||
clean_content = content.split("```json")[1].split("```")[0].strip()
|
||||
elif "```" in content:
|
||||
clean_content = content.split("```")[1].split("```")[0].strip()
|
||||
if response_format is not None:
|
||||
# Models may wrap JSON in markdown code blocks
|
||||
clean_content = content
|
||||
if "```json" in content:
|
||||
clean_content = content.split("```json")[1].split("```")[0].strip()
|
||||
elif "```" in content:
|
||||
clean_content = content.split("```")[1].split("```")[0].strip()
|
||||
|
||||
try:
|
||||
json_data = json.loads(clean_content)
|
||||
except json.JSONDecodeError:
|
||||
# Fallback to parsing raw content if markdown stripping failed
|
||||
json_data = json.loads(content)
|
||||
try:
|
||||
json_data = json.loads(clean_content)
|
||||
except json.JSONDecodeError:
|
||||
# Fallback to parsing raw content if markdown stripping failed
|
||||
json_data = json.loads(content)
|
||||
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
else:
|
||||
result = response_format.model_validate(json_data)
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
else:
|
||||
result = content
|
||||
result = response_format.model_validate(json_data)
|
||||
else:
|
||||
result = content
|
||||
|
||||
# Record metrics and log slow calls
|
||||
duration = time.time() - start_time
|
||||
response_usage = _usage_from_anthropic_response(response)
|
||||
input_tokens = response_usage.input_tokens
|
||||
output_tokens = response_usage.output_tokens
|
||||
input_tokens = response.usage.input_tokens or 0 if response.usage else 0
|
||||
output_tokens = response.usage.output_tokens or 0 if response.usage else 0
|
||||
total_tokens = input_tokens + output_tokens
|
||||
cached_tokens = response_usage.cached_tokens
|
||||
cached_tokens = getattr(response.usage, "cache_read_input_tokens", 0) or 0 if response.usage else 0
|
||||
|
||||
# Record LLM metrics
|
||||
metrics = get_metrics_collector()
|
||||
@@ -466,7 +415,6 @@ class AnthropicLLM(LLMInterface):
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.messages.create(**call_params)
|
||||
stash_response_usage(_usage_from_anthropic_response(response))
|
||||
|
||||
# Extract content and tool calls
|
||||
content_parts = []
|
||||
|
||||
@@ -16,7 +16,6 @@ from typing import Any
|
||||
from pydantic import ValidationError
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
|
||||
@@ -119,14 +118,12 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
Raises:
|
||||
RuntimeError: If the connection test fails.
|
||||
"""
|
||||
from ...config import get_config
|
||||
|
||||
try:
|
||||
test_messages = [{"role": "user", "content": "test"}]
|
||||
await self.call(
|
||||
messages=test_messages,
|
||||
max_completion_tokens=10,
|
||||
temperature=get_config().llm_temperature_verification,
|
||||
temperature=0.0,
|
||||
scope="verification",
|
||||
max_retries=0,
|
||||
)
|
||||
@@ -229,16 +226,6 @@ class ClaudeCodeLLM(LLMInterface):
|
||||
if isinstance(block, TextBlock):
|
||||
full_text += block.text
|
||||
|
||||
# The Claude Agent SDK doesn't report exact counts; stash the same
|
||||
# char/4 estimate the success path traces so a later parse/validate
|
||||
# failure records consistent (estimated) tokens, not zero (#2387).
|
||||
stash_response_usage(
|
||||
LLMResponseUsage(
|
||||
input_tokens=sum(len(m.get("content", "")) for m in messages) // 4,
|
||||
output_tokens=len(full_text) // 4,
|
||||
)
|
||||
)
|
||||
|
||||
# Handle structured output
|
||||
if response_format is not None:
|
||||
# Models may wrap JSON in markdown
|
||||
|
||||
@@ -60,22 +60,6 @@ _CODEX_TERMINAL_REFRESH_ERROR_CODES = frozenset(
|
||||
)
|
||||
|
||||
|
||||
def default_codex_auth_file() -> Path:
|
||||
"""Return the path to Codex's ``auth.json``.
|
||||
|
||||
Honors the ``CODEX_HOME`` environment variable — the same variable the
|
||||
canonical ``@openai/codex`` CLI uses to relocate its config/credentials
|
||||
directory — and falls back to ``~/.codex`` when it is unset or empty.
|
||||
|
||||
Resolved lazily on each call (rather than cached at import time) so that
|
||||
the environment is read at the point of use.
|
||||
"""
|
||||
codex_home = os.environ.get("CODEX_HOME")
|
||||
if codex_home:
|
||||
return Path(codex_home) / "auth.json"
|
||||
return Path.home() / ".codex" / "auth.json"
|
||||
|
||||
|
||||
class CodexRefreshExpiredError(RuntimeError):
|
||||
"""Raised when the Codex refresh_token itself is no longer valid.
|
||||
|
||||
@@ -102,7 +86,7 @@ class CodexAuthManager:
|
||||
The OAuth refresh token. May be ``None`` when the auth file omits it;
|
||||
the provider still works as a one-shot loader in that case.
|
||||
auth_file:
|
||||
Path to the Codex ``auth.json``. Used for re-reading the refresh token
|
||||
Path to ``~/.codex/auth.json``. Used for re-reading the refresh token
|
||||
on demand and for atomic persistence of rotated credentials.
|
||||
"""
|
||||
|
||||
@@ -131,8 +115,7 @@ class CodexAuthManager:
|
||||
Parameters
|
||||
----------
|
||||
auth_file:
|
||||
Defaults to ``$CODEX_HOME/auth.json`` (or ``~/.codex/auth.json``
|
||||
when ``CODEX_HOME`` is unset).
|
||||
Defaults to ``~/.codex/auth.json``.
|
||||
|
||||
Raises
|
||||
------
|
||||
@@ -143,7 +126,7 @@ class CodexAuthManager:
|
||||
``auth_mode``.
|
||||
"""
|
||||
if auth_file is None:
|
||||
auth_file = default_codex_auth_file()
|
||||
auth_file = Path.home() / ".codex" / "auth.json"
|
||||
|
||||
if not auth_file.exists():
|
||||
raise FileNotFoundError(f"Codex auth file not found: {auth_file}. Run 'codex auth login' to authenticate.")
|
||||
|
||||
@@ -2,9 +2,8 @@
|
||||
OpenAI Codex LLM provider using ChatGPT Plus/Pro OAuth authentication.
|
||||
|
||||
This provider enables using ChatGPT Plus/Pro subscriptions for API calls
|
||||
without separate OpenAI Platform API credits. It uses OAuth tokens from the
|
||||
Codex ``auth.json`` (``$CODEX_HOME/auth.json``, or ``~/.codex/auth.json`` when
|
||||
``CODEX_HOME`` is unset) and communicates with the ChatGPT backend API.
|
||||
without separate OpenAI Platform API credits. It uses OAuth tokens from
|
||||
~/.codex/auth.json and communicates with the ChatGPT backend API.
|
||||
|
||||
Tokens are refreshed automatically: the provider decodes the access_token
|
||||
JWT's ``exp`` claim and proactively refreshes via
|
||||
@@ -26,7 +25,6 @@ from typing import Any
|
||||
import httpx
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
|
||||
@@ -37,7 +35,6 @@ from .codex_auth import (
|
||||
_CODEX_TOKEN_REFRESH_SKEW_SECONDS,
|
||||
CodexAuthManager,
|
||||
CodexRefreshExpiredError,
|
||||
default_codex_auth_file,
|
||||
)
|
||||
|
||||
# Re-export for backward compatibility (tests import from this module).
|
||||
@@ -58,15 +55,14 @@ class CodexLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using OpenAI Codex OAuth authentication.
|
||||
|
||||
Authenticates using ChatGPT Plus/Pro credentials stored in the Codex
|
||||
``auth.json`` (honoring ``CODEX_HOME``, default ``~/.codex``) and makes API
|
||||
calls to chatgpt.com/backend-api/codex/responses.
|
||||
Authenticates using ChatGPT Plus/Pro credentials stored in ~/.codex/auth.json
|
||||
and makes API calls to chatgpt.com/backend-api/codex/responses.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str, # Will be ignored, reads from the Codex auth.json (CODEX_HOME or ~/.codex)
|
||||
api_key: str, # Will be ignored, reads from ~/.codex/auth.json
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
@@ -85,14 +81,12 @@ class CodexLLM(LLMInterface):
|
||||
refresh_token = self._load_codex_refresh_token()
|
||||
logger.info(f"Loaded Codex OAuth credentials for account: {account_id}")
|
||||
except Exception as e:
|
||||
auth_file = default_codex_auth_file()
|
||||
raise RuntimeError(
|
||||
f"Failed to load Codex OAuth credentials from {auth_file}: {e}\n\n"
|
||||
f"Failed to load Codex OAuth credentials from ~/.codex/auth.json: {e}\n\n"
|
||||
"To set up Codex authentication:\n"
|
||||
"1. Install Codex CLI: npm install -g @openai/codex\n"
|
||||
"2. Login: codex auth login\n"
|
||||
f"3. Verify: ls {auth_file}\n\n"
|
||||
"(Set CODEX_HOME to use a credentials directory other than ~/.codex.)\n\n"
|
||||
"3. Verify: ls ~/.codex/auth.json\n\n"
|
||||
"Or use a different provider (openai, anthropic, gemini) with API keys."
|
||||
) from e
|
||||
|
||||
@@ -100,7 +94,7 @@ class CodexLLM(LLMInterface):
|
||||
access_token=access_token,
|
||||
account_id=account_id,
|
||||
refresh_token=refresh_token,
|
||||
auth_file=default_codex_auth_file(),
|
||||
auth_file=Path.home() / ".codex" / "auth.json",
|
||||
)
|
||||
|
||||
# Use ChatGPT backend API endpoint. Codex auth is tied to
|
||||
@@ -162,7 +156,7 @@ class CodexLLM(LLMInterface):
|
||||
|
||||
def _load_codex_auth(self) -> tuple[str, str]:
|
||||
"""
|
||||
Load OAuth credentials from the Codex ``auth.json`` (CODEX_HOME or ~/.codex).
|
||||
Load OAuth credentials from ~/.codex/auth.json.
|
||||
|
||||
Returns:
|
||||
Tuple of (access_token, account_id).
|
||||
@@ -171,7 +165,7 @@ class CodexLLM(LLMInterface):
|
||||
FileNotFoundError: If auth file doesn't exist.
|
||||
ValueError: If auth file is invalid.
|
||||
"""
|
||||
auth_file = default_codex_auth_file()
|
||||
auth_file = Path.home() / ".codex" / "auth.json"
|
||||
|
||||
if not auth_file.exists():
|
||||
raise FileNotFoundError(
|
||||
@@ -203,7 +197,9 @@ class CodexLLM(LLMInterface):
|
||||
pre- and post-``__init__`` because it does not depend on
|
||||
``_auth_manager`` being constructed yet.
|
||||
"""
|
||||
auth_file = self._auth_manager._auth_file if hasattr(self, "_auth_manager") else default_codex_auth_file()
|
||||
auth_file = (
|
||||
self._auth_manager._auth_file if hasattr(self, "_auth_manager") else Path.home() / ".codex" / "auth.json"
|
||||
)
|
||||
return CodexAuthManager.load_refresh_token_from_file(auth_file)
|
||||
|
||||
@staticmethod
|
||||
@@ -415,16 +411,6 @@ class CodexLLM(LLMInterface):
|
||||
# Parse SSE stream
|
||||
content = await self._parse_sse_stream(response)
|
||||
|
||||
# Codex SSE carries no usage block; stash the same char/4 estimate
|
||||
# the success path traces so a later parse/validate failure records
|
||||
# consistent (estimated) token counts rather than zero (#2387).
|
||||
stash_response_usage(
|
||||
LLMResponseUsage(
|
||||
input_tokens=sum(len(m.get("content", "")) for m in messages) // 4,
|
||||
output_tokens=len(content) // 4,
|
||||
)
|
||||
)
|
||||
|
||||
# Handle structured output
|
||||
if response_format is not None:
|
||||
# Models may wrap JSON in markdown
|
||||
|
||||
@@ -20,7 +20,6 @@ from google.genai import errors as genai_errors
|
||||
from google.genai import types as genai_types
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.llm_wrapper import parse_llm_json
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
@@ -51,18 +50,6 @@ def _to_int(value: Any) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
def _usage_from_gemini_response(response: Any) -> LLMResponseUsage:
|
||||
"""Extract prompt/candidate/cached token counts from a Gemini usage_metadata block."""
|
||||
usage = getattr(response, "usage_metadata", None)
|
||||
if not usage:
|
||||
return LLMResponseUsage()
|
||||
return LLMResponseUsage(
|
||||
input_tokens=usage.prompt_token_count or 0,
|
||||
output_tokens=usage.candidates_token_count or 0,
|
||||
cached_tokens=getattr(usage, "cached_content_token_count", 0) or 0,
|
||||
)
|
||||
|
||||
|
||||
class GeminiLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider for Google Gemini and Vertex AI.
|
||||
@@ -89,7 +76,6 @@ class GeminiLLM(LLMInterface):
|
||||
|
||||
# Safety settings: None means use Gemini's defaults
|
||||
self._safety_settings: list | None = kwargs.get("gemini_safety_settings")
|
||||
self._service_tier: str | None = kwargs.get("gemini_service_tier")
|
||||
|
||||
# User-configured extra params merged into the GenerateContentConfig of
|
||||
# every call. Gemini's request body nests generation params, so we expose
|
||||
@@ -120,16 +106,6 @@ class GeminiLLM(LLMInterface):
|
||||
self._client = genai.Client(api_key=self.api_key)
|
||||
logger.info(f"Gemini API: model={self.model}")
|
||||
|
||||
def _apply_service_tier(self, config_kwargs: dict[str, Any]) -> None:
|
||||
if not self._service_tier:
|
||||
return
|
||||
|
||||
http_options = dict(config_kwargs.get("http_options") or {})
|
||||
extra_body = dict(http_options.get("extra_body") or {})
|
||||
extra_body.setdefault("service_tier", self._service_tier)
|
||||
http_options["extra_body"] = extra_body
|
||||
config_kwargs["http_options"] = http_options
|
||||
|
||||
def _init_vertexai(self, **kwargs: Any) -> None:
|
||||
"""Initialize Vertex AI client with project, region, and credentials."""
|
||||
# Extract Vertex AI config from kwargs
|
||||
@@ -271,13 +247,16 @@ class GeminiLLM(LLMInterface):
|
||||
else:
|
||||
gemini_contents.append(genai_types.Content(role="user", parts=[genai_types.Part(text=content)]))
|
||||
|
||||
def _system_instruction_with_schema() -> str:
|
||||
# Add the JSON schema as a textual hint in the system_instruction (matching
|
||||
# the normal uncached path). Structured output is still enforced via
|
||||
# response_schema regardless; this is just guidance text.
|
||||
if response_format is not None and hasattr(response_format, "model_json_schema"):
|
||||
schema = response_format.model_json_schema()
|
||||
schema_msg = (
|
||||
f"\n\nYou must respond with valid JSON matching this schema:\n"
|
||||
f"{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
)
|
||||
return (system_instruction + schema_msg) if system_instruction else schema_msg
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2, ensure_ascii=False)}"
|
||||
if system_instruction:
|
||||
system_instruction += schema_msg
|
||||
else:
|
||||
system_instruction = schema_msg
|
||||
|
||||
# Apply safety settings: context var (per-request bank override) takes precedence over instance default
|
||||
effective_safety_settings = _safety_settings_ctx.get()
|
||||
@@ -294,18 +273,11 @@ class GeminiLLM(LLMInterface):
|
||||
def _build_generation_config(use_cache: bool) -> "genai_types.GenerateContentConfig | None":
|
||||
# Seed with user-configured extra params; explicit settings below win.
|
||||
config_kwargs: dict[str, Any] = dict(self._extra_body)
|
||||
self._apply_service_tier(config_kwargs)
|
||||
if use_cache:
|
||||
config_kwargs["cached_content"] = cached_prefix
|
||||
elif (
|
||||
use_schema_prompt_fallback
|
||||
and response_format is not None
|
||||
and hasattr(response_format, "model_json_schema")
|
||||
):
|
||||
config_kwargs["system_instruction"] = _system_instruction_with_schema()
|
||||
elif system_instruction:
|
||||
config_kwargs["system_instruction"] = system_instruction
|
||||
if response_format is not None and not use_schema_prompt_fallback:
|
||||
if response_format is not None:
|
||||
config_kwargs["response_mime_type"] = "application/json"
|
||||
config_kwargs["response_schema"] = response_format
|
||||
if temperature is not None:
|
||||
@@ -323,7 +295,6 @@ class GeminiLLM(LLMInterface):
|
||||
return genai_types.GenerateContentConfig(**config_kwargs) if config_kwargs else None
|
||||
|
||||
cache_active = using_cache
|
||||
use_schema_prompt_fallback = False
|
||||
generation_config = _build_generation_config(cache_active)
|
||||
|
||||
last_exception = None
|
||||
@@ -340,9 +311,6 @@ class GeminiLLM(LLMInterface):
|
||||
),
|
||||
timeout=90.0, # Safety net for network hangs; valid slow responses are <90s
|
||||
)
|
||||
# Stash usage before parse/validate, which may raise locally
|
||||
# even though the provider charged for these tokens (#2387).
|
||||
stash_response_usage(_usage_from_gemini_response(response))
|
||||
|
||||
content = response.text
|
||||
|
||||
@@ -444,26 +412,12 @@ class GeminiLLM(LLMInterface):
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=input_tokens + output_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
return result, token_usage
|
||||
return result
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
last_exception = e
|
||||
if (
|
||||
attempt < max_retries
|
||||
and response_format is not None
|
||||
and hasattr(response_format, "model_json_schema")
|
||||
and not cache_active
|
||||
and not use_schema_prompt_fallback
|
||||
):
|
||||
logger.warning("Gemini returned invalid JSON, retrying with prompt-side schema guidance...")
|
||||
cache_active = False
|
||||
use_schema_prompt_fallback = True
|
||||
generation_config = _build_generation_config(cache_active)
|
||||
await asyncio.sleep(min(initial_backoff * (2**attempt), max_backoff))
|
||||
continue
|
||||
if attempt < max_retries:
|
||||
logger.warning("Gemini returned invalid JSON, retrying...")
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
@@ -650,7 +604,6 @@ class GeminiLLM(LLMInterface):
|
||||
def _build_tools_config(use_cache: bool) -> "genai_types.GenerateContentConfig":
|
||||
# Seed with user-configured extra params; explicit settings below win.
|
||||
config_kwargs: dict[str, Any] = dict(self._extra_body)
|
||||
self._apply_service_tier(config_kwargs)
|
||||
if use_cache:
|
||||
config_kwargs["cached_content"] = cached_prefix
|
||||
else:
|
||||
@@ -709,7 +662,6 @@ class GeminiLLM(LLMInterface):
|
||||
),
|
||||
timeout=90.0, # Safety net for network hangs; valid slow responses are <90s
|
||||
)
|
||||
stash_response_usage(_usage_from_gemini_response(response))
|
||||
|
||||
# Extract content and tool calls
|
||||
content = None
|
||||
@@ -797,8 +749,6 @@ class GeminiLLM(LLMInterface):
|
||||
finish_reason=finish_reason,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cached_tokens=cached_input_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
|
||||
except genai_errors.APIError as e:
|
||||
|
||||
@@ -15,15 +15,10 @@ is handled automatically by LiteLLM.
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from litellm.exceptions import Timeout as LiteLLMTimeout
|
||||
|
||||
from hindsight_api.config import DEFAULT_LLM_TIMEOUT, ENV_LLM_TIMEOUT
|
||||
from hindsight_api.engine.llm_interface import LLMInterface, OutputTooLongError
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
from hindsight_api.worker.stage import set_stage
|
||||
@@ -31,22 +26,6 @@ from hindsight_api.worker.stage import set_stage
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _usage_from_litellm_response(response: Any) -> LLMResponseUsage:
|
||||
"""Extract prompt/completion/cached token counts from a LiteLLM (OpenAI-shaped) usage block."""
|
||||
usage = getattr(response, "usage", None)
|
||||
if not usage:
|
||||
return LLMResponseUsage()
|
||||
cached_tokens = 0
|
||||
details = getattr(usage, "prompt_tokens_details", None)
|
||||
if details:
|
||||
cached_tokens = getattr(details, "cached_tokens", 0) or 0
|
||||
return LLMResponseUsage(
|
||||
input_tokens=getattr(usage, "prompt_tokens", 0) or 0,
|
||||
output_tokens=getattr(usage, "completion_tokens", 0) or 0,
|
||||
cached_tokens=cached_tokens,
|
||||
)
|
||||
|
||||
|
||||
class LiteLLMLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using the LiteLLM SDK for universal model support.
|
||||
@@ -68,16 +47,13 @@ class LiteLLMLLM(LLMInterface):
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
timeout: float | None = None,
|
||||
timeout: float = 300.0,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
bedrock_service_tier: str | None = None,
|
||||
default_headers: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(provider, api_key, base_url, model, reasoning_effort, **kwargs)
|
||||
# ``None`` falls back to HINDSIGHT_API_LLM_TIMEOUT, then DEFAULT_LLM_TIMEOUT — never None,
|
||||
# so the hard ``asyncio.wait_for`` backstop in ``call`` is always bounded.
|
||||
self.timeout = timeout if timeout is not None else float(os.getenv(ENV_LLM_TIMEOUT, str(DEFAULT_LLM_TIMEOUT)))
|
||||
self.timeout = timeout
|
||||
self._litellm: Any = None
|
||||
# User-configured extra params merged as top-level kwargs into every
|
||||
# completion call so LiteLLM normalizes them per-provider (e.g. maps
|
||||
@@ -85,13 +61,6 @@ class LiteLLMLLM(LLMInterface):
|
||||
# drops any the target model rejects (litellm.drop_params=True below).
|
||||
# Sourced from llm_extra_body (env: HINDSIGHT_API_LLM_EXTRA_BODY).
|
||||
self._extra_body: dict[str, Any] = extra_body or {}
|
||||
# Operator-configured default headers forwarded to litellm.acompletion as
|
||||
# ``extra_headers`` (used by deployments routing through proxies / request-
|
||||
# tracing middleware). Mirrors the Anthropic provider's default_headers
|
||||
# wiring. Sourced from llm_default_headers (env: HINDSIGHT_API_LLM_DEFAULT_HEADERS).
|
||||
# Copied so a caller-owned dict can't be mutated through us, and a fresh
|
||||
# copy is handed to each call below to avoid cross-request contamination.
|
||||
self._default_headers: dict[str, Any] = dict(default_headers or {})
|
||||
self.bedrock_service_tier = bedrock_service_tier
|
||||
|
||||
try:
|
||||
@@ -108,14 +77,12 @@ class LiteLLMLLM(LLMInterface):
|
||||
raise RuntimeError("LiteLLM SDK not installed. Run: uv add litellm or pip install litellm") from e
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
from ...config import get_config
|
||||
|
||||
try:
|
||||
test_messages = [{"role": "user", "content": "test"}]
|
||||
await self.call(
|
||||
messages=test_messages,
|
||||
max_completion_tokens=50,
|
||||
temperature=get_config().llm_temperature_verification,
|
||||
temperature=0.0,
|
||||
scope="verification",
|
||||
max_retries=0,
|
||||
)
|
||||
@@ -154,13 +121,6 @@ class LiteLLMLLM(LLMInterface):
|
||||
for key, value in self._extra_body.items():
|
||||
kwargs.setdefault(key, value)
|
||||
|
||||
# Forward operator-configured default headers as ``extra_headers`` so they
|
||||
# reach the provider behind LiteLLM (proxies / request-tracing middleware).
|
||||
# ``setdefault`` keeps any explicit per-call ``extra_headers`` authoritative;
|
||||
# a per-call copy prevents LiteLLM/downstream from mutating the stored dict.
|
||||
if self._default_headers:
|
||||
kwargs.setdefault("extra_headers", dict(self._default_headers))
|
||||
|
||||
# Bedrock service tier: flex (50% cheaper), priority, or reserved
|
||||
if self.model.startswith("bedrock/") and self.bedrock_service_tier is not None:
|
||||
kwargs["service_tier"] = self.bedrock_service_tier
|
||||
@@ -249,14 +209,7 @@ class LiteLLMLLM(LLMInterface):
|
||||
if attempt > 0:
|
||||
set_stage(f"llm.{self._stage_label}.{scope}.attempt={attempt + 1}/{max_retries + 1}")
|
||||
try:
|
||||
response = await asyncio.wait_for(
|
||||
self._acompletion(**call_kwargs),
|
||||
timeout=self.timeout,
|
||||
)
|
||||
# Stash usage before the length check and parse/validate below,
|
||||
# which may raise locally even though the provider charged for
|
||||
# these tokens (#2387).
|
||||
stash_response_usage(_usage_from_litellm_response(response))
|
||||
response = await self._acompletion(**call_kwargs)
|
||||
|
||||
content = response.choices[0].message.content or ""
|
||||
finish_reason = response.choices[0].finish_reason
|
||||
@@ -287,9 +240,8 @@ class LiteLLMLLM(LLMInterface):
|
||||
result = content
|
||||
|
||||
# Extract usage
|
||||
response_usage = _usage_from_litellm_response(response)
|
||||
input_tokens = response_usage.input_tokens
|
||||
output_tokens = response_usage.output_tokens
|
||||
input_tokens = getattr(response.usage, "prompt_tokens", 0) or 0
|
||||
output_tokens = getattr(response.usage, "completion_tokens", 0) or 0
|
||||
total_tokens = input_tokens + output_tokens
|
||||
|
||||
# Record metrics
|
||||
@@ -352,25 +304,6 @@ class LiteLLMLLM(LLMInterface):
|
||||
logger.error(f"LiteLLM returned invalid JSON after {max_retries + 1} attempts")
|
||||
raise
|
||||
|
||||
except (TimeoutError, asyncio.TimeoutError, LiteLLMTimeout) as e:
|
||||
# litellm/httpx don't always honor their own ``timeout=`` (e.g. a connection held
|
||||
# open with no token progress), so ``wait_for`` is the hard cap that cancels a hung
|
||||
# call regardless — otherwise one straggler pins a worker slot and stalls its gather.
|
||||
last_exception = e
|
||||
exc_name = type(e).__name__
|
||||
if attempt < max_retries:
|
||||
logger.warning(
|
||||
f"LiteLLM call exceeded timeout={self.timeout}s ({exc_name}, scope={scope}), retrying..."
|
||||
)
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
logger.error(
|
||||
f"LiteLLM call timed out after {self.timeout}s on {attempt + 1} attempts "
|
||||
f"({exc_name}, scope={scope})"
|
||||
)
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
error_str = str(e).lower()
|
||||
# Fast fail on auth errors
|
||||
@@ -421,18 +354,7 @@ class LiteLLMLLM(LLMInterface):
|
||||
if attempt > 0:
|
||||
set_stage(f"llm.{self._stage_label}.tools.attempt={attempt + 1}/{max_retries + 1}")
|
||||
try:
|
||||
response = await asyncio.wait_for(
|
||||
self._acompletion(**call_kwargs),
|
||||
timeout=self.timeout,
|
||||
)
|
||||
# Stash usage before the tool-call argument parse below, which
|
||||
# can raise json.JSONDecodeError locally even though the provider
|
||||
# already billed for these tokens; without this the error trace
|
||||
# records 0/0 tokens (#2387). Mirrors call() and the anthropic/
|
||||
# gemini call_with_tools paths so the litellm tool path (and the
|
||||
# LiteLLMRouterLLM subclass that inherits this method) completes
|
||||
# the #2396 usage-on-error coverage.
|
||||
stash_response_usage(_usage_from_litellm_response(response))
|
||||
response = await self._acompletion(**call_kwargs)
|
||||
|
||||
message = response.choices[0].message
|
||||
content = message.content
|
||||
@@ -502,23 +424,6 @@ class LiteLLMLLM(LLMInterface):
|
||||
output_tokens=output_tokens,
|
||||
)
|
||||
|
||||
except (TimeoutError, asyncio.TimeoutError, LiteLLMTimeout) as e:
|
||||
# See ``call`` — hard cap so a hung completion cannot block
|
||||
# forever and pin a worker slot / concurrency permit.
|
||||
last_exception = e
|
||||
exc_name = type(e).__name__
|
||||
if attempt < max_retries:
|
||||
logger.warning(
|
||||
f"LiteLLM tool call exceeded timeout={self.timeout}s ({exc_name}, scope={scope}), retrying..."
|
||||
)
|
||||
await asyncio.sleep(min(initial_backoff * (2**attempt), max_backoff))
|
||||
continue
|
||||
logger.error(
|
||||
f"LiteLLM tool call timed out after {self.timeout}s on {attempt + 1} attempts "
|
||||
f"({exc_name}, scope={scope})"
|
||||
)
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
error_str = str(e).lower()
|
||||
if "401" in error_str or "403" in error_str or "unauthorized" in error_str:
|
||||
|
||||
@@ -67,7 +67,7 @@ class LiteLLMRouterLLM(LiteLLMLLM):
|
||||
model: str,
|
||||
config: dict[str, Any],
|
||||
reasoning_effort: str = "low",
|
||||
timeout: float | None = None,
|
||||
timeout: float = 300.0,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(
|
||||
@@ -146,28 +146,16 @@ class LiteLLMRouterLLM(LiteLLMLLM):
|
||||
kwargs["max_completion_tokens"] = self._cap_max_completion_tokens(max_completion_tokens)
|
||||
if temperature is not None:
|
||||
kwargs["temperature"] = temperature
|
||||
|
||||
# Forward operator-configured default headers as ``extra_headers`` so they
|
||||
# reach the provider behind the Router (proxies / request-tracing middleware).
|
||||
# This override deliberately omits api_key/base_url/extra_body (those live in
|
||||
# the per-deployment Router config), but headers are a cross-cutting operator
|
||||
# concern, so we inject them here too — mirroring the base provider.
|
||||
# ``setdefault`` keeps any explicit per-call ``extra_headers`` authoritative;
|
||||
# a per-call copy prevents LiteLLM/downstream from mutating the stored dict.
|
||||
if self._default_headers:
|
||||
kwargs.setdefault("extra_headers", dict(self._default_headers))
|
||||
return kwargs
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
from hindsight_api.engine.llm_interface import OutputTooLongError
|
||||
|
||||
from ...config import get_config
|
||||
|
||||
try:
|
||||
await self.call(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
max_completion_tokens=50,
|
||||
temperature=get_config().llm_temperature_verification,
|
||||
temperature=0.0,
|
||||
scope="verification",
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
@@ -101,7 +101,7 @@ class MockLLM(LLMInterface):
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
response_format: Optional Pydantic model for structured output.
|
||||
max_completion_tokens: Not used in mock.
|
||||
temperature: Recorded on the call record for test assertions.
|
||||
temperature: Not used in mock.
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Not used in mock.
|
||||
initial_backoff: Not used in mock.
|
||||
@@ -123,9 +123,6 @@ class MockLLM(LLMInterface):
|
||||
if response_format and hasattr(response_format, "__name__")
|
||||
else str(response_format),
|
||||
"scope": scope,
|
||||
# Record the temperature so tests can assert per-operation temperature
|
||||
# wiring (None means the parameter was omitted from the call).
|
||||
"temperature": temperature,
|
||||
}
|
||||
self._mock_calls.append(call_record)
|
||||
logger.debug(f"Mock LLM call recorded: scope={scope}, model={self.model}")
|
||||
@@ -211,7 +208,7 @@ class MockLLM(LLMInterface):
|
||||
messages: List of message dicts. Can include tool results with role='tool'.
|
||||
tools: List of tool definitions in OpenAI format.
|
||||
max_completion_tokens: Not used in mock.
|
||||
temperature: Recorded on the call record for test assertions.
|
||||
temperature: Not used in mock.
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Not used in mock.
|
||||
initial_backoff: Not used in mock.
|
||||
@@ -228,9 +225,6 @@ class MockLLM(LLMInterface):
|
||||
"messages": messages,
|
||||
"tools": [t.get("function", {}).get("name") for t in tools],
|
||||
"scope": scope,
|
||||
# Record the temperature so tests can assert per-operation temperature
|
||||
# wiring (None means the parameter was omitted from the call).
|
||||
"temperature": temperature,
|
||||
}
|
||||
self._mock_calls.append(call_record)
|
||||
|
||||
|
||||
@@ -26,8 +26,6 @@ import logging
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from email.utils import parsedate_to_datetime
|
||||
from typing import Any
|
||||
from urllib.parse import parse_qs, urlparse, urlunparse
|
||||
|
||||
@@ -36,8 +34,7 @@ from openai import APIConnectionError, APIStatusError, AsyncOpenAI, LengthFinish
|
||||
|
||||
from hindsight_api.config import DEFAULT_LLM_TIMEOUT, ENV_LLM_TIMEOUT
|
||||
from hindsight_api.engine.bank_attribution import apply_bank_attribution
|
||||
from hindsight_api.engine.llm_interface import LLMInterface, OutputTooLongError, ProviderRateLimitResetError
|
||||
from hindsight_api.engine.llm_trace import LLMResponseUsage, stash_response_usage
|
||||
from hindsight_api.engine.llm_interface import LLMInterface, OutputTooLongError
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
from hindsight_api.worker.stage import set_stage
|
||||
@@ -86,49 +83,6 @@ def _strip_code_fences(content: str) -> str:
|
||||
return content
|
||||
|
||||
|
||||
# Reasoning/thinking tags emitted by extended-thinking models. Some providers
|
||||
# (e.g. MiniMax-M3) leak the chain-of-thought wrapped in these tags into the
|
||||
# response body instead of a separate reasoning_content field. Each entry is
|
||||
# (open_tag, close_tag); the open tag also matches when the close tag is missing
|
||||
# (truncated output) so a dangling block is removed to end-of-string.
|
||||
_REASONING_TAG_PAIRS: tuple[tuple[str, str], ...] = (
|
||||
("<think>", "</think>"),
|
||||
("<thinking>", "</thinking>"),
|
||||
("<thought>", "</thought>"),
|
||||
("<reasoning>", "</reasoning>"),
|
||||
("|startthink|", "|endthink|"),
|
||||
)
|
||||
|
||||
|
||||
def _strip_reasoning_tags(text: str) -> str:
|
||||
"""Strip extended-thinking/reasoning blocks from an LLM response.
|
||||
|
||||
Removes the full set of tag styles emitted by reasoning models:
|
||||
``<think>``, ``<thinking>``, ``<thought>``, ``<reasoning>`` and the
|
||||
``|startthink|...|endthink|`` markers. Both the structured (JSON) path and
|
||||
the free-form path must call this — otherwise a non-structured response
|
||||
(e.g. a mental-model markdown blob from MiniMax-M3) leaks the raw
|
||||
``<think>...</think>`` verbatim into stored memories.
|
||||
|
||||
Handles two cases:
|
||||
1. Closed blocks: ``<think>...</think>`` removed wherever they appear.
|
||||
2. Unclosed blocks: a dangling ``<think>`` with no closing tag (model output
|
||||
truncated mid-thought) is removed from the open tag to end-of-string.
|
||||
|
||||
Returns the input unchanged (modulo surrounding whitespace) when no tags are
|
||||
present.
|
||||
"""
|
||||
if not text:
|
||||
return text
|
||||
for open_tag, close_tag in _REASONING_TAG_PAIRS:
|
||||
open_re = re.escape(open_tag)
|
||||
close_re = re.escape(close_tag)
|
||||
# Closed blocks first, then any remaining unclosed (truncated) block.
|
||||
text = re.sub(rf"{open_re}.*?{close_re}", "", text, flags=re.DOTALL)
|
||||
text = re.sub(rf"{open_re}.*", "", text, flags=re.DOTALL)
|
||||
return text.strip()
|
||||
|
||||
|
||||
def _response_get(response: Any, key: str, default: Any = None) -> Any:
|
||||
if isinstance(response, dict):
|
||||
return response.get(key, default)
|
||||
@@ -233,21 +187,6 @@ def _content_or_error(response: Any, *, provider: str, model: str, scope: str) -
|
||||
return content, choice
|
||||
|
||||
|
||||
def _usage_from_openai_response(response: Any) -> LLMResponseUsage:
|
||||
"""Extract prompt/completion/cached token counts from an OpenAI-shaped usage block."""
|
||||
usage = getattr(response, "usage", None)
|
||||
input_tokens = (usage.prompt_tokens or 0) if usage else 0
|
||||
output_tokens = (usage.completion_tokens or 0) if usage else 0
|
||||
cached_tokens = 0
|
||||
if usage and getattr(usage, "prompt_tokens_details", None):
|
||||
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
|
||||
return LLMResponseUsage(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
)
|
||||
|
||||
|
||||
def _ensure_json_word_in_user_message(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
"""Some OpenAI-compatible gateways require 'json' in a user message for json_object mode."""
|
||||
|
||||
@@ -295,122 +234,6 @@ def _summarize_status_error(e: APIStatusError, body_max: int = 400) -> str:
|
||||
return f"HTTP {e.status_code}: {body_str or '<no body>'}"
|
||||
|
||||
|
||||
_RATE_LIMIT_RESET_AT_RE = re.compile(
|
||||
r"\breset at\s+"
|
||||
r"(?P<reset_at>\d{4}-\d{2}-\d{2}[ T]\d{2}:\d{2}:\d{2}(?:\s*(?:Z|[+-]\d{2}:?\d{2}))?)",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
_RATE_LIMIT_WINDOW_RE = re.compile(
|
||||
r"\b(?:for|in)\s+(?P<amount>\d+)\s*(?P<unit>second|minute|hour|day)s?\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _status_error_body_text(e: APIStatusError) -> str:
|
||||
body: Any = getattr(e, "body", None)
|
||||
if body is None:
|
||||
try:
|
||||
body = e.response.text
|
||||
except Exception:
|
||||
body = None
|
||||
if isinstance(body, (dict, list)):
|
||||
try:
|
||||
return json.dumps(body, default=str, ensure_ascii=False)
|
||||
except Exception:
|
||||
return str(body)
|
||||
return str(body or "").strip()
|
||||
|
||||
|
||||
def _parse_retry_after_header(value: str | None, now: datetime) -> datetime | None:
|
||||
if not value:
|
||||
return None
|
||||
raw = value.strip()
|
||||
try:
|
||||
seconds = float(raw)
|
||||
except ValueError:
|
||||
seconds = -1.0
|
||||
if seconds >= 0:
|
||||
return now + timedelta(seconds=seconds)
|
||||
|
||||
try:
|
||||
parsed = parsedate_to_datetime(raw)
|
||||
except (TypeError, ValueError, IndexError, OverflowError):
|
||||
return None
|
||||
if parsed.tzinfo is None:
|
||||
parsed = parsed.replace(tzinfo=UTC)
|
||||
return parsed.astimezone(UTC)
|
||||
|
||||
|
||||
def _parse_reset_at_datetime(value: str) -> datetime | None:
|
||||
raw = value.strip().replace(" ", "T")
|
||||
if raw.endswith("Z"):
|
||||
raw = f"{raw[:-1]}+00:00"
|
||||
elif re.search(r"[+-]\d{4}$", raw):
|
||||
raw = f"{raw[:-2]}:{raw[-2:]}"
|
||||
try:
|
||||
parsed = datetime.fromisoformat(raw)
|
||||
except ValueError:
|
||||
return None
|
||||
if parsed.tzinfo is None:
|
||||
# Some providers (z.ai included) return a wall-clock reset timestamp
|
||||
# without a zone. Interpret it in the host's local zone so logs, status
|
||||
# pages, and the queued next_retry_at describe the same operator-facing
|
||||
# clock instead of silently shifting by UTC offset.
|
||||
parsed = parsed.astimezone()
|
||||
return parsed.astimezone(UTC)
|
||||
|
||||
|
||||
def _rate_limit_retry_at(e: APIStatusError) -> datetime | None:
|
||||
now = datetime.now(UTC)
|
||||
response = getattr(e, "response", None)
|
||||
headers = getattr(response, "headers", None)
|
||||
if headers is not None:
|
||||
retry_at = _parse_retry_after_header(headers.get("retry-after") or headers.get("Retry-After"), now)
|
||||
if retry_at is not None and retry_at > now:
|
||||
return retry_at
|
||||
|
||||
body_text = _status_error_body_text(e)
|
||||
reset_match = _RATE_LIMIT_RESET_AT_RE.search(body_text)
|
||||
if reset_match:
|
||||
retry_at = _parse_reset_at_datetime(reset_match.group("reset_at"))
|
||||
if retry_at is not None and retry_at > now:
|
||||
return retry_at
|
||||
|
||||
window_match = _RATE_LIMIT_WINDOW_RE.search(body_text)
|
||||
if not window_match:
|
||||
return None
|
||||
amount = int(window_match.group("amount"))
|
||||
unit = window_match.group("unit").lower()
|
||||
if unit == "second":
|
||||
seconds = amount
|
||||
elif unit == "minute":
|
||||
seconds = amount * 60
|
||||
elif unit == "hour":
|
||||
seconds = amount * 3600
|
||||
else:
|
||||
seconds = amount * 86400
|
||||
return now + timedelta(seconds=seconds)
|
||||
|
||||
|
||||
def _raise_provider_quota_defer(
|
||||
e: APIStatusError, *, provider: str, model: str, scope: str, max_backoff: float
|
||||
) -> None:
|
||||
if e.status_code != 429:
|
||||
return
|
||||
retry_at = _rate_limit_retry_at(e)
|
||||
if retry_at is None:
|
||||
return
|
||||
if (retry_at - datetime.now(UTC)).total_seconds() <= max_backoff:
|
||||
return
|
||||
summary = _summarize_status_error(e)
|
||||
raise ProviderRateLimitResetError(
|
||||
retry_at=retry_at,
|
||||
message=(
|
||||
f"Provider quota exhausted ({provider}/{model}, scope={scope}); retry at {retry_at.isoformat()}: {summary}"
|
||||
),
|
||||
) from e
|
||||
|
||||
|
||||
class OpenAICompatibleLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider for OpenAI-compatible APIs.
|
||||
@@ -446,7 +269,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
base_url: Base URL for the API (uses defaults for groq/ollama/lmstudio if empty).
|
||||
model: Model name.
|
||||
reasoning_effort: Reasoning effort level for supported models ("low", "medium", "high").
|
||||
timeout: Request timeout in seconds (uses env var or 120s default).
|
||||
timeout: Request timeout in seconds (uses env var or 300s default).
|
||||
groq_service_tier: Groq service tier ("on_demand", "flex", "auto").
|
||||
extra_body: Extra body params merged into every API call.
|
||||
**kwargs: Additional provider-specific parameters.
|
||||
@@ -465,10 +288,8 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
"deepseek",
|
||||
"volcano",
|
||||
"openrouter",
|
||||
"requesty",
|
||||
"zai",
|
||||
"opencode-go",
|
||||
"atlas",
|
||||
"fireworks",
|
||||
]
|
||||
if self.provider not in valid_providers:
|
||||
@@ -490,14 +311,10 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
self.base_url = "https://api.deepseek.com"
|
||||
elif self.provider == "openrouter":
|
||||
self.base_url = "https://openrouter.ai/api/v1"
|
||||
elif self.provider == "requesty":
|
||||
self.base_url = "https://router.requesty.ai/v1"
|
||||
elif self.provider == "zai":
|
||||
self.base_url = "https://api.z.ai/api/coding/paas/v4"
|
||||
elif self.provider == "opencode-go":
|
||||
self.base_url = "https://opencode.ai/zen/go/v1"
|
||||
elif self.provider == "atlas":
|
||||
self.base_url = "https://api.atlascloud.ai/v1"
|
||||
elif self.provider == "fireworks":
|
||||
# OpenAI-compatible inference host (online path). The batch API
|
||||
# lives on a separate control-plane host — see FireworksLLM.
|
||||
@@ -516,10 +333,8 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
"minimax",
|
||||
"deepseek",
|
||||
"openrouter",
|
||||
"requesty",
|
||||
"zai",
|
||||
"opencode-go",
|
||||
"atlas",
|
||||
"ollama-cloud",
|
||||
)
|
||||
and not self.api_key
|
||||
@@ -794,9 +609,6 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
try:
|
||||
if response_format is not None:
|
||||
response = await self._client.chat.completions.create(**call_params)
|
||||
# Stash usage before parse/validate, which may raise locally
|
||||
# even though the provider charged for these tokens (#2387).
|
||||
stash_response_usage(_usage_from_openai_response(response))
|
||||
|
||||
content, first_choice = _content_or_error(
|
||||
response,
|
||||
@@ -805,10 +617,15 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
scope=scope,
|
||||
)
|
||||
|
||||
# Strip reasoning model thinking tags (closed and unclosed).
|
||||
# Strip reasoning model thinking tags
|
||||
# Supports: <think>, <thinking>, <thought>, <reasoning>, |startthink|/|endthink|
|
||||
original_len = len(content)
|
||||
content = _strip_reasoning_tags(content)
|
||||
content = re.sub(r"<think>.*?</think>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"<thinking>.*?</thinking>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"<thought>.*?</thought>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"<reasoning>.*?</reasoning>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"\|startthink\|.*?\|endthink\|", "", content, flags=re.DOTALL)
|
||||
content = content.strip()
|
||||
if len(content) < original_len:
|
||||
logger.debug(f"Stripped {original_len - len(content)} chars of reasoning tokens")
|
||||
|
||||
@@ -850,7 +667,6 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
result = response_format.model_validate(json_data)
|
||||
else:
|
||||
response = await self._client.chat.completions.create(**call_params)
|
||||
stash_response_usage(_usage_from_openai_response(response))
|
||||
result, first_choice = _content_or_error(
|
||||
response,
|
||||
provider=self.provider,
|
||||
@@ -858,33 +674,15 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
scope=scope,
|
||||
)
|
||||
|
||||
# Free-form (non-structured) output also leaks reasoning tags:
|
||||
# reasoning models like MiniMax-M3 wrap their chain-of-thought
|
||||
# in <think>...</think> in the response body. Without this strip
|
||||
# a mental-model markdown blob is stored verbatim with the raw
|
||||
# thinking tags. Mirrors the structured-output path above.
|
||||
result = _strip_reasoning_tags(result)
|
||||
|
||||
# Record token usage metrics
|
||||
duration = time.time() - start_time
|
||||
usage = response.usage
|
||||
response_usage = _usage_from_openai_response(response)
|
||||
input_tokens = response_usage.input_tokens
|
||||
output_tokens = response_usage.output_tokens
|
||||
input_tokens = usage.prompt_tokens or 0 if usage else 0
|
||||
output_tokens = usage.completion_tokens or 0 if usage else 0
|
||||
total_tokens = usage.total_tokens or 0 if usage else 0
|
||||
cached_tokens = response_usage.cached_tokens
|
||||
thoughts_tokens = 0
|
||||
if usage and getattr(usage, "completion_tokens_details", None):
|
||||
thoughts_tokens = getattr(usage.completion_tokens_details, "reasoning_tokens", 0) or 0
|
||||
# OpenAI-compatible providers fold reasoning tokens into
|
||||
# ``completion_tokens`` (and thus ``total_tokens``), but the
|
||||
# TokenUsage contract — and the Gemini provider — treat
|
||||
# ``output_tokens``/``total_tokens`` as visible-only, surfacing
|
||||
# reasoning separately in ``thoughts_tokens``. Subtract so the
|
||||
# two fields don't double-count reasoning (cost over-attribution).
|
||||
if thoughts_tokens:
|
||||
output_tokens = max(0, output_tokens - thoughts_tokens)
|
||||
total_tokens = max(0, total_tokens - thoughts_tokens)
|
||||
cached_tokens = 0
|
||||
if usage and getattr(usage, "prompt_tokens_details", None):
|
||||
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
|
||||
|
||||
# Record LLM metrics
|
||||
metrics = get_metrics_collector()
|
||||
@@ -933,7 +731,6 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
return result, token_usage
|
||||
return result
|
||||
@@ -964,10 +761,6 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
logger.error(f"Auth error (HTTP {e.status_code}), not retrying: {str(e)}")
|
||||
raise
|
||||
|
||||
_raise_provider_quota_defer(
|
||||
e, provider=self.provider, model=self.model, scope=scope, max_backoff=max_backoff
|
||||
)
|
||||
|
||||
# Handle tool_use_failed error - model outputted in tool call format
|
||||
if e.status_code == 400 and response_format is not None:
|
||||
try:
|
||||
@@ -1021,6 +814,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
f"scope={scope}): {_summarize_status_error(e)}"
|
||||
)
|
||||
raise
|
||||
|
||||
except ProviderResponseError as e:
|
||||
last_exception = e
|
||||
if e.retryable and attempt < max_retries:
|
||||
@@ -1184,17 +978,6 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
usage = response.usage
|
||||
input_tokens = usage.prompt_tokens or 0 if usage else 0
|
||||
output_tokens = usage.completion_tokens or 0 if usage else 0
|
||||
cached_tokens = 0
|
||||
if usage and getattr(usage, "prompt_tokens_details", None):
|
||||
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
|
||||
thoughts_tokens = 0
|
||||
if usage and getattr(usage, "completion_tokens_details", None):
|
||||
thoughts_tokens = getattr(usage.completion_tokens_details, "reasoning_tokens", 0) or 0
|
||||
# See ``call()``: OpenAI-compatible ``completion_tokens`` includes
|
||||
# reasoning, so make ``output_tokens`` visible-only to avoid
|
||||
# double-counting it against ``thoughts_tokens``.
|
||||
if thoughts_tokens:
|
||||
output_tokens = max(0, output_tokens - thoughts_tokens)
|
||||
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
@@ -1237,8 +1020,6 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
finish_reason=finish_reason,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
thoughts_tokens=thoughts_tokens,
|
||||
)
|
||||
|
||||
except APIConnectionError as e:
|
||||
@@ -1266,10 +1047,6 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
f"not retrying: {_summarize_status_error(e)}"
|
||||
)
|
||||
raise
|
||||
_raise_provider_quota_defer(
|
||||
e, provider=self.provider, model=self.model, scope=scope, max_backoff=max_backoff
|
||||
)
|
||||
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
logger.warning(
|
||||
@@ -1283,6 +1060,7 @@ class OpenAICompatibleLLM(LLMInterface):
|
||||
f"({self.provider}/{self.model}, scope={scope}): {_summarize_status_error(e)}"
|
||||
)
|
||||
raise
|
||||
|
||||
except Exception:
|
||||
raise
|
||||
|
||||
|
||||
@@ -6,17 +6,12 @@ structured information like temporal constraints.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from hindsight_api.engine.temporal_periods import (
|
||||
NO_TEMPORAL_CONSTRAINT,
|
||||
extract_period,
|
||||
is_embedded_cjk_dateparser_match,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -128,12 +123,9 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
|
||||
# Check for period expressions first (these need special handling)
|
||||
query_lower = query.lower()
|
||||
period_result = extract_period(query_lower, reference_date)
|
||||
if period_result is NO_TEMPORAL_CONSTRAINT:
|
||||
return QueryAnalysis(temporal_constraint=None)
|
||||
if isinstance(period_result, tuple):
|
||||
start_date, end_date = period_result
|
||||
return QueryAnalysis(temporal_constraint=TemporalConstraint(start_date=start_date, end_date=end_date))
|
||||
period_result = self._extract_period(query_lower, reference_date)
|
||||
if period_result is not None:
|
||||
return QueryAnalysis(temporal_constraint=period_result)
|
||||
|
||||
# Lazy load dateparser (only imports on first call, then cached)
|
||||
self.load()
|
||||
@@ -166,12 +158,7 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
|
||||
# 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)
|
||||
and not is_embedded_cjk_dateparser_match(query, text)
|
||||
]
|
||||
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)
|
||||
@@ -185,6 +172,127 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
|
||||
return QueryAnalysis(temporal_constraint=TemporalConstraint(start_date=start_date, end_date=end_date))
|
||||
|
||||
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),
|
||||
)
|
||||
|
||||
# Yesterday patterns (English, Spanish, Italian, French, German)
|
||||
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):
|
||||
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):
|
||||
# "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):
|
||||
# "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):
|
||||
# "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):
|
||||
# "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):
|
||||
# "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):
|
||||
# "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,
|
||||
):
|
||||
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,
|
||||
):
|
||||
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,
|
||||
):
|
||||
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,
|
||||
):
|
||||
days_since_sat = (reference_date.weekday() + 2) % 7
|
||||
if days_since_sat == 0:
|
||||
days_since_sat = 7
|
||||
sat = reference_date - timedelta(days=days_since_sat)
|
||||
return constraint(sat, sat + timedelta(days=1))
|
||||
|
||||
# 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,
|
||||
}
|
||||
|
||||
for pattern, month_num in month_patterns.items():
|
||||
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)
|
||||
if month_num == 12:
|
||||
end = datetime(year, 12, 31)
|
||||
else:
|
||||
end = datetime(year, month_num + 1, 1) - timedelta(days=1)
|
||||
return constraint(start, end)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
class TransformerQueryAnalyzer(QueryAnalyzer):
|
||||
"""
|
||||
|
||||
@@ -15,7 +15,7 @@ import time
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable
|
||||
|
||||
from ...config import get_config
|
||||
from .models import DirectiveInfo, LLMCall, ReflectAgentResult, StructuredOutputResult, TokenUsageSummary, ToolCall
|
||||
from .models import DirectiveInfo, LLMCall, ReflectAgentResult, TokenUsageSummary, ToolCall
|
||||
from .prompts import (
|
||||
_extract_directive_rules,
|
||||
build_final_prompt,
|
||||
@@ -90,87 +90,12 @@ _LEAKED_JSON_SUFFIX = re.compile(
|
||||
r'\s*```(?:json)?\s*\{[^}]*(?:"(?:observation_ids|memory_ids|mental_model_ids)"|\})\s*```\s*$',
|
||||
re.DOTALL | re.IGNORECASE,
|
||||
)
|
||||
_LEAKED_JSON_OBJECT = re.compile(
|
||||
r'\s*\{[^{]*"(?:observation_ids|memory_ids|mental_model_ids|answer)"[^}]*\}\s*$', re.DOTALL
|
||||
)
|
||||
_TRAILING_IDS_PATTERN = re.compile(
|
||||
r"\s*(?:observation_ids|memory_ids|mental_model_ids)\s*[=:]\s*\[.*?\]\s*$", re.DOTALL | re.IGNORECASE
|
||||
)
|
||||
_JSON_CODE_FENCE_PATTERN = re.compile(r"^\s*```(?:json)?\s*(\{.*\})\s*```\s*$", re.DOTALL | re.IGNORECASE)
|
||||
|
||||
_DONE_ARGUMENT_KEYS = frozenset(
|
||||
{
|
||||
"answer",
|
||||
"directive_compliance",
|
||||
"memory_ids",
|
||||
"mental_model_ids",
|
||||
"observation_ids",
|
||||
"model_ids",
|
||||
}
|
||||
)
|
||||
_DONE_ARGUMENT_MARKER_KEYS = _DONE_ARGUMENT_KEYS - {"answer"}
|
||||
_LEAKED_JSON_ID_KEYS = frozenset({"memory_ids", "mental_model_ids", "observation_ids", "model_ids"})
|
||||
|
||||
|
||||
def _unwrap_leaked_done_arguments(text: str) -> str | None:
|
||||
"""Return the answer when a done tool call was rendered as JSON text.
|
||||
|
||||
Some providers leak the done tool's argument object instead of surfacing it
|
||||
as a native tool call, e.g. {"answer": "...", "memory_ids": [...]}. Only
|
||||
unwrap objects that match the done argument shape so normal JSON answers
|
||||
stay intact.
|
||||
"""
|
||||
candidate = text.strip()
|
||||
if not candidate:
|
||||
return None
|
||||
|
||||
fenced = _JSON_CODE_FENCE_PATTERN.match(candidate)
|
||||
if fenced:
|
||||
candidate = fenced.group(1).strip()
|
||||
|
||||
try:
|
||||
payload = json.loads(candidate)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
return None
|
||||
answer = payload.get("answer")
|
||||
if not isinstance(answer, str) or not answer.strip():
|
||||
return None
|
||||
|
||||
keys = set(payload)
|
||||
if not keys.intersection(_DONE_ARGUMENT_MARKER_KEYS):
|
||||
return None
|
||||
if not keys.issubset(_DONE_ARGUMENT_KEYS):
|
||||
return None
|
||||
|
||||
for key in ("memory_ids", "mental_model_ids", "observation_ids", "model_ids"):
|
||||
value = payload.get(key)
|
||||
if value is not None and not isinstance(value, list):
|
||||
return None
|
||||
|
||||
return answer.strip()
|
||||
|
||||
|
||||
def _strip_trailing_id_json_object(text: str) -> str:
|
||||
stripped = text.rstrip()
|
||||
if not stripped.endswith("}"):
|
||||
return text.strip()
|
||||
|
||||
start = stripped.rfind("{")
|
||||
if start < 0:
|
||||
return text.strip()
|
||||
|
||||
try:
|
||||
payload = json.loads(stripped[start:])
|
||||
except json.JSONDecodeError:
|
||||
return text.strip()
|
||||
|
||||
if not isinstance(payload, dict) or not payload:
|
||||
return text.strip()
|
||||
keys = set(payload)
|
||||
if not keys.issubset(_LEAKED_JSON_ID_KEYS):
|
||||
return text.strip()
|
||||
|
||||
return stripped[:start].strip()
|
||||
|
||||
|
||||
def _clean_answer_text(text: str) -> str:
|
||||
@@ -179,10 +104,6 @@ def _clean_answer_text(text: str) -> str:
|
||||
Some LLMs output the done() call as text instead of a proper tool call.
|
||||
This strips out patterns like: done({"answer": "...", ...})
|
||||
"""
|
||||
unwrapped = _unwrap_leaked_done_arguments(text)
|
||||
if unwrapped is not None:
|
||||
return unwrapped
|
||||
|
||||
# Remove done() call pattern from the end of the text
|
||||
cleaned = _DONE_CALL_PATTERN.sub("", text).strip()
|
||||
return cleaned if cleaned else text
|
||||
@@ -201,17 +122,13 @@ def _clean_done_answer(text: str) -> str:
|
||||
if not text:
|
||||
return text
|
||||
|
||||
unwrapped = _unwrap_leaked_done_arguments(text)
|
||||
if unwrapped is not None:
|
||||
return unwrapped
|
||||
|
||||
cleaned = text
|
||||
|
||||
# Remove leaked JSON in code blocks at the end
|
||||
cleaned = _LEAKED_JSON_SUFFIX.sub("", cleaned).strip()
|
||||
|
||||
# Remove leaked raw JSON objects at the end
|
||||
cleaned = _strip_trailing_id_json_object(cleaned)
|
||||
cleaned = _LEAKED_JSON_OBJECT.sub("", cleaned).strip()
|
||||
|
||||
# Remove trailing ID patterns
|
||||
cleaned = _TRAILING_IDS_PATTERN.sub("", cleaned).strip()
|
||||
@@ -224,7 +141,7 @@ async def _generate_structured_output(
|
||||
response_schema: dict,
|
||||
llm_config: "LLMProvider",
|
||||
reflect_id: str,
|
||||
) -> StructuredOutputResult:
|
||||
) -> tuple[dict[str, Any] | None, int, int]:
|
||||
"""Generate structured output from an answer using the provided JSON schema.
|
||||
|
||||
Args:
|
||||
@@ -234,8 +151,8 @@ async def _generate_structured_output(
|
||||
reflect_id: Reflect ID for logging
|
||||
|
||||
Returns:
|
||||
A StructuredOutputResult carrying the structured output (None if
|
||||
generation fails) and the call's token usage.
|
||||
Tuple of (structured_output, input_tokens, output_tokens).
|
||||
structured_output is None if generation fails.
|
||||
"""
|
||||
try:
|
||||
from typing import Any as TypingAny
|
||||
@@ -269,7 +186,7 @@ async def _generate_structured_output(
|
||||
|
||||
if not fields:
|
||||
logger.warning(f"[REFLECT {reflect_id}] No fields found in response_schema, skipping structured output")
|
||||
return StructuredOutputResult()
|
||||
return None, 0, 0
|
||||
|
||||
DynamicModel = create_model("StructuredResponse", **fields)
|
||||
|
||||
@@ -322,9 +239,6 @@ OUTPUT:"""
|
||||
],
|
||||
response_format=DynamicModel,
|
||||
scope="reflect_structured",
|
||||
max_retries=1,
|
||||
initial_backoff=0.25,
|
||||
max_backoff=1.0,
|
||||
skip_validation=True, # We'll handle the dict ourselves
|
||||
return_usage=True,
|
||||
)
|
||||
@@ -345,17 +259,11 @@ OUTPUT:"""
|
||||
logger.warning(f"[REFLECT {reflect_id}] Required field '{field_name}' is empty in structured output")
|
||||
|
||||
logger.info(f"[REFLECT {reflect_id}] Generated structured output with {len(structured_output)} fields")
|
||||
return StructuredOutputResult(
|
||||
structured_output=structured_output,
|
||||
input_tokens=usage.input_tokens,
|
||||
output_tokens=usage.output_tokens,
|
||||
cached_tokens=usage.cached_tokens,
|
||||
thoughts_tokens=usage.thoughts_tokens,
|
||||
)
|
||||
return structured_output, usage.input_tokens, usage.output_tokens
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"[REFLECT {reflect_id}] Failed to generate structured output: {e}")
|
||||
return StructuredOutputResult()
|
||||
return None, 0, 0
|
||||
|
||||
|
||||
def _count_messages_tokens(messages: list[dict[str, Any]]) -> int:
|
||||
@@ -527,14 +435,9 @@ async def run_reflect_agent(
|
||||
llm_trace: list[dict[str, Any]] = []
|
||||
context_history: list[dict[str, Any]] = [] # For final prompt fallback
|
||||
|
||||
# Token usage tracking - accumulate across all LLM calls.
|
||||
# cached_tokens and thoughts_tokens are surfaced for cost attribution
|
||||
# and prompt-cache tuning. Both are subsets of (or parallel to) the
|
||||
# input/output counts and are NOT double-counted in total_tokens.
|
||||
# Token usage tracking - accumulate across all LLM calls
|
||||
total_input_tokens = 0
|
||||
total_output_tokens = 0
|
||||
total_cached_tokens = 0
|
||||
total_thoughts_tokens = 0
|
||||
|
||||
# Track available IDs for validation (prevents hallucinated citations)
|
||||
available_memory_ids: set[str] = set()
|
||||
@@ -557,8 +460,6 @@ async def run_reflect_agent(
|
||||
input_tokens=total_input_tokens,
|
||||
output_tokens=total_output_tokens,
|
||||
total_tokens=total_input_tokens + total_output_tokens,
|
||||
cached_tokens=total_cached_tokens,
|
||||
thoughts_tokens=total_thoughts_tokens,
|
||||
)
|
||||
|
||||
def _log_completion(answer: str, iterations: int, forced: bool = False):
|
||||
@@ -625,8 +526,6 @@ async def run_reflect_agent(
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
total_cached_tokens += getattr(usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
@@ -640,12 +539,11 @@ async def run_reflect_agent(
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
@@ -690,8 +588,6 @@ async def run_reflect_agent(
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
total_cached_tokens += getattr(usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
@@ -704,12 +600,11 @@ async def run_reflect_agent(
|
||||
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
@@ -766,8 +661,6 @@ async def run_reflect_agent(
|
||||
consecutive_errors = 0
|
||||
total_input_tokens += result.input_tokens
|
||||
total_output_tokens += result.output_tokens
|
||||
total_cached_tokens += getattr(result, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(result, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": f"agent_{iteration + 1}",
|
||||
@@ -816,8 +709,6 @@ async def run_reflect_agent(
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
total_cached_tokens += getattr(usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
@@ -831,12 +722,11 @@ async def run_reflect_agent(
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
@@ -893,8 +783,6 @@ async def run_reflect_agent(
|
||||
)
|
||||
total_input_tokens += rewrite_usage.input_tokens
|
||||
total_output_tokens += rewrite_usage.output_tokens
|
||||
total_cached_tokens += getattr(rewrite_usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(rewrite_usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final_rewrite",
|
||||
@@ -908,12 +796,11 @@ async def run_reflect_agent(
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
|
||||
_log_completion(answer, iteration + 1)
|
||||
return ReflectAgentResult(
|
||||
@@ -948,8 +835,6 @@ async def run_reflect_agent(
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
total_cached_tokens += getattr(usage, "cached_tokens", 0) or 0
|
||||
total_thoughts_tokens += getattr(usage, "thoughts_tokens", 0) or 0
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
@@ -963,12 +848,11 @@ async def run_reflect_agent(
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
structured_output = struct.structured_output
|
||||
total_input_tokens += struct.input_tokens
|
||||
total_output_tokens += struct.output_tokens
|
||||
total_cached_tokens += struct.cached_tokens
|
||||
total_thoughts_tokens += struct.thoughts_tokens
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
@@ -1263,15 +1147,14 @@ async def _process_done_tool(
|
||||
structured_output = None
|
||||
final_usage = usage
|
||||
if response_schema and llm_config and answer:
|
||||
struct = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
structured_output = struct.structured_output
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
# Add structured output tokens to usage
|
||||
final_usage = TokenUsageSummary(
|
||||
input_tokens=usage.input_tokens + struct.input_tokens,
|
||||
output_tokens=usage.output_tokens + struct.output_tokens,
|
||||
total_tokens=usage.total_tokens + struct.input_tokens + struct.output_tokens,
|
||||
cached_tokens=usage.cached_tokens + struct.cached_tokens,
|
||||
thoughts_tokens=usage.thoughts_tokens + struct.thoughts_tokens,
|
||||
input_tokens=usage.input_tokens + struct_in,
|
||||
output_tokens=usage.output_tokens + struct_out,
|
||||
total_tokens=usage.total_tokens + struct_in + struct_out,
|
||||
)
|
||||
|
||||
log_completion(answer, iterations)
|
||||
|
||||
@@ -26,13 +26,10 @@ or stay the same per refresh, never get worse.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Annotated, Any, Literal, Union
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
|
||||
from hindsight_api.engine.llm_wrapper import parse_llm_json
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from .structured_doc import (
|
||||
Block,
|
||||
@@ -147,27 +144,6 @@ Operation = Annotated[
|
||||
Field(discriminator="op"),
|
||||
]
|
||||
|
||||
_OPERATION_ADAPTER: TypeAdapter[Operation] = TypeAdapter(Operation)
|
||||
|
||||
|
||||
def _validate_operations_list(raw_ops: Any) -> tuple[list[Operation], list[dict[str, Any]]]:
|
||||
"""Validate each operation independently; drop invalid ops instead of failing the batch."""
|
||||
if not isinstance(raw_ops, list):
|
||||
raise TypeError(f"operations must be a list, got {type(raw_ops)!r}")
|
||||
valid: list[Operation] = []
|
||||
skipped: list[dict[str, Any]] = []
|
||||
for i, item in enumerate(raw_ops):
|
||||
try:
|
||||
valid.append(_OPERATION_ADAPTER.validate_python(item))
|
||||
except ValidationError as exc:
|
||||
skipped.append({"index": i, "op": item, "error": exc.errors(include_url=False)})
|
||||
logger.warning(
|
||||
"[STRUCTURED_DELTA] skipping invalid operation at index %s: %s",
|
||||
i,
|
||||
exc.errors(include_url=False),
|
||||
)
|
||||
return valid, skipped
|
||||
|
||||
|
||||
class DeltaOperationList(BaseModel):
|
||||
"""Container for the operations produced by an LLM delta call."""
|
||||
@@ -176,104 +152,6 @@ class DeltaOperationList(BaseModel):
|
||||
operations: list[Operation] = Field(default_factory=list)
|
||||
|
||||
|
||||
class DeltaAllOpsInvalidError(ValueError):
|
||||
"""Raised when the model emitted operations but none survived validation.
|
||||
|
||||
Distinct from an empty ``operations`` array (a legitimate no-op): here every
|
||||
op was malformed, so returning zero valid ops would make the caller apply
|
||||
nothing and silently drop this refresh's new facts. Raising instead lets the
|
||||
caller fall back to a full rewrite, which still integrates the new facts.
|
||||
"""
|
||||
|
||||
|
||||
def _finalize_operations(valid: list[Operation], skipped: list[dict[str, Any]]) -> DeltaOperationList:
|
||||
"""Build the result, but refuse a wholesale validation failure as a silent no-op."""
|
||||
if skipped and not valid:
|
||||
raise DeltaAllOpsInvalidError(f"all {len(skipped)} delta operation(s) failed validation")
|
||||
return DeltaOperationList(operations=valid)
|
||||
|
||||
|
||||
def _extract_balanced_json_object(text: str) -> str | None:
|
||||
"""Return the first top-level ``{...}`` slice, ignoring trailing junk."""
|
||||
start = text.find("{")
|
||||
if start < 0:
|
||||
return None
|
||||
depth = 0
|
||||
in_string = False
|
||||
escape = False
|
||||
for i in range(start, len(text)):
|
||||
ch = text[i]
|
||||
if in_string:
|
||||
if escape:
|
||||
escape = False
|
||||
elif ch == "\\":
|
||||
escape = True
|
||||
elif ch == '"':
|
||||
in_string = False
|
||||
continue
|
||||
if ch == '"':
|
||||
in_string = True
|
||||
elif ch == "{":
|
||||
depth += 1
|
||||
elif ch == "}":
|
||||
depth -= 1
|
||||
if depth == 0:
|
||||
return text[start : i + 1]
|
||||
return None
|
||||
|
||||
|
||||
def parse_delta_operation_list(raw: Any) -> DeltaOperationList:
|
||||
"""Parse structured-delta LLM output into a validated operation list."""
|
||||
if isinstance(raw, DeltaOperationList):
|
||||
return raw
|
||||
if isinstance(raw, dict):
|
||||
ops_raw = raw.get("operations", [])
|
||||
valid, skipped = _validate_operations_list(ops_raw)
|
||||
if skipped:
|
||||
logger.info(
|
||||
"[STRUCTURED_DELTA] parsed %s op(s), skipped %s invalid op(s) from dict payload",
|
||||
len(valid),
|
||||
len(skipped),
|
||||
)
|
||||
return _finalize_operations(valid, skipped)
|
||||
|
||||
text = (raw or "").strip()
|
||||
if not text:
|
||||
return DeltaOperationList()
|
||||
|
||||
candidates: list[str] = [text]
|
||||
extracted = _extract_balanced_json_object(text)
|
||||
if extracted and extracted != text:
|
||||
candidates.append(extracted)
|
||||
|
||||
last_error: Exception | None = None
|
||||
for candidate in candidates:
|
||||
try:
|
||||
payload = parse_llm_json(candidate)
|
||||
except json.JSONDecodeError as exc:
|
||||
last_error = exc
|
||||
continue
|
||||
if not isinstance(payload, dict) or "operations" not in payload:
|
||||
last_error = ValueError("delta payload must be an object with an operations array")
|
||||
continue
|
||||
try:
|
||||
valid, skipped = _validate_operations_list(payload["operations"])
|
||||
except TypeError as exc:
|
||||
last_error = exc
|
||||
continue
|
||||
if skipped:
|
||||
logger.info(
|
||||
"[STRUCTURED_DELTA] parsed %s op(s), skipped %s invalid op(s)",
|
||||
len(valid),
|
||||
len(skipped),
|
||||
)
|
||||
return _finalize_operations(valid, skipped)
|
||||
|
||||
if last_error is not None:
|
||||
raise last_error
|
||||
return DeltaOperationList()
|
||||
|
||||
|
||||
# Application ---------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
@@ -78,32 +78,9 @@ class DirectiveInfo(BaseModel):
|
||||
class TokenUsageSummary(BaseModel):
|
||||
"""Total token usage across all LLM calls."""
|
||||
|
||||
input_tokens: int = Field(default=0, description="Total input tokens used (includes any cached prefix tokens)")
|
||||
output_tokens: int = Field(default=0, description="Total visible output tokens used (excludes reasoning/thoughts)")
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output, excludes thoughts)")
|
||||
cached_tokens: int = Field(
|
||||
default=0,
|
||||
description="Cached/cache-read prompt tokens summed across calls. Subset of input_tokens.",
|
||||
)
|
||||
thoughts_tokens: int = Field(
|
||||
default=0,
|
||||
description=(
|
||||
"Reasoning/thinking tokens summed across calls. Billed at the output rate by some providers "
|
||||
"but not part of visible output."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class StructuredOutputResult(BaseModel):
|
||||
"""Result of structured-output generation, including token usage for the call."""
|
||||
|
||||
structured_output: dict[str, Any] | None = Field(
|
||||
default=None, description="Generated structured output, or None if generation failed"
|
||||
)
|
||||
input_tokens: int = Field(default=0, description="Input tokens used")
|
||||
output_tokens: int = Field(default=0, description="Visible output tokens used")
|
||||
cached_tokens: int = Field(default=0, description="Cached prefix tokens. Subset of input_tokens.")
|
||||
thoughts_tokens: int = Field(default=0, description="Reasoning/thinking tokens, when reported by the provider")
|
||||
input_tokens: int = Field(default=0, description="Total input tokens used")
|
||||
output_tokens: int = Field(default=0, description="Total output tokens used")
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output)")
|
||||
|
||||
|
||||
class ReflectAgentResult(BaseModel):
|
||||
|
||||
@@ -734,65 +734,7 @@ Examples
|
||||
``{"operations": [{"op": "replace_block", "section_id": "overview",
|
||||
"index": 0, "block": {"type": "paragraph", "text": "Updated summary."}}]}``
|
||||
- Remove an obsolete block →
|
||||
``{"operations": [{"op": "remove_block", "section_id": "status", "index": 2}]}``
|
||||
|
||||
JSON STRING RULES (critical)
|
||||
- Every ``text`` and ``items`` string must be valid JSON: escape ``"`` as ``\\"``,
|
||||
backslashes as ``\\\\``, and newlines as ``\\n``. Do not use raw backticks inside
|
||||
strings unless needed; prefer plain quotes for file paths.
|
||||
- ``replace_block``, ``insert_block``, and ``remove_block`` MUST include ``index`` (0-based block position in that section). Use ``replace_section_blocks`` only when replacing every block in a section.
|
||||
|
||||
- Do not append extra ``]`` or ``}`` after the closing ``}`` of the root object."""
|
||||
|
||||
_STRUCTURED_DELTA_DEFAULT_MAX_INPUT_TOKENS = 24_000
|
||||
|
||||
|
||||
def _truncate_cl100k(text: str, max_tokens: int) -> str:
|
||||
"""Truncate text to at most max_tokens using cl100k_base."""
|
||||
if max_tokens <= 0:
|
||||
return ""
|
||||
from .tokenization import count_cl100k_tokens
|
||||
|
||||
if count_cl100k_tokens(text) <= max_tokens:
|
||||
return text
|
||||
enc = __import__("tiktoken").get_encoding("cl100k_base")
|
||||
return enc.decode(enc.encode(text)[:max_tokens])
|
||||
|
||||
|
||||
def _fit_structured_delta_prompt_parts(
|
||||
*,
|
||||
source_query: str,
|
||||
current_document_json: str,
|
||||
candidate_markdown: str,
|
||||
facts_block: str,
|
||||
budget_hint: str,
|
||||
task_footer: str,
|
||||
max_input_tokens: int,
|
||||
) -> tuple[str, str, str, bool]:
|
||||
"""Shrink large prompt sections to fit within max_input_tokens (cl100k estimate)."""
|
||||
from .tokenization import count_cl100k_tokens
|
||||
|
||||
fixed = (
|
||||
f"## Topic\n{source_query}\n\n"
|
||||
f"## CURRENT DOCUMENT (apply ops to this; reference section ids as listed)\n"
|
||||
f"```json\n\n```\n\n"
|
||||
f"## NEW INFORMATION SYNTHESIS (context for how new facts relate to the topic)\n"
|
||||
f"```markdown\n\n```\n\n"
|
||||
f"## SUPPORTING FACTS (new since last refresh — integrate these)\n"
|
||||
f"{budget_hint}\n\n"
|
||||
f"{task_footer}"
|
||||
)
|
||||
facts_header = "## SUPPORTING FACTS (new since last refresh — integrate these)\n"
|
||||
facts_prefix_tokens = count_cl100k_tokens(facts_header)
|
||||
reserved_facts = min(4096, max(512, max_input_tokens // 8))
|
||||
doc_budget = max(1024, (max_input_tokens - count_cl100k_tokens(fixed) - reserved_facts) * 55 // 100)
|
||||
cand_budget = max(512, (max_input_tokens - count_cl100k_tokens(fixed) - reserved_facts) * 30 // 100)
|
||||
facts_budget = max(256, reserved_facts - facts_prefix_tokens)
|
||||
doc_json = _truncate_cl100k(current_document_json, doc_budget)
|
||||
candidate = _truncate_cl100k(candidate_markdown, cand_budget)
|
||||
facts_body = _truncate_cl100k(facts_block, facts_budget)
|
||||
truncated = doc_json != current_document_json or candidate != candidate_markdown or facts_body != facts_block
|
||||
return doc_json, candidate, facts_body, truncated
|
||||
``{"operations": [{"op": "remove_block", "section_id": "status", "index": 2}]}``"""
|
||||
|
||||
|
||||
def build_structured_delta_prompt(
|
||||
@@ -802,7 +744,6 @@ def build_structured_delta_prompt(
|
||||
supporting_facts: list[dict[str, Any]],
|
||||
source_query: str,
|
||||
max_output_tokens: int | None = None,
|
||||
max_input_tokens: int | None = None,
|
||||
) -> str:
|
||||
"""Build the user prompt for a structured-delta mental model refresh.
|
||||
|
||||
@@ -833,39 +774,19 @@ def build_structured_delta_prompt(
|
||||
"block-level ops) so the response always parses as valid JSON."
|
||||
)
|
||||
|
||||
task_footer = (
|
||||
return (
|
||||
f"## Topic\n{source_query}\n\n"
|
||||
f"## CURRENT DOCUMENT (apply ops to this; reference section ids as listed)\n"
|
||||
f"```json\n{current_document_json}\n```\n\n"
|
||||
f"## NEW INFORMATION SYNTHESIS (context for how new facts relate to the topic)\n"
|
||||
f"```markdown\n{candidate_markdown}\n```\n\n"
|
||||
f"## SUPPORTING FACTS (new since last refresh — integrate these)\n{facts_block}"
|
||||
f"{budget_hint}\n\n"
|
||||
"## Task\n"
|
||||
"Output a JSON object matching the operations schema. Integrate the new "
|
||||
"supporting facts into CURRENT DOCUMENT. Add, update, or remove content "
|
||||
"as needed. Preserve unchanged sections and blocks by not mentioning them."
|
||||
)
|
||||
input_cap = max_input_tokens if max_input_tokens is not None else _STRUCTURED_DELTA_DEFAULT_MAX_INPUT_TOKENS
|
||||
doc_json, candidate, facts_body, input_truncated = _fit_structured_delta_prompt_parts(
|
||||
source_query=source_query,
|
||||
current_document_json=current_document_json,
|
||||
candidate_markdown=candidate_markdown,
|
||||
facts_block=facts_block,
|
||||
budget_hint=budget_hint,
|
||||
task_footer=task_footer,
|
||||
max_input_tokens=input_cap,
|
||||
)
|
||||
truncation_note = ""
|
||||
if input_truncated:
|
||||
truncation_note = (
|
||||
"\n\n*Note: Document, synthesis, or facts were truncated to fit the model "
|
||||
"context window. Prefer minimal, high-leverage operations.*"
|
||||
)
|
||||
|
||||
return (
|
||||
f"## Topic\n{source_query}\n\n"
|
||||
f"## CURRENT DOCUMENT (apply ops to this; reference section ids as listed)\n"
|
||||
f"```json\n{doc_json}\n```\n\n"
|
||||
f"## NEW INFORMATION SYNTHESIS (context for how new facts relate to the topic)\n"
|
||||
f"```markdown\n{candidate}\n```\n\n"
|
||||
f"## SUPPORTING FACTS (new since last refresh — integrate these)\n{facts_body}"
|
||||
f"{budget_hint}{truncation_note}\n\n"
|
||||
f"{task_footer}"
|
||||
)
|
||||
|
||||
|
||||
DELTA_SYSTEM_PROMPT = """You are performing a surgical delta update to an existing mental model document.
|
||||
|
||||
@@ -31,20 +31,8 @@ class LLMToolCallResult(BaseModel):
|
||||
content: str | None = Field(default=None, description="Text content if any")
|
||||
tool_calls: list[LLMToolCall] = Field(default_factory=list, description="Tool calls requested by the LLM")
|
||||
finish_reason: str | None = Field(default=None, description="Reason the LLM stopped: 'stop', 'tool_calls', etc.")
|
||||
input_tokens: int = Field(
|
||||
default=0,
|
||||
description="Input tokens used in this call (includes any cached prefix tokens reported by the provider)",
|
||||
)
|
||||
output_tokens: int = Field(
|
||||
default=0, description="Visible output tokens used in this call (excludes reasoning/thoughts)"
|
||||
)
|
||||
cached_tokens: int = Field(
|
||||
default=0, description="Cached prefix tokens, when reported by the provider. Subset of input_tokens."
|
||||
)
|
||||
thoughts_tokens: int = Field(
|
||||
default=0,
|
||||
description="Reasoning/thinking tokens. Billed at the output rate by some providers but not part of visible output.",
|
||||
)
|
||||
input_tokens: int = Field(default=0, description="Input tokens used in this call")
|
||||
output_tokens: int = Field(default=0, description="Output tokens used in this call")
|
||||
|
||||
|
||||
class ToolCallTrace(BaseModel):
|
||||
@@ -103,18 +91,9 @@ class TokenUsage(BaseModel):
|
||||
)
|
||||
|
||||
input_tokens: int = Field(default=0, description="Number of input/prompt tokens consumed")
|
||||
output_tokens: int = Field(
|
||||
default=0, description="Number of visible output/completion tokens generated (excludes reasoning/thoughts)"
|
||||
)
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output, excludes thoughts)")
|
||||
output_tokens: int = Field(default=0, description="Number of output/completion tokens generated")
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output)")
|
||||
cached_tokens: int = Field(default=0, description="Cached/cache-read prompt tokens, when reported by the provider")
|
||||
thoughts_tokens: int = Field(
|
||||
default=0,
|
||||
description=(
|
||||
"Reasoning/thinking tokens generated by the model. Billed at the output rate by some providers "
|
||||
"(e.g. Gemini 2.5+ family) but not surfaced in the visible response."
|
||||
),
|
||||
)
|
||||
|
||||
def __add__(self, other: "TokenUsage") -> "TokenUsage":
|
||||
"""Allow aggregating token usage from multiple calls."""
|
||||
@@ -123,38 +102,9 @@ class TokenUsage(BaseModel):
|
||||
output_tokens=self.output_tokens + other.output_tokens,
|
||||
total_tokens=self.total_tokens + other.total_tokens,
|
||||
cached_tokens=self.cached_tokens + other.cached_tokens,
|
||||
thoughts_tokens=self.thoughts_tokens + other.thoughts_tokens,
|
||||
)
|
||||
|
||||
|
||||
class ExtractedFact(BaseModel):
|
||||
"""A single candidate fact produced by dry-run extraction (no resolution/links/persistence).
|
||||
|
||||
A deliberate subset of the persisted memory-unit shape — only the fields a fresh extraction
|
||||
yields. Storage/consolidation/curation fields (id, document_id, chunk_id, proof_count, state, …)
|
||||
are omitted because nothing is stored. Entities are raw, unresolved names.
|
||||
"""
|
||||
|
||||
text: str = Field(description="The extracted fact text.")
|
||||
fact_type: str = Field(description="Perspective classification: 'world' or 'experience'.")
|
||||
occurred_start: str | None = Field(default=None, description="ISO timestamp the fact's event started, if dated.")
|
||||
occurred_end: str | None = Field(default=None, description="ISO timestamp the fact's event ended, if dated.")
|
||||
entities: list[str] = Field(
|
||||
default_factory=list, description="Raw (unresolved) entity names mentioned in the fact."
|
||||
)
|
||||
|
||||
|
||||
class DryRunExtractionResult(BaseModel):
|
||||
"""Result of dry-run fact extraction: candidate facts plus aggregated LLM token usage."""
|
||||
|
||||
facts: list[ExtractedFact] = Field(
|
||||
default_factory=list, description="Candidate facts the retain step would extract."
|
||||
)
|
||||
usage: TokenUsage = Field(
|
||||
default_factory=TokenUsage, description="Aggregated token usage across the extraction LLM calls."
|
||||
)
|
||||
|
||||
|
||||
class DispositionTraits(BaseModel):
|
||||
"""
|
||||
Disposition traits for a memory bank.
|
||||
@@ -172,47 +122,6 @@ class DispositionTraits(BaseModel):
|
||||
model_config = ConfigDict(json_schema_extra={"example": {"skepticism": 3, "literalism": 3, "empathy": 3}})
|
||||
|
||||
|
||||
class RecallScores(BaseModel):
|
||||
"""Per-result recall scores from different stages of the pipeline.
|
||||
|
||||
``final`` is the value results are ranked by. The others are diagnostic and
|
||||
can be filtered on via the recall ``min_scores`` request parameter. ``semantic``
|
||||
and ``keyword`` are the raw per-strategy retrieval scores (``None`` when that
|
||||
strategy did not surface this result); ``reranker`` is the cross-encoder's
|
||||
normalized relevance.
|
||||
"""
|
||||
|
||||
final: float = Field(description="Final ranking score (combined reranker + recency/temporal/proof boosts)")
|
||||
reranker: float | None = Field(
|
||||
default=None,
|
||||
description="Cross-encoder relevance, normalized 0-1. None when the reranker is a passthrough (rrf/interleave modes).",
|
||||
)
|
||||
semantic: float | None = Field(
|
||||
default=None, description="Vector cosine similarity (0-1). None if this result was not surfaced semantically."
|
||||
)
|
||||
keyword: float | None = Field(
|
||||
default=None,
|
||||
description="Keyword/full-text (BM25) score (>= 0, unbounded). None if this result was not surfaced by keyword search.",
|
||||
)
|
||||
|
||||
|
||||
class MinScores(BaseModel):
|
||||
"""Optional per-stage score floors for recall (all inclusive, AND-ed).
|
||||
|
||||
``semantic`` and ``keyword`` are **retrieval-level** cutoffs pushed into the SQL
|
||||
arms (overriding the global ``semantic_min_similarity`` / ``bm25_min_score``
|
||||
config for this request), so they prune weak matches before fusion. ``reranker``
|
||||
and ``final`` are **post-query** filters applied to the scored results after
|
||||
reranking. Any field left None imposes no floor; all-None (the default) means
|
||||
no score filtering.
|
||||
"""
|
||||
|
||||
semantic: float | None = Field(default=None, description="Retrieval-level: minimum vector similarity (0-1).")
|
||||
keyword: float | None = Field(default=None, description="Retrieval-level: minimum keyword/full-text (BM25) score.")
|
||||
reranker: float | None = Field(default=None, description="Post-query: minimum normalized reranker score (0-1).")
|
||||
final: float | None = Field(default=None, description="Post-query: minimum final ranking score.")
|
||||
|
||||
|
||||
class MemoryFact(BaseModel):
|
||||
"""
|
||||
A single memory fact returned by search or think operations.
|
||||
@@ -243,7 +152,7 @@ class MemoryFact(BaseModel):
|
||||
|
||||
id: str = Field(description="Unique identifier for the memory fact")
|
||||
text: str = Field(description="The actual text content of the memory")
|
||||
fact_type: str = Field(description="Type of fact: 'world', 'experience', or 'observation'")
|
||||
fact_type: str = Field(description="Type of fact: 'world', 'experience', 'opinion', or 'observation'")
|
||||
entities: list[str] | None = Field(None, description="Entity names mentioned in this fact")
|
||||
context: str | None = Field(None, description="Additional context for the memory")
|
||||
occurred_start: str | None = Field(None, description="ISO format date when the event started occurring")
|
||||
@@ -272,10 +181,6 @@ class MemoryFact(BaseModel):
|
||||
None,
|
||||
description="IDs of source facts this observation was derived from (observation type only, when source_facts is enabled)",
|
||||
)
|
||||
scores: RecallScores | None = Field(
|
||||
None,
|
||||
description="Recall scores from each pipeline stage (final/reranker/semantic/keyword). Not returned for source facts.",
|
||||
)
|
||||
|
||||
|
||||
class ChunkInfo(BaseModel):
|
||||
@@ -374,8 +279,7 @@ class ReflectResult(BaseModel):
|
||||
],
|
||||
"experience": [],
|
||||
"opinion": [],
|
||||
"observation": [],
|
||||
"mental-models": [],
|
||||
"mental_models": [],
|
||||
"directives": [
|
||||
{
|
||||
"id": "directive-123",
|
||||
@@ -392,7 +296,7 @@ class ReflectResult(BaseModel):
|
||||
|
||||
text: str = Field(description="The formulated answer text")
|
||||
based_on: dict[str, Any] = Field(
|
||||
description="Facts used to formulate the answer, organized by type (world, experience, observation, mental-models, directives)"
|
||||
description="Facts used to formulate the answer, organized by type (world, experience, mental_models, directives)"
|
||||
)
|
||||
structured_output: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
|
||||
@@ -14,7 +14,6 @@ from typing import Any, Literal, cast
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, create_model, field_validator
|
||||
|
||||
from ..llm_interface import ProviderRateLimitResetError
|
||||
from ..llm_wrapper import LLMConfig, OutputTooLongError, sanitize_llm_output
|
||||
from ..operation_metadata import RetainExtractionErrors
|
||||
from ..response_models import TokenUsage
|
||||
@@ -193,7 +192,7 @@ class ExtractedFact(BaseModel):
|
||||
occurred_start: str | None = Field(default=None, description="ISO timestamp for events")
|
||||
occurred_end: str | None = Field(default=None, description="ISO timestamp for event end")
|
||||
fact_type: Literal["world", "assistant"] = Field(
|
||||
description="'world' = objective/external facts, including user preferences, rules, corrections, and constraints even when stated during a conversation. 'assistant' = actions, experiences, or observations the assistant/agent actually performed."
|
||||
description="'world' = objective/external facts. 'assistant' = first-person actions, experiences, or observations by the speaker."
|
||||
)
|
||||
entities: list[Entity] | None = Field(default=None, description="People, places, concepts")
|
||||
causal_relations: list[FactCausalRelation] | None = Field(
|
||||
@@ -296,7 +295,7 @@ class ExtractedFactVerbose(BaseModel):
|
||||
)
|
||||
|
||||
fact_type: Literal["world", "assistant"] = Field(
|
||||
description="'world' = objective/external facts about the user, other people, events, general knowledge, preferences, rules, corrections, or constraints. 'assistant' = actions, experiences, or observations the assistant/agent actually performed (e.g., 'I changed X', 'I discovered Y')."
|
||||
description="'world' = objective/external facts about other people, events, general knowledge. 'assistant' = first-person actions, experiences, or observations by the speaker (e.g., 'I changed X', 'I discovered Y')."
|
||||
)
|
||||
|
||||
entities: list[Entity] | None = Field(
|
||||
@@ -346,7 +345,7 @@ class ExtractedFactNoCausal(BaseModel):
|
||||
occurred_start: str | None = Field(default=None, description="WHEN the event happened (ISO timestamp).")
|
||||
occurred_end: str | None = Field(default=None, description="WHEN the event ended (ISO timestamp).")
|
||||
fact_type: Literal["world", "assistant"] = Field(
|
||||
description="'world' = about the user/others, including user preferences, rules, corrections, and constraints. 'assistant' = actions or experiences the assistant/agent actually performed."
|
||||
description="'world' = about the user/others. 'assistant' = experience with assistant."
|
||||
)
|
||||
entities: list[Entity] | None = Field(
|
||||
default=None,
|
||||
@@ -421,14 +420,19 @@ _RECURSIVE_TEXT_SEPARATORS = [
|
||||
"", # Characters (last resort)
|
||||
]
|
||||
|
||||
# A single structured unit (a JSONL line or a conversation turn) is kept whole
|
||||
# even when it overflows the budget — but only up to this multiple. Beyond it,
|
||||
# the unit is split as text rather than handed to the LLM wildly over budget
|
||||
# (the extractor has no second re-chunk pass; an oversized chunk just errors).
|
||||
_CHUNK_OVERFLOW_FACTOR = 1.5
|
||||
|
||||
|
||||
def _split_oversized_unit(text: str, max_chars: int) -> list[str]:
|
||||
"""Sentence-aware split of a single unit that overflowed the budget.
|
||||
|
||||
Used when one JSONL line / conversation turn is so large it can't be kept
|
||||
whole within the configured structured-chunk limit. The resulting fragments
|
||||
are no longer valid JSON, but the fact extractor treats every chunk as plain
|
||||
text.
|
||||
whole within ``_CHUNK_OVERFLOW_FACTOR``. The resulting fragments are no
|
||||
longer valid JSON, but the fact extractor treats every chunk as plain text.
|
||||
"""
|
||||
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
||||
|
||||
@@ -442,26 +446,18 @@ def _split_oversized_unit(text: str, max_chars: int) -> list[str]:
|
||||
return splitter.split_text(text)
|
||||
|
||||
|
||||
def chunk_text(text: str, max_chars: int, structured_chunk_size: int | None = None) -> list[str]:
|
||||
def chunk_text(text: str, max_chars: int) -> list[str]:
|
||||
"""
|
||||
Split text into chunks, preserving conversation structure when possible.
|
||||
|
||||
For JSON conversation arrays (user/assistant turns) and JSONL (newline-delimited
|
||||
JSON objects), splits at turn/line boundaries so no object is split across chunks.
|
||||
A single turn/line that overflows ``max_chars`` is kept whole only up to
|
||||
``structured_chunk_size``. When unset, that limit defaults to ``max_chars``.
|
||||
For plain text, uses sentence-aware splitting.
|
||||
|
||||
The result is idempotent: re-chunking any chunk this returns yields that chunk
|
||||
unchanged. The streaming retain pipeline pre-chunks each document once and then
|
||||
re-chunks every piece during extraction; if a piece re-split, its sub-chunks
|
||||
would inherit one chunk_index and collide on ``chunk_id`` (issue #2301).
|
||||
A single turn/line that overflows is kept whole up to ``_CHUNK_OVERFLOW_FACTOR``×
|
||||
the budget, then split as text. For plain text, uses sentence-aware splitting.
|
||||
|
||||
Args:
|
||||
text: Input text to chunk (plain text, JSON conversation, or JSONL)
|
||||
max_chars: Target maximum characters per chunk
|
||||
structured_chunk_size: Maximum characters for a single JSONL line or
|
||||
conversation turn to keep whole. Defaults to ``max_chars``.
|
||||
max_chars: Maximum characters per chunk (default 120k ≈ 30k tokens)
|
||||
|
||||
Returns:
|
||||
List of text chunks, roughly under max_chars
|
||||
@@ -470,31 +466,17 @@ def chunk_text(text: str, max_chars: int, structured_chunk_size: int | None = No
|
||||
if len(text) <= max_chars:
|
||||
return [text]
|
||||
|
||||
structured_limit = structured_chunk_size if structured_chunk_size is not None else max_chars
|
||||
|
||||
# Try to parse as JSON conversation array
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
if isinstance(parsed, list) and all(isinstance(turn, dict) for turn in parsed):
|
||||
# This looks like a conversation - chunk at turn boundaries
|
||||
return _chunk_conversation(parsed, max_chars)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
parsed = None
|
||||
|
||||
if isinstance(parsed, list) and all(isinstance(turn, dict) for turn in parsed):
|
||||
# This looks like a conversation - chunk at turn boundaries
|
||||
return _chunk_conversation(parsed, max_chars, structured_limit)
|
||||
|
||||
if isinstance(parsed, dict):
|
||||
# A single JSON object — e.g. one JSONL line handed back to the extractor
|
||||
# after the producer already pre-chunked it. It is one structured unit:
|
||||
# keep it whole up to the structured limit, else split it as text within
|
||||
# the chunk budget. Without this, a lone object (one line, so _chunk_jsonl
|
||||
# declines) would fall through to plain-text splitting and re-split a chunk
|
||||
# the producer deliberately kept whole — breaking idempotency (issue #2301).
|
||||
if len(text) <= structured_limit:
|
||||
return [text]
|
||||
return _split_oversized_unit(text, max_chars)
|
||||
pass
|
||||
|
||||
# Try to parse as JSONL (newline-delimited JSON objects, e.g. session logs)
|
||||
jsonl_chunks = _chunk_jsonl(text, max_chars, structured_limit)
|
||||
jsonl_chunks = _chunk_jsonl(text, max_chars)
|
||||
if jsonl_chunks is not None:
|
||||
return jsonl_chunks
|
||||
|
||||
@@ -502,19 +484,20 @@ def chunk_text(text: str, max_chars: int, structured_chunk_size: int | None = No
|
||||
return _split_oversized_unit(text, max_chars)
|
||||
|
||||
|
||||
def _chunk_conversation(turns: list[dict], max_chars: int, structured_limit: int) -> list[str]:
|
||||
def _chunk_conversation(turns: list[dict], max_chars: int) -> list[str]:
|
||||
"""
|
||||
Chunk a conversation array at turn boundaries, preserving complete turns.
|
||||
|
||||
Args:
|
||||
turns: List of conversation turn dicts (with 'role' and 'content' keys)
|
||||
max_chars: Maximum characters per chunk
|
||||
structured_limit: Maximum characters for a single turn to keep whole
|
||||
|
||||
Returns:
|
||||
List of JSON-serialized chunks, each containing complete turns
|
||||
"""
|
||||
|
||||
overflow_limit = int(max_chars * _CHUNK_OVERFLOW_FACTOR)
|
||||
|
||||
chunks = []
|
||||
current_chunk = []
|
||||
current_size = 2 # Account for "[]"
|
||||
@@ -529,16 +512,13 @@ def _chunk_conversation(turns: list[dict], max_chars: int, structured_limit: int
|
||||
for turn in turns:
|
||||
# Estimate size of this turn when serialized (with comma separator)
|
||||
turn_json = json.dumps(turn, ensure_ascii=False)
|
||||
turn_unit_size = len(turn_json)
|
||||
turn_size = turn_unit_size + 1 # +1 for comma
|
||||
turn_size = len(turn_json) + 1 # +1 for comma
|
||||
|
||||
# A turn too large to keep whole even alone: flush, then split it as
|
||||
# text. Fragment within min(structured_limit, max_chars) so no fragment
|
||||
# exceeds the chunk budget — otherwise a downstream re-chunk would split
|
||||
# it again and collide on chunk_id (issue #2301).
|
||||
if turn_unit_size > structured_limit:
|
||||
# text so no chunk runs far over budget (the extractor won't re-chunk).
|
||||
if turn_size > overflow_limit:
|
||||
_flush()
|
||||
chunks.extend(_split_oversized_unit(turn_json, min(structured_limit, max_chars)))
|
||||
chunks.extend(_split_oversized_unit(turn_json, max_chars))
|
||||
continue
|
||||
|
||||
# If adding this turn would exceed limit and we have turns, save current chunk
|
||||
@@ -555,20 +535,18 @@ def _chunk_conversation(turns: list[dict], max_chars: int, structured_limit: int
|
||||
return chunks if chunks else [json.dumps(turns, ensure_ascii=False)]
|
||||
|
||||
|
||||
def _chunk_jsonl(text: str, max_chars: int, structured_limit: int) -> list[str] | None:
|
||||
def _chunk_jsonl(text: str, max_chars: int) -> list[str] | None:
|
||||
"""Chunk newline-delimited JSON (JSONL) at line boundaries.
|
||||
|
||||
Detects JSONL — two or more non-empty lines, each a complete JSON object —
|
||||
and packs whole lines into chunks so no line is split across chunks (multiple
|
||||
short lines may share a chunk). A line that overflows ``max_chars`` is kept
|
||||
whole only up to ``structured_limit``. Returns ``None`` if the input is not
|
||||
JSONL, so the caller falls back to plain-text splitting.
|
||||
short lines may share a chunk). A line that overflows is kept whole up to
|
||||
``_CHUNK_OVERFLOW_FACTOR``× the budget, then split as text. Returns ``None``
|
||||
if the input is not JSONL, so the caller falls back to plain-text splitting.
|
||||
|
||||
Args:
|
||||
text: Input text to inspect/chunk.
|
||||
max_chars: Maximum characters per chunk.
|
||||
structured_limit: Maximum characters for a single JSONL line to
|
||||
keep whole.
|
||||
|
||||
Returns:
|
||||
List of JSONL chunks (lines joined by newline), or ``None`` if not JSONL.
|
||||
@@ -585,6 +563,8 @@ def _chunk_jsonl(text: str, max_chars: int, structured_limit: int) -> list[str]
|
||||
if not isinstance(obj, dict):
|
||||
return None
|
||||
|
||||
overflow_limit = int(max_chars * _CHUNK_OVERFLOW_FACTOR)
|
||||
|
||||
chunks: list[str] = []
|
||||
current_chunk: list[str] = []
|
||||
current_size = 0
|
||||
@@ -597,20 +577,17 @@ def _chunk_jsonl(text: str, max_chars: int, structured_limit: int) -> list[str]
|
||||
current_size = 0
|
||||
|
||||
for line in lines:
|
||||
line_unit_size = len(line)
|
||||
line_size = len(line) + 1 # +1 for the joining newline
|
||||
|
||||
# A line too large to keep whole even alone: flush, then split it as
|
||||
# text. Fragment within min(structured_limit, max_chars) so no fragment
|
||||
# exceeds the chunk budget — otherwise a downstream re-chunk would split
|
||||
# it again and collide on chunk_id (issue #2301).
|
||||
if line_unit_size > structured_limit:
|
||||
# text so no chunk runs far over budget (the extractor won't re-chunk).
|
||||
if line_size > overflow_limit:
|
||||
_flush()
|
||||
chunks.extend(_split_oversized_unit(line, min(structured_limit, max_chars)))
|
||||
chunks.extend(_split_oversized_unit(line, max_chars))
|
||||
continue
|
||||
|
||||
# If adding this line would exceed the limit and we have lines, flush.
|
||||
# A line up to structured_limit is kept whole (a bounded overflow).
|
||||
# A line up to overflow_limit is kept whole (a small, bounded overflow).
|
||||
if current_size + line_size > max_chars and current_chunk:
|
||||
_flush()
|
||||
|
||||
@@ -663,8 +640,8 @@ fact_kind:
|
||||
- "conversation": Ongoing state, preference, trait (no dates)
|
||||
|
||||
fact_type:
|
||||
- "world": Objective/external facts, including the user's preferences, rules, corrections, constraints, plans, traits, or context. These stay "world" even when the user states them during an assistant interaction (e.g., "User prefers browser_navigate over web_search", "User corrected the project deadline").
|
||||
- "assistant": Actions, experiences, or observations the assistant/agent actually performed (e.g., "I changed X", "I discovered Y", "I debugged Z"). Use this for the assistant/agent doing, trying, learning, deciding, recommending, or responding — not merely for user facts mentioned in conversation.
|
||||
- "world": About other people, external events, general knowledge, objective facts
|
||||
- "assistant": First-person actions, experiences, or observations by the speaker/author (e.g., "I changed X", "I discovered Y", "I debugged Z"). Also includes interactions with the user (requests, recommendations). If the narrator describes something they did, tried, learned, or decided — use "assistant".
|
||||
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
TEMPORAL HANDLING
|
||||
@@ -766,7 +743,7 @@ RULES:
|
||||
- Extract all entities (people, places, organizations, objects, concepts).
|
||||
- Extract temporal information (occurred_start, occurred_end, fact_kind, when).
|
||||
- Extract location (where) and people (who).
|
||||
- fact_type: use "world" for user preferences, rules, corrections, constraints, traits, and other objective facts, even when stated during an assistant interaction. Use "assistant" only for actions or experiences the assistant/agent actually performed."""
|
||||
- fact_type: use "world" unless the content is clearly an interaction with the assistant."""
|
||||
|
||||
VERBATIM_FACT_EXTRACTION_PROMPT = _BASE_FACT_EXTRACTION_PROMPT.format(
|
||||
retain_mission_section="{retain_mission_section}",
|
||||
@@ -867,8 +844,8 @@ For CONVERSATIONS (fact_kind="conversation"):
|
||||
FACT TYPE
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
- **world**: User's life, preferences, rules, corrections, constraints, other people, and events (facts that would exist without this conversation)
|
||||
- **assistant**: Actions or experiences the assistant/agent actually performed while helping the user (requests, recommendations, help)
|
||||
- **world**: User's life, other people, events (would exist without this conversation)
|
||||
- **assistant**: Interactions with assistant (requests, recommendations, help)
|
||||
⚠️ CRITICAL for assistant facts: ALWAYS capture the user's request/question in the fact!
|
||||
Include: what the user asked, what problem they wanted solved, what context they provided
|
||||
|
||||
@@ -1203,17 +1180,9 @@ def _build_request_body(llm_config, config, prompt: str, user_message: str, resp
|
||||
request_body = {
|
||||
"model": llm_config.model,
|
||||
"messages": [{"role": "system", "content": prompt}, {"role": "user", "content": user_message}],
|
||||
"temperature": 0.1,
|
||||
}
|
||||
|
||||
# Honour the configured retain temperature. ``None`` omits the parameter
|
||||
# entirely (for models like Azure GPT-5.5 that reject explicit temperatures),
|
||||
# mirroring LLMProvider.call, which drops temperature when it is None. The
|
||||
# batch path builds the request body directly instead of going through
|
||||
# LLMProvider.call (#2469 only de-hardcoded the streaming path), so it must
|
||||
# apply the same rule here.
|
||||
if config.llm_temperature_retain is not None:
|
||||
request_body["temperature"] = config.llm_temperature_retain
|
||||
|
||||
# Add max_completion_tokens if configured
|
||||
if config.retain_max_completion_tokens:
|
||||
request_body["max_completion_tokens"] = config.retain_max_completion_tokens
|
||||
@@ -1322,7 +1291,7 @@ async def _extract_facts_from_chunk(
|
||||
messages=[{"role": "system", "content": prompt}, {"role": "user", "content": user_message}],
|
||||
response_format=response_schema,
|
||||
scope="retain_extract_facts",
|
||||
temperature=config.llm_temperature_retain,
|
||||
temperature=0.1,
|
||||
max_completion_tokens=config.retain_max_completion_tokens,
|
||||
max_retries=llm_max_retries,
|
||||
initial_backoff=initial_backoff,
|
||||
@@ -1769,11 +1738,7 @@ async def extract_facts_from_text(
|
||||
- chunks: List of tuples (chunk_text, fact_count) for each chunk
|
||||
- usage: Aggregated token usage across all LLM calls
|
||||
"""
|
||||
chunks = chunk_text(
|
||||
text,
|
||||
max_chars=config.retain_chunk_size,
|
||||
structured_chunk_size=config.retain_structured_chunk_size,
|
||||
)
|
||||
chunks = chunk_text(text, max_chars=config.retain_chunk_size)
|
||||
|
||||
# Log chunk count before starting LLM requests
|
||||
total_chars = sum(len(c) for c in chunks)
|
||||
@@ -1822,28 +1787,10 @@ async def extract_facts_from_text(
|
||||
total_usage = total_usage + chunk_usage
|
||||
|
||||
if failed_chunks:
|
||||
# Include the exception message — not just the type — so operators
|
||||
# can tell a structured-JSON parse failure apart from a rate limit
|
||||
# apart from a network 5xx, all of which can surface as the same
|
||||
# exception types. The error_message we propagate to the
|
||||
# async_operations row is the only inspection surface a worker-side
|
||||
# failure leaves behind, and a bare "chunk 0: RuntimeError" is not
|
||||
# actionable.
|
||||
failed_summary = ", ".join(f"chunk {idx}: {type(err).__name__}: {err}" for idx, err in failed_chunks[:5])
|
||||
quota_errors = [err for _, err in failed_chunks if isinstance(err, ProviderRateLimitResetError)]
|
||||
if quota_errors and len(quota_errors) == len(failed_chunks):
|
||||
retry_at = max(err.retry_at for err in quota_errors)
|
||||
raise ProviderRateLimitResetError(
|
||||
retry_at=retry_at,
|
||||
message=(
|
||||
f"Fact extraction deferred by provider quota: {len(failed_chunks)}/{len(chunks)} chunks failed. "
|
||||
f"First failures: {failed_summary}. Provider detail: {quota_errors[0]}"
|
||||
),
|
||||
) from quota_errors[0]
|
||||
|
||||
# Fail the entire retain — partial extraction is not acceptable.
|
||||
# All successfully extracted facts are discarded because the transaction
|
||||
# hasn't committed yet. The worker poller will retry the entire task.
|
||||
failed_summary = ", ".join(f"chunk {idx}: {type(err).__name__}" for idx, err in failed_chunks[:5])
|
||||
raise RuntimeError(
|
||||
f"Fact extraction failed: {len(failed_chunks)}/{len(chunks)} chunks failed. "
|
||||
f"First failures: {failed_summary}"
|
||||
@@ -1974,11 +1921,7 @@ async def extract_facts_from_contents_batch_api(
|
||||
prompt, response_schema = _build_extraction_prompt_and_schema(config)
|
||||
|
||||
for content_index, item in enumerate(contents):
|
||||
chunks = chunk_text(
|
||||
item.content,
|
||||
max_chars=config.retain_chunk_size,
|
||||
structured_chunk_size=config.retain_structured_chunk_size,
|
||||
)
|
||||
chunks = chunk_text(item.content, max_chars=config.retain_chunk_size)
|
||||
|
||||
for chunk_index_in_content, chunk in enumerate(chunks):
|
||||
all_chunks_info.append((chunk, content_index, chunk_index_in_content, item.event_date, item.context))
|
||||
@@ -2399,11 +2342,7 @@ def _extract_facts_chunks(
|
||||
global_chunk_idx = 0
|
||||
|
||||
for content_index, content in enumerate(contents):
|
||||
chunks = chunk_text(
|
||||
content.content,
|
||||
config.retain_chunk_size,
|
||||
structured_chunk_size=config.retain_structured_chunk_size,
|
||||
)
|
||||
chunks = chunk_text(content.content, config.retain_chunk_size)
|
||||
for chunk in chunks:
|
||||
chunks_metadata.append(
|
||||
ChunkMetadata(
|
||||
|
||||
@@ -89,29 +89,7 @@ async def _fire_memory_defense_webhook(
|
||||
if webhook_manager is None:
|
||||
return
|
||||
try:
|
||||
from ...webhooks import (
|
||||
MemoryDefenseEventData,
|
||||
MemoryDefenseHit,
|
||||
WebhookEvent,
|
||||
WebhookEventType,
|
||||
)
|
||||
|
||||
# Translate per-match raw dicts on the decision into MemoryDefenseHit
|
||||
# entries on the wire. The decision's hits list is already fingerprinted
|
||||
# by apply_redaction (the raw value never lands in hits, by contract),
|
||||
# so this is purely a shape conversion. None when no per-hit data is
|
||||
# available so receivers can distinguish "no preview info" from
|
||||
# "scanned, nothing matched" (the latter wouldn't be a webhook delivery
|
||||
# in the first place).
|
||||
decision_hits = getattr(decision, "hits", None) or []
|
||||
hits: list[MemoryDefenseHit] | None = [
|
||||
MemoryDefenseHit(
|
||||
detector=str(h.get("detector") or ""),
|
||||
preview=str(h.get("preview") or ""),
|
||||
)
|
||||
for h in decision_hits
|
||||
if h.get("detector") and h.get("preview")
|
||||
] or None
|
||||
from ...webhooks import MemoryDefenseEventData, WebhookEvent, WebhookEventType
|
||||
|
||||
event = WebhookEvent(
|
||||
event=WebhookEventType.MEMORY_DEFENSE_TRIGGERED,
|
||||
@@ -125,17 +103,6 @@ async def _fire_memory_defense_webhook(
|
||||
document_id=document_id,
|
||||
matched_types=decision.matched_types or None,
|
||||
message=decision.message or None,
|
||||
hits=hits,
|
||||
# Optional SIEM-enrichment fields populated by downstream
|
||||
# extensions (e.g. hindsight-cloud's _CloudDefenseDecision
|
||||
# subclass). Read via getattr so OSS doesn't need to know
|
||||
# about extension subclasses. Combined with the manager's
|
||||
# exclude_none serialization, missing values stay absent
|
||||
# from the wire entirely rather than appearing as null.
|
||||
severity=getattr(decision, "severity", None),
|
||||
api_key_name=getattr(decision, "api_key_name", None),
|
||||
memory_unit_id=getattr(decision, "memory_unit_id", None),
|
||||
receipt_uri=getattr(decision, "receipt_uri", None),
|
||||
),
|
||||
)
|
||||
await webhook_manager.fire_event_with_conn(event, conn, schema=schema)
|
||||
@@ -834,28 +801,6 @@ async def retain_batch(
|
||||
if first.get("tags"):
|
||||
existing_content["tags"] = first["tags"]
|
||||
contents_dicts = [existing_content, *contents_dicts]
|
||||
# Merge JSON arrays to keep original_text valid (#2409).
|
||||
# Without this, combined_content joins items with "\n", producing
|
||||
# "[...]\n[...]" which is not valid JSON. On the next append cycle
|
||||
# chunk_text() fails to parse it and falls through to sentence-
|
||||
# boundary text splitting, breaking speaker attribution.
|
||||
try:
|
||||
_merged = []
|
||||
for _item in contents_dicts:
|
||||
_parsed = json.loads(_item.get("content", ""))
|
||||
if isinstance(_parsed, list) and all(isinstance(_e, dict) for _e in _parsed):
|
||||
_merged.extend(_parsed)
|
||||
else:
|
||||
_merged = None
|
||||
break
|
||||
if _merged is not None:
|
||||
contents_dicts = [{"content": json.dumps(_merged, ensure_ascii=False)}]
|
||||
if first.get("context"):
|
||||
contents_dicts[0]["context"] = first["context"]
|
||||
if first.get("tags"):
|
||||
contents_dicts[0]["tags"] = first["tags"]
|
||||
except (json.JSONDecodeError, ValueError, TypeError):
|
||||
pass
|
||||
# Rebuild contents list to match
|
||||
contents = _build_contents(contents_dicts, document_tags)
|
||||
log_buffer.append(
|
||||
@@ -920,15 +865,10 @@ async def retain_batch(
|
||||
# retain code paths.
|
||||
chunk_batch_size = getattr(config, "retain_chunk_batch_size", 100)
|
||||
chunk_size = getattr(config, "retain_chunk_size", 3000)
|
||||
structured_chunk_size = getattr(config, "retain_structured_chunk_size", None)
|
||||
all_pre_chunks: list[str] = []
|
||||
chunk_to_content: list[int] = [] # maps chunk index -> index into contents
|
||||
for content_idx, content in enumerate(contents):
|
||||
content_chunks = fact_extraction.chunk_text(
|
||||
content.content,
|
||||
chunk_size,
|
||||
structured_chunk_size=structured_chunk_size,
|
||||
)
|
||||
content_chunks = fact_extraction.chunk_text(content.content, chunk_size)
|
||||
all_pre_chunks.extend(content_chunks)
|
||||
chunk_to_content.extend([content_idx] * len(content_chunks))
|
||||
|
||||
@@ -1637,19 +1577,8 @@ async def _streaming_retain_batch(
|
||||
# Check if facts are already committed (recovery from previous crash).
|
||||
# If so, skip extraction+writes and jump straight to final ANN pass.
|
||||
# ---------------------------------------------------------------------------
|
||||
# Only the call that starts a document at chunk 0 may take the whole-document
|
||||
# skip. When an oversized single item is split into several sequential
|
||||
# sub-batches that SHARE one document_id AND one operation_id (see
|
||||
# _split_contents_into_sub_batches), the first sub-batch commits its chunks
|
||||
# and stamps effective_doc_id into result_metadata.facts_committed_document_ids.
|
||||
# Without the offset gate, every later sub-batch (chunk_index_offset > 0) would
|
||||
# then see its own document already "committed" and skip extraction, dropping
|
||||
# all chunks past the first slice. A non-zero offset inherently means this call
|
||||
# continues a document another sub-batch already started, so it must always do
|
||||
# its work — crash-safety for those chunks still comes from the per-chunk hash
|
||||
# recovery (existing_chunk_hashes) below.
|
||||
facts_already_committed = False
|
||||
if operation_id and chunk_index_offset == 0:
|
||||
if operation_id:
|
||||
try:
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
row = await conn.fetchrow(
|
||||
@@ -2319,14 +2248,9 @@ def _chunk_contents_for_delta(contents: list[RetainContent], config) -> dict[int
|
||||
"""
|
||||
result = {}
|
||||
global_chunk_idx = 0
|
||||
chunk_size = getattr(config, "retain_chunk_size", 3000)
|
||||
structured_chunk_size = getattr(config, "retain_structured_chunk_size", None)
|
||||
for content in contents:
|
||||
chunks = fact_extraction.chunk_text(
|
||||
content.content,
|
||||
chunk_size,
|
||||
structured_chunk_size=structured_chunk_size,
|
||||
)
|
||||
chunk_size = getattr(config, "retain_chunk_size", 3000)
|
||||
chunks = fact_extraction.chunk_text(content.content, chunk_size)
|
||||
for chunk_text in chunks:
|
||||
result[global_chunk_idx] = chunk_text
|
||||
global_chunk_idx += 1
|
||||
|
||||
@@ -24,9 +24,7 @@ class RetainContentDict(TypedDict, total=False):
|
||||
tags: Visibility scope tags for this content item (optional)
|
||||
observation_scopes: How to scope observations for consolidation (optional).
|
||||
"per_tag" runs one pass per individual tag; "combined" (default) runs a
|
||||
single pass with all tags; "shared" runs a single pass over one global,
|
||||
untagged scope so memories consolidate together regardless of tags;
|
||||
a list[list[str]] specifies exact passes.
|
||||
single pass with all tags; a list[list[str]] specifies exact passes.
|
||||
update_mode: How to handle existing documents with the same document_id (optional).
|
||||
"replace" (default) deletes old data and reprocesses. "append" concatenates
|
||||
new content to the existing document and reprocesses.
|
||||
@@ -40,7 +38,7 @@ class RetainContentDict(TypedDict, total=False):
|
||||
entities: list[dict[str, str]] # [{"text": "...", "type": "..."}]
|
||||
tags: list[str] # Visibility scope tags
|
||||
observation_scopes: (
|
||||
Literal["per_tag", "combined", "all_combinations", "shared"] | list[list[str]]
|
||||
Literal["per_tag", "combined", "all_combinations"] | list[list[str]]
|
||||
) # Observation scopes for consolidation
|
||||
update_mode: Literal["replace", "append"]
|
||||
|
||||
@@ -59,7 +57,7 @@ class RetainContent:
|
||||
metadata: dict[str, str] = field(default_factory=dict)
|
||||
entities: list[dict[str, str]] = field(default_factory=list) # User-provided entities
|
||||
tags: list[str] = field(default_factory=list) # Visibility scope tags
|
||||
observation_scopes: Literal["per_tag", "combined", "all_combinations", "shared"] | list[list[str]] | None = (
|
||||
observation_scopes: Literal["per_tag", "combined", "all_combinations"] | list[list[str]] | None = (
|
||||
None # Observation scopes
|
||||
)
|
||||
|
||||
@@ -126,7 +124,7 @@ class ExtractedFact:
|
||||
mentioned_at: datetime | None = None
|
||||
metadata: dict[str, str] = field(default_factory=dict)
|
||||
tags: list[str] = field(default_factory=list) # Visibility scope tags
|
||||
observation_scopes: Literal["per_tag", "combined", "all_combinations", "shared"] | list[list[str]] | None = (
|
||||
observation_scopes: Literal["per_tag", "combined", "all_combinations"] | list[list[str]] | None = (
|
||||
None # Observation scopes
|
||||
)
|
||||
|
||||
@@ -178,7 +176,7 @@ class ProcessedFact:
|
||||
tags: list[str] = field(default_factory=list)
|
||||
|
||||
# Observation scopes for consolidation
|
||||
observation_scopes: Literal["per_tag", "combined", "all_combinations", "shared"] | list[list[str]] | None = None
|
||||
observation_scopes: Literal["per_tag", "combined", "all_combinations"] | list[list[str]] | None = None
|
||||
|
||||
@property
|
||||
def is_duplicate(self) -> bool:
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
Helper functions for hybrid search (semantic + BM25 + graph).
|
||||
"""
|
||||
|
||||
from .types import ArmScores, MergedCandidate, RetrievalResult
|
||||
from .types import MergedCandidate, RetrievalResult
|
||||
|
||||
|
||||
def cap_per_source(results: list[RetrievalResult], cap: int) -> list[RetrievalResult]:
|
||||
@@ -51,7 +51,6 @@ def reciprocal_rank_fusion(result_lists: list[list[RetrievalResult]], k: int = 6
|
||||
rrf_scores = {}
|
||||
source_ranks = {} # Track rank from each source for each doc_id
|
||||
all_retrievals = {} # Store the actual RetrievalResult (use first occurrence)
|
||||
arm_scores: dict[str, ArmScores] = {} # doc_id -> raw per-strategy scores across arms
|
||||
|
||||
source_names = ["semantic", "bm25", "graph", "temporal"]
|
||||
|
||||
@@ -80,29 +79,17 @@ def reciprocal_rank_fusion(result_lists: list[list[RetrievalResult]], k: int = 6
|
||||
if doc_id not in rrf_scores:
|
||||
rrf_scores[doc_id] = 0.0
|
||||
source_ranks[doc_id] = {}
|
||||
arm_scores[doc_id] = ArmScores()
|
||||
|
||||
rrf_scores[doc_id] += 1.0 / (k + rank)
|
||||
source_ranks[doc_id][f"{source_name}_rank"] = rank
|
||||
|
||||
# Capture this arm's raw score for the doc (the merged RetrievalResult
|
||||
# below keeps only the first arm's score, so record each arm here).
|
||||
if source_name == "semantic" and retrieval.similarity is not None:
|
||||
arm_scores[doc_id].semantic = retrieval.similarity
|
||||
elif source_name == "bm25" and retrieval.bm25_score is not None:
|
||||
arm_scores[doc_id].keyword = retrieval.bm25_score
|
||||
|
||||
# Combine into final results with metadata
|
||||
merged_results = []
|
||||
for rrf_rank, (doc_id, rrf_score) in enumerate(
|
||||
sorted(rrf_scores.items(), key=lambda x: x[1], reverse=True), start=1
|
||||
):
|
||||
merged_candidate = MergedCandidate(
|
||||
retrieval=all_retrievals[doc_id],
|
||||
rrf_score=rrf_score,
|
||||
rrf_rank=rrf_rank,
|
||||
source_ranks=source_ranks[doc_id],
|
||||
arm_scores=arm_scores[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)
|
||||
|
||||
@@ -131,7 +118,6 @@ def interleave_fusion(result_lists: list[list[RetrievalResult]]) -> list[MergedC
|
||||
source_names = ["semantic", "bm25", "graph", "temporal"]
|
||||
source_ranks: dict[str, dict[str, int]] = {}
|
||||
all_retrievals: dict[str, RetrievalResult] = {}
|
||||
arm_scores: dict[str, ArmScores] = {}
|
||||
|
||||
for source_idx, results in enumerate(result_lists):
|
||||
source_name = source_names[source_idx] if source_idx < len(source_names) else f"source_{source_idx}"
|
||||
@@ -143,11 +129,6 @@ def interleave_fusion(result_lists: list[list[RetrievalResult]]) -> list[MergedC
|
||||
doc_id = retrieval.id
|
||||
all_retrievals.setdefault(doc_id, retrieval)
|
||||
source_ranks.setdefault(doc_id, {})[f"{source_name}_rank"] = rank
|
||||
arm = arm_scores.setdefault(doc_id, ArmScores())
|
||||
if source_name == "semantic" and retrieval.similarity is not None:
|
||||
arm.semantic = retrieval.similarity
|
||||
elif source_name == "bm25" and retrieval.bm25_score is not None:
|
||||
arm.keyword = retrieval.bm25_score
|
||||
|
||||
# Round-robin pick across arms in priority order: all #1s, then all #2s, ...
|
||||
ordered_ids: list[str] = []
|
||||
@@ -170,7 +151,6 @@ def interleave_fusion(result_lists: list[list[RetrievalResult]]) -> list[MergedC
|
||||
rrf_score=float(n - pos),
|
||||
rrf_rank=pos + 1,
|
||||
source_ranks=source_ranks[doc_id],
|
||||
arm_scores=arm_scores[doc_id],
|
||||
)
|
||||
for pos, doc_id in enumerate(ordered_ids)
|
||||
]
|
||||
|
||||
@@ -251,10 +251,8 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
result.activation = row["score"]
|
||||
results.append(result)
|
||||
|
||||
# filter_results_by_tags is a no-op when no filter applies (tags falsy and not
|
||||
# the exact-empty/global scope), so call it unconditionally — gating on `if tags:`
|
||||
# would skip the untagged-only filter for tags=[] + tags_match="exact".
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
if tags:
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
if tag_groups:
|
||||
results = filter_results_by_tag_groups(results, tag_groups)
|
||||
|
||||
@@ -16,44 +16,6 @@ _RECENCY_ALPHA: float = 0.2
|
||||
_TEMPORAL_ALPHA: float = 0.2
|
||||
_PROOF_COUNT_ALPHA: float = 0.1 # Conservative: max ±5% for evidence strength
|
||||
|
||||
# Recency decay: maps a memory's age (days) onto a freshness signal in [0, 1]
|
||||
# where 0.5 is neutral (no boost). The signal is then folded into the
|
||||
# multiplicative recency_boost via `1 + recency_alpha * (recency - 0.5)`.
|
||||
#
|
||||
# "linear" — straight line from 1.0 (today) to a floor of 0.1, reaching
|
||||
# the floor at `linear_window_days`. The historical default.
|
||||
# "exponential" — 0.5 ** (days_ago / halflife_days). The half-life is the age
|
||||
# at which the signal is exactly neutral (0.5): younger
|
||||
# memories are boosted, older ones penalised, with a smooth
|
||||
# asymptote toward 0 (no hard cutoff).
|
||||
# "none" — always neutral (0.5), disabling the recency boost entirely.
|
||||
# The validated set of names lives in config.RECENCY_DECAY_FUNCTIONS.
|
||||
_RECENCY_DECAY_FUNCTION: str = "linear"
|
||||
_RECENCY_DECAY_LINEAR_WINDOW_DAYS: float = 365.0
|
||||
_RECENCY_DECAY_HALFLIFE_DAYS: float = 90.0
|
||||
|
||||
|
||||
def compute_recency_decay(
|
||||
days_ago: float,
|
||||
function: str = _RECENCY_DECAY_FUNCTION,
|
||||
linear_window_days: float = _RECENCY_DECAY_LINEAR_WINDOW_DAYS,
|
||||
halflife_days: float = _RECENCY_DECAY_HALFLIFE_DAYS,
|
||||
) -> float:
|
||||
"""Map a memory's age in days to a freshness signal in [0, 1] (neutral 0.5).
|
||||
|
||||
Future-dated memories (negative ``days_ago``) clamp to the maximum freshness
|
||||
so they are never penalised. See ``RECENCY_DECAY_FUNCTIONS`` for the shapes.
|
||||
"""
|
||||
if function == "none":
|
||||
return 0.5
|
||||
if function == "exponential":
|
||||
if halflife_days <= 0:
|
||||
return 0.5
|
||||
return min(1.0, 0.5 ** (days_ago / halflife_days))
|
||||
# "linear" (default): straight decay to a 0.1 floor over the window.
|
||||
window = linear_window_days if linear_window_days > 0 else _RECENCY_DECAY_LINEAR_WINDOW_DAYS
|
||||
return max(0.1, min(1.0, 1.0 - (days_ago / window)))
|
||||
|
||||
|
||||
def apply_combined_scoring(
|
||||
scored_results: list[ScoredResult],
|
||||
@@ -62,9 +24,6 @@ def apply_combined_scoring(
|
||||
temporal_alpha: float = _TEMPORAL_ALPHA,
|
||||
proof_count_alpha: float = _PROOF_COUNT_ALPHA,
|
||||
is_passthrough_reranker: bool = False,
|
||||
recency_decay_function: str = _RECENCY_DECAY_FUNCTION,
|
||||
recency_decay_linear_window_days: float = _RECENCY_DECAY_LINEAR_WINDOW_DAYS,
|
||||
recency_decay_halflife_days: float = _RECENCY_DECAY_HALFLIFE_DAYS,
|
||||
) -> None:
|
||||
"""Apply combined scoring to a list of ScoredResults in-place.
|
||||
|
||||
@@ -98,12 +57,6 @@ def apply_combined_scoring(
|
||||
recency_alpha: Max relative recency adjustment (default 0.2 → ±10%).
|
||||
temporal_alpha: Max relative temporal adjustment (default 0.2 → ±10%).
|
||||
proof_count_alpha: Max relative proof count adjustment (default 0.1 → ±5%).
|
||||
recency_decay_function: Age→freshness curve — "linear" (default),
|
||||
"exponential", or "none". See compute_recency_decay.
|
||||
recency_decay_linear_window_days: Days over which the linear curve
|
||||
decays to its floor (default 365).
|
||||
recency_decay_halflife_days: For the exponential curve, the age at which
|
||||
the recency signal is neutral (0.5) (default 90).
|
||||
"""
|
||||
if now.tzinfo is None:
|
||||
now = now.replace(tzinfo=UTC)
|
||||
@@ -145,26 +98,14 @@ def apply_combined_scoring(
|
||||
sr.cross_encoder_score_normalized = 1.0 - (0.9 * new_rank / denom)
|
||||
|
||||
for sr in scored_results:
|
||||
# Recency: configurable decay (linear default; see compute_recency_decay)
|
||||
# → [0.0, 1.0]; neutral 0.5 if no date.
|
||||
# Use the unit's effective time (occurred_start, then mentioned_at, then
|
||||
# occurred_end) — the same COALESCE order as retrieval._coalesce_date — so a
|
||||
# memory that carries only a mentioned_at / occurred_end (e.g. conversation
|
||||
# facts or ongoing states that intentionally lack occurred_start) still gets
|
||||
# correct recency ordering instead of a flat neutral 0.5.
|
||||
# Recency: linear decay over 365 days → [0.1, 1.0]; neutral 0.5 if no date.
|
||||
sr.recency = 0.5
|
||||
effective = sr.retrieval.occurred_start or sr.retrieval.mentioned_at or sr.retrieval.occurred_end
|
||||
if effective:
|
||||
occurred = effective
|
||||
if sr.retrieval.occurred_start:
|
||||
occurred = sr.retrieval.occurred_start
|
||||
if occurred.tzinfo is None:
|
||||
occurred = occurred.replace(tzinfo=UTC)
|
||||
days_ago = (now - occurred).total_seconds() / 86400
|
||||
sr.recency = compute_recency_decay(
|
||||
days_ago,
|
||||
recency_decay_function,
|
||||
recency_decay_linear_window_days,
|
||||
recency_decay_halflife_days,
|
||||
)
|
||||
sr.recency = max(0.1, min(1.0, 1.0 - (days_ago / 365)))
|
||||
|
||||
# Temporal proximity: meaningful only for temporal queries; neutral otherwise.
|
||||
sr.temporal = sr.retrieval.temporal_proximity if sr.retrieval.temporal_proximity is not None else 0.5
|
||||
@@ -177,9 +118,6 @@ def apply_combined_scoring(
|
||||
else:
|
||||
# Neutral baseline is precisely 0.5, ensuring neutral multiplier (1.0)
|
||||
proof_norm = 0.5
|
||||
# Surface the proof signal so the trace can show the proof_count_boost
|
||||
# factor (otherwise the reranked breakdown can't reconcile CE × boosts).
|
||||
sr.proof_norm = proof_norm
|
||||
|
||||
# RRF: kept at 0.0 for trace continuity but excluded from scoring.
|
||||
# RRF is batch-relative (min-max normalised) and redundant after reranking.
|
||||
|
||||
@@ -104,8 +104,6 @@ async def retrieve_semantic_bm25_combined(
|
||||
tag_groups: list[TagGroup] | None = None,
|
||||
created_after: datetime | None = None,
|
||||
created_before: datetime | None = None,
|
||||
min_semantic: float | None = None,
|
||||
min_keyword: float | None = None,
|
||||
) -> dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]]:
|
||||
"""
|
||||
Combined semantic + BM25 retrieval for multiple fact types in a single query.
|
||||
@@ -145,12 +143,6 @@ async def retrieve_semantic_bm25_combined(
|
||||
config = get_config()
|
||||
tokens = tokenize_query(query_text)
|
||||
|
||||
# Per-request retrieval-level score floors (recall min_scores.semantic / .keyword)
|
||||
# override the global config defaults for this query, pruning weak matches in
|
||||
# the SQL arms before fusion.
|
||||
sem_min = min_semantic if min_semantic is not None else config.semantic_min_similarity
|
||||
bm25_min = min_keyword if min_keyword is not None else config.bm25_min_score
|
||||
|
||||
# Over-fetch for HNSW approximation; semantic results trimmed to limit in Python.
|
||||
hnsw_fetch = max(limit * 5, 100)
|
||||
|
||||
@@ -211,7 +203,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
embedding_param="$1",
|
||||
bank_id_param="$2",
|
||||
fetch_limit=hnsw_fetch,
|
||||
min_similarity=sem_min,
|
||||
min_similarity=config.semantic_min_similarity,
|
||||
tags_clause=tags_clause,
|
||||
groups_clause=groups_clause,
|
||||
extra_where=created_range_clause,
|
||||
@@ -237,7 +229,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
arm_index=i,
|
||||
text_search_extension=text_ext,
|
||||
bm25_language=config.text_search_extension_native_language,
|
||||
bm25_min_score=bm25_min,
|
||||
bm25_min_score=config.bm25_min_score,
|
||||
extra_where=created_range_clause,
|
||||
)
|
||||
)
|
||||
@@ -285,7 +277,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
embedding_param="$1",
|
||||
bank_id_param="$2",
|
||||
fetch_limit=hnsw_fetch,
|
||||
min_similarity=sem_min,
|
||||
min_similarity=config.semantic_min_similarity,
|
||||
tags_clause=fb_tags_clause,
|
||||
groups_clause=fb_groups_clause,
|
||||
extra_where=fb_created_clause,
|
||||
@@ -714,8 +706,6 @@ async def retrieve_all_fact_types_parallel(
|
||||
tag_groups: list[TagGroup] | None = None,
|
||||
created_after: datetime | None = None,
|
||||
created_before: datetime | None = None,
|
||||
min_semantic: float | None = None,
|
||||
min_keyword: float | None = None,
|
||||
) -> MultiFactTypeRetrievalResult:
|
||||
"""
|
||||
Optimized retrieval for multiple fact types using batched queries.
|
||||
@@ -776,8 +766,6 @@ async def retrieve_all_fact_types_parallel(
|
||||
tag_groups=tag_groups,
|
||||
created_after=created_after,
|
||||
created_before=created_before,
|
||||
min_semantic=min_semantic,
|
||||
min_keyword=min_keyword,
|
||||
)
|
||||
semantic_bm25_time = time.time() - semantic_bm25_start
|
||||
|
||||
@@ -793,7 +781,7 @@ async def retrieve_all_fact_types_parallel(
|
||||
tc_start,
|
||||
tc_end,
|
||||
budget=thinking_budget,
|
||||
semantic_threshold=min_semantic if min_semantic is not None else 0.1,
|
||||
semantic_threshold=0.1,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
tag_groups=tag_groups,
|
||||
|
||||
@@ -14,12 +14,6 @@ AND matching (all/all_strict): Memory matches if ALL request tags are present in
|
||||
EXACT matching: Memory matches only if its tag set EQUALS the request tag set (order-
|
||||
independent). Used for observation "scope" filtering, where each observation lives
|
||||
under exactly one scope (its full tag set) and "scope [a]" must not match "[a, b]".
|
||||
An EMPTY request scope (no tags — ``[]`` or ``None``) is the global/untagged scope and
|
||||
matches only untagged memories — the scope that ``observation_scopes="shared"``
|
||||
consolidation writes to. This is the one mode where absent tags filter rather than
|
||||
meaning "no filter"; all other modes treat empty/absent tags as "no filtering". This
|
||||
mirrors the ``GET .../graph`` endpoint, where ``tags_match="exact"`` with no tags also
|
||||
selects the global scope.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -88,16 +82,11 @@ def build_tags_where_clause(
|
||||
>>> clause, params, next_offset = build_tags_where_clause(['user_a'], 3, 'mu.', 'any_strict')
|
||||
>>> print(clause) # "AND mu.tags IS NOT NULL AND mu.tags != '{}' AND mu.tags && $3"
|
||||
"""
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
|
||||
if match == "exact" and not tags:
|
||||
# Empty/absent scope = global/untagged: match only untagged rows. No bind param
|
||||
# needed (callers gate the param on truthy `tags`, so none is appended).
|
||||
return f"AND ({column} IS NULL OR {column} = '{{}}')", [], param_offset
|
||||
|
||||
if not tags:
|
||||
return "", [], param_offset
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
|
||||
if match == "exact":
|
||||
# Set equality (order-independent): superset AND subset. Untagged rows
|
||||
# (empty array) never satisfy `@>` of a non-empty scope, so they're excluded.
|
||||
@@ -137,16 +126,11 @@ def build_tags_where_clause_simple(
|
||||
Returns:
|
||||
SQL clause string or empty string.
|
||||
"""
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
|
||||
if match == "exact" and not tags:
|
||||
# Empty/absent scope = global/untagged: match only untagged rows. No bind param
|
||||
# needed (callers gate the param on truthy `tags`, so none is appended).
|
||||
return f"AND ({column} IS NULL OR {column} = '{{}}')"
|
||||
|
||||
if not tags:
|
||||
return ""
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
|
||||
if match == "exact":
|
||||
# Set equality (order-independent): superset AND subset. Untagged rows
|
||||
# (empty array) never satisfy `@>` of a non-empty scope, so they're excluded.
|
||||
@@ -180,10 +164,6 @@ def filter_results_by_tags(
|
||||
Returns:
|
||||
Filtered list of results.
|
||||
"""
|
||||
if match == "exact" and not tags:
|
||||
# Empty/absent scope = global/untagged: keep only untagged results.
|
||||
return [r for r in results if not getattr(r, "tags", None)]
|
||||
|
||||
if not tags:
|
||||
return results
|
||||
|
||||
@@ -287,9 +267,6 @@ def _build_group_clause(
|
||||
if isinstance(group, TagGroupLeaf):
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
if group.match == "exact":
|
||||
if len(group.tags) == 0:
|
||||
# Empty scope = global/untagged: match only untagged rows (no bind param).
|
||||
return f"({column} IS NULL OR {column} = '{{}}')", [], param_offset
|
||||
clause = f"({column} @> ${param_offset} AND {column} <@ ${param_offset})"
|
||||
return clause, [group.tags], param_offset + 1
|
||||
operator, include_untagged = _parse_tags_match(group.match)
|
||||
@@ -392,9 +369,6 @@ def _match_group(result: object, group: TagGroup) -> bool:
|
||||
if isinstance(group, TagGroupLeaf):
|
||||
result_tags = getattr(result, "tags", None)
|
||||
is_untagged = result_tags is None or len(result_tags) == 0
|
||||
if group.match == "exact" and len(group.tags) == 0:
|
||||
# Empty scope = global/untagged: match only untagged results.
|
||||
return is_untagged
|
||||
_, include_untagged = _parse_tags_match(group.match)
|
||||
is_any_match = group.match in ("any", "any_strict")
|
||||
tags_set = set(group.tags)
|
||||
|
||||
@@ -5,7 +5,6 @@ Think operation utilities for formulating answers based on agent and world facts
|
||||
import logging
|
||||
from datetime import datetime
|
||||
|
||||
from ...config import get_config
|
||||
from ..response_models import DispositionTraits, MemoryFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -252,7 +251,7 @@ async def reflect(
|
||||
answer_text = await llm_config.call(
|
||||
messages=[{"role": "system", "content": system_message}, {"role": "user", "content": prompt}],
|
||||
scope="memory_think",
|
||||
temperature=get_config().llm_temperature_reflect,
|
||||
temperature=0.9,
|
||||
max_completion_tokens=1000,
|
||||
)
|
||||
|
||||
|
||||
@@ -392,7 +392,7 @@ class SearchTracer:
|
||||
|
||||
# Extract score components (only include non-None values)
|
||||
# Keys from ScoredResult.to_dict(): cross_encoder_score, cross_encoder_score_normalized,
|
||||
# rrf_normalized, temporal, recency, proof_norm, combined_score, weight
|
||||
# rrf_normalized, temporal, recency, combined_score, weight
|
||||
score_components = {}
|
||||
for key in [
|
||||
"cross_encoder_score",
|
||||
@@ -401,7 +401,6 @@ class SearchTracer:
|
||||
"rrf_normalized",
|
||||
"temporal",
|
||||
"recency",
|
||||
"proof_norm",
|
||||
"combined_score",
|
||||
]:
|
||||
if key in result and result[key] is not None:
|
||||
|
||||
@@ -82,20 +82,6 @@ class RetrievalResult:
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ArmScores:
|
||||
"""Raw per-strategy retrieval scores for a single doc, aggregated across arms.
|
||||
|
||||
Fusion keeps only the first-seen RetrievalResult per doc, so its per-arm score
|
||||
fields reflect just one arm. This captures each arm's raw score for the same doc
|
||||
so the recall response can report them (and ``min_scores`` can filter on them).
|
||||
``None`` means the doc was not surfaced by that arm.
|
||||
"""
|
||||
|
||||
semantic: float | None = None # cosine similarity from the semantic arm
|
||||
keyword: float | None = None # BM25 / full-text score from the keyword arm
|
||||
|
||||
|
||||
@dataclass
|
||||
class MergedCandidate:
|
||||
"""
|
||||
@@ -111,7 +97,6 @@ class MergedCandidate:
|
||||
rrf_score: float
|
||||
rrf_rank: int = 0
|
||||
source_ranks: dict[str, int] = field(default_factory=dict) # method_name -> rank
|
||||
arm_scores: "ArmScores" = field(default_factory=lambda: ArmScores()) # raw per-strategy scores
|
||||
|
||||
@property
|
||||
def id(self) -> str:
|
||||
@@ -138,7 +123,6 @@ class ScoredResult:
|
||||
rrf_normalized: float = 0.0
|
||||
recency: float = 0.5
|
||||
temporal: float = 0.5
|
||||
proof_norm: float = 0.5 # log-normalized proof count (neutral 0.5); drives proof_count_boost
|
||||
|
||||
# Final combined score
|
||||
combined_score: float = 0.0
|
||||
@@ -195,7 +179,6 @@ class ScoredResult:
|
||||
result["rrf_normalized"] = self.rrf_normalized
|
||||
result["temporal"] = self.temporal
|
||||
result["recency"] = self.recency
|
||||
result["proof_norm"] = self.proof_norm
|
||||
result["combined_score"] = self.combined_score
|
||||
result["weight"] = self.weight
|
||||
result["activation"] = self.weight # Legacy field
|
||||
|
||||
@@ -1,155 +0,0 @@
|
||||
"""Explicit period extraction helpers for DateparserQueryAnalyzer.
|
||||
|
||||
This module keeps the public period-extraction API and the non-Chinese period
|
||||
rules. Chinese rules live in chinese_temporal_periods.py because that rule set is
|
||||
substantially larger and has different boundary behavior from whitespace-based
|
||||
languages.
|
||||
"""
|
||||
|
||||
import calendar
|
||||
import re
|
||||
import unicodedata
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
DateRange = tuple[datetime, datetime]
|
||||
|
||||
|
||||
class NoTemporalConstraintSentinel:
|
||||
pass
|
||||
|
||||
|
||||
NO_TEMPORAL_CONSTRAINT = NoTemporalConstraintSentinel()
|
||||
|
||||
__all__ = [
|
||||
"NO_TEMPORAL_CONSTRAINT",
|
||||
"extract_period",
|
||||
"is_embedded_cjk_dateparser_match",
|
||||
]
|
||||
|
||||
|
||||
def _is_cjk_character(char: str) -> bool:
|
||||
return "\u4e00" <= char <= "\u9fff"
|
||||
|
||||
|
||||
def is_embedded_cjk_dateparser_match(query: str, matched_text: str) -> bool:
|
||||
from hindsight_api.engine.chinese_temporal_periods import (
|
||||
is_embedded_cjk_dateparser_match as chinese_is_embedded_cjk_dateparser_match,
|
||||
)
|
||||
|
||||
return chinese_is_embedded_cjk_dateparser_match(query, matched_text)
|
||||
|
||||
|
||||
def _constraint(start: datetime, end: datetime) -> DateRange:
|
||||
return (
|
||||
start.replace(hour=0, minute=0, second=0, microsecond=0),
|
||||
end.replace(hour=23, minute=59, second=59, microsecond=999999),
|
||||
)
|
||||
|
||||
|
||||
def _month_end(year: int, month: int) -> datetime:
|
||||
return datetime(year, month, calendar.monthrange(year, month)[1])
|
||||
|
||||
|
||||
def _extract_non_chinese_period(query: str, reference_date: datetime) -> DateRange | None:
|
||||
if re.search(r"\b(yesterday|ayer|ieri|hier|gestern)\b", query, re.IGNORECASE):
|
||||
d = reference_date - timedelta(days=1)
|
||||
return _constraint(d, d)
|
||||
|
||||
if re.search(r"\b(today|hoy|oggi|aujourd\'?hui|heute)\b", query, re.IGNORECASE):
|
||||
return _constraint(reference_date, reference_date)
|
||||
|
||||
if re.search(r"\b(a\s+)?couple\s+(of\s+)?days?\s+ago\b", query, re.IGNORECASE):
|
||||
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):
|
||||
return _constraint(reference_date - timedelta(days=5), reference_date - timedelta(days=2))
|
||||
|
||||
if re.search(r"\b(a\s+)?couple\s+(of\s+)?weeks?\s+ago\b", query, re.IGNORECASE):
|
||||
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):
|
||||
return _constraint(reference_date - timedelta(weeks=5), reference_date - timedelta(weeks=2))
|
||||
|
||||
if re.search(r"\b(a\s+)?couple\s+(of\s+)?months?\s+ago\b", query, re.IGNORECASE):
|
||||
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):
|
||||
return _constraint(reference_date - timedelta(days=150), reference_date - timedelta(days=60))
|
||||
|
||||
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))
|
||||
|
||||
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)
|
||||
|
||||
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))
|
||||
|
||||
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
|
||||
sat = reference_date - timedelta(days=days_since_sat)
|
||||
return _constraint(sat, sat + timedelta(days=1))
|
||||
|
||||
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,
|
||||
}
|
||||
for pattern, month_num in month_patterns.items():
|
||||
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)
|
||||
return _constraint(start, _month_end(year, month_num))
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def extract_period(query: str, reference_date: datetime) -> DateRange | NoTemporalConstraintSentinel | None:
|
||||
"""Extract explicit period-based temporal expressions.
|
||||
|
||||
Non-Chinese rules are kept here. Chinese rules are delegated to
|
||||
chinese_temporal_periods.py and are skipped entirely for non-CJK queries.
|
||||
"""
|
||||
query = unicodedata.normalize("NFKC", query)
|
||||
|
||||
if any(_is_cjk_character(char) for char in query):
|
||||
from hindsight_api.engine.chinese_temporal_periods import extract_chinese_period
|
||||
|
||||
chinese_result = extract_chinese_period(query, reference_date)
|
||||
if chinese_result is not None:
|
||||
return chinese_result
|
||||
|
||||
return _extract_non_chinese_period(query, reference_date)
|
||||
@@ -22,7 +22,7 @@ from pydantic import BaseModel, Field
|
||||
# Bump when the archive layout changes in a backward-incompatible way.
|
||||
SCHEMA_VERSION = 1
|
||||
|
||||
ObservationScopes = Literal["per_tag", "combined", "all_combinations", "shared"] | list[list[str]]
|
||||
ObservationScopes = Literal["per_tag", "combined", "all_combinations"] | list[list[str]]
|
||||
|
||||
|
||||
class TransferCausalRelation(BaseModel):
|
||||
|
||||
@@ -39,9 +39,7 @@ from hindsight_api.extensions.operation_validator import (
|
||||
BankListContext,
|
||||
BankListResult,
|
||||
BankReadContext,
|
||||
BankReadOperation,
|
||||
BankWriteContext,
|
||||
BankWriteOperation,
|
||||
# Consolidation operation
|
||||
ConsolidateContext,
|
||||
ConsolidateResult,
|
||||
@@ -56,7 +54,6 @@ from hindsight_api.extensions.operation_validator import (
|
||||
OperationValidationError,
|
||||
OperationValidatorExtension,
|
||||
PrecheckContext,
|
||||
PrecheckOperation,
|
||||
RecallContext,
|
||||
RecallResult,
|
||||
ReflectContext,
|
||||
@@ -90,7 +87,6 @@ __all__ = [
|
||||
"OperationValidationError",
|
||||
"OperationValidatorExtension",
|
||||
"PrecheckContext",
|
||||
"PrecheckOperation",
|
||||
"RecallContext",
|
||||
"RecallResult",
|
||||
"ReflectContext",
|
||||
@@ -102,9 +98,7 @@ __all__ = [
|
||||
"BankListContext",
|
||||
"BankListResult",
|
||||
"BankReadContext",
|
||||
"BankReadOperation",
|
||||
"BankWriteContext",
|
||||
"BankWriteOperation",
|
||||
# Operation Validator - Consolidation
|
||||
"ConsolidateContext",
|
||||
"ConsolidateResult",
|
||||
|
||||
@@ -52,5 +52,4 @@ class MemoryDefenseRegexExtension(MemoryDefenseExtension):
|
||||
message=f"Sensitive data pattern matched: {', '.join(result.matched_types)}",
|
||||
redacted_content=result.content if rule.action is DefenseAction.REDACT else None,
|
||||
matched_types=result.matched_types,
|
||||
hits=result.hits,
|
||||
)
|
||||
|
||||
@@ -29,14 +29,9 @@ class DefenseAction(str, Enum):
|
||||
|
||||
_VALID_ACTIONS = {a.value for a in DefenseAction}
|
||||
|
||||
# ``policy.rules[*].on`` names a detector. The OSS extension only screens for
|
||||
# ``sensitive_data``; any other name is a silent no-op here and is dispatched
|
||||
# by whichever extension is loaded (e.g. hindsight-cloud screens cloud-only
|
||||
# detectors). The parser therefore does NOT validate ``on`` against a fixed
|
||||
# list — pinning the OSS roster to cloud's would force an OSS bump for every
|
||||
# new cloud detector just to avoid 422-ing a write it never interprets. We
|
||||
# only require ``on`` to be a non-empty string; entitlement and dispatch are
|
||||
# the loaded extension's ``screen()`` job.
|
||||
# Detector identifiers valid as ``policy.rules[*].on``. The OSS extension only
|
||||
# screens for sensitive data (secrets/PII), so that's the only accepted value.
|
||||
_VALID_DETECTORS = {"sensitive_data"}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -58,58 +53,19 @@ class DefenseDecision:
|
||||
message: str = ""
|
||||
redacted_content: str | None = None
|
||||
matched_types: list[str] = field(default_factory=list)
|
||||
# Per-match fingerprinted previews. Each entry is
|
||||
# ``{"detector": <pattern label>, "preview": <fingerprinted value>}``.
|
||||
# The preview is *never* the raw value — see :func:`_fingerprint_value`.
|
||||
# OSS populates this from ``apply_redaction``; downstream extensions
|
||||
# populate it from their own detectors. Optional: empty when the
|
||||
# match path didn't capture per-hit values.
|
||||
hits: list[dict] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RedactionResult:
|
||||
content: str
|
||||
matched_types: list[str]
|
||||
# Same shape as ``DefenseDecision.hits`` — one entry per matched value
|
||||
# (so a single content with two GitHub tokens produces two entries).
|
||||
hits: list[dict] = field(default_factory=list)
|
||||
|
||||
|
||||
def _fingerprint_value(value: str) -> str:
|
||||
"""Return a redaction-identifiable preview of a matched value.
|
||||
|
||||
The preview keeps the prefix and a short suffix so a SIEM operator can
|
||||
correlate against their credential inventory (the prefix names the
|
||||
provider; the suffix disambiguates specific instances) without the raw
|
||||
secret crossing the wire. Length-aware so short values don't accidentally
|
||||
leak material:
|
||||
|
||||
- Length < 6: redact entirely (return a fixed-length mask). Catches
|
||||
noise like a single ``-----BEGIN...`` marker line.
|
||||
- Length 6-15: keep the first 2 + last 2 around an ellipsis.
|
||||
- Length > 15: keep the first 4 + last 4 around an ellipsis.
|
||||
|
||||
Examples::
|
||||
|
||||
_fingerprint_value("ghp_AAAA...AAAA" + "A" * 36) -> "ghp_...AAAA"
|
||||
_fingerprint_value("AKIA" + "B" * 16) -> "AKIA...BBBB"
|
||||
_fingerprint_value("123-45-6789") -> "12...89"
|
||||
_fingerprint_value("abc") -> "[redacted]"
|
||||
"""
|
||||
n = len(value)
|
||||
if n < 6:
|
||||
return "[redacted]"
|
||||
if n <= 15:
|
||||
return f"{value[:2]}...{value[-2:]}"
|
||||
return f"{value[:4]}...{value[-4:]}"
|
||||
|
||||
|
||||
def parse_policy(raw: dict | None) -> DefensePolicy:
|
||||
"""Parse a raw bank-config dict into a frozen DefensePolicy.
|
||||
|
||||
Raises ValueError for a missing/empty ``on`` or an unknown action; the
|
||||
HTTP layer converts those into a 422 response.
|
||||
Raises ValueError for unknown detectors or actions; the HTTP layer
|
||||
converts those into a 422 response.
|
||||
"""
|
||||
if raw is None:
|
||||
return DefensePolicy()
|
||||
@@ -117,8 +73,8 @@ def parse_policy(raw: dict | None) -> DefensePolicy:
|
||||
rules: list[PolicyRule] = []
|
||||
for item in raw.get("rules", []) or []:
|
||||
on_raw = item.get("on")
|
||||
if not isinstance(on_raw, str) or not on_raw:
|
||||
raise ValueError(f"invalid on {on_raw!r}; must be a non-empty string")
|
||||
if on_raw not in _VALID_DETECTORS:
|
||||
raise ValueError(f"invalid on {on_raw!r}; must be one of {sorted(_VALID_DETECTORS)}")
|
||||
action_raw = item.get("action")
|
||||
if action_raw not in _VALID_ACTIONS:
|
||||
raise ValueError(f"invalid action {action_raw!r}; must be one of {sorted(_VALID_ACTIONS)}")
|
||||
@@ -210,42 +166,16 @@ _COMPILED_REDACTIONS: list[tuple[str, re.Pattern]] = [
|
||||
def apply_redaction(content: str) -> RedactionResult:
|
||||
"""Scrub known secret/PII patterns from content with [REDACTED:type] markers.
|
||||
|
||||
Returns the (possibly unchanged) content alongside:
|
||||
- ``matched_types``: pattern labels that matched (deduplicated, in
|
||||
first-occurrence order). Empty when nothing matched.
|
||||
- ``hits``: per-match fingerprinted previews — one entry per matched
|
||||
substring (so two GitHub tokens in the same content produce two
|
||||
entries). Each entry is ``{"detector": label, "preview": fingerprint}``
|
||||
where ``preview`` is a length-aware redaction of the original value.
|
||||
The raw secret never appears in ``hits``.
|
||||
|
||||
The two-pass shape (find matches first, then substitute) lets us capture
|
||||
raw values for fingerprinting before they're replaced by ``[REDACTED:type]``
|
||||
markers. A single-pass approach would lose the originals.
|
||||
Returns the (possibly unchanged) content alongside the list of pattern
|
||||
labels that matched (empty when nothing matched).
|
||||
"""
|
||||
matched: list[str] = []
|
||||
hits: list[dict] = []
|
||||
for label, pattern in _COMPILED_REDACTIONS:
|
||||
raw_hits = pattern.findall(content)
|
||||
if not raw_hits:
|
||||
continue
|
||||
if label not in matched:
|
||||
new_content = pattern.sub(f"[REDACTED:{label}]", content)
|
||||
if new_content != content:
|
||||
matched.append(label)
|
||||
for raw in raw_hits:
|
||||
# findall returns either a string or a tuple of capture groups
|
||||
# depending on the pattern. The redaction-pattern catalog uses a
|
||||
# mix; coerce to the matched substring as best we can.
|
||||
if isinstance(raw, tuple):
|
||||
# Pick the longest non-empty group as the canonical match.
|
||||
non_empty = [g for g in raw if g]
|
||||
raw_str = max(non_empty, key=len) if non_empty else ""
|
||||
else:
|
||||
raw_str = raw
|
||||
if not raw_str:
|
||||
continue
|
||||
hits.append({"detector": label, "preview": _fingerprint_value(raw_str)})
|
||||
content = pattern.sub(f"[REDACTED:{label}]", content)
|
||||
return RedactionResult(content=content, matched_types=matched, hits=hits)
|
||||
content = new_content
|
||||
return RedactionResult(content=content, matched_types=matched)
|
||||
|
||||
|
||||
class MemoryDefenseExtension(Extension, ABC):
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from enum import StrEnum
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from hindsight_api.extensions.base import Extension
|
||||
@@ -83,18 +82,6 @@ class ValidationResult:
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class PrecheckOperation(StrEnum):
|
||||
"""Route operation names passed to the pre-body-parse precheck hook."""
|
||||
|
||||
DRY_RUN_EXTRACT = "dry_run_extract"
|
||||
FILES_RETAIN = "files_retain"
|
||||
MENTAL_MODEL_CREATE = "mental_model_create"
|
||||
MENTAL_MODEL_REFRESH = "mental_model_refresh"
|
||||
RECALL = "recall"
|
||||
REFLECT = "reflect"
|
||||
RETAIN = "retain"
|
||||
|
||||
|
||||
@dataclass
|
||||
class PrecheckContext:
|
||||
"""Context for a pre-body-parse precheck on an operation.
|
||||
@@ -104,14 +91,12 @@ class PrecheckContext:
|
||||
therefore intentionally carries only the cheap, already-resolved
|
||||
pieces of request state:
|
||||
|
||||
- ``operation``: a short string-compatible enum identifying the route.
|
||||
- ``operation``: a short string identifying the route, e.g. ``"retain"``,
|
||||
``"recall"``, ``"reflect"``, ``"files_retain"``, ``"mental_model_create"``,
|
||||
``"mental_model_refresh"``.
|
||||
- ``bank_id``: parsed from the URL path.
|
||||
- ``request_context``: the authenticated :class:`RequestContext` (tenant
|
||||
already resolved by the tenant extension).
|
||||
- ``content_length``: value of the ``Content-Length`` request header as an
|
||||
int, or ``None`` when the header is absent or unparseable (e.g. chunked
|
||||
transfer encoding). Lets a precheck make size-aware decisions — such as
|
||||
an upper-bound cost estimate — without reading or deserialising the body.
|
||||
|
||||
Implementations should keep precheck cheap and side-effect-free. The
|
||||
full per-request validators (``validate_retain`` / ``validate_recall``
|
||||
@@ -119,10 +104,9 @@ class PrecheckContext:
|
||||
the source of truth for the precise per-call cost / quota arithmetic.
|
||||
"""
|
||||
|
||||
operation: PrecheckOperation
|
||||
operation: str
|
||||
bank_id: str
|
||||
request_context: "RequestContext"
|
||||
content_length: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -219,16 +203,6 @@ class RetainResult:
|
||||
llm_input_tokens: int | None = None
|
||||
llm_output_tokens: int | None = None
|
||||
llm_total_tokens: int | None = None
|
||||
# Diagnostic token splits surfaced for cost attribution and prompt-cache
|
||||
# tuning. ``llm_cached_input_tokens`` is the subset of llm_input_tokens
|
||||
# served from the provider's prompt cache (e.g. Gemini's
|
||||
# cached_content_token_count). ``llm_thoughts_tokens`` is reasoning tokens
|
||||
# that are billed at the output rate by some providers (Gemini 2.5+) but
|
||||
# are not part of the visible response. Both default to None when the
|
||||
# engine/provider didn't report them; downstream metering extensions
|
||||
# should treat None as 0.
|
||||
llm_cached_input_tokens: int | None = None
|
||||
llm_thoughts_tokens: int | None = None
|
||||
# Content tokens the retain pipeline actually processed, after
|
||||
# chunk-level content-hash deduplication. Semantics:
|
||||
# None — no dedup signal available (e.g. a first-time retain or a
|
||||
@@ -314,77 +288,12 @@ class ConsolidateResult:
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class BankReadOperation(StrEnum):
|
||||
"""Bank-scoped read operation names passed to validate_bank_read."""
|
||||
|
||||
GET_BANK_CONFIG = "get_bank_config"
|
||||
GET_BANK_PROFILE = "get_bank_profile"
|
||||
GET_BANK_STATS = "get_bank_stats"
|
||||
GET_CHUNK = "get_chunk"
|
||||
GET_DIRECTIVE = "get_directive"
|
||||
GET_DOCUMENT = "get_document"
|
||||
GET_ENTITY = "get_entity"
|
||||
GET_ENTITY_GRAPH = "get_entity_graph"
|
||||
GET_ENTITY_STATE = "get_entity_state"
|
||||
GET_GRAPH_DATA = "get_graph_data"
|
||||
GET_MEMORIES_TIMESERIES = "get_memories_timeseries"
|
||||
GET_MEMORY_UNIT = "get_memory_unit"
|
||||
GET_OBSERVATION_HISTORY = "get_observation_history"
|
||||
GET_OPERATION_STATUS = "get_operation_status"
|
||||
LIST_DIRECTIVES = "list_directives"
|
||||
LIST_DOCUMENT_CHUNKS = "list_document_chunks"
|
||||
LIST_DOCUMENTS = "list_documents"
|
||||
LIST_ENTITIES = "list_entities"
|
||||
LIST_MEMORY_UNITS = "list_memory_units"
|
||||
LIST_MENTAL_MODEL_TAGS = "list_mental_model_tags"
|
||||
LIST_MENTAL_MODELS = "list_mental_models"
|
||||
LIST_OBSERVATION_SCOPES = "list_observation_scopes"
|
||||
LIST_OPERATIONS = "list_operations"
|
||||
LIST_TAGS = "list_tags"
|
||||
LIST_WEBHOOK_DELIVERIES = "list_webhook_deliveries"
|
||||
LIST_WEBHOOKS = "list_webhooks"
|
||||
|
||||
|
||||
class BankWriteOperation(StrEnum):
|
||||
"""Bank-scoped write operation names passed to validate_bank_write."""
|
||||
|
||||
CANCEL_OPERATION = "cancel_operation"
|
||||
CLEAR_MENTAL_MODEL = "clear_mental_model"
|
||||
CLEAR_OBSERVATIONS = "clear_observations"
|
||||
CLEAR_OBSERVATIONS_FOR_MEMORY = "clear_observations_for_memory"
|
||||
CREATE_DIRECTIVE = "create_directive"
|
||||
CREATE_MENTAL_MODEL = "create_mental_model"
|
||||
CREATE_WEBHOOK = "create_webhook"
|
||||
DELETE_BANK = "delete_bank"
|
||||
DELETE_DIRECTIVE = "delete_directive"
|
||||
DELETE_DOCUMENT = "delete_document"
|
||||
DELETE_MENTAL_MODEL = "delete_mental_model"
|
||||
DELETE_WEBHOOK = "delete_webhook"
|
||||
MERGE_BANK_MISSION = "merge_bank_mission"
|
||||
REPROCESS_DOCUMENT = "reprocess_document"
|
||||
RESET_BANK_CONFIG = "reset_bank_config"
|
||||
RETRY_FAILED_CONSOLIDATION = "retry_failed_consolidation"
|
||||
RETRY_OPERATION = "retry_operation"
|
||||
RUN_CONSOLIDATION = "run_consolidation"
|
||||
SET_BANK_MISSION = "set_bank_mission"
|
||||
SUBMIT_ASYNC_CONSOLIDATION = "submit_async_consolidation"
|
||||
SUBMIT_ASYNC_GRAPH_MAINTENANCE = "submit_async_graph_maintenance"
|
||||
UPDATE_BANK = "update_bank"
|
||||
UPDATE_BANK_CONFIG = "update_bank_config"
|
||||
UPDATE_BANK_DISPOSITION = "update_bank_disposition"
|
||||
UPDATE_DIRECTIVE = "update_directive"
|
||||
UPDATE_DOCUMENT = "update_document"
|
||||
UPDATE_MEMORY_UNIT = "update_memory_unit"
|
||||
UPDATE_MENTAL_MODEL = "update_mental_model"
|
||||
UPDATE_WEBHOOK = "update_webhook"
|
||||
|
||||
|
||||
@dataclass
|
||||
class BankReadContext:
|
||||
"""Context for a bank read operation validation (pre-operation)."""
|
||||
|
||||
bank_id: str
|
||||
operation: BankReadOperation
|
||||
operation: str # "get_bank_profile", "get_bank_stats"
|
||||
request_context: "RequestContext"
|
||||
|
||||
|
||||
@@ -393,7 +302,7 @@ class BankWriteContext:
|
||||
"""Context for a bank write operation validation (pre-operation)."""
|
||||
|
||||
bank_id: str
|
||||
operation: BankWriteOperation
|
||||
operation: str # "delete_bank", "update_bank", "update_bank_disposition", "set_bank_mission", "merge_bank_mission", "clear_observations", "clear_observations_for_memory"
|
||||
request_context: "RequestContext"
|
||||
|
||||
|
||||
|
||||
@@ -12,7 +12,6 @@ from datetime import datetime, timezone
|
||||
from typing import Any, Callable
|
||||
|
||||
from fastmcp import FastMCP
|
||||
from mcp.types import ToolAnnotations
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
@@ -22,7 +21,7 @@ from hindsight_api.config import (
|
||||
)
|
||||
from hindsight_api.engine.audit import AuditEntry, AuditLogger
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES, MinScores
|
||||
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
|
||||
from hindsight_api.engine.search.tags import TagGroup
|
||||
from hindsight_api.extensions import OperationValidationError
|
||||
from hindsight_api.models import RequestContext
|
||||
@@ -200,47 +199,6 @@ def build_content_dict(
|
||||
return content_dict, None
|
||||
|
||||
|
||||
# MCP tool annotations. Hindsight is a closed memory store (no open-world / internet
|
||||
# access), so openWorldHint=False throughout. readOnlyHint lets clients group and
|
||||
# auto-approve safe reads; destructiveHint flags tools that delete or clear memory.
|
||||
_READ_ONLY_TOOLS = {
|
||||
"recall",
|
||||
"reflect",
|
||||
"list_banks",
|
||||
"get_bank",
|
||||
"get_bank_stats",
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"list_directives",
|
||||
"list_memories",
|
||||
"get_memory",
|
||||
"list_documents",
|
||||
"get_document",
|
||||
"list_operations",
|
||||
"get_operation",
|
||||
"list_tags",
|
||||
}
|
||||
_DESTRUCTIVE_TOOLS = {
|
||||
"delete_bank",
|
||||
"clear_memories",
|
||||
"clear_mental_model",
|
||||
"delete_mental_model",
|
||||
"delete_directive",
|
||||
"delete_document",
|
||||
"invalidate_memory",
|
||||
}
|
||||
|
||||
|
||||
def _tool_annotations(name: str) -> ToolAnnotations:
|
||||
if name in _READ_ONLY_TOOLS:
|
||||
return ToolAnnotations(readOnlyHint=True, openWorldHint=False)
|
||||
if name in _DESTRUCTIVE_TOOLS:
|
||||
return ToolAnnotations(readOnlyHint=False, destructiveHint=True, openWorldHint=False)
|
||||
# Everything else writes but does not destructively delete/clear memory
|
||||
# (retain, create_*, update_*, refresh_mental_model, cancel_operation).
|
||||
return ToolAnnotations(readOnlyHint=False, destructiveHint=False, openWorldHint=False)
|
||||
|
||||
|
||||
def register_mcp_tools(
|
||||
mcp: FastMCP,
|
||||
memory: MemoryEngine,
|
||||
@@ -594,7 +552,7 @@ def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(description=description, annotations=_tool_annotations("retain"))
|
||||
@mcp.tool(description=description)
|
||||
async def retain(
|
||||
content: str,
|
||||
context: str = "general",
|
||||
@@ -650,7 +608,7 @@ def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(description=description, annotations=_tool_annotations("retain"))
|
||||
@mcp.tool(description=description)
|
||||
async def retain(
|
||||
content: str,
|
||||
context: str = "general",
|
||||
@@ -708,7 +666,7 @@ def _register_sync_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("sync_retain"))
|
||||
@mcp.tool()
|
||||
async def sync_retain(
|
||||
content: str,
|
||||
context: str = "general",
|
||||
@@ -766,7 +724,7 @@ def _register_sync_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("sync_retain"))
|
||||
@mcp.tool()
|
||||
async def sync_retain(
|
||||
content: str,
|
||||
context: str = "general",
|
||||
@@ -827,18 +785,16 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(description=description, annotations=_tool_annotations("recall"))
|
||||
@mcp.tool(description=description)
|
||||
async def recall(
|
||||
query: str,
|
||||
max_tokens: int = 4096,
|
||||
budget: str = "high",
|
||||
types: list[str] | None = None,
|
||||
prefer_observations: bool = False,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str = "any",
|
||||
tag_groups: list[dict] | None = None,
|
||||
query_timestamp: str | None = None,
|
||||
min_scores: dict | None = None,
|
||||
bank_id: str | None = None,
|
||||
) -> str | dict:
|
||||
"""
|
||||
@@ -847,10 +803,6 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
max_tokens: Maximum tokens to return in results (default: 4096)
|
||||
budget: Search budget - 'low', 'mid', or 'high' (default: 'high'). Higher budgets search more thoroughly.
|
||||
types: Fact types to include (e.g., ['world', 'experience']). Default: all types.
|
||||
prefer_observations: When recalling raw facts together with 'observation', drop any raw fact
|
||||
that a returned observation was consolidated from, so the observation supersedes it (no
|
||||
duplicate content). Disabled by default; set true to enable. No effect unless
|
||||
'observation' and a raw type are both in types. Default: False.
|
||||
tags: Optional tags to filter results by (e.g., ['project:alpha']). Mutually exclusive with tag_groups.
|
||||
tags_match: How to match tags - 'any' (match any tag) or 'all' (match all tags). Default: 'any'
|
||||
tag_groups: Compound tag filter using boolean groups (AND-ed together). Each group is a leaf
|
||||
@@ -859,11 +811,6 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
Mutually exclusive with tags.
|
||||
query_timestamp: Temporal context for the query (ISO format, e.g., '2024-01-15T10:30:00Z').
|
||||
Anchors relative temporal expressions and recency scoring.
|
||||
min_scores: Optional per-stage score floors as an object with any of: "semantic", "keyword"
|
||||
(retrieval-level cutoffs), "reranker", "final" (post-ranking). E.g. {"reranker": 0.5}.
|
||||
All inclusive and AND-ed; omit for no score filtering. The reranker's absolute scores are
|
||||
not calibrated across queries, so only threshold against scores you've calibrated for your
|
||||
own data.
|
||||
bank_id: Optional bank to search in (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
@@ -884,7 +831,6 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
"bank_id": target_bank,
|
||||
"query": query,
|
||||
"fact_type": fact_types,
|
||||
"prefer_observations": prefer_observations,
|
||||
"budget": budget_enum,
|
||||
"max_tokens": max_tokens,
|
||||
"request_context": _get_request_context(config),
|
||||
@@ -896,8 +842,6 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
recall_kwargs["tag_groups"] = _TAG_GROUP_LIST_ADAPTER.validate_python(tag_groups)
|
||||
if query_timestamp is not None:
|
||||
recall_kwargs["question_date"] = parse_timestamp(query_timestamp)
|
||||
if min_scores is not None:
|
||||
recall_kwargs["min_scores"] = MinScores.model_validate(min_scores)
|
||||
|
||||
recall_result = await memory.recall_async(**recall_kwargs)
|
||||
|
||||
@@ -913,18 +857,16 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(description=description, annotations=_tool_annotations("recall"))
|
||||
@mcp.tool(description=description)
|
||||
async def recall(
|
||||
query: str,
|
||||
max_tokens: int = 4096,
|
||||
budget: str = "high",
|
||||
types: list[str] | None = None,
|
||||
prefer_observations: bool = False,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str = "any",
|
||||
tag_groups: list[dict] | None = None,
|
||||
query_timestamp: str | None = None,
|
||||
min_scores: dict | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Args:
|
||||
@@ -932,10 +874,6 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
max_tokens: Maximum tokens to return in results (default: 4096)
|
||||
budget: Search budget - 'low', 'mid', or 'high' (default: 'high'). Higher budgets search more thoroughly.
|
||||
types: Fact types to include (e.g., ['world', 'experience']). Default: all types.
|
||||
prefer_observations: When recalling raw facts together with 'observation', drop any raw fact
|
||||
that a returned observation was consolidated from, so the observation supersedes it (no
|
||||
duplicate content). Disabled by default; set true to enable. No effect unless
|
||||
'observation' and a raw type are both in types. Default: False.
|
||||
tags: Optional tags to filter results by (e.g., ['project:alpha']). Mutually exclusive with tag_groups.
|
||||
tags_match: How to match tags - 'any' (match any tag) or 'all' (match all tags). Default: 'any'
|
||||
tag_groups: Compound tag filter using boolean groups (AND-ed together). Each group is a leaf
|
||||
@@ -944,11 +882,6 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
Mutually exclusive with tags.
|
||||
query_timestamp: Temporal context for the query (ISO format, e.g., '2024-01-15T10:30:00Z').
|
||||
Anchors relative temporal expressions and recency scoring.
|
||||
min_scores: Optional per-stage score floors as an object with any of: "semantic", "keyword"
|
||||
(retrieval-level cutoffs), "reranker", "final" (post-ranking). E.g. {"reranker": 0.5}.
|
||||
All inclusive and AND-ed; omit for no score filtering. The reranker's absolute scores are
|
||||
not calibrated across queries, so only threshold against scores you've calibrated for your
|
||||
own data.
|
||||
"""
|
||||
try:
|
||||
target_bank = config.bank_id_resolver()
|
||||
@@ -968,7 +901,6 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
"bank_id": target_bank,
|
||||
"query": query,
|
||||
"fact_type": fact_types,
|
||||
"prefer_observations": prefer_observations,
|
||||
"budget": budget_enum,
|
||||
"max_tokens": max_tokens,
|
||||
"request_context": _get_request_context(config),
|
||||
@@ -980,8 +912,6 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
|
||||
recall_kwargs["tag_groups"] = _TAG_GROUP_LIST_ADAPTER.validate_python(tag_groups)
|
||||
if query_timestamp is not None:
|
||||
recall_kwargs["question_date"] = parse_timestamp(query_timestamp)
|
||||
if min_scores is not None:
|
||||
recall_kwargs["min_scores"] = MinScores.model_validate(min_scores)
|
||||
|
||||
recall_result = await memory.recall_async(**recall_kwargs)
|
||||
|
||||
@@ -1001,7 +931,7 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("reflect"))
|
||||
@mcp.tool()
|
||||
async def reflect(
|
||||
query: str,
|
||||
context: str | None = None,
|
||||
@@ -1011,7 +941,6 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str = "any",
|
||||
include_based_on: bool = False,
|
||||
include_trace: bool = False,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
@@ -1042,7 +971,6 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
|
||||
tags: Optional tags to filter memories by (e.g., ['project:alpha'])
|
||||
tags_match: How to match tags - 'any' (match any tag) or 'all' (match all tags). Default: 'any'
|
||||
include_based_on: Include source facts used for synthesis. Defaults to false because broad reflections can exceed MCP client result limits.
|
||||
include_trace: Include the reflection's internal trace fields (tool_trace/llm_trace and directives_applied). Defaults to false because the trace can be tens of KB and overflow MCP client context; enable only for debugging.
|
||||
bank_id: Optional bank to reflect in (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
@@ -1072,15 +1000,6 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
|
||||
result_data = json.loads(reflect_result.model_dump_json(indent=2))
|
||||
if not include_based_on:
|
||||
result_data.pop("based_on", None)
|
||||
if not include_trace:
|
||||
# The agentic reflect loop's trace fields can be tens of KB (full
|
||||
# mental-model text) and silently overflow MCP client context; the
|
||||
# REST API omits them by default too. directives_applied is built by
|
||||
# the engine "for the trace" and carries full directive content, so it
|
||||
# belongs with tool_trace/llm_trace here. Opt in via include_trace.
|
||||
result_data.pop("tool_trace", None)
|
||||
result_data.pop("llm_trace", None)
|
||||
result_data.pop("directives_applied", None)
|
||||
if response_schema is not None and hasattr(reflect_result, "structured_output"):
|
||||
result_data["structured_output"] = reflect_result.structured_output
|
||||
return json.dumps(result_data, indent=2)
|
||||
@@ -1093,7 +1012,7 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("reflect"))
|
||||
@mcp.tool()
|
||||
async def reflect(
|
||||
query: str,
|
||||
context: str | None = None,
|
||||
@@ -1103,7 +1022,6 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str = "any",
|
||||
include_based_on: bool = False,
|
||||
include_trace: bool = False,
|
||||
) -> dict:
|
||||
"""
|
||||
Generate thoughtful analysis by synthesizing stored memories with the bank's personality.
|
||||
@@ -1133,7 +1051,6 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
|
||||
tags: Optional tags to filter memories by (e.g., ['project:alpha'])
|
||||
tags_match: How to match tags - 'any' (match any tag) or 'all' (match all tags). Default: 'any'
|
||||
include_based_on: Include source facts used for synthesis. Defaults to false because broad reflections can exceed MCP client result limits.
|
||||
include_trace: Include the reflection's internal trace fields (tool_trace/llm_trace and directives_applied). Defaults to false because the trace can be tens of KB and overflow MCP client context; enable only for debugging.
|
||||
"""
|
||||
try:
|
||||
target_bank = config.bank_id_resolver()
|
||||
@@ -1162,15 +1079,6 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
|
||||
result_data = reflect_result.model_dump()
|
||||
if not include_based_on:
|
||||
result_data.pop("based_on", None)
|
||||
if not include_trace:
|
||||
# The agentic reflect loop's trace fields can be tens of KB (full
|
||||
# mental-model text) and silently overflow MCP client context; the
|
||||
# REST API omits them by default too. directives_applied is built by
|
||||
# the engine "for the trace" and carries full directive content, so it
|
||||
# belongs with tool_trace/llm_trace here. Opt in via include_trace.
|
||||
result_data.pop("tool_trace", None)
|
||||
result_data.pop("llm_trace", None)
|
||||
result_data.pop("directives_applied", None)
|
||||
if response_schema is not None and hasattr(reflect_result, "structured_output"):
|
||||
result_data["structured_output"] = reflect_result.structured_output
|
||||
return result_data
|
||||
@@ -1185,7 +1093,7 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
|
||||
def _register_list_banks(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the list_banks tool."""
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_banks"))
|
||||
@mcp.tool()
|
||||
async def list_banks() -> str:
|
||||
"""
|
||||
List all available memory banks.
|
||||
@@ -1210,7 +1118,7 @@ def _register_list_banks(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCon
|
||||
def _register_create_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the create_bank tool."""
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("create_bank"))
|
||||
@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.
|
||||
@@ -1274,7 +1182,7 @@ def _register_list_mental_models(mcp: FastMCP, memory: MemoryEngine, config: MCP
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_mental_models"))
|
||||
@mcp.tool()
|
||||
async def list_mental_models(
|
||||
tags: list[str] | None = None,
|
||||
detail: str = "full",
|
||||
@@ -1313,7 +1221,7 @@ def _register_list_mental_models(mcp: FastMCP, memory: MemoryEngine, config: MCP
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_mental_models"))
|
||||
@mcp.tool()
|
||||
async def list_mental_models(
|
||||
tags: list[str] | None = None,
|
||||
detail: str = "full",
|
||||
@@ -1354,7 +1262,7 @@ def _register_get_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPTo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_mental_model"))
|
||||
@mcp.tool()
|
||||
async def get_mental_model(
|
||||
mental_model_id: str,
|
||||
detail: str = "full",
|
||||
@@ -1394,7 +1302,7 @@ def _register_get_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPTo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_mental_model"))
|
||||
@mcp.tool()
|
||||
async def get_mental_model(
|
||||
mental_model_id: str,
|
||||
detail: str = "full",
|
||||
@@ -1436,7 +1344,7 @@ def _register_create_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MC
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("create_mental_model"))
|
||||
@mcp.tool()
|
||||
async def create_mental_model(
|
||||
name: str,
|
||||
source_query: str,
|
||||
@@ -1520,7 +1428,7 @@ def _register_create_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MC
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("create_mental_model"))
|
||||
@mcp.tool()
|
||||
async def create_mental_model(
|
||||
name: str,
|
||||
source_query: str,
|
||||
@@ -1602,7 +1510,7 @@ def _register_update_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MC
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("update_mental_model"))
|
||||
@mcp.tool()
|
||||
async def update_mental_model(
|
||||
mental_model_id: str,
|
||||
name: str | None = None,
|
||||
@@ -1663,7 +1571,7 @@ def _register_update_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MC
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("update_mental_model"))
|
||||
@mcp.tool()
|
||||
async def update_mental_model(
|
||||
mental_model_id: str,
|
||||
name: str | None = None,
|
||||
@@ -1726,7 +1634,7 @@ def _register_delete_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MC
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("delete_mental_model"))
|
||||
@mcp.tool()
|
||||
async def delete_mental_model(
|
||||
mental_model_id: str,
|
||||
bank_id: str | None = None,
|
||||
@@ -1762,7 +1670,7 @@ def _register_delete_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MC
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("delete_mental_model"))
|
||||
@mcp.tool()
|
||||
async def delete_mental_model(
|
||||
mental_model_id: str,
|
||||
) -> dict:
|
||||
@@ -1800,7 +1708,7 @@ def _register_refresh_mental_model(mcp: FastMCP, memory: MemoryEngine, config: M
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("refresh_mental_model"))
|
||||
@mcp.tool()
|
||||
async def refresh_mental_model(
|
||||
mental_model_id: str,
|
||||
bank_id: str | None = None,
|
||||
@@ -1844,7 +1752,7 @@ def _register_refresh_mental_model(mcp: FastMCP, memory: MemoryEngine, config: M
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("refresh_mental_model"))
|
||||
@mcp.tool()
|
||||
async def refresh_mental_model(
|
||||
mental_model_id: str,
|
||||
) -> dict:
|
||||
@@ -1888,7 +1796,7 @@ def _register_clear_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCP
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("clear_mental_model"))
|
||||
@mcp.tool()
|
||||
async def clear_mental_model(
|
||||
mental_model_id: str,
|
||||
bank_id: str | None = None,
|
||||
@@ -1934,7 +1842,7 @@ def _register_clear_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCP
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("clear_mental_model"))
|
||||
@mcp.tool()
|
||||
async def clear_mental_model(
|
||||
mental_model_id: str,
|
||||
) -> dict:
|
||||
@@ -1985,7 +1893,7 @@ def _register_list_directives(mcp: FastMCP, memory: MemoryEngine, config: MCPToo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_directives"))
|
||||
@mcp.tool()
|
||||
async def list_directives(
|
||||
tags: list[str] | None = None,
|
||||
active_only: bool = True,
|
||||
@@ -2023,7 +1931,7 @@ def _register_list_directives(mcp: FastMCP, memory: MemoryEngine, config: MCPToo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_directives"))
|
||||
@mcp.tool()
|
||||
async def list_directives(
|
||||
tags: list[str] | None = None,
|
||||
active_only: bool = True,
|
||||
@@ -2063,7 +1971,7 @@ def _register_create_directive(mcp: FastMCP, memory: MemoryEngine, config: MCPTo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("create_directive"))
|
||||
@mcp.tool()
|
||||
async def create_directive(
|
||||
name: str,
|
||||
content: str,
|
||||
@@ -2109,7 +2017,7 @@ def _register_create_directive(mcp: FastMCP, memory: MemoryEngine, config: MCPTo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("create_directive"))
|
||||
@mcp.tool()
|
||||
async def create_directive(
|
||||
name: str,
|
||||
content: str,
|
||||
@@ -2157,7 +2065,7 @@ def _register_delete_directive(mcp: FastMCP, memory: MemoryEngine, config: MCPTo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("delete_directive"))
|
||||
@mcp.tool()
|
||||
async def delete_directive(
|
||||
directive_id: str,
|
||||
bank_id: str | None = None,
|
||||
@@ -2193,7 +2101,7 @@ def _register_delete_directive(mcp: FastMCP, memory: MemoryEngine, config: MCPTo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("delete_directive"))
|
||||
@mcp.tool()
|
||||
async def delete_directive(
|
||||
directive_id: str,
|
||||
) -> dict:
|
||||
@@ -2236,7 +2144,7 @@ def _register_list_memories(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_memories"))
|
||||
@mcp.tool()
|
||||
async def list_memories(
|
||||
type: str | None = None,
|
||||
q: str | None = None,
|
||||
@@ -2251,7 +2159,7 @@ def _register_list_memories(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
browse/search without relevance ranking.
|
||||
|
||||
Args:
|
||||
type: Filter by fact type: 'world', 'experience', or 'observation'
|
||||
type: Filter by fact type: 'world', 'experience', or 'opinion'
|
||||
q: Optional text search query to filter memories
|
||||
limit: Maximum number of results (default: 100)
|
||||
offset: Pagination offset (default: 0)
|
||||
@@ -2280,7 +2188,7 @@ def _register_list_memories(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_memories"))
|
||||
@mcp.tool()
|
||||
async def list_memories(
|
||||
type: str | None = None,
|
||||
q: str | None = None,
|
||||
@@ -2294,7 +2202,7 @@ def _register_list_memories(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
browse/search without relevance ranking.
|
||||
|
||||
Args:
|
||||
type: Filter by fact type: 'world', 'experience', or 'observation'
|
||||
type: Filter by fact type: 'world', 'experience', or 'opinion'
|
||||
q: Optional text search query to filter memories
|
||||
limit: Maximum number of results (default: 100)
|
||||
offset: Pagination offset (default: 0)
|
||||
@@ -2326,7 +2234,7 @@ def _register_get_memory(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCon
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_memory"))
|
||||
@mcp.tool()
|
||||
async def get_memory(
|
||||
memory_id: str,
|
||||
bank_id: str | None = None,
|
||||
@@ -2362,7 +2270,7 @@ def _register_get_memory(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCon
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_memory"))
|
||||
@mcp.tool()
|
||||
async def get_memory(
|
||||
memory_id: str,
|
||||
) -> dict:
|
||||
@@ -2413,7 +2321,7 @@ def _register_update_memory(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(description=_EDIT_DOC, annotations=_tool_annotations("update_memory"))
|
||||
@mcp.tool()
|
||||
async def update_memory(
|
||||
memory_id: str,
|
||||
text: str | None = None,
|
||||
@@ -2424,7 +2332,7 @@ def _register_update_memory(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
entities: list[str] | None = None,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
f"""{_EDIT_DOC}
|
||||
Args:
|
||||
memory_id: The ID of the memory unit to edit.
|
||||
bank_id: Optional bank (defaults to session bank). Use for cross-bank operations.
|
||||
@@ -2459,7 +2367,7 @@ def _register_update_memory(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(description=_EDIT_DOC, annotations=_tool_annotations("update_memory"))
|
||||
@mcp.tool()
|
||||
async def update_memory(
|
||||
memory_id: str,
|
||||
text: str | None = None,
|
||||
@@ -2469,7 +2377,7 @@ def _register_update_memory(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
fact_type: str | None = None,
|
||||
entities: list[str] | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
f"""{_EDIT_DOC}
|
||||
Args:
|
||||
memory_id: The ID of the memory unit to edit.
|
||||
"""
|
||||
@@ -2518,14 +2426,14 @@ def _register_invalidate_memory(mcp: FastMCP, memory: MemoryEngine, config: MCPT
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(description=_INVALIDATE_DOC, annotations=_tool_annotations("invalidate_memory"))
|
||||
@mcp.tool()
|
||||
async def invalidate_memory(
|
||||
memory_id: str,
|
||||
reason: str | None = None,
|
||||
restore: bool = False,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
f"""{_INVALIDATE_DOC}
|
||||
Args:
|
||||
memory_id: The ID of the memory unit to retire (or restore).
|
||||
reason: Optional free-text reason recorded when invalidating.
|
||||
@@ -2558,13 +2466,13 @@ def _register_invalidate_memory(mcp: FastMCP, memory: MemoryEngine, config: MCPT
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(description=_INVALIDATE_DOC, annotations=_tool_annotations("invalidate_memory"))
|
||||
@mcp.tool()
|
||||
async def invalidate_memory(
|
||||
memory_id: str,
|
||||
reason: str | None = None,
|
||||
restore: bool = False,
|
||||
) -> dict:
|
||||
"""
|
||||
f"""{_INVALIDATE_DOC}
|
||||
Args:
|
||||
memory_id: The ID of the memory unit to retire (or restore).
|
||||
reason: Optional free-text reason recorded when invalidating.
|
||||
@@ -2605,7 +2513,7 @@ def _register_list_documents(mcp: FastMCP, memory: MemoryEngine, config: MCPTool
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_documents"))
|
||||
@mcp.tool()
|
||||
async def list_documents(
|
||||
q: str | None = None,
|
||||
limit: int = 100,
|
||||
@@ -2643,7 +2551,7 @@ def _register_list_documents(mcp: FastMCP, memory: MemoryEngine, config: MCPTool
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_documents"))
|
||||
@mcp.tool()
|
||||
async def list_documents(
|
||||
q: str | None = None,
|
||||
limit: int = 100,
|
||||
@@ -2683,7 +2591,7 @@ def _register_get_document(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsC
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_document"))
|
||||
@mcp.tool()
|
||||
async def get_document(
|
||||
document_id: str,
|
||||
bank_id: str | None = None,
|
||||
@@ -2719,7 +2627,7 @@ def _register_get_document(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsC
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_document"))
|
||||
@mcp.tool()
|
||||
async def get_document(
|
||||
document_id: str,
|
||||
) -> dict:
|
||||
@@ -2757,7 +2665,7 @@ def _register_delete_document(mcp: FastMCP, memory: MemoryEngine, config: MCPToo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("delete_document"))
|
||||
@mcp.tool()
|
||||
async def delete_document(
|
||||
document_id: str,
|
||||
bank_id: str | None = None,
|
||||
@@ -2791,7 +2699,7 @@ def _register_delete_document(mcp: FastMCP, memory: MemoryEngine, config: MCPToo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("delete_document"))
|
||||
@mcp.tool()
|
||||
async def delete_document(
|
||||
document_id: str,
|
||||
) -> dict:
|
||||
@@ -2832,7 +2740,7 @@ def _register_list_operations(mcp: FastMCP, memory: MemoryEngine, config: MCPToo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_operations"))
|
||||
@mcp.tool()
|
||||
async def list_operations(
|
||||
status: str | None = None,
|
||||
limit: int = 20,
|
||||
@@ -2869,7 +2777,7 @@ def _register_list_operations(mcp: FastMCP, memory: MemoryEngine, config: MCPToo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_operations"))
|
||||
@mcp.tool()
|
||||
async def list_operations(
|
||||
status: str | None = None,
|
||||
limit: int = 20,
|
||||
@@ -2908,7 +2816,7 @@ def _register_get_operation(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_operation"))
|
||||
@mcp.tool()
|
||||
async def get_operation(
|
||||
operation_id: str,
|
||||
bank_id: str | None = None,
|
||||
@@ -2942,7 +2850,7 @@ def _register_get_operation(mcp: FastMCP, memory: MemoryEngine, config: MCPTools
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_operation"))
|
||||
@mcp.tool()
|
||||
async def get_operation(
|
||||
operation_id: str,
|
||||
) -> dict:
|
||||
@@ -2978,7 +2886,7 @@ def _register_cancel_operation(mcp: FastMCP, memory: MemoryEngine, config: MCPTo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("cancel_operation"))
|
||||
@mcp.tool()
|
||||
async def cancel_operation(
|
||||
operation_id: str,
|
||||
bank_id: str | None = None,
|
||||
@@ -3010,7 +2918,7 @@ def _register_cancel_operation(mcp: FastMCP, memory: MemoryEngine, config: MCPTo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("cancel_operation"))
|
||||
@mcp.tool()
|
||||
async def cancel_operation(
|
||||
operation_id: str,
|
||||
) -> dict:
|
||||
@@ -3049,7 +2957,7 @@ def _register_list_tags(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConf
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_tags"))
|
||||
@mcp.tool()
|
||||
async def list_tags(
|
||||
q: str | None = None,
|
||||
limit: int = 100,
|
||||
@@ -3086,7 +2994,7 @@ def _register_list_tags(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConf
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("list_tags"))
|
||||
@mcp.tool()
|
||||
async def list_tags(
|
||||
q: str | None = None,
|
||||
limit: int = 100,
|
||||
@@ -3125,7 +3033,7 @@ def _register_get_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfi
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_bank"))
|
||||
@mcp.tool()
|
||||
async def get_bank(
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
@@ -3158,7 +3066,7 @@ def _register_get_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfi
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_bank"))
|
||||
@mcp.tool()
|
||||
async def get_bank() -> dict:
|
||||
"""
|
||||
Get the profile of this memory bank.
|
||||
@@ -3188,7 +3096,7 @@ def _register_get_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfi
|
||||
def _register_get_bank_stats(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the get_bank_stats tool (multi-bank only)."""
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("get_bank_stats"))
|
||||
@mcp.tool()
|
||||
async def get_bank_stats(
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
@@ -3261,7 +3169,7 @@ def _register_update_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("update_bank"))
|
||||
@mcp.tool()
|
||||
async def update_bank(
|
||||
name: str | None = None,
|
||||
mission: str | None = None,
|
||||
@@ -3283,8 +3191,7 @@ def _register_update_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo
|
||||
- retain_mission: Steers what gets extracted during retain().
|
||||
- retain_extraction_mode: 'concise' (default), 'verbose', or 'custom'.
|
||||
- retain_custom_instructions: Custom extraction prompt (active when mode is 'custom').
|
||||
- retain_chunk_size: Target maximum characters for each content chunk.
|
||||
- retain_structured_chunk_size: Maximum characters for a single JSONL line or conversation turn to keep whole.
|
||||
- retain_chunk_size: Maximum token size for each content chunk.
|
||||
- retain_chunk_batch_size: Number of chunks to process in parallel.
|
||||
- enable_observations: Toggle observation consolidation after retain().
|
||||
- observations_mission: Controls observation synthesis rules.
|
||||
@@ -3322,7 +3229,7 @@ def _register_update_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("update_bank"))
|
||||
@mcp.tool()
|
||||
async def update_bank(
|
||||
name: str | None = None,
|
||||
mission: str | None = None,
|
||||
@@ -3343,8 +3250,7 @@ def _register_update_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo
|
||||
- retain_mission: Steers what gets extracted during retain().
|
||||
- retain_extraction_mode: 'concise' (default), 'verbose', or 'custom'.
|
||||
- retain_custom_instructions: Custom extraction prompt (active when mode is 'custom').
|
||||
- retain_chunk_size: Target maximum characters for each content chunk.
|
||||
- retain_structured_chunk_size: Maximum characters for a single JSONL line or conversation turn to keep whole.
|
||||
- retain_chunk_size: Maximum token size for each content chunk.
|
||||
- retain_chunk_batch_size: Number of chunks to process in parallel.
|
||||
- enable_observations: Toggle observation consolidation after retain().
|
||||
- observations_mission: Controls observation synthesis rules.
|
||||
@@ -3385,7 +3291,7 @@ def _register_delete_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("delete_bank"))
|
||||
@mcp.tool()
|
||||
async def delete_bank(
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
@@ -3417,7 +3323,7 @@ def _register_delete_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("delete_bank"))
|
||||
@mcp.tool()
|
||||
async def delete_bank() -> dict:
|
||||
"""
|
||||
Delete this memory bank and all its data.
|
||||
@@ -3448,7 +3354,7 @@ def _register_clear_memories(mcp: FastMCP, memory: MemoryEngine, config: MCPTool
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("clear_memories"))
|
||||
@mcp.tool()
|
||||
async def clear_memories(
|
||||
type: str | None = None,
|
||||
bank_id: str | None = None,
|
||||
@@ -3459,7 +3365,7 @@ def _register_clear_memories(mcp: FastMCP, memory: MemoryEngine, config: MCPTool
|
||||
Optionally filter by fact type to only clear specific kinds of memories.
|
||||
|
||||
Args:
|
||||
type: Optional fact type filter: 'world', 'experience', or 'observation'. If not specified, clears all.
|
||||
type: Optional fact type filter: 'world', 'experience', or 'opinion'. If not specified, clears all.
|
||||
bank_id: Optional bank (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
@@ -3483,7 +3389,7 @@ def _register_clear_memories(mcp: FastMCP, memory: MemoryEngine, config: MCPTool
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool(annotations=_tool_annotations("clear_memories"))
|
||||
@mcp.tool()
|
||||
async def clear_memories(
|
||||
type: str | None = None,
|
||||
) -> dict:
|
||||
@@ -3493,7 +3399,7 @@ def _register_clear_memories(mcp: FastMCP, memory: MemoryEngine, config: MCPTool
|
||||
Optionally filter by fact type to only clear specific kinds of memories.
|
||||
|
||||
Args:
|
||||
type: Optional fact type filter: 'world', 'experience', or 'observation'. If not specified, clears all.
|
||||
type: Optional fact type filter: 'world', 'experience', or 'opinion'. If not specified, clears all.
|
||||
"""
|
||||
try:
|
||||
target_bank = config.bank_id_resolver()
|
||||
|
||||
@@ -11,17 +11,15 @@ This module provides metrics for:
|
||||
- Database connection pool metrics
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import importlib
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
|
||||
_resource_mod = importlib.import_module("resource") if importlib.util.find_spec("resource") else None
|
||||
import threading
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from typing import TYPE_CHECKING, Callable, NamedTuple
|
||||
from typing import TYPE_CHECKING, Callable
|
||||
|
||||
from opentelemetry import metrics
|
||||
from opentelemetry.exporter.prometheus import PrometheusMetricReader
|
||||
@@ -41,32 +39,6 @@ def _get_tenant() -> str:
|
||||
return get_current_schema()
|
||||
|
||||
|
||||
def _is_client_cancellation(exc: BaseException) -> bool:
|
||||
"""Whether *exc* is a client-disconnect cancellation rather than a failure.
|
||||
|
||||
An abandoned recall/reflect raises OperationCancelledError (issue #2122);
|
||||
the HTTP layer re-raises it as ``HTTPException(499) from exc`` (see
|
||||
api/http.py run_cancellable_on_disconnect). The exception itself, or any
|
||||
link in its ``__cause__`` chain, being an OperationCancelledError marks it
|
||||
as a cancellation. Matching on the cause chain rather than a bare status
|
||||
code avoids misclassifying an unrelated 499 as a cancellation. Per the
|
||||
engine contract a cancellation is "not a failure to retry or report"
|
||||
(cancellation.OperationCancelledError), so it must not be counted against
|
||||
``hindsight.operation.total``.
|
||||
"""
|
||||
# Imported lazily to avoid import-time coupling (cf. _get_tenant above).
|
||||
from hindsight_api.cancellation import OperationCancelledError
|
||||
|
||||
cause: BaseException | None = exc
|
||||
seen: set[int] = set() # guard against a cyclic __cause__ chain
|
||||
while cause is not None and id(cause) not in seen:
|
||||
if isinstance(cause, OperationCancelledError):
|
||||
return True
|
||||
seen.add(id(cause))
|
||||
cause = cause.__cause__
|
||||
return False
|
||||
|
||||
|
||||
# Custom bucket boundaries for operation duration (in seconds)
|
||||
# Fine granularity in 0-30s range where most operations complete
|
||||
DURATION_BUCKETS = (0.1, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 20.0, 30.0, 60.0, 120.0)
|
||||
@@ -77,28 +49,6 @@ LLM_DURATION_BUCKETS = (0.1, 0.25, 0.5, 1.0, 2.0, 3.0, 5.0, 10.0, 15.0, 30.0, 60
|
||||
# HTTP request duration buckets (millisecond-level for fast endpoints)
|
||||
HTTP_DURATION_BUCKETS = (0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0)
|
||||
|
||||
# How often the backlog / queue-depth gauge caches are refreshed (seconds).
|
||||
# The counts are aggregate COUNT queries, so a background task refreshes a
|
||||
# cache and the observable gauges read from it — keeping the /metrics scrape
|
||||
# path synchronous (the same reason the db-pool gauges read cached state).
|
||||
BACKLOG_METRICS_REFRESH_SECONDS = 30
|
||||
|
||||
|
||||
class _AsyncOpKey(NamedTuple):
|
||||
"""Cache / label key for the async-operation queue gauge."""
|
||||
|
||||
tenant: str
|
||||
operation_type: str
|
||||
status: str
|
||||
bank_id: str | None
|
||||
|
||||
|
||||
class _BacklogKey(NamedTuple):
|
||||
"""Cache / label key for the consolidation backlog and failed gauges."""
|
||||
|
||||
tenant: str
|
||||
bank_id: str | None
|
||||
|
||||
|
||||
def get_token_bucket(token_count: int) -> str:
|
||||
"""
|
||||
@@ -137,27 +87,6 @@ def get_token_bucket(token_count: int) -> str:
|
||||
return "50k+"
|
||||
|
||||
|
||||
# Template unbounded id segments before a path is used as the low-cardinality
|
||||
# "endpoint" metric label. A raw per-bank path segment (e.g. user-123) would
|
||||
# otherwise create one never-evicted OTel series per bank.
|
||||
_METRIC_BANK_SEGMENT_RE = re.compile(r"(/banks/)[^/]+")
|
||||
_METRIC_UUID_RE = re.compile(r"/[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}")
|
||||
_METRIC_NUMERIC_ID_RE = re.compile(r"/\d+(?=/|$)")
|
||||
|
||||
|
||||
def normalize_http_endpoint(path: str) -> str:
|
||||
"""Template high-cardinality id segments in an HTTP path for safe metric labeling.
|
||||
|
||||
Collapses the "/banks/<id>" segment (any bank id, including non-numeric ones like
|
||||
"user-123"), UUIDs, and numeric ids to placeholders so the "endpoint" metric label
|
||||
has bounded cardinality. Analogous to get_token_bucket for token counts.
|
||||
"""
|
||||
path = _METRIC_BANK_SEGMENT_RE.sub(r"\g<1>{bank_id}", path)
|
||||
path = _METRIC_UUID_RE.sub("/{id}", path)
|
||||
path = _METRIC_NUMERIC_ID_RE.sub("/{id}", path)
|
||||
return path
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Global meter instance
|
||||
@@ -246,19 +175,6 @@ class MetricsCollectorBase:
|
||||
"""Context manager to record operation duration and status."""
|
||||
raise NotImplementedError
|
||||
|
||||
def record_operation_result(
|
||||
self,
|
||||
operation: str,
|
||||
bank_id: str,
|
||||
success: bool,
|
||||
duration: float,
|
||||
source: str = "api",
|
||||
budget: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
):
|
||||
"""Record a single completed operation with an explicit success label."""
|
||||
raise NotImplementedError
|
||||
|
||||
def record_llm_call(
|
||||
self,
|
||||
provider: str,
|
||||
@@ -312,19 +228,6 @@ class NoOpMetricsCollector(MetricsCollectorBase):
|
||||
"""No-op context manager."""
|
||||
yield
|
||||
|
||||
def record_operation_result(
|
||||
self,
|
||||
operation: str,
|
||||
bank_id: str,
|
||||
success: bool,
|
||||
duration: float,
|
||||
source: str = "api",
|
||||
budget: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
):
|
||||
"""No-op operation result recording."""
|
||||
pass
|
||||
|
||||
def record_llm_call(
|
||||
self,
|
||||
provider: str,
|
||||
@@ -432,13 +335,6 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
# DB pool metrics holder (set via set_db_pool)
|
||||
self._db_pool: "asyncpg.Pool | None" = None
|
||||
|
||||
# Backlog / queue-depth gauge caches, refreshed by a background task
|
||||
# (see _setup_backlog_metrics) so the scrape path stays synchronous.
|
||||
self._async_ops_counts: dict[_AsyncOpKey, int] = {}
|
||||
self._consolidation_backlog: dict[_BacklogKey, int] = {}
|
||||
self._consolidation_failed: dict[_BacklogKey, int] = {}
|
||||
self._backlog_task: "asyncio.Task | None" = None
|
||||
|
||||
@contextmanager
|
||||
def record_operation(
|
||||
self,
|
||||
@@ -464,51 +360,6 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
max_tokens: Optional max tokens for the operation
|
||||
"""
|
||||
start_time = time.time()
|
||||
success = True
|
||||
cancelled = False
|
||||
try:
|
||||
yield
|
||||
except Exception as exc:
|
||||
# A client disconnect cancels the operation cooperatively (#2122),
|
||||
# raised as OperationCancelledError and re-raised by the HTTP layer
|
||||
# as HTTPException(499) from it. An abandoned request is neither a
|
||||
# success nor a failure, so it is excluded from the metric entirely
|
||||
# rather than inflating either the failure or the success rate on
|
||||
# hindsight.operation.total.
|
||||
if _is_client_cancellation(exc):
|
||||
cancelled = True
|
||||
else:
|
||||
success = False
|
||||
raise
|
||||
finally:
|
||||
if not cancelled:
|
||||
self.record_operation_result(
|
||||
operation,
|
||||
bank_id,
|
||||
success=success,
|
||||
duration=time.time() - start_time,
|
||||
source=source,
|
||||
budget=budget,
|
||||
max_tokens=max_tokens,
|
||||
)
|
||||
|
||||
def record_operation_result(
|
||||
self,
|
||||
operation: str,
|
||||
bank_id: str,
|
||||
success: bool,
|
||||
duration: float,
|
||||
source: str = "api",
|
||||
budget: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
):
|
||||
"""Record a single completed operation (duration + count) with a success label.
|
||||
|
||||
Direct (non-context-manager) recording for code paths that need explicit
|
||||
success control rather than the exception-based ``record_operation`` — e.g.
|
||||
the async worker, where deferrals/retries are not terminal outcomes and must
|
||||
not be counted as completions.
|
||||
"""
|
||||
attributes = {
|
||||
"operation": operation,
|
||||
"source": source,
|
||||
@@ -520,13 +371,22 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
attributes["budget"] = budget
|
||||
if max_tokens:
|
||||
attributes["max_tokens"] = str(max_tokens)
|
||||
attributes["success"] = str(success).lower()
|
||||
|
||||
# Record duration
|
||||
self.operation_duration.record(duration, attributes)
|
||||
success = True
|
||||
try:
|
||||
yield
|
||||
except Exception:
|
||||
success = False
|
||||
raise
|
||||
finally:
|
||||
duration = time.time() - start_time
|
||||
attributes["success"] = str(success).lower()
|
||||
|
||||
# Record operation count
|
||||
self.operation_total.add(1, attributes)
|
||||
# Record duration
|
||||
self.operation_duration.record(duration, attributes)
|
||||
|
||||
# Record operation count
|
||||
self.operation_total.add(1, attributes)
|
||||
|
||||
def record_llm_call(
|
||||
self,
|
||||
@@ -731,10 +591,6 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
"""
|
||||
self._db_pool = pool
|
||||
self._setup_db_pool_metrics()
|
||||
from .config import get_config
|
||||
|
||||
if get_config().metrics_backlog_enabled:
|
||||
self._setup_backlog_metrics()
|
||||
|
||||
def _setup_db_pool_metrics(self):
|
||||
"""Set up observable gauges for database pool metrics."""
|
||||
@@ -800,192 +656,6 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
unit="{connections}",
|
||||
)
|
||||
|
||||
def _setup_backlog_metrics(self):
|
||||
"""Observable gauges for the async-operation queue and the
|
||||
consolidation backlog.
|
||||
|
||||
These mirror fields the bank-stats endpoint already computes
|
||||
(``operations_by_status``, ``pending_consolidation``,
|
||||
``failed_consolidation``) but expose them as scrapable gauges, so
|
||||
queue depth and backlog can be trended and alerted on instead of only
|
||||
polled per-bank over HTTP. The two motivating questions both come for
|
||||
free here: "is the worker keeping up?" (async-op queue) and "is the
|
||||
knowledge base caught up?" (consolidation backlog) — including the
|
||||
``processing`` state, which is the only signal that surfaces a hung
|
||||
operation stuck holding a worker slot.
|
||||
|
||||
Counts are aggregate ``COUNT`` queries, so a background task refreshes
|
||||
a cache every ``BACKLOG_METRICS_REFRESH_SECONDS`` and these callbacks
|
||||
read it — keeping the scrape path synchronous, the same approach as
|
||||
the db-pool gauges above.
|
||||
"""
|
||||
if self._backlog_task is not None:
|
||||
return # already started for this collector
|
||||
|
||||
def get_async_operations(_options):
|
||||
for key, value in list(self._async_ops_counts.items()):
|
||||
attrs = {"tenant": key.tenant, "operation_type": key.operation_type, "status": key.status}
|
||||
if key.bank_id is not None:
|
||||
attrs["bank_id"] = key.bank_id
|
||||
yield metrics.Observation(value, attrs)
|
||||
|
||||
def get_consolidation_backlog(_options):
|
||||
for key, value in list(self._consolidation_backlog.items()):
|
||||
attrs = {"tenant": key.tenant}
|
||||
if key.bank_id is not None:
|
||||
attrs["bank_id"] = key.bank_id
|
||||
yield metrics.Observation(value, attrs)
|
||||
|
||||
def get_consolidation_failed(_options):
|
||||
for key, value in list(self._consolidation_failed.items()):
|
||||
attrs = {"tenant": key.tenant}
|
||||
if key.bank_id is not None:
|
||||
attrs["bank_id"] = key.bank_id
|
||||
yield metrics.Observation(value, attrs)
|
||||
|
||||
self.meter.create_observable_gauge(
|
||||
name="hindsight.async_operations",
|
||||
callbacks=[get_async_operations],
|
||||
description="Async operations in a non-terminal state, by operation_type and status "
|
||||
"(pending=queued backlog, processing=in-flight, failed=stranded)",
|
||||
unit="{operations}",
|
||||
)
|
||||
self.meter.create_observable_gauge(
|
||||
name="hindsight.consolidation.backlog",
|
||||
callbacks=[get_consolidation_backlog],
|
||||
description="Source memories (experience/world) not yet consolidated into observations",
|
||||
unit="{memories}",
|
||||
)
|
||||
self.meter.create_observable_gauge(
|
||||
name="hindsight.consolidation.failed",
|
||||
callbacks=[get_consolidation_failed],
|
||||
description="Source memories whose consolidation permanently failed "
|
||||
"(recoverable via the consolidation recovery endpoint)",
|
||||
unit="{memories}",
|
||||
)
|
||||
|
||||
# Drive the caches from a background task on the running loop.
|
||||
# set_db_pool runs during async startup, so a loop is normally present;
|
||||
# if not, the gauges simply stay empty rather than crashing collection.
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
logger.warning("No running event loop; backlog metrics disabled")
|
||||
return
|
||||
# Process-lifetime task: there is no collector teardown hook to cancel it
|
||||
# on, so it's torn down with the event loop at process shutdown. If a
|
||||
# shutdown path is ever added, cancel self._backlog_task there.
|
||||
self._backlog_task = loop.create_task(self._backlog_refresh_loop())
|
||||
|
||||
async def _backlog_refresh_loop(self):
|
||||
"""Periodically refresh the backlog / queue-depth caches."""
|
||||
while True:
|
||||
try:
|
||||
await self._refresh_backlog()
|
||||
except Exception:
|
||||
logger.debug("Backlog metrics refresh failed", exc_info=True)
|
||||
await asyncio.sleep(BACKLOG_METRICS_REFRESH_SECONDS)
|
||||
|
||||
async def _refresh_backlog(self):
|
||||
"""Recount the async-operation queue and consolidation backlog across
|
||||
every provisioned Hindsight schema.
|
||||
|
||||
Per-bank labels are gated behind ``metrics_include_bank_id`` (off by
|
||||
default) to keep cardinality bounded; when off, counts are aggregated
|
||||
per tenant/schema. All SQL here is PostgreSQL-specific (``FILTER``,
|
||||
``information_schema``), which is consistent with this collector
|
||||
already being bound to an asyncpg pool.
|
||||
"""
|
||||
if self._db_pool is None:
|
||||
return
|
||||
|
||||
async_ops: dict[_AsyncOpKey, int] = {}
|
||||
backlog: dict[_BacklogKey, int] = {}
|
||||
failed: dict[_BacklogKey, int] = {}
|
||||
per_bank = self._include_bank_id
|
||||
bank_sel = "bank_id, " if per_bank else ""
|
||||
bank_grp = " GROUP BY bank_id" if per_bank else ""
|
||||
|
||||
async with self._db_pool.acquire() as conn:
|
||||
# memory_units is the central per-tenant table; its presence marks a
|
||||
# provisioned Hindsight schema.
|
||||
schema_rows = await conn.fetch(
|
||||
"SELECT table_schema FROM information_schema.tables WHERE table_name = 'memory_units'"
|
||||
)
|
||||
for schema_row in schema_rows:
|
||||
schema = schema_row["table_schema"]
|
||||
|
||||
# Worker queue depth — mirrors operations_by_status, split by
|
||||
# operation_type. Terminal states (completed/cancelled) are
|
||||
# excluded on purpose: a gauge of finished work grows without
|
||||
# bound and says nothing about current load.
|
||||
# Index: idx_async_operations_status.
|
||||
ops_grp = "operation_type, status" + (", bank_id" if per_bank else "")
|
||||
try:
|
||||
rows = await conn.fetch(
|
||||
f"SELECT operation_type, status, {bank_sel}COUNT(*) AS count "
|
||||
f'FROM "{schema}".async_operations '
|
||||
"WHERE status IN ('pending', 'processing', 'failed') "
|
||||
f"GROUP BY {ops_grp}"
|
||||
)
|
||||
for row in rows:
|
||||
bank = row["bank_id"] if per_bank else None
|
||||
key = _AsyncOpKey(schema, row["operation_type"] or "unknown", row["status"], bank)
|
||||
async_ops[key] = async_ops.get(key, 0) + int(row["count"])
|
||||
except Exception:
|
||||
logger.debug("Async-ops queue query failed for schema %s", schema, exc_info=True)
|
||||
|
||||
# Consolidation backlog + stranded counts. Two separate COUNT(*)
|
||||
# queries rather than one with two FILTERs — each WHERE matches a
|
||||
# partial-index predicate exactly:
|
||||
# idx_memory_units_unconsolidated WHERE consolidated_at IS NULL ...
|
||||
# idx_memory_units_consolidation_failed WHERE consolidation_failed_at IS NOT NULL ...
|
||||
# GROUP BY bank_id still composes — bank_id is each index's lead column.
|
||||
#
|
||||
# The backlog count runs with seqscan disabled in a scoped
|
||||
# transaction. The partial index matches its predicate, but
|
||||
# `consolidated_at IS NULL` is true for a large fraction of the
|
||||
# table (every observation has a null consolidated_at), so the
|
||||
# planner misjudges selectivity and otherwise seq-scans the whole
|
||||
# (largest) table on every refresh — verified on a 114k-row table
|
||||
# via EXPLAIN: seq scan ~92 ms vs index scan ~0.1 ms. SET LOCAL
|
||||
# forces the index path and resets at transaction end. The failed
|
||||
# count below needs no such nudge: `consolidation_failed_at IS NOT
|
||||
# NULL` is rare, so its index is chosen on cost.
|
||||
try:
|
||||
async with conn.transaction():
|
||||
await conn.execute("SET LOCAL enable_seqscan = off")
|
||||
rows = await conn.fetch(
|
||||
f"SELECT {bank_sel}COUNT(*) AS count "
|
||||
f'FROM "{schema}".memory_units '
|
||||
"WHERE consolidated_at IS NULL AND fact_type IN ('experience', 'world')"
|
||||
f"{bank_grp}"
|
||||
)
|
||||
for row in rows:
|
||||
bank = row["bank_id"] if per_bank else None
|
||||
key = _BacklogKey(schema, bank)
|
||||
backlog[key] = backlog.get(key, 0) + int(row["count"])
|
||||
except Exception:
|
||||
logger.debug("Consolidation backlog query failed for schema %s", schema, exc_info=True)
|
||||
|
||||
try:
|
||||
rows = await conn.fetch(
|
||||
f"SELECT {bank_sel}COUNT(*) AS count "
|
||||
f'FROM "{schema}".memory_units '
|
||||
"WHERE consolidation_failed_at IS NOT NULL AND fact_type IN ('experience', 'world')"
|
||||
f"{bank_grp}"
|
||||
)
|
||||
for row in rows:
|
||||
bank = row["bank_id"] if per_bank else None
|
||||
key = _BacklogKey(schema, bank)
|
||||
failed[key] = failed.get(key, 0) + int(row["count"])
|
||||
except Exception:
|
||||
logger.debug("Consolidation failed query failed for schema %s", schema, exc_info=True)
|
||||
|
||||
self._async_ops_counts = async_ops
|
||||
self._consolidation_backlog = backlog
|
||||
self._consolidation_failed = failed
|
||||
|
||||
|
||||
# Global metrics collector instance (defaults to no-op)
|
||||
_metrics_collector: MetricsCollectorBase = NoOpMetricsCollector()
|
||||
|
||||
@@ -27,12 +27,10 @@ from alembic.config import Config
|
||||
from alembic.script.revision import ResolutionError
|
||||
from alembic.util.exc import CommandError
|
||||
from sqlalchemy import Connection, create_engine, text
|
||||
from sqlalchemy.pool import NullPool
|
||||
|
||||
from ._pg_search import normalize_pg_search_tokenizer, pg_search_bm25_columns
|
||||
from ._vector_index import (
|
||||
bootstrap_extension,
|
||||
configured_vector_extension,
|
||||
detect_vector_extension,
|
||||
index_type_keyword,
|
||||
index_using_clause,
|
||||
@@ -61,86 +59,6 @@ def _detect_vector_extension(conn, vector_extension: str = "pgvector") -> str:
|
||||
return detect_vector_extension(conn, vector_extension)
|
||||
|
||||
|
||||
def _ensure_pgvector_extension_in_public(conn: Connection) -> None:
|
||||
"""Ensure pgvector is installed before pgvector-backed migrations run."""
|
||||
logger.debug("Checking pgvector extension availability...")
|
||||
|
||||
# First, check if extension already exists
|
||||
ext_check = conn.execute(
|
||||
text(
|
||||
"SELECT extname, nspname FROM pg_extension e "
|
||||
"JOIN pg_namespace n ON e.extnamespace = n.oid "
|
||||
"WHERE extname = 'vector'"
|
||||
)
|
||||
).fetchone()
|
||||
|
||||
if ext_check:
|
||||
# Extension exists - check if in correct schema
|
||||
ext_schema = ext_check[1]
|
||||
if ext_schema == "public":
|
||||
logger.info("pgvector extension found in public schema - ready to use")
|
||||
else:
|
||||
# Extension in wrong schema - try to fix if we have permissions
|
||||
logger.warning(
|
||||
f"pgvector extension found in schema '{ext_schema}' instead of 'public'. Attempting to relocate..."
|
||||
)
|
||||
try:
|
||||
conn.execute(text("DROP EXTENSION vector CASCADE"))
|
||||
conn.execute(text("SET search_path TO public"))
|
||||
conn.execute(text("CREATE EXTENSION vector"))
|
||||
conn.commit()
|
||||
logger.info("pgvector extension relocated to public schema")
|
||||
except Exception as e:
|
||||
# Failed to relocate - log but don't fail if extension exists somewhere
|
||||
logger.warning(
|
||||
f"Could not relocate pgvector extension to public schema: {e}. "
|
||||
f"Continuing with extension in '{ext_schema}' schema."
|
||||
)
|
||||
conn.rollback()
|
||||
else:
|
||||
# Extension doesn't exist - try to install
|
||||
logger.info("pgvector extension not found, attempting to install...")
|
||||
try:
|
||||
conn.execute(text("SET search_path TO public"))
|
||||
conn.execute(text("CREATE EXTENSION vector"))
|
||||
conn.commit()
|
||||
logger.info("pgvector extension installed in public schema")
|
||||
except Exception as e:
|
||||
# Installation failed - this is only fatal if extension truly doesn't exist
|
||||
# Check one more time in case another process installed it
|
||||
conn.rollback()
|
||||
ext_recheck = conn.execute(
|
||||
text(
|
||||
"SELECT nspname FROM pg_extension e "
|
||||
"JOIN pg_namespace n ON e.extnamespace = n.oid "
|
||||
"WHERE extname = 'vector'"
|
||||
)
|
||||
).fetchone()
|
||||
|
||||
if ext_recheck:
|
||||
logger.warning(
|
||||
f"Could not install pgvector extension (permission denied?), "
|
||||
f"but extension exists in '{ext_recheck[0]}' schema. Continuing..."
|
||||
)
|
||||
else:
|
||||
# Extension truly doesn't exist and we can't install it
|
||||
logger.error(
|
||||
f"pgvector extension is not installed and cannot be installed: {e}. "
|
||||
f"Please ensure pgvector is installed by a database administrator. "
|
||||
f"See: https://github.com/pgvector/pgvector#installation"
|
||||
)
|
||||
raise RuntimeError(
|
||||
"pgvector extension is required but not installed. Please install it with: CREATE EXTENSION vector;"
|
||||
) from e
|
||||
|
||||
|
||||
def _bootstrap_vector_extension_for_migrations(conn: Connection, vector_extension: str) -> None:
|
||||
"""Bootstrap the configured vector backend before schema migrations run."""
|
||||
if vector_extension == "pgvector":
|
||||
_ensure_pgvector_extension_in_public(conn)
|
||||
bootstrap_extension(conn, vector_extension)
|
||||
|
||||
|
||||
def _drop_per_bank_vector_indexes(conn: Connection, schema_name: str) -> None:
|
||||
"""Drop per-bank partial memory_units vector indexes after global ScaNN is ready."""
|
||||
rows = conn.execute(
|
||||
@@ -329,14 +247,7 @@ def run_migrations(
|
||||
# 2. After acquiring the lock, COMMIT the transaction on the advisory-lock
|
||||
# connection itself before running migrations. pg_advisory_lock is
|
||||
# session-level, so the lock survives the COMMIT.
|
||||
# NullPool: do not retain the connection in a pool after the migration.
|
||||
# Each schema migration opens a few short-lived engines (here plus the
|
||||
# ensure_* steps); with the default QueuePool those connections linger
|
||||
# until GC, and running many schemas in parallel (migration_concurrency)
|
||||
# multiplies that footprint and exhausts max_connections — observed as
|
||||
# "FATAL: sorry, too many clients already" sweeping 20k schemas at
|
||||
# concurrency 12. NullPool closes the connection on return.
|
||||
engine = create_engine(migration_url, poolclass=NullPool)
|
||||
engine = create_engine(migration_url)
|
||||
with engine.connect() as conn:
|
||||
logger.debug(f"Acquiring migration advisory lock for schema '{schema_name}' (id={lock_id})...")
|
||||
while True:
|
||||
@@ -356,8 +267,83 @@ def run_migrations(
|
||||
logger.debug("Migration advisory lock acquired")
|
||||
|
||||
try:
|
||||
vector_extension = configured_vector_extension()
|
||||
_bootstrap_vector_extension_for_migrations(conn, vector_extension)
|
||||
# Ensure pgvector extension is installed globally BEFORE schema migrations
|
||||
# This is critical: the extension must exist database-wide before any schema
|
||||
# migrations run, otherwise custom schemas won't have access to vector types
|
||||
logger.debug("Checking pgvector extension availability...")
|
||||
|
||||
# First, check if extension already exists
|
||||
ext_check = conn.execute(
|
||||
text(
|
||||
"SELECT extname, nspname FROM pg_extension e "
|
||||
"JOIN pg_namespace n ON e.extnamespace = n.oid "
|
||||
"WHERE extname = 'vector'"
|
||||
)
|
||||
).fetchone()
|
||||
|
||||
if ext_check:
|
||||
# Extension exists - check if in correct schema
|
||||
ext_schema = ext_check[1]
|
||||
if ext_schema == "public":
|
||||
logger.info("pgvector extension found in public schema - ready to use")
|
||||
else:
|
||||
# Extension in wrong schema - try to fix if we have permissions
|
||||
logger.warning(
|
||||
f"pgvector extension found in schema '{ext_schema}' instead of 'public'. "
|
||||
f"Attempting to relocate..."
|
||||
)
|
||||
try:
|
||||
conn.execute(text("DROP EXTENSION vector CASCADE"))
|
||||
conn.execute(text("SET search_path TO public"))
|
||||
conn.execute(text("CREATE EXTENSION vector"))
|
||||
conn.commit()
|
||||
logger.info("pgvector extension relocated to public schema")
|
||||
except Exception as e:
|
||||
# Failed to relocate - log but don't fail if extension exists somewhere
|
||||
logger.warning(
|
||||
f"Could not relocate pgvector extension to public schema: {e}. "
|
||||
f"Continuing with extension in '{ext_schema}' schema."
|
||||
)
|
||||
conn.rollback()
|
||||
else:
|
||||
# Extension doesn't exist - try to install
|
||||
logger.info("pgvector extension not found, attempting to install...")
|
||||
try:
|
||||
conn.execute(text("SET search_path TO public"))
|
||||
conn.execute(text("CREATE EXTENSION vector"))
|
||||
conn.commit()
|
||||
logger.info("pgvector extension installed in public schema")
|
||||
except Exception as e:
|
||||
# Installation failed - this is only fatal if extension truly doesn't exist
|
||||
# Check one more time in case another process installed it
|
||||
conn.rollback()
|
||||
ext_recheck = conn.execute(
|
||||
text(
|
||||
"SELECT nspname FROM pg_extension e "
|
||||
"JOIN pg_namespace n ON e.extnamespace = n.oid "
|
||||
"WHERE extname = 'vector'"
|
||||
)
|
||||
).fetchone()
|
||||
|
||||
if ext_recheck:
|
||||
logger.warning(
|
||||
f"Could not install pgvector extension (permission denied?), "
|
||||
f"but extension exists in '{ext_recheck[0]}' schema. Continuing..."
|
||||
)
|
||||
else:
|
||||
# Extension truly doesn't exist and we can't install it
|
||||
logger.error(
|
||||
f"pgvector extension is not installed and cannot be installed: {e}. "
|
||||
f"Please ensure pgvector is installed by a database administrator. "
|
||||
f"See: https://github.com/pgvector/pgvector#installation"
|
||||
)
|
||||
raise RuntimeError(
|
||||
"pgvector extension is required but not installed. "
|
||||
"Please install it with: CREATE EXTENSION vector;"
|
||||
) from e
|
||||
|
||||
vector_extension = os.getenv("HINDSIGHT_API_VECTOR_EXTENSION", "pgvector").lower()
|
||||
bootstrap_extension(conn, vector_extension)
|
||||
|
||||
# Commit any pending transaction on the advisory-lock connection
|
||||
# before running migrations. Some code paths above (e.g., the
|
||||
@@ -414,7 +400,7 @@ def check_migration_status(
|
||||
return None, None
|
||||
|
||||
# Get current revision from database
|
||||
engine = create_engine(to_libpq_url(database_url), poolclass=NullPool)
|
||||
engine = create_engine(to_libpq_url(database_url))
|
||||
with engine.connect() as connection:
|
||||
context = MigrationContext.configure(connection)
|
||||
current_rev = context.get_current_revision()
|
||||
@@ -587,7 +573,7 @@ def ensure_embedding_dimension(
|
||||
"""
|
||||
schema_name = schema or "public"
|
||||
|
||||
engine = create_engine(to_libpq_url(database_url), poolclass=NullPool)
|
||||
engine = create_engine(to_libpq_url(database_url))
|
||||
with engine.connect() as conn:
|
||||
# Check if memory_units table exists (proxy for schema being initialized)
|
||||
table_exists = conn.execute(
|
||||
@@ -610,10 +596,6 @@ def ensure_embedding_dimension(
|
||||
|
||||
_migrate_table_embedding_dimension(conn, schema_name, "memory_units", required_dimension, vector_ext)
|
||||
_migrate_table_embedding_dimension(conn, schema_name, "mental_models", required_dimension, vector_ext)
|
||||
# NOTE: invalidated_memory_units is deliberately omitted. The curation archive has no
|
||||
# embedding column at all (dropped in migration d4f6a8c2e1b3) — invalidate stores no
|
||||
# embedding and revert recomputes one — so there is no archive vector to re-dimension
|
||||
# and a model switch can't trip a dimension mismatch there (#2209).
|
||||
|
||||
|
||||
def ensure_vector_extension(
|
||||
@@ -640,7 +622,7 @@ def ensure_vector_extension(
|
||||
"""
|
||||
schema_name = schema or "public"
|
||||
|
||||
engine = create_engine(to_libpq_url(database_url), poolclass=NullPool)
|
||||
engine = create_engine(to_libpq_url(database_url))
|
||||
with engine.connect() as conn:
|
||||
# Detect which vector extension should be used
|
||||
target_ext = _detect_vector_extension(conn, vector_extension)
|
||||
@@ -692,20 +674,24 @@ def ensure_vector_extension(
|
||||
|
||||
if not current_index_info:
|
||||
if table_name == "memory_units" and uses_per_bank_vector_indexes(target_ext):
|
||||
# Per-bank backends never use a GLOBAL memory_units vector index.
|
||||
# Every vector search is bank + fact_type scoped and served by the
|
||||
# per-(bank, fact_type) partial indexes created at bank-creation time
|
||||
# (bank_utils.create_bank_vector_indexes); the planner never picks a
|
||||
# global index when bank_id is in the WHERE clause, which is exactly
|
||||
# why migration d5e6f7a8b9c0 drops it for these backends. So don't
|
||||
# create one here either — not even on an empty schema with no per-bank
|
||||
# indexes yet (those are built when the first bank is created). Verified
|
||||
# via EXPLAIN: the query uses idx_mu_emb_* whether or not the global
|
||||
# index exists, so creating it is dead weight.
|
||||
logger.debug(
|
||||
f"Per-bank vector backend ({target_ext}); skipping global {index_name} creation on {table_name}"
|
||||
)
|
||||
continue
|
||||
# Check whether per-bank partial vector indexes already cover this table
|
||||
# (created by the bank_utils lifecycle — no global index needed in that case)
|
||||
per_bank_index_count = conn.execute(
|
||||
text("""
|
||||
SELECT COUNT(*)
|
||||
FROM pg_indexes
|
||||
WHERE schemaname = :schema
|
||||
AND tablename = :table_name
|
||||
AND indexname LIKE 'idx_mu_emb_%'
|
||||
"""),
|
||||
{"schema": schema_name, "table_name": table_name},
|
||||
).scalar()
|
||||
if per_bank_index_count and per_bank_index_count > 0:
|
||||
logger.debug(
|
||||
f"No global embedding index on {table_name}, but {per_bank_index_count} "
|
||||
f"per-bank partial vector indexes exist — skipping global index creation"
|
||||
)
|
||||
continue
|
||||
logger.warning(f"No embedding index found for {table_name}, will create it if safe")
|
||||
mismatched_tables.append((table_name, index_name, None, row_count))
|
||||
continue
|
||||
@@ -850,7 +836,7 @@ def ensure_text_search_extension(
|
||||
schema_name = schema or "public"
|
||||
pg_search_tokenizer = normalize_pg_search_tokenizer(pg_search_tokenizer)
|
||||
|
||||
engine = create_engine(to_libpq_url(database_url), poolclass=NullPool)
|
||||
engine = create_engine(to_libpq_url(database_url))
|
||||
with engine.connect() as conn:
|
||||
# Tables with search_vector columns to check
|
||||
tables_to_check = [
|
||||
@@ -1143,129 +1129,3 @@ def ensure_text_search_extension(
|
||||
|
||||
conn.commit()
|
||||
logger.info(f"Successfully migrated text search to {text_search_extension}")
|
||||
|
||||
|
||||
def _migrate_one_schema_pg(
|
||||
database_url: str,
|
||||
schema: str,
|
||||
*,
|
||||
migration_database_url: str | None,
|
||||
embedding_dimension: int | None,
|
||||
vector_extension: str,
|
||||
text_search_extension: str,
|
||||
pg_search_tokenizer: str | None,
|
||||
ensure_extensions: bool,
|
||||
) -> str:
|
||||
"""Run migrations + post-migration extension setup for a SINGLE PG schema.
|
||||
|
||||
Module-level (not a closure) so it is picklable and can run inside a
|
||||
``ProcessPoolExecutor`` worker. The steps run strictly in order — this is
|
||||
the per-tenant sequential unit; parallelism happens only *across* schemas.
|
||||
Returns the schema name on success; raises on the first failing step so the
|
||||
caller can attribute the failure back to this schema.
|
||||
"""
|
||||
run_migrations(database_url, schema=schema, migration_database_url=migration_database_url)
|
||||
if embedding_dimension is not None:
|
||||
ensure_embedding_dimension(
|
||||
database_url,
|
||||
embedding_dimension,
|
||||
schema=schema,
|
||||
vector_extension=vector_extension,
|
||||
)
|
||||
if ensure_extensions:
|
||||
ensure_vector_extension(database_url, vector_extension=vector_extension, schema=schema)
|
||||
ensure_text_search_extension(
|
||||
database_url,
|
||||
text_search_extension=text_search_extension,
|
||||
schema=schema,
|
||||
pg_search_tokenizer=pg_search_tokenizer,
|
||||
)
|
||||
return schema
|
||||
|
||||
|
||||
def _make_migration_executor(max_workers: int):
|
||||
"""Build the executor that runs per-schema migrations in parallel.
|
||||
|
||||
Each schema must run in its OWN process — Alembic's ``command.upgrade()``
|
||||
uses non-thread-safe module globals (serialized in-process by
|
||||
``_alembic_lock``), so a thread pool would not actually run two upgrades at
|
||||
once. ``spawn`` gives every worker a clean interpreter on all platforms,
|
||||
avoiding the fork-of-a-multithreaded-process deadlock hazard (the API server
|
||||
holds threads/pools when migrations run on startup).
|
||||
|
||||
Factored out so tests can substitute an in-process executor.
|
||||
"""
|
||||
import multiprocessing
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
|
||||
return ProcessPoolExecutor(max_workers=max_workers, mp_context=multiprocessing.get_context("spawn"))
|
||||
|
||||
|
||||
def run_migrations_for_schemas(
|
||||
database_url: str,
|
||||
schemas: list[str],
|
||||
*,
|
||||
concurrency: int = 1,
|
||||
migration_database_url: str | None = None,
|
||||
embedding_dimension: int | None = None,
|
||||
vector_extension: str = "pgvector",
|
||||
text_search_extension: str = "native",
|
||||
pg_search_tokenizer: str | None = None,
|
||||
ensure_extensions: bool = True,
|
||||
) -> None:
|
||||
"""Run PostgreSQL migrations for many schemas, up to ``concurrency`` at once.
|
||||
|
||||
Within a schema the work is always sequential (migrate → embedding dim →
|
||||
vector ext → text-search ext). Across schemas, when ``concurrency > 1`` each
|
||||
schema is migrated in its OWN process: Alembic's ``command.upgrade()`` relies
|
||||
on non-thread-safe module-level globals (serialized in-process by
|
||||
``_alembic_lock``), so threads would gain nothing — separate interpreters
|
||||
each get a clean Alembic context. Per-schema advisory locks
|
||||
(``_get_schema_lock_id``) keep concurrent processes from colliding on the
|
||||
same schema across replicas.
|
||||
|
||||
``database_url`` must already be resolved (e.g. an embedded ``pg0`` instance
|
||||
started in the parent) — workers receive it verbatim and only connect.
|
||||
|
||||
Failures are collected per schema and re-raised together so one bad tenant
|
||||
does not hide the status of the others.
|
||||
"""
|
||||
if not schemas:
|
||||
return
|
||||
|
||||
worker_kwargs = dict(
|
||||
migration_database_url=migration_database_url,
|
||||
embedding_dimension=embedding_dimension,
|
||||
vector_extension=vector_extension,
|
||||
text_search_extension=text_search_extension,
|
||||
pg_search_tokenizer=pg_search_tokenizer,
|
||||
ensure_extensions=ensure_extensions,
|
||||
)
|
||||
|
||||
effective = max(1, min(concurrency, len(schemas)))
|
||||
if effective == 1:
|
||||
# Inline, in-process — no subprocess overhead for the common single
|
||||
# tenant / sequential case (and keeps embedded pg0 dev simple).
|
||||
for schema in schemas:
|
||||
_migrate_one_schema_pg(database_url, schema, **worker_kwargs)
|
||||
return
|
||||
|
||||
logger.info("Migrating %d schema(s) with concurrency=%d", len(schemas), effective)
|
||||
errors: dict[str, BaseException] = {}
|
||||
with _make_migration_executor(effective) as executor:
|
||||
futures = {
|
||||
executor.submit(_migrate_one_schema_pg, database_url, schema, **worker_kwargs): schema for schema in schemas
|
||||
}
|
||||
for future in futures:
|
||||
schema = futures[future]
|
||||
try:
|
||||
future.result()
|
||||
except Exception as exc: # noqa: BLE001 — aggregate per-schema, re-raise below
|
||||
errors[schema] = exc
|
||||
logger.error("Migration failed for schema '%s': %s", schema, exc)
|
||||
|
||||
if errors:
|
||||
failed = ", ".join(sorted(errors))
|
||||
raise RuntimeError(
|
||||
f"Database migrations failed for {len(errors)} of {len(schemas)} schema(s): {failed}"
|
||||
) from next(iter(errors.values()))
|
||||
|
||||
@@ -1,58 +1,6 @@
|
||||
import logging
|
||||
import os
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
|
||||
def detect_container_runtime() -> str | None:
|
||||
"""Detect whether the process is running inside a container.
|
||||
|
||||
Returns "kubernetes", "docker", or None. Used to warn operators that the
|
||||
default ``socket.gethostname()`` worker id is unstable across container
|
||||
recreation (the random container id changes on restart, so tasks stuck in
|
||||
'processing' under the old id are never recovered).
|
||||
"""
|
||||
if os.getenv("KUBERNETES_SERVICE_HOST"):
|
||||
return "kubernetes"
|
||||
# Docker (and most OCI runtimes) create this marker file in every container.
|
||||
if os.path.exists("/.dockerenv"):
|
||||
return "docker"
|
||||
# cgroup v1 fallback for runtimes that don't write /.dockerenv.
|
||||
try:
|
||||
with open("/proc/1/cgroup", encoding="utf-8") as f:
|
||||
if any(token in f.read() for token in ("docker", "containerd", "kubepods")):
|
||||
return "docker"
|
||||
except OSError:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def warn_if_container_default_worker_id(worker_id: str | None) -> None:
|
||||
"""Warn when worker id will fall back to an unstable container hostname."""
|
||||
if worker_id:
|
||||
return
|
||||
|
||||
runtime = detect_container_runtime()
|
||||
if not runtime:
|
||||
return
|
||||
|
||||
logging.warning(
|
||||
"\n"
|
||||
"============================================================\n"
|
||||
" WARNING: HINDSIGHT_API_WORKER_ID is not set and Hindsight\n"
|
||||
f" appears to be running inside {runtime}.\n"
|
||||
"\n"
|
||||
" The worker id is defaulting to the container hostname,\n"
|
||||
" which CHANGES every time the container is recreated.\n"
|
||||
" When that happens, tasks left in 'processing' under the\n"
|
||||
" old hostname are never recovered — consolidation and other\n"
|
||||
" async operations can get stuck indefinitely.\n"
|
||||
"\n"
|
||||
" Set HINDSIGHT_API_WORKER_ID to a STABLE value (e.g. the\n"
|
||||
" compose service name or StatefulSet pod name) to avoid this.\n"
|
||||
"============================================================"
|
||||
)
|
||||
|
||||
|
||||
def mask_network_location(url):
|
||||
if not url:
|
||||
return url
|
||||
|
||||
@@ -4,7 +4,6 @@ from .manager import WebhookManager
|
||||
from .models import (
|
||||
ConsolidationEventData,
|
||||
MemoryDefenseEventData,
|
||||
MemoryDefenseHit,
|
||||
RetainEventData,
|
||||
WebhookConfig,
|
||||
WebhookEvent,
|
||||
@@ -18,6 +17,5 @@ __all__ = [
|
||||
"WebhookEventType",
|
||||
"ConsolidationEventData",
|
||||
"MemoryDefenseEventData",
|
||||
"MemoryDefenseHit",
|
||||
"RetainEventData",
|
||||
]
|
||||
|
||||
@@ -70,10 +70,7 @@ class WebhookManager:
|
||||
webhook_table = _fq_table("webhooks", schema)
|
||||
ops_table = _fq_table("async_operations", schema)
|
||||
now = datetime.now(timezone.utc)
|
||||
# Drop null fields so receivers don't see promised-but-unfilled keys.
|
||||
# OSS leaves SIEM-enrichment fields (severity, api_key_name, etc.) None
|
||||
# because it doesn't have the data; cloud populates them when it does.
|
||||
payload_str = event.model_dump_json(exclude_none=True)
|
||||
payload_str = event.model_dump_json()
|
||||
|
||||
try:
|
||||
async with self._backend.acquire() as conn:
|
||||
@@ -153,10 +150,7 @@ class WebhookManager:
|
||||
webhook_table = _fq_table("webhooks", schema)
|
||||
ops_table = _fq_table("async_operations", schema)
|
||||
now = datetime.now(timezone.utc)
|
||||
# Drop null fields so receivers don't see promised-but-unfilled keys.
|
||||
# OSS leaves SIEM-enrichment fields (severity, api_key_name, etc.) None
|
||||
# because it doesn't have the data; cloud populates them when it does.
|
||||
payload_str = event.model_dump_json(exclude_none=True)
|
||||
payload_str = event.model_dump_json()
|
||||
|
||||
try:
|
||||
rows = await self._backend.ops.get_webhooks_for_dispatch(
|
||||
|
||||
@@ -24,43 +24,14 @@ class RetainEventData(BaseModel):
|
||||
tags: list[str] | None = None
|
||||
|
||||
|
||||
class MemoryDefenseHit(BaseModel):
|
||||
"""A single secret match inside a non-allow decision.
|
||||
|
||||
``preview`` is a fingerprinted, redaction-identifiable rendering of the
|
||||
matched value (e.g. ``ghp_AAAA...BBBB``) so SIEM operators can correlate
|
||||
against their credential inventory WITHOUT the raw secret crossing the
|
||||
network. Implementations must never put the raw value here.
|
||||
"""
|
||||
|
||||
detector: str # the inner detector that matched (e.g. "GitHub Token")
|
||||
preview: str # fingerprinted value, never the raw secret
|
||||
|
||||
|
||||
class MemoryDefenseEventData(BaseModel):
|
||||
"""Payload for a memory_defense.triggered event (one item, one non-allow decision).
|
||||
|
||||
The four base fields (``action``/``detector``/``document_id``/``message``)
|
||||
plus ``matched_types`` are populated by every implementation including OSS's
|
||||
built-in regex defense. The remaining fields are optional SIEM-enrichment
|
||||
surfaces that downstream extensions (e.g. hindsight-cloud) populate when
|
||||
they have richer per-decision context — severity classification, the API
|
||||
key that submitted the retain, fingerprinted hit previews for SIEM
|
||||
correlation, and pointers into the audit trail. OSS leaves them ``None``;
|
||||
receivers should treat absence as "not provided" rather than "no match".
|
||||
"""
|
||||
"""Payload for a memory_defense.triggered event (one item, one non-allow decision)."""
|
||||
|
||||
action: str # "redact" or "block"
|
||||
detector: str | None = None # e.g. "sensitive_data"
|
||||
document_id: str | None = None
|
||||
matched_types: list[str] | None = None # redaction pattern labels that fired
|
||||
message: str | None = None
|
||||
# --- Optional SIEM enrichment (populated by extensions, not OSS) ---
|
||||
severity: str | None = None # "low" / "medium" / "high" / "critical"
|
||||
api_key_name: str | None = None # human-readable name of the submitting API key
|
||||
hits: list[MemoryDefenseHit] | None = None # per-match fingerprints for correlation
|
||||
memory_unit_id: str | None = None # drill-down pointer (when the decision was REDACT)
|
||||
receipt_uri: str | None = None # storage pointer for the audit trail entry
|
||||
|
||||
|
||||
class WebhookEvent(BaseModel):
|
||||
|
||||
@@ -136,7 +136,7 @@ def main():
|
||||
# Worker options
|
||||
parser.add_argument(
|
||||
"--worker-id",
|
||||
default=config.worker_id,
|
||||
default=config.worker_id or socket.gethostname(),
|
||||
help="Worker identifier (default: hostname, env: HINDSIGHT_API_WORKER_ID)",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -178,17 +178,10 @@ def main():
|
||||
# Configure logging
|
||||
config.configure_logging()
|
||||
|
||||
from ..utils import warn_if_container_default_worker_id
|
||||
|
||||
warn_if_container_default_worker_id(args.worker_id)
|
||||
worker_id = args.worker_id or socket.gethostname()
|
||||
worker_id_source = "HINDSIGHT_API_WORKER_ID/--worker-id" if args.worker_id else "hostname (default)"
|
||||
logger.info(f"Worker id: {worker_id} (source: {worker_id_source})")
|
||||
|
||||
# Import MemoryEngine here to avoid circular imports
|
||||
from .. import MemoryEngine
|
||||
|
||||
print(f"Starting Hindsight Worker: {worker_id}")
|
||||
print(f"Starting Hindsight Worker: {args.worker_id}")
|
||||
print(f" Poll interval: {args.poll_interval}ms")
|
||||
print(f" Max retries: {args.max_retries}")
|
||||
print(f" Max slots: {config.worker_max_slots}")
|
||||
@@ -256,7 +249,7 @@ def main():
|
||||
schema = None if config.database_schema == DEFAULT_DATABASE_SCHEMA else config.database_schema
|
||||
poller = WorkerPoller(
|
||||
backend=memory._backend,
|
||||
worker_id=worker_id,
|
||||
worker_id=args.worker_id,
|
||||
executor=memory.execute_task,
|
||||
poll_interval_ms=args.poll_interval,
|
||||
schema=schema,
|
||||
|
||||
@@ -20,23 +20,9 @@ from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from ..engine.schema import fq_table_explicit as fq_table
|
||||
from ..metrics import get_metrics_collector
|
||||
from .exceptions import DeferOperation, RetryTaskAt
|
||||
from .stage import StageHolder, bind_holder
|
||||
|
||||
# Map DB operation_type -> metric `operation` label, collapsing the retain
|
||||
# variants onto "retain" so async worker completions land on the same
|
||||
# operation="retain" series the synchronous API path emits. Unknown types
|
||||
# pass through unchanged.
|
||||
_RETAIN_OP_TYPES = {"retain", "batch_retain", "file_convert_retain"}
|
||||
|
||||
|
||||
def _metric_operation_label(operation_type: str | None) -> str:
|
||||
if operation_type in _RETAIN_OP_TYPES:
|
||||
return "retain"
|
||||
return operation_type or "unknown"
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from hindsight_api.engine.db.base import DatabaseBackend, DatabaseConnection
|
||||
from hindsight_api.extensions.tenant import TenantExtension
|
||||
@@ -240,23 +226,38 @@ class WorkerPoller:
|
||||
"""
|
||||
async with self._backend.acquire() as conn:
|
||||
if await self._optional_routines.is_installed(conn, "schemas_with_pending_work"):
|
||||
# The routine IS the authority on where work exists: every schema
|
||||
# it returns is claimable, and every schema it does NOT return is
|
||||
# treated as having nothing to do this cycle. That is the entire
|
||||
# point of installing it — one round-trip replaces N per-schema
|
||||
# EXISTS probes. We deliberately do NOT re-verify the omitted
|
||||
# schemas with a per-schema scan: that re-runs the exact queries
|
||||
# the routine exists to avoid, on every idle poll, silently
|
||||
# negating the optimisation.
|
||||
#
|
||||
# Because the result is trusted wholesale, the routine is only
|
||||
# appropriate for multi-tenant deployments. A single-schema
|
||||
# (default/public only) install should NOT create it and instead
|
||||
# falls through to the per-schema path below — a single cheap
|
||||
# EXISTS check that cannot starve. See
|
||||
# ``hindsight_api.engine.db.optional_routines``.
|
||||
rows = await conn.fetch("SELECT * FROM public.schemas_with_pending_work()")
|
||||
return {self._normalize_poll_schema(r[0]) for r in rows}
|
||||
routine_active = {self._normalize_poll_schema(r[0]) for r in rows}
|
||||
known_schemas = set(schemas)
|
||||
active = routine_active & known_schemas
|
||||
unknown = routine_active - known_schemas
|
||||
if unknown:
|
||||
logger.warning(
|
||||
"Optional PG routine public.schemas_with_pending_work() returned schema(s) "
|
||||
"not present in tenant discovery: %s",
|
||||
sorted(str(s) for s in unknown),
|
||||
)
|
||||
|
||||
# The optional routine returns PostgreSQL schema names, but the poller uses
|
||||
# None for the default schema. Older operator-supplied implementations also
|
||||
# commonly scan tenant_% only; when the default schema is in scope but absent
|
||||
# from the routine result, verify via the fully-correct per-schema fallback so
|
||||
# public single-tenant deployments cannot silently starve.
|
||||
should_verify_with_fallback = (None in known_schemas and None not in active) or (
|
||||
bool(routine_active) and not active
|
||||
)
|
||||
if not should_verify_with_fallback:
|
||||
return active
|
||||
|
||||
fallback_active = await self._scan_active_schemas_by_exists(conn, schemas)
|
||||
missed = fallback_active - active
|
||||
if missed:
|
||||
logger.warning(
|
||||
"Optional PG routine public.schemas_with_pending_work() missed claimable schema(s) %s; "
|
||||
"using per-schema fallback for this poll",
|
||||
sorted(str(s) for s in missed),
|
||||
)
|
||||
return fallback_active
|
||||
|
||||
return await self._scan_active_schemas_by_exists(conn, schemas)
|
||||
|
||||
@@ -715,24 +716,6 @@ class WorkerPoller:
|
||||
"""
|
||||
task_type = task.task_dict.get("type", "unknown")
|
||||
bank_id = task.task_dict.get("bank_id", "unknown")
|
||||
# Operation metric (source="worker"): record on terminal outcomes only, so
|
||||
# async worker throughput and latency (retain, consolidation and the other
|
||||
# worker task types) are visible in Prometheus. Prefer the DB-authoritative
|
||||
# operation_type.
|
||||
#
|
||||
# success semantics are deliberately narrow: success=false means the task
|
||||
# raised out to the poller (an unexpected error, or retry-exhausted). It does
|
||||
# NOT capture deterministic failures that the executor handles itself and
|
||||
# returns from normally (file_convert_retain, non-retryable errors via
|
||||
# memory_engine.execute_task) — those record success=true here. Treat this as
|
||||
# a completion-throughput signal, not a failure-rate one: for authoritative
|
||||
# failure visibility use the hindsight_async_operations{status="failed"} gauge,
|
||||
# which reads each operation's final DB status.
|
||||
op_label = _metric_operation_label(task.task_dict.get("operation_type") or task_type)
|
||||
op_start = time.time()
|
||||
metrics = get_metrics_collector()
|
||||
# None = not a terminal outcome (deferred/retried) → no metric.
|
||||
terminal_success: bool | None = None
|
||||
|
||||
# Bind the stage holder in this task's own contextvar scope so engine
|
||||
# code running under us can update it via stage.set_stage(). If holder
|
||||
@@ -749,28 +732,14 @@ class WorkerPoller:
|
||||
task.task_dict["_schema"] = task.schema
|
||||
await self._executor(task.task_dict)
|
||||
logger.debug(f"Task {task.operation_id} execution finished")
|
||||
terminal_success = True
|
||||
except DeferOperation as e:
|
||||
# Deferral is not a terminal outcome — do not record a completion.
|
||||
await self._defer_operation(task.operation_id, e.exec_date, e.reason, task.schema)
|
||||
except RetryTaskAt as e:
|
||||
# Retry is not a terminal outcome — do not record a completion.
|
||||
await self._schedule_retry(task.operation_id, e.retry_at, str(e), task.schema)
|
||||
except Exception as e:
|
||||
logger.error(f"Task {task.operation_id} failed: {e}")
|
||||
traceback.print_exc()
|
||||
await self._mark_failed(task.operation_id, str(e), task.schema)
|
||||
terminal_success = False
|
||||
|
||||
# Record the metric outside the executor's exception scope so a metrics
|
||||
# reporting failure can never be mistaken for a task failure and flip terminal state.
|
||||
if terminal_success is not None:
|
||||
try:
|
||||
metrics.record_operation_result(
|
||||
op_label, bank_id, success=terminal_success, duration=time.time() - op_start, source="worker"
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(f"Failed to record worker operation metric for {task.operation_id}", exc_info=True)
|
||||
|
||||
async def recover_own_tasks(self) -> int:
|
||||
"""
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "hindsight-api-slim"
|
||||
version = "0.8.4"
|
||||
version = "0.8.1"
|
||||
description = "Hindsight: Agent Memory That Works Like Human Memory"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
@@ -60,10 +60,10 @@ dependencies = [
|
||||
"pyasn1>=0.6.3", # DoS vulnerability fix
|
||||
"urllib3>=2.7.0", # Decompression-bomb safeguards bypass + sensitive header forwarding fixes
|
||||
"langchain-core>=1.2.22", # Path traversal in legacy load_prompt functions fix
|
||||
"langsmith>=0.8.18", # GHSA-f4xh-w4cj-qxq8: arbitrary server-side file read in TracingMiddleware fix (supersedes >=0.6.3 SSRF tracing-header-injection floor)
|
||||
"langsmith>=0.6.3", # SSRF via tracing header injection fix
|
||||
"protobuf>=6.33.5", # JSON recursion depth bypass fix
|
||||
"pillow>=12.1.1", # Out-of-bounds write in PSD image loading fix
|
||||
"cryptography>=48.0.1", # GHSA-537c-gmf6-5ccf: bundled-OpenSSL OOB read fix needs >=48.0.1. Prior <47 cap (47.0.0 SIGILL on ARM64 Docker/Podman, pyca/cryptography#14733) lifted — 47/48/49 verified importing + RSA sign/verify cleanly on linux/arm64 (Docker on Apple Silicon) and native arm64 macOS; upstream issue closed unconfirmed.
|
||||
"cryptography>=46.0.6,<47", # Incomplete DNS name constraint enforcement fix; cap <47 — 47.0.0 SIGILLs on some ARM64 Linux VMs (Docker/Podman on Apple Silicon), pyca/cryptography#14733
|
||||
"filelock>=3.20.1", # TOCTOU race condition fix
|
||||
"authlib>=1.6.9", # Account takeover/JWS header injection vulnerability fix
|
||||
"pyjwt>=2.12.0", # Accepts unknown crit header extensions fix
|
||||
@@ -74,7 +74,6 @@ dependencies = [
|
||||
"pygments>=2.20.0", # ReDoS via inefficient GUID regex fix
|
||||
"claude-agent-sdk>=0.2.82",
|
||||
"boto3>=1.42.74",
|
||||
"croniter>=2.0.0", # Cron parsing for scheduled mental model refresh
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
@@ -201,11 +200,12 @@ select = [
|
||||
"W", # pycodestyle warnings
|
||||
"F", # Pyflakes
|
||||
"I", # isort
|
||||
"B021", # flake8-bugbear: f-string used as docstring (leaves __doc__ None)
|
||||
]
|
||||
ignore = [
|
||||
"E501", # line too long (handled by formatter)
|
||||
"E402", # module import not at top of file
|
||||
"F401", # unused import (too noisy during development)
|
||||
"F841", # unused variable (too noisy during development)
|
||||
"F811", # redefined while unused
|
||||
"F821", # undefined name (forward references in type hints)
|
||||
]
|
||||
|
||||
@@ -11,71 +11,11 @@ import pytest
|
||||
import pytest_asyncio
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# Force torch to initialize exactly once, in the main thread, at conftest import
|
||||
# time — before any fixture spins up an event loop or sentence-transformers'
|
||||
# thread pools. torch's C-level `_add_docstr(_has_torch_function, ...)` in
|
||||
# torch/overrides.py is not re-entrancy-safe: when the first `import torch`
|
||||
# happens lazily from inside concurrent/async code (e.g.
|
||||
# embeddings.initialize() -> sentence_transformers -> transformers -> torch, or
|
||||
# cross_encoder's ThreadPoolExecutor), torch/overrides.py can execute twice and
|
||||
# raise "RuntimeError: function '_has_torch_function' already has a docstring",
|
||||
# failing collection of every test on the pytest-xdist shard. Importing it here
|
||||
# (single-threaded, before any concurrency) makes that registration happen once
|
||||
# per worker process. Guarded so slim/no-torch environments still collect.
|
||||
try:
|
||||
import torch # noqa: F401 # eager one-time init; see comment above
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
from hindsight_api import LLMConfig, LocalSTEmbeddings, MemoryEngine, RequestContext
|
||||
from hindsight_api.engine.cross_encoder import LocalSTCrossEncoder
|
||||
from hindsight_api.engine.query_analyzer import DateparserQueryAnalyzer
|
||||
from hindsight_api.engine.task_backend import SyncTaskBackend
|
||||
from hindsight_api.pg0 import EmbeddedPostgres
|
||||
from hindsight_api.tracing import unregister_span_recorder
|
||||
|
||||
|
||||
async def _teardown_memory_engine(mem: MemoryEngine) -> None:
|
||||
"""Tear down a test MemoryEngine, guaranteeing its span recorder is unregistered.
|
||||
|
||||
LLM-trace recorders live in a process-global registry; ``MemoryEngine.close()`` is
|
||||
the only thing that removes the engine's recorder from it. If close() is skipped
|
||||
(pool already closing) or raises before that step, the recorder leaks and a later
|
||||
test's LLM calls get recorded into the shared DB — the flaky
|
||||
test_llm_trace::test_disabled_writes_no_rows (#2229). Unregister unconditionally;
|
||||
it's a no-op when close() already did it.
|
||||
"""
|
||||
try:
|
||||
if mem._pool and not mem._pool._closing:
|
||||
await mem.close()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
unregister_span_recorder(mem._llm_recorder)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _cleanup_leaked_span_recorders():
|
||||
"""Fail-safe for the process-global LLM-trace recorder registry (#2229).
|
||||
|
||||
``MemoryEngine.__init__`` registers its recorder in the shared registry, and
|
||||
only ``close()`` removes it. Tests that construct an engine directly (without
|
||||
``_teardown_memory_engine``/``close()``) leak an *enabled* recorder; a later
|
||||
test's LLM calls then get recorded into the shared DB, flaking
|
||||
``test_llm_trace::test_disabled_writes_no_rows`` (it observes rows for its
|
||||
bank even though its own recorder is disabled). ``_teardown_memory_engine``
|
||||
guards the fixtures; this guards everything else by dropping any recorder a
|
||||
test added to the registry.
|
||||
"""
|
||||
from hindsight_api.tracing import get_span_recorder
|
||||
|
||||
recorders = get_span_recorder()._recorders
|
||||
before = {id(r) for r in recorders}
|
||||
yield
|
||||
for recorder in list(recorders):
|
||||
if id(recorder) not in before:
|
||||
recorders.remove(recorder)
|
||||
|
||||
|
||||
# Default pg0 instance configuration for tests
|
||||
DEFAULT_PG0_INSTANCE_NAME = "hindsight-test"
|
||||
@@ -84,13 +24,11 @@ DEFAULT_PG0_PORT = int(os.environ.get("HINDSIGHT_TEST_PG_PORT", "5556"))
|
||||
# Keep the background MaintenanceLoop from auto-starting during tests. In
|
||||
# production it sweeps retention and re-schedules consolidation, but its timers
|
||||
# would race shared-pg0 test data (e.g. delete llm_requests/audit_log rows a test
|
||||
# just inserted). Disabling the reconcile interval, the mental-model refresh tick
|
||||
# and llm-trace retention — with audit retention already off by default — leaves
|
||||
# no job enabled, so the loop never starts. Tests that exercise it call
|
||||
# MaintenanceLoop methods (_run_reconcile / _run_scheduled_mm_refresh /
|
||||
# _purge_expired) directly.
|
||||
# just inserted). Disabling the reconcile interval and llm-trace retention — with
|
||||
# audit retention already off by default — leaves no job enabled, so the loop
|
||||
# never starts. Tests that exercise it call MaintenanceLoop methods
|
||||
# (_run_reconcile / _purge_expired) directly.
|
||||
os.environ.setdefault("HINDSIGHT_API_CONSOLIDATION_RECONCILE_INTERVAL_SECONDS", "0")
|
||||
os.environ.setdefault("HINDSIGHT_API_MENTAL_MODEL_REFRESH_TICK_SECONDS", "0")
|
||||
os.environ.setdefault("HINDSIGHT_API_LLM_TRACE_RETENTION_DAYS", "-1")
|
||||
|
||||
|
||||
@@ -404,7 +342,10 @@ async def oracle_memory(oracle_db_url, embeddings, cross_encoder, query_analyzer
|
||||
)
|
||||
await mem.initialize()
|
||||
yield mem
|
||||
await _teardown_memory_engine(mem)
|
||||
try:
|
||||
await mem.close()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
# Restore original env var and clear config cache
|
||||
if old_backend is None:
|
||||
@@ -522,7 +463,11 @@ async def memory(pg0_db_url, embeddings, cross_encoder, query_analyzer):
|
||||
)
|
||||
await mem.initialize()
|
||||
yield mem
|
||||
await _teardown_memory_engine(mem)
|
||||
try:
|
||||
if mem._pool and not mem._pool._closing:
|
||||
await mem.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@pytest_asyncio.fixture(scope="function")
|
||||
@@ -551,7 +496,11 @@ async def memory_real_llm(pg0_db_url, embeddings, cross_encoder, query_analyzer)
|
||||
)
|
||||
await mem.initialize()
|
||||
yield mem
|
||||
await _teardown_memory_engine(mem)
|
||||
try:
|
||||
if mem._pool and not mem._pool._closing:
|
||||
await mem.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@pytest_asyncio.fixture(scope="function")
|
||||
@@ -578,7 +527,11 @@ async def memory_no_llm_verify(pg0_db_url, embeddings, cross_encoder, query_anal
|
||||
)
|
||||
await mem.initialize()
|
||||
yield mem
|
||||
await _teardown_memory_engine(mem)
|
||||
try:
|
||||
if mem._pool and not mem._pool._closing:
|
||||
await mem.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
|
||||
@@ -88,10 +88,7 @@ async def test_backup_tables_covers_entire_schema(backup_test_schema):
|
||||
await conn.close()
|
||||
|
||||
# alembic_version is migration bookkeeping, not data — never backed up.
|
||||
# bank_stats_cache is a derived TTL cache of get_bank_stats results: it has no
|
||||
# FK to banks (so the restore cascade never touches it) and repopulates itself
|
||||
# on demand, so it is deliberately not backed up — a restore starts it cold.
|
||||
schema_tables = {r["table_name"] for r in rows} - {"alembic_version", "bank_stats_cache"}
|
||||
schema_tables = {r["table_name"] for r in rows} - {"alembic_version"}
|
||||
backup_tables = set(BACKUP_TABLES)
|
||||
|
||||
missing = schema_tables - backup_tables
|
||||
@@ -550,36 +547,3 @@ async def test_run_migration_with_schema_only_runs_requested_schema(monkeypatch)
|
||||
assert calls["run_migrations"] == [("resolved::postgresql://test", "tenant_demo")]
|
||||
assert calls["ensure_vector_extension"] == [("resolved::postgresql://test", "pgvector", "tenant_demo")]
|
||||
assert calls["ensure_text_search_extension"] == [("resolved::postgresql://test", "native", "", "tenant_demo")]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("ensure_extensions", "expected"),
|
||||
[(True, True), (False, False)],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_migration_threads_ensure_extensions_flag(monkeypatch, ensure_extensions, expected):
|
||||
"""The --skip-extension-reconcile flag (ensure_extensions=False) must reach run_migrations_for_schemas.
|
||||
|
||||
The post-migration vector/text-search reconcile only does work on a backend change, so operators
|
||||
can skip it on a no-change re-migration over many tenant schemas. Verify the flag is threaded through
|
||||
rather than silently dropped.
|
||||
"""
|
||||
monkeypatch.setenv("HINDSIGHT_API_DATABASE_URL", "postgresql://test")
|
||||
captured: dict = {}
|
||||
|
||||
async def fake_resolve_database_url(db_url: str) -> str:
|
||||
return f"resolved::{db_url}"
|
||||
|
||||
def fake_run_migrations_for_schemas(database_url, schemas, **kwargs):
|
||||
captured["ensure_extensions"] = kwargs.get("ensure_extensions")
|
||||
|
||||
monkeypatch.setattr(admin_cli, "load_extension", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(admin_cli, "resolve_database_url", fake_resolve_database_url)
|
||||
|
||||
from hindsight_api import migrations as migrations_module
|
||||
|
||||
monkeypatch.setattr(migrations_module, "run_migrations_for_schemas", fake_run_migrations_for_schemas)
|
||||
|
||||
await admin_cli._run_migration("postgresql://test", schema="tenant_demo", ensure_extensions=ensure_extensions)
|
||||
|
||||
assert captured["ensure_extensions"] is expected
|
||||
|
||||
@@ -1,111 +0,0 @@
|
||||
"""Regression tests for issue #1002 — Anthropic structured output via forced tool_use.
|
||||
|
||||
When strict_schema=True, AnthropicLLM.call() must request the schema through a single
|
||||
forced tool_use tool (tool_choice={"type":"tool",...}) and read the validated args from
|
||||
the tool_use block, NOT inject the schema as text and json.loads() the reply (which caused
|
||||
a ~1:1 invalid-JSON retry storm / OOM in production).
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class _Decision(BaseModel):
|
||||
action: str
|
||||
reason: str
|
||||
|
||||
|
||||
def _make_anthropic_provider():
|
||||
with patch("anthropic.AsyncAnthropic") as mock_client_cls:
|
||||
mock_client_cls.return_value = MagicMock()
|
||||
from hindsight_api.engine.providers.anthropic_llm import AnthropicLLM
|
||||
|
||||
provider = AnthropicLLM(
|
||||
provider="anthropic",
|
||||
api_key="fake-key",
|
||||
base_url="",
|
||||
model="claude-sonnet-4-20250514",
|
||||
)
|
||||
provider._client = MagicMock()
|
||||
return provider
|
||||
|
||||
|
||||
def _tool_use_response(args: dict):
|
||||
block = MagicMock()
|
||||
block.type = "tool_use"
|
||||
block.name = "structured_response"
|
||||
block.input = args
|
||||
resp = MagicMock()
|
||||
resp.content = [block]
|
||||
resp.usage = MagicMock(input_tokens=5, output_tokens=2, cache_read_input_tokens=0)
|
||||
resp.stop_reason = "tool_use"
|
||||
return resp
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strict_schema_uses_forced_tool_choice():
|
||||
"""strict_schema=True ⇒ a single tool is defined and tool_choice forces it (no schema text-injection)."""
|
||||
provider = _make_anthropic_provider()
|
||||
provider._client.messages.create = AsyncMock(return_value=_tool_use_response({"action": "skip", "reason": "dup"}))
|
||||
with patch("hindsight_api.engine.providers.anthropic_llm.get_metrics_collector"):
|
||||
result = await provider.call(
|
||||
messages=[{"role": "user", "content": "decide"}],
|
||||
response_format=_Decision,
|
||||
strict_schema=True,
|
||||
scope="test",
|
||||
max_retries=0,
|
||||
)
|
||||
kwargs = provider._client.messages.create.call_args.kwargs
|
||||
# forced tool_use requested
|
||||
assert "tools" in kwargs and len(kwargs["tools"]) == 1
|
||||
assert kwargs["tool_choice"] == {"type": "tool", "name": "structured_response"}
|
||||
# schema NOT injected as text into the system prompt
|
||||
assert "valid JSON matching this schema" not in (kwargs.get("system") or "")
|
||||
# validated model returned straight from tool_use.input
|
||||
assert isinstance(result, _Decision)
|
||||
assert result.action == "skip"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strict_schema_tool_use_never_hits_json_retry_loop():
|
||||
"""A tool_use response is structurally valid → no second messages.create call (no retry storm)."""
|
||||
provider = _make_anthropic_provider()
|
||||
create = AsyncMock(return_value=_tool_use_response({"action": "keep", "reason": "novel"}))
|
||||
provider._client.messages.create = create
|
||||
with patch("hindsight_api.engine.providers.anthropic_llm.get_metrics_collector"):
|
||||
await provider.call(
|
||||
messages=[{"role": "user", "content": "x"}],
|
||||
response_format=_Decision,
|
||||
strict_schema=True,
|
||||
scope="test",
|
||||
max_retries=10, # would allow 11 attempts on the old text-parse path
|
||||
)
|
||||
assert create.await_count == 1 # exactly one call — the bug was N retries on malformed text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_strict_keeps_text_injection_fallback():
|
||||
"""strict_schema=False (default) preserves the legacy schema-in-prompt behavior."""
|
||||
provider = _make_anthropic_provider()
|
||||
block = MagicMock()
|
||||
block.type = "text"
|
||||
block.text = '{"action":"skip","reason":"d"}'
|
||||
resp = MagicMock()
|
||||
resp.content = [block]
|
||||
resp.usage = MagicMock(input_tokens=5, output_tokens=2, cache_read_input_tokens=0)
|
||||
resp.stop_reason = "end_turn"
|
||||
provider._client.messages.create = AsyncMock(return_value=resp)
|
||||
with patch("hindsight_api.engine.providers.anthropic_llm.get_metrics_collector"):
|
||||
result = await provider.call(
|
||||
messages=[{"role": "user", "content": "decide"}],
|
||||
response_format=_Decision,
|
||||
strict_schema=False,
|
||||
scope="test",
|
||||
max_retries=0,
|
||||
)
|
||||
kwargs = provider._client.messages.create.call_args.kwargs
|
||||
assert "tools" not in kwargs # no forced tool when not strict
|
||||
assert "valid JSON matching this schema" in (kwargs.get("system") or "")
|
||||
assert isinstance(result, _Decision)
|
||||
@@ -1,90 +0,0 @@
|
||||
"""Regression test: submitting an async op for a bank that doesn't exist must
|
||||
raise a clean validation error, not a raw asyncpg `ForeignKeyViolationError`.
|
||||
|
||||
`_submit_async_operation` inserts into `async_operations`, which has an FK to
|
||||
`banks.bank_id`. If a caller submits for a missing bank (typo, race against a
|
||||
deletion, integration that derives bank IDs before the bank is created), the
|
||||
INSERT raises `asyncpg.exceptions.ForeignKeyViolationError`. The FastAPI
|
||||
endpoint's broad `except Exception` then surfaces it as a 500 — but this is
|
||||
a client error, not a server error, and should be a 404.
|
||||
|
||||
This test exercises the call directly via `MemoryEngine.submit_async_*` so
|
||||
the failure mode is observable without spinning up the HTTP layer.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
|
||||
from hindsight_api.extensions.operation_validator import OperationValidationError
|
||||
|
||||
pytestmark = pytest.mark.xdist_group("async_submit_bank_not_found_tests")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def no_inline_execution(memory):
|
||||
"""Prevent SyncTaskBackend from running the submitted op inline so we
|
||||
only test the submit-path failure, not downstream execution."""
|
||||
|
||||
async def _noop(_payload):
|
||||
return None
|
||||
|
||||
original = memory._task_backend.submit_task
|
||||
memory._task_backend.submit_task = _noop
|
||||
yield
|
||||
memory._task_backend.submit_task = original
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consolidation_submit_on_missing_bank_raises_validation_error(
|
||||
memory, request_context, no_inline_execution
|
||||
):
|
||||
"""A `/consolidate` submit against a bank that doesn't exist must raise
|
||||
OperationValidationError(404), not a raw asyncpg FK violation that bubbles
|
||||
out as a 500 from the API."""
|
||||
missing_bank = f"does-not-exist-{uuid.uuid4().hex[:8]}"
|
||||
|
||||
with pytest.raises(OperationValidationError) as exc_info:
|
||||
await memory.submit_async_consolidation(
|
||||
bank_id=missing_bank,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert missing_bank in exc_info.value.reason
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scoped_consolidation_submit_on_missing_bank_raises_validation_error(
|
||||
memory, request_context, no_inline_execution
|
||||
):
|
||||
"""Scoped consolidates (with `observation_scopes`) take the
|
||||
`dedupe_by_bank=False` branch, which historically skipped the bank lock
|
||||
entirely and went straight to the FK-violating INSERT. Same 404 contract."""
|
||||
missing_bank = f"does-not-exist-{uuid.uuid4().hex[:8]}"
|
||||
|
||||
with pytest.raises(OperationValidationError) as exc_info:
|
||||
await memory.submit_async_consolidation(
|
||||
bank_id=missing_bank,
|
||||
request_context=request_context,
|
||||
observation_scopes=[{"tag": "anything"}],
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert missing_bank in exc_info.value.reason
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_graph_maintenance_on_missing_bank_short_circuits(memory, request_context, no_inline_execution):
|
||||
"""`submit_async_graph_maintenance` has its own short-circuit that checks
|
||||
the per-bank queue before calling `_submit_async_operation`. A missing
|
||||
bank means an empty queue, so it returns `no_work=True` without reaching
|
||||
the FK-violating INSERT. This test pins that behaviour."""
|
||||
missing_bank = f"does-not-exist-{uuid.uuid4().hex[:8]}"
|
||||
|
||||
result = await memory.submit_async_graph_maintenance(
|
||||
bank_id=missing_bank,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
assert result == {"operation_id": None, "no_work": True}
|
||||
@@ -1,244 +0,0 @@
|
||||
"""
|
||||
Tests for the async-operation queue and consolidation backlog gauges
|
||||
(``_setup_backlog_metrics`` / ``_refresh_backlog`` in metrics.py).
|
||||
|
||||
These gauges expose, as scrapable time-series, the same counts the bank-stats
|
||||
endpoint already returns per bank (``operations_by_status``,
|
||||
``pending_consolidation``, ``failed_consolidation``):
|
||||
|
||||
- ``hindsight_async_operations{operation_type,status}`` — worker queue depth
|
||||
(pending=backlog, processing=in-flight, failed=stranded)
|
||||
- ``hindsight_consolidation_backlog`` — source memories not yet consolidated
|
||||
- ``hindsight_consolidation_failed`` — source memories permanently failed
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from hindsight_api.metrics import MetricsCollector, _AsyncOpKey, _BacklogKey
|
||||
|
||||
|
||||
class _FakeTxn:
|
||||
async def __aenter__(self):
|
||||
return None
|
||||
|
||||
async def __aexit__(self, *exc):
|
||||
return False
|
||||
|
||||
|
||||
class _FakeConn:
|
||||
"""asyncpg-like connection whose fetch() is dispatched by SQL substring."""
|
||||
|
||||
def __init__(self, fetch_fn):
|
||||
self._fetch_fn = fetch_fn
|
||||
self.executed = []
|
||||
|
||||
async def fetch(self, sql, *args):
|
||||
return self._fetch_fn(sql, *args)
|
||||
|
||||
async def execute(self, sql, *args):
|
||||
self.executed.append(sql)
|
||||
|
||||
def transaction(self):
|
||||
return _FakeTxn()
|
||||
|
||||
|
||||
class _FakeAcquire:
|
||||
def __init__(self, conn):
|
||||
self._conn = conn
|
||||
|
||||
async def __aenter__(self):
|
||||
return self._conn
|
||||
|
||||
async def __aexit__(self, *exc):
|
||||
return False
|
||||
|
||||
|
||||
class _FakePool:
|
||||
def __init__(self, fetch_fn):
|
||||
self._conn = _FakeConn(fetch_fn)
|
||||
|
||||
def acquire(self):
|
||||
return _FakeAcquire(self._conn)
|
||||
|
||||
|
||||
def _collector(include_bank_id=False):
|
||||
mock_config = MagicMock()
|
||||
mock_config.metrics_include_bank_id = include_bank_id
|
||||
with (
|
||||
patch("hindsight_api.metrics.get_meter", return_value=MagicMock()),
|
||||
patch("hindsight_api.config.get_config", return_value=mock_config),
|
||||
):
|
||||
return MetricsCollector()
|
||||
|
||||
|
||||
def _set_db_pool_with_backlog_enabled(collector, pool):
|
||||
"""Call set_db_pool with the backlog flag forced on (it's off by default)."""
|
||||
mock_config = MagicMock()
|
||||
mock_config.metrics_backlog_enabled = True
|
||||
with patch("hindsight_api.config.get_config", return_value=mock_config):
|
||||
collector.set_db_pool(pool)
|
||||
|
||||
|
||||
def _rows_for(sql):
|
||||
"""Canned results, keyed off distinctive substrings of each query."""
|
||||
if "information_schema.tables" in sql:
|
||||
return [{"table_schema": "public"}]
|
||||
if "async_operations" in sql:
|
||||
return [
|
||||
{"operation_type": "retain", "status": "pending", "count": 5},
|
||||
{"operation_type": "consolidation", "status": "pending", "count": 12},
|
||||
{"operation_type": "consolidation", "status": "processing", "count": 1},
|
||||
{"operation_type": "consolidation", "status": "failed", "count": 2},
|
||||
]
|
||||
if "memory_units" in sql and "consolidated_at IS NULL" in sql:
|
||||
return [{"count": 42}]
|
||||
if "memory_units" in sql and "consolidation_failed_at IS NOT NULL" in sql:
|
||||
return [{"count": 3}]
|
||||
return []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_backlog_aggregates_queue_and_consolidation():
|
||||
collector = _collector(include_bank_id=False)
|
||||
collector._db_pool = _FakePool(lambda sql, *a: _rows_for(sql))
|
||||
|
||||
await collector._refresh_backlog()
|
||||
|
||||
# Worker queue depth keyed by (schema, operation_type, status, bank=None)
|
||||
assert collector._async_ops_counts[("public", "retain", "pending", None)] == 5
|
||||
assert collector._async_ops_counts[("public", "consolidation", "pending", None)] == 12
|
||||
assert collector._async_ops_counts[("public", "consolidation", "processing", None)] == 1
|
||||
assert collector._async_ops_counts[("public", "consolidation", "failed", None)] == 2
|
||||
# Consolidation backlog (source memories), keyed by (schema, bank=None)
|
||||
assert collector._consolidation_backlog[("public", None)] == 42
|
||||
assert collector._consolidation_failed[("public", None)] == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_backlog_uses_index_matched_predicates_not_filter_scan():
|
||||
"""Backlog/failed must be two separate COUNT(*) queries whose WHERE matches
|
||||
a partial-index predicate exactly (no FILTER over a full-table scan), and
|
||||
the queue query must exclude terminal statuses."""
|
||||
captured = []
|
||||
collector = _collector()
|
||||
collector._db_pool = _FakePool(lambda sql, *a: (captured.append(sql), _rows_for(sql))[1])
|
||||
await collector._refresh_backlog()
|
||||
|
||||
mem_queries = [s for s in captured if "memory_units" in s and "COUNT(*)" in s]
|
||||
assert len(mem_queries) == 2 # split, not a single two-FILTER aggregate
|
||||
assert all("FILTER" not in s for s in mem_queries)
|
||||
assert any("consolidated_at IS NULL AND fact_type IN ('experience', 'world')" in s for s in mem_queries)
|
||||
assert any("consolidation_failed_at IS NOT NULL AND fact_type IN ('experience', 'world')" in s for s in mem_queries)
|
||||
|
||||
ops_sql = next(s for s in captured if "async_operations" in s and "GROUP BY" in s)
|
||||
assert "status IN ('pending', 'processing', 'failed')" in ops_sql
|
||||
assert "completed" not in ops_sql and "cancelled" not in ops_sql
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_backlog_count_runs_with_seqscan_disabled():
|
||||
"""`consolidated_at IS NULL` is true for a large fraction of the table, so
|
||||
the planner misjudges selectivity and won't use the partial index without a
|
||||
nudge — the backlog count must issue SET LOCAL enable_seqscan=off."""
|
||||
collector = _collector()
|
||||
pool = _FakePool(lambda sql, *a: _rows_for(sql))
|
||||
collector._db_pool = pool
|
||||
await collector._refresh_backlog()
|
||||
|
||||
assert any("enable_seqscan" in s.lower() and "off" in s.lower() for s in pool._conn.executed)
|
||||
# the result is still correct under the nudge
|
||||
assert collector._consolidation_backlog[("public", None)] == 42
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_backlog_per_bank_labels_and_group_by_when_enabled():
|
||||
"""With metrics_include_bank_id on, bank_id enters the cache key and the
|
||||
SQL switches to GROUP BY bank_id."""
|
||||
captured = []
|
||||
|
||||
def fetch(sql, *a):
|
||||
captured.append(sql)
|
||||
if "information_schema.tables" in sql:
|
||||
return [{"table_schema": "public"}]
|
||||
if "async_operations" in sql:
|
||||
return [{"operation_type": "retain", "status": "pending", "bank_id": "bankA", "count": 4}]
|
||||
if "memory_units" in sql and "consolidated_at IS NULL" in sql:
|
||||
return [{"bank_id": "bankA", "count": 11}]
|
||||
if "memory_units" in sql and "consolidation_failed_at IS NOT NULL" in sql:
|
||||
return [{"bank_id": "bankA", "count": 2}]
|
||||
return []
|
||||
|
||||
collector = _collector(include_bank_id=True)
|
||||
collector._db_pool = _FakePool(fetch)
|
||||
await collector._refresh_backlog()
|
||||
|
||||
assert collector._async_ops_counts[("public", "retain", "pending", "bankA")] == 4
|
||||
assert collector._consolidation_backlog[("public", "bankA")] == 11
|
||||
assert collector._consolidation_failed[("public", "bankA")] == 2
|
||||
# bank_id must be grouped in every per-bank count query
|
||||
assert all("GROUP BY bank_id" in s for s in captured if "memory_units" in s and "COUNT(*)" in s)
|
||||
|
||||
|
||||
def test_gauges_register_and_emit_cached_values_without_bank_id():
|
||||
collector = _collector(include_bank_id=False)
|
||||
# Sync call: no running loop, so gauges register but no background task spawns.
|
||||
_set_db_pool_with_backlog_enabled(collector, MagicMock())
|
||||
|
||||
gauges = {
|
||||
c.kwargs["name"]: c.kwargs["callbacks"][0]
|
||||
for c in collector.meter.create_observable_gauge.call_args_list
|
||||
if "callbacks" in c.kwargs
|
||||
}
|
||||
assert "hindsight.async_operations" in gauges
|
||||
assert "hindsight.consolidation.backlog" in gauges
|
||||
assert "hindsight.consolidation.failed" in gauges
|
||||
|
||||
collector._async_ops_counts = {
|
||||
_AsyncOpKey("public", "retain", "pending", None): 7,
|
||||
_AsyncOpKey("public", "consolidation", "processing", None): 1,
|
||||
}
|
||||
collector._consolidation_backlog = {_BacklogKey("public", None): 9}
|
||||
|
||||
obs = list(gauges["hindsight.async_operations"](None))
|
||||
by_label = {(o.attributes["operation_type"], o.attributes["status"]): o.value for o in obs}
|
||||
assert by_label[("retain", "pending")] == 7
|
||||
assert by_label[("consolidation", "processing")] == 1
|
||||
assert all("bank_id" not in o.attributes for o in obs) # cardinality guard
|
||||
|
||||
backlog_obs = list(gauges["hindsight.consolidation.backlog"](None))
|
||||
assert backlog_obs[0].value == 9
|
||||
assert backlog_obs[0].attributes["tenant"] == "public"
|
||||
|
||||
|
||||
def test_gauge_emits_bank_id_attribute_when_present():
|
||||
collector = _collector(include_bank_id=True)
|
||||
_set_db_pool_with_backlog_enabled(collector, MagicMock())
|
||||
gauges = {
|
||||
c.kwargs["name"]: c.kwargs["callbacks"][0]
|
||||
for c in collector.meter.create_observable_gauge.call_args_list
|
||||
if "callbacks" in c.kwargs
|
||||
}
|
||||
collector._consolidation_backlog = {_BacklogKey("public", "bankA"): 4}
|
||||
obs = list(gauges["hindsight.consolidation.backlog"](None))
|
||||
assert obs[0].value == 4
|
||||
assert obs[0].attributes["bank_id"] == "bankA"
|
||||
|
||||
|
||||
def test_backlog_gauges_not_registered_when_flag_disabled():
|
||||
"""Backlog metrics are off by default: set_db_pool must not register the
|
||||
gauges unless metrics_backlog_enabled is set."""
|
||||
collector = _collector()
|
||||
mock_config = MagicMock()
|
||||
mock_config.metrics_backlog_enabled = False
|
||||
with patch("hindsight_api.config.get_config", return_value=mock_config):
|
||||
collector.set_db_pool(MagicMock())
|
||||
|
||||
names = [
|
||||
c.kwargs.get("name") for c in collector.meter.create_observable_gauge.call_args_list if "callbacks" in c.kwargs
|
||||
]
|
||||
assert "hindsight.async_operations" not in names
|
||||
assert "hindsight.consolidation.backlog" not in names
|
||||
assert "hindsight.consolidation.failed" not in names
|
||||
assert collector._backlog_task is None
|
||||
@@ -82,8 +82,7 @@ async def test_bank_llm_not_configured(api_client, memory, monkeypatch):
|
||||
monkeypatch.setattr(cfg, "provider", "none")
|
||||
body = (await api_client.post("/v1/default/banks/llm-none/health/llm")).json()
|
||||
assert all(op["status"] == "not_configured" and op["ok"] is False for op in body["operations"])
|
||||
# latency_ms is null when not configured; responses omit null fields, so use .get().
|
||||
assert all(op.get("latency_ms") is None for op in body["operations"])
|
||||
assert all(op["latency_ms"] is None for op in body["operations"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -186,43 +186,6 @@ async def test_invalidate_drops_entry() -> None:
|
||||
assert calls[0] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_detaches_in_flight_loader() -> None:
|
||||
cache = BankStatsCache(ttl_seconds=60, max_entries=100)
|
||||
stale_started = asyncio.Event()
|
||||
release_stale = asyncio.Event()
|
||||
fresh_started = asyncio.Event()
|
||||
|
||||
async def stale_loader() -> dict[str, Any]:
|
||||
stale_started.set()
|
||||
await release_stale.wait()
|
||||
return {"v": "stale"}
|
||||
|
||||
async def fresh_loader() -> dict[str, Any]:
|
||||
fresh_started.set()
|
||||
return {"v": "fresh"}
|
||||
|
||||
stale_task = asyncio.create_task(cache.get_or_load("schema", "bank", stale_loader))
|
||||
await stale_started.wait()
|
||||
await cache.invalidate("schema", "bank")
|
||||
|
||||
# A request after invalidation must start a new load instead of joining the
|
||||
# pre-invalidation query, which may contain data from before a bank write.
|
||||
fresh_result = await asyncio.wait_for(cache.get_or_load("schema", "bank", fresh_loader), timeout=1)
|
||||
assert fresh_started.is_set()
|
||||
assert fresh_result == {"v": "fresh"}
|
||||
|
||||
release_stale.set()
|
||||
assert await stale_task == {"v": "stale"}
|
||||
|
||||
# The stale loader completed last, but must not overwrite the fresh value.
|
||||
async def should_not_run() -> dict[str, Any]:
|
||||
raise AssertionError("fresh value was not cached")
|
||||
|
||||
cached = await cache.get_or_load("schema", "bank", should_not_run)
|
||||
assert cached == {"v": "fresh"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clear_drops_all_entries() -> None:
|
||||
cache = BankStatsCache(ttl_seconds=60, max_entries=100)
|
||||
@@ -236,27 +199,3 @@ async def test_clear_drops_all_entries() -> None:
|
||||
await cache.get_or_load("s", "a", loader)
|
||||
await cache.get_or_load("s", "b", loader)
|
||||
assert calls[0] == 4
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clear_detaches_in_flight_loaders() -> None:
|
||||
cache = BankStatsCache(ttl_seconds=60, max_entries=100)
|
||||
stale_started = asyncio.Event()
|
||||
release_stale = asyncio.Event()
|
||||
|
||||
async def stale_loader() -> dict[str, Any]:
|
||||
stale_started.set()
|
||||
await release_stale.wait()
|
||||
return {"v": "stale"}
|
||||
|
||||
async def fresh_loader() -> dict[str, Any]:
|
||||
return {"v": "fresh"}
|
||||
|
||||
stale_task = asyncio.create_task(cache.get_or_load("schema", "bank", stale_loader))
|
||||
await stale_started.wait()
|
||||
await cache.clear()
|
||||
assert await cache.get_or_load("schema", "bank", fresh_loader) == {"v": "fresh"}
|
||||
|
||||
release_stale.set()
|
||||
assert await stale_task == {"v": "stale"}
|
||||
assert await cache.get_or_load("schema", "bank", fresh_loader) == {"v": "fresh"}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user