Compare commits
154
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c732864144 | ||
|
|
0b752448f7 | ||
|
|
888b50de12 | ||
|
|
fb7be3eced | ||
|
|
4499254f6d | ||
|
|
9943957fb7 | ||
|
|
03f47e29c8 | ||
|
|
1240b82629 | ||
|
|
08f1cda3bf | ||
|
|
a3a9d7b37d | ||
|
|
c2607d7699 | ||
|
|
e99ee0f243 | ||
|
|
c568094b8c | ||
|
|
5179d5f77d | ||
|
|
981cf6057f | ||
|
|
d90588b3e1 | ||
|
|
d0f67c9f8b | ||
|
|
fedfb494ee | ||
|
|
0430588e32 | ||
|
|
2af0e08dba | ||
|
|
f64817814a | ||
|
|
fa4cbf7ef2 | ||
|
|
2109397028 | ||
|
|
c4ef090a20 | ||
|
|
96f487213c | ||
|
|
0d8d805832 | ||
|
|
1cd836229b | ||
|
|
90ad003c46 | ||
|
|
278718dd84 | ||
|
|
093ecff48d | ||
|
|
85b9074f43 | ||
|
|
7e339e1677 | ||
|
|
dd621a69d0 | ||
|
|
7097716204 | ||
|
|
d3302c95b9 | ||
|
|
665877bb01 | ||
|
|
a43d208e93 | ||
|
|
34d9188e13 | ||
|
|
9a776e9f58 | ||
|
|
d02affd8f2 | ||
|
|
6b346925e2 | ||
|
|
63e2964a4c | ||
|
|
d5403a4b29 | ||
|
|
a24941f83b | ||
|
|
21b25fe8fe | ||
|
|
794a7435a9 | ||
|
|
038a9c2313 | ||
|
|
749478d9f9 | ||
|
|
96f0e54efa | ||
|
|
382550690a | ||
|
|
6c7f057e9d | ||
|
|
539190b69e | ||
|
|
1499ce5549 | ||
|
|
8564135b2a | ||
|
|
44d912533c | ||
|
|
35127d5f8b | ||
|
|
86c733c10e | ||
|
|
cb7ebe80bb | ||
|
|
615509011e | ||
|
|
af6bd1b5e1 | ||
|
|
579b10b53d | ||
|
|
4b57b82301 | ||
|
|
9c3fda74e2 | ||
|
|
f0cb1925ec | ||
|
|
039944cae2 | ||
|
|
ef9d3a15cb | ||
|
|
d788a55e28 | ||
|
|
c8ae82d62f | ||
|
|
27498f99d0 | ||
|
|
1530c09120 | ||
|
|
1163b1f6a6 | ||
|
|
fe88bdf704 | ||
|
|
cbb8fc6723 | ||
|
|
c33b9b8bb2 | ||
|
|
b364bc3402 | ||
|
|
35f0984b72 | ||
|
|
5dc45194c9 | ||
|
|
ff47814422 | ||
|
|
1ba70f81c8 | ||
|
|
fe15b5ec87 | ||
|
|
10e21f7302 | ||
|
|
7d3ac5ddb9 | ||
|
|
f4f86e3842 | ||
|
|
728ce13cea | ||
|
|
ecc590cb79 | ||
|
|
381c96c093 | ||
|
|
ab5e31f203 | ||
|
|
0da77ce2c9 | ||
|
|
d57e8639c5 | ||
|
|
03bf13e9e3 | ||
|
|
ff20bf9dc7 | ||
|
|
751f99a82f | ||
|
|
49ae55af03 | ||
|
|
c2ac7d0440 | ||
|
|
657fe023b2 | ||
|
|
9c95a1ac1d | ||
|
|
15540075b2 | ||
|
|
3f211f0729 | ||
|
|
8781c9fbfe | ||
|
|
12e9a3d305 | ||
|
|
c16ccc2c22 | ||
|
|
a7c094d436 | ||
|
|
b8f06a09fb | ||
|
|
b43ef98686 | ||
|
|
f17703fb37 | ||
|
|
cfcc23c152 | ||
|
|
7300d5be4b | ||
|
|
81c82d9b93 | ||
|
|
7551e65e55 | ||
|
|
94cc0a1270 | ||
|
|
67c47881cb | ||
|
|
2b72e1fd68 | ||
|
|
d2b797fff8 | ||
|
|
fccbdfef16 | ||
|
|
20f2b92069 | ||
|
|
1bf90358c3 | ||
|
|
2118d0a7cd | ||
|
|
e5fc6eedb6 | ||
|
|
bb0e0316a7 | ||
|
|
3172e99cab | ||
|
|
1c9a7a0d5e | ||
|
|
90e370ef35 | ||
|
|
084242a6dd | ||
|
|
83f44c4b41 | ||
|
|
7bdb8fc2e3 | ||
|
|
5b52a84fff | ||
|
|
f3c5a9c1c2 | ||
|
|
5832b907c6 | ||
|
|
50fa2ed090 | ||
|
|
522b71aab8 | ||
|
|
31b5c5845d | ||
|
|
c0ca9b027e | ||
|
|
1d4879a206 | ||
|
|
8e39cb7bc8 | ||
|
|
b378f6852f | ||
|
|
9c2df9d89f | ||
|
|
ec2231799e | ||
|
|
aebef9408b | ||
|
|
66abad61b8 | ||
|
|
9db64ecda3 | ||
|
|
ddaa5f5f1b | ||
|
|
87d4a36509 | ||
|
|
0bf85a3435 | ||
|
|
16b85a4faa | ||
|
|
4c792400c1 | ||
|
|
0284595909 | ||
|
|
fe4ed1db73 | ||
|
|
bac4b24e30 | ||
|
|
3290f4bfff | ||
|
|
63a65d0723 | ||
|
|
870cfccabb | ||
|
|
4476a10aa3 | ||
|
|
4f2833873c | ||
|
|
1eeced3116 |
+9
-1
@@ -2,7 +2,7 @@
|
||||
# Copy this file to .env and fill in your values
|
||||
|
||||
# LLM Configuration (Required)
|
||||
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio
|
||||
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio, vertexai
|
||||
HINDSIGHT_API_LLM_PROVIDER=openai
|
||||
HINDSIGHT_API_LLM_API_KEY=your-api-key-here
|
||||
HINDSIGHT_API_LLM_MODEL=o3-mini
|
||||
@@ -13,6 +13,13 @@ HINDSIGHT_API_LLM_BASE_URL=https://api.openai.com/v1
|
||||
# HINDSIGHT_API_LLM_API_KEY=your-anthropic-api-key
|
||||
# HINDSIGHT_API_LLM_MODEL=claude-sonnet-4-20250514
|
||||
|
||||
# Example: Google Vertex AI configuration
|
||||
# HINDSIGHT_API_LLM_PROVIDER=vertexai
|
||||
# HINDSIGHT_API_LLM_MODEL=google/gemini-2.0-flash-001
|
||||
# HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=your-gcp-project-id
|
||||
# HINDSIGHT_API_LLM_VERTEXAI_REGION=us-central1
|
||||
# HINDSIGHT_API_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY=/path/to/service-account-key.json # Optional, uses ADC if not set
|
||||
|
||||
# Example: LM Studio local configuration (Qwen 2.5 32B recommended)
|
||||
# HINDSIGHT_API_LLM_PROVIDER=lmstudio
|
||||
# HINDSIGHT_API_LLM_API_KEY=lmstudio
|
||||
@@ -26,6 +33,7 @@ HINDSIGHT_API_LOG_LEVEL=info
|
||||
|
||||
# Database (Optional - uses embedded pg0 by default)
|
||||
# HINDSIGHT_API_DATABASE_URL=postgresql://user:pass@host:5432/db
|
||||
# HINDSIGHT_API_DATABASE_SCHEMA=public # PostgreSQL schema name (default: public)
|
||||
|
||||
# Embeddings Configuration (Optional - uses local by default)
|
||||
# Provider: "local" (default) or "tei" (HuggingFace Text Embeddings Inference)
|
||||
|
||||
@@ -139,6 +139,104 @@ jobs:
|
||||
path: hindsight-clients/typescript/*.tgz
|
||||
retention-days: 1
|
||||
|
||||
release-openclaw-integration:
|
||||
runs-on: ubuntu-latest
|
||||
environment: npm
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
registry-url: 'https://registry.npmjs.org'
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/openclaw
|
||||
run: npm ci
|
||||
|
||||
- name: Build
|
||||
working-directory: ./hindsight-integrations/openclaw
|
||||
run: npm run build
|
||||
|
||||
- name: Publish to npm
|
||||
working-directory: ./hindsight-integrations/openclaw
|
||||
run: |
|
||||
set +e
|
||||
OUTPUT=$(npm publish --access public 2>&1)
|
||||
EXIT_CODE=$?
|
||||
echo "$OUTPUT"
|
||||
if [ $EXIT_CODE -ne 0 ]; then
|
||||
if echo "$OUTPUT" | grep -q "cannot publish over"; then
|
||||
echo "Package version already published, skipping..."
|
||||
exit 0
|
||||
fi
|
||||
exit $EXIT_CODE
|
||||
fi
|
||||
env:
|
||||
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
|
||||
|
||||
- name: Pack for GitHub release
|
||||
working-directory: ./hindsight-integrations/openclaw
|
||||
run: npm pack
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: openclaw-integration
|
||||
path: hindsight-integrations/openclaw/*.tgz
|
||||
retention-days: 1
|
||||
|
||||
release-ai-sdk-integration:
|
||||
runs-on: ubuntu-latest
|
||||
environment: npm
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
registry-url: 'https://registry.npmjs.org'
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/ai-sdk
|
||||
run: npm ci
|
||||
|
||||
- name: Build
|
||||
working-directory: ./hindsight-integrations/ai-sdk
|
||||
run: npm run build
|
||||
|
||||
- name: Publish to npm
|
||||
working-directory: ./hindsight-integrations/ai-sdk
|
||||
run: |
|
||||
set +e
|
||||
OUTPUT=$(npm publish --access public 2>&1)
|
||||
EXIT_CODE=$?
|
||||
echo "$OUTPUT"
|
||||
if [ $EXIT_CODE -ne 0 ]; then
|
||||
if echo "$OUTPUT" | grep -q "cannot publish over"; then
|
||||
echo "Package version already published, skipping..."
|
||||
exit 0
|
||||
fi
|
||||
exit $EXIT_CODE
|
||||
fi
|
||||
env:
|
||||
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
|
||||
|
||||
- name: Pack for GitHub release
|
||||
working-directory: ./hindsight-integrations/ai-sdk
|
||||
run: npm pack
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ai-sdk-integration
|
||||
path: hindsight-integrations/ai-sdk/*.tgz
|
||||
retention-days: 1
|
||||
|
||||
release-control-plane:
|
||||
runs-on: ubuntu-latest
|
||||
environment: npm
|
||||
@@ -242,6 +340,7 @@ jobs:
|
||||
retention-days: 1
|
||||
|
||||
release-docker-images:
|
||||
name: Release Docker (${{ matrix.image_name }}${{ matrix.tag_suffix }})
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
@@ -251,10 +350,28 @@ jobs:
|
||||
include:
|
||||
- target: api-only
|
||||
image_name: hindsight-api
|
||||
tag_suffix: ""
|
||||
build_args: ""
|
||||
- target: api-only
|
||||
image_name: hindsight-api
|
||||
tag_suffix: "-slim"
|
||||
build_args: |
|
||||
INCLUDE_LOCAL_MODELS=false
|
||||
PRELOAD_ML_MODELS=false
|
||||
- target: cp-only
|
||||
image_name: hindsight-control-plane
|
||||
tag_suffix: ""
|
||||
build_args: ""
|
||||
- target: standalone
|
||||
image_name: hindsight
|
||||
tag_suffix: ""
|
||||
build_args: ""
|
||||
- target: standalone
|
||||
image_name: hindsight
|
||||
tag_suffix: "-slim"
|
||||
build_args: |
|
||||
INCLUDE_LOCAL_MODELS=false
|
||||
PRELOAD_ML_MODELS=false
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
@@ -292,6 +409,9 @@ jobs:
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ghcr.io/${{ github.repository_owner }}/${{ matrix.image_name }}
|
||||
flavor: |
|
||||
latest=auto
|
||||
suffix=${{ matrix.tag_suffix }}
|
||||
tags: |
|
||||
type=semver,pattern={{version}},value=${{ steps.get_version.outputs.VERSION }}
|
||||
type=semver,pattern={{major}}.{{minor}},value=${{ steps.get_version.outputs.VERSION }}
|
||||
@@ -317,7 +437,7 @@ jobs:
|
||||
# - name: Smoke test - verify container starts
|
||||
# env:
|
||||
# GROQ_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
# run: ./scripts/docker-smoke-test.sh "${{ matrix.image_name }}:test" "${{ matrix.target }}"
|
||||
# run: ./docker/test-image.sh "${{ matrix.image_name }}:test" "${{ matrix.target }}"
|
||||
|
||||
# Build multi-platform and push to release tags
|
||||
- name: Build and push release images
|
||||
@@ -326,6 +446,7 @@ jobs:
|
||||
context: .
|
||||
file: docker/standalone/Dockerfile
|
||||
target: ${{ matrix.target }}
|
||||
build-args: ${{ matrix.build_args }}
|
||||
push: true
|
||||
platforms: linux/amd64,linux/arm64
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
@@ -366,7 +487,7 @@ jobs:
|
||||
|
||||
create-github-release:
|
||||
runs-on: ubuntu-latest
|
||||
needs: [release-python-packages, release-typescript-client, release-control-plane, release-rust-cli, release-docker-images, release-helm-chart]
|
||||
needs: [release-python-packages, release-typescript-client, release-openclaw-integration, release-ai-sdk-integration, release-control-plane, release-rust-cli, release-docker-images, release-helm-chart]
|
||||
permissions:
|
||||
contents: write
|
||||
|
||||
@@ -389,6 +510,18 @@ jobs:
|
||||
name: typescript-client
|
||||
path: ./artifacts/typescript-client
|
||||
|
||||
- name: Download OpenClaw Integration
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: openclaw-integration
|
||||
path: ./artifacts/openclaw-integration
|
||||
|
||||
- name: Download AI SDK Integration
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: ai-sdk-integration
|
||||
path: ./artifacts/ai-sdk-integration
|
||||
|
||||
- name: Download Control Plane
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
@@ -430,6 +563,10 @@ jobs:
|
||||
cp artifacts/python-packages/hindsight-embed/dist/* release-assets/ || true
|
||||
# TypeScript client
|
||||
cp artifacts/typescript-client/*.tgz release-assets/ || true
|
||||
# OpenClaw Integration
|
||||
cp artifacts/openclaw-integration/*.tgz release-assets/ || true
|
||||
# AI SDK Integration
|
||||
cp artifacts/ai-sdk-integration/*.tgz release-assets/ || true
|
||||
# Control Plane
|
||||
cp artifacts/control-plane/*.tgz release-assets/ || true
|
||||
# Rust CLI binaries
|
||||
|
||||
+274
-69
@@ -9,42 +9,11 @@ concurrency:
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
build-python-packages:
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
matrix:
|
||||
include:
|
||||
- name: hindsight-all
|
||||
path: hindsight
|
||||
- name: hindsight-api
|
||||
path: hindsight-api
|
||||
- name: hindsight-client
|
||||
path: hindsight-clients/python
|
||||
- name: hindsight-embed
|
||||
path: hindsight-embed
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build ${{ matrix.name }}
|
||||
working-directory: ./${{ matrix.path }}
|
||||
run: uv build
|
||||
|
||||
build-api-python-versions:
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
matrix:
|
||||
python-version: ['3.11', '3.12', '3.13']
|
||||
python-version: ['3.11', '3.12', '3.13', '3.14']
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
@@ -82,6 +51,52 @@ jobs:
|
||||
- name: Build TypeScript client
|
||||
run: npm run build --workspace=hindsight-clients/typescript
|
||||
|
||||
build-openclaw-integration:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/openclaw
|
||||
run: npm ci
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/openclaw
|
||||
run: npm test
|
||||
|
||||
- name: Build
|
||||
working-directory: ./hindsight-integrations/openclaw
|
||||
run: npm run build
|
||||
|
||||
build-ai-sdk-integration:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/ai-sdk
|
||||
run: npm ci
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/ai-sdk
|
||||
run: npm test
|
||||
|
||||
- name: Build
|
||||
working-directory: ./hindsight-integrations/ai-sdk
|
||||
run: npm run build
|
||||
|
||||
build-control-plane:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
@@ -153,8 +168,15 @@ jobs:
|
||||
- name: Build docs
|
||||
run: npm run build --workspace=hindsight-docs
|
||||
|
||||
build-rust-cli:
|
||||
test-rust-cli:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
HINDSIGHT_API_LLM_PROVIDER: groq
|
||||
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
|
||||
HINDSIGHT_API_URL: http://localhost:8888
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
@@ -171,6 +193,10 @@ jobs:
|
||||
hindsight-cli/target
|
||||
key: ${{ runner.os }}-cargo-${{ hashFiles('**/Cargo.lock') }}
|
||||
|
||||
- name: Run unit tests
|
||||
working-directory: hindsight-cli
|
||||
run: cargo test
|
||||
|
||||
- name: Build CLI
|
||||
working-directory: hindsight-cli
|
||||
run: cargo build --release
|
||||
@@ -182,29 +208,6 @@ jobs:
|
||||
path: hindsight-cli/target/release/hindsight
|
||||
retention-days: 1
|
||||
|
||||
test-rust-cli:
|
||||
runs-on: ubuntu-latest
|
||||
needs: build-rust-cli
|
||||
env:
|
||||
HINDSIGHT_API_LLM_PROVIDER: groq
|
||||
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
|
||||
HINDSIGHT_API_URL: http://localhost:8888
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Download CLI artifact
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: hindsight-cli
|
||||
path: /tmp/cli
|
||||
|
||||
- name: Make CLI executable
|
||||
run: chmod +x /tmp/cli/hindsight
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
@@ -251,7 +254,7 @@ jobs:
|
||||
|
||||
- name: Run CLI smoke test
|
||||
run: |
|
||||
HINDSIGHT_CLI=/tmp/cli/hindsight ./hindsight-cli/smoke-test.sh
|
||||
HINDSIGHT_CLI=hindsight-cli/target/release/hindsight ./hindsight-cli/smoke-test.sh
|
||||
|
||||
- name: Show API server logs
|
||||
if: always()
|
||||
@@ -274,16 +277,35 @@ jobs:
|
||||
run: helm lint helm/hindsight
|
||||
|
||||
build-docker-images:
|
||||
name: Build Docker (${{ matrix.name }})
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
matrix:
|
||||
include:
|
||||
- target: api-only
|
||||
name: api
|
||||
variant: full
|
||||
build_args: ""
|
||||
- target: api-only
|
||||
name: api-slim
|
||||
variant: slim
|
||||
build_args: |
|
||||
INCLUDE_LOCAL_MODELS=false
|
||||
PRELOAD_ML_MODELS=false
|
||||
- target: cp-only
|
||||
name: control-plane
|
||||
variant: full
|
||||
build_args: ""
|
||||
- target: standalone
|
||||
name: standalone
|
||||
variant: full
|
||||
build_args: ""
|
||||
- target: standalone
|
||||
name: standalone-slim
|
||||
variant: slim
|
||||
build_args: |
|
||||
INCLUDE_LOCAL_MODELS=false
|
||||
PRELOAD_ML_MODELS=false
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
@@ -302,20 +324,31 @@ jobs:
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Build ${{ matrix.name }} image
|
||||
- name: Build ${{ matrix.name }} image (${{ matrix.variant }})
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: docker/standalone/Dockerfile
|
||||
target: ${{ matrix.target }}
|
||||
build-args: ${{ matrix.build_args }}
|
||||
push: false
|
||||
load: false
|
||||
load: ${{ matrix.variant == 'slim' }}
|
||||
tags: hindsight-${{ matrix.name }}:test
|
||||
# Removed GitHub Actions cache (type=gha) - it frequently returns 502 errors
|
||||
# causing buildx to fail with "failed to parse error response 502"
|
||||
# Build will be slower but more reliable
|
||||
|
||||
# TODO: Re-enable smoke test when disk space issue is resolved
|
||||
# - name: Smoke test - verify container starts
|
||||
# env:
|
||||
# GROQ_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
# run: ./scripts/docker-smoke-test.sh "hindsight-${{ matrix.name }}:test" "${{ matrix.target }}"
|
||||
# Only test slim variants to save disk space (they're much smaller)
|
||||
# Slim variants require external embedding providers
|
||||
- name: Smoke test - verify container starts
|
||||
if: matrix.variant == 'slim'
|
||||
env:
|
||||
GROQ_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_API_EMBEDDINGS_PROVIDER: openai
|
||||
HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
|
||||
HINDSIGHT_API_RERANKER_PROVIDER: cohere
|
||||
HINDSIGHT_API_COHERE_API_KEY: ${{ secrets.COHERE_API_KEY }}
|
||||
run: ./docker/test-image.sh "hindsight-${{ matrix.name }}:test" "${{ matrix.target }}"
|
||||
|
||||
test-api:
|
||||
runs-on: ubuntu-latest
|
||||
@@ -738,9 +771,9 @@ jobs:
|
||||
test-embed:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
HINDSIGHT_EMBED_LLM_PROVIDER: groq
|
||||
HINDSIGHT_EMBED_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_EMBED_LLM_MODEL: openai/gpt-oss-20b
|
||||
HINDSIGHT_API_LLM_PROVIDER: groq
|
||||
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
|
||||
# Prefer CPU-only PyTorch in CI
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
@@ -771,13 +804,65 @@ jobs:
|
||||
${{ runner.os }}-huggingface-embed-
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Run unit and integration tests
|
||||
working-directory: ./hindsight-embed
|
||||
run: uv run pytest tests/ -v
|
||||
|
||||
- name: Run smoke test
|
||||
working-directory: ./hindsight-embed
|
||||
run: ./test.sh
|
||||
|
||||
test-hindsight-all:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
HINDSIGHT_API_LLM_PROVIDER: groq
|
||||
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
|
||||
# For test_server_integration.py compatibility
|
||||
HINDSIGHT_LLM_PROVIDER: groq
|
||||
HINDSIGHT_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_LLM_MODEL: openai/gpt-oss-20b
|
||||
# Prefer CPU-only PyTorch in CI
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build hindsight-all
|
||||
working-directory: ./hindsight
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight
|
||||
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-all-${{ hashFiles('hindsight/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-all-
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Run unit tests
|
||||
working-directory: ./hindsight
|
||||
run: uv run pytest tests/ -v
|
||||
|
||||
test-doc-examples:
|
||||
runs-on: ubuntu-latest
|
||||
needs: build-rust-cli
|
||||
needs: test-rust-cli
|
||||
env:
|
||||
HINDSIGHT_API_LLM_PROVIDER: groq
|
||||
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
@@ -887,6 +972,78 @@ jobs:
|
||||
echo "=== API Server Logs ==="
|
||||
cat /tmp/api-server.log || echo "No API server log found"
|
||||
|
||||
test-upgrade:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
HINDSIGHT_API_LLM_PROVIDER: groq
|
||||
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0 # Full history needed for git clone of tags
|
||||
|
||||
- name: Fetch tags
|
||||
run: git fetch --tags
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Install hindsight-dev dependencies
|
||||
working-directory: ./hindsight-dev
|
||||
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install current hindsight-api
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --index-strategy unsafe-best-match
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
print('Downloading embedding model...')
|
||||
SentenceTransformer('BAAI/bge-small-en-v1.5')
|
||||
print('Downloading cross-encoder model...')
|
||||
CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2')
|
||||
print('Models downloaded successfully')
|
||||
"
|
||||
|
||||
- name: Run upgrade tests
|
||||
working-directory: ./hindsight-dev
|
||||
run: uv run pytest upgrade_tests/ -v --tb=short
|
||||
|
||||
- name: Show upgrade test logs
|
||||
if: always()
|
||||
run: |
|
||||
echo "=== Upgrade Test Server Logs ==="
|
||||
for log in /tmp/upgrade-test-*.log; do
|
||||
if [ -f "$log" ]; then
|
||||
echo ""
|
||||
echo "--- $log ---"
|
||||
tail -500 "$log"
|
||||
fi
|
||||
done
|
||||
|
||||
verify-generated-files:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
@@ -957,4 +1114,52 @@ jobs:
|
||||
git diff --stat
|
||||
exit 1
|
||||
fi
|
||||
echo "✓ All generated files are up to date"
|
||||
echo "✓ All generated files are up to date"
|
||||
|
||||
check-openapi-compatibility:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0 # Fetch full git history to access base branch
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Install hindsight-dev dependencies
|
||||
run: |
|
||||
cd hindsight-dev && uv sync --frozen --index-strategy unsafe-best-match
|
||||
|
||||
- name: Check OpenAPI compatibility with base branch
|
||||
run: |
|
||||
# Get the base branch (usually main)
|
||||
BASE_BRANCH="${{ github.base_ref }}"
|
||||
|
||||
if [ -z "$BASE_BRANCH" ]; then
|
||||
echo "⚠️ Warning: No base branch found (not a PR?). Skipping compatibility check."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
echo "Checking OpenAPI compatibility against base branch: $BASE_BRANCH"
|
||||
|
||||
# Extract the old OpenAPI spec from base branch
|
||||
git show "origin/$BASE_BRANCH:hindsight-docs/static/openapi.json" > /tmp/old-openapi.json
|
||||
|
||||
if [ ! -s /tmp/old-openapi.json ]; then
|
||||
echo "⚠️ Warning: Could not find OpenAPI spec in base branch. Skipping compatibility check."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Check compatibility using our tool
|
||||
cd hindsight-dev
|
||||
uv run check-openapi-compatibility /tmp/old-openapi.json ../hindsight-docs/static/openapi.json
|
||||
+7
-2
@@ -29,7 +29,7 @@ nltk_data/
|
||||
|
||||
# Monitoring stack (Prometheus/Grafana binaries and data)
|
||||
.monitoring/
|
||||
.pgbouncer
|
||||
.pgbouncer/
|
||||
|
||||
# Large benchmark datasets (will be downloaded automatically)
|
||||
**/longmemeval_s_cleaned.json
|
||||
@@ -45,9 +45,14 @@ hindsight-docs/static/llms-full.txt
|
||||
|
||||
hindsight-dev/benchmarks/locomo/results/
|
||||
hindsight-dev/benchmarks/longmemeval/results/
|
||||
hindsight-dev/benchmarks/consolidation/results/
|
||||
benchmarks/results/
|
||||
hindsight-cli/target
|
||||
hindsight-clients/rust/target
|
||||
.claude
|
||||
whats-next.md
|
||||
TASK.md
|
||||
CHANGELOG.md
|
||||
# Changelog is now tracked in hindsight-docs/src/pages/changelog.md
|
||||
# CHANGELOG.md
|
||||
|
||||
blog-post*
|
||||
@@ -1,153 +1,3 @@
|
||||
# AGENTS.md
|
||||
|
||||
This document captures architectural decisions and coding conventions for the Hindsight project.
|
||||
|
||||
## Documentation
|
||||
|
||||
- **Main documentation**: [hindsight-docs/docs/developer/](./hindsight-docs/docs/developer/)
|
||||
- **Use case patterns**: [hindsight-docs/docs/cookbook/](./hindsight-docs/docs/cookbook/)
|
||||
- **API reference**: Auto-generated from OpenAPI spec
|
||||
|
||||
## Project Structure
|
||||
|
||||
```
|
||||
hindsight/ # Python package for embedded usage
|
||||
hindsight-api/ # FastAPI server (core memory engine)
|
||||
hindsight-cli/ # Rust CLI client
|
||||
hindsight-embed/ # Embedded CLI (no server needed)
|
||||
hindsight-control-plane/ # Next.js admin UI
|
||||
hindsight-docs/ # Docusaurus documentation site
|
||||
hindsight-dev/ # Development tools and benchmarks
|
||||
hindsight-integrations/ # Framework integrations (LangChain, etc.)
|
||||
hindsight-clients/ # Generated API clients (Python, TypeScript, Rust)
|
||||
```
|
||||
|
||||
## Core Concepts
|
||||
|
||||
### Memory Banks
|
||||
- Each bank is an isolated memory store (like a "brain" for one user/agent)
|
||||
- Banks contain: memory units (facts), entities, documents, entity links
|
||||
- Banks have a **disposition** (personality traits) and **background** (context)
|
||||
- Bank isolation is strict - no cross-bank data leakage
|
||||
|
||||
### Memory Types
|
||||
- **World facts**: General knowledge ("The sky is blue")
|
||||
- **Experience facts**: Personal experiences ("I visited Paris in 2023")
|
||||
- **Opinion facts**: Beliefs with confidence scores ("Paris is beautiful" - 0.9 confidence)
|
||||
|
||||
### Operations
|
||||
- **Retain**: Store new memories (extracts facts, entities, relationships)
|
||||
- **Recall**: Retrieve memories (semantic, BM25, graph, temporal search)
|
||||
- **Reflect**: Deep analysis to form new insights/opinions
|
||||
|
||||
## API Design Decisions
|
||||
|
||||
### Single Bank Per Request
|
||||
- All API endpoints (`recall`, `reflect`, `retain`) operate on a single bank
|
||||
- Multi-bank queries are the **client/agent's responsibility** to orchestrate
|
||||
- This keeps the API simple and the isolation model clear
|
||||
|
||||
### Disposition Traits (3-trait system)
|
||||
- **Skepticism** (1-5): How skeptical vs trusting when forming opinions
|
||||
- **Literalism** (1-5): How literally to interpret information
|
||||
- **Empathy** (1-5): How much to consider emotional context
|
||||
- These influence the `reflect` operation, not `recall`
|
||||
- Background info also only affects `reflect` (opinion formation)
|
||||
|
||||
## Multi-Bank Architecture Patterns
|
||||
|
||||
See [hindsight-docs/docs/cookbook/](./hindsight-docs/docs/cookbook/) for detailed guides:
|
||||
|
||||
- **Per-User Memory**: One bank per user, simplest pattern
|
||||
- **Support Agent + Shared Knowledge**: User bank + shared docs bank, client orchestrates
|
||||
|
||||
## Developer Guide
|
||||
|
||||
### Running the API Server
|
||||
|
||||
```bash
|
||||
# From project root
|
||||
./scripts/dev/start-api.sh
|
||||
|
||||
# With options
|
||||
./scripts/dev/start-api.sh --reload --port 8888 --log-level debug
|
||||
```
|
||||
|
||||
### Running Tests
|
||||
|
||||
```bash
|
||||
# API tests
|
||||
cd hindsight-api
|
||||
uv run pytest tests/
|
||||
|
||||
# Specific test
|
||||
uv run pytest tests/test_http_api_integration.py -v
|
||||
```
|
||||
|
||||
### Generating OpenAPI Spec
|
||||
|
||||
After changing API endpoints, regenerate the OpenAPI spec and docs:
|
||||
|
||||
```bash
|
||||
./scripts/generate-openapi.sh
|
||||
```
|
||||
|
||||
This will:
|
||||
1. Generate `openapi.json` at project root
|
||||
2. Copy to `hindsight-docs/openapi.json`
|
||||
3. Regenerate API reference documentation
|
||||
|
||||
### Generating API Clients
|
||||
|
||||
After updating the OpenAPI spec, regenerate all clients:
|
||||
|
||||
```bash
|
||||
./scripts/generate-clients.sh
|
||||
```
|
||||
|
||||
This generates:
|
||||
- **Rust client**: `hindsight-clients/rust/` (via progenitor in build.rs)
|
||||
- **Python client**: `hindsight-clients/python/` (via openapi-generator Docker)
|
||||
- **TypeScript client**: `hindsight-clients/typescript/` (via @hey-api/openapi-ts)
|
||||
|
||||
Note: The maintained wrapper `hindsight_client.py` and `README.md` are preserved during regeneration.
|
||||
|
||||
### Running the Documentation Site
|
||||
|
||||
```bash
|
||||
./scripts/dev/start-docs.sh
|
||||
```
|
||||
|
||||
### Running the Control Plane
|
||||
|
||||
```bash
|
||||
./scripts/dev/start-control-plane.sh
|
||||
```
|
||||
|
||||
## Code Style
|
||||
|
||||
### Python (hindsight-api)
|
||||
- Use `uv` for package management
|
||||
- Async throughout (asyncpg, async FastAPI endpoints)
|
||||
- Pydantic models for request/response validation
|
||||
- No py files at project root - maintain clean directory structure
|
||||
|
||||
### TypeScript (control-plane, clients)
|
||||
- Next.js with App Router for control plane
|
||||
- Tailwind CSS with shadcn/ui components
|
||||
|
||||
### Rust (CLI)
|
||||
- Async with tokio
|
||||
- reqwest for HTTP client
|
||||
- progenitor for API client generation
|
||||
|
||||
## Database
|
||||
|
||||
- PostgreSQL with pgvector extension
|
||||
- Schema managed via Alembic migrations in `hindsight-api/alembic/`, db migrations happen during api startup, no manual commands
|
||||
- Key tables: `banks`, `memory_units`, `documents`, `entities`, `entity_links`
|
||||
|
||||
# Branding
|
||||
## Colors
|
||||
- Primary: gradient from #0074d9 to #009296
|
||||
|
||||
See [CLAUDE.md](./CLAUDE.md) for project documentation and coding conventions.
|
||||
|
||||
@@ -7,8 +7,7 @@ This file provides guidance to Claude Code (claude.ai/code) when working with co
|
||||
Hindsight is an agent memory system that provides long-term memory for AI agents using biomimetic data structures. Memories are organized as:
|
||||
- **World facts**: General knowledge ("The sky is blue")
|
||||
- **Experience facts**: Personal experiences ("I visited Paris in 2023")
|
||||
- **Opinion facts**: Beliefs with confidence scores ("Paris is beautiful" - 0.9 confidence)
|
||||
- **Observations**: Complex mental models derived from reflection
|
||||
- **Mental models**: Consolidated knowledge synthesized from facts ("User prefers functional programming patterns")
|
||||
|
||||
## Development Commands
|
||||
|
||||
@@ -101,7 +100,7 @@ cd hindsight-control-plane && npm run dev
|
||||
Main operations:
|
||||
- **Retain**: Store memories, extracts facts/entities/relationships
|
||||
- **Recall**: Retrieve memories via 4 parallel strategies (semantic, BM25, graph, temporal) + reranking
|
||||
- **Reflect**: Deep analysis forming new opinions/observations (disposition-aware)
|
||||
- **Reflect**: Disposition-aware reasoning using memories and mental models.
|
||||
|
||||
### Database
|
||||
PostgreSQL with pgvector. Schema managed via Alembic migrations in `hindsight-api/hindsight_api/alembic/`. Migrations run automatically on API startup.
|
||||
@@ -199,6 +198,38 @@ When adding or modifying parameters in the dataplane API (hindsight-api), you mu
|
||||
- Pydantic models for request/response
|
||||
- Ruff for linting (line-length 120)
|
||||
- No Python files at project root - maintain clean directory structure
|
||||
- **Never use multi-item tuple return values** - prefer dataclass or Pydantic model for structured returns
|
||||
|
||||
### Type Safety with Pydantic Models
|
||||
**NEVER use raw `dict` types for structured data.** Always use Pydantic models:
|
||||
- Use Pydantic `BaseModel` for all data structures passed between functions
|
||||
- Add `@field_validator` for type coercion (e.g., ensuring datetimes are timezone-aware)
|
||||
- Avoid `dict.get()` patterns - use typed model attributes instead
|
||||
- Parse external data (JSON, API responses) into Pydantic models at the boundary
|
||||
- This catches type errors at parse time, not deep in business logic
|
||||
|
||||
```python
|
||||
# BAD - error-prone dict access
|
||||
def process(data: dict) -> str:
|
||||
return data.get("name", "") # No validation, silent failures
|
||||
|
||||
# GOOD - typed and validated
|
||||
class UserData(BaseModel):
|
||||
name: str
|
||||
created_at: datetime
|
||||
|
||||
@field_validator("created_at", mode="before")
|
||||
@classmethod
|
||||
def ensure_tz_aware(cls, v):
|
||||
if isinstance(v, str):
|
||||
v = datetime.fromisoformat(v.replace("Z", "+00:00"))
|
||||
if v.tzinfo is None:
|
||||
return v.replace(tzinfo=timezone.utc)
|
||||
return v
|
||||
|
||||
def process(data: UserData) -> str:
|
||||
return data.name # Type-safe, validated at construction
|
||||
```
|
||||
|
||||
### TypeScript Style
|
||||
- Next.js App Router for control plane
|
||||
|
||||
@@ -93,6 +93,34 @@ uv run ty check hindsight_api # Type check
|
||||
3. Run tests to ensure nothing breaks
|
||||
4. Submit a PR with a clear description of changes
|
||||
|
||||
## Release Process
|
||||
|
||||
The project uses `scripts/release.sh` for creating releases. This script automates the entire release workflow:
|
||||
|
||||
1. Bumps version in all components (API, clients, CLI, control plane, Helm)
|
||||
2. **Regenerates OpenAPI spec and client SDKs** (Python, TypeScript, Rust)
|
||||
3. Updates documentation versioning
|
||||
4. Creates a commit and git tag
|
||||
5. Pushes to GitHub (triggers CI/CD to publish packages)
|
||||
|
||||
### Usage
|
||||
|
||||
```bash
|
||||
./scripts/release.sh <version>
|
||||
```
|
||||
|
||||
**Example:**
|
||||
```bash
|
||||
./scripts/release.sh 0.5.0
|
||||
```
|
||||
|
||||
### Important for Developers
|
||||
|
||||
- During development, version bumps in `__init__.py` do NOT require client regeneration
|
||||
- Clients are only regenerated during releases
|
||||
- Do not manually run `./scripts/generate-clients.sh` unless testing generation changes
|
||||
- Client version comments will reflect the API version from the latest release
|
||||
|
||||
## Reporting Issues
|
||||
|
||||
Open an issue on GitHub with:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
<div align="center">
|
||||
|
||||

|
||||

|
||||
|
||||
[Documentation](https://hindsight.vectorize.io) • [Paper](https://arxiv.org/abs/2512.12818) • [Cookbook](https://hindsight.vectorize.io/cookbook) • [Hindsight Cloud](https://vectorize.io/hindsight/cloud)
|
||||
|
||||
@@ -17,76 +17,76 @@
|
||||
|
||||
## What is Hindsight?
|
||||
|
||||
Hindsight™ is an agent memory system built to create smarter agents that learn over time. It eliminates the shortcomings of alternative techniques such as RAG and knowledge graph and delivers state-of-the-art performance on long term memory tasks.
|
||||
Hindsight™ is an agent memory system built to create smarter agents that learn over time. Most agent memory systems focus on recalling conversation history. Hindsight is focused on making agents that learn, not just remember.
|
||||
|
||||
Hindsight addresses common challenges that have frustrated AI engineers building agents to automate tasks and assist users with conversational interfaces. Many of these challenges stem directly from a lack of memory.
|
||||
|
||||
- **Inconsistency:** Agents complete tasks successfully one time, then fail when asked to complete the same task again. Memory gives the agent a mechanism to remember what worked and what didn't and to use that information to reduce errors and improve consistency.
|
||||
- **Hallucinations:** Long term memory can be seeded with external knowledge to ground agent behavior in reliable sources to augment training data.
|
||||
- **Cognitive Overload:** As workflows get complex, retrievals, tool calls, user messages and agent responses can grow to fill the context window leading to context rot. Short term memory optimization allows agents to reduce tokens and focus context by removing irrelevant details.
|
||||
<video src="https://github.com/user-attachments/assets/923b798d-3581-4897-bb62-9cfa5a931682" controls></video>
|
||||
|
||||
## How is Hindsight Different From Other Memory Systems?
|
||||
|
||||

|
||||
|
||||
Most agent memory implementation rely on basic vector search or sometimes use a knowledge graph. Hindsight uses biomimetic data structures to organize agent memories in a way that is more like how human memory works:
|
||||
|
||||
- **World:** Facts about the world ("The stove gets hot")
|
||||
- **Experiences:** Agent's own experiences ("I touched the stove and it really hurt")
|
||||
- **Opinion:** Beliefs with confidence scores ("I shouldn't touch the stove again" - .99 confidence)
|
||||
- **Observation:** Complex mental models derived by reflecting on facts and experiences ("Curling irons, ovens, and fire are also hot. I shouldn't touch those either.")
|
||||
|
||||
Memories in Hindsight are stored in banks (i.e. memory banks). When memories are added to Hindsight, they are pushed into either the world facts or experiences memory pathway. They are then represented as a combination of entities, relationships, and time series with sparse/dense vector representations to aid in later recall.
|
||||
|
||||
Hindsight provides three simple methods to interact with the system:
|
||||
|
||||
- **Retain:** Provide information to Hindsight that you want it to remember
|
||||
- **Recall:** Retrieve memories from Hindsight
|
||||
- **Reflect:** Reflect on memories and experiences to generate new observations and insights from existing memories.
|
||||
|
||||
### Agent Memory That Learns
|
||||
|
||||
A key goal of Hindsight is to build agent memory that enables agents to learn and improve over time. This is the role of the `reflect` operation which provides the agent to form broader opinions and observations over time.
|
||||
|
||||
For example, imagine a product support agent that is helping a user troubleshoot a problem. It uses a `search-documentation` tool it found on an MCP server. Later in the conversation, the agent discovers that the documentation returned from the tool wasn't for the product the user was asking about. The agent now has an experience in its memory bank. And just like humans, we want that agent to learn from its experience.
|
||||
|
||||
As the agent gains more experiences, `reflect` allows the agent to form observations about what worked, what didn't, and what to do differently the next time it encounters a similar task.
|
||||
|
||||
---
|
||||
It eliminates the shortcomings of alternative techniques such as RAG and knowledge graph and delivers state-of-the-art performance on long term memory tasks.
|
||||
|
||||
## Memory Performance & Accuracy
|
||||
|
||||
Hindsight has achieved state-of-the-art performance on the LongMemEval benchmark, widely used to assess memory system performance across a variety of conversational
|
||||
AI scenarios. The current reported performance of Hindsight and other agent memory solutions as of December 2025 is shown here:
|
||||
Hindsight is the most accurate agent memory system ever tested according to benchmark performance. It has achieved state-of-the-art performance on the LongMemEval benchmark, widely used to assess memory system performance across a variety of conversational AI scenarios. The current reported performance of Hindsight and other agent memory solutions as of January 2026 is shown here:
|
||||
|
||||

|
||||
|
||||
The benchmark performance data for Hindsight and GPT-4o (full context) have been reproduced by research collaborators at the Virginia Tech [Sanghani Center for Artificial Intelligence and Data Analytics](https://sanghani.cs.vt.edu/) and The Washington Post. Other scores are self-reported by software vendors.
|
||||
The benchmark performance data for Hindsight has been independently reproduced by research collaborators at the Virginia Tech [Sanghani Center for Artificial Intelligence and Data Analytics](https://sanghani.cs.vt.edu/) and The Washington Post. Other scores are self-reported by software vendors.
|
||||
|
||||
A thorough examination of the techniques implemented in Hindsight and detailed breakdowns of benchmark performance are [available on arXiv](https://arxiv.org/abs/2512.12818). This research is currently being prepared for conference submission and the wider peer review process.
|
||||
Hindsight is being used in production at Fortune 500 enterprises and by a growing number of AI startups.
|
||||
|
||||
## Adding Hindsight to Your AI Agents
|
||||
|
||||
The easiest way use Hindsight with an existing agent is with the LLM Wrapper. You can add memory to your agent with 2 lines of code. That will swap your current LLM client out with the Hindsight wrapper. After that, memories will be stored and retrieved automatically as you make LLM calls.
|
||||
|
||||
If you need more control over how and when your agent stores and recalls memories, there's also a simple API you can integrate with using the SDKs or directly via HTTP.
|
||||
|
||||

|
||||
|
||||
---
|
||||
|
||||
> 🤖 **Using a coding agent?** Install the Hindsight documentation skill for instant access to docs while you code:
|
||||
> ```bash
|
||||
> npx skills add https://github.com/vectorize-io/hindsight --skill hindsight-docs
|
||||
> ```
|
||||
> Works with Claude Code, Cursor, and other AI coding assistants.
|
||||
|
||||
---
|
||||
|
||||
The benchmark results from this research can be inspected in our [visual benchmark explorer](https://hindsight-benchmarks.vercel.app). As additional improvements are made to Hindsight, new benchmark data will be available for review using this same tool.
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Docker (recommended)
|
||||
|
||||
```bash
|
||||
export OPENAI_API_KEY=your-key
|
||||
export OPENAI_API_KEY=sk-xxx
|
||||
|
||||
docker run --rm -it --pull always -p 8888:8888 -p 9999:9999 \
|
||||
-e HINDSIGHT_API_LLM_API_KEY=$OPENAI_API_KEY \
|
||||
-e HINDSIGHT_API_LLM_MODEL=o3-mini \
|
||||
-v $HOME/.hindsight-docker:/home/hindsight/.pg0 \
|
||||
ghcr.io/vectorize-io/hindsight:latest
|
||||
```
|
||||
|
||||
>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`, and `lmstudio`. The documentation provides more details on [supported models](https://hindsight.vectorize.io/developer/models).
|
||||
|
||||
API: http://localhost:8888
|
||||
UI: http://localhost:9999
|
||||
|
||||
Install client:
|
||||
|
||||
### Docker (external PostgreSQL)
|
||||
|
||||
```bash
|
||||
export OPENAI_API_KEY=sk-xxx
|
||||
export HINDSIGHT_DB_PASSWORD=choose-a-password
|
||||
cd docker/docker-compose
|
||||
docker compose up
|
||||
```
|
||||
|
||||
|
||||
>API: http://localhost:8888
|
||||
>UI: http://localhost:9999
|
||||
|
||||
### Client
|
||||
|
||||
```bash
|
||||
pip install hindsight-client -U
|
||||
@@ -94,7 +94,7 @@ pip install hindsight-client -U
|
||||
npm install @vectorize-io/hindsight-client
|
||||
```
|
||||
|
||||
Python example:
|
||||
#### Python
|
||||
|
||||
```python
|
||||
from hindsight_client import Hindsight
|
||||
@@ -111,7 +111,29 @@ client.recall(bank_id="my-bank", query="What does Alice do?")
|
||||
client.reflect(bank_id="my-bank", query="Tell me about Alice")
|
||||
```
|
||||
|
||||
### Python (embedded, no Docker)
|
||||
#### Node.js / TypeScript
|
||||
|
||||
```bash
|
||||
npm install @vectorize-io/hindsight-client
|
||||
```
|
||||
|
||||
```javascript
|
||||
const { HindsightClient } = require('@vectorize-io/hindsight-client');
|
||||
|
||||
const main = async () => {
|
||||
const client = new HindsightClient({ baseUrl: 'http://localhost:8888' });
|
||||
|
||||
await client.retain('my-bank', 'Alice loves hiking in Yosemite');
|
||||
|
||||
const results = await client.recall('my-bank', 'What does Alice like?');
|
||||
console.log(results);
|
||||
}
|
||||
|
||||
main();
|
||||
```
|
||||
|
||||
|
||||
### Python Embedded (no server required)
|
||||
|
||||
```bash
|
||||
pip install hindsight-all -U
|
||||
@@ -131,25 +153,48 @@ with HindsightServer(
|
||||
results = client.recall(bank_id="my-bank", query="Where does Alice work?")
|
||||
```
|
||||
|
||||
### Node.js / TypeScript
|
||||
|
||||
```bash
|
||||
npm install @vectorize-io/hindsight-client
|
||||
```
|
||||
---
|
||||
|
||||
```javascript
|
||||
const { HindsightClient } = require('@vectorize-io/hindsight-client');
|
||||
## Use Cases
|
||||
|
||||
const client = new HindsightClient({ baseUrl: 'http://localhost:8888' });
|
||||
|
||||
await client.retain('my-bank', 'Alice loves hiking in Yosemite');
|
||||
await client.recall('my-bank', 'What does Alice like?');
|
||||
```
|
||||
Hindsight is built to support conversational AI agents as well as agents that are intended to perform tasks autonomously. The ideal use case for Hindsight are agents that require a blend of these features such as AI employees that need to handle open-ended tasks, change behavior based on user feedback, and learn to perform complex tasks to automate work at a level that approximates a human work. Hindsight can be used with simple AI workflows like those built with n8n and other similar tools, but may be overkill for such applications.
|
||||
|
||||
### Per-User Memories and Chat History
|
||||
|
||||
One of the simpler use cases you can use Hindsight for is to personalize AI chatbots and other conversational agents by storing and recalling memories associated with individual users.
|
||||
|
||||
The requirements for this use case usually look something like this:
|
||||
|
||||

|
||||
|
||||
<video src="https://github.com/user-attachments/assets/4805e8e1-e7d1-47c6-a4f8-2344a5ec8906" controls></video>
|
||||
|
||||
Satisfying these requirements in Hindsight is straightforward. When new user inputs and tool calls are ingested into Hindsight using the retain operation, custom metadata can be used to enrich the new memories. Metadata provides a convenient way to isolate memories that need to be restricted to a given user. Once these are fed into the retain operation, any raw memories and mental models that get created can be filtered when retrieving relevant memories.
|
||||
|
||||

|
||||
|
||||
---
|
||||
|
||||
## Architecture & Operations
|
||||
|
||||

|
||||
|
||||
Most agent memory implementation rely on basic vector search or sometimes use a knowledge graph. Hindsight uses biomimetic data structures to organize agent memories in a way that is more like how human memory works:
|
||||
|
||||
- **World:** Facts about the world ("The stove gets hot")
|
||||
- **Experiences:** Agent's own experiences ("I touched the stove and it really hurt")
|
||||
- **Mental Models:** Learned understanding of the agent's world formed by reflecting on raw memories and experiences.
|
||||
|
||||
Memories in Hindsight are stored in banks (i.e. memory banks). When memories are added to Hindsight, they are pushed into either the world facts or experiences memory pathway. They are then represented as a combination of entities, relationships, and time series with sparse/dense vector representations to aid in later recall.
|
||||
|
||||
Hindsight provides three simple methods to interact with the system:
|
||||
|
||||
- **Retain:** Provide information to Hindsight that you want it to remember
|
||||
- **Recall:** Retrieve memories from Hindsight
|
||||
- **Reflect:** Reflect on memories and experiences to generate new observations and insights from existing memories.
|
||||
|
||||
### Retain
|
||||
|
||||
The `retain` operation is used to push new memories into Hindsight. It tells Hindsight to _retain_ the information you pass in as an input.
|
||||
@@ -208,7 +253,7 @@ The final output is trimmed as needed to fit within the token limit.
|
||||
|
||||
### Reflect
|
||||
|
||||
The reflect operation is used to perform a more thorough analysis of existing memories. This allows the agent to form new connections between memories which are then persisted as opinions and/or observations. When building agents, the reflect operation is a key capability to enable the agent to learn from its experiences.
|
||||
The reflect operation is used to perform a more thorough analysis of existing memories. This allows the agent to form new connections between memories and build a more thorough understanding of its world.
|
||||
|
||||
For example, the `reflect` operation can be used to support use cases such as:
|
||||
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
# Docker Compose file for Hindsight with PostgreSQL and pgvector
|
||||
#
|
||||
# Make sure to set the required environment variables before running:
|
||||
# - HINDSIGHT_DB_PASSWORD: Password for the PostgreSQL user
|
||||
# - Configure LLM provider variables as needed (see below in the hindsight service)
|
||||
#
|
||||
# Usage:
|
||||
# docker compose up -d
|
||||
#
|
||||
# Optional environment variables with defaults:
|
||||
# - HINDSIGHT_VERSION: Hindsight application version (default: latest)
|
||||
# - HINDSIGHT_DB_USER: PostgreSQL user (default: hindsight_user)
|
||||
# - HINDSIGHT_DB_NAME: PostgreSQL database name (default: hindsight_db)
|
||||
# - HINDSIGHT_DB_VERSION: PostgreSQL version (default: 18)
|
||||
|
||||
services:
|
||||
db:
|
||||
# Use a PostgreSQL-Image with pgvector extension pre-installed
|
||||
# see https://hub.docker.com/r/pgvector/pgvector
|
||||
image: pgvector/pgvector:pg${HINDSIGHT_DB_VERSION:-18}
|
||||
container_name: hindsight-db
|
||||
restart: always
|
||||
# Expose PostgreSQL port
|
||||
# ports:
|
||||
# - "5432:5432"
|
||||
environment:
|
||||
POSTGRES_USER: ${HINDSIGHT_DB_USER:-hindsight_user}
|
||||
POSTGRES_PASSWORD: ${HINDSIGHT_DB_PASSWORD:?Please set the HINDSIGHT_DB_PASSWORD env variable}
|
||||
POSTGRES_DB: ${HINDSIGHT_DB_NAME:-hindsight_db}
|
||||
volumes:
|
||||
- pg_data:/var/lib/postgresql/${HINDSIGHT_DB_VERSION:-18}/docker
|
||||
networks:
|
||||
- hindsight-net
|
||||
|
||||
hindsight:
|
||||
image: ghcr.io/vectorize-io/hindsight:${HINDSIGHT_VERSION:-latest}
|
||||
container_name: hindsight-app
|
||||
ports:
|
||||
- "8888:8888"
|
||||
- "9999:9999"
|
||||
environment:
|
||||
- HINDSIGHT_API_LLM_API_KEY=${OPENAI_API_KEY?Please set the OPENAI_API_KEY env variable}
|
||||
- HINDSIGHT_API_DATABASE_URL=postgresql://${HINDSIGHT_DB_USER:-hindsight_user}:${HINDSIGHT_DB_PASSWORD:?Please set the HINDSIGHT_DB_PASSWORD env variable}@db:5432/${HINDSIGHT_DB_NAME:-hindsight_db}
|
||||
depends_on:
|
||||
- db
|
||||
networks:
|
||||
- hindsight-net
|
||||
|
||||
networks:
|
||||
hindsight-net:
|
||||
driver: bridge
|
||||
|
||||
volumes:
|
||||
pg_data:
|
||||
@@ -169,16 +169,34 @@ ENV PATH="/app/api/.venv/bin:${PATH}"
|
||||
|
||||
# Pre-download ML models to avoid runtime download (conditional)
|
||||
# Only runs if both PRELOAD_ML_MODELS=true AND INCLUDE_LOCAL_MODELS=true
|
||||
# Includes retry logic with exponential backoff for transient network failures
|
||||
ARG PRELOAD_ML_MODELS
|
||||
ARG INCLUDE_LOCAL_MODELS
|
||||
ENV HF_HUB_DOWNLOAD_TIMEOUT=600
|
||||
RUN if [ "$PRELOAD_ML_MODELS" = "true" ] && [ "$INCLUDE_LOCAL_MODELS" = "true" ]; then \
|
||||
/app/api/.venv/bin/python -c "\
|
||||
MAX_RETRIES=3; \
|
||||
RETRY_DELAY=10; \
|
||||
for i in $(seq 1 $MAX_RETRIES); do \
|
||||
echo "Attempt $i/$MAX_RETRIES: Downloading ML models..."; \
|
||||
/app/api/.venv/bin/python -c "\
|
||||
import os; os.environ['HF_HUB_DOWNLOAD_TIMEOUT'] = '600'; \
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder; \
|
||||
print('Downloading embedding model...'); \
|
||||
SentenceTransformer('BAAI/bge-small-en-v1.5'); \
|
||||
print('Downloading cross-encoder model...'); \
|
||||
CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2'); \
|
||||
print('Models cached successfully')"; \
|
||||
print('Downloading tiktoken encoding...'); import tiktoken; tiktoken.get_encoding('cl100k_base'); \
|
||||
print('Models cached successfully')" && break; \
|
||||
if [ $i -lt $MAX_RETRIES ]; then \
|
||||
echo "Attempt $i failed, retrying in ${RETRY_DELAY}s..."; \
|
||||
sleep $RETRY_DELAY; \
|
||||
RETRY_DELAY=$((RETRY_DELAY * 2)); \
|
||||
fi; \
|
||||
done; \
|
||||
if [ $i -eq $MAX_RETRIES ] && ! /app/api/.venv/bin/python -c "from sentence_transformers import SentenceTransformer; SentenceTransformer('BAAI/bge-small-en-v1.5')" 2>/dev/null; then \
|
||||
echo "ERROR: Failed to download models after $MAX_RETRIES attempts"; \
|
||||
exit 1; \
|
||||
fi; \
|
||||
elif [ "$INCLUDE_LOCAL_MODELS" != "true" ]; then echo "Skipping ML model preload (local-models not included)"; \
|
||||
else echo "Skipping ML model preload"; fi
|
||||
|
||||
@@ -190,6 +208,10 @@ ENV HINDSIGHT_API_LOG_LEVEL=info
|
||||
ENV HINDSIGHT_ENABLE_API=true
|
||||
ENV HINDSIGHT_ENABLE_CP=false
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
# Suppress verbose transformers/HuggingFace model loading warnings
|
||||
ENV TRANSFORMERS_VERBOSITY=error
|
||||
ENV HF_HUB_VERBOSITY=error
|
||||
ENV TOKENIZERS_PARALLELISM=false
|
||||
|
||||
CMD ["/app/start-all.sh"]
|
||||
|
||||
@@ -277,16 +299,34 @@ ENV PATH="/app/api/.venv/bin:${PATH}"
|
||||
|
||||
# Pre-download ML models to avoid runtime download (conditional)
|
||||
# Only runs if both PRELOAD_ML_MODELS=true AND INCLUDE_LOCAL_MODELS=true
|
||||
# Includes retry logic with exponential backoff for transient network failures
|
||||
ARG PRELOAD_ML_MODELS
|
||||
ARG INCLUDE_LOCAL_MODELS
|
||||
ENV HF_HUB_DOWNLOAD_TIMEOUT=600
|
||||
RUN if [ "$PRELOAD_ML_MODELS" = "true" ] && [ "$INCLUDE_LOCAL_MODELS" = "true" ]; then \
|
||||
/app/api/.venv/bin/python -c "\
|
||||
MAX_RETRIES=3; \
|
||||
RETRY_DELAY=10; \
|
||||
for i in $(seq 1 $MAX_RETRIES); do \
|
||||
echo "Attempt $i/$MAX_RETRIES: Downloading ML models..."; \
|
||||
/app/api/.venv/bin/python -c "\
|
||||
import os; os.environ['HF_HUB_DOWNLOAD_TIMEOUT'] = '600'; \
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder; \
|
||||
print('Downloading embedding model...'); \
|
||||
SentenceTransformer('BAAI/bge-small-en-v1.5'); \
|
||||
print('Downloading cross-encoder model...'); \
|
||||
CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2'); \
|
||||
print('Models cached successfully')"; \
|
||||
print('Downloading tiktoken encoding...'); import tiktoken; tiktoken.get_encoding('cl100k_base'); \
|
||||
print('Models cached successfully')" && break; \
|
||||
if [ $i -lt $MAX_RETRIES ]; then \
|
||||
echo "Attempt $i failed, retrying in ${RETRY_DELAY}s..."; \
|
||||
sleep $RETRY_DELAY; \
|
||||
RETRY_DELAY=$((RETRY_DELAY * 2)); \
|
||||
fi; \
|
||||
done; \
|
||||
if [ $i -eq $MAX_RETRIES ] && ! /app/api/.venv/bin/python -c "from sentence_transformers import SentenceTransformer; SentenceTransformer('BAAI/bge-small-en-v1.5')" 2>/dev/null; then \
|
||||
echo "ERROR: Failed to download models after $MAX_RETRIES attempts"; \
|
||||
exit 1; \
|
||||
fi; \
|
||||
elif [ "$INCLUDE_LOCAL_MODELS" != "true" ]; then echo "Skipping ML model preload (local-models not included)"; \
|
||||
else echo "Skipping ML model preload"; fi
|
||||
|
||||
@@ -300,6 +340,10 @@ ENV HINDSIGHT_CP_DATAPLANE_API_URL=http://localhost:8888
|
||||
ENV HINDSIGHT_ENABLE_API=true
|
||||
ENV HINDSIGHT_ENABLE_CP=true
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
# Suppress verbose transformers/HuggingFace model loading warnings
|
||||
ENV TRANSFORMERS_VERBOSITY=error
|
||||
ENV HF_HUB_VERBOSITY=error
|
||||
ENV TOKENIZERS_PARALLELISM=false
|
||||
|
||||
CMD ["/app/start-all.sh"]
|
||||
|
||||
|
||||
@@ -6,28 +6,40 @@
|
||||
# Can be run locally or in CI pipelines.
|
||||
#
|
||||
# Usage:
|
||||
# ./scripts/docker-smoke-test.sh <image> [target]
|
||||
# ./docker/test-image.sh <image> [target]
|
||||
#
|
||||
# Arguments:
|
||||
# image - Docker image to test (e.g., hindsight-api:test, ghcr.io/vectorize-io/hindsight:latest)
|
||||
# target - Optional: 'cp-only' for control plane, otherwise assumes API image (default: api)
|
||||
#
|
||||
# Environment variables:
|
||||
# GROQ_API_KEY - Required for API/standalone images (LLM verification)
|
||||
# HINDSIGHT_API_LLM_PROVIDER - LLM provider (default: groq)
|
||||
# HINDSIGHT_API_LLM_MODEL - LLM model (default: llama-3.3-70b-versatile)
|
||||
# SMOKE_TEST_TIMEOUT - Timeout in seconds (default: 120)
|
||||
# SMOKE_TEST_CONTAINER_NAME - Container name (default: hindsight-smoke-test)
|
||||
# GROQ_API_KEY - Required for API/standalone images (LLM verification)
|
||||
# HINDSIGHT_API_LLM_PROVIDER - LLM provider (default: groq)
|
||||
# HINDSIGHT_API_LLM_MODEL - LLM model (default: llama-3.3-70b-versatile)
|
||||
# HINDSIGHT_API_EMBEDDINGS_PROVIDER - Embeddings provider (optional, for slim images: openai, cohere, tei)
|
||||
# HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY - OpenAI API key for embeddings (optional)
|
||||
# HINDSIGHT_API_RERANKER_PROVIDER - Reranker provider (optional, for slim images: cohere, tei)
|
||||
# HINDSIGHT_API_COHERE_API_KEY - Cohere API key for reranking (optional)
|
||||
# SMOKE_TEST_TIMEOUT - Timeout in seconds (default: 120)
|
||||
# SMOKE_TEST_CONTAINER_NAME - Container name (default: hindsight-smoke-test)
|
||||
#
|
||||
# Examples:
|
||||
# # Test a locally built image
|
||||
# ./scripts/docker-smoke-test.sh hindsight-api:test
|
||||
# # Test a locally built full image
|
||||
# ./docker/test-image.sh hindsight-api:test
|
||||
#
|
||||
# # Test a released image
|
||||
# ./scripts/docker-smoke-test.sh ghcr.io/vectorize-io/hindsight:latest
|
||||
# ./docker/test-image.sh ghcr.io/vectorize-io/hindsight:latest
|
||||
#
|
||||
# # Test control plane image
|
||||
# ./scripts/docker-smoke-test.sh hindsight-control-plane:test cp-only
|
||||
# ./docker/test-image.sh hindsight-control-plane:test cp-only
|
||||
#
|
||||
# # Test slim image with external providers
|
||||
# export GROQ_API_KEY=gsk_xxx
|
||||
# export HINDSIGHT_API_EMBEDDINGS_PROVIDER=openai
|
||||
# export HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=sk-xxx
|
||||
# export HINDSIGHT_API_RERANKER_PROVIDER=cohere
|
||||
# export HINDSIGHT_API_COHERE_API_KEY=xxx
|
||||
# ./docker/test-image.sh hindsight-slim:test
|
||||
#
|
||||
# Exit codes:
|
||||
# 0 - Success (container healthy)
|
||||
@@ -108,12 +120,32 @@ if [ "$TARGET" = "cp-only" ]; then
|
||||
-p "${HEALTH_PORT}:${HEALTH_PORT}" \
|
||||
"$IMAGE"
|
||||
else
|
||||
docker run -d --name "$CONTAINER_NAME" \
|
||||
-e HINDSIGHT_API_LLM_PROVIDER="$LLM_PROVIDER" \
|
||||
-e HINDSIGHT_API_LLM_API_KEY="${GROQ_API_KEY}" \
|
||||
-e HINDSIGHT_API_LLM_MODEL="$LLM_MODEL" \
|
||||
-p "${HEALTH_PORT}:${HEALTH_PORT}" \
|
||||
"$IMAGE"
|
||||
# Build docker run command with required and optional env vars
|
||||
DOCKER_CMD="docker run -d --name $CONTAINER_NAME"
|
||||
DOCKER_CMD="$DOCKER_CMD -e HINDSIGHT_API_LLM_PROVIDER=$LLM_PROVIDER"
|
||||
DOCKER_CMD="$DOCKER_CMD -e HINDSIGHT_API_LLM_API_KEY=${GROQ_API_KEY}"
|
||||
DOCKER_CMD="$DOCKER_CMD -e HINDSIGHT_API_LLM_MODEL=$LLM_MODEL"
|
||||
|
||||
# Add optional embeddings provider config
|
||||
if [ -n "${HINDSIGHT_API_EMBEDDINGS_PROVIDER:-}" ]; then
|
||||
DOCKER_CMD="$DOCKER_CMD -e HINDSIGHT_API_EMBEDDINGS_PROVIDER=${HINDSIGHT_API_EMBEDDINGS_PROVIDER}"
|
||||
fi
|
||||
if [ -n "${HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY:-}" ]; then
|
||||
DOCKER_CMD="$DOCKER_CMD -e HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=${HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY}"
|
||||
fi
|
||||
|
||||
# Add optional reranker provider config
|
||||
if [ -n "${HINDSIGHT_API_RERANKER_PROVIDER:-}" ]; then
|
||||
DOCKER_CMD="$DOCKER_CMD -e HINDSIGHT_API_RERANKER_PROVIDER=${HINDSIGHT_API_RERANKER_PROVIDER}"
|
||||
fi
|
||||
if [ -n "${HINDSIGHT_API_COHERE_API_KEY:-}" ]; then
|
||||
DOCKER_CMD="$DOCKER_CMD -e HINDSIGHT_API_COHERE_API_KEY=${HINDSIGHT_API_COHERE_API_KEY}"
|
||||
fi
|
||||
|
||||
DOCKER_CMD="$DOCKER_CMD -p ${HEALTH_PORT}:${HEALTH_PORT}"
|
||||
DOCKER_CMD="$DOCKER_CMD $IMAGE"
|
||||
|
||||
eval $DOCKER_CMD
|
||||
fi
|
||||
|
||||
# Wait for health endpoint
|
||||
Executable
+51
@@ -0,0 +1,51 @@
|
||||
#!/bin/bash
|
||||
#
|
||||
# Local Test Script for Slim Docker Images
|
||||
#
|
||||
# This script makes it easy to test slim images locally with external providers.
|
||||
# It expects API keys to be set in environment variables.
|
||||
#
|
||||
# Usage:
|
||||
# export GROQ_API_KEY=gsk_xxx
|
||||
# export OPENAI_API_KEY=sk-xxx
|
||||
# export COHERE_API_KEY=xxx
|
||||
# ./docker/test-slim-local.sh
|
||||
#
|
||||
# Or inline:
|
||||
# GROQ_API_KEY=gsk_xxx OPENAI_API_KEY=sk_xxx COHERE_API_KEY=xxx ./docker/test-slim-local.sh
|
||||
#
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# Check for required API keys
|
||||
if [ -z "${GROQ_API_KEY:-}" ]; then
|
||||
echo "❌ Error: GROQ_API_KEY environment variable is required"
|
||||
echo "Set it with: export GROQ_API_KEY=gsk_xxx"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
if [ -z "${OPENAI_API_KEY:-}" ]; then
|
||||
echo "❌ Error: OPENAI_API_KEY environment variable is required"
|
||||
echo "Set it with: export OPENAI_API_KEY=sk-xxx"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
if [ -z "${COHERE_API_KEY:-}" ]; then
|
||||
echo "❌ Error: COHERE_API_KEY environment variable is required"
|
||||
echo "Set it with: export COHERE_API_KEY=xxx"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Configuration
|
||||
IMAGE="${1:-hindsight-slim:test}"
|
||||
echo "Testing image: $IMAGE"
|
||||
echo ""
|
||||
|
||||
# Set up external providers
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=openai
|
||||
export HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=$OPENAI_API_KEY
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=$COHERE_API_KEY
|
||||
|
||||
# Run the test
|
||||
exec "$(dirname "$0")/test-image.sh" "$IMAGE" standalone
|
||||
@@ -2,8 +2,8 @@ apiVersion: v2
|
||||
name: hindsight
|
||||
description: Hindsight helm chart
|
||||
type: application
|
||||
version: 0.3.0
|
||||
appVersion: "0.3.0"
|
||||
version: 0.4.10
|
||||
appVersion: "0.4.10"
|
||||
keywords:
|
||||
- ai
|
||||
- memory
|
||||
|
||||
@@ -80,6 +80,22 @@ Control plane selector labels
|
||||
app.kubernetes.io/component: control-plane
|
||||
{{- end }}
|
||||
|
||||
{{/*
|
||||
Worker labels
|
||||
*/}}
|
||||
{{- define "hindsight.worker.labels" -}}
|
||||
{{ include "hindsight.labels" . }}
|
||||
app.kubernetes.io/component: worker
|
||||
{{- end }}
|
||||
|
||||
{{/*
|
||||
Worker selector labels
|
||||
*/}}
|
||||
{{- define "hindsight.worker.selectorLabels" -}}
|
||||
{{ include "hindsight.selectorLabels" . }}
|
||||
app.kubernetes.io/component: worker
|
||||
{{- end }}
|
||||
|
||||
{{/*
|
||||
Create the name of the service account to use
|
||||
*/}}
|
||||
|
||||
@@ -33,7 +33,7 @@ spec:
|
||||
- name: api
|
||||
securityContext:
|
||||
{{- toYaml .Values.securityContext | nindent 10 }}
|
||||
image: "{{ .Values.api.image.repository }}:{{ .Values.api.image.tag | default .Values.version }}"
|
||||
image: "{{ .Values.api.image.repository }}:{{ .Values.api.image.tag | default .Values.version | default .Chart.AppVersion }}"
|
||||
imagePullPolicy: {{ .Values.api.image.pullPolicy }}
|
||||
ports:
|
||||
- name: http
|
||||
@@ -55,6 +55,14 @@ spec:
|
||||
{{- end }}
|
||||
- name: HINDSIGHT_API_DATABASE_URL
|
||||
value: {{ include "hindsight.databaseUrl" . | quote }}
|
||||
{{- /* Disable internal worker when dedicated workers are enabled */}}
|
||||
{{- if .Values.worker.enabled }}
|
||||
- name: HINDSIGHT_API_WORKER_ENABLED
|
||||
value: "false"
|
||||
{{- end }}
|
||||
{{- /* Explicitly set port to override K8s service discovery env var (HINDSIGHT_API_PORT) */}}
|
||||
- name: HINDSIGHT_API_PORT
|
||||
value: {{ .Values.api.service.targetPort | quote }}
|
||||
{{- range $key, $value := .Values.api.env }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
@@ -79,7 +87,7 @@ spec:
|
||||
nodeSelector:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.affinity }}
|
||||
{{- with (.Values.api.affinity | default .Values.affinity) }}
|
||||
affinity:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
|
||||
@@ -33,7 +33,7 @@ spec:
|
||||
- name: control-plane
|
||||
securityContext:
|
||||
{{- toYaml .Values.securityContext | nindent 10 }}
|
||||
image: "{{ .Values.controlPlane.image.repository }}:{{ .Values.controlPlane.image.tag | default .Values.version }}"
|
||||
image: "{{ .Values.controlPlane.image.repository }}:{{ .Values.controlPlane.image.tag | default .Values.version | default .Chart.AppVersion }}"
|
||||
imagePullPolicy: {{ .Values.controlPlane.image.pullPolicy }}
|
||||
ports:
|
||||
- name: http
|
||||
@@ -71,7 +71,7 @@ spec:
|
||||
nodeSelector:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.affinity }}
|
||||
{{- with (.Values.controlPlane.affinity | default .Values.affinity) }}
|
||||
affinity:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
{{- if and .Values.api.enabled .Values.api.podDisruptionBudget.enabled }}
|
||||
apiVersion: policy/v1
|
||||
kind: PodDisruptionBudget
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-api
|
||||
labels:
|
||||
{{- include "hindsight.api.labels" . | nindent 4 }}
|
||||
spec:
|
||||
{{- if .Values.api.podDisruptionBudget.minAvailable }}
|
||||
minAvailable: {{ .Values.api.podDisruptionBudget.minAvailable }}
|
||||
{{- end }}
|
||||
{{- if .Values.api.podDisruptionBudget.maxUnavailable }}
|
||||
maxUnavailable: {{ .Values.api.podDisruptionBudget.maxUnavailable }}
|
||||
{{- end }}
|
||||
selector:
|
||||
matchLabels:
|
||||
{{- include "hindsight.api.selectorLabels" . | nindent 6 }}
|
||||
{{- end }}
|
||||
---
|
||||
{{- if and .Values.controlPlane.enabled .Values.controlPlane.podDisruptionBudget.enabled }}
|
||||
apiVersion: policy/v1
|
||||
kind: PodDisruptionBudget
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-control-plane
|
||||
labels:
|
||||
{{- include "hindsight.controlPlane.labels" . | nindent 4 }}
|
||||
spec:
|
||||
{{- if .Values.controlPlane.podDisruptionBudget.minAvailable }}
|
||||
minAvailable: {{ .Values.controlPlane.podDisruptionBudget.minAvailable }}
|
||||
{{- end }}
|
||||
{{- if .Values.controlPlane.podDisruptionBudget.maxUnavailable }}
|
||||
maxUnavailable: {{ .Values.controlPlane.podDisruptionBudget.maxUnavailable }}
|
||||
{{- end }}
|
||||
selector:
|
||||
matchLabels:
|
||||
{{- include "hindsight.controlPlane.selectorLabels" . | nindent 6 }}
|
||||
{{- end }}
|
||||
---
|
||||
{{- if and .Values.worker.enabled .Values.worker.podDisruptionBudget.enabled }}
|
||||
apiVersion: policy/v1
|
||||
kind: PodDisruptionBudget
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-worker
|
||||
labels:
|
||||
{{- include "hindsight.worker.labels" . | nindent 4 }}
|
||||
spec:
|
||||
{{- if .Values.worker.podDisruptionBudget.minAvailable }}
|
||||
minAvailable: {{ .Values.worker.podDisruptionBudget.minAvailable }}
|
||||
{{- end }}
|
||||
{{- if .Values.worker.podDisruptionBudget.maxUnavailable }}
|
||||
maxUnavailable: {{ .Values.worker.podDisruptionBudget.maxUnavailable }}
|
||||
{{- end }}
|
||||
selector:
|
||||
matchLabels:
|
||||
{{- include "hindsight.worker.selectorLabels" . | nindent 6 }}
|
||||
{{- end }}
|
||||
@@ -0,0 +1,25 @@
|
||||
{{- if .Values.worker.enabled }}
|
||||
apiVersion: v1
|
||||
kind: Service
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-worker
|
||||
labels:
|
||||
{{- include "hindsight.worker.labels" . | nindent 4 }}
|
||||
{{- if .Values.podAnnotations }}
|
||||
annotations:
|
||||
{{- /* Common Prometheus annotations for metrics scraping */}}
|
||||
prometheus.io/scrape: "true"
|
||||
prometheus.io/port: {{ .Values.worker.service.port | quote }}
|
||||
prometheus.io/path: "/metrics"
|
||||
{{- end }}
|
||||
spec:
|
||||
# Headless service for StatefulSet (enables stable DNS names like worker-0.worker.namespace)
|
||||
clusterIP: None
|
||||
ports:
|
||||
- port: {{ .Values.worker.service.port }}
|
||||
targetPort: {{ .Values.worker.service.targetPort }}
|
||||
protocol: TCP
|
||||
name: http
|
||||
selector:
|
||||
{{- include "hindsight.worker.selectorLabels" . | nindent 4 }}
|
||||
{{- end }}
|
||||
@@ -0,0 +1,110 @@
|
||||
{{- if .Values.worker.enabled }}
|
||||
apiVersion: apps/v1
|
||||
kind: StatefulSet
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-worker
|
||||
labels:
|
||||
{{- include "hindsight.worker.labels" . | nindent 4 }}
|
||||
spec:
|
||||
serviceName: {{ include "hindsight.fullname" . }}-worker
|
||||
replicas: {{ .Values.worker.replicaCount }}
|
||||
selector:
|
||||
matchLabels:
|
||||
{{- include "hindsight.worker.selectorLabels" . | nindent 6 }}
|
||||
template:
|
||||
metadata:
|
||||
annotations:
|
||||
{{- if not .Values.existingSecret }}
|
||||
checksum/secret: {{ include (print $.Template.BasePath "/secret.yaml") . | sha256sum }}
|
||||
{{- end }}
|
||||
{{- with .Values.podAnnotations }}
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
labels:
|
||||
{{- include "hindsight.worker.selectorLabels" . | nindent 8 }}
|
||||
spec:
|
||||
{{- if .Values.serviceAccount.create }}
|
||||
serviceAccountName: {{ include "hindsight.serviceAccountName" . }}
|
||||
{{- end }}
|
||||
securityContext:
|
||||
{{- toYaml .Values.podSecurityContext | nindent 8 }}
|
||||
containers:
|
||||
- name: worker
|
||||
securityContext:
|
||||
{{- toYaml .Values.securityContext | nindent 10 }}
|
||||
image: "{{ .Values.worker.image.repository }}:{{ .Values.worker.image.tag | default .Values.version | default .Chart.AppVersion }}"
|
||||
imagePullPolicy: {{ .Values.worker.image.pullPolicy }}
|
||||
command: ["hindsight-worker"]
|
||||
ports:
|
||||
- name: http
|
||||
containerPort: {{ .Values.worker.service.targetPort }}
|
||||
protocol: TCP
|
||||
{{- if .Values.existingSecret }}
|
||||
envFrom:
|
||||
- secretRef:
|
||||
name: {{ .Values.existingSecret }}
|
||||
{{- end }}
|
||||
env:
|
||||
{{- /* POSTGRES_PASSWORD must be defined before DATABASE_URL for $(VAR) interpolation */}}
|
||||
{{- if not .Values.postgresql.enabled }}
|
||||
- name: POSTGRES_PASSWORD
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: {{ include "hindsight.secretName" . }}
|
||||
key: postgres-password
|
||||
{{- end }}
|
||||
- name: HINDSIGHT_API_DATABASE_URL
|
||||
value: {{ include "hindsight.databaseUrl" . | quote }}
|
||||
{{- /* Worker ID uses pod name (StatefulSet provides stable names like worker-0, worker-1) */}}
|
||||
- name: HINDSIGHT_API_WORKER_ID
|
||||
valueFrom:
|
||||
fieldRef:
|
||||
fieldPath: metadata.name
|
||||
{{- /* Inherit LLM config from api.env */}}
|
||||
{{- range $key, $value := .Values.api.env }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
{{- /* Worker-specific env vars */}}
|
||||
{{- range $key, $value := .Values.worker.env }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
{{- /* Only use secrets when not using existingSecret */}}
|
||||
{{- if not .Values.existingSecret }}
|
||||
{{- /* Inherit secrets from api.secrets */}}
|
||||
{{- range $key, $value := .Values.api.secrets }}
|
||||
- name: {{ $key }}
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: {{ include "hindsight.secretName" $ }}
|
||||
key: {{ $key }}
|
||||
{{- end }}
|
||||
{{- /* Worker-specific secrets (can override api.secrets) */}}
|
||||
{{- range $key, $value := .Values.worker.secrets }}
|
||||
- name: {{ $key }}
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: {{ include "hindsight.secretName" $ }}
|
||||
key: {{ $key }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
livenessProbe:
|
||||
{{- toYaml .Values.worker.livenessProbe | nindent 10 }}
|
||||
readinessProbe:
|
||||
{{- toYaml .Values.worker.readinessProbe | nindent 10 }}
|
||||
resources:
|
||||
{{- toYaml .Values.worker.resources | nindent 10 }}
|
||||
{{- with .Values.nodeSelector }}
|
||||
nodeSelector:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with (.Values.worker.affinity | default .Values.affinity) }}
|
||||
affinity:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.tolerations }}
|
||||
tolerations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
@@ -1,7 +1,8 @@
|
||||
# Default values for hindsight
|
||||
|
||||
# Chart version - use this to set a consistent image tag across all components
|
||||
version: "0.1.1"
|
||||
# Global version override - use this to set a consistent image tag across all components
|
||||
# If not set, defaults to Chart.appVersion from Chart.yaml
|
||||
# version: ""
|
||||
|
||||
# Use an existing secret instead of creating one from values
|
||||
# When set, all keys from this secret are injected as environment variables via envFrom
|
||||
@@ -57,6 +58,15 @@ api:
|
||||
timeoutSeconds: 3
|
||||
failureThreshold: 3
|
||||
|
||||
# Pod disruption budget
|
||||
podDisruptionBudget:
|
||||
enabled: false
|
||||
minAvailable: 1
|
||||
# maxUnavailable: 1
|
||||
|
||||
# Pod affinity/anti-affinity (overrides global affinity for this component)
|
||||
# affinity: {}
|
||||
|
||||
# Environment variables
|
||||
env:
|
||||
#HINDSIGHT_API_LLM_PROVIDER: "groq"
|
||||
@@ -67,6 +77,72 @@ api:
|
||||
# HINDSIGHT_API_LLM_API_KEY: "your-api-key"
|
||||
# HINDSIGHT_API_LLM_BASE_URL: "https://api.groq.com/openai/v1"
|
||||
|
||||
# Worker settings (distributed task processing)
|
||||
# When enabled, dedicated worker pods process tasks and the API's internal worker is disabled
|
||||
worker:
|
||||
enabled: false
|
||||
replicaCount: 2
|
||||
image:
|
||||
repository: ghcr.io/vectorize-io/hindsight-api
|
||||
pullPolicy: IfNotPresent
|
||||
# tag: "" # defaults to .Values.version, then Chart.appVersion if not specified
|
||||
|
||||
service:
|
||||
# Service for metrics scraping (headless for StatefulSet)
|
||||
port: 8889
|
||||
targetPort: 8889
|
||||
|
||||
# Resource limits and requests
|
||||
resources:
|
||||
limits:
|
||||
cpu: 2000m
|
||||
memory: 4Gi
|
||||
requests:
|
||||
cpu: 500m
|
||||
memory: 1Gi
|
||||
|
||||
# Liveness and readiness probes
|
||||
livenessProbe:
|
||||
httpGet:
|
||||
path: /health
|
||||
port: 8889
|
||||
initialDelaySeconds: 30
|
||||
periodSeconds: 10
|
||||
timeoutSeconds: 5
|
||||
failureThreshold: 3
|
||||
|
||||
readinessProbe:
|
||||
httpGet:
|
||||
path: /health
|
||||
port: 8889
|
||||
initialDelaySeconds: 10
|
||||
periodSeconds: 5
|
||||
timeoutSeconds: 3
|
||||
failureThreshold: 3
|
||||
|
||||
# Worker-specific environment variables
|
||||
env:
|
||||
# Poll interval in milliseconds (how often to check for new tasks)
|
||||
HINDSIGHT_API_WORKER_POLL_INTERVAL_MS: "500"
|
||||
# Number of tasks to claim per poll cycle
|
||||
HINDSIGHT_API_WORKER_BATCH_SIZE: "10"
|
||||
# Max retries before marking a task as failed
|
||||
HINDSIGHT_API_WORKER_MAX_RETRIES: "3"
|
||||
# HTTP port for metrics/health (matches service.targetPort)
|
||||
HINDSIGHT_API_WORKER_HTTP_PORT: "8889"
|
||||
|
||||
# Pod disruption budget
|
||||
podDisruptionBudget:
|
||||
enabled: false
|
||||
minAvailable: 1
|
||||
# maxUnavailable: 1
|
||||
|
||||
# Pod affinity/anti-affinity (overrides global affinity for this component)
|
||||
# affinity: {}
|
||||
|
||||
# Secret environment variables (inherited from api.secrets if not specified)
|
||||
secrets: {}
|
||||
|
||||
# Image settings for control plane
|
||||
controlPlane:
|
||||
enabled: true
|
||||
@@ -107,6 +183,15 @@ controlPlane:
|
||||
timeoutSeconds: 3
|
||||
failureThreshold: 3
|
||||
|
||||
# Pod disruption budget
|
||||
podDisruptionBudget:
|
||||
enabled: false
|
||||
minAvailable: 1
|
||||
# maxUnavailable: 1
|
||||
|
||||
# Pod affinity/anti-affinity (overrides global affinity for this component)
|
||||
# affinity: {}
|
||||
|
||||
# Environment variables
|
||||
env:
|
||||
NODE_ENV: "production"
|
||||
@@ -205,7 +290,7 @@ nodeSelector: {}
|
||||
# Tolerations
|
||||
tolerations: []
|
||||
|
||||
# Affinity
|
||||
# Affinity (applied to all components unless overridden per-component)
|
||||
affinity: {}
|
||||
|
||||
# Autoscaling
|
||||
|
||||
@@ -46,4 +46,4 @@ __all__ = [
|
||||
"RemoteTEICrossEncoder",
|
||||
"LLMConfig",
|
||||
]
|
||||
__version__ = "0.1.0"
|
||||
__version__ = "0.4.10"
|
||||
|
||||
@@ -244,6 +244,65 @@ def run_db_migration(
|
||||
typer.echo("Database migrations completed successfully")
|
||||
|
||||
|
||||
async def _decommission_worker(db_url: str, worker_id: str, schema: str = "public") -> int:
|
||||
"""Release all tasks owned by a worker, setting them back to pending status."""
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
|
||||
conn = await asyncpg.connect(resolved_url)
|
||||
try:
|
||||
table = _fq_table("async_operations", schema)
|
||||
result = await conn.fetch(
|
||||
f"""
|
||||
UPDATE {table}
|
||||
SET status = 'pending', worker_id = NULL, claimed_at = NULL, updated_at = now()
|
||||
WHERE worker_id = $1 AND status = 'processing'
|
||||
RETURNING operation_id
|
||||
""",
|
||||
worker_id,
|
||||
)
|
||||
return len(result)
|
||||
finally:
|
||||
await conn.close()
|
||||
|
||||
|
||||
@app.command(name="decommission-worker")
|
||||
def decommission_worker(
|
||||
worker_id: str = typer.Argument(..., help="Worker ID to decommission"),
|
||||
schema: str = typer.Option("public", "--schema", "-s", help="Database schema"),
|
||||
yes: bool = typer.Option(False, "--yes", "-y", help="Skip confirmation prompt"),
|
||||
):
|
||||
"""Release all tasks owned by a worker (sets status back to pending).
|
||||
|
||||
Use this command when a worker has crashed or been removed without graceful shutdown.
|
||||
All tasks that were being processed by the worker will be released back to the queue
|
||||
so other workers can pick them up.
|
||||
"""
|
||||
config = HindsightConfig.from_env()
|
||||
|
||||
if not config.database_url:
|
||||
typer.echo("Error: Database URL not configured.", err=True)
|
||||
typer.echo("Set HINDSIGHT_API_DATABASE_URL environment variable.", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
if not yes:
|
||||
typer.confirm(
|
||||
f"This will release all tasks owned by worker '{worker_id}' back to pending. Continue?",
|
||||
abort=True,
|
||||
)
|
||||
|
||||
typer.echo(f"Decommissioning worker '{worker_id}' (schema: {schema})...")
|
||||
|
||||
count = asyncio.run(_decommission_worker(config.database_url, worker_id, schema))
|
||||
|
||||
if count > 0:
|
||||
typer.echo(f"Released {count} task(s) from worker '{worker_id}'")
|
||||
else:
|
||||
typer.echo(f"No tasks found for worker '{worker_id}'")
|
||||
|
||||
|
||||
def main():
|
||||
app()
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ from collections.abc import Sequence
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from pgvector.sqlalchemy import Vector
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
@@ -23,8 +24,21 @@ depends_on: str | Sequence[str] | None = None
|
||||
def upgrade() -> None:
|
||||
"""Upgrade schema - create all tables from scratch."""
|
||||
|
||||
# Enable required extensions
|
||||
op.execute("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
# Note: pgvector extension is installed globally BEFORE migrations run
|
||||
# See migrations.py:run_migrations() - this ensures the extension is available
|
||||
# to all schemas, not just the one being migrated
|
||||
|
||||
# We keep this here as a fallback for backwards compatibility
|
||||
# This may fail if user lacks permissions, which is fine if extension already exists
|
||||
try:
|
||||
op.execute("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
except Exception:
|
||||
# Extension might already exist or user lacks permissions - verify it exists
|
||||
conn = op.get_bind()
|
||||
result = conn.execute(text("SELECT 1 FROM pg_extension WHERE extname = 'vector'")).fetchone()
|
||||
if not result:
|
||||
# Extension truly doesn't exist - re-raise the error
|
||||
raise
|
||||
|
||||
# Create banks table
|
||||
op.create_table(
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
"""mental_models_v4
|
||||
|
||||
Revision ID: h3c4d5e6f7g8
|
||||
Revises: g2a3b4c5d6e7
|
||||
Create Date: 2026-01-08 00:00:00.000000
|
||||
|
||||
This migration implements the v4 mental models system:
|
||||
1. Deletes existing observation memory_units (observations now in mental models)
|
||||
2. Adds mission column to banks (replacing background)
|
||||
3. Creates mental_models table with final schema
|
||||
|
||||
Mental models can reference entities when an entity is "promoted" to a mental model.
|
||||
Summary content is stored as JSONB observations with per-observation fact attribution.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "h3c4d5e6f7g8"
|
||||
down_revision: str | Sequence[str] | None = "g2a3b4c5d6e7"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Apply mental models v4 changes."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Step 1: Delete observation memory_units (cascades to unit_entities links)
|
||||
# Observations are now handled through mental models, not memory_units
|
||||
op.execute(f"DELETE FROM {schema}memory_units WHERE fact_type = 'observation'")
|
||||
|
||||
# Step 2: Drop observation-specific index (if it exists)
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_observation_date")
|
||||
|
||||
# Step 3: Add mission column to banks (replacing background)
|
||||
op.execute(f"ALTER TABLE {schema}banks ADD COLUMN IF NOT EXISTS mission TEXT")
|
||||
|
||||
# Migrate: copy background to mission if background column exists
|
||||
# Use DO block to check column existence first (idempotent for re-runs)
|
||||
schema_name = context.config.get_main_option("target_schema") or "public"
|
||||
op.execute(f"""
|
||||
DO $$
|
||||
BEGIN
|
||||
IF EXISTS (
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_schema = '{schema_name}' AND table_name = 'banks' AND column_name = 'background'
|
||||
) THEN
|
||||
UPDATE {schema}banks
|
||||
SET mission = background
|
||||
WHERE mission IS NULL;
|
||||
END IF;
|
||||
END $$;
|
||||
""")
|
||||
|
||||
# Remove background column (replaced by mission)
|
||||
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS background")
|
||||
|
||||
# Step 4: Create mental_models table with final v4 schema (if not exists)
|
||||
op.execute(f"""
|
||||
CREATE TABLE IF NOT EXISTS {schema}mental_models (
|
||||
id VARCHAR(64) NOT NULL,
|
||||
bank_id VARCHAR(64) NOT NULL,
|
||||
subtype VARCHAR(32) NOT NULL,
|
||||
name VARCHAR(256) NOT NULL,
|
||||
description TEXT NOT NULL,
|
||||
entity_id UUID,
|
||||
observations JSONB DEFAULT '{{"observations": []}}'::jsonb,
|
||||
links VARCHAR[],
|
||||
tags VARCHAR[] DEFAULT '{{}}',
|
||||
last_updated TIMESTAMP WITH TIME ZONE,
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
|
||||
PRIMARY KEY (id, bank_id),
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE,
|
||||
FOREIGN KEY (entity_id) REFERENCES {schema}entities(id) ON DELETE SET NULL,
|
||||
CONSTRAINT ck_mental_models_subtype CHECK (subtype IN ('structural', 'emergent', 'pinned', 'learned'))
|
||||
)
|
||||
""")
|
||||
|
||||
# Step 5: Create indexes for efficient queries (if not exist)
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_bank_id ON {schema}mental_models(bank_id)")
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_subtype ON {schema}mental_models(bank_id, subtype)")
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_entity_id ON {schema}mental_models(entity_id)")
|
||||
# GIN index for efficient tags array filtering
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_tags ON {schema}mental_models USING GIN(tags)")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Revert mental models v4 changes."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop mental_models table (cascades to indexes)
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}mental_models CASCADE")
|
||||
|
||||
# Add back background column to banks
|
||||
op.execute(f"ALTER TABLE {schema}banks ADD COLUMN IF NOT EXISTS background TEXT")
|
||||
|
||||
# Migrate mission back to background
|
||||
op.execute(f"UPDATE {schema}banks SET background = mission WHERE background IS NULL")
|
||||
|
||||
# Remove mission column
|
||||
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS mission")
|
||||
|
||||
# Note: Cannot restore deleted observations - they are lost on downgrade
|
||||
@@ -0,0 +1,41 @@
|
||||
"""delete_opinions
|
||||
|
||||
Revision ID: i4d5e6f7g8h9
|
||||
Revises: h3c4d5e6f7g8
|
||||
Create Date: 2026-01-15 00:00:00.000000
|
||||
|
||||
This migration removes opinion facts from memory_units.
|
||||
Opinions are no longer a separate fact type - they are now represented
|
||||
through mental model observations with confidence scores.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "i4d5e6f7g8h9"
|
||||
down_revision: str | Sequence[str] | None = "h3c4d5e6f7g8"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Delete opinion memory_units."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Delete opinion memory_units (cascades to unit_entities links)
|
||||
# Opinions are now handled through mental model observations
|
||||
op.execute(f"DELETE FROM {schema}memory_units WHERE fact_type = 'opinion'")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Cannot restore deleted opinions."""
|
||||
# Note: Cannot restore deleted opinions - they are lost on downgrade
|
||||
pass
|
||||
@@ -0,0 +1,95 @@
|
||||
"""mental_model_versions
|
||||
|
||||
Revision ID: j5e6f7g8h9i0
|
||||
Revises: i4d5e6f7g8h9
|
||||
Create Date: 2026-01-16 00:00:00.000000
|
||||
|
||||
This migration adds versioning support for mental models:
|
||||
1. Creates mental_model_versions table to store observation snapshots
|
||||
2. Adds version column to mental_models for tracking current version
|
||||
|
||||
This enables changelog/diff functionality for mental model observations.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "j5e6f7g8h9i0"
|
||||
down_revision: str | Sequence[str] | None = "i4d5e6f7g8h9"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create mental_model_versions table and add version tracking."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Create mental_model_versions table for storing observation snapshots
|
||||
op.execute(f"""
|
||||
CREATE TABLE {schema}mental_model_versions (
|
||||
id SERIAL PRIMARY KEY,
|
||||
mental_model_id VARCHAR(64) NOT NULL,
|
||||
bank_id VARCHAR(64) NOT NULL,
|
||||
version INT NOT NULL,
|
||||
observations JSONB NOT NULL DEFAULT '{{"observations": []}}'::jsonb,
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
|
||||
FOREIGN KEY (mental_model_id, bank_id)
|
||||
REFERENCES {schema}mental_models(id, bank_id) ON DELETE CASCADE,
|
||||
UNIQUE (mental_model_id, bank_id, version)
|
||||
)
|
||||
""")
|
||||
|
||||
# Index for efficient version queries (get latest, list versions)
|
||||
op.execute(f"""
|
||||
CREATE INDEX idx_mental_model_versions_lookup
|
||||
ON {schema}mental_model_versions(mental_model_id, bank_id, version DESC)
|
||||
""")
|
||||
|
||||
# Add version column to mental_models to track current version
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
ADD COLUMN IF NOT EXISTS version INT NOT NULL DEFAULT 0
|
||||
""")
|
||||
|
||||
# Migrate existing mental models: create version 1 for any that have observations
|
||||
op.execute(f"""
|
||||
INSERT INTO {schema}mental_model_versions (mental_model_id, bank_id, version, observations, created_at)
|
||||
SELECT id, bank_id, 1, observations, COALESCE(last_updated, created_at)
|
||||
FROM {schema}mental_models
|
||||
WHERE observations IS NOT NULL
|
||||
AND observations != '{{"observations": []}}'::jsonb
|
||||
AND (observations->'observations') IS NOT NULL
|
||||
AND jsonb_array_length(observations->'observations') > 0
|
||||
""")
|
||||
|
||||
# Update version to 1 for migrated mental models
|
||||
op.execute(f"""
|
||||
UPDATE {schema}mental_models
|
||||
SET version = 1
|
||||
WHERE observations IS NOT NULL
|
||||
AND observations != '{{"observations": []}}'::jsonb
|
||||
AND (observations->'observations') IS NOT NULL
|
||||
AND jsonb_array_length(observations->'observations') > 0
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove mental_model_versions table and version column."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop index
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_mental_model_versions_lookup")
|
||||
|
||||
# Drop versions table
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}mental_model_versions")
|
||||
|
||||
# Remove version column from mental_models
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP COLUMN IF EXISTS version")
|
||||
@@ -0,0 +1,58 @@
|
||||
"""add_directive_subtype
|
||||
|
||||
Revision ID: k6f7g8h9i0j1
|
||||
Revises: j5e6f7g8h9i0
|
||||
Create Date: 2026-01-16 00:00:00.000000
|
||||
|
||||
This migration adds 'directive' to the mental_models subtype constraint.
|
||||
Directives are hard rules with user-provided observations that the reflect agent must follow.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "k6f7g8h9i0j1"
|
||||
down_revision: str | Sequence[str] | None = "j5e6f7g8h9i0"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add 'directive' to mental_models subtype constraint."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop existing constraint
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS ck_mental_models_subtype")
|
||||
|
||||
# Create new constraint with 'directive' added
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
ADD CONSTRAINT ck_mental_models_subtype
|
||||
CHECK (subtype IN ('structural', 'emergent', 'pinned', 'learned', 'directive'))
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove 'directive' from mental_models subtype constraint."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# First delete any directives (cannot downgrade if they exist)
|
||||
op.execute(f"DELETE FROM {schema}mental_models WHERE subtype = 'directive'")
|
||||
|
||||
# Drop constraint with directive
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS ck_mental_models_subtype")
|
||||
|
||||
# Recreate original constraint without directive
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
ADD CONSTRAINT ck_mental_models_subtype
|
||||
CHECK (subtype IN ('structural', 'emergent', 'pinned', 'learned'))
|
||||
""")
|
||||
@@ -0,0 +1,109 @@
|
||||
"""add_worker_columns
|
||||
|
||||
Revision ID: l7g8h9i0j1k2
|
||||
Revises: k6f7g8h9i0j1
|
||||
Create Date: 2026-01-19 00:00:00.000000
|
||||
|
||||
This migration adds columns to async_operations for distributed worker support:
|
||||
- worker_id: ID of the worker that claimed the task
|
||||
- claimed_at: When the task was claimed
|
||||
- retry_count: Number of retry attempts
|
||||
- task_payload: The serialized task dictionary
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import context, op
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "l7g8h9i0j1k2"
|
||||
down_revision: str | Sequence[str] | None = "k6f7g8h9i0j1"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add worker columns to async_operations."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Add worker_id column (ID of worker that claimed the task)
|
||||
op.add_column(
|
||||
"async_operations",
|
||||
sa.Column("worker_id", sa.Text(), nullable=True),
|
||||
schema=context.config.get_main_option("target_schema") or None,
|
||||
)
|
||||
|
||||
# Add claimed_at column (when task was claimed by worker)
|
||||
op.add_column(
|
||||
"async_operations",
|
||||
sa.Column("claimed_at", postgresql.TIMESTAMP(timezone=True), nullable=True),
|
||||
schema=context.config.get_main_option("target_schema") or None,
|
||||
)
|
||||
|
||||
# Add retry_count column (number of retry attempts)
|
||||
op.add_column(
|
||||
"async_operations",
|
||||
sa.Column("retry_count", sa.Integer(), server_default="0", nullable=False),
|
||||
schema=context.config.get_main_option("target_schema") or None,
|
||||
)
|
||||
|
||||
# Add task_payload column (serialized task dictionary)
|
||||
op.add_column(
|
||||
"async_operations",
|
||||
sa.Column(
|
||||
"task_payload",
|
||||
postgresql.JSONB(astext_type=sa.Text()),
|
||||
nullable=True,
|
||||
),
|
||||
schema=context.config.get_main_option("target_schema") or None,
|
||||
)
|
||||
|
||||
# Add index for efficient worker polling (pending tasks ordered by creation time)
|
||||
op.execute(
|
||||
f"CREATE INDEX idx_async_operations_pending_claim ON {schema}async_operations (status, created_at) "
|
||||
f"WHERE status = 'pending' AND task_payload IS NOT NULL"
|
||||
)
|
||||
|
||||
# Add index for finding tasks by worker_id (for decommissioning)
|
||||
op.execute(
|
||||
f"CREATE INDEX idx_async_operations_worker_id ON {schema}async_operations (worker_id) WHERE worker_id IS NOT NULL"
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove worker columns from async_operations."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop indexes
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_async_operations_pending_claim")
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_async_operations_worker_id")
|
||||
|
||||
# Drop columns
|
||||
op.drop_column(
|
||||
"async_operations",
|
||||
"task_payload",
|
||||
schema=context.config.get_main_option("target_schema") or None,
|
||||
)
|
||||
op.drop_column(
|
||||
"async_operations",
|
||||
"retry_count",
|
||||
schema=context.config.get_main_option("target_schema") or None,
|
||||
)
|
||||
op.drop_column(
|
||||
"async_operations",
|
||||
"claimed_at",
|
||||
schema=context.config.get_main_option("target_schema") or None,
|
||||
)
|
||||
op.drop_column(
|
||||
"async_operations",
|
||||
"worker_id",
|
||||
schema=context.config.get_main_option("target_schema") or None,
|
||||
)
|
||||
@@ -0,0 +1,41 @@
|
||||
"""mental_model_id_to_text
|
||||
|
||||
Revision ID: m8h9i0j1k2l3
|
||||
Revises: l7g8h9i0j1k2
|
||||
Create Date: 2026-01-19 00:00:00.000000
|
||||
|
||||
This migration changes the mental_models.id column from VARCHAR(64) to TEXT
|
||||
to support longer model IDs (e.g., entity names that exceed 64 characters).
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "m8h9i0j1k2l3"
|
||||
down_revision: str | Sequence[str] | None = "l7g8h9i0j1k2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Change mental_models.id from VARCHAR(64) to TEXT."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Alter the id column type from VARCHAR(64) to TEXT
|
||||
op.execute(f"ALTER TABLE {schema}mental_models ALTER COLUMN id TYPE TEXT")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Revert mental_models.id from TEXT to VARCHAR(64)."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Note: This may fail if any id values exceed 64 characters
|
||||
op.execute(f"ALTER TABLE {schema}mental_models ALTER COLUMN id TYPE VARCHAR(64)")
|
||||
+134
@@ -0,0 +1,134 @@
|
||||
"""learnings_and_pinned_reflections
|
||||
|
||||
Revision ID: n9i0j1k2l3m4
|
||||
Revises: m8h9i0j1k2l3
|
||||
Create Date: 2026-01-21 00:00:00.000000
|
||||
|
||||
This migration:
|
||||
1. Creates the 'learnings' table for automatic bottom-up consolidation
|
||||
2. Creates the 'pinned_reflections' table for user-curated living documents
|
||||
3. Adds consolidation tracking columns to the 'banks' table
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "n9i0j1k2l3m4"
|
||||
down_revision: str | Sequence[str] | None = "m8h9i0j1k2l3"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create learnings and pinned_reflections tables."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# 1. Create learnings table
|
||||
op.execute(f"""
|
||||
CREATE TABLE {schema}learnings (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
bank_id VARCHAR(64) NOT NULL,
|
||||
text TEXT NOT NULL,
|
||||
proof_count INT NOT NULL DEFAULT 1,
|
||||
history JSONB DEFAULT '[]'::jsonb,
|
||||
mission_context VARCHAR(64),
|
||||
pre_mission_change BOOLEAN DEFAULT FALSE,
|
||||
embedding vector(384),
|
||||
tags VARCHAR[] DEFAULT ARRAY[]::VARCHAR[],
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
|
||||
updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
|
||||
)
|
||||
""")
|
||||
|
||||
# Add foreign key constraint
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}learnings
|
||||
ADD CONSTRAINT fk_learnings_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
|
||||
""")
|
||||
|
||||
# Indexes for learnings
|
||||
op.execute(f"CREATE INDEX idx_learnings_bank_id ON {schema}learnings(bank_id)")
|
||||
op.execute(f"""
|
||||
CREATE INDEX idx_learnings_embedding ON {schema}learnings
|
||||
USING hnsw (embedding vector_cosine_ops)
|
||||
""")
|
||||
op.execute(f"CREATE INDEX idx_learnings_tags ON {schema}learnings USING GIN(tags)")
|
||||
|
||||
# Full-text search for learnings
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}learnings ADD COLUMN search_vector tsvector
|
||||
GENERATED ALWAYS AS (to_tsvector('english', text)) STORED
|
||||
""")
|
||||
op.execute(f"CREATE INDEX idx_learnings_text_search ON {schema}learnings USING gin(search_vector)")
|
||||
|
||||
# 2. Create pinned_reflections table
|
||||
op.execute(f"""
|
||||
CREATE TABLE {schema}pinned_reflections (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
bank_id VARCHAR(64) NOT NULL,
|
||||
name VARCHAR(256) NOT NULL,
|
||||
source_query TEXT NOT NULL,
|
||||
content TEXT NOT NULL,
|
||||
embedding vector(384),
|
||||
tags VARCHAR[] DEFAULT ARRAY[]::VARCHAR[],
|
||||
last_refreshed_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
|
||||
)
|
||||
""")
|
||||
|
||||
# Add foreign key constraint
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}pinned_reflections
|
||||
ADD CONSTRAINT fk_pinned_reflections_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
|
||||
""")
|
||||
|
||||
# Indexes for pinned_reflections
|
||||
op.execute(f"CREATE INDEX idx_pinned_reflections_bank_id ON {schema}pinned_reflections(bank_id)")
|
||||
op.execute(f"""
|
||||
CREATE INDEX idx_pinned_reflections_embedding ON {schema}pinned_reflections
|
||||
USING hnsw (embedding vector_cosine_ops)
|
||||
""")
|
||||
op.execute(f"CREATE INDEX idx_pinned_reflections_tags ON {schema}pinned_reflections USING GIN(tags)")
|
||||
|
||||
# Full-text search for pinned_reflections
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}pinned_reflections ADD COLUMN search_vector tsvector
|
||||
GENERATED ALWAYS AS (to_tsvector('english', COALESCE(name, '') || ' ' || content)) STORED
|
||||
""")
|
||||
op.execute(f"""
|
||||
CREATE INDEX idx_pinned_reflections_text_search ON {schema}pinned_reflections
|
||||
USING gin(search_vector)
|
||||
""")
|
||||
|
||||
# 3. Add consolidation tracking columns to banks table
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}banks
|
||||
ADD COLUMN IF NOT EXISTS last_consolidated_at TIMESTAMP WITH TIME ZONE
|
||||
""")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}banks
|
||||
ADD COLUMN IF NOT EXISTS mission_changed_at TIMESTAMP WITH TIME ZONE
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop learnings and pinned_reflections tables."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop tables
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}learnings CASCADE")
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}pinned_reflections CASCADE")
|
||||
|
||||
# Remove columns from banks
|
||||
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS last_consolidated_at")
|
||||
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS mission_changed_at")
|
||||
+113
@@ -0,0 +1,113 @@
|
||||
"""migrate_mental_models_data
|
||||
|
||||
Revision ID: o0j1k2l3m4n5
|
||||
Revises: n9i0j1k2l3m4
|
||||
Create Date: 2026-01-21 00:00:00.000000
|
||||
|
||||
This migration:
|
||||
1. Migrates existing 'pinned' mental models to the new 'pinned_reflections' table
|
||||
2. Migrates existing 'learned' mental models to the new 'learnings' table
|
||||
3. Deletes non-directive mental models (structural, emergent, pinned, learned)
|
||||
4. Drops the mental_model_versions table (no longer used)
|
||||
5. Adds a CHECK constraint that only 'directive' subtype is allowed
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "o0j1k2l3m4n5"
|
||||
down_revision: str | Sequence[str] | None = "n9i0j1k2l3m4"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Migrate data and clean up old mental models."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# 1. Migrate 'pinned' mental models to pinned_reflections
|
||||
# For pinned models, the first observation's content becomes the pinned reflection content
|
||||
op.execute(f"""
|
||||
INSERT INTO {schema}pinned_reflections (bank_id, name, source_query, content, tags, created_at)
|
||||
SELECT
|
||||
bank_id,
|
||||
name,
|
||||
description AS source_query,
|
||||
COALESCE(
|
||||
observations->'observations'->0->>'content',
|
||||
description,
|
||||
''
|
||||
) AS content,
|
||||
tags,
|
||||
created_at
|
||||
FROM {schema}mental_models
|
||||
WHERE subtype = 'pinned'
|
||||
ON CONFLICT DO NOTHING
|
||||
""")
|
||||
|
||||
# 2. Migrate 'learned' mental models to learnings
|
||||
# Each observation in a learned model becomes a separate learning
|
||||
op.execute(f"""
|
||||
INSERT INTO {schema}learnings (bank_id, text, proof_count, tags, created_at)
|
||||
SELECT
|
||||
mm.bank_id,
|
||||
obs->>'content' AS text,
|
||||
GREATEST(1, COALESCE(jsonb_array_length(obs->'evidence'), 1)) AS proof_count,
|
||||
mm.tags,
|
||||
mm.created_at
|
||||
FROM {schema}mental_models mm,
|
||||
LATERAL jsonb_array_elements(mm.observations->'observations') AS obs
|
||||
WHERE mm.subtype = 'learned'
|
||||
AND obs->>'content' IS NOT NULL
|
||||
AND obs->>'content' != ''
|
||||
ON CONFLICT DO NOTHING
|
||||
""")
|
||||
|
||||
# 3. Delete all non-directive mental models (they've been migrated or are obsolete)
|
||||
op.execute(f"""
|
||||
DELETE FROM {schema}mental_models
|
||||
WHERE subtype != 'directive'
|
||||
""")
|
||||
|
||||
# 4. Drop the mental_model_versions table (no longer used)
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}mental_model_versions CASCADE")
|
||||
|
||||
# 5. Drop old constraints and add new one that only allows 'directive'
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS ck_mental_models_subtype")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
ADD CONSTRAINT ck_mental_models_subtype CHECK (subtype = 'directive')
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Reverse the migration (data migration is one-way, so this just removes constraints)."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Remove the directive-only constraint
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS ck_mental_models_subtype")
|
||||
|
||||
# Re-create mental_model_versions table
|
||||
op.execute(f"""
|
||||
CREATE TABLE IF NOT EXISTS {schema}mental_model_versions (
|
||||
id SERIAL PRIMARY KEY,
|
||||
bank_id VARCHAR(64) NOT NULL,
|
||||
model_id VARCHAR(128) NOT NULL,
|
||||
version INT NOT NULL,
|
||||
observations JSONB NOT NULL,
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
|
||||
)
|
||||
""")
|
||||
op.execute(
|
||||
f"CREATE INDEX IF NOT EXISTS idx_mm_versions_lookup ON {schema}mental_model_versions(bank_id, model_id, version DESC)"
|
||||
)
|
||||
|
||||
# Note: Data migration cannot be reversed - pinned_reflections and learnings data remains
|
||||
+194
@@ -0,0 +1,194 @@
|
||||
"""new_knowledge_architecture
|
||||
|
||||
Revision ID: p1k2l3m4n5o6
|
||||
Revises: o0j1k2l3m4n5
|
||||
Create Date: 2026-01-21 00:00:00.000000
|
||||
|
||||
This migration implements the new knowledge architecture:
|
||||
1. Drops the 'learnings' table (mental models are now in memory_units)
|
||||
2. Renames 'pinned_reflections' to 'reflections'
|
||||
3. Drops the 'mental_models' table completely
|
||||
4. Creates 'directives' table for hard rules
|
||||
5. Adds mental model support columns to 'memory_units' (proof_count, source_memory_ids, history)
|
||||
|
||||
The new architecture:
|
||||
- Directives: Hard rules in their own table
|
||||
- Mental Models: Stored in memory_units with fact_type='mental_model'
|
||||
- Reflections: User-curated documents (renamed from pinned_reflections)
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "p1k2l3m4n5o6"
|
||||
down_revision: str | Sequence[str] | None = "o0j1k2l3m4n5"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Implement new knowledge architecture."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# 1. Drop the learnings table (mental models will be in memory_units)
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}learnings CASCADE")
|
||||
|
||||
# 2. Rename pinned_reflections to reflections
|
||||
op.execute(f"ALTER TABLE IF EXISTS {schema}pinned_reflections RENAME TO reflections")
|
||||
|
||||
# Rename indexes for reflections
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_pinned_reflections_bank_id RENAME TO idx_reflections_bank_id")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_pinned_reflections_embedding RENAME TO idx_reflections_embedding")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_pinned_reflections_tags RENAME TO idx_reflections_tags")
|
||||
op.execute(
|
||||
f"ALTER INDEX IF EXISTS {schema}idx_pinned_reflections_text_search RENAME TO idx_reflections_text_search"
|
||||
)
|
||||
|
||||
# Rename foreign key constraint
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}reflections
|
||||
DROP CONSTRAINT IF EXISTS fk_pinned_reflections_bank_id
|
||||
""")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}reflections
|
||||
ADD CONSTRAINT fk_reflections_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
|
||||
""")
|
||||
|
||||
# 3. Drop the mental_models table completely
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}mental_models CASCADE")
|
||||
|
||||
# 4. Create directives table
|
||||
op.execute(f"""
|
||||
CREATE TABLE {schema}directives (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
bank_id VARCHAR(64) NOT NULL,
|
||||
name VARCHAR(256) NOT NULL,
|
||||
content TEXT NOT NULL,
|
||||
priority INT NOT NULL DEFAULT 0,
|
||||
is_active BOOLEAN NOT NULL DEFAULT TRUE,
|
||||
tags VARCHAR[] DEFAULT ARRAY[]::VARCHAR[],
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
|
||||
updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
|
||||
)
|
||||
""")
|
||||
|
||||
# Add foreign key and indexes for directives
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}directives
|
||||
ADD CONSTRAINT fk_directives_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
|
||||
""")
|
||||
op.execute(f"CREATE INDEX idx_directives_bank_id ON {schema}directives(bank_id)")
|
||||
op.execute(f"CREATE INDEX idx_directives_bank_active ON {schema}directives(bank_id, is_active)")
|
||||
op.execute(f"CREATE INDEX idx_directives_tags ON {schema}directives USING GIN(tags)")
|
||||
|
||||
# 5. Add mental model support columns to memory_units
|
||||
# proof_count: Number of memories that support this mental model
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD COLUMN IF NOT EXISTS proof_count INT DEFAULT 1
|
||||
""")
|
||||
|
||||
# source_memory_ids: Array of memory IDs that consolidated into this mental model
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD COLUMN IF NOT EXISTS source_memory_ids UUID[] DEFAULT ARRAY[]::UUID[]
|
||||
""")
|
||||
|
||||
# history: JSONB array tracking changes to mental models
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD COLUMN IF NOT EXISTS history JSONB DEFAULT '[]'::jsonb
|
||||
""")
|
||||
|
||||
# Add index for finding mental models
|
||||
op.execute(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_units_mental_models
|
||||
ON {schema}memory_units(bank_id, fact_type)
|
||||
WHERE fact_type = 'mental_model'
|
||||
""")
|
||||
|
||||
# 6. Update fact_type check constraint to include 'mental_model'
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD CONSTRAINT memory_units_fact_type_check
|
||||
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation', 'mental_model'))
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Reverse the migration."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Restore original fact_type check constraint (without 'mental_model')
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD CONSTRAINT memory_units_fact_type_check
|
||||
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation'))
|
||||
""")
|
||||
|
||||
# Drop mental model columns from memory_units
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS proof_count")
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS source_memory_ids")
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS history")
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_mental_models")
|
||||
|
||||
# Drop directives table
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}directives CASCADE")
|
||||
|
||||
# Rename reflections back to pinned_reflections
|
||||
op.execute(f"ALTER TABLE IF EXISTS {schema}reflections RENAME TO pinned_reflections")
|
||||
|
||||
# Restore indexes
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_bank_id RENAME TO idx_pinned_reflections_bank_id")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_embedding RENAME TO idx_pinned_reflections_embedding")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_tags RENAME TO idx_pinned_reflections_tags")
|
||||
op.execute(
|
||||
f"ALTER INDEX IF EXISTS {schema}idx_reflections_text_search RENAME TO idx_pinned_reflections_text_search"
|
||||
)
|
||||
|
||||
# Restore foreign key
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}pinned_reflections
|
||||
DROP CONSTRAINT IF EXISTS fk_reflections_bank_id
|
||||
""")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}pinned_reflections
|
||||
ADD CONSTRAINT fk_pinned_reflections_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
|
||||
""")
|
||||
|
||||
# Re-create learnings table
|
||||
op.execute(f"""
|
||||
CREATE TABLE IF NOT EXISTS {schema}learnings (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
bank_id VARCHAR(64) NOT NULL,
|
||||
text TEXT NOT NULL,
|
||||
proof_count INT NOT NULL DEFAULT 1,
|
||||
history JSONB DEFAULT '[]'::jsonb,
|
||||
mission_context VARCHAR(64),
|
||||
pre_mission_change BOOLEAN DEFAULT FALSE,
|
||||
embedding vector(384),
|
||||
tags VARCHAR[] DEFAULT ARRAY[]::VARCHAR[],
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
|
||||
updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now()
|
||||
)
|
||||
""")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}learnings
|
||||
ADD CONSTRAINT fk_learnings_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
|
||||
""")
|
||||
|
||||
# Note: mental_models table recreation is complex and would need separate handling
|
||||
+50
@@ -0,0 +1,50 @@
|
||||
"""fix_mental_model_fact_type
|
||||
|
||||
Revision ID: q2l3m4n5o6p7
|
||||
Revises: p1k2l3m4n5o6
|
||||
Create Date: 2026-01-21 13:30:00.000000
|
||||
|
||||
Fix the fact_type check constraint to include 'mental_model'.
|
||||
This is a fix for p1k2l3m4n5o6 which should have included this change.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "q2l3m4n5o6p7"
|
||||
down_revision: str | Sequence[str] | None = "p1k2l3m4n5o6"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add 'mental_model' to the fact_type check constraint."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop the old constraint and add the new one with mental_model included
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD CONSTRAINT memory_units_fact_type_check
|
||||
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation', 'mental_model'))
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove 'mental_model' from the fact_type check constraint."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD CONSTRAINT memory_units_fact_type_check
|
||||
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation'))
|
||||
""")
|
||||
+47
@@ -0,0 +1,47 @@
|
||||
"""Add reflect_response JSONB column to reflections
|
||||
|
||||
Revision ID: r3m4n5o6p7q8
|
||||
Revises: q2l3m4n5o6p7
|
||||
Create Date: 2026-01-21
|
||||
|
||||
This migration adds a reflect_response JSONB column to store the full
|
||||
reflect API response payload, including based_on facts and trace data.
|
||||
|
||||
Note: Table was renamed from pinned_reflections to reflections in p1k2l3m4n5o6.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "r3m4n5o6p7q8"
|
||||
down_revision: str | Sequence[str] | None = "q2l3m4n5o6p7"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add reflect_response JSONB column to reflections."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Add reflect_response column to store the full reflect API response
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}reflections
|
||||
ADD COLUMN IF NOT EXISTS reflect_response JSONB
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove reflect_response column from reflections."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}reflections
|
||||
DROP COLUMN IF EXISTS reflect_response
|
||||
""")
|
||||
+53
@@ -0,0 +1,53 @@
|
||||
"""Add consolidated_at column to memory_units for incremental consolidation tracking.
|
||||
|
||||
This allows consolidation to track progress at the memory level rather than
|
||||
using a bank-level watermark. If consolidation crashes, already-processed
|
||||
memories won't be reprocessed.
|
||||
|
||||
Revision ID: s4n5o6p7q8r9
|
||||
Revises: r3m4n5o6p7q8
|
||||
Create Date: 2025-01-22
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "s4n5o6p7q8r9"
|
||||
down_revision: str | Sequence[str] | None = "r3m4n5o6p7q8"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Add consolidated_at column to memory_units
|
||||
op.execute(
|
||||
f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD COLUMN IF NOT EXISTS consolidated_at TIMESTAMPTZ DEFAULT NULL
|
||||
"""
|
||||
)
|
||||
|
||||
# Create index for efficient querying of unconsolidated memories
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_units_unconsolidated
|
||||
ON {schema}memory_units (bank_id, created_at)
|
||||
WHERE consolidated_at IS NULL AND fact_type IN ('experience', 'world')
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_unconsolidated")
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS consolidated_at")
|
||||
+134
@@ -0,0 +1,134 @@
|
||||
"""Rename mental_model fact_type to observation and reflections table to mental_models
|
||||
|
||||
Revision ID: t5o6p7q8r9s0
|
||||
Revises: s4n5o6p7q8r9
|
||||
Create Date: 2026-01-26
|
||||
|
||||
This migration implements the terminology rename:
|
||||
1. mental_model (fact_type in memory_units) -> observation
|
||||
2. reflections table -> mental_models table
|
||||
|
||||
The new terminology:
|
||||
- Observations: Consolidated knowledge synthesized from facts (was mental_model)
|
||||
- Mental Models: Stored reflect responses (was reflections)
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "t5o6p7q8r9s0"
|
||||
down_revision: str | Sequence[str] | None = "s4n5o6p7q8r9"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Rename mental_model -> observation and reflections -> mental_models."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# 1. Update fact_type values: mental_model -> observation
|
||||
op.execute(f"""
|
||||
UPDATE {schema}memory_units
|
||||
SET fact_type = 'observation'
|
||||
WHERE fact_type = 'mental_model'
|
||||
""")
|
||||
|
||||
# 2. Update the CHECK constraint - remove mental_model, keep observation
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD CONSTRAINT memory_units_fact_type_check
|
||||
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation'))
|
||||
""")
|
||||
|
||||
# 3. Rename the index for observations (was for mental_models)
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_mental_models")
|
||||
op.execute(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_units_observations
|
||||
ON {schema}memory_units(bank_id, fact_type)
|
||||
WHERE fact_type = 'observation'
|
||||
""")
|
||||
|
||||
# 4. Update the unconsolidated index to not filter by fact_type since observations
|
||||
# are now the consolidated type
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_unconsolidated")
|
||||
op.execute(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_units_unconsolidated
|
||||
ON {schema}memory_units (bank_id, created_at)
|
||||
WHERE consolidated_at IS NULL AND fact_type IN ('experience', 'world')
|
||||
""")
|
||||
|
||||
# 5. Rename reflections table to mental_models
|
||||
op.execute(f"ALTER TABLE IF EXISTS {schema}reflections RENAME TO mental_models")
|
||||
|
||||
# 6. Rename indexes for mental_models (was reflections)
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_bank_id RENAME TO idx_mental_models_bank_id")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_embedding RENAME TO idx_mental_models_embedding")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_tags RENAME TO idx_mental_models_tags")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_reflections_text_search RENAME TO idx_mental_models_text_search")
|
||||
|
||||
# 7. Rename foreign key constraint
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
DROP CONSTRAINT IF EXISTS fk_reflections_bank_id
|
||||
""")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
ADD CONSTRAINT fk_mental_models_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Reverse: observation -> mental_model and mental_models -> reflections."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# 1. Rename mental_models table back to reflections
|
||||
op.execute(f"ALTER TABLE IF EXISTS {schema}mental_models RENAME TO reflections")
|
||||
|
||||
# 2. Rename indexes back
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_mental_models_bank_id RENAME TO idx_reflections_bank_id")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_mental_models_embedding RENAME TO idx_reflections_embedding")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_mental_models_tags RENAME TO idx_reflections_tags")
|
||||
op.execute(f"ALTER INDEX IF EXISTS {schema}idx_mental_models_text_search RENAME TO idx_reflections_text_search")
|
||||
|
||||
# 3. Rename foreign key back
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}reflections
|
||||
DROP CONSTRAINT IF EXISTS fk_mental_models_bank_id
|
||||
""")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}reflections
|
||||
ADD CONSTRAINT fk_reflections_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE
|
||||
""")
|
||||
|
||||
# 4. Update fact_type values: observation -> mental_model
|
||||
op.execute(f"""
|
||||
UPDATE {schema}memory_units
|
||||
SET fact_type = 'mental_model'
|
||||
WHERE fact_type = 'observation'
|
||||
""")
|
||||
|
||||
# 5. Update the CHECK constraint back
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP CONSTRAINT IF EXISTS memory_units_fact_type_check")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}memory_units
|
||||
ADD CONSTRAINT memory_units_fact_type_check
|
||||
CHECK (fact_type IN ('world', 'experience', 'opinion', 'observation', 'mental_model'))
|
||||
""")
|
||||
|
||||
# 6. Rename index back
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_observations")
|
||||
op.execute(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_memory_units_mental_models
|
||||
ON {schema}memory_units(bank_id, fact_type)
|
||||
WHERE fact_type = 'mental_model'
|
||||
""")
|
||||
@@ -0,0 +1,41 @@
|
||||
"""Change mental_models.id from UUID to TEXT
|
||||
|
||||
Revision ID: u6p7q8r9s0t1
|
||||
Revises: t5o6p7q8r9s0
|
||||
Create Date: 2026-01-27
|
||||
|
||||
This migration changes the mental_models.id column from UUID to TEXT
|
||||
to support user-defined text identifiers like 'team-communication' instead of UUIDs.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "u6p7q8r9s0t1"
|
||||
down_revision: str | Sequence[str] | None = "t5o6p7q8r9s0"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Change mental_models.id from UUID to TEXT."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Change the id column type from UUID to TEXT
|
||||
# Existing UUIDs will be converted to their string representation
|
||||
op.execute(f"ALTER TABLE {schema}mental_models ALTER COLUMN id TYPE TEXT USING id::TEXT")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Revert mental_models.id from TEXT to UUID."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Note: This will fail if any id values are not valid UUIDs
|
||||
op.execute(f"ALTER TABLE {schema}mental_models ALTER COLUMN id TYPE UUID USING id::UUID")
|
||||
+50
@@ -0,0 +1,50 @@
|
||||
"""Add max_tokens and trigger columns to mental_models
|
||||
|
||||
Revision ID: v7q8r9s0t1u2
|
||||
Revises: u6p7q8r9s0t1
|
||||
Create Date: 2026-01-27
|
||||
|
||||
This migration adds:
|
||||
- max_tokens column: token limit for content generation during refresh
|
||||
- trigger column: JSONB for trigger settings (e.g., refresh_after_consolidation)
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "v7q8r9s0t1u2"
|
||||
down_revision: str | Sequence[str] | None = "u6p7q8r9s0t1"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add max_tokens and trigger columns to mental_models."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
ADD COLUMN IF NOT EXISTS max_tokens INT NOT NULL DEFAULT 2048
|
||||
""")
|
||||
|
||||
# trigger column stores trigger settings as JSONB
|
||||
# Default: refresh_after_consolidation = false (not "real time")
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
ADD COLUMN IF NOT EXISTS trigger JSONB NOT NULL DEFAULT '{{"refresh_after_consolidation": false}}'::jsonb
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove max_tokens and trigger columns from mental_models."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP COLUMN IF EXISTS max_tokens")
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP COLUMN IF EXISTS trigger")
|
||||
+60
@@ -0,0 +1,60 @@
|
||||
"""Fix mental_models primary key to be scoped per bank
|
||||
|
||||
Revision ID: w8r9s0t1u2v3
|
||||
Revises: v7q8r9s0t1u2
|
||||
Create Date: 2026-02-05
|
||||
|
||||
This migration fixes a critical bank isolation bug where mental_models.id was
|
||||
globally unique across all banks instead of being scoped per bank. This caused
|
||||
conflicts when different banks tried to use the same custom ID.
|
||||
|
||||
CRITICAL FIX: Changes primary key from (id) to (bank_id, id) to ensure proper isolation.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "w8r9s0t1u2v3"
|
||||
down_revision: str | Sequence[str] | None = "v7q8r9s0t1u2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Change mental_models primary key from (id) to (bank_id, id) for proper bank isolation."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop the old primary key constraint (just id)
|
||||
# Note: The constraint might be named differently on different DBs
|
||||
# Try both old names (pinned_reflections_pkey from original, mental_models_pkey from rename)
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS pinned_reflections_pkey")
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS mental_models_pkey")
|
||||
|
||||
# Create the new composite primary key (bank_id, id)
|
||||
# This ensures IDs are scoped per bank, not globally
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
ADD CONSTRAINT mental_models_pkey PRIMARY KEY (bank_id, id)
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Revert mental_models primary key from (bank_id, id) to (id)."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop the composite primary key
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP CONSTRAINT IF EXISTS mental_models_pkey")
|
||||
|
||||
# Restore the old primary key (just id)
|
||||
# WARNING: This downgrade will fail if there are duplicate IDs across banks
|
||||
op.execute(f"""
|
||||
ALTER TABLE {schema}mental_models
|
||||
ADD CONSTRAINT mental_models_pkey PRIMARY KEY (id)
|
||||
""")
|
||||
@@ -72,22 +72,24 @@ def create_app(
|
||||
|
||||
# Mount MCP server and chain its lifespan if enabled
|
||||
if mcp_app is not None:
|
||||
# Get the MCP app's underlying Starlette app for lifespan access
|
||||
mcp_starlette_app = mcp_app.mcp_app
|
||||
# Get both MCP apps' underlying Starlette apps for lifespan access
|
||||
multi_bank_starlette_app = mcp_app.multi_bank_app
|
||||
single_bank_starlette_app = mcp_app.single_bank_app
|
||||
|
||||
# Store the original lifespan
|
||||
original_lifespan = app.router.lifespan_context
|
||||
|
||||
@asynccontextmanager
|
||||
async def chained_lifespan(app_instance: FastAPI):
|
||||
"""Chain the MCP lifespan with the main app lifespan."""
|
||||
# Start MCP lifespan first
|
||||
async with mcp_starlette_app.router.lifespan_context(mcp_starlette_app):
|
||||
logger.info("MCP lifespan started")
|
||||
# Then start the original app lifespan
|
||||
async with original_lifespan(app_instance):
|
||||
yield
|
||||
logger.info("MCP lifespan stopped")
|
||||
"""Chain both MCP lifespans with the main app lifespan."""
|
||||
# Start both MCP lifespans (multi-bank and single-bank)
|
||||
async with multi_bank_starlette_app.router.lifespan_context(multi_bank_starlette_app):
|
||||
async with single_bank_starlette_app.router.lifespan_context(single_bank_starlette_app):
|
||||
logger.info("MCP lifespans started (multi-bank and single-bank)")
|
||||
# Then start the original app lifespan
|
||||
async with original_lifespan(app_instance):
|
||||
yield
|
||||
logger.info("MCP lifespans stopped")
|
||||
|
||||
# Replace the app's lifespan with the chained version
|
||||
app.router.lifespan_context = chained_lifespan
|
||||
|
||||
+1257
-115
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,4 @@
|
||||
"""Hindsight MCP Server implementation using FastMCP."""
|
||||
"""Hindsight MCP Server implementation using FastMCP (HTTP transport)."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
@@ -8,7 +8,10 @@ from contextvars import ContextVar
|
||||
from fastmcp import FastMCP
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
|
||||
from hindsight_api.engine.memory_engine import _current_schema
|
||||
from hindsight_api.extensions import MCPExtension, load_extension
|
||||
from hindsight_api.extensions.tenant import AuthenticationError
|
||||
from hindsight_api.mcp_tools import MCPToolsConfig, register_mcp_tools
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
# Configure logging from HINDSIGHT_API_LOG_LEVEL environment variable
|
||||
@@ -30,21 +33,49 @@ logger = logging.getLogger(__name__)
|
||||
# Default bank_id from environment variable
|
||||
DEFAULT_BANK_ID = os.environ.get("HINDSIGHT_MCP_BANK_ID", "default")
|
||||
|
||||
# Legacy MCP authentication token (for backwards compatibility)
|
||||
# If set, this token is checked first before TenantExtension auth
|
||||
MCP_AUTH_TOKEN = os.environ.get("HINDSIGHT_API_MCP_AUTH_TOKEN")
|
||||
|
||||
# Context variable to hold the current bank_id
|
||||
_current_bank_id: ContextVar[str | None] = ContextVar("current_bank_id", default=None)
|
||||
|
||||
# Context variable to hold the current API key (for tenant auth propagation)
|
||||
_current_api_key: ContextVar[str | None] = ContextVar("current_api_key", default=None)
|
||||
|
||||
# Context variables for tenant_id and api_key_id (set by authenticate, used by usage metering)
|
||||
_current_tenant_id: ContextVar[str | None] = ContextVar("current_tenant_id", default=None)
|
||||
_current_api_key_id: ContextVar[str | None] = ContextVar("current_api_key_id", default=None)
|
||||
|
||||
|
||||
def get_current_bank_id() -> str | None:
|
||||
"""Get the current bank_id from context."""
|
||||
return _current_bank_id.get()
|
||||
|
||||
|
||||
def create_mcp_server(memory: MemoryEngine) -> FastMCP:
|
||||
def get_current_api_key() -> str | None:
|
||||
"""Get the current API key from context."""
|
||||
return _current_api_key.get()
|
||||
|
||||
|
||||
def get_current_tenant_id() -> str | None:
|
||||
"""Get the current tenant_id from context."""
|
||||
return _current_tenant_id.get()
|
||||
|
||||
|
||||
def get_current_api_key_id() -> str | None:
|
||||
"""Get the current api_key_id from context."""
|
||||
return _current_api_key_id.get()
|
||||
|
||||
|
||||
def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||
"""
|
||||
Create and configure the Hindsight MCP server.
|
||||
|
||||
Args:
|
||||
memory: MemoryEngine instance (required)
|
||||
multi_bank: If True, expose all tools with bank_id parameters (default).
|
||||
If False, only expose bank-scoped tools without bank_id parameters.
|
||||
|
||||
Returns:
|
||||
Configured FastMCP server instance with stateless_http enabled
|
||||
@@ -52,218 +83,81 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
|
||||
# Use stateless_http=True for Claude Code compatibility
|
||||
mcp = FastMCP("hindsight-mcp-server", stateless_http=True)
|
||||
|
||||
@mcp.tool()
|
||||
async def retain(
|
||||
content: str,
|
||||
context: str = "general",
|
||||
async_processing: bool = True,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Store important information to long-term memory.
|
||||
# Configure and register tools using shared module
|
||||
config = MCPToolsConfig(
|
||||
bank_id_resolver=get_current_bank_id,
|
||||
api_key_resolver=get_current_api_key, # Propagate API key for tenant auth
|
||||
tenant_id_resolver=get_current_tenant_id, # Propagate tenant_id for usage metering
|
||||
api_key_id_resolver=get_current_api_key_id, # Propagate api_key_id for usage metering
|
||||
include_bank_id_param=multi_bank,
|
||||
tools=None if multi_bank else {"retain", "recall", "reflect"}, # Scoped tools for single-bank mode
|
||||
retain_fire_and_forget=False, # HTTP MCP supports sync/async modes
|
||||
)
|
||||
|
||||
Use this tool PROACTIVELY whenever the user shares:
|
||||
- Personal facts, preferences, or interests
|
||||
- Important events or milestones
|
||||
- User history, experiences, or background
|
||||
- Decisions, opinions, or stated preferences
|
||||
- Goals, plans, or future intentions
|
||||
- Relationships or people mentioned
|
||||
- Work context, projects, or responsibilities
|
||||
register_mcp_tools(mcp, memory, config)
|
||||
|
||||
Args:
|
||||
content: The fact/memory to store (be specific and include relevant details)
|
||||
context: Category for the memory (e.g., 'preferences', 'work', 'hobbies', 'family'). Default: 'general'
|
||||
async_processing: If True, queue for background processing and return immediately. If False, wait for completion. Default: True
|
||||
bank_id: Optional bank to store in (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or get_current_bank_id()
|
||||
if target_bank is None:
|
||||
return "Error: No bank_id configured"
|
||||
contents = [{"content": content, "context": context}]
|
||||
if async_processing:
|
||||
# Queue for background processing and return immediately
|
||||
result = await memory.submit_async_retain(
|
||||
bank_id=target_bank, contents=contents, request_context=RequestContext()
|
||||
)
|
||||
return f"Memory queued for background processing (operation_id: {result.get('operation_id', 'N/A')})"
|
||||
else:
|
||||
# Wait for completion
|
||||
await memory.retain_batch_async(
|
||||
bank_id=target_bank,
|
||||
contents=contents,
|
||||
request_context=RequestContext(),
|
||||
)
|
||||
return f"Memory stored successfully in bank '{target_bank}'"
|
||||
except Exception as e:
|
||||
logger.error(f"Error storing memory: {e}", exc_info=True)
|
||||
return f"Error: {str(e)}"
|
||||
|
||||
@mcp.tool()
|
||||
async def recall(query: str, max_tokens: int = 4096, bank_id: str | None = None) -> str:
|
||||
"""
|
||||
Search memories to provide personalized, context-aware responses.
|
||||
|
||||
Use this tool PROACTIVELY to:
|
||||
- Check user's preferences before making suggestions
|
||||
- Recall user's history to provide continuity
|
||||
- Remember user's goals and context
|
||||
- Personalize responses based on past interactions
|
||||
|
||||
Args:
|
||||
query: Natural language search query (e.g., "user's food preferences", "what projects is user working on")
|
||||
max_tokens: Maximum tokens in the response (default: 4096)
|
||||
bank_id: Optional bank to search in (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or get_current_bank_id()
|
||||
if target_bank is None:
|
||||
return "Error: No bank_id configured"
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
|
||||
recall_result = await memory.recall_async(
|
||||
bank_id=target_bank,
|
||||
query=query,
|
||||
fact_type=list(VALID_RECALL_FACT_TYPES),
|
||||
budget=Budget.HIGH,
|
||||
max_tokens=max_tokens,
|
||||
request_context=RequestContext(),
|
||||
)
|
||||
|
||||
# Use model's JSON serialization
|
||||
return recall_result.model_dump_json(indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"Error searching: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}", "results": []}}'
|
||||
|
||||
@mcp.tool()
|
||||
async def reflect(query: str, context: str | None = None, budget: str = "low", bank_id: str | None = None) -> str:
|
||||
"""
|
||||
Generate thoughtful analysis by synthesizing stored memories with the bank's personality.
|
||||
|
||||
WHEN TO USE THIS TOOL:
|
||||
Use reflect when you need reasoned analysis, not just fact retrieval. This tool
|
||||
thinks through the question using everything the bank knows and its personality traits.
|
||||
|
||||
EXAMPLES OF GOOD QUERIES:
|
||||
- "What patterns have emerged in how I approach debugging?"
|
||||
- "Based on my past decisions, what architectural style do I prefer?"
|
||||
- "What might be the best approach for this problem given what you know about me?"
|
||||
- "How should I prioritize these tasks based on my goals?"
|
||||
|
||||
HOW IT DIFFERS FROM RECALL:
|
||||
- recall: Returns raw facts matching your search (fast lookup)
|
||||
- reflect: Reasons across memories to form a synthesized answer (deeper analysis)
|
||||
|
||||
Use recall for "what did I say about X?" and reflect for "what should I do about X?"
|
||||
|
||||
Args:
|
||||
query: The question or topic to reflect on
|
||||
context: Optional context about why this reflection is needed
|
||||
budget: Search budget - 'low', 'mid', or 'high' (default: 'low')
|
||||
bank_id: Optional bank to reflect in (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or get_current_bank_id()
|
||||
if target_bank is None:
|
||||
return "Error: No bank_id configured"
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
|
||||
# Map string budget to enum
|
||||
budget_map = {"low": Budget.LOW, "mid": Budget.MID, "high": Budget.HIGH}
|
||||
budget_enum = budget_map.get(budget.lower(), Budget.LOW)
|
||||
|
||||
reflect_result = await memory.reflect_async(
|
||||
bank_id=target_bank,
|
||||
query=query,
|
||||
budget=budget_enum,
|
||||
context=context,
|
||||
request_context=RequestContext(),
|
||||
)
|
||||
|
||||
return reflect_result.model_dump_json(indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"Error reflecting: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}", "text": ""}}'
|
||||
|
||||
@mcp.tool()
|
||||
async def list_banks() -> str:
|
||||
"""
|
||||
List all available memory banks.
|
||||
|
||||
Use this tool to discover what memory banks exist in the system.
|
||||
Each bank is an isolated memory store (like a separate "brain").
|
||||
|
||||
Returns:
|
||||
JSON list of banks with their IDs, names, dispositions, and backgrounds.
|
||||
"""
|
||||
try:
|
||||
banks = await memory.list_banks(request_context=RequestContext())
|
||||
return json.dumps({"banks": banks}, indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"Error listing banks: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}", "banks": []}}'
|
||||
|
||||
@mcp.tool()
|
||||
async def create_bank(bank_id: str, name: str | None = None, background: str | None = None) -> str:
|
||||
"""
|
||||
Create a new memory bank or get an existing one.
|
||||
|
||||
Memory banks are isolated stores - each one is like a separate "brain" for a user/agent.
|
||||
Banks are auto-created with default settings if they don't exist.
|
||||
|
||||
Args:
|
||||
bank_id: Unique identifier for the bank (e.g., 'user-123', 'agent-alpha')
|
||||
name: Optional human-friendly name for the bank
|
||||
background: Optional background context about the bank's owner/purpose
|
||||
"""
|
||||
try:
|
||||
# get_bank_profile auto-creates bank if it doesn't exist
|
||||
profile = await memory.get_bank_profile(bank_id, request_context=RequestContext())
|
||||
|
||||
# Update name/background if provided
|
||||
if name is not None or background is not None:
|
||||
await memory.update_bank(
|
||||
bank_id,
|
||||
name=name,
|
||||
background=background,
|
||||
request_context=RequestContext(),
|
||||
)
|
||||
# Fetch updated profile
|
||||
profile = await memory.get_bank_profile(bank_id, request_context=RequestContext())
|
||||
|
||||
# Serialize disposition if it's a Pydantic model
|
||||
if "disposition" in profile and hasattr(profile["disposition"], "model_dump"):
|
||||
profile["disposition"] = profile["disposition"].model_dump()
|
||||
return json.dumps(profile, indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating bank: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}"}}'
|
||||
# Load and register additional tools from MCP extension if configured
|
||||
mcp_extension = load_extension("MCP", MCPExtension)
|
||||
if mcp_extension:
|
||||
logger.info(f"Loading MCP extension: {mcp_extension.__class__.__name__}")
|
||||
mcp_extension.register_tools(mcp, memory)
|
||||
|
||||
return mcp
|
||||
|
||||
|
||||
class MCPMiddleware:
|
||||
"""ASGI middleware that extracts bank_id from header or path and sets context.
|
||||
"""ASGI middleware that handles authentication and routes to appropriate MCP server.
|
||||
|
||||
Bank ID can be provided via:
|
||||
1. X-Bank-Id header (recommended for Claude Code)
|
||||
2. URL path: /mcp/{bank_id}/
|
||||
3. Environment variable HINDSIGHT_MCP_BANK_ID (fallback default)
|
||||
Authentication:
|
||||
1. If HINDSIGHT_API_MCP_AUTH_TOKEN is set (legacy), validates against that token
|
||||
2. Otherwise, uses TenantExtension.authenticate_mcp() from the MemoryEngine
|
||||
- DefaultTenantExtension: no auth required (local dev)
|
||||
- ApiKeyTenantExtension: validates against env var
|
||||
|
||||
For Claude Code, configure with:
|
||||
Two modes based on URL structure:
|
||||
|
||||
1. Multi-bank mode (for /mcp/ root endpoint):
|
||||
- Exposes all tools: retain, recall, reflect, list_banks, create_bank
|
||||
- All tools include optional bank_id parameter for cross-bank operations
|
||||
- Bank ID from: X-Bank-Id header or HINDSIGHT_MCP_BANK_ID env var
|
||||
|
||||
2. Single-bank mode (for /mcp/{bank_id}/ endpoints):
|
||||
- Exposes bank-scoped tools only: retain, recall, reflect
|
||||
- No bank_id parameter (comes from URL)
|
||||
- No bank management tools (list_banks, create_bank)
|
||||
- Recommended for agent isolation
|
||||
|
||||
Examples:
|
||||
# Single-bank mode (recommended for agent isolation)
|
||||
claude mcp add --transport http my-agent http://localhost:8888/mcp/my-agent-bank/ \\
|
||||
--header "Authorization: Bearer <token>"
|
||||
|
||||
# Multi-bank mode (for cross-bank operations)
|
||||
claude mcp add --transport http hindsight http://localhost:8888/mcp \\
|
||||
--header "X-Bank-Id: my-bank"
|
||||
--header "X-Bank-Id: my-bank" --header "Authorization: Bearer <token>"
|
||||
"""
|
||||
|
||||
def __init__(self, app, memory: MemoryEngine):
|
||||
self.app = app
|
||||
self.memory = memory
|
||||
self.mcp_server = create_mcp_server(memory)
|
||||
self.mcp_app = self.mcp_server.http_app(path="/")
|
||||
# Expose the lifespan for the parent app to chain
|
||||
self.lifespan = self.mcp_app.lifespan_handler if hasattr(self.mcp_app, "lifespan_handler") else None
|
||||
self.tenant_extension = memory._tenant_extension
|
||||
|
||||
# Create two server instances:
|
||||
# 1. Multi-bank server (for /mcp/ root endpoint)
|
||||
self.multi_bank_server = create_mcp_server(memory, multi_bank=True)
|
||||
self.multi_bank_app = self.multi_bank_server.http_app(path="/")
|
||||
|
||||
# 2. Single-bank server (for /mcp/{bank_id}/ endpoints)
|
||||
self.single_bank_server = create_mcp_server(memory, multi_bank=False)
|
||||
self.single_bank_app = self.single_bank_server.http_app(path="/")
|
||||
|
||||
# Backward compatibility: expose multi_bank_app as mcp_app
|
||||
self.mcp_app = self.multi_bank_app
|
||||
|
||||
# Expose the lifespan for the parent app to chain (use multi-bank as default)
|
||||
self.lifespan = (
|
||||
self.multi_bank_app.lifespan_handler if hasattr(self.multi_bank_app, "lifespan_handler") else None
|
||||
)
|
||||
|
||||
def _get_header(self, scope: dict, name: str) -> str | None:
|
||||
"""Extract a header value from ASGI scope."""
|
||||
@@ -275,9 +169,47 @@ class MCPMiddleware:
|
||||
|
||||
async def __call__(self, scope, receive, send):
|
||||
if scope["type"] != "http":
|
||||
await self.mcp_app(scope, receive, send)
|
||||
await self.multi_bank_app(scope, receive, send)
|
||||
return
|
||||
|
||||
# Extract auth token from header (for tenant auth propagation)
|
||||
auth_header = self._get_header(scope, "Authorization")
|
||||
auth_token: str | None = None
|
||||
if auth_header:
|
||||
# Support both "Bearer <token>" and direct token
|
||||
auth_token = auth_header[7:].strip() if auth_header.startswith("Bearer ") else auth_header.strip()
|
||||
|
||||
# Authenticate: check legacy MCP_AUTH_TOKEN first, then TenantExtension
|
||||
tenant_context = None
|
||||
auth_tenant_id: str | None = None
|
||||
auth_api_key_id: str | None = None
|
||||
if MCP_AUTH_TOKEN:
|
||||
# Legacy authentication mode - validate against static token
|
||||
if not auth_token:
|
||||
await self._send_error(send, 401, "Authorization header required")
|
||||
return
|
||||
if auth_token != MCP_AUTH_TOKEN:
|
||||
await self._send_error(send, 401, "Invalid authentication token")
|
||||
return
|
||||
# Legacy mode doesn't use tenant schemas
|
||||
tenant_context = None
|
||||
else:
|
||||
# Use TenantExtension.authenticate_mcp() for auth
|
||||
try:
|
||||
auth_context = RequestContext(api_key=auth_token)
|
||||
tenant_context = await self.tenant_extension.authenticate_mcp(auth_context)
|
||||
# Capture tenant_id and api_key_id set by authenticate() for usage metering
|
||||
auth_tenant_id = auth_context.tenant_id
|
||||
auth_api_key_id = auth_context.api_key_id
|
||||
except AuthenticationError as e:
|
||||
await self._send_error(send, 401, str(e))
|
||||
return
|
||||
|
||||
# Set schema from tenant context so downstream DB queries use the correct schema
|
||||
schema_token = (
|
||||
_current_schema.set(tenant_context.schema_name) if tenant_context and tenant_context.schema_name else None
|
||||
)
|
||||
|
||||
path = scope.get("path", "")
|
||||
|
||||
# Strip any mount prefix (e.g., /mcp) that FastAPI might not have stripped
|
||||
@@ -291,8 +223,13 @@ class MCPMiddleware:
|
||||
elif path == "/mcp":
|
||||
path = "/"
|
||||
|
||||
# Ensure path has leading slash (needed after stripping mount path)
|
||||
if path and not path.startswith("/"):
|
||||
path = "/" + path
|
||||
|
||||
# Try to get bank_id from header first (for Claude Code compatibility)
|
||||
bank_id = self._get_header(scope, "X-Bank-Id")
|
||||
bank_id_from_path = False
|
||||
|
||||
# MCP endpoint paths that should not be treated as bank_ids
|
||||
MCP_ENDPOINTS = {"sse", "messages"}
|
||||
@@ -305,6 +242,7 @@ class MCPMiddleware:
|
||||
if parts[0] and parts[0] not in MCP_ENDPOINTS:
|
||||
# First segment looks like a bank_id
|
||||
bank_id = parts[0]
|
||||
bank_id_from_path = True
|
||||
new_path = "/" + parts[1] if len(parts) > 1 else "/"
|
||||
|
||||
# Fall back to default bank_id
|
||||
@@ -312,8 +250,18 @@ class MCPMiddleware:
|
||||
bank_id = DEFAULT_BANK_ID
|
||||
logger.debug(f"Using default bank_id: {bank_id}")
|
||||
|
||||
# Set bank_id context
|
||||
token = _current_bank_id.set(bank_id)
|
||||
# Select the appropriate MCP app based on how bank_id was provided:
|
||||
# - Path-based bank_id → single-bank app (no bank_id param, scoped tools)
|
||||
# - Header/env bank_id → multi-bank app (bank_id param, all tools)
|
||||
target_app = self.single_bank_app if bank_id_from_path else self.multi_bank_app
|
||||
|
||||
# Set bank_id, api_key, tenant_id, and api_key_id context
|
||||
bank_id_token = _current_bank_id.set(bank_id)
|
||||
# Store the auth token for tenant extension to validate
|
||||
api_key_token = _current_api_key.set(auth_token) if auth_token else None
|
||||
# Store tenant_id and api_key_id from authentication for usage metering
|
||||
tenant_id_token = _current_tenant_id.set(auth_tenant_id) if auth_tenant_id else None
|
||||
api_key_id_token = _current_api_key_id.set(auth_api_key_id) if auth_api_key_id else None
|
||||
try:
|
||||
new_scope = scope.copy()
|
||||
new_scope["path"] = new_path
|
||||
@@ -322,7 +270,7 @@ class MCPMiddleware:
|
||||
|
||||
# Wrap send to rewrite the SSE endpoint URL to include bank_id if using path-based routing
|
||||
async def send_wrapper(message):
|
||||
if message["type"] == "http.response.body":
|
||||
if message["type"] == "http.response.body" and bank_id_from_path:
|
||||
body = message.get("body", b"")
|
||||
if body and b"/messages" in body:
|
||||
# Rewrite /messages to /{bank_id}/messages in SSE endpoint event
|
||||
@@ -330,9 +278,17 @@ class MCPMiddleware:
|
||||
message = {**message, "body": body}
|
||||
await send(message)
|
||||
|
||||
await self.mcp_app(new_scope, receive, send_wrapper)
|
||||
await target_app(new_scope, receive, send_wrapper)
|
||||
finally:
|
||||
_current_bank_id.reset(token)
|
||||
_current_bank_id.reset(bank_id_token)
|
||||
if api_key_token is not None:
|
||||
_current_api_key.reset(api_key_token)
|
||||
if tenant_id_token is not None:
|
||||
_current_tenant_id.reset(tenant_id_token)
|
||||
if api_key_id_token is not None:
|
||||
_current_api_key_id.reset(api_key_id_token)
|
||||
if schema_token is not None:
|
||||
_current_schema.reset(schema_token)
|
||||
|
||||
async def _send_error(self, send, status: int, message: str):
|
||||
"""Send an error response."""
|
||||
@@ -354,12 +310,23 @@ class MCPMiddleware:
|
||||
|
||||
def create_mcp_app(memory: MemoryEngine):
|
||||
"""
|
||||
Create an ASGI app that handles MCP requests.
|
||||
Create an ASGI app that handles MCP requests with dynamic tool exposure.
|
||||
|
||||
Bank ID can be provided via:
|
||||
1. X-Bank-Id header: claude mcp add --transport http hindsight http://localhost:8888/mcp --header "X-Bank-Id: my-bank"
|
||||
2. URL path: /mcp/{bank_id}/
|
||||
3. Environment variable HINDSIGHT_MCP_BANK_ID (fallback, default: "default")
|
||||
Authentication:
|
||||
Uses the TenantExtension from the MemoryEngine (same auth as REST API).
|
||||
|
||||
Two modes based on URL structure:
|
||||
|
||||
1. Single-bank mode (recommended for agent isolation):
|
||||
- URL: /mcp/{bank_id}/
|
||||
- Tools: retain, recall, reflect (no bank_id parameter)
|
||||
- Example: claude mcp add --transport http my-agent http://localhost:8888/mcp/my-agent-bank/
|
||||
|
||||
2. Multi-bank mode (for cross-bank operations):
|
||||
- URL: /mcp/
|
||||
- Tools: retain, recall, reflect, list_banks, create_bank (all with bank_id parameter)
|
||||
- Bank ID from: X-Bank-Id header or HINDSIGHT_MCP_BANK_ID env var (default: "default")
|
||||
- Example: claude mcp add --transport http hindsight http://localhost:8888/mcp --header "X-Bank-Id: my-bank"
|
||||
|
||||
Args:
|
||||
memory: MemoryEngine instance
|
||||
|
||||
@@ -4,6 +4,8 @@ Banner display for Hindsight API startup.
|
||||
Shows the logo and tagline with gradient colors.
|
||||
"""
|
||||
|
||||
from .utils import mask_network_location
|
||||
|
||||
# Gradient colors: #0074d9 -> #009296
|
||||
GRADIENT_START = (0, 116, 217) # #0074d9
|
||||
GRADIENT_END = (0, 146, 150) # #009296
|
||||
@@ -83,11 +85,14 @@ def print_startup_info(
|
||||
embeddings_provider: str,
|
||||
reranker_provider: str,
|
||||
mcp_enabled: bool = False,
|
||||
version: str | None = None,
|
||||
):
|
||||
"""Print styled startup information."""
|
||||
print(color_start("Starting Hindsight API..."))
|
||||
if version:
|
||||
print(f" {dim('Version:')} {color(f'v{version}', 0.1)}")
|
||||
print(f" {dim('URL:')} {color(f'http://{host}:{port}', 0.2)}")
|
||||
print(f" {dim('Database:')} {color(database_url, 0.4)}")
|
||||
print(f" {dim('Database:')} {color(mask_network_location(database_url), 0.4)}")
|
||||
print(f" {dim('LLM:')} {color(f'{llm_provider} / {llm_model}', 0.6)}")
|
||||
print(f" {dim('Embeddings:')} {color(embeddings_provider, 0.8)}")
|
||||
print(f" {dim('Reranker:')} {color(reranker_provider, 1.0)}")
|
||||
|
||||
@@ -4,9 +4,12 @@ Centralized configuration for Hindsight API.
|
||||
All environment variables and their defaults are defined here.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from dotenv import find_dotenv, load_dotenv
|
||||
|
||||
@@ -17,11 +20,15 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# Environment variable names
|
||||
ENV_DATABASE_URL = "HINDSIGHT_API_DATABASE_URL"
|
||||
ENV_DATABASE_SCHEMA = "HINDSIGHT_API_DATABASE_SCHEMA"
|
||||
ENV_LLM_PROVIDER = "HINDSIGHT_API_LLM_PROVIDER"
|
||||
ENV_LLM_API_KEY = "HINDSIGHT_API_LLM_API_KEY"
|
||||
ENV_LLM_MODEL = "HINDSIGHT_API_LLM_MODEL"
|
||||
ENV_LLM_BASE_URL = "HINDSIGHT_API_LLM_BASE_URL"
|
||||
ENV_LLM_MAX_CONCURRENT = "HINDSIGHT_API_LLM_MAX_CONCURRENT"
|
||||
ENV_LLM_MAX_RETRIES = "HINDSIGHT_API_LLM_MAX_RETRIES"
|
||||
ENV_LLM_INITIAL_BACKOFF = "HINDSIGHT_API_LLM_INITIAL_BACKOFF"
|
||||
ENV_LLM_MAX_BACKOFF = "HINDSIGHT_API_LLM_MAX_BACKOFF"
|
||||
ENV_LLM_TIMEOUT = "HINDSIGHT_API_LLM_TIMEOUT"
|
||||
ENV_LLM_GROQ_SERVICE_TIER = "HINDSIGHT_API_LLM_GROQ_SERVICE_TIER"
|
||||
|
||||
@@ -30,14 +37,35 @@ ENV_RETAIN_LLM_PROVIDER = "HINDSIGHT_API_RETAIN_LLM_PROVIDER"
|
||||
ENV_RETAIN_LLM_API_KEY = "HINDSIGHT_API_RETAIN_LLM_API_KEY"
|
||||
ENV_RETAIN_LLM_MODEL = "HINDSIGHT_API_RETAIN_LLM_MODEL"
|
||||
ENV_RETAIN_LLM_BASE_URL = "HINDSIGHT_API_RETAIN_LLM_BASE_URL"
|
||||
ENV_RETAIN_LLM_MAX_CONCURRENT = "HINDSIGHT_API_RETAIN_LLM_MAX_CONCURRENT"
|
||||
ENV_RETAIN_LLM_MAX_RETRIES = "HINDSIGHT_API_RETAIN_LLM_MAX_RETRIES"
|
||||
ENV_RETAIN_LLM_INITIAL_BACKOFF = "HINDSIGHT_API_RETAIN_LLM_INITIAL_BACKOFF"
|
||||
ENV_RETAIN_LLM_MAX_BACKOFF = "HINDSIGHT_API_RETAIN_LLM_MAX_BACKOFF"
|
||||
ENV_RETAIN_LLM_TIMEOUT = "HINDSIGHT_API_RETAIN_LLM_TIMEOUT"
|
||||
|
||||
ENV_REFLECT_LLM_PROVIDER = "HINDSIGHT_API_REFLECT_LLM_PROVIDER"
|
||||
ENV_REFLECT_LLM_API_KEY = "HINDSIGHT_API_REFLECT_LLM_API_KEY"
|
||||
ENV_REFLECT_LLM_MODEL = "HINDSIGHT_API_REFLECT_LLM_MODEL"
|
||||
ENV_REFLECT_LLM_BASE_URL = "HINDSIGHT_API_REFLECT_LLM_BASE_URL"
|
||||
ENV_REFLECT_LLM_MAX_CONCURRENT = "HINDSIGHT_API_REFLECT_LLM_MAX_CONCURRENT"
|
||||
ENV_REFLECT_LLM_MAX_RETRIES = "HINDSIGHT_API_REFLECT_LLM_MAX_RETRIES"
|
||||
ENV_REFLECT_LLM_INITIAL_BACKOFF = "HINDSIGHT_API_REFLECT_LLM_INITIAL_BACKOFF"
|
||||
ENV_REFLECT_LLM_MAX_BACKOFF = "HINDSIGHT_API_REFLECT_LLM_MAX_BACKOFF"
|
||||
ENV_REFLECT_LLM_TIMEOUT = "HINDSIGHT_API_REFLECT_LLM_TIMEOUT"
|
||||
|
||||
ENV_CONSOLIDATION_LLM_PROVIDER = "HINDSIGHT_API_CONSOLIDATION_LLM_PROVIDER"
|
||||
ENV_CONSOLIDATION_LLM_API_KEY = "HINDSIGHT_API_CONSOLIDATION_LLM_API_KEY"
|
||||
ENV_CONSOLIDATION_LLM_MODEL = "HINDSIGHT_API_CONSOLIDATION_LLM_MODEL"
|
||||
ENV_CONSOLIDATION_LLM_BASE_URL = "HINDSIGHT_API_CONSOLIDATION_LLM_BASE_URL"
|
||||
ENV_CONSOLIDATION_LLM_MAX_CONCURRENT = "HINDSIGHT_API_CONSOLIDATION_LLM_MAX_CONCURRENT"
|
||||
ENV_CONSOLIDATION_LLM_MAX_RETRIES = "HINDSIGHT_API_CONSOLIDATION_LLM_MAX_RETRIES"
|
||||
ENV_CONSOLIDATION_LLM_INITIAL_BACKOFF = "HINDSIGHT_API_CONSOLIDATION_LLM_INITIAL_BACKOFF"
|
||||
ENV_CONSOLIDATION_LLM_MAX_BACKOFF = "HINDSIGHT_API_CONSOLIDATION_LLM_MAX_BACKOFF"
|
||||
ENV_CONSOLIDATION_LLM_TIMEOUT = "HINDSIGHT_API_CONSOLIDATION_LLM_TIMEOUT"
|
||||
|
||||
ENV_EMBEDDINGS_PROVIDER = "HINDSIGHT_API_EMBEDDINGS_PROVIDER"
|
||||
ENV_EMBEDDINGS_LOCAL_MODEL = "HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL"
|
||||
ENV_EMBEDDINGS_LOCAL_FORCE_CPU = "HINDSIGHT_API_EMBEDDINGS_LOCAL_FORCE_CPU"
|
||||
ENV_EMBEDDINGS_TEI_URL = "HINDSIGHT_API_EMBEDDINGS_TEI_URL"
|
||||
ENV_EMBEDDINGS_OPENAI_API_KEY = "HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY"
|
||||
ENV_EMBEDDINGS_OPENAI_MODEL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL"
|
||||
@@ -57,6 +85,7 @@ ENV_RERANKER_LITELLM_MODEL = "HINDSIGHT_API_RERANKER_LITELLM_MODEL"
|
||||
|
||||
ENV_RERANKER_PROVIDER = "HINDSIGHT_API_RERANKER_PROVIDER"
|
||||
ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
|
||||
ENV_RERANKER_LOCAL_FORCE_CPU = "HINDSIGHT_API_RERANKER_LOCAL_FORCE_CPU"
|
||||
ENV_RERANKER_LOCAL_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT"
|
||||
ENV_RERANKER_TEI_URL = "HINDSIGHT_API_RERANKER_TEI_URL"
|
||||
ENV_RERANKER_TEI_BATCH_SIZE = "HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE"
|
||||
@@ -68,6 +97,7 @@ ENV_RERANKER_FLASHRANK_CACHE_DIR = "HINDSIGHT_API_RERANKER_FLASHRANK_CACHE_DIR"
|
||||
ENV_HOST = "HINDSIGHT_API_HOST"
|
||||
ENV_PORT = "HINDSIGHT_API_PORT"
|
||||
ENV_LOG_LEVEL = "HINDSIGHT_API_LOG_LEVEL"
|
||||
ENV_LOG_FORMAT = "HINDSIGHT_API_LOG_FORMAT"
|
||||
ENV_WORKERS = "HINDSIGHT_API_WORKERS"
|
||||
ENV_MCP_ENABLED = "HINDSIGHT_API_MCP_ENABLED"
|
||||
ENV_GRAPH_RETRIEVER = "HINDSIGHT_API_GRAPH_RETRIEVER"
|
||||
@@ -76,17 +106,24 @@ ENV_RECALL_MAX_CONCURRENT = "HINDSIGHT_API_RECALL_MAX_CONCURRENT"
|
||||
ENV_RECALL_CONNECTION_BUDGET = "HINDSIGHT_API_RECALL_CONNECTION_BUDGET"
|
||||
ENV_MCP_LOCAL_BANK_ID = "HINDSIGHT_API_MCP_LOCAL_BANK_ID"
|
||||
ENV_MCP_INSTRUCTIONS = "HINDSIGHT_API_MCP_INSTRUCTIONS"
|
||||
ENV_MENTAL_MODEL_REFRESH_CONCURRENCY = "HINDSIGHT_API_MENTAL_MODEL_REFRESH_CONCURRENCY"
|
||||
|
||||
# Observation thresholds
|
||||
ENV_OBSERVATION_MIN_FACTS = "HINDSIGHT_API_OBSERVATION_MIN_FACTS"
|
||||
ENV_OBSERVATION_TOP_ENTITIES = "HINDSIGHT_API_OBSERVATION_TOP_ENTITIES"
|
||||
# Vertex AI configuration
|
||||
ENV_LLM_VERTEXAI_PROJECT_ID = "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID"
|
||||
ENV_LLM_VERTEXAI_REGION = "HINDSIGHT_API_LLM_VERTEXAI_REGION"
|
||||
ENV_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY = "HINDSIGHT_API_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY"
|
||||
|
||||
# Retain settings
|
||||
ENV_RETAIN_MAX_COMPLETION_TOKENS = "HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS"
|
||||
ENV_RETAIN_CHUNK_SIZE = "HINDSIGHT_API_RETAIN_CHUNK_SIZE"
|
||||
ENV_RETAIN_EXTRACT_CAUSAL_LINKS = "HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS"
|
||||
ENV_RETAIN_EXTRACTION_MODE = "HINDSIGHT_API_RETAIN_EXTRACTION_MODE"
|
||||
ENV_RETAIN_OBSERVATIONS_ASYNC = "HINDSIGHT_API_RETAIN_OBSERVATIONS_ASYNC"
|
||||
ENV_RETAIN_CUSTOM_INSTRUCTIONS = "HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS"
|
||||
|
||||
# Observations settings (consolidated knowledge from facts)
|
||||
ENV_ENABLE_OBSERVATIONS = "HINDSIGHT_API_ENABLE_OBSERVATIONS"
|
||||
ENV_CONSOLIDATION_BATCH_SIZE = "HINDSIGHT_API_CONSOLIDATION_BATCH_SIZE"
|
||||
ENV_CONSOLIDATION_MAX_TOKENS = "HINDSIGHT_API_CONSOLIDATION_MAX_TOKENS"
|
||||
|
||||
# Optimization flags
|
||||
ENV_SKIP_LLM_VERIFICATION = "HINDSIGHT_API_SKIP_LLM_VERIFICATION"
|
||||
@@ -101,25 +138,57 @@ ENV_DB_POOL_MAX_SIZE = "HINDSIGHT_API_DB_POOL_MAX_SIZE"
|
||||
ENV_DB_COMMAND_TIMEOUT = "HINDSIGHT_API_DB_COMMAND_TIMEOUT"
|
||||
ENV_DB_ACQUIRE_TIMEOUT = "HINDSIGHT_API_DB_ACQUIRE_TIMEOUT"
|
||||
|
||||
# Background task processing
|
||||
ENV_TASK_BACKEND = "HINDSIGHT_API_TASK_BACKEND"
|
||||
ENV_TASK_BACKEND_MEMORY_BATCH_SIZE = "HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_SIZE"
|
||||
ENV_TASK_BACKEND_MEMORY_BATCH_INTERVAL = "HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_INTERVAL"
|
||||
# Worker configuration (distributed task processing)
|
||||
ENV_WORKER_ENABLED = "HINDSIGHT_API_WORKER_ENABLED"
|
||||
ENV_WORKER_ID = "HINDSIGHT_API_WORKER_ID"
|
||||
ENV_WORKER_POLL_INTERVAL_MS = "HINDSIGHT_API_WORKER_POLL_INTERVAL_MS"
|
||||
ENV_WORKER_MAX_RETRIES = "HINDSIGHT_API_WORKER_MAX_RETRIES"
|
||||
ENV_WORKER_HTTP_PORT = "HINDSIGHT_API_WORKER_HTTP_PORT"
|
||||
ENV_WORKER_MAX_SLOTS = "HINDSIGHT_API_WORKER_MAX_SLOTS"
|
||||
ENV_WORKER_CONSOLIDATION_MAX_SLOTS = "HINDSIGHT_API_WORKER_CONSOLIDATION_MAX_SLOTS"
|
||||
|
||||
# Reflect agent settings
|
||||
ENV_REFLECT_MAX_ITERATIONS = "HINDSIGHT_API_REFLECT_MAX_ITERATIONS"
|
||||
|
||||
# Default values
|
||||
DEFAULT_DATABASE_URL = "pg0"
|
||||
DEFAULT_DATABASE_SCHEMA = "public"
|
||||
DEFAULT_LLM_PROVIDER = "openai"
|
||||
DEFAULT_LLM_MODEL = "gpt-5-mini"
|
||||
|
||||
# Provider-specific default models
|
||||
PROVIDER_DEFAULT_MODELS = {
|
||||
"openai": "o3-mini",
|
||||
"anthropic": "claude-haiku-4-5-20251001",
|
||||
"gemini": "gemini-2.5-flash",
|
||||
"groq": "openai/gpt-oss-120b",
|
||||
"ollama": "gemma3:12b",
|
||||
"lmstudio": "local-model",
|
||||
"vertexai": "gemini-2.0-flash-001",
|
||||
"openai-codex": "gpt-5.2-codex",
|
||||
"claude-code": "claude-sonnet-4-5-20250929",
|
||||
"mock": "mock-model",
|
||||
}
|
||||
DEFAULT_LLM_MODEL = "o3-mini" # Fallback if provider not in table
|
||||
DEFAULT_LLM_MAX_CONCURRENT = 32
|
||||
DEFAULT_LLM_MAX_RETRIES = 10 # Max retry attempts for LLM API calls
|
||||
DEFAULT_LLM_INITIAL_BACKOFF = 1.0 # Initial backoff in seconds for retry exponential backoff
|
||||
DEFAULT_LLM_MAX_BACKOFF = 60.0 # Max backoff cap in seconds for retry exponential backoff
|
||||
DEFAULT_LLM_TIMEOUT = 120.0 # seconds
|
||||
|
||||
# Vertex AI defaults
|
||||
DEFAULT_LLM_VERTEXAI_PROJECT_ID = None # Required for Vertex AI
|
||||
DEFAULT_LLM_VERTEXAI_REGION = "us-central1"
|
||||
DEFAULT_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY = None # Optional, uses ADC if not set
|
||||
|
||||
DEFAULT_EMBEDDINGS_PROVIDER = "local"
|
||||
DEFAULT_EMBEDDINGS_LOCAL_MODEL = "BAAI/bge-small-en-v1.5"
|
||||
DEFAULT_EMBEDDINGS_LOCAL_FORCE_CPU = False # Force CPU mode for local embeddings (avoids MPS/XPC issues on macOS)
|
||||
DEFAULT_EMBEDDINGS_OPENAI_MODEL = "text-embedding-3-small"
|
||||
DEFAULT_EMBEDDING_DIMENSION = 384
|
||||
|
||||
DEFAULT_RERANKER_PROVIDER = "local"
|
||||
DEFAULT_RERANKER_LOCAL_MODEL = "cross-encoder/ms-marco-MiniLM-L-6-v2"
|
||||
DEFAULT_RERANKER_LOCAL_FORCE_CPU = False # Force CPU mode for local reranker (avoids MPS/XPC issues on macOS)
|
||||
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT = 4 # Limit concurrent CPU-bound reranking to prevent thrashing
|
||||
DEFAULT_RERANKER_TEI_BATCH_SIZE = 128
|
||||
DEFAULT_RERANKER_TEI_MAX_CONCURRENT = 8
|
||||
@@ -138,6 +207,7 @@ DEFAULT_RERANKER_LITELLM_MODEL = "cohere/rerank-english-v3.0"
|
||||
DEFAULT_HOST = "0.0.0.0"
|
||||
DEFAULT_PORT = 8888
|
||||
DEFAULT_LOG_LEVEL = "info"
|
||||
DEFAULT_LOG_FORMAT = "text" # Options: "text", "json"
|
||||
DEFAULT_WORKERS = 1
|
||||
DEFAULT_MCP_ENABLED = True
|
||||
DEFAULT_GRAPH_RETRIEVER = "link_expansion" # Options: "link_expansion", "mpfp", "bfs"
|
||||
@@ -145,18 +215,20 @@ DEFAULT_MPFP_TOP_K_NEIGHBORS = 20 # Fan-out limit per node in MPFP graph traver
|
||||
DEFAULT_RECALL_MAX_CONCURRENT = 32 # Max concurrent recall operations per worker
|
||||
DEFAULT_RECALL_CONNECTION_BUDGET = 4 # Max concurrent DB connections per recall operation
|
||||
DEFAULT_MCP_LOCAL_BANK_ID = "mcp"
|
||||
|
||||
# Observation thresholds
|
||||
DEFAULT_OBSERVATION_MIN_FACTS = 5 # Min facts required to generate entity observations
|
||||
DEFAULT_OBSERVATION_TOP_ENTITIES = 5 # Max entities to process per retain batch
|
||||
DEFAULT_MENTAL_MODEL_REFRESH_CONCURRENCY = 8 # Max concurrent mental model refreshes
|
||||
|
||||
# Retain settings
|
||||
DEFAULT_RETAIN_MAX_COMPLETION_TOKENS = 64000 # Max tokens for fact extraction LLM call
|
||||
DEFAULT_RETAIN_CHUNK_SIZE = 3000 # Max chars per chunk for fact extraction
|
||||
DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS = True # Extract causal links between facts
|
||||
DEFAULT_RETAIN_EXTRACTION_MODE = "concise" # Extraction mode: "concise" or "verbose"
|
||||
RETAIN_EXTRACTION_MODES = ("concise", "verbose") # Allowed extraction modes
|
||||
DEFAULT_RETAIN_OBSERVATIONS_ASYNC = False # Run observation generation async (after retain completes)
|
||||
DEFAULT_RETAIN_EXTRACTION_MODE = "concise" # Extraction mode: "concise", "verbose", or "custom"
|
||||
RETAIN_EXTRACTION_MODES = ("concise", "verbose", "custom") # Allowed extraction modes
|
||||
DEFAULT_RETAIN_CUSTOM_INSTRUCTIONS = None # Custom extraction guidelines (only used when mode="custom")
|
||||
|
||||
# Observations defaults (consolidated knowledge from facts)
|
||||
DEFAULT_ENABLE_OBSERVATIONS = True # Observations enabled by default
|
||||
DEFAULT_CONSOLIDATION_BATCH_SIZE = 50 # Memories to load per batch (internal memory optimization)
|
||||
DEFAULT_CONSOLIDATION_MAX_TOKENS = 1024 # Max tokens for recall when finding related observations
|
||||
|
||||
# Database migrations
|
||||
DEFAULT_RUN_MIGRATIONS_ON_STARTUP = True
|
||||
@@ -167,10 +239,17 @@ DEFAULT_DB_POOL_MAX_SIZE = 100
|
||||
DEFAULT_DB_COMMAND_TIMEOUT = 60 # seconds
|
||||
DEFAULT_DB_ACQUIRE_TIMEOUT = 30 # seconds
|
||||
|
||||
# Background task processing
|
||||
DEFAULT_TASK_BACKEND = "memory" # Options: "memory", "noop"
|
||||
DEFAULT_TASK_BACKEND_MEMORY_BATCH_SIZE = 10
|
||||
DEFAULT_TASK_BACKEND_MEMORY_BATCH_INTERVAL = 1.0 # seconds
|
||||
# Worker configuration (distributed task processing)
|
||||
DEFAULT_WORKER_ENABLED = True # API runs worker by default (standalone mode)
|
||||
DEFAULT_WORKER_ID = None # Will use hostname if not specified
|
||||
DEFAULT_WORKER_POLL_INTERVAL_MS = 500 # Poll database every 500ms
|
||||
DEFAULT_WORKER_MAX_RETRIES = 3 # Max retries before marking task failed
|
||||
DEFAULT_WORKER_HTTP_PORT = 8889 # HTTP port for worker metrics/health
|
||||
DEFAULT_WORKER_MAX_SLOTS = 10 # Total concurrent tasks per worker
|
||||
DEFAULT_WORKER_CONSOLIDATION_MAX_SLOTS = 2 # Max concurrent consolidation tasks per worker
|
||||
|
||||
# Reflect agent settings
|
||||
DEFAULT_REFLECT_MAX_ITERATIONS = 10 # Max tool call iterations before forcing response
|
||||
|
||||
# Default MCP tool descriptions (can be customized via env vars)
|
||||
DEFAULT_MCP_RETAIN_DESCRIPTION = """Store important information to long-term memory.
|
||||
@@ -196,6 +275,36 @@ Use this tool PROACTIVELY to:
|
||||
EMBEDDING_DIMENSION = DEFAULT_EMBEDDING_DIMENSION
|
||||
|
||||
|
||||
class JsonFormatter(logging.Formatter):
|
||||
"""JSON formatter for structured logging.
|
||||
|
||||
Outputs logs in JSON format with a 'severity' field that cloud logging
|
||||
systems (GCP, AWS CloudWatch, etc.) can parse to correctly categorize log levels.
|
||||
"""
|
||||
|
||||
SEVERITY_MAP = {
|
||||
logging.DEBUG: "DEBUG",
|
||||
logging.INFO: "INFO",
|
||||
logging.WARNING: "WARNING",
|
||||
logging.ERROR: "ERROR",
|
||||
logging.CRITICAL: "CRITICAL",
|
||||
}
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
log_entry = {
|
||||
"severity": self.SEVERITY_MAP.get(record.levelno, "DEFAULT"),
|
||||
"message": record.getMessage(),
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"logger": record.name,
|
||||
}
|
||||
|
||||
# Add exception info if present
|
||||
if record.exc_info:
|
||||
log_entry["exception"] = self.formatException(record.exc_info)
|
||||
|
||||
return json.dumps(log_entry)
|
||||
|
||||
|
||||
def _validate_extraction_mode(mode: str) -> str:
|
||||
"""Validate and normalize extraction mode."""
|
||||
mode_lower = mode.lower()
|
||||
@@ -208,12 +317,18 @@ def _validate_extraction_mode(mode: str) -> str:
|
||||
return mode_lower
|
||||
|
||||
|
||||
def _get_default_model_for_provider(provider: str) -> str:
|
||||
"""Get the default model for a given provider."""
|
||||
return PROVIDER_DEFAULT_MODELS.get(provider.lower(), DEFAULT_LLM_MODEL)
|
||||
|
||||
|
||||
@dataclass
|
||||
class HindsightConfig:
|
||||
"""Configuration container for Hindsight API."""
|
||||
|
||||
# Database
|
||||
database_url: str
|
||||
database_schema: str
|
||||
|
||||
# LLM (default, used as fallback for per-operation config)
|
||||
llm_provider: str
|
||||
@@ -221,22 +336,51 @@ class HindsightConfig:
|
||||
llm_model: str
|
||||
llm_base_url: str | None
|
||||
llm_max_concurrent: int
|
||||
llm_max_retries: int
|
||||
llm_initial_backoff: float
|
||||
llm_max_backoff: float
|
||||
llm_timeout: float
|
||||
|
||||
# Vertex AI configuration
|
||||
llm_vertexai_project_id: str | None
|
||||
llm_vertexai_region: str
|
||||
llm_vertexai_service_account_key: str | None
|
||||
|
||||
# Per-operation LLM configuration (None = use default LLM config)
|
||||
retain_llm_provider: str | None
|
||||
retain_llm_api_key: str | None
|
||||
retain_llm_model: str | None
|
||||
retain_llm_base_url: str | None
|
||||
retain_llm_max_concurrent: int | None
|
||||
retain_llm_max_retries: int | None
|
||||
retain_llm_initial_backoff: float | None
|
||||
retain_llm_max_backoff: float | None
|
||||
retain_llm_timeout: float | None
|
||||
|
||||
reflect_llm_provider: str | None
|
||||
reflect_llm_api_key: str | None
|
||||
reflect_llm_model: str | None
|
||||
reflect_llm_base_url: str | None
|
||||
reflect_llm_max_concurrent: int | None
|
||||
reflect_llm_max_retries: int | None
|
||||
reflect_llm_initial_backoff: float | None
|
||||
reflect_llm_max_backoff: float | None
|
||||
reflect_llm_timeout: float | None
|
||||
|
||||
consolidation_llm_provider: str | None
|
||||
consolidation_llm_api_key: str | None
|
||||
consolidation_llm_model: str | None
|
||||
consolidation_llm_base_url: str | None
|
||||
consolidation_llm_max_concurrent: int | None
|
||||
consolidation_llm_max_retries: int | None
|
||||
consolidation_llm_initial_backoff: float | None
|
||||
consolidation_llm_max_backoff: float | None
|
||||
consolidation_llm_timeout: float | None
|
||||
|
||||
# Embeddings
|
||||
embeddings_provider: str
|
||||
embeddings_local_model: str
|
||||
embeddings_local_force_cpu: bool
|
||||
embeddings_tei_url: str | None
|
||||
embeddings_openai_base_url: str | None
|
||||
embeddings_cohere_base_url: str | None
|
||||
@@ -244,6 +388,8 @@ class HindsightConfig:
|
||||
# Reranker
|
||||
reranker_provider: str
|
||||
reranker_local_model: str
|
||||
reranker_local_force_cpu: bool
|
||||
reranker_local_max_concurrent: int
|
||||
reranker_tei_url: str | None
|
||||
reranker_tei_batch_size: int
|
||||
reranker_tei_max_concurrent: int
|
||||
@@ -254,6 +400,7 @@ class HindsightConfig:
|
||||
host: str
|
||||
port: int
|
||||
log_level: str
|
||||
log_format: str
|
||||
mcp_enabled: bool
|
||||
|
||||
# Recall
|
||||
@@ -261,17 +408,19 @@ class HindsightConfig:
|
||||
mpfp_top_k_neighbors: int
|
||||
recall_max_concurrent: int
|
||||
recall_connection_budget: int
|
||||
|
||||
# Observation thresholds
|
||||
observation_min_facts: int
|
||||
observation_top_entities: int
|
||||
mental_model_refresh_concurrency: int
|
||||
|
||||
# Retain settings
|
||||
retain_max_completion_tokens: int
|
||||
retain_chunk_size: int
|
||||
retain_extract_causal_links: bool
|
||||
retain_extraction_mode: str
|
||||
retain_observations_async: bool
|
||||
retain_custom_instructions: str | None
|
||||
|
||||
# Observations settings (consolidated knowledge from facts)
|
||||
enable_observations: bool
|
||||
consolidation_batch_size: int
|
||||
consolidation_max_tokens: int
|
||||
|
||||
# Optimization flags
|
||||
skip_llm_verification: bool
|
||||
@@ -286,42 +435,151 @@ class HindsightConfig:
|
||||
db_command_timeout: int
|
||||
db_acquire_timeout: int
|
||||
|
||||
# Background task processing
|
||||
task_backend: str
|
||||
task_backend_memory_batch_size: int
|
||||
task_backend_memory_batch_interval: float
|
||||
# Worker configuration (distributed task processing)
|
||||
worker_enabled: bool
|
||||
worker_id: str | None
|
||||
worker_poll_interval_ms: int
|
||||
worker_max_retries: int
|
||||
worker_http_port: int
|
||||
worker_max_slots: int
|
||||
worker_consolidation_max_slots: int
|
||||
|
||||
# Reflect agent settings
|
||||
reflect_max_iterations: int
|
||||
|
||||
def validate(self) -> None:
|
||||
"""Validate configuration values and raise errors for invalid combinations."""
|
||||
# RETAIN_MAX_COMPLETION_TOKENS must be greater than RETAIN_CHUNK_SIZE
|
||||
# to ensure the LLM has enough output capacity to extract facts from chunks
|
||||
if 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})"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> "HindsightConfig":
|
||||
"""Create configuration from environment variables."""
|
||||
return cls(
|
||||
# Get provider first to determine default model
|
||||
llm_provider = os.getenv(ENV_LLM_PROVIDER, DEFAULT_LLM_PROVIDER)
|
||||
llm_model = os.getenv(ENV_LLM_MODEL) or _get_default_model_for_provider(llm_provider)
|
||||
|
||||
config = cls(
|
||||
# Database
|
||||
database_url=os.getenv(ENV_DATABASE_URL, DEFAULT_DATABASE_URL),
|
||||
database_schema=os.getenv(ENV_DATABASE_SCHEMA, DEFAULT_DATABASE_SCHEMA),
|
||||
# LLM
|
||||
llm_provider=os.getenv(ENV_LLM_PROVIDER, DEFAULT_LLM_PROVIDER),
|
||||
llm_provider=llm_provider,
|
||||
llm_api_key=os.getenv(ENV_LLM_API_KEY),
|
||||
llm_model=os.getenv(ENV_LLM_MODEL, DEFAULT_LLM_MODEL),
|
||||
llm_model=llm_model,
|
||||
llm_base_url=os.getenv(ENV_LLM_BASE_URL) or None,
|
||||
llm_max_concurrent=int(os.getenv(ENV_LLM_MAX_CONCURRENT, str(DEFAULT_LLM_MAX_CONCURRENT))),
|
||||
llm_max_retries=int(os.getenv(ENV_LLM_MAX_RETRIES, str(DEFAULT_LLM_MAX_RETRIES))),
|
||||
llm_initial_backoff=float(os.getenv(ENV_LLM_INITIAL_BACKOFF, str(DEFAULT_LLM_INITIAL_BACKOFF))),
|
||||
llm_max_backoff=float(os.getenv(ENV_LLM_MAX_BACKOFF, str(DEFAULT_LLM_MAX_BACKOFF))),
|
||||
llm_timeout=float(os.getenv(ENV_LLM_TIMEOUT, str(DEFAULT_LLM_TIMEOUT))),
|
||||
# Vertex AI
|
||||
llm_vertexai_project_id=os.getenv(ENV_LLM_VERTEXAI_PROJECT_ID) or DEFAULT_LLM_VERTEXAI_PROJECT_ID,
|
||||
llm_vertexai_region=os.getenv(ENV_LLM_VERTEXAI_REGION, DEFAULT_LLM_VERTEXAI_REGION),
|
||||
llm_vertexai_service_account_key=os.getenv(ENV_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY)
|
||||
or DEFAULT_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY,
|
||||
# Per-operation LLM config (None = use default)
|
||||
retain_llm_provider=os.getenv(ENV_RETAIN_LLM_PROVIDER) or None,
|
||||
retain_llm_api_key=os.getenv(ENV_RETAIN_LLM_API_KEY) or None,
|
||||
retain_llm_model=os.getenv(ENV_RETAIN_LLM_MODEL) or None,
|
||||
retain_llm_model=os.getenv(ENV_RETAIN_LLM_MODEL)
|
||||
or (
|
||||
_get_default_model_for_provider(os.getenv(ENV_RETAIN_LLM_PROVIDER))
|
||||
if os.getenv(ENV_RETAIN_LLM_PROVIDER)
|
||||
else None
|
||||
),
|
||||
retain_llm_base_url=os.getenv(ENV_RETAIN_LLM_BASE_URL) or None,
|
||||
retain_llm_max_concurrent=int(os.getenv(ENV_RETAIN_LLM_MAX_CONCURRENT))
|
||||
if os.getenv(ENV_RETAIN_LLM_MAX_CONCURRENT)
|
||||
else None,
|
||||
retain_llm_max_retries=int(os.getenv(ENV_RETAIN_LLM_MAX_RETRIES))
|
||||
if os.getenv(ENV_RETAIN_LLM_MAX_RETRIES)
|
||||
else None,
|
||||
retain_llm_initial_backoff=float(os.getenv(ENV_RETAIN_LLM_INITIAL_BACKOFF))
|
||||
if os.getenv(ENV_RETAIN_LLM_INITIAL_BACKOFF)
|
||||
else None,
|
||||
retain_llm_max_backoff=float(os.getenv(ENV_RETAIN_LLM_MAX_BACKOFF))
|
||||
if os.getenv(ENV_RETAIN_LLM_MAX_BACKOFF)
|
||||
else None,
|
||||
retain_llm_timeout=float(os.getenv(ENV_RETAIN_LLM_TIMEOUT)) if os.getenv(ENV_RETAIN_LLM_TIMEOUT) else None,
|
||||
reflect_llm_provider=os.getenv(ENV_REFLECT_LLM_PROVIDER) or None,
|
||||
reflect_llm_api_key=os.getenv(ENV_REFLECT_LLM_API_KEY) or None,
|
||||
reflect_llm_model=os.getenv(ENV_REFLECT_LLM_MODEL) or None,
|
||||
reflect_llm_model=os.getenv(ENV_REFLECT_LLM_MODEL)
|
||||
or (
|
||||
_get_default_model_for_provider(os.getenv(ENV_REFLECT_LLM_PROVIDER))
|
||||
if os.getenv(ENV_REFLECT_LLM_PROVIDER)
|
||||
else None
|
||||
),
|
||||
reflect_llm_base_url=os.getenv(ENV_REFLECT_LLM_BASE_URL) or None,
|
||||
reflect_llm_max_concurrent=int(os.getenv(ENV_REFLECT_LLM_MAX_CONCURRENT))
|
||||
if os.getenv(ENV_REFLECT_LLM_MAX_CONCURRENT)
|
||||
else None,
|
||||
reflect_llm_max_retries=int(os.getenv(ENV_REFLECT_LLM_MAX_RETRIES))
|
||||
if os.getenv(ENV_REFLECT_LLM_MAX_RETRIES)
|
||||
else None,
|
||||
reflect_llm_initial_backoff=float(os.getenv(ENV_REFLECT_LLM_INITIAL_BACKOFF))
|
||||
if os.getenv(ENV_REFLECT_LLM_INITIAL_BACKOFF)
|
||||
else None,
|
||||
reflect_llm_max_backoff=float(os.getenv(ENV_REFLECT_LLM_MAX_BACKOFF))
|
||||
if os.getenv(ENV_REFLECT_LLM_MAX_BACKOFF)
|
||||
else None,
|
||||
reflect_llm_timeout=float(os.getenv(ENV_REFLECT_LLM_TIMEOUT))
|
||||
if os.getenv(ENV_REFLECT_LLM_TIMEOUT)
|
||||
else None,
|
||||
consolidation_llm_provider=os.getenv(ENV_CONSOLIDATION_LLM_PROVIDER) or None,
|
||||
consolidation_llm_api_key=os.getenv(ENV_CONSOLIDATION_LLM_API_KEY) or None,
|
||||
consolidation_llm_model=os.getenv(ENV_CONSOLIDATION_LLM_MODEL)
|
||||
or (
|
||||
_get_default_model_for_provider(os.getenv(ENV_CONSOLIDATION_LLM_PROVIDER))
|
||||
if os.getenv(ENV_CONSOLIDATION_LLM_PROVIDER)
|
||||
else None
|
||||
),
|
||||
consolidation_llm_base_url=os.getenv(ENV_CONSOLIDATION_LLM_BASE_URL) or None,
|
||||
consolidation_llm_max_concurrent=int(os.getenv(ENV_CONSOLIDATION_LLM_MAX_CONCURRENT))
|
||||
if os.getenv(ENV_CONSOLIDATION_LLM_MAX_CONCURRENT)
|
||||
else None,
|
||||
consolidation_llm_max_retries=int(os.getenv(ENV_CONSOLIDATION_LLM_MAX_RETRIES))
|
||||
if os.getenv(ENV_CONSOLIDATION_LLM_MAX_RETRIES)
|
||||
else None,
|
||||
consolidation_llm_initial_backoff=float(os.getenv(ENV_CONSOLIDATION_LLM_INITIAL_BACKOFF))
|
||||
if os.getenv(ENV_CONSOLIDATION_LLM_INITIAL_BACKOFF)
|
||||
else None,
|
||||
consolidation_llm_max_backoff=float(os.getenv(ENV_CONSOLIDATION_LLM_MAX_BACKOFF))
|
||||
if os.getenv(ENV_CONSOLIDATION_LLM_MAX_BACKOFF)
|
||||
else None,
|
||||
consolidation_llm_timeout=float(os.getenv(ENV_CONSOLIDATION_LLM_TIMEOUT))
|
||||
if os.getenv(ENV_CONSOLIDATION_LLM_TIMEOUT)
|
||||
else None,
|
||||
# Embeddings
|
||||
embeddings_provider=os.getenv(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER),
|
||||
embeddings_local_model=os.getenv(ENV_EMBEDDINGS_LOCAL_MODEL, DEFAULT_EMBEDDINGS_LOCAL_MODEL),
|
||||
embeddings_local_force_cpu=os.getenv(
|
||||
ENV_EMBEDDINGS_LOCAL_FORCE_CPU, str(DEFAULT_EMBEDDINGS_LOCAL_FORCE_CPU)
|
||||
).lower()
|
||||
in ("true", "1"),
|
||||
embeddings_tei_url=os.getenv(ENV_EMBEDDINGS_TEI_URL),
|
||||
embeddings_openai_base_url=os.getenv(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None,
|
||||
embeddings_cohere_base_url=os.getenv(ENV_EMBEDDINGS_COHERE_BASE_URL) or None,
|
||||
# Reranker
|
||||
reranker_provider=os.getenv(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER),
|
||||
reranker_local_model=os.getenv(ENV_RERANKER_LOCAL_MODEL, DEFAULT_RERANKER_LOCAL_MODEL),
|
||||
reranker_local_force_cpu=os.getenv(
|
||||
ENV_RERANKER_LOCAL_FORCE_CPU, str(DEFAULT_RERANKER_LOCAL_FORCE_CPU)
|
||||
).lower()
|
||||
in ("true", "1"),
|
||||
reranker_local_max_concurrent=int(
|
||||
os.getenv(ENV_RERANKER_LOCAL_MAX_CONCURRENT, str(DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT))
|
||||
),
|
||||
reranker_tei_url=os.getenv(ENV_RERANKER_TEI_URL),
|
||||
reranker_tei_batch_size=int(os.getenv(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE))),
|
||||
reranker_tei_max_concurrent=int(
|
||||
@@ -333,6 +591,7 @@ class HindsightConfig:
|
||||
host=os.getenv(ENV_HOST, DEFAULT_HOST),
|
||||
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
|
||||
log_level=os.getenv(ENV_LOG_LEVEL, DEFAULT_LOG_LEVEL),
|
||||
log_format=os.getenv(ENV_LOG_FORMAT, DEFAULT_LOG_FORMAT).lower(),
|
||||
mcp_enabled=os.getenv(ENV_MCP_ENABLED, str(DEFAULT_MCP_ENABLED)).lower() == "true",
|
||||
# Recall
|
||||
graph_retriever=os.getenv(ENV_GRAPH_RETRIEVER, DEFAULT_GRAPH_RETRIEVER),
|
||||
@@ -341,14 +600,12 @@ class HindsightConfig:
|
||||
recall_connection_budget=int(
|
||||
os.getenv(ENV_RECALL_CONNECTION_BUDGET, str(DEFAULT_RECALL_CONNECTION_BUDGET))
|
||||
),
|
||||
mental_model_refresh_concurrency=int(
|
||||
os.getenv(ENV_MENTAL_MODEL_REFRESH_CONCURRENCY, str(DEFAULT_MENTAL_MODEL_REFRESH_CONCURRENCY))
|
||||
),
|
||||
# Optimization flags
|
||||
skip_llm_verification=os.getenv(ENV_SKIP_LLM_VERIFICATION, "false").lower() == "true",
|
||||
lazy_reranker=os.getenv(ENV_LAZY_RERANKER, "false").lower() == "true",
|
||||
# Observation thresholds
|
||||
observation_min_facts=int(os.getenv(ENV_OBSERVATION_MIN_FACTS, str(DEFAULT_OBSERVATION_MIN_FACTS))),
|
||||
observation_top_entities=int(
|
||||
os.getenv(ENV_OBSERVATION_TOP_ENTITIES, str(DEFAULT_OBSERVATION_TOP_ENTITIES))
|
||||
),
|
||||
# Retain settings
|
||||
retain_max_completion_tokens=int(
|
||||
os.getenv(ENV_RETAIN_MAX_COMPLETION_TOKENS, str(DEFAULT_RETAIN_MAX_COMPLETION_TOKENS))
|
||||
@@ -361,10 +618,15 @@ class HindsightConfig:
|
||||
retain_extraction_mode=_validate_extraction_mode(
|
||||
os.getenv(ENV_RETAIN_EXTRACTION_MODE, DEFAULT_RETAIN_EXTRACTION_MODE)
|
||||
),
|
||||
retain_observations_async=os.getenv(
|
||||
ENV_RETAIN_OBSERVATIONS_ASYNC, str(DEFAULT_RETAIN_OBSERVATIONS_ASYNC)
|
||||
).lower()
|
||||
== "true",
|
||||
retain_custom_instructions=os.getenv(ENV_RETAIN_CUSTOM_INSTRUCTIONS) or DEFAULT_RETAIN_CUSTOM_INSTRUCTIONS,
|
||||
# Observations settings (consolidated knowledge from facts)
|
||||
enable_observations=os.getenv(ENV_ENABLE_OBSERVATIONS, str(DEFAULT_ENABLE_OBSERVATIONS)).lower() == "true",
|
||||
consolidation_batch_size=int(
|
||||
os.getenv(ENV_CONSOLIDATION_BATCH_SIZE, str(DEFAULT_CONSOLIDATION_BATCH_SIZE))
|
||||
),
|
||||
consolidation_max_tokens=int(
|
||||
os.getenv(ENV_CONSOLIDATION_MAX_TOKENS, str(DEFAULT_CONSOLIDATION_MAX_TOKENS))
|
||||
),
|
||||
# Database migrations
|
||||
run_migrations_on_startup=os.getenv(ENV_RUN_MIGRATIONS_ON_STARTUP, "true").lower() == "true",
|
||||
# Database connection pool
|
||||
@@ -372,15 +634,21 @@ class HindsightConfig:
|
||||
db_pool_max_size=int(os.getenv(ENV_DB_POOL_MAX_SIZE, str(DEFAULT_DB_POOL_MAX_SIZE))),
|
||||
db_command_timeout=int(os.getenv(ENV_DB_COMMAND_TIMEOUT, str(DEFAULT_DB_COMMAND_TIMEOUT))),
|
||||
db_acquire_timeout=int(os.getenv(ENV_DB_ACQUIRE_TIMEOUT, str(DEFAULT_DB_ACQUIRE_TIMEOUT))),
|
||||
# Background task processing
|
||||
task_backend=os.getenv(ENV_TASK_BACKEND, DEFAULT_TASK_BACKEND),
|
||||
task_backend_memory_batch_size=int(
|
||||
os.getenv(ENV_TASK_BACKEND_MEMORY_BATCH_SIZE, str(DEFAULT_TASK_BACKEND_MEMORY_BATCH_SIZE))
|
||||
),
|
||||
task_backend_memory_batch_interval=float(
|
||||
os.getenv(ENV_TASK_BACKEND_MEMORY_BATCH_INTERVAL, str(DEFAULT_TASK_BACKEND_MEMORY_BATCH_INTERVAL))
|
||||
# Worker configuration
|
||||
worker_enabled=os.getenv(ENV_WORKER_ENABLED, str(DEFAULT_WORKER_ENABLED)).lower() == "true",
|
||||
worker_id=os.getenv(ENV_WORKER_ID) or DEFAULT_WORKER_ID,
|
||||
worker_poll_interval_ms=int(os.getenv(ENV_WORKER_POLL_INTERVAL_MS, str(DEFAULT_WORKER_POLL_INTERVAL_MS))),
|
||||
worker_max_retries=int(os.getenv(ENV_WORKER_MAX_RETRIES, str(DEFAULT_WORKER_MAX_RETRIES))),
|
||||
worker_http_port=int(os.getenv(ENV_WORKER_HTTP_PORT, str(DEFAULT_WORKER_HTTP_PORT))),
|
||||
worker_max_slots=int(os.getenv(ENV_WORKER_MAX_SLOTS, str(DEFAULT_WORKER_MAX_SLOTS))),
|
||||
worker_consolidation_max_slots=int(
|
||||
os.getenv(ENV_WORKER_CONSOLIDATION_MAX_SLOTS, str(DEFAULT_WORKER_CONSOLIDATION_MAX_SLOTS))
|
||||
),
|
||||
# Reflect agent settings
|
||||
reflect_max_iterations=int(os.getenv(ENV_REFLECT_MAX_ITERATIONS, str(DEFAULT_REFLECT_MAX_ITERATIONS))),
|
||||
)
|
||||
config.validate()
|
||||
return config
|
||||
|
||||
def get_llm_base_url(self) -> str:
|
||||
"""Get the LLM base URL, with provider-specific defaults."""
|
||||
@@ -410,16 +678,32 @@ class HindsightConfig:
|
||||
return log_level_map.get(self.log_level.lower(), logging.INFO)
|
||||
|
||||
def configure_logging(self) -> None:
|
||||
"""Configure Python logging based on the log level."""
|
||||
logging.basicConfig(
|
||||
level=self.get_python_log_level(),
|
||||
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
|
||||
force=True, # Override any existing configuration
|
||||
)
|
||||
"""Configure Python logging based on the log level and format.
|
||||
|
||||
When log_format is "json", outputs structured JSON logs with a severity
|
||||
field that GCP Cloud Logging can parse for proper log level categorization.
|
||||
"""
|
||||
root_logger = logging.getLogger()
|
||||
root_logger.setLevel(self.get_python_log_level())
|
||||
|
||||
# Remove existing handlers
|
||||
for handler in root_logger.handlers[:]:
|
||||
root_logger.removeHandler(handler)
|
||||
|
||||
# Create handler writing to stdout (GCP treats stderr as ERROR)
|
||||
handler = logging.StreamHandler(sys.stdout)
|
||||
handler.setLevel(self.get_python_log_level())
|
||||
|
||||
if self.log_format == "json":
|
||||
handler.setFormatter(JsonFormatter())
|
||||
else:
|
||||
handler.setFormatter(logging.Formatter("%(asctime)s - %(levelname)s - %(name)s - %(message)s"))
|
||||
|
||||
root_logger.addHandler(handler)
|
||||
|
||||
def log_config(self) -> None:
|
||||
"""Log the current configuration (without sensitive values)."""
|
||||
logger.info(f"Database: {self.database_url}")
|
||||
logger.info(f"Database: {self.database_url} (schema: {self.database_schema})")
|
||||
logger.info(f"LLM: provider={self.llm_provider}, model={self.llm_model}")
|
||||
if self.retain_llm_provider or self.retain_llm_model:
|
||||
retain_provider = self.retain_llm_provider or self.llm_provider
|
||||
@@ -429,6 +713,10 @@ class HindsightConfig:
|
||||
reflect_provider = self.reflect_llm_provider or self.llm_provider
|
||||
reflect_model = self.reflect_llm_model or self.llm_model
|
||||
logger.info(f"LLM (reflect): provider={reflect_provider}, model={reflect_model}")
|
||||
if self.consolidation_llm_provider or self.consolidation_llm_model:
|
||||
consolidation_provider = self.consolidation_llm_provider or self.llm_provider
|
||||
consolidation_model = self.consolidation_llm_model or self.llm_model
|
||||
logger.info(f"LLM (consolidation): provider={consolidation_provider}, model={consolidation_model}")
|
||||
logger.info(f"Embeddings: provider={self.embeddings_provider}")
|
||||
logger.info(f"Reranker: provider={self.reranker_provider}")
|
||||
logger.info(f"Graph retriever: {self.graph_retriever}")
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
"""
|
||||
Daemon mode support for Hindsight API.
|
||||
|
||||
Provides idle timeout and lockfile management for running as a background daemon.
|
||||
Provides idle timeout for running as a background daemon.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import fcntl
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
@@ -15,10 +14,11 @@ from pathlib import Path
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Default daemon configuration
|
||||
DEFAULT_DAEMON_PORT = 8889
|
||||
DEFAULT_DAEMON_PORT = 8888
|
||||
DEFAULT_IDLE_TIMEOUT = 0 # 0 = no auto-exit (hindsight-embed passes its own timeout)
|
||||
LOCKFILE_PATH = Path.home() / ".hindsight" / "daemon.lock"
|
||||
DAEMON_LOG_PATH = Path.home() / ".hindsight" / "daemon.log"
|
||||
|
||||
# Allow override via environment variable for profile-specific logs
|
||||
DAEMON_LOG_PATH = Path(os.getenv("HINDSIGHT_API_DAEMON_LOG", str(Path.home() / ".hindsight" / "daemon.log")))
|
||||
|
||||
|
||||
class IdleTimeoutMiddleware:
|
||||
@@ -52,82 +52,10 @@ class IdleTimeoutMiddleware:
|
||||
logger.info(f"Idle timeout reached ({self.idle_timeout}s), shutting down daemon")
|
||||
# Give a moment for any in-flight requests
|
||||
await asyncio.sleep(1)
|
||||
os._exit(0)
|
||||
# Send SIGTERM to ourselves to trigger graceful shutdown
|
||||
import signal
|
||||
|
||||
|
||||
class DaemonLock:
|
||||
"""
|
||||
File-based lock to prevent multiple daemon instances.
|
||||
|
||||
Uses fcntl.flock for atomic locking on Unix systems.
|
||||
"""
|
||||
|
||||
def __init__(self, lockfile: Path = LOCKFILE_PATH):
|
||||
self.lockfile = lockfile
|
||||
self._fd = None
|
||||
|
||||
def acquire(self) -> bool:
|
||||
"""
|
||||
Try to acquire the daemon lock.
|
||||
|
||||
Returns True if lock acquired, False if another daemon is running.
|
||||
"""
|
||||
self.lockfile.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
try:
|
||||
self._fd = open(self.lockfile, "w")
|
||||
fcntl.flock(self._fd.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||
# Write PID for debugging
|
||||
self._fd.write(str(os.getpid()))
|
||||
self._fd.flush()
|
||||
return True
|
||||
except (IOError, OSError):
|
||||
# Lock is held by another process
|
||||
if self._fd:
|
||||
self._fd.close()
|
||||
self._fd = None
|
||||
return False
|
||||
|
||||
def release(self):
|
||||
"""Release the daemon lock."""
|
||||
if self._fd:
|
||||
try:
|
||||
fcntl.flock(self._fd.fileno(), fcntl.LOCK_UN)
|
||||
self._fd.close()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
self._fd = None
|
||||
# Remove lockfile
|
||||
try:
|
||||
self.lockfile.unlink()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def is_locked(self) -> bool:
|
||||
"""Check if the lock is held by another process."""
|
||||
if not self.lockfile.exists():
|
||||
return False
|
||||
|
||||
try:
|
||||
fd = open(self.lockfile, "r")
|
||||
fcntl.flock(fd.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||
# We got the lock, so no one else has it
|
||||
fcntl.flock(fd.fileno(), fcntl.LOCK_UN)
|
||||
fd.close()
|
||||
return False
|
||||
except (IOError, OSError):
|
||||
return True
|
||||
|
||||
def get_pid(self) -> int | None:
|
||||
"""Get the PID of the daemon holding the lock."""
|
||||
if not self.lockfile.exists():
|
||||
return None
|
||||
try:
|
||||
with open(self.lockfile, "r") as f:
|
||||
return int(f.read().strip())
|
||||
except (ValueError, IOError):
|
||||
return None
|
||||
os.kill(os.getpid(), signal.SIGTERM)
|
||||
|
||||
|
||||
def daemonize():
|
||||
@@ -136,16 +64,21 @@ def daemonize():
|
||||
|
||||
Uses double-fork technique to properly detach from terminal.
|
||||
"""
|
||||
# First fork
|
||||
pid = os.fork()
|
||||
if pid > 0:
|
||||
# Parent exits
|
||||
sys.exit(0)
|
||||
# First fork - detach from parent
|
||||
try:
|
||||
pid = os.fork()
|
||||
if pid > 0:
|
||||
sys.exit(0)
|
||||
except OSError as e:
|
||||
sys.stderr.write(f"fork #1 failed: {e}\n")
|
||||
sys.exit(1)
|
||||
|
||||
# Create new session
|
||||
# Decouple from parent environment
|
||||
os.chdir("/")
|
||||
os.setsid()
|
||||
os.umask(0)
|
||||
|
||||
# Second fork to prevent zombie processes
|
||||
# Second fork - prevent zombie
|
||||
pid = os.fork()
|
||||
if pid > 0:
|
||||
sys.exit(0)
|
||||
@@ -178,27 +111,3 @@ def check_daemon_running(port: int = DEFAULT_DAEMON_PORT) -> bool:
|
||||
return result == 0
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def stop_daemon(port: int = DEFAULT_DAEMON_PORT) -> bool:
|
||||
"""Stop a running daemon by sending SIGTERM to the process."""
|
||||
lock = DaemonLock()
|
||||
pid = lock.get_pid()
|
||||
|
||||
if pid is None:
|
||||
return False
|
||||
|
||||
try:
|
||||
import signal
|
||||
|
||||
os.kill(pid, signal.SIGTERM)
|
||||
# Wait for process to exit
|
||||
for _ in range(50): # Wait up to 5 seconds
|
||||
time.sleep(0.1)
|
||||
try:
|
||||
os.kill(pid, 0) # Check if process exists
|
||||
except OSError:
|
||||
return True # Process exited
|
||||
return False
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Consolidation engine for automatic learning creation from memories."""
|
||||
|
||||
from .consolidator import run_consolidation_job
|
||||
|
||||
__all__ = ["run_consolidation_job"]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,85 @@
|
||||
"""Prompts for the consolidation engine."""
|
||||
|
||||
CONSOLIDATION_SYSTEM_PROMPT = """You are a memory consolidation system. Your job is to convert facts into durable knowledge (observations) and merge with existing knowledge when appropriate.
|
||||
|
||||
You must output ONLY valid JSON with no markdown code blocks or additional text. However, the "text" field within each observation should use markdown formatting (headers, lists, bold, etc.) for clarity and readability.
|
||||
|
||||
## EXTRACT DURABLE KNOWLEDGE, NOT EPHEMERAL STATE
|
||||
Facts often describe events or actions. Extract the DURABLE KNOWLEDGE implied by the fact, not the transient state.
|
||||
|
||||
Examples of extracting durable knowledge:
|
||||
- "User moved to Room 203" -> "Room 203 exists" (location exists, not where user is now)
|
||||
- "User visited Acme Corp at Room 105" -> "Acme Corp is located in Room 105"
|
||||
- "User took the elevator to floor 3" -> "Floor 3 is accessible by elevator"
|
||||
- "User met Sarah at the lobby" -> "Sarah can be found at the lobby"
|
||||
|
||||
DO NOT track current user position/state as knowledge - that changes constantly.
|
||||
DO track permanent facts learned from the user's actions.
|
||||
|
||||
## PRESERVE SPECIFIC DETAILS
|
||||
Keep names, locations, numbers, and other specifics. Do NOT:
|
||||
- Abstract into general principles
|
||||
- Generate business insights
|
||||
- Make knowledge generic
|
||||
|
||||
GOOD examples:
|
||||
- Fact: "John likes pizza" -> "John likes pizza"
|
||||
- Fact: "Alice works at Google" -> "Alice works at Google"
|
||||
|
||||
BAD examples:
|
||||
- "John likes pizza" -> "Understanding dietary preferences helps..." (TOO ABSTRACT)
|
||||
- "User is at Room 203" -> "User is currently at Room 203" (EPHEMERAL STATE)
|
||||
|
||||
## MERGE RULES (when comparing to existing observations):
|
||||
1. REDUNDANT: Same information worded differently → update existing
|
||||
2. CONTRADICTION: Opposite information about same topic → update with temporal markers showing change
|
||||
Example: "Alex used to love pizza but now hates it" OR "Alex's pizza preference changed from love to hate"
|
||||
3. UPDATE: New state replacing old state → update showing the transition with "used to", "now", "changed from X to Y"
|
||||
|
||||
## CRITICAL RULES:
|
||||
- NEVER merge facts about DIFFERENT people
|
||||
- NEVER merge unrelated topics (food preferences vs work vs hobbies)
|
||||
- When merging contradictions, the "text" field MUST capture BOTH states with temporal markers:
|
||||
* Use "used to X, now Y" OR "changed from X to Y" OR "X but now Y"
|
||||
* DO NOT just state the new fact - you MUST show the change
|
||||
- Keep observations focused on ONE specific topic per person
|
||||
- The "text" field MUST contain durable knowledge, not ephemeral state
|
||||
- Do NOT include "tags" in output - tags are handled automatically"""
|
||||
|
||||
CONSOLIDATION_USER_PROMPT = """Analyze this new fact and consolidate into knowledge.
|
||||
{mission_section}
|
||||
NEW FACT: {fact_text}
|
||||
|
||||
EXISTING OBSERVATIONS (JSON array with source memories and dates):
|
||||
{observations_text}
|
||||
|
||||
Each observation includes:
|
||||
- id: unique identifier for updating
|
||||
- text: the observation content
|
||||
- proof_count: number of supporting memories
|
||||
- tags: visibility scope (handled automatically)
|
||||
- created_at/updated_at: when observation was created/modified
|
||||
- occurred_start/occurred_end: temporal range of source facts
|
||||
- source_memories: array of supporting facts with their text and dates
|
||||
|
||||
Instructions:
|
||||
1. Extract DURABLE KNOWLEDGE from the new fact (not ephemeral state)
|
||||
2. Review source_memories in existing observations to understand evidence
|
||||
3. Check dates to detect contradictions or updates
|
||||
4. Compare with observations:
|
||||
- Same topic → UPDATE with learning_id
|
||||
- New topic → CREATE new observation
|
||||
- Purely ephemeral → return []
|
||||
|
||||
Output JSON array of actions (the "text" field should use markdown formatting for structure):
|
||||
[
|
||||
{{"action": "update", "learning_id": "uuid-from-observations", "text": "## Updated Knowledge\n\n**Key point**: details here\n\n- Supporting detail 1\n- Supporting detail 2", "reason": "..."}},
|
||||
{{"action": "create", "text": "## New Durable Knowledge\n\nDescription with **emphasis** and proper structure", "reason": "..."}}
|
||||
]
|
||||
|
||||
Return [] if fact contains no durable knowledge.
|
||||
|
||||
IMPORTANT: Format the "text" field with markdown for better readability:
|
||||
- Use headers, lists, bold/italic, tables where appropriate
|
||||
- CRITICAL: Add blank lines before and after block elements (tables, code blocks, lists)
|
||||
- Ensure proper spacing for markdown to render correctly"""
|
||||
@@ -9,6 +9,7 @@ Configuration via environment variables - see hindsight_api.config for all env v
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import warnings
|
||||
from abc import ABC, abstractmethod
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
@@ -20,6 +21,7 @@ from ..config import (
|
||||
DEFAULT_RERANKER_FLASHRANK_CACHE_DIR,
|
||||
DEFAULT_RERANKER_FLASHRANK_MODEL,
|
||||
DEFAULT_RERANKER_LITELLM_MODEL,
|
||||
DEFAULT_RERANKER_LOCAL_FORCE_CPU,
|
||||
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||
DEFAULT_RERANKER_LOCAL_MODEL,
|
||||
DEFAULT_RERANKER_PROVIDER,
|
||||
@@ -33,6 +35,7 @@ from ..config import (
|
||||
ENV_RERANKER_FLASHRANK_CACHE_DIR,
|
||||
ENV_RERANKER_FLASHRANK_MODEL,
|
||||
ENV_RERANKER_LITELLM_MODEL,
|
||||
ENV_RERANKER_LOCAL_FORCE_CPU,
|
||||
ENV_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||
ENV_RERANKER_LOCAL_MODEL,
|
||||
ENV_RERANKER_PROVIDER,
|
||||
@@ -99,7 +102,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
_executor: ThreadPoolExecutor | None = None
|
||||
_max_concurrent: int = 4 # Limit concurrent CPU-bound reranking calls
|
||||
|
||||
def __init__(self, model_name: str | None = None, max_concurrent: int = 4):
|
||||
def __init__(self, model_name: str | None = None, max_concurrent: int = 4, force_cpu: bool = False):
|
||||
"""
|
||||
Initialize local SentenceTransformers cross-encoder.
|
||||
|
||||
@@ -108,8 +111,11 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
Default: cross-encoder/ms-marco-MiniLM-L-6-v2
|
||||
max_concurrent: Maximum concurrent reranking calls (default: 2).
|
||||
Higher values may cause CPU thrashing under load.
|
||||
force_cpu: Force CPU mode (avoids MPS/XPC issues on macOS in daemon mode).
|
||||
Default: False
|
||||
"""
|
||||
self.model_name = model_name or DEFAULT_RERANKER_LOCAL_MODEL
|
||||
self.force_cpu = force_cpu
|
||||
self._model = None
|
||||
LocalSTCrossEncoder._max_concurrent = max_concurrent
|
||||
|
||||
@@ -130,13 +136,55 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
"Install it with: pip install sentence-transformers"
|
||||
)
|
||||
|
||||
# Note: We use CPU even when GPU/MPS is available because:
|
||||
# 1. The reranker model (MiniLM) is tiny (~22M params)
|
||||
# 2. Batch sizes are small (~100-200 pairs)
|
||||
# 3. Data transfer overhead to GPU outweighs compute benefit
|
||||
# 4. CPU inference is actually faster for this workload
|
||||
logger.info(f"Reranker: initializing local provider with model {self.model_name}")
|
||||
self._model = CrossEncoder(self.model_name)
|
||||
|
||||
# Determine device based on hardware availability.
|
||||
# We always set low_cpu_mem_usage=False to prevent lazy loading (meta tensors)
|
||||
# which can cause issues when accelerate is installed but no GPU is available.
|
||||
# Note: We do NOT use device_map because CrossEncoder internally calls .to(device)
|
||||
# after loading, which conflicts with accelerate's device_map handling.
|
||||
import torch
|
||||
|
||||
# Force CPU mode if configured (used in daemon mode to avoid MPS/XPC issues on macOS)
|
||||
if self.force_cpu:
|
||||
device = "cpu"
|
||||
logger.info("Reranker: forcing CPU mode (HINDSIGHT_API_RERANKER_LOCAL_FORCE_CPU=1)")
|
||||
else:
|
||||
# Check for GPU (CUDA) or Apple Silicon (MPS)
|
||||
# Wrap in try-except to gracefully handle any device detection issues
|
||||
# (e.g., in CI environments or when PyTorch is built without GPU support)
|
||||
device = "cpu" # Default to CPU
|
||||
try:
|
||||
has_gpu = torch.cuda.is_available() or (
|
||||
hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
|
||||
)
|
||||
if has_gpu:
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to detect GPU/MPS, falling back to CPU: {e}")
|
||||
|
||||
# Suppress verbose transformers warnings during model loading
|
||||
# This suppresses the "UNEXPECTED" warnings from CrossEncoder which are harmless
|
||||
# but look alarming to users (e.g., "embeddings.position_ids | UNEXPECTED")
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore", category=UserWarning)
|
||||
warnings.filterwarnings("ignore", message=".*was not found in model state dict.*")
|
||||
warnings.filterwarnings("ignore", message=".*UNEXPECTED.*")
|
||||
|
||||
# Also suppress transformers library logging temporarily
|
||||
transformers_logger = logging.getLogger("transformers")
|
||||
original_level = transformers_logger.level
|
||||
transformers_logger.setLevel(logging.ERROR)
|
||||
|
||||
try:
|
||||
self._model = CrossEncoder(
|
||||
self.model_name,
|
||||
device=device,
|
||||
model_kwargs={"low_cpu_mem_usage": False},
|
||||
)
|
||||
finally:
|
||||
# Restore original logging level
|
||||
transformers_logger.setLevel(original_level)
|
||||
|
||||
# Initialize shared executor (limited workers naturally limits concurrency)
|
||||
if LocalSTCrossEncoder._executor is None:
|
||||
@@ -148,6 +196,11 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
else:
|
||||
logger.info("Reranker: local provider initialized (using existing executor)")
|
||||
|
||||
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""Synchronous prediction wrapper for thread pool execution."""
|
||||
scores = self._model.predict(pairs, show_progress_bar=False)
|
||||
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs for relevance.
|
||||
@@ -165,11 +218,11 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
|
||||
# Use dedicated executor - limited workers naturally limits concurrency
|
||||
loop = asyncio.get_event_loop()
|
||||
scores = await loop.run_in_executor(
|
||||
return await loop.run_in_executor(
|
||||
LocalSTCrossEncoder._executor,
|
||||
lambda: self._model.predict(pairs, show_progress_bar=False),
|
||||
self._predict_sync,
|
||||
pairs,
|
||||
)
|
||||
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
|
||||
|
||||
|
||||
class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
@@ -579,7 +632,7 @@ class FlashRankCrossEncoder(CrossEncoderModel):
|
||||
return
|
||||
|
||||
try:
|
||||
from flashrank import Ranker # type: ignore[import-untyped]
|
||||
from flashrank import Ranker
|
||||
except ImportError:
|
||||
raise ImportError("flashrank is required for FlashRankCrossEncoder. Install it with: pip install flashrank")
|
||||
|
||||
@@ -606,7 +659,7 @@ class FlashRankCrossEncoder(CrossEncoderModel):
|
||||
|
||||
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""Synchronous predict - processes each query group."""
|
||||
from flashrank import RerankRequest # type: ignore[import-untyped]
|
||||
from flashrank import RerankRequest
|
||||
|
||||
if not pairs:
|
||||
return []
|
||||
@@ -768,29 +821,33 @@ class LiteLLMCrossEncoder(CrossEncoderModel):
|
||||
|
||||
def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
"""
|
||||
Create a CrossEncoderModel instance based on environment variables.
|
||||
Create a CrossEncoderModel instance based on configuration.
|
||||
|
||||
See hindsight_api.config for environment variable names and defaults.
|
||||
Reads configuration via get_config() to ensure consistency across the codebase.
|
||||
|
||||
Returns:
|
||||
Configured CrossEncoderModel instance
|
||||
"""
|
||||
provider = os.environ.get(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER).lower()
|
||||
from ..config import get_config
|
||||
|
||||
config = get_config()
|
||||
provider = config.reranker_provider.lower()
|
||||
|
||||
if provider == "tei":
|
||||
url = os.environ.get(ENV_RERANKER_TEI_URL)
|
||||
url = config.reranker_tei_url
|
||||
if not url:
|
||||
raise ValueError(f"{ENV_RERANKER_TEI_URL} is required when {ENV_RERANKER_PROVIDER} is 'tei'")
|
||||
batch_size = int(os.environ.get(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE)))
|
||||
max_concurrent = int(os.environ.get(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT)))
|
||||
return RemoteTEICrossEncoder(base_url=url, batch_size=batch_size, max_concurrent=max_concurrent)
|
||||
return RemoteTEICrossEncoder(
|
||||
base_url=url,
|
||||
batch_size=config.reranker_tei_batch_size,
|
||||
max_concurrent=config.reranker_tei_max_concurrent,
|
||||
)
|
||||
elif provider == "local":
|
||||
model = os.environ.get(ENV_RERANKER_LOCAL_MODEL)
|
||||
model_name = model or DEFAULT_RERANKER_LOCAL_MODEL
|
||||
max_concurrent = int(
|
||||
os.environ.get(ENV_RERANKER_LOCAL_MAX_CONCURRENT, str(DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT))
|
||||
return LocalSTCrossEncoder(
|
||||
model_name=config.reranker_local_model,
|
||||
max_concurrent=config.reranker_local_max_concurrent,
|
||||
force_cpu=config.reranker_local_force_cpu,
|
||||
)
|
||||
return LocalSTCrossEncoder(model_name=model_name, max_concurrent=max_concurrent)
|
||||
elif provider == "cohere":
|
||||
api_key = os.environ.get(ENV_COHERE_API_KEY)
|
||||
if not api_key:
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Directives module for hard rules injected into prompts."""
|
||||
|
||||
from .models import Directive
|
||||
|
||||
__all__ = ["Directive"]
|
||||
@@ -0,0 +1,37 @@
|
||||
"""Pydantic models for directives."""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class Directive(BaseModel):
|
||||
"""A directive is a hard rule injected into prompts.
|
||||
|
||||
Directives are user-defined rules that guide agent behavior. Unlike mental models
|
||||
which are automatically consolidated from memories, directives are explicit
|
||||
instructions that are always included in relevant prompts.
|
||||
|
||||
Examples:
|
||||
- "Always respond in formal English"
|
||||
- "Never share personal data with third parties"
|
||||
- "Prefer conservative investment recommendations"
|
||||
"""
|
||||
|
||||
id: UUID = Field(description="Unique identifier")
|
||||
bank_id: str = Field(description="Bank this directive belongs to")
|
||||
name: str = Field(description="Human-readable name")
|
||||
content: str = Field(description="The directive text to inject into prompts")
|
||||
priority: int = Field(default=0, description="Higher priority directives are injected first")
|
||||
is_active: bool = Field(default=True, description="Whether this directive is currently active")
|
||||
tags: list[str] = Field(default_factory=list, description="Tags for filtering")
|
||||
created_at: datetime = Field(
|
||||
default_factory=lambda: datetime.now(timezone.utc), description="When this directive was created"
|
||||
)
|
||||
updated_at: datetime = Field(
|
||||
default_factory=lambda: datetime.now(timezone.utc), description="When this directive was last updated"
|
||||
)
|
||||
|
||||
class Config:
|
||||
from_attributes = True
|
||||
@@ -11,6 +11,7 @@ Configuration via environment variables - see hindsight_api.config for all env v
|
||||
|
||||
import logging
|
||||
import os
|
||||
import warnings
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import httpx
|
||||
@@ -18,6 +19,7 @@ import httpx
|
||||
from ..config import (
|
||||
DEFAULT_EMBEDDINGS_COHERE_MODEL,
|
||||
DEFAULT_EMBEDDINGS_LITELLM_MODEL,
|
||||
DEFAULT_EMBEDDINGS_LOCAL_FORCE_CPU,
|
||||
DEFAULT_EMBEDDINGS_LOCAL_MODEL,
|
||||
DEFAULT_EMBEDDINGS_OPENAI_MODEL,
|
||||
DEFAULT_EMBEDDINGS_PROVIDER,
|
||||
@@ -26,6 +28,7 @@ from ..config import (
|
||||
ENV_EMBEDDINGS_COHERE_BASE_URL,
|
||||
ENV_EMBEDDINGS_COHERE_MODEL,
|
||||
ENV_EMBEDDINGS_LITELLM_MODEL,
|
||||
ENV_EMBEDDINGS_LOCAL_FORCE_CPU,
|
||||
ENV_EMBEDDINGS_LOCAL_MODEL,
|
||||
ENV_EMBEDDINGS_OPENAI_API_KEY,
|
||||
ENV_EMBEDDINGS_OPENAI_BASE_URL,
|
||||
@@ -92,15 +95,18 @@ class LocalSTEmbeddings(Embeddings):
|
||||
The embedding dimension is auto-detected from the model.
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str | None = None):
|
||||
def __init__(self, model_name: str | None = None, force_cpu: bool = False):
|
||||
"""
|
||||
Initialize local SentenceTransformers embeddings.
|
||||
|
||||
Args:
|
||||
model_name: Name of the SentenceTransformer model to use.
|
||||
Default: BAAI/bge-small-en-v1.5
|
||||
force_cpu: Force CPU mode (avoids MPS/XPC issues on macOS in daemon mode).
|
||||
Default: False
|
||||
"""
|
||||
self.model_name = model_name or DEFAULT_EMBEDDINGS_LOCAL_MODEL
|
||||
self.force_cpu = force_cpu
|
||||
self._model = None
|
||||
self._dimension: int | None = None
|
||||
|
||||
@@ -128,12 +134,52 @@ class LocalSTEmbeddings(Embeddings):
|
||||
)
|
||||
|
||||
logger.info(f"Embeddings: initializing local provider with model {self.model_name}")
|
||||
# Disable lazy loading (meta tensors) which causes issues with newer transformers/accelerate
|
||||
# Setting low_cpu_mem_usage=False and device_map=None ensures tensors are fully materialized
|
||||
self._model = SentenceTransformer(
|
||||
self.model_name,
|
||||
model_kwargs={"low_cpu_mem_usage": False, "device_map": None},
|
||||
)
|
||||
|
||||
# Determine device based on hardware availability.
|
||||
# We always set low_cpu_mem_usage=False to prevent lazy loading (meta tensors)
|
||||
# which can cause issues when accelerate is installed but no GPU is available.
|
||||
import torch
|
||||
|
||||
# Force CPU mode if configured (used in daemon mode to avoid MPS/XPC issues on macOS)
|
||||
if self.force_cpu:
|
||||
device = "cpu"
|
||||
logger.info("Embeddings: forcing CPU mode")
|
||||
else:
|
||||
# Check for GPU (CUDA) or Apple Silicon (MPS)
|
||||
# Wrap in try-except to gracefully handle any device detection issues
|
||||
# (e.g., in CI environments or when PyTorch is built without GPU support)
|
||||
device = "cpu" # Default to CPU
|
||||
try:
|
||||
has_gpu = torch.cuda.is_available() or (
|
||||
hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
|
||||
)
|
||||
if has_gpu:
|
||||
device = None # Let sentence-transformers auto-detect GPU/MPS
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to detect GPU/MPS, falling back to CPU: {e}")
|
||||
|
||||
# Suppress verbose transformers warnings during model loading
|
||||
# This suppresses the "UNEXPECTED" warnings from BertModel which are harmless
|
||||
# but look alarming to users (e.g., "embeddings.position_ids | UNEXPECTED")
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore", category=UserWarning)
|
||||
warnings.filterwarnings("ignore", message=".*was not found in model state dict.*")
|
||||
warnings.filterwarnings("ignore", message=".*UNEXPECTED.*")
|
||||
|
||||
# Also suppress transformers library logging temporarily
|
||||
transformers_logger = logging.getLogger("transformers")
|
||||
original_level = transformers_logger.level
|
||||
transformers_logger.setLevel(logging.ERROR)
|
||||
|
||||
try:
|
||||
self._model = SentenceTransformer(
|
||||
self.model_name,
|
||||
device=device,
|
||||
model_kwargs={"low_cpu_mem_usage": False},
|
||||
)
|
||||
finally:
|
||||
# Restore original logging level
|
||||
transformers_logger.setLevel(original_level)
|
||||
|
||||
self._dimension = self._model.get_sentence_embedding_dimension()
|
||||
logger.info(f"Embeddings: local provider initialized (dim: {self._dimension})")
|
||||
@@ -150,6 +196,7 @@ class LocalSTEmbeddings(Embeddings):
|
||||
"""
|
||||
if self._model is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
|
||||
embeddings = self._model.encode(texts, convert_to_numpy=True, show_progress_bar=False)
|
||||
return [emb.tolist() for emb in embeddings]
|
||||
|
||||
@@ -516,7 +563,7 @@ class CohereEmbeddings(Embeddings):
|
||||
model=self.model,
|
||||
input_type=self.input_type,
|
||||
)
|
||||
if response.embeddings:
|
||||
if response.embeddings and isinstance(response.embeddings, list):
|
||||
self._dimension = len(response.embeddings[0])
|
||||
|
||||
logger.info(f"Embeddings: Cohere provider initialized (model: {self.model}, dim: {self._dimension})")
|
||||
@@ -673,24 +720,28 @@ class LiteLLMEmbeddings(Embeddings):
|
||||
|
||||
def create_embeddings_from_env() -> Embeddings:
|
||||
"""
|
||||
Create an Embeddings instance based on environment variables.
|
||||
Create an Embeddings instance based on configuration.
|
||||
|
||||
See hindsight_api.config for environment variable names and defaults.
|
||||
Reads configuration via get_config() to ensure consistency across the codebase.
|
||||
|
||||
Returns:
|
||||
Configured Embeddings instance
|
||||
"""
|
||||
provider = os.environ.get(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER).lower()
|
||||
from ..config import get_config
|
||||
|
||||
config = get_config()
|
||||
provider = config.embeddings_provider.lower()
|
||||
|
||||
if provider == "tei":
|
||||
url = os.environ.get(ENV_EMBEDDINGS_TEI_URL)
|
||||
url = config.embeddings_tei_url
|
||||
if not url:
|
||||
raise ValueError(f"{ENV_EMBEDDINGS_TEI_URL} is required when {ENV_EMBEDDINGS_PROVIDER} is 'tei'")
|
||||
return RemoteTEIEmbeddings(base_url=url)
|
||||
elif provider == "local":
|
||||
model = os.environ.get(ENV_EMBEDDINGS_LOCAL_MODEL)
|
||||
model_name = model or DEFAULT_EMBEDDINGS_LOCAL_MODEL
|
||||
return LocalSTEmbeddings(model_name=model_name)
|
||||
return LocalSTEmbeddings(
|
||||
model_name=config.embeddings_local_model,
|
||||
force_cpu=config.embeddings_local_force_cpu,
|
||||
)
|
||||
elif provider == "openai":
|
||||
# Use dedicated embeddings API key, or fall back to LLM API key
|
||||
api_key = os.environ.get(ENV_EMBEDDINGS_OPENAI_API_KEY) or os.environ.get(ENV_LLM_API_KEY)
|
||||
|
||||
@@ -160,14 +160,14 @@ class MemoryEngineInterface(ABC):
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Get bank profile including disposition and background.
|
||||
Get bank profile including disposition and mission.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Bank profile dict.
|
||||
Bank profile dict with bank_id, name, disposition, and mission.
|
||||
"""
|
||||
...
|
||||
|
||||
@@ -190,25 +190,44 @@ class MemoryEngineInterface(ABC):
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def merge_bank_background(
|
||||
async def merge_bank_mission(
|
||||
self,
|
||||
bank_id: str,
|
||||
new_info: str,
|
||||
*,
|
||||
update_disposition: bool = True,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Merge new background information into bank profile.
|
||||
Merge new mission information into bank profile.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
new_info: New background information to merge.
|
||||
update_disposition: Whether to infer disposition from background.
|
||||
new_info: New mission information to merge.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Updated background info.
|
||||
Updated mission info.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def set_bank_mission(
|
||||
self,
|
||||
bank_id: str,
|
||||
mission: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Set the bank's mission (replaces existing).
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
mission: The mission text.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with bank_id and mission.
|
||||
"""
|
||||
...
|
||||
|
||||
@@ -423,49 +442,6 @@ class MemoryEngineInterface(ABC):
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def get_entity_observations(
|
||||
self,
|
||||
bank_id: str,
|
||||
entity_id: str,
|
||||
*,
|
||||
limit: int = 10,
|
||||
request_context: "RequestContext",
|
||||
) -> list[Any]:
|
||||
"""
|
||||
Get observations for an entity.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
entity_id: The entity ID.
|
||||
limit: Maximum observations.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
List of EntityObservation objects.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def regenerate_entity_observations(
|
||||
self,
|
||||
bank_id: str,
|
||||
entity_id: str,
|
||||
entity_name: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> None:
|
||||
"""
|
||||
Regenerate observations for an entity.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
entity_id: The entity ID.
|
||||
entity_name: The entity's canonical name.
|
||||
request_context: Request context for authentication.
|
||||
"""
|
||||
...
|
||||
|
||||
# =========================================================================
|
||||
# Statistics & Operations
|
||||
# =========================================================================
|
||||
@@ -518,7 +494,7 @@ class MemoryEngineInterface(ABC):
|
||||
bank_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
List async operations for a bank.
|
||||
|
||||
@@ -527,7 +503,7 @@ class MemoryEngineInterface(ABC):
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
List of operation dicts with id, task_type, status, etc.
|
||||
Dict with 'total' (int) and 'operations' (list of operation dicts).
|
||||
"""
|
||||
...
|
||||
|
||||
@@ -561,16 +537,16 @@ class MemoryEngineInterface(ABC):
|
||||
bank_id: str,
|
||||
*,
|
||||
name: str | None = None,
|
||||
background: str | None = None,
|
||||
mission: str | None = None,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Update bank name and/or background.
|
||||
Update bank name and/or mission.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
name: New bank name (optional).
|
||||
background: New background text (optional, replaces existing).
|
||||
mission: New mission text (optional, replaces existing).
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
"""
|
||||
Abstract interface for LLM providers.
|
||||
|
||||
This module defines the interface that all LLM providers must implement,
|
||||
enabling support for multiple LLM backends (OpenAI, Anthropic, Gemini, Codex, etc.)
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
|
||||
from .response_models import LLMToolCallResult, TokenUsage
|
||||
|
||||
|
||||
class LLMInterface(ABC):
|
||||
"""
|
||||
Abstract interface for LLM providers.
|
||||
|
||||
All LLM provider implementations must inherit from this class and implement
|
||||
the required methods.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Initialize LLM provider.
|
||||
|
||||
Args:
|
||||
provider: Provider name (e.g., "openai", "codex", "anthropic", "gemini").
|
||||
api_key: API key or authentication token.
|
||||
base_url: Base URL for the API.
|
||||
model: Model name.
|
||||
reasoning_effort: Reasoning effort level for supported providers.
|
||||
**kwargs: Additional provider-specific parameters.
|
||||
"""
|
||||
self.provider = provider.lower()
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.reasoning_effort = reasoning_effort
|
||||
|
||||
@abstractmethod
|
||||
async def verify_connection(self) -> None:
|
||||
"""
|
||||
Verify that the LLM provider is configured correctly by making a simple test call.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the connection test fails.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def call(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any | None = None,
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 60.0,
|
||||
skip_validation: bool = False,
|
||||
strict_schema: bool = False,
|
||||
return_usage: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
Make an LLM API call with retry logic.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
response_format: Optional Pydantic model for structured output.
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
skip_validation: Return raw JSON without Pydantic validation.
|
||||
strict_schema: Use strict JSON schema enforcement (OpenAI only).
|
||||
return_usage: If True, return tuple (result, TokenUsage) instead of just result.
|
||||
|
||||
Returns:
|
||||
If return_usage=False: Parsed response if response_format is provided, otherwise text content.
|
||||
If return_usage=True: Tuple of (result, TokenUsage) with token counts.
|
||||
|
||||
Raises:
|
||||
OutputTooLongError: If output exceeds token limits.
|
||||
Exception: Re-raises API errors after retries exhausted.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def call_with_tools(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "tools",
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make an LLM API call with tool/function calling support.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts. Can include tool results with role='tool'.
|
||||
tools: List of tool definitions in OpenAI format.
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or specific function.
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def cleanup(self) -> None:
|
||||
"""Clean up resources (close connections, etc.)."""
|
||||
pass
|
||||
|
||||
|
||||
class OutputTooLongError(Exception):
|
||||
"""
|
||||
Bridge exception raised when LLM output exceeds token limits.
|
||||
|
||||
This wraps provider-specific errors (e.g., OpenAI's LengthFinishReasonError)
|
||||
to allow callers to handle output length issues without depending on
|
||||
provider-specific implementations.
|
||||
"""
|
||||
|
||||
pass
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,14 @@
|
||||
"""
|
||||
Mental models module for Hindsight.
|
||||
|
||||
Mental models contain directives - hard rules that are injected into reflect prompts.
|
||||
Directives are user-defined and their observations are user-provided (not LLM-generated).
|
||||
|
||||
Other types of consolidated knowledge are handled by:
|
||||
- Learnings: Automatic bottom-up consolidation from facts
|
||||
- Pinned Reflections: User-curated living documents
|
||||
"""
|
||||
|
||||
from .models import MentalModel, MentalModelSubtype
|
||||
|
||||
__all__ = ["MentalModel", "MentalModelSubtype"]
|
||||
@@ -0,0 +1,53 @@
|
||||
"""
|
||||
Pydantic models for mental models.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class MentalModelSubtype(str, Enum):
|
||||
"""Subtype of mental model.
|
||||
|
||||
Currently only DIRECTIVE is supported. Other types of consolidated knowledge
|
||||
are handled by:
|
||||
- Learnings: Automatic bottom-up consolidation from facts
|
||||
- Pinned Reflections: User-curated living documents
|
||||
"""
|
||||
|
||||
DIRECTIVE = "directive" # User-defined hard rules, observations user-provided
|
||||
|
||||
|
||||
class MentalModel(BaseModel):
|
||||
"""
|
||||
A mental model representing synthesized understanding.
|
||||
|
||||
Mental models are the agent's consolidated knowledge. Unlike raw facts,
|
||||
mental models provide:
|
||||
- A one-liner description for quick scanning/retrieval
|
||||
- A full summary for deep understanding
|
||||
- Links to related mental models
|
||||
"""
|
||||
|
||||
id: str = Field(description="Unique identifier within the bank")
|
||||
bank_id: str = Field(description="Bank this mental model belongs to")
|
||||
subtype: MentalModelSubtype = Field(description="How this model was created")
|
||||
name: str = Field(description="Human-readable name")
|
||||
description: str = Field(description="One-liner for quick scanning and retrieval matching")
|
||||
summary: str | None = Field(default=None, description="Full synthesized understanding")
|
||||
|
||||
# References
|
||||
entity_id: str | None = Field(default=None, description="Reference to entities table when type=entity")
|
||||
source_facts: list[str] = Field(default_factory=list, description="Fact IDs used to generate summary")
|
||||
links: list[str] = Field(default_factory=list, description="Related mental model IDs")
|
||||
|
||||
# Tags for scoped visibility (similar to document tags)
|
||||
tags: list[str] = Field(default_factory=list, description="Tags for scoped visibility filtering")
|
||||
|
||||
# Timestamps
|
||||
last_updated: datetime | None = Field(default=None, description="When summary was last regenerated")
|
||||
created_at: datetime = Field(
|
||||
default_factory=lambda: datetime.now(timezone.utc), description="When this model was created"
|
||||
)
|
||||
@@ -0,0 +1,14 @@
|
||||
"""
|
||||
LLM provider implementations.
|
||||
|
||||
This package contains concrete implementations of the LLMInterface for various providers.
|
||||
"""
|
||||
|
||||
from .anthropic_llm import AnthropicLLM
|
||||
from .claude_code_llm import ClaudeCodeLLM
|
||||
from .codex_llm import CodexLLM
|
||||
from .gemini_llm import GeminiLLM
|
||||
from .mock_llm import MockLLM
|
||||
from .openai_compatible_llm import OpenAICompatibleLLM
|
||||
|
||||
__all__ = ["AnthropicLLM", "ClaudeCodeLLM", "CodexLLM", "GeminiLLM", "MockLLM", "OpenAICompatibleLLM"]
|
||||
@@ -0,0 +1,434 @@
|
||||
"""
|
||||
Anthropic LLM provider using the Anthropic Python SDK.
|
||||
|
||||
This provider enables using Claude models from Anthropic with support for:
|
||||
- Structured JSON output
|
||||
- Tool/function calling with proper format conversion
|
||||
- Extended thinking mode
|
||||
- Retry logic with exponential backoff
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
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
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AnthropicLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using Anthropic's Claude models.
|
||||
|
||||
Supports structured output, tool calling, and extended thinking mode.
|
||||
Handles format conversion between OpenAI-style messages and Anthropic's format.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
timeout: float = 300.0,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Initialize Anthropic LLM provider.
|
||||
|
||||
Args:
|
||||
provider: Provider name (should be "anthropic").
|
||||
api_key: Anthropic API key.
|
||||
base_url: Base URL for the API (optional, uses Anthropic default if empty).
|
||||
model: Model name (e.g., "claude-sonnet-4-20250514").
|
||||
reasoning_effort: Reasoning effort level (not used by Anthropic).
|
||||
timeout: Request timeout in seconds.
|
||||
**kwargs: Additional provider-specific parameters.
|
||||
"""
|
||||
super().__init__(provider, api_key, base_url, model, reasoning_effort, **kwargs)
|
||||
|
||||
if not self.api_key:
|
||||
raise ValueError("API key is required for Anthropic provider")
|
||||
|
||||
# Import and initialize Anthropic client
|
||||
try:
|
||||
from anthropic import AsyncAnthropic
|
||||
|
||||
client_kwargs: dict[str, Any] = {"api_key": self.api_key}
|
||||
if self.base_url:
|
||||
client_kwargs["base_url"] = self.base_url
|
||||
if timeout:
|
||||
client_kwargs["timeout"] = timeout
|
||||
|
||||
self._client = AsyncAnthropic(**client_kwargs)
|
||||
logger.info(f"Anthropic client initialized for model: {self.model}")
|
||||
except ImportError as e:
|
||||
raise RuntimeError("Anthropic SDK not installed. Run: uv add anthropic or pip install anthropic") from e
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""
|
||||
Verify that the Anthropic provider is configured correctly by making a simple test call.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the connection test fails.
|
||||
"""
|
||||
try:
|
||||
test_messages = [{"role": "user", "content": "test"}]
|
||||
await self.call(
|
||||
messages=test_messages,
|
||||
max_completion_tokens=10,
|
||||
temperature=0.0,
|
||||
scope="test",
|
||||
max_retries=0,
|
||||
)
|
||||
logger.info("Anthropic connection verified successfully")
|
||||
except Exception as e:
|
||||
logger.error(f"Anthropic connection verification failed: {e}")
|
||||
raise RuntimeError(f"Failed to verify Anthropic connection: {e}") from e
|
||||
|
||||
async def call(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any | None = None,
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 60.0,
|
||||
skip_validation: bool = False,
|
||||
strict_schema: bool = False,
|
||||
return_usage: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
Make an LLM API call with retry logic.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
response_format: Optional Pydantic model for structured output.
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
skip_validation: Return raw JSON without Pydantic validation.
|
||||
strict_schema: Use strict JSON schema enforcement (not supported by Anthropic).
|
||||
return_usage: If True, return tuple (result, TokenUsage) instead of just result.
|
||||
|
||||
Returns:
|
||||
If return_usage=False: Parsed response if response_format is provided, otherwise text content.
|
||||
If return_usage=True: Tuple of (result, TokenUsage) with token counts.
|
||||
|
||||
Raises:
|
||||
OutputTooLongError: If output exceeds token limits.
|
||||
Exception: Re-raises API errors after retries exhausted.
|
||||
"""
|
||||
from anthropic import APIConnectionError, APIStatusError, RateLimitError
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# Convert OpenAI-style messages to Anthropic format
|
||||
system_prompt = None
|
||||
anthropic_messages = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == "system":
|
||||
if system_prompt:
|
||||
system_prompt += "\n\n" + content
|
||||
else:
|
||||
system_prompt = content
|
||||
else:
|
||||
anthropic_messages.append({"role": role, "content": content})
|
||||
|
||||
# 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()
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2)}"
|
||||
if system_prompt:
|
||||
system_prompt += schema_msg
|
||||
else:
|
||||
system_prompt = schema_msg
|
||||
|
||||
# Prepare parameters
|
||||
call_params: dict[str, Any] = {
|
||||
"model": self.model,
|
||||
"messages": anthropic_messages,
|
||||
"max_tokens": max_completion_tokens if max_completion_tokens is not None else 4096,
|
||||
}
|
||||
|
||||
if system_prompt:
|
||||
call_params["system"] = system_prompt
|
||||
|
||||
if temperature is not None:
|
||||
call_params["temperature"] = temperature
|
||||
|
||||
last_exception = None
|
||||
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.messages.create(**call_params)
|
||||
|
||||
# Anthropic response content is a list of blocks
|
||||
content = ""
|
||||
for block in response.content:
|
||||
if block.type == "text":
|
||||
content += block.text
|
||||
|
||||
if response_format is not None:
|
||||
# Models may wrap JSON in markdown code blocks
|
||||
clean_content = content
|
||||
if "```json" in content:
|
||||
clean_content = content.split("```json")[1].split("```")[0].strip()
|
||||
elif "```" in content:
|
||||
clean_content = content.split("```")[1].split("```")[0].strip()
|
||||
|
||||
try:
|
||||
json_data = json.loads(clean_content)
|
||||
except json.JSONDecodeError:
|
||||
# Fallback to parsing raw content if markdown stripping failed
|
||||
json_data = json.loads(content)
|
||||
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
else:
|
||||
result = response_format.model_validate(json_data)
|
||||
else:
|
||||
result = content
|
||||
|
||||
# Record metrics and log slow calls
|
||||
duration = time.time() - start_time
|
||||
input_tokens = response.usage.input_tokens or 0 if response.usage else 0
|
||||
output_tokens = response.usage.output_tokens or 0 if response.usage else 0
|
||||
total_tokens = input_tokens + output_tokens
|
||||
|
||||
# Record LLM metrics
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
)
|
||||
|
||||
# Log slow calls
|
||||
if duration > 10.0:
|
||||
logger.info(
|
||||
f"slow llm call: scope={scope}, model={self.provider}/{self.model}, "
|
||||
f"input_tokens={input_tokens}, output_tokens={output_tokens}, "
|
||||
f"time={duration:.3f}s"
|
||||
)
|
||||
|
||||
if return_usage:
|
||||
token_usage = TokenUsage(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
)
|
||||
return result, token_usage
|
||||
return result
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
logger.warning("Anthropic returned invalid JSON, retrying...")
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
logger.error(f"Anthropic returned invalid JSON after {max_retries + 1} attempts")
|
||||
raise
|
||||
|
||||
except (APIConnectionError, RateLimitError, APIStatusError) as e:
|
||||
# Fast fail on 401/403
|
||||
if isinstance(e, APIStatusError) and e.status_code in (401, 403):
|
||||
logger.error(f"Anthropic auth error (HTTP {e.status_code}), not retrying: {str(e)}")
|
||||
raise
|
||||
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
# Check if it's a rate limit or server error
|
||||
should_retry = isinstance(e, (APIConnectionError, RateLimitError)) or (
|
||||
isinstance(e, APIStatusError) and e.status_code >= 500
|
||||
)
|
||||
|
||||
if should_retry:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
jitter = backoff * 0.2 * (2 * (time.time() % 1) - 1)
|
||||
await asyncio.sleep(backoff + jitter)
|
||||
continue
|
||||
|
||||
logger.error(f"Anthropic API error after {max_retries + 1} attempts: {str(e)}")
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected error during Anthropic call: {type(e).__name__}: {str(e)}")
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("Anthropic call failed after all retries")
|
||||
|
||||
async def call_with_tools(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "tools",
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make an LLM API call with tool/function calling support.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts. Can include tool results with role='tool'.
|
||||
tools: List of tool definitions in OpenAI format.
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or specific function.
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
"""
|
||||
from anthropic import APIConnectionError, APIStatusError
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# Convert OpenAI tool format to Anthropic format
|
||||
anthropic_tools = []
|
||||
for tool in tools:
|
||||
func = tool.get("function", {})
|
||||
anthropic_tools.append(
|
||||
{
|
||||
"name": func.get("name", ""),
|
||||
"description": func.get("description", ""),
|
||||
"input_schema": func.get("parameters", {"type": "object", "properties": {}}),
|
||||
}
|
||||
)
|
||||
|
||||
# Convert messages - handle tool results
|
||||
system_prompt = None
|
||||
anthropic_messages = []
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == "system":
|
||||
system_prompt = (system_prompt + "\n\n" + content) if system_prompt else content
|
||||
elif role == "tool":
|
||||
# Anthropic uses tool_result blocks
|
||||
anthropic_messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "tool_result", "tool_use_id": msg.get("tool_call_id", ""), "content": content}
|
||||
],
|
||||
}
|
||||
)
|
||||
elif role == "assistant" and msg.get("tool_calls"):
|
||||
# Convert assistant tool calls
|
||||
tool_use_blocks = []
|
||||
for tc in msg["tool_calls"]:
|
||||
tool_use_blocks.append(
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": tc.get("id", ""),
|
||||
"name": tc.get("function", {}).get("name", ""),
|
||||
"input": json.loads(tc.get("function", {}).get("arguments", "{}")),
|
||||
}
|
||||
)
|
||||
anthropic_messages.append({"role": "assistant", "content": tool_use_blocks})
|
||||
else:
|
||||
anthropic_messages.append({"role": role, "content": content})
|
||||
|
||||
call_params: dict[str, Any] = {
|
||||
"model": self.model,
|
||||
"messages": anthropic_messages,
|
||||
"tools": anthropic_tools,
|
||||
"max_tokens": max_completion_tokens or 4096,
|
||||
}
|
||||
if system_prompt:
|
||||
call_params["system"] = system_prompt
|
||||
|
||||
if temperature is not None:
|
||||
call_params["temperature"] = temperature
|
||||
|
||||
last_exception = None
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.messages.create(**call_params)
|
||||
|
||||
# Extract content and tool calls
|
||||
content_parts = []
|
||||
tool_calls: list[LLMToolCall] = []
|
||||
|
||||
for block in response.content:
|
||||
if block.type == "text":
|
||||
content_parts.append(block.text)
|
||||
elif block.type == "tool_use":
|
||||
tool_calls.append(LLMToolCall(id=block.id, name=block.name, arguments=block.input or {}))
|
||||
|
||||
content = "".join(content_parts) if content_parts else None
|
||||
finish_reason = "tool_calls" if tool_calls else "stop"
|
||||
|
||||
# Extract token usage
|
||||
input_tokens = response.usage.input_tokens or 0
|
||||
output_tokens = response.usage.output_tokens or 0
|
||||
|
||||
# Record metrics
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=time.time() - start_time,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
)
|
||||
|
||||
return LLMToolCallResult(
|
||||
content=content,
|
||||
tool_calls=tool_calls,
|
||||
finish_reason=finish_reason,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
)
|
||||
|
||||
except (APIConnectionError, APIStatusError) as e:
|
||||
if isinstance(e, APIStatusError) and e.status_code in (401, 403):
|
||||
raise
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
await asyncio.sleep(min(initial_backoff * (2**attempt), max_backoff))
|
||||
continue
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("Anthropic tool call failed")
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
"""Clean up resources (close Anthropic client connections)."""
|
||||
if hasattr(self, "_client") and self._client:
|
||||
await self._client.close()
|
||||
@@ -0,0 +1,493 @@
|
||||
"""
|
||||
Claude Code LLM provider using Claude Agent SDK.
|
||||
|
||||
This provider enables using Claude Pro/Max subscriptions for API calls
|
||||
via the Claude CLI authentication. It uses the Claude Agent SDK which
|
||||
automatically handles authentication via `claude auth login` credentials.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
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
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ClaudeCodeLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using Claude Code authentication.
|
||||
|
||||
Authenticates using Claude Pro/Max credentials via `claude auth login`
|
||||
and makes API calls through the Claude Agent SDK.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str, # Will be ignored, uses CLI auth
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""Initialize Claude Code LLM provider."""
|
||||
super().__init__(provider, api_key, base_url, model, reasoning_effort, **kwargs)
|
||||
|
||||
# Verify Claude Agent SDK is available
|
||||
try:
|
||||
self._verify_claude_code_available()
|
||||
logger.info("Claude Code: Using Claude Agent SDK (authentication via claude auth login)")
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Failed to initialize Claude Code provider: {e}\n\n"
|
||||
"To set up Claude Code authentication:\n"
|
||||
"1. Install Claude Code CLI: npm install -g @anthropics/claude-code\n"
|
||||
"2. Login with your Pro/Max plan: claude auth login\n"
|
||||
"3. Verify authentication: claude --version\n\n"
|
||||
"Or use a different provider (anthropic, openai, gemini) with API keys."
|
||||
) from e
|
||||
|
||||
# Metrics collector is imported at module level
|
||||
|
||||
def _verify_claude_code_available(self) -> None:
|
||||
"""
|
||||
Verify that Claude Agent SDK can be imported and is properly configured.
|
||||
|
||||
Raises:
|
||||
ImportError: If Claude Agent SDK is not installed.
|
||||
RuntimeError: If Claude Code is not authenticated.
|
||||
"""
|
||||
try:
|
||||
# Import Claude Agent SDK
|
||||
# Reduce Claude Agent SDK logging verbosity
|
||||
import logging as sdk_logging
|
||||
|
||||
from claude_agent_sdk import query # noqa: F401
|
||||
|
||||
sdk_logging.getLogger("claude_agent_sdk").setLevel(sdk_logging.WARNING)
|
||||
sdk_logging.getLogger("claude_agent_sdk._internal").setLevel(sdk_logging.WARNING)
|
||||
|
||||
logger.debug("Claude Agent SDK imported successfully")
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"Claude Agent SDK not installed. Run: uv add claude-agent-sdk or pip install claude-agent-sdk"
|
||||
) from e
|
||||
|
||||
# SDK will automatically check for authentication when first used
|
||||
# No need to verify here - let it fail gracefully on first call with helpful error
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""
|
||||
Verify that the Claude Code provider is configured correctly by making a simple test call.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the connection test fails.
|
||||
"""
|
||||
try:
|
||||
test_messages = [{"role": "user", "content": "test"}]
|
||||
await self.call(
|
||||
messages=test_messages,
|
||||
max_completion_tokens=10,
|
||||
temperature=0.0,
|
||||
scope="test",
|
||||
max_retries=0,
|
||||
)
|
||||
logger.info("Claude Code connection verified successfully")
|
||||
except Exception as e:
|
||||
logger.error(f"Claude Code connection verification failed: {e}")
|
||||
raise RuntimeError(f"Failed to verify Claude Code connection: {e}") from e
|
||||
|
||||
async def call(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any | None = None,
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 60.0,
|
||||
skip_validation: bool = False,
|
||||
strict_schema: bool = False,
|
||||
return_usage: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
Make an LLM API call with retry logic.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
response_format: Optional Pydantic model for structured output.
|
||||
max_completion_tokens: Maximum tokens in response (ignored by Claude Agent SDK).
|
||||
temperature: Sampling temperature (ignored by Claude Agent SDK).
|
||||
scope: Scope identifier for tracking.
|
||||
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: Use strict JSON schema enforcement (not supported).
|
||||
return_usage: If True, return tuple (result, TokenUsage) instead of just result.
|
||||
|
||||
Returns:
|
||||
If return_usage=False: Parsed response if response_format is provided, otherwise text content.
|
||||
If return_usage=True: Tuple of (result, TokenUsage) with estimated token counts.
|
||||
|
||||
Raises:
|
||||
OutputTooLongError: If output exceeds token limits (not supported by Claude Agent SDK).
|
||||
Exception: Re-raises API errors after retries exhausted.
|
||||
"""
|
||||
from claude_agent_sdk import AssistantMessage, ClaudeAgentOptions, TextBlock, query
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# Build system prompt
|
||||
system_prompt = ""
|
||||
user_content = ""
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == "system":
|
||||
system_prompt += ("\n\n" + content) if system_prompt else content
|
||||
elif role == "user":
|
||||
user_content += ("\n\n" + content) if user_content else content
|
||||
elif role == "assistant":
|
||||
# Claude Agent SDK doesn't support multi-turn easily in query()
|
||||
# For now, prepend assistant messages to user content
|
||||
user_content += f"\n\n[Previous assistant response: {content}]"
|
||||
|
||||
# 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()
|
||||
schema_instruction = (
|
||||
f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2)}\n\n"
|
||||
"Respond with ONLY the JSON, no markdown formatting."
|
||||
)
|
||||
user_content += schema_instruction
|
||||
|
||||
# Configure SDK options
|
||||
options = ClaudeAgentOptions(
|
||||
system_prompt=system_prompt if system_prompt else None,
|
||||
max_turns=1, # Single-turn for API-style interactions
|
||||
allowed_tools=[], # Disable tools for standard LLM calls
|
||||
)
|
||||
|
||||
# Call Claude Agent SDK
|
||||
last_exception = None
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
# Collect streaming response
|
||||
full_text = ""
|
||||
|
||||
async for message in query(prompt=user_content, options=options):
|
||||
if isinstance(message, AssistantMessage):
|
||||
for block in message.content:
|
||||
if isinstance(block, TextBlock):
|
||||
full_text += block.text
|
||||
|
||||
# Handle structured output
|
||||
if response_format is not None:
|
||||
# Models may wrap JSON in markdown
|
||||
clean_text = full_text
|
||||
if "```json" in full_text:
|
||||
clean_text = full_text.split("```json")[1].split("```")[0].strip()
|
||||
elif "```" in full_text:
|
||||
clean_text = full_text.split("```")[1].split("```")[0].strip()
|
||||
|
||||
try:
|
||||
json_data = json.loads(clean_text)
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f"Claude Code JSON parse error (attempt {attempt + 1}/{max_retries + 1}): {e}")
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
last_exception = e
|
||||
continue
|
||||
raise
|
||||
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
else:
|
||||
result = response_format.model_validate(json_data)
|
||||
else:
|
||||
result = full_text
|
||||
|
||||
# Record metrics
|
||||
duration = time.time() - start_time
|
||||
metrics = get_metrics_collector()
|
||||
|
||||
# Estimate token usage (Claude Agent SDK doesn't report exact counts)
|
||||
# Use character count / 4 as rough estimate (1 token ≈ 4 characters)
|
||||
estimated_input = sum(len(m.get("content", "")) for m in messages) // 4
|
||||
estimated_output = len(full_text) // 4
|
||||
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=estimated_input,
|
||||
output_tokens=estimated_output,
|
||||
success=True,
|
||||
)
|
||||
|
||||
# Log slow calls
|
||||
if duration > 10.0:
|
||||
logger.info(
|
||||
f"slow llm call: scope={scope}, model={self.provider}/{self.model}, time={duration:.3f}s"
|
||||
)
|
||||
|
||||
if return_usage:
|
||||
token_usage = TokenUsage(
|
||||
input_tokens=estimated_input,
|
||||
output_tokens=estimated_output,
|
||||
total_tokens=estimated_input + estimated_output,
|
||||
)
|
||||
return result, token_usage
|
||||
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
last_exception = e
|
||||
|
||||
# Check for authentication errors
|
||||
error_str = str(e).lower()
|
||||
if "auth" in error_str or "login" in error_str or "credential" in error_str:
|
||||
logger.error(f"Claude Code authentication error: {e}")
|
||||
raise RuntimeError(
|
||||
f"Claude Code authentication failed: {e}\n\n"
|
||||
"Run 'claude auth login' to authenticate with Claude Pro/Max."
|
||||
) from e
|
||||
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
logger.warning(f"Claude Code error (attempt {attempt + 1}/{max_retries + 1}): {e}")
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
logger.error(f"Claude Code error after {max_retries + 1} attempts: {e}")
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("Claude Code call failed after all retries")
|
||||
|
||||
async def call_with_tools(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "tools",
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make an LLM API call with tool/function calling support using Claude Agent SDK.
|
||||
|
||||
This implementation uses ClaudeSDKClient (not query()) because custom tools via
|
||||
SDK MCP servers are only supported with the client. Tools are converted from OpenAI
|
||||
format to SDK MCP tools, and tool names are formatted as mcp__hindsight_tools__{name}.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts. Can include tool results with role='tool'.
|
||||
tools: List of tool definitions in OpenAI format.
|
||||
max_completion_tokens: Maximum tokens in response (not used by Claude Agent SDK).
|
||||
temperature: Sampling temperature (not used by Claude Agent SDK).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools (not used by Claude Agent SDK).
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
"""
|
||||
from claude_agent_sdk import (
|
||||
AssistantMessage,
|
||||
ClaudeAgentOptions,
|
||||
ClaudeSDKClient,
|
||||
SdkMcpTool,
|
||||
TextBlock,
|
||||
ToolUseBlock,
|
||||
create_sdk_mcp_server,
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# Convert OpenAI tool format to Claude Agent SDK SdkMcpTool format
|
||||
sdk_tools: list[SdkMcpTool] = []
|
||||
tool_names: list[str] = []
|
||||
|
||||
for tool in tools:
|
||||
func = tool.get("function", {})
|
||||
tool_name = func.get("name", "")
|
||||
tool_description = func.get("description", "")
|
||||
parameters = func.get("parameters", {})
|
||||
|
||||
# Create a handler with proper closure to avoid transport issues
|
||||
def make_handler(name: str):
|
||||
async def handler(args: dict[str, Any]) -> dict[str, Any]:
|
||||
# Return immediately with success - tool execution happens externally
|
||||
return {
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": f"[Tool {name} called successfully]",
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
return handler
|
||||
|
||||
sdk_tools.append(
|
||||
SdkMcpTool(
|
||||
name=tool_name,
|
||||
description=tool_description,
|
||||
input_schema=parameters,
|
||||
handler=make_handler(tool_name),
|
||||
)
|
||||
)
|
||||
tool_names.append(tool_name)
|
||||
|
||||
# Create an MCP server with the tools
|
||||
mcp_server = create_sdk_mcp_server(
|
||||
name="hindsight_tools",
|
||||
version="1.0.0",
|
||||
tools=sdk_tools if sdk_tools else None,
|
||||
)
|
||||
|
||||
# Build system prompt and user content from messages
|
||||
system_prompt = ""
|
||||
user_content = ""
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == "system":
|
||||
system_prompt += ("\n\n" + content) if system_prompt else content
|
||||
elif role == "user":
|
||||
user_content += ("\n\n" + content) if user_content else content
|
||||
elif role == "assistant":
|
||||
# Include previous assistant messages as context
|
||||
user_content += f"\n\n[Previous assistant response: {content}]"
|
||||
elif role == "tool":
|
||||
# Tool results are already in tool_results_map, append to user context
|
||||
tool_call_id = msg.get("tool_call_id", "")
|
||||
user_content += f"\n\n[Tool result for {tool_call_id}: {content}]"
|
||||
|
||||
# Format tool names for SDK MCP servers: mcp__{server_name}__{tool_name}
|
||||
# This is required by the Claude Agent SDK for MCP server tools
|
||||
allowed_tool_names = [f"mcp__hindsight_tools__{name}" for name in tool_names]
|
||||
|
||||
# Configure SDK options with MCP server
|
||||
options = ClaudeAgentOptions(
|
||||
system_prompt=system_prompt if system_prompt else None,
|
||||
max_turns=1, # Single-turn for API-style interactions
|
||||
mcp_servers={"hindsight_tools": mcp_server} if sdk_tools else {},
|
||||
allowed_tools=allowed_tool_names if allowed_tool_names else [],
|
||||
)
|
||||
|
||||
# Call Claude Agent SDK with retry logic
|
||||
last_exception = None
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
full_text = ""
|
||||
tool_calls: list[LLMToolCall] = []
|
||||
|
||||
# Use ClaudeSDKClient for tool calling support
|
||||
# Note: query() does NOT support custom tools, only ClaudeSDKClient does
|
||||
async with ClaudeSDKClient(options=options) as client:
|
||||
# Send the query
|
||||
await client.query(user_content)
|
||||
|
||||
# Receive response
|
||||
async for message in client.receive_response():
|
||||
if isinstance(message, AssistantMessage):
|
||||
for block in message.content:
|
||||
if isinstance(block, TextBlock):
|
||||
full_text += block.text
|
||||
elif isinstance(block, ToolUseBlock):
|
||||
# SDK returns tool names with MCP prefix (mcp__hindsight_tools__{name})
|
||||
# Strip the prefix to return original tool name expected by caller
|
||||
tool_name = block.name
|
||||
if tool_name.startswith("mcp__hindsight_tools__"):
|
||||
tool_name = tool_name.replace("mcp__hindsight_tools__", "", 1)
|
||||
|
||||
tool_calls.append(
|
||||
LLMToolCall(
|
||||
id=block.id,
|
||||
name=tool_name,
|
||||
arguments=block.input,
|
||||
)
|
||||
)
|
||||
|
||||
# Record metrics
|
||||
duration = time.time() - start_time
|
||||
metrics = get_metrics_collector()
|
||||
|
||||
# Estimate token usage (Claude Agent SDK doesn't report exact counts)
|
||||
estimated_input = sum(len(m.get("content", "")) for m in messages) // 4
|
||||
estimated_output = len(full_text) // 4
|
||||
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=estimated_input,
|
||||
output_tokens=estimated_output,
|
||||
success=True,
|
||||
)
|
||||
|
||||
# Log slow calls
|
||||
if duration > 10.0:
|
||||
logger.info(
|
||||
f"slow llm call: scope={scope}, model={self.provider}/{self.model}, time={duration:.3f}s"
|
||||
)
|
||||
|
||||
return LLMToolCallResult(
|
||||
content=full_text if full_text else None,
|
||||
tool_calls=tool_calls,
|
||||
finish_reason="tool_calls" if tool_calls else "stop",
|
||||
input_tokens=estimated_input,
|
||||
output_tokens=estimated_output,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
last_exception = e
|
||||
|
||||
# Check for authentication errors
|
||||
error_str = str(e).lower()
|
||||
if "auth" in error_str or "login" in error_str or "credential" in error_str:
|
||||
logger.error(f"Claude Code authentication error: {e}")
|
||||
raise RuntimeError(
|
||||
f"Claude Code authentication failed: {e}\n\n"
|
||||
"Run 'claude auth login' to authenticate with Claude Pro/Max."
|
||||
) from e
|
||||
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
logger.warning(f"Claude Code tool call error (attempt {attempt + 1}/{max_retries + 1}): {e}")
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
logger.error(f"Claude Code tool call error after {max_retries + 1} attempts: {e}")
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("Claude Code tool call failed after all retries")
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
"""Clean up resources (no HTTP client to close for Claude Agent SDK)."""
|
||||
pass
|
||||
@@ -0,0 +1,578 @@
|
||||
"""
|
||||
OpenAI Codex LLM provider using ChatGPT Plus/Pro OAuth authentication.
|
||||
|
||||
This provider enables using ChatGPT Plus/Pro subscriptions for API calls
|
||||
without separate OpenAI Platform API credits. It uses OAuth tokens from
|
||||
~/.codex/auth.json and communicates with the ChatGPT backend API.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
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
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CodexLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider using OpenAI Codex OAuth authentication.
|
||||
|
||||
Authenticates using ChatGPT Plus/Pro credentials stored in ~/.codex/auth.json
|
||||
and makes API calls to chatgpt.com/backend-api/codex/responses.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str, # Will be ignored, reads from ~/.codex/auth.json
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""Initialize Codex LLM provider."""
|
||||
super().__init__(provider, api_key, base_url, model, reasoning_effort, **kwargs)
|
||||
|
||||
# Load Codex OAuth credentials
|
||||
try:
|
||||
self.access_token, self.account_id = self._load_codex_auth()
|
||||
logger.info(f"Loaded Codex OAuth credentials for account: {self.account_id}")
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
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"
|
||||
"3. Verify: ls ~/.codex/auth.json\n\n"
|
||||
"Or use a different provider (openai, anthropic, gemini) with API keys."
|
||||
) from e
|
||||
|
||||
# Use ChatGPT backend API endpoint
|
||||
if not self.base_url:
|
||||
self.base_url = "https://chatgpt.com/backend-api"
|
||||
|
||||
# Normalize model name (strip openai/ prefix if present)
|
||||
if self.model.startswith("openai/"):
|
||||
self.model = self.model[len("openai/") :]
|
||||
|
||||
# Map reasoning effort to Codex reasoning summary format
|
||||
# Codex supports: "auto", "concise", "detailed"
|
||||
self.reasoning_summary = self._map_reasoning_effort(reasoning_effort)
|
||||
|
||||
# HTTP client for SSE streaming
|
||||
self._client = httpx.AsyncClient(timeout=120.0)
|
||||
|
||||
def _load_codex_auth(self) -> tuple[str, str]:
|
||||
"""
|
||||
Load OAuth credentials from ~/.codex/auth.json.
|
||||
|
||||
Returns:
|
||||
Tuple of (access_token, account_id).
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If auth file doesn't exist.
|
||||
ValueError: If auth file is invalid.
|
||||
"""
|
||||
auth_file = Path.home() / ".codex" / "auth.json"
|
||||
|
||||
if not auth_file.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Codex auth file not found: {auth_file}\nRun 'codex auth login' to authenticate with ChatGPT Plus/Pro."
|
||||
)
|
||||
|
||||
with open(auth_file) as f:
|
||||
data = json.load(f)
|
||||
|
||||
# Validate auth structure
|
||||
auth_mode = data.get("auth_mode")
|
||||
if auth_mode != "chatgpt":
|
||||
raise ValueError(f"Expected auth_mode='chatgpt', got: {auth_mode}")
|
||||
|
||||
tokens = data.get("tokens", {})
|
||||
access_token = tokens.get("access_token")
|
||||
account_id = tokens.get("account_id")
|
||||
|
||||
if not access_token:
|
||||
raise ValueError("No access_token found in Codex auth file. Run 'codex auth login' again.")
|
||||
|
||||
return access_token, account_id
|
||||
|
||||
def _map_reasoning_effort(self, effort: str) -> str:
|
||||
"""
|
||||
Map standard reasoning effort to Codex reasoning summary format.
|
||||
|
||||
Args:
|
||||
effort: Standard effort level ("low", "medium", "high", "xhigh").
|
||||
|
||||
Returns:
|
||||
Codex reasoning summary: "concise", "detailed", or "auto".
|
||||
"""
|
||||
mapping = {
|
||||
"low": "concise",
|
||||
"medium": "auto",
|
||||
"high": "detailed",
|
||||
"xhigh": "detailed",
|
||||
}
|
||||
return mapping.get(effort.lower(), "auto")
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""Verify Codex connection by making a simple test call."""
|
||||
try:
|
||||
logger.info(f"Verifying Codex LLM: model={self.model}, account={self.account_id}...")
|
||||
await self.call(
|
||||
messages=[{"role": "user", "content": "Say 'ok'"}],
|
||||
max_completion_tokens=10,
|
||||
max_retries=2,
|
||||
initial_backoff=0.5,
|
||||
max_backoff=2.0,
|
||||
)
|
||||
logger.info(f"Codex LLM verified: {self.model}")
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Codex LLM connection verification failed for {self.model}: {e}") from e
|
||||
|
||||
async def call(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any | None = None,
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 60.0,
|
||||
skip_validation: bool = False,
|
||||
strict_schema: bool = False,
|
||||
return_usage: bool = False,
|
||||
) -> Any:
|
||||
"""Make API call to Codex backend with SSE streaming."""
|
||||
start_time = time.time()
|
||||
|
||||
# Prepare system instructions
|
||||
system_instruction = ""
|
||||
user_messages = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == "system":
|
||||
system_instruction += ("\n\n" + content) if system_instruction else content
|
||||
else:
|
||||
user_messages.append(msg)
|
||||
|
||||
# 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()
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2)}"
|
||||
system_instruction += schema_msg
|
||||
|
||||
# gpt-5.2-codex only supports "detailed" reasoning summary
|
||||
reasoning_summary = "detailed" if "5.2" in self.model else self.reasoning_summary
|
||||
|
||||
# Build Codex request payload
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"instructions": system_instruction,
|
||||
"input": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": msg.get("role", "user"),
|
||||
"content": msg.get("content", ""),
|
||||
}
|
||||
for msg in user_messages
|
||||
],
|
||||
"tools": [],
|
||||
"tool_choice": "auto",
|
||||
"parallel_tool_calls": True,
|
||||
"reasoning": {"summary": reasoning_summary},
|
||||
"store": False, # Codex uses stateless mode
|
||||
"stream": True, # SSE streaming
|
||||
"include": ["reasoning.encrypted_content"],
|
||||
"prompt_cache_key": str(uuid.uuid4()),
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.access_token}",
|
||||
"Content-Type": "application/json",
|
||||
"OpenAI-Account-ID": self.account_id,
|
||||
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7)",
|
||||
"Origin": "https://chatgpt.com",
|
||||
}
|
||||
|
||||
url = f"{self.base_url}/codex/responses"
|
||||
last_exception = None
|
||||
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.post(url, json=payload, headers=headers, timeout=120.0)
|
||||
response.raise_for_status()
|
||||
|
||||
# Parse SSE stream
|
||||
content = await self._parse_sse_stream(response)
|
||||
|
||||
# Handle structured output
|
||||
if response_format is not None:
|
||||
# Models may wrap JSON in markdown
|
||||
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 as e:
|
||||
logger.warning(f"Codex JSON parse error (attempt {attempt + 1}/{max_retries + 1}): {e}")
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
last_exception = e
|
||||
continue
|
||||
raise
|
||||
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
else:
|
||||
result = response_format.model_validate(json_data)
|
||||
else:
|
||||
result = content
|
||||
|
||||
# Record metrics
|
||||
duration = time.time() - start_time
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=0, # Codex doesn't report token counts in SSE
|
||||
output_tokens=0,
|
||||
success=True,
|
||||
)
|
||||
|
||||
if return_usage:
|
||||
# Codex doesn't provide token counts, estimate based on content
|
||||
estimated_input = sum(len(m.get("content", "")) for m in messages) // 4
|
||||
estimated_output = len(content) // 4
|
||||
token_usage = TokenUsage(
|
||||
input_tokens=estimated_input,
|
||||
output_tokens=estimated_output,
|
||||
total_tokens=estimated_input + estimated_output,
|
||||
)
|
||||
return result, token_usage
|
||||
|
||||
return result
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
last_exception = e
|
||||
status_code = e.response.status_code
|
||||
|
||||
# Fast fail on auth errors
|
||||
if status_code in (401, 403):
|
||||
logger.error(f"Codex auth error (HTTP {status_code}): {e.response.text[:200]}")
|
||||
raise RuntimeError(
|
||||
"Codex authentication failed. Your OAuth token may have expired.\n"
|
||||
"Run 'codex auth login' to re-authenticate."
|
||||
) from e
|
||||
|
||||
# Log the actual error message from the API
|
||||
error_detail = e.response.text[:500] if hasattr(e.response, "text") else str(e)
|
||||
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
logger.warning(
|
||||
f"Codex HTTP error {status_code} (attempt {attempt + 1}/{max_retries + 1}): {error_detail}"
|
||||
)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
logger.error(
|
||||
f"Codex HTTP error after {max_retries + 1} attempts: Status {status_code}, Detail: {error_detail}"
|
||||
)
|
||||
raise
|
||||
|
||||
except httpx.RequestError as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
logger.warning(f"Codex connection error (attempt {attempt + 1}/{max_retries + 1}): {e}")
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
logger.error(f"Codex connection error after {max_retries + 1} attempts: {e}")
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected Codex error: {type(e).__name__}: {e}")
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("Codex call failed after all retries")
|
||||
|
||||
async def _parse_sse_stream(self, response: httpx.Response) -> str:
|
||||
"""
|
||||
Parse Server-Sent Events (SSE) stream from Codex API.
|
||||
|
||||
Args:
|
||||
response: HTTP response with SSE stream.
|
||||
|
||||
Returns:
|
||||
Extracted text content from stream.
|
||||
"""
|
||||
full_text = ""
|
||||
event_type = None
|
||||
|
||||
async for line in response.aiter_lines():
|
||||
if not line:
|
||||
continue
|
||||
|
||||
# Track event type
|
||||
if line.startswith("event: "):
|
||||
event_type = line[7:]
|
||||
|
||||
# Parse data
|
||||
elif line.startswith("data: "):
|
||||
data_str = line[6:]
|
||||
if data_str == "[DONE]":
|
||||
break
|
||||
|
||||
try:
|
||||
data = json.loads(data_str)
|
||||
|
||||
# Extract content based on event type
|
||||
if event_type == "response.text.delta" and "delta" in data:
|
||||
full_text += data["delta"]
|
||||
elif event_type == "response.content_part.delta" and "delta" in data:
|
||||
full_text += data["delta"]
|
||||
# Check for item content
|
||||
elif "item" in data:
|
||||
item = data["item"]
|
||||
if "content" in item:
|
||||
content = item["content"]
|
||||
if isinstance(content, list):
|
||||
for part in content:
|
||||
if isinstance(part, dict) and "text" in part:
|
||||
full_text += part["text"]
|
||||
elif isinstance(content, str):
|
||||
full_text += content
|
||||
|
||||
except json.JSONDecodeError:
|
||||
# Skip malformed JSON events
|
||||
pass
|
||||
|
||||
return full_text
|
||||
|
||||
async def call_with_tools(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "tools",
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make API call with tool calling support.
|
||||
|
||||
Parses Codex SSE stream to extract tool calls from response.output_item.done events.
|
||||
Tools are converted from OpenAI format to Codex format (flat structure at top level).
|
||||
|
||||
Args:
|
||||
messages: List of message dicts. Can include tool results with role='tool'.
|
||||
tools: List of tool definitions in OpenAI format.
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature.
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or specific function.
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
# Prepare system instructions
|
||||
system_instruction = ""
|
||||
user_messages = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == "system":
|
||||
system_instruction += ("\n\n" + content) if system_instruction else content
|
||||
elif role == "tool":
|
||||
# Handle tool results
|
||||
user_messages.append(
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": f"Tool result: {content}",
|
||||
}
|
||||
)
|
||||
else:
|
||||
user_messages.append(
|
||||
{
|
||||
"type": "message",
|
||||
"role": role,
|
||||
"content": content,
|
||||
}
|
||||
)
|
||||
|
||||
# Convert tools to Codex format
|
||||
# Codex expects tools with type and name/description/parameters at top level
|
||||
codex_tools = []
|
||||
for tool in tools:
|
||||
func = tool.get("function", {})
|
||||
codex_tools.append(
|
||||
{
|
||||
"type": "function",
|
||||
"name": func.get("name", ""),
|
||||
"description": func.get("description", ""),
|
||||
"parameters": func.get("parameters", {}),
|
||||
}
|
||||
)
|
||||
|
||||
# gpt-5.2-codex only supports "detailed" reasoning summary
|
||||
reasoning_summary = "detailed" if "5.2" in self.model else self.reasoning_summary
|
||||
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"instructions": system_instruction,
|
||||
"input": user_messages,
|
||||
"tools": codex_tools,
|
||||
"tool_choice": tool_choice,
|
||||
"parallel_tool_calls": True,
|
||||
"reasoning": {"summary": reasoning_summary},
|
||||
"store": False,
|
||||
"stream": True,
|
||||
"include": ["reasoning.encrypted_content"],
|
||||
"prompt_cache_key": str(uuid.uuid4()),
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.access_token}",
|
||||
"Content-Type": "application/json",
|
||||
"OpenAI-Account-ID": self.account_id,
|
||||
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7)",
|
||||
"Origin": "https://chatgpt.com",
|
||||
}
|
||||
|
||||
url = f"{self.base_url}/codex/responses"
|
||||
|
||||
# Debug logging for troubleshooting
|
||||
logger.debug(f"Codex tool call request: url={url}, model={payload['model']}, tools={len(codex_tools)}")
|
||||
|
||||
try:
|
||||
response = await self._client.post(url, json=payload, headers=headers, timeout=120.0)
|
||||
|
||||
# Log response details on error
|
||||
if response.status_code != 200:
|
||||
logger.error(f"Codex API error {response.status_code}: {response.text[:500]}")
|
||||
|
||||
response.raise_for_status()
|
||||
|
||||
# Parse SSE for tool calls and content
|
||||
content, tool_calls = await self._parse_sse_tool_stream(response)
|
||||
|
||||
duration = time.time() - start_time
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
success=True,
|
||||
)
|
||||
|
||||
return LLMToolCallResult(
|
||||
content=content,
|
||||
tool_calls=tool_calls,
|
||||
finish_reason="tool_calls" if tool_calls else "stop",
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Codex tool call error: {e}")
|
||||
raise
|
||||
|
||||
async def _parse_sse_tool_stream(self, response: httpx.Response) -> tuple[str | None, list[LLMToolCall]]:
|
||||
"""
|
||||
Parse SSE stream for tool calls and content.
|
||||
|
||||
Returns:
|
||||
Tuple of (content, tool_calls).
|
||||
"""
|
||||
content = ""
|
||||
tool_calls: list[LLMToolCall] = []
|
||||
event_type = None
|
||||
|
||||
async for line in response.aiter_lines():
|
||||
if not line:
|
||||
continue
|
||||
|
||||
if line.startswith("event: "):
|
||||
event_type = line[7:]
|
||||
|
||||
elif line.startswith("data: "):
|
||||
data_str = line[6:]
|
||||
if data_str == "[DONE]":
|
||||
break
|
||||
|
||||
try:
|
||||
data = json.loads(data_str)
|
||||
|
||||
# Extract text content
|
||||
if event_type == "response.text.delta" and "delta" in data:
|
||||
content += data["delta"]
|
||||
|
||||
# Extract completed tool calls from response.output_item.done
|
||||
elif event_type == "response.output_item.done":
|
||||
item = data.get("item", {})
|
||||
if item.get("type") == "function_call" and item.get("status") == "completed":
|
||||
tool_name = item.get("name", "")
|
||||
arguments_str = item.get("arguments", "{}")
|
||||
call_id = item.get("call_id", "")
|
||||
|
||||
try:
|
||||
arguments = json.loads(arguments_str)
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(f"Failed to parse tool arguments: {arguments_str}")
|
||||
arguments = {}
|
||||
|
||||
tool_calls.append(
|
||||
LLMToolCall(
|
||||
id=call_id,
|
||||
name=tool_name,
|
||||
arguments=arguments,
|
||||
)
|
||||
)
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f"Failed to parse SSE data: {e}, data_str: {data_str[:200]}")
|
||||
|
||||
return content if content else None, tool_calls
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
"""Clean up HTTP client."""
|
||||
await self._client.aclose()
|
||||
@@ -0,0 +1,502 @@
|
||||
"""
|
||||
Google Gemini/VertexAI LLM provider.
|
||||
|
||||
This provider supports both:
|
||||
1. Gemini API (api.generativeai.google.com) with API key authentication
|
||||
2. Vertex AI with service account or Application Default Credentials (ADC)
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from google import genai
|
||||
from google.genai import errors as genai_errors
|
||||
from google.genai import types as genai_types
|
||||
|
||||
from hindsight_api.engine.llm_interface import LLMInterface, OutputTooLongError
|
||||
from hindsight_api.engine.response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
from hindsight_api.metrics import get_metrics_collector
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Vertex AI imports (optional)
|
||||
try:
|
||||
import google.auth
|
||||
from google.oauth2 import service_account
|
||||
|
||||
VERTEXAI_AVAILABLE = True
|
||||
except ImportError:
|
||||
VERTEXAI_AVAILABLE = False
|
||||
|
||||
|
||||
class GeminiLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider for Google Gemini and Vertex AI.
|
||||
|
||||
Supports:
|
||||
- Gemini API: provider="gemini", requires api_key
|
||||
- Vertex AI: provider="vertexai", requires project_id and region, uses ADC or service account
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""Initialize Gemini/VertexAI LLM provider."""
|
||||
super().__init__(provider, api_key, base_url, model, reasoning_effort, **kwargs)
|
||||
|
||||
self._client = None
|
||||
self._is_vertexai = self.provider == "vertexai"
|
||||
|
||||
if self._is_vertexai:
|
||||
self._init_vertexai(**kwargs)
|
||||
else:
|
||||
self._init_gemini()
|
||||
|
||||
def _init_gemini(self) -> None:
|
||||
"""Initialize Gemini API client."""
|
||||
if not self.api_key:
|
||||
raise ValueError("Gemini provider requires api_key")
|
||||
|
||||
self._client = genai.Client(api_key=self.api_key)
|
||||
logger.info(f"Gemini API: model={self.model}")
|
||||
|
||||
def _init_vertexai(self, **kwargs: Any) -> None:
|
||||
"""Initialize Vertex AI client with project, region, and credentials."""
|
||||
# Extract Vertex AI config from kwargs
|
||||
project_id = kwargs.get("vertexai_project_id")
|
||||
region = kwargs.get("vertexai_region", "us-central1")
|
||||
service_account_key = kwargs.get("vertexai_service_account_key")
|
||||
credentials = kwargs.get("vertexai_credentials") # Pre-loaded credentials object
|
||||
|
||||
if not project_id:
|
||||
raise ValueError(
|
||||
"HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID is required for Vertex AI provider. "
|
||||
"Set it to your GCP project ID."
|
||||
)
|
||||
|
||||
auth_method = "ADC"
|
||||
|
||||
# Use pre-loaded credentials if provided (passed from LLMProvider)
|
||||
if credentials is not None:
|
||||
auth_method = "service_account"
|
||||
# Otherwise, load explicit service account credentials if path provided
|
||||
elif service_account_key:
|
||||
if not VERTEXAI_AVAILABLE:
|
||||
raise ValueError(
|
||||
"Vertex AI service account auth requires 'google-auth' package. "
|
||||
"Install with: pip install google-auth"
|
||||
)
|
||||
credentials = service_account.Credentials.from_service_account_file(
|
||||
service_account_key,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"],
|
||||
)
|
||||
auth_method = "service_account"
|
||||
logger.info(f"Vertex AI: Using service account key: {service_account_key}")
|
||||
|
||||
# Strip google/ prefix from model name — native SDK uses bare names
|
||||
# e.g. "google/gemini-2.0-flash-lite-001" -> "gemini-2.0-flash-lite-001"
|
||||
if self.model.startswith("google/"):
|
||||
self.model = self.model[len("google/") :]
|
||||
|
||||
# Create Vertex AI client
|
||||
client_kwargs: dict[str, Any] = {
|
||||
"vertexai": True,
|
||||
"project": project_id,
|
||||
"location": region,
|
||||
}
|
||||
if credentials is not None:
|
||||
client_kwargs["credentials"] = credentials
|
||||
|
||||
self._client = genai.Client(**client_kwargs)
|
||||
|
||||
logger.info(f"Vertex AI: project={project_id}, region={region}, model={self.model}, auth={auth_method}")
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""
|
||||
Verify that the Gemini/VertexAI provider is configured correctly.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the connection test fails.
|
||||
"""
|
||||
try:
|
||||
logger.info(f"Verifying {self.provider.upper()}: model={self.model}...")
|
||||
await self.call(
|
||||
messages=[{"role": "user", "content": "Say 'ok'"}],
|
||||
max_completion_tokens=100,
|
||||
max_retries=2,
|
||||
initial_backoff=0.5,
|
||||
max_backoff=2.0,
|
||||
)
|
||||
logger.info(f"{self.provider.upper()} connection verified successfully")
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to verify {self.provider.upper()} connection: {e}") from e
|
||||
|
||||
async def call(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any | None = None,
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 60.0,
|
||||
skip_validation: bool = False,
|
||||
strict_schema: bool = False,
|
||||
return_usage: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
Make a Gemini/VertexAI API call with retry logic.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
response_format: Optional Pydantic model for structured output.
|
||||
max_completion_tokens: Maximum tokens in response (not supported by Gemini).
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
skip_validation: Return raw JSON without Pydantic validation.
|
||||
strict_schema: Use strict JSON schema enforcement (not supported by Gemini).
|
||||
return_usage: If True, return tuple (result, TokenUsage).
|
||||
|
||||
Returns:
|
||||
If return_usage=False: Parsed response if response_format provided, else text.
|
||||
If return_usage=True: Tuple of (result, TokenUsage).
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
# Convert OpenAI-style messages to Gemini format
|
||||
system_instruction = None
|
||||
gemini_contents = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == "system":
|
||||
if system_instruction:
|
||||
system_instruction += "\n\n" + content
|
||||
else:
|
||||
system_instruction = content
|
||||
elif role == "assistant":
|
||||
gemini_contents.append(genai_types.Content(role="model", parts=[genai_types.Part(text=content)]))
|
||||
else:
|
||||
gemini_contents.append(genai_types.Content(role="user", parts=[genai_types.Part(text=content)]))
|
||||
|
||||
# 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()
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2)}"
|
||||
if system_instruction:
|
||||
system_instruction += schema_msg
|
||||
else:
|
||||
system_instruction = schema_msg
|
||||
|
||||
# Build generation config
|
||||
config_kwargs: dict[str, Any] = {}
|
||||
if system_instruction:
|
||||
config_kwargs["system_instruction"] = system_instruction
|
||||
if response_format is not None:
|
||||
config_kwargs["response_mime_type"] = "application/json"
|
||||
config_kwargs["response_schema"] = response_format
|
||||
if temperature is not None:
|
||||
config_kwargs["temperature"] = temperature
|
||||
|
||||
generation_config = genai_types.GenerateContentConfig(**config_kwargs) if config_kwargs else None
|
||||
|
||||
last_exception = None
|
||||
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.aio.models.generate_content(
|
||||
model=self.model,
|
||||
contents=gemini_contents,
|
||||
config=generation_config,
|
||||
)
|
||||
|
||||
content = response.text
|
||||
|
||||
# Handle empty response
|
||||
if content is None:
|
||||
block_reason = None
|
||||
if hasattr(response, "candidates") and response.candidates:
|
||||
candidate = response.candidates[0]
|
||||
if hasattr(candidate, "finish_reason"):
|
||||
block_reason = candidate.finish_reason
|
||||
|
||||
if attempt < max_retries:
|
||||
logger.warning(f"Gemini returned empty response (reason: {block_reason}), retrying...")
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
raise RuntimeError(f"Gemini returned empty response after {max_retries + 1} attempts")
|
||||
|
||||
# Parse structured output if requested
|
||||
if response_format is not None:
|
||||
json_data = json.loads(content)
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
else:
|
||||
result = response_format.model_validate(json_data)
|
||||
else:
|
||||
result = content
|
||||
|
||||
# Extract token usage
|
||||
input_tokens = 0
|
||||
output_tokens = 0
|
||||
if hasattr(response, "usage_metadata") and response.usage_metadata:
|
||||
usage = response.usage_metadata
|
||||
input_tokens = usage.prompt_token_count or 0
|
||||
output_tokens = usage.candidates_token_count or 0
|
||||
|
||||
# Record metrics
|
||||
duration = time.time() - start_time
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
)
|
||||
|
||||
# Log slow calls
|
||||
if duration > 10.0 and input_tokens > 0:
|
||||
logger.info(
|
||||
f"slow llm call: scope={scope}, model={self.provider}/{self.model}, "
|
||||
f"input_tokens={input_tokens}, output_tokens={output_tokens}, "
|
||||
f"time={duration:.3f}s"
|
||||
)
|
||||
|
||||
if return_usage:
|
||||
token_usage = TokenUsage(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=input_tokens + output_tokens,
|
||||
)
|
||||
return result, token_usage
|
||||
return result
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
logger.warning("Gemini returned invalid JSON, retrying...")
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
logger.error(f"Gemini returned invalid JSON after {max_retries + 1} attempts")
|
||||
raise
|
||||
|
||||
except genai_errors.APIError as e:
|
||||
# Fast fail on auth errors - these won't recover with retries
|
||||
if e.code in (401, 403):
|
||||
logger.error(f"Gemini auth error (HTTP {e.code}), not retrying: {str(e)}")
|
||||
raise
|
||||
|
||||
# Retry on retryable errors (rate limits, server errors, client errors)
|
||||
if e.code in (400, 429, 500, 502, 503, 504) or (e.code and e.code >= 500):
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
jitter = backoff * 0.2 * (2 * (time.time() % 1) - 1)
|
||||
await asyncio.sleep(backoff + jitter)
|
||||
else:
|
||||
logger.error(f"Gemini API error after {max_retries + 1} attempts: {str(e)}")
|
||||
raise
|
||||
else:
|
||||
logger.error(f"Gemini API error: {type(e).__name__}: {str(e)}")
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected error during Gemini call: {type(e).__name__}: {str(e)}")
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("Gemini call failed after all retries")
|
||||
|
||||
async def call_with_tools(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "tools",
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make a Gemini/VertexAI API call with tool/function calling support.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts. Can include tool results with role='tool'.
|
||||
tools: List of tool definitions in OpenAI format.
|
||||
max_completion_tokens: Maximum tokens (not supported by Gemini).
|
||||
temperature: Sampling temperature.
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools (Gemini uses "auto" only).
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
# Convert tools to Gemini format
|
||||
gemini_tools = []
|
||||
for tool in tools:
|
||||
func = tool.get("function", {})
|
||||
gemini_tools.append(
|
||||
genai_types.Tool(
|
||||
function_declarations=[
|
||||
genai_types.FunctionDeclaration(
|
||||
name=func.get("name", ""),
|
||||
description=func.get("description", ""),
|
||||
parameters=func.get("parameters"),
|
||||
)
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
# Convert messages
|
||||
system_instruction = None
|
||||
gemini_contents = []
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == "system":
|
||||
system_instruction = (system_instruction + "\n\n" + content) if system_instruction else content
|
||||
elif role == "tool":
|
||||
# Gemini uses function_response
|
||||
gemini_contents.append(
|
||||
genai_types.Content(
|
||||
role="user",
|
||||
parts=[
|
||||
genai_types.Part(
|
||||
function_response=genai_types.FunctionResponse(
|
||||
name=msg.get("name", ""),
|
||||
response={"result": content},
|
||||
)
|
||||
)
|
||||
],
|
||||
)
|
||||
)
|
||||
elif role == "assistant":
|
||||
gemini_contents.append(genai_types.Content(role="model", parts=[genai_types.Part(text=content)]))
|
||||
else:
|
||||
gemini_contents.append(genai_types.Content(role="user", parts=[genai_types.Part(text=content)]))
|
||||
|
||||
config_kwargs: dict[str, Any] = {"tools": gemini_tools}
|
||||
if system_instruction:
|
||||
config_kwargs["system_instruction"] = system_instruction
|
||||
if temperature is not None:
|
||||
config_kwargs["temperature"] = temperature
|
||||
|
||||
config = genai_types.GenerateContentConfig(**config_kwargs)
|
||||
|
||||
last_exception = None
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.aio.models.generate_content(
|
||||
model=self.model,
|
||||
contents=gemini_contents,
|
||||
config=config,
|
||||
)
|
||||
|
||||
# Extract content and tool calls
|
||||
content = None
|
||||
tool_calls: list[LLMToolCall] = []
|
||||
|
||||
if response.candidates and response.candidates[0].content:
|
||||
parts = response.candidates[0].content.parts
|
||||
if parts:
|
||||
for part in parts:
|
||||
if hasattr(part, "text") and part.text:
|
||||
content = part.text
|
||||
if hasattr(part, "function_call") and part.function_call:
|
||||
fc = part.function_call
|
||||
tool_calls.append(
|
||||
LLMToolCall(
|
||||
id=f"gemini_{len(tool_calls)}",
|
||||
name=fc.name,
|
||||
arguments=dict(fc.args) if fc.args else {},
|
||||
)
|
||||
)
|
||||
|
||||
finish_reason = "tool_calls" if tool_calls else "stop"
|
||||
|
||||
# Extract token usage
|
||||
input_tokens = 0
|
||||
output_tokens = 0
|
||||
if response.usage_metadata:
|
||||
input_tokens = response.usage_metadata.prompt_token_count or 0
|
||||
output_tokens = response.usage_metadata.candidates_token_count or 0
|
||||
|
||||
# Record metrics
|
||||
duration = time.time() - start_time
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
)
|
||||
|
||||
return LLMToolCallResult(
|
||||
content=content,
|
||||
tool_calls=tool_calls,
|
||||
finish_reason=finish_reason,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
)
|
||||
|
||||
except genai_errors.APIError as e:
|
||||
# Fast fail on auth errors
|
||||
if e.code in (401, 403):
|
||||
logger.error(f"Gemini auth error (HTTP {e.code}), not retrying: {str(e)}")
|
||||
raise
|
||||
|
||||
# Retry on retryable errors
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected error during Gemini tool call: {type(e).__name__}: {str(e)}")
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("Gemini tool call failed")
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
"""Clean up resources (close connections, etc.)."""
|
||||
# Gemini client doesn't require explicit cleanup
|
||||
pass
|
||||
@@ -0,0 +1,254 @@
|
||||
"""
|
||||
Mock LLM provider for testing.
|
||||
|
||||
This provider allows tests to record LLM calls and return configurable mock responses
|
||||
without making actual API calls to external LLM services.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from ..llm_interface import LLMInterface
|
||||
from ..response_models import LLMToolCall, LLMToolCallResult, TokenUsage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MockLLM(LLMInterface):
|
||||
"""
|
||||
Mock LLM provider for testing.
|
||||
|
||||
This provider records all calls and returns configurable mock responses,
|
||||
enabling tests to verify LLM interactions without making real API calls.
|
||||
|
||||
Example:
|
||||
# Create mock provider
|
||||
mock_llm = MockLLM(provider="mock", api_key="", base_url="", model="mock-model")
|
||||
|
||||
# Set mock response
|
||||
mock_llm.set_mock_response({"answer": "test"})
|
||||
|
||||
# Make calls
|
||||
result = await mock_llm.call(
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
response_format=MyResponseModel
|
||||
)
|
||||
|
||||
# Verify calls
|
||||
calls = mock_llm.get_mock_calls()
|
||||
assert len(calls) == 1
|
||||
assert calls[0]["scope"] == "memory"
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Initialize mock LLM provider.
|
||||
|
||||
Args:
|
||||
provider: Provider name (should be "mock").
|
||||
api_key: Not used for mock provider.
|
||||
base_url: Not used for mock provider.
|
||||
model: Model name for tracking.
|
||||
reasoning_effort: Not used for mock provider.
|
||||
**kwargs: Additional parameters (not used).
|
||||
"""
|
||||
super().__init__(provider, api_key, base_url, model, reasoning_effort, **kwargs)
|
||||
|
||||
# Storage for test verification
|
||||
self._mock_calls: list[dict] = []
|
||||
self._mock_response: Any = None
|
||||
self._mock_exception: Exception | None = None
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""
|
||||
Verify mock provider (always succeeds).
|
||||
|
||||
Mock provider doesn't need connection verification since it doesn't
|
||||
make real API calls.
|
||||
"""
|
||||
logger.debug("Mock LLM: connection verification (always succeeds)")
|
||||
|
||||
async def call(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any | None = None,
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 60.0,
|
||||
skip_validation: bool = False,
|
||||
strict_schema: bool = False,
|
||||
return_usage: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
Make a mock LLM API call.
|
||||
|
||||
Records the call for test verification and returns the configured mock response.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
response_format: Optional Pydantic model for structured output.
|
||||
max_completion_tokens: Not used in mock.
|
||||
temperature: Not used in mock.
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Not used in mock.
|
||||
initial_backoff: Not used in mock.
|
||||
max_backoff: Not used in mock.
|
||||
skip_validation: Return raw JSON without Pydantic validation.
|
||||
strict_schema: Not used in mock.
|
||||
return_usage: If True, return tuple (result, TokenUsage) instead of just result.
|
||||
|
||||
Returns:
|
||||
If return_usage=False: Parsed response if response_format is provided, otherwise text content.
|
||||
If return_usage=True: Tuple of (result, TokenUsage) with mock token counts.
|
||||
"""
|
||||
# Record the call for test verification
|
||||
call_record = {
|
||||
"provider": self.provider,
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
"response_format": response_format.__name__
|
||||
if response_format and hasattr(response_format, "__name__")
|
||||
else str(response_format),
|
||||
"scope": scope,
|
||||
}
|
||||
self._mock_calls.append(call_record)
|
||||
logger.debug(f"Mock LLM call recorded: scope={scope}, model={self.model}")
|
||||
|
||||
# Raise mock exception if configured
|
||||
if self._mock_exception is not None:
|
||||
raise self._mock_exception
|
||||
|
||||
# Return mock response
|
||||
if self._mock_response is not None:
|
||||
result = self._mock_response
|
||||
elif response_format is not None:
|
||||
# Try to create a minimal valid instance of the response format
|
||||
try:
|
||||
# For Pydantic models, try to create with minimal valid data
|
||||
result = {"mock": True}
|
||||
except Exception:
|
||||
result = {"mock": True}
|
||||
else:
|
||||
result = "mock response"
|
||||
|
||||
if return_usage:
|
||||
token_usage = TokenUsage(input_tokens=10, output_tokens=5, total_tokens=15)
|
||||
return result, token_usage
|
||||
return result
|
||||
|
||||
async def call_with_tools(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "tools",
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make a mock LLM API call with tool/function calling support.
|
||||
|
||||
Records the call for test verification and returns the configured mock response.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts. Can include tool results with role='tool'.
|
||||
tools: List of tool definitions in OpenAI format.
|
||||
max_completion_tokens: Not used in mock.
|
||||
temperature: Not used in mock.
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Not used in mock.
|
||||
initial_backoff: Not used in mock.
|
||||
max_backoff: Not used in mock.
|
||||
tool_choice: Not used in mock.
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
"""
|
||||
# Record the call for test verification
|
||||
call_record = {
|
||||
"provider": self.provider,
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
"tools": [t.get("function", {}).get("name") for t in tools],
|
||||
"scope": scope,
|
||||
}
|
||||
self._mock_calls.append(call_record)
|
||||
|
||||
# Raise mock exception if configured
|
||||
if self._mock_exception is not None:
|
||||
raise self._mock_exception
|
||||
|
||||
if self._mock_response is not None:
|
||||
if isinstance(self._mock_response, LLMToolCallResult):
|
||||
return self._mock_response
|
||||
# Allow setting just tool calls as a list
|
||||
if isinstance(self._mock_response, list):
|
||||
return LLMToolCallResult(
|
||||
tool_calls=[
|
||||
LLMToolCall(id=f"mock_{i}", name=tc["name"], arguments=tc.get("arguments", {}))
|
||||
for i, tc in enumerate(self._mock_response)
|
||||
],
|
||||
finish_reason="tool_calls",
|
||||
)
|
||||
|
||||
return LLMToolCallResult(content="mock response", finish_reason="stop")
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
"""Clean up resources (no-op for mock provider)."""
|
||||
pass
|
||||
|
||||
def set_mock_response(self, response: Any) -> None:
|
||||
"""
|
||||
Set the response to return from mock calls.
|
||||
|
||||
Args:
|
||||
response: The response to return. Can be:
|
||||
- A dict/Pydantic model for regular calls
|
||||
- An LLMToolCallResult for tool calls
|
||||
- A list of tool call dicts for tool calls
|
||||
- Any other value to return as-is
|
||||
"""
|
||||
self._mock_response = response
|
||||
|
||||
def set_mock_exception(self, exception: Exception) -> None:
|
||||
"""
|
||||
Set an exception to raise from mock calls.
|
||||
|
||||
Args:
|
||||
exception: The exception to raise on the next call.
|
||||
After raising, the exception is cleared.
|
||||
"""
|
||||
self._mock_exception = exception
|
||||
|
||||
def get_mock_calls(self) -> list[dict]:
|
||||
"""
|
||||
Get the list of recorded mock calls.
|
||||
|
||||
Returns:
|
||||
List of call records, each containing:
|
||||
- provider: Provider name
|
||||
- model: Model name
|
||||
- messages: Messages sent
|
||||
- response_format/tools: Format or tools used
|
||||
- scope: Call scope
|
||||
"""
|
||||
return self._mock_calls
|
||||
|
||||
def clear_mock_calls(self) -> None:
|
||||
"""Clear the recorded mock calls and any set exception."""
|
||||
self._mock_calls = []
|
||||
self._mock_exception = None
|
||||
@@ -0,0 +1,745 @@
|
||||
"""
|
||||
OpenAI-compatible LLM provider supporting OpenAI, Groq, Ollama, and LMStudio.
|
||||
|
||||
This provider handles all OpenAI API-compatible models including:
|
||||
- OpenAI: GPT-4, GPT-4o, GPT-5, o1, o3 (reasoning models)
|
||||
- Groq: Fast inference with seed control and service tiers
|
||||
- Ollama: Local models with native streaming API support
|
||||
- LMStudio: Local models with OpenAI-compatible API
|
||||
|
||||
Features:
|
||||
- Reasoning models with extended thinking (o1, o3, GPT-5 families)
|
||||
- Strict JSON schema enforcement (OpenAI)
|
||||
- Provider-specific parameters (Groq seed, service tier)
|
||||
- Native Ollama streaming for better structured output
|
||||
- Automatic token limit handling per model family
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from openai import APIConnectionError, APIStatusError, AsyncOpenAI, LengthFinishReasonError
|
||||
|
||||
from hindsight_api.config import DEFAULT_LLM_TIMEOUT, ENV_LLM_TIMEOUT
|
||||
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
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Seed applied to every Groq request for deterministic behavior
|
||||
DEFAULT_LLM_SEED = 4242
|
||||
|
||||
|
||||
class OpenAICompatibleLLM(LLMInterface):
|
||||
"""
|
||||
LLM provider for OpenAI-compatible APIs.
|
||||
|
||||
Supports:
|
||||
- OpenAI: Standard models (GPT-4, GPT-4o) and reasoning models (o1, o3, GPT-5)
|
||||
- Groq: Fast inference with seed control and service tiers
|
||||
- Ollama: Local models with native streaming API for better structured output
|
||||
- LMStudio: Local models with OpenAI-compatible API
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
reasoning_effort: str = "low",
|
||||
timeout: float | None = None,
|
||||
groq_service_tier: str | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Initialize OpenAI-compatible LLM provider.
|
||||
|
||||
Args:
|
||||
provider: Provider name ("openai", "groq", "ollama", "lmstudio").
|
||||
api_key: API key (optional for ollama/lmstudio).
|
||||
base_url: Base URL for the API (uses defaults for groq/ollama/lmstudio if empty).
|
||||
model: Model name.
|
||||
reasoning_effort: Reasoning effort level for supported models ("low", "medium", "high").
|
||||
timeout: Request timeout in seconds (uses env var or 300s default).
|
||||
groq_service_tier: Groq service tier ("on_demand", "flex", "auto").
|
||||
**kwargs: Additional provider-specific parameters.
|
||||
"""
|
||||
super().__init__(provider, api_key, base_url, model, reasoning_effort, **kwargs)
|
||||
|
||||
# Validate provider
|
||||
valid_providers = ["openai", "groq", "ollama", "lmstudio"]
|
||||
if self.provider not in valid_providers:
|
||||
raise ValueError(f"OpenAICompatibleLLM only supports: {', '.join(valid_providers)}. Got: {self.provider}")
|
||||
|
||||
# Set default base URLs
|
||||
if not self.base_url:
|
||||
if self.provider == "groq":
|
||||
self.base_url = "https://api.groq.com/openai/v1"
|
||||
elif self.provider == "ollama":
|
||||
self.base_url = "http://localhost:11434/v1"
|
||||
elif self.provider == "lmstudio":
|
||||
self.base_url = "http://localhost:1234/v1"
|
||||
|
||||
# For ollama/lmstudio, use dummy key if not provided
|
||||
if self.provider in ("ollama", "lmstudio") and not self.api_key:
|
||||
self.api_key = "local"
|
||||
|
||||
# Validate API key for cloud providers
|
||||
if self.provider in ("openai", "groq") and not self.api_key:
|
||||
raise ValueError(f"API key is required for {self.provider}")
|
||||
|
||||
# Groq service tier configuration
|
||||
self.groq_service_tier = groq_service_tier or os.getenv("HINDSIGHT_API_LLM_GROQ_SERVICE_TIER", "auto")
|
||||
|
||||
# Get timeout config
|
||||
self.timeout = timeout or float(os.getenv(ENV_LLM_TIMEOUT, str(DEFAULT_LLM_TIMEOUT)))
|
||||
|
||||
# Create OpenAI client
|
||||
client_kwargs: dict[str, Any] = {"api_key": self.api_key, "max_retries": 0}
|
||||
if self.base_url:
|
||||
client_kwargs["base_url"] = self.base_url
|
||||
if self.timeout:
|
||||
client_kwargs["timeout"] = self.timeout
|
||||
|
||||
self._client = AsyncOpenAI(**client_kwargs)
|
||||
logger.info(
|
||||
f"OpenAI-compatible client initialized: provider={self.provider}, model={self.model}, "
|
||||
f"base_url={self.base_url or 'default'}"
|
||||
)
|
||||
|
||||
async def verify_connection(self) -> None:
|
||||
"""
|
||||
Verify that the provider is configured correctly by making a simple test call.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the connection test fails.
|
||||
"""
|
||||
try:
|
||||
logger.info(f"Verifying connection: {self.provider}/{self.model}")
|
||||
await self.call(
|
||||
messages=[{"role": "user", "content": "Say 'ok'"}],
|
||||
max_completion_tokens=100,
|
||||
max_retries=2,
|
||||
initial_backoff=0.5,
|
||||
max_backoff=2.0,
|
||||
)
|
||||
logger.info(f"Connection verified: {self.provider}/{self.model}")
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Connection verification failed for {self.provider}/{self.model}: {e}") from e
|
||||
|
||||
def _supports_reasoning_model(self) -> bool:
|
||||
"""Check if the current model is a reasoning model (o1, o3, GPT-5, DeepSeek)."""
|
||||
model_lower = self.model.lower()
|
||||
return any(x in model_lower for x in ["gpt-5", "o1", "o3", "deepseek"])
|
||||
|
||||
def _get_max_reasoning_tokens(self) -> int | None:
|
||||
"""Get max reasoning tokens for reasoning models."""
|
||||
model_lower = self.model.lower()
|
||||
|
||||
# GPT-4 and GPT-4.1 models have different caps
|
||||
if any(x in model_lower for x in ["gpt-4.1", "gpt-4-"]):
|
||||
return 32000
|
||||
elif "gpt-4o" in model_lower:
|
||||
return 16384
|
||||
|
||||
return None
|
||||
|
||||
async def call(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any | None = None,
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 60.0,
|
||||
skip_validation: bool = False,
|
||||
strict_schema: bool = False,
|
||||
return_usage: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
Make an LLM API call with retry logic.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
response_format: Optional Pydantic model for structured output.
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
skip_validation: Return raw JSON without Pydantic validation.
|
||||
strict_schema: Use strict JSON schema enforcement (OpenAI only).
|
||||
return_usage: If True, return tuple (result, TokenUsage) instead of just result.
|
||||
|
||||
Returns:
|
||||
If return_usage=False: Parsed response if response_format is provided, otherwise text content.
|
||||
If return_usage=True: Tuple of (result, TokenUsage) with token counts.
|
||||
|
||||
Raises:
|
||||
OutputTooLongError: If output exceeds token limits.
|
||||
Exception: Re-raises API errors after retries exhausted.
|
||||
"""
|
||||
# Handle Ollama with native API for structured output (better schema enforcement)
|
||||
if self.provider == "ollama" and response_format is not None:
|
||||
return await self._call_ollama_native(
|
||||
messages=messages,
|
||||
response_format=response_format,
|
||||
max_completion_tokens=max_completion_tokens,
|
||||
temperature=temperature,
|
||||
max_retries=max_retries,
|
||||
initial_backoff=initial_backoff,
|
||||
max_backoff=max_backoff,
|
||||
skip_validation=skip_validation,
|
||||
scope=scope,
|
||||
return_usage=return_usage,
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# Build call parameters
|
||||
call_params: dict[str, Any] = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
}
|
||||
|
||||
# Check if model supports reasoning parameter
|
||||
is_reasoning_model = self._supports_reasoning_model()
|
||||
|
||||
# Apply model-specific token limits
|
||||
if max_completion_tokens is not None:
|
||||
max_tokens_cap = self._get_max_reasoning_tokens()
|
||||
if max_tokens_cap and max_completion_tokens > max_tokens_cap:
|
||||
max_completion_tokens = max_tokens_cap
|
||||
# For reasoning models, enforce minimum to ensure space for reasoning + output
|
||||
if is_reasoning_model and max_completion_tokens < 16000:
|
||||
max_completion_tokens = 16000
|
||||
call_params["max_completion_tokens"] = max_completion_tokens
|
||||
|
||||
# Temperature - reasoning models don't support custom temperature
|
||||
if temperature is not None and not is_reasoning_model:
|
||||
call_params["temperature"] = temperature
|
||||
|
||||
# Set reasoning_effort for reasoning models
|
||||
if is_reasoning_model:
|
||||
call_params["reasoning_effort"] = self.reasoning_effort
|
||||
|
||||
# Provider-specific parameters
|
||||
if self.provider == "groq":
|
||||
call_params["seed"] = DEFAULT_LLM_SEED
|
||||
extra_body: dict[str, Any] = {}
|
||||
# Add service_tier if configured
|
||||
if self.groq_service_tier:
|
||||
extra_body["service_tier"] = self.groq_service_tier
|
||||
# Add reasoning parameters for reasoning models
|
||||
if is_reasoning_model:
|
||||
extra_body["include_reasoning"] = False
|
||||
if extra_body:
|
||||
call_params["extra_body"] = extra_body
|
||||
|
||||
# Prepare response format ONCE before retry loop
|
||||
if response_format is not None:
|
||||
schema = None
|
||||
if hasattr(response_format, "model_json_schema"):
|
||||
schema = response_format.model_json_schema()
|
||||
|
||||
if strict_schema and schema is not None:
|
||||
# Use OpenAI's strict JSON schema enforcement
|
||||
call_params["response_format"] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "response",
|
||||
"strict": True,
|
||||
"schema": schema,
|
||||
},
|
||||
}
|
||||
else:
|
||||
# Soft enforcement: add schema to prompt and use json_object mode
|
||||
if schema is not None:
|
||||
schema_msg = (
|
||||
f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2)}"
|
||||
)
|
||||
|
||||
if call_params["messages"] and call_params["messages"][0].get("role") == "system":
|
||||
first_msg = call_params["messages"][0]
|
||||
if isinstance(first_msg, dict) and isinstance(first_msg.get("content"), str):
|
||||
first_msg["content"] += schema_msg
|
||||
elif call_params["messages"]:
|
||||
first_msg = call_params["messages"][0]
|
||||
if isinstance(first_msg, dict) and isinstance(first_msg.get("content"), str):
|
||||
first_msg["content"] = schema_msg + "\n\n" + first_msg["content"]
|
||||
if self.provider not in ("lmstudio", "ollama"):
|
||||
# LM Studio and Ollama don't support json_object response format reliably
|
||||
call_params["response_format"] = {"type": "json_object"}
|
||||
|
||||
last_exception = None
|
||||
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
if response_format is not None:
|
||||
response = await self._client.chat.completions.create(**call_params)
|
||||
|
||||
content = response.choices[0].message.content
|
||||
|
||||
# Strip reasoning model thinking tags
|
||||
# Supports: <think>, <thinking>, <reasoning>, |startthink|/|endthink|
|
||||
if content:
|
||||
original_len = len(content)
|
||||
content = re.sub(r"<think>.*?</think>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"<thinking>.*?</thinking>", "", content, flags=re.DOTALL)
|
||||
content = re.sub(r"<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")
|
||||
|
||||
# For local models, they may wrap JSON in markdown code blocks
|
||||
if self.provider in ("lmstudio", "ollama"):
|
||||
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
|
||||
json_data = json.loads(content)
|
||||
else:
|
||||
# Log raw LLM response for debugging JSON parse issues
|
||||
try:
|
||||
json_data = json.loads(content)
|
||||
except json.JSONDecodeError as json_err:
|
||||
# Truncate content for logging
|
||||
content_preview = content[:500] if content else "<empty>"
|
||||
if content and len(content) > 700:
|
||||
content_preview = f"{content[:500]}...TRUNCATED...{content[-200:]}"
|
||||
logger.warning(
|
||||
f"JSON parse error from LLM response (attempt {attempt + 1}/{max_retries + 1}): {json_err}\n"
|
||||
f" Model: {self.provider}/{self.model}\n"
|
||||
f" Content length: {len(content) if content else 0} chars\n"
|
||||
f" Content preview: {content_preview!r}\n"
|
||||
f" Finish reason: {response.choices[0].finish_reason if response.choices else 'unknown'}"
|
||||
)
|
||||
# Retry on JSON parse errors
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
last_exception = json_err
|
||||
continue
|
||||
else:
|
||||
logger.error(f"JSON parse error after {max_retries + 1} attempts, giving up")
|
||||
raise
|
||||
|
||||
if skip_validation:
|
||||
result = json_data
|
||||
else:
|
||||
result = response_format.model_validate(json_data)
|
||||
else:
|
||||
response = await self._client.chat.completions.create(**call_params)
|
||||
result = response.choices[0].message.content
|
||||
|
||||
# Record token usage metrics
|
||||
duration = time.time() - start_time
|
||||
usage = response.usage
|
||||
input_tokens = usage.prompt_tokens or 0 if usage else 0
|
||||
output_tokens = usage.completion_tokens or 0 if usage else 0
|
||||
total_tokens = usage.total_tokens or 0 if usage else 0
|
||||
|
||||
# Record LLM metrics
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
)
|
||||
|
||||
# Log slow calls
|
||||
if duration > 10.0 and usage:
|
||||
ratio = max(1, output_tokens) / max(1, input_tokens)
|
||||
cached_tokens = 0
|
||||
if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details:
|
||||
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
|
||||
cache_info = f", cached_tokens={cached_tokens}" if cached_tokens > 0 else ""
|
||||
logger.info(
|
||||
f"slow llm call: scope={scope}, model={self.provider}/{self.model}, "
|
||||
f"input_tokens={input_tokens}, output_tokens={output_tokens}, "
|
||||
f"total_tokens={total_tokens}{cache_info}, time={duration:.3f}s, ratio out/in={ratio:.2f}"
|
||||
)
|
||||
|
||||
if return_usage:
|
||||
token_usage = TokenUsage(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
)
|
||||
return result, token_usage
|
||||
return result
|
||||
|
||||
except LengthFinishReasonError as e:
|
||||
logger.warning(f"LLM output exceeded token limits: {str(e)}")
|
||||
raise OutputTooLongError(
|
||||
"LLM output exceeded token limits. Input may need to be split into smaller chunks."
|
||||
) from e
|
||||
|
||||
except APIConnectionError as e:
|
||||
last_exception = e
|
||||
status_code = getattr(e, "status_code", None) or getattr(
|
||||
getattr(e, "response", None), "status_code", None
|
||||
)
|
||||
logger.warning(f"APIConnectionError (HTTP {status_code}), attempt {attempt + 1}: {str(e)[:200]}")
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
logger.error(f"Connection error after {max_retries + 1} attempts: {str(e)}")
|
||||
raise
|
||||
|
||||
except APIStatusError as e:
|
||||
# Fast fail only on 401 (unauthorized) and 403 (forbidden)
|
||||
if e.status_code in (401, 403):
|
||||
logger.error(f"Auth error (HTTP {e.status_code}), not retrying: {str(e)}")
|
||||
raise
|
||||
|
||||
# Handle tool_use_failed error - model outputted in tool call format
|
||||
if e.status_code == 400 and response_format is not None:
|
||||
try:
|
||||
error_body = e.body if hasattr(e, "body") else {}
|
||||
if isinstance(error_body, dict):
|
||||
error_info: dict[str, Any] = error_body.get("error") or {}
|
||||
if error_info.get("code") == "tool_use_failed":
|
||||
failed_gen = error_info.get("failed_generation", "")
|
||||
if failed_gen:
|
||||
# Parse tool call format and convert to expected format
|
||||
tool_call = json.loads(failed_gen)
|
||||
tool_name = tool_call.get("name", "")
|
||||
tool_args = tool_call.get("arguments", {})
|
||||
converted = {"actions": [{"tool": tool_name, **tool_args}]}
|
||||
if skip_validation:
|
||||
result = converted
|
||||
else:
|
||||
result = response_format.model_validate(converted)
|
||||
|
||||
# Record metrics
|
||||
duration = time.time() - start_time
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
success=True,
|
||||
)
|
||||
if return_usage:
|
||||
return result, TokenUsage(input_tokens=0, output_tokens=0, total_tokens=0)
|
||||
return result
|
||||
except (json.JSONDecodeError, KeyError, TypeError):
|
||||
pass # Failed to parse tool_use_failed, continue with normal retry
|
||||
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
jitter = backoff * 0.2 * (2 * (time.time() % 1) - 1)
|
||||
sleep_time = backoff + jitter
|
||||
await asyncio.sleep(sleep_time)
|
||||
else:
|
||||
logger.error(f"API error after {max_retries + 1} attempts: {str(e)}")
|
||||
raise
|
||||
|
||||
except Exception:
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("LLM call failed after all retries with no exception captured")
|
||||
|
||||
async def call_with_tools(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]],
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "tools",
|
||||
max_retries: int = 5,
|
||||
initial_backoff: float = 1.0,
|
||||
max_backoff: float = 30.0,
|
||||
tool_choice: str | dict[str, Any] = "auto",
|
||||
) -> LLMToolCallResult:
|
||||
"""
|
||||
Make an LLM API call with tool/function calling support.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts. Can include tool results with role='tool'.
|
||||
tools: List of tool definitions in OpenAI format.
|
||||
max_completion_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature (0.0-2.0).
|
||||
scope: Scope identifier for tracking.
|
||||
max_retries: Maximum retry attempts.
|
||||
initial_backoff: Initial backoff time in seconds.
|
||||
max_backoff: Maximum backoff time in seconds.
|
||||
tool_choice: How to choose tools - "auto", "none", "required", or specific function.
|
||||
|
||||
Returns:
|
||||
LLMToolCallResult with content and/or tool_calls.
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
# Build call parameters
|
||||
call_params: dict[str, Any] = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
"tools": tools,
|
||||
"tool_choice": tool_choice,
|
||||
}
|
||||
|
||||
if max_completion_tokens is not None:
|
||||
call_params["max_completion_tokens"] = max_completion_tokens
|
||||
if temperature is not None:
|
||||
call_params["temperature"] = temperature
|
||||
|
||||
# Provider-specific parameters
|
||||
if self.provider == "groq":
|
||||
call_params["seed"] = DEFAULT_LLM_SEED
|
||||
|
||||
last_exception = None
|
||||
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await self._client.chat.completions.create(**call_params)
|
||||
|
||||
message = response.choices[0].message
|
||||
finish_reason = response.choices[0].finish_reason
|
||||
|
||||
# Extract tool calls if present
|
||||
tool_calls: list[LLMToolCall] = []
|
||||
if message.tool_calls:
|
||||
for tc in message.tool_calls:
|
||||
try:
|
||||
args = json.loads(tc.function.arguments) if tc.function.arguments else {}
|
||||
except json.JSONDecodeError:
|
||||
args = {"_raw": tc.function.arguments}
|
||||
tool_calls.append(LLMToolCall(id=tc.id, name=tc.function.name, arguments=args))
|
||||
|
||||
content = message.content
|
||||
|
||||
# Record metrics
|
||||
duration = time.time() - start_time
|
||||
usage = response.usage
|
||||
input_tokens = usage.prompt_tokens or 0 if usage else 0
|
||||
output_tokens = usage.completion_tokens or 0 if usage else 0
|
||||
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
)
|
||||
|
||||
return LLMToolCallResult(
|
||||
content=content,
|
||||
tool_calls=tool_calls,
|
||||
finish_reason=finish_reason,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
)
|
||||
|
||||
except APIConnectionError as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
await asyncio.sleep(min(initial_backoff * (2**attempt), max_backoff))
|
||||
continue
|
||||
raise
|
||||
|
||||
except APIStatusError as e:
|
||||
if e.status_code in (401, 403):
|
||||
raise
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
await asyncio.sleep(min(initial_backoff * (2**attempt), max_backoff))
|
||||
continue
|
||||
raise
|
||||
|
||||
except Exception:
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("Tool call failed after all retries")
|
||||
|
||||
async def _call_ollama_native(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any,
|
||||
max_completion_tokens: int | None,
|
||||
temperature: float | None,
|
||||
max_retries: int,
|
||||
initial_backoff: float,
|
||||
max_backoff: float,
|
||||
skip_validation: bool,
|
||||
scope: str = "memory",
|
||||
return_usage: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
Call Ollama using native API with JSON schema enforcement.
|
||||
|
||||
Ollama's native API supports passing a full JSON schema in the 'format' parameter,
|
||||
which provides better structured output control than the OpenAI-compatible API.
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
# Get the JSON schema from the Pydantic model
|
||||
schema = response_format.model_json_schema() if hasattr(response_format, "model_json_schema") else None
|
||||
|
||||
# Build the base URL for Ollama's native API
|
||||
# Default OpenAI-compatible URL is http://localhost:11434/v1
|
||||
# Native API is at http://localhost:11434/api/chat
|
||||
base_url = self.base_url or "http://localhost:11434/v1"
|
||||
if base_url.endswith("/v1"):
|
||||
native_url = base_url[:-3] + "/api/chat"
|
||||
else:
|
||||
native_url = base_url.rstrip("/") + "/api/chat"
|
||||
|
||||
# Build request payload
|
||||
payload: dict[str, Any] = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
"stream": False,
|
||||
}
|
||||
|
||||
# Add schema as format parameter for structured output
|
||||
if schema:
|
||||
payload["format"] = schema
|
||||
|
||||
# Add optional parameters with optimized defaults for Ollama
|
||||
options: dict[str, Any] = {
|
||||
"num_ctx": 16384, # 16k context window for larger prompts
|
||||
"num_batch": 512, # Optimal batch size for prompt processing
|
||||
}
|
||||
if max_completion_tokens:
|
||||
options["num_predict"] = max_completion_tokens
|
||||
if temperature is not None:
|
||||
options["temperature"] = temperature
|
||||
payload["options"] = options
|
||||
|
||||
last_exception = None
|
||||
|
||||
async with httpx.AsyncClient(timeout=300.0) as client:
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
response = await client.post(native_url, json=payload)
|
||||
response.raise_for_status()
|
||||
|
||||
result = response.json()
|
||||
content = result.get("message", {}).get("content", "")
|
||||
|
||||
# Parse JSON response
|
||||
try:
|
||||
json_data = json.loads(content)
|
||||
except json.JSONDecodeError as json_err:
|
||||
content_preview = content[:500] if content else "<empty>"
|
||||
if content and len(content) > 700:
|
||||
content_preview = f"{content[:500]}...TRUNCATED...{content[-200:]}"
|
||||
logger.warning(
|
||||
f"Ollama JSON parse error (attempt {attempt + 1}/{max_retries + 1}): {json_err}\n"
|
||||
f" Model: ollama/{self.model}\n"
|
||||
f" Content length: {len(content) if content else 0} chars\n"
|
||||
f" Content preview: {content_preview!r}"
|
||||
)
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
last_exception = json_err
|
||||
continue
|
||||
else:
|
||||
raise
|
||||
|
||||
# Extract token usage from Ollama response
|
||||
duration = time.time() - start_time
|
||||
input_tokens = result.get("prompt_eval_count", 0) or 0
|
||||
output_tokens = result.get("eval_count", 0) or 0
|
||||
total_tokens = input_tokens + output_tokens
|
||||
|
||||
# Record LLM metrics
|
||||
metrics = get_metrics_collector()
|
||||
metrics.record_llm_call(
|
||||
provider=self.provider,
|
||||
model=self.model,
|
||||
scope=scope,
|
||||
duration=duration,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
)
|
||||
|
||||
# Validate against Pydantic model or return raw JSON
|
||||
if skip_validation:
|
||||
validated_result = json_data
|
||||
else:
|
||||
validated_result = response_format.model_validate(json_data)
|
||||
|
||||
if return_usage:
|
||||
token_usage = TokenUsage(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
)
|
||||
return validated_result, token_usage
|
||||
return validated_result
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
logger.warning(
|
||||
f"Ollama HTTP error (attempt {attempt + 1}/{max_retries + 1}): {e.response.status_code}"
|
||||
)
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
logger.error(f"Ollama HTTP error after {max_retries + 1} attempts: {e}")
|
||||
raise
|
||||
|
||||
except httpx.RequestError as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
logger.warning(f"Ollama connection error (attempt {attempt + 1}/{max_retries + 1}): {e}")
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
logger.error(f"Ollama connection error after {max_retries + 1} attempts: {e}")
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected error during Ollama call: {type(e).__name__}: {e}")
|
||||
raise
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError("Ollama call failed after all retries")
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
"""Clean up resources (close OpenAI client connections)."""
|
||||
if hasattr(self, "_client") and self._client:
|
||||
await self._client.close()
|
||||
@@ -0,0 +1,18 @@
|
||||
"""
|
||||
Reflect agent module for agentic reflection with tools.
|
||||
|
||||
The reflect agent uses an iterative loop with tools to:
|
||||
1. Lookup mental models (existing knowledge)
|
||||
2. Recall facts (semantic + temporal search)
|
||||
3. Expand memories (get chunk/document context)
|
||||
"""
|
||||
|
||||
from .agent import ReflectAgentResult, run_reflect_agent
|
||||
from .models import ReflectAction, ReflectActionBatch
|
||||
|
||||
__all__ = [
|
||||
"run_reflect_agent",
|
||||
"ReflectAgentResult",
|
||||
"ReflectAction",
|
||||
"ReflectActionBatch",
|
||||
]
|
||||
@@ -0,0 +1,933 @@
|
||||
"""
|
||||
Reflect agent - agentic loop for reflection with native tool calling.
|
||||
|
||||
Uses hierarchical retrieval:
|
||||
1. search_mental_models - User-curated summaries (highest quality)
|
||||
2. search_observations - Consolidated knowledge with freshness
|
||||
3. recall - Raw facts as ground truth
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable
|
||||
|
||||
from .models import DirectiveInfo, LLMCall, ReflectAgentResult, TokenUsageSummary, ToolCall
|
||||
from .prompts import FINAL_SYSTEM_PROMPT, _extract_directive_rules, build_final_prompt, build_system_prompt_for_tools
|
||||
from .tools_schema import get_reflect_tools
|
||||
|
||||
|
||||
def _build_directives_applied(directives: list[dict[str, Any]] | None) -> list[DirectiveInfo]:
|
||||
"""Build list of DirectiveInfo from directive mental models.
|
||||
|
||||
Handles multiple directive formats:
|
||||
1. New format: directives have direct 'content' field
|
||||
2. Fallback: directives have 'description' field
|
||||
"""
|
||||
if not directives:
|
||||
return []
|
||||
|
||||
result = []
|
||||
for directive in directives:
|
||||
directive_id = directive.get("id", "")
|
||||
directive_name = directive.get("name", "")
|
||||
|
||||
# Get content from 'content' field or fallback to 'description'
|
||||
content = directive.get("content", "") or directive.get("description", "")
|
||||
|
||||
result.append(DirectiveInfo(id=directive_id, name=directive_name, content=content))
|
||||
|
||||
return result
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..llm_wrapper import LLMProvider
|
||||
from ..response_models import LLMToolCall
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_MAX_ITERATIONS = 10
|
||||
|
||||
|
||||
def _normalize_tool_name(name: str) -> str:
|
||||
"""Normalize tool name from various LLM output formats.
|
||||
|
||||
Some LLMs output tool names in non-standard formats:
|
||||
- 'functions.done' (OpenAI-style prefix)
|
||||
- 'call=functions.done' (some models)
|
||||
- 'call=done' (some models)
|
||||
- 'done<|channel|>commentary' (malformed special tokens appended)
|
||||
|
||||
Returns the normalized tool name (e.g., 'done', 'recall', etc.)
|
||||
"""
|
||||
# Handle 'call=functions.name' or 'call=name' format
|
||||
if name.startswith("call="):
|
||||
name = name[len("call=") :]
|
||||
|
||||
# Handle 'functions.name' format
|
||||
if name.startswith("functions."):
|
||||
name = name[len("functions.") :]
|
||||
|
||||
# Handle malformed special tokens appended to tool name
|
||||
# e.g., 'done<|channel|>commentary' -> 'done'
|
||||
if "<|" in name:
|
||||
name = name.split("<|")[0]
|
||||
|
||||
return name
|
||||
|
||||
|
||||
def _is_done_tool(name: str) -> bool:
|
||||
"""Check if the tool name represents the 'done' tool."""
|
||||
return _normalize_tool_name(name) == "done"
|
||||
|
||||
|
||||
# Pattern to match done() call as text - handles done({...}) with nested JSON
|
||||
_DONE_CALL_PATTERN = re.compile(r"done\s*\(\s*\{.*$", re.DOTALL)
|
||||
|
||||
# Patterns for leaked structured output in the answer field
|
||||
_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
|
||||
)
|
||||
|
||||
|
||||
def _clean_answer_text(text: str) -> str:
|
||||
"""Clean up answer text by removing any done() tool call syntax.
|
||||
|
||||
Some LLMs output the done() call as text instead of a proper tool call.
|
||||
This strips out patterns like: done({"answer": "...", ...})
|
||||
"""
|
||||
# Remove done() call pattern from the end of the text
|
||||
cleaned = _DONE_CALL_PATTERN.sub("", text).strip()
|
||||
return cleaned if cleaned else text
|
||||
|
||||
|
||||
def _clean_done_answer(text: str) -> str:
|
||||
"""Clean up the answer field from a done() tool call.
|
||||
|
||||
Some LLMs leak structured output patterns into the answer text, such as:
|
||||
- JSON code blocks with observation_ids/memory_ids at the end
|
||||
- Raw JSON objects with these fields
|
||||
- Plain text like "observation_ids: [...]"
|
||||
|
||||
This cleans those patterns while preserving the actual answer content.
|
||||
"""
|
||||
if not text:
|
||||
return text
|
||||
|
||||
cleaned = text
|
||||
|
||||
# Remove leaked JSON in code blocks at the end
|
||||
cleaned = _LEAKED_JSON_SUFFIX.sub("", cleaned).strip()
|
||||
|
||||
# Remove leaked raw JSON objects at the end
|
||||
cleaned = _LEAKED_JSON_OBJECT.sub("", cleaned).strip()
|
||||
|
||||
# Remove trailing ID patterns
|
||||
cleaned = _TRAILING_IDS_PATTERN.sub("", cleaned).strip()
|
||||
|
||||
return cleaned if cleaned else text
|
||||
|
||||
|
||||
async def _generate_structured_output(
|
||||
answer: str,
|
||||
response_schema: dict,
|
||||
llm_config: "LLMProvider",
|
||||
reflect_id: str,
|
||||
) -> tuple[dict[str, Any] | None, int, int]:
|
||||
"""Generate structured output from an answer using the provided JSON schema.
|
||||
|
||||
Args:
|
||||
answer: The text answer to extract structured data from
|
||||
response_schema: JSON Schema for the expected output structure
|
||||
llm_config: LLM provider for making the extraction call
|
||||
reflect_id: Reflect ID for logging
|
||||
|
||||
Returns:
|
||||
Tuple of (structured_output, input_tokens, output_tokens).
|
||||
structured_output is None if generation fails.
|
||||
"""
|
||||
try:
|
||||
from typing import Any as TypingAny
|
||||
|
||||
from pydantic import create_model
|
||||
|
||||
def _json_schema_type_to_python(field_schema: dict) -> type:
|
||||
"""Map JSON schema type to Python type for better LLM guidance."""
|
||||
json_type = field_schema.get("type", "string")
|
||||
if json_type == "array":
|
||||
return list
|
||||
elif json_type == "object":
|
||||
return dict
|
||||
elif json_type == "integer":
|
||||
return int
|
||||
elif json_type == "number":
|
||||
return float
|
||||
elif json_type == "boolean":
|
||||
return bool
|
||||
else:
|
||||
return str
|
||||
|
||||
# Build fields from JSON schema properties
|
||||
schema_props = response_schema.get("properties", {})
|
||||
required_fields = set(response_schema.get("required", []))
|
||||
fields: dict[str, TypingAny] = {}
|
||||
for field_name, field_schema in schema_props.items():
|
||||
field_type = _json_schema_type_to_python(field_schema)
|
||||
default = ... if field_name in required_fields else None
|
||||
fields[field_name] = (field_type, default)
|
||||
|
||||
if not fields:
|
||||
logger.warning(f"[REFLECT {reflect_id}] No fields found in response_schema, skipping structured output")
|
||||
return None, 0, 0
|
||||
|
||||
DynamicModel = create_model("StructuredResponse", **fields)
|
||||
|
||||
# Include the full schema in the prompt for better LLM guidance
|
||||
schema_str = json.dumps(response_schema, indent=2)
|
||||
|
||||
# Build field descriptions for the prompt
|
||||
field_descriptions = []
|
||||
for field_name, field_schema in schema_props.items():
|
||||
field_type = field_schema.get("type", "string")
|
||||
field_desc = field_schema.get("description", "")
|
||||
is_required = field_name in required_fields
|
||||
req_marker = " (REQUIRED)" if is_required else " (optional)"
|
||||
field_descriptions.append(f"- {field_name} ({field_type}){req_marker}: {field_desc}")
|
||||
fields_text = "\n".join(field_descriptions)
|
||||
|
||||
# Call LLM with the answer to extract structured data
|
||||
structured_prompt = f"""Your task is to extract specific information from the answer below and format it as JSON.
|
||||
|
||||
ANSWER TO EXTRACT FROM:
|
||||
\"\"\"
|
||||
{answer}
|
||||
\"\"\"
|
||||
|
||||
REQUIRED OUTPUT FORMAT - Extract the following fields from the answer above:
|
||||
{fields_text}
|
||||
|
||||
JSON Schema:
|
||||
```json
|
||||
{schema_str}
|
||||
```
|
||||
|
||||
INSTRUCTIONS:
|
||||
1. Read the answer carefully and identify the information that matches each field
|
||||
2. Extract the ACTUAL content from the answer - do NOT leave fields empty if information is present
|
||||
3. For string fields: use the exact text or a clear summary from the answer
|
||||
4. For array fields: return a JSON array (e.g., ["item1", "item2"]), NOT a string
|
||||
5. For required fields: you MUST provide a value extracted from the answer
|
||||
6. Return ONLY the JSON object, no explanation
|
||||
|
||||
OUTPUT:"""
|
||||
|
||||
structured_result, usage = await llm_config.call(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a precise data extraction assistant. Extract information from text and return it as valid JSON matching the provided schema. Always extract actual content - never return empty strings for required fields if information is available.",
|
||||
},
|
||||
{"role": "user", "content": structured_prompt},
|
||||
],
|
||||
response_format=DynamicModel,
|
||||
scope="reflect_structured",
|
||||
skip_validation=True, # We'll handle the dict ourselves
|
||||
return_usage=True,
|
||||
)
|
||||
|
||||
# Convert to dict
|
||||
if hasattr(structured_result, "model_dump"):
|
||||
structured_output = structured_result.model_dump()
|
||||
elif isinstance(structured_result, dict):
|
||||
structured_output = structured_result
|
||||
else:
|
||||
# Try to parse as JSON
|
||||
structured_output = json.loads(str(structured_result))
|
||||
|
||||
# Validate that required fields have non-empty values
|
||||
for field_name in required_fields:
|
||||
value = structured_output.get(field_name)
|
||||
if value is None or value == "" or value == []:
|
||||
logger.warning(f"[REFLECT {reflect_id}] Required field '{field_name}' is empty in structured output")
|
||||
|
||||
logger.info(f"[REFLECT {reflect_id}] Generated structured output with {len(structured_output)} fields")
|
||||
return structured_output, usage.input_tokens, usage.output_tokens
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"[REFLECT {reflect_id}] Failed to generate structured output: {e}")
|
||||
return None, 0, 0
|
||||
|
||||
|
||||
async def run_reflect_agent(
|
||||
llm_config: "LLMProvider",
|
||||
bank_id: str,
|
||||
query: str,
|
||||
bank_profile: dict[str, Any],
|
||||
search_mental_models_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
search_observations_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
recall_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
|
||||
context: str | None = None,
|
||||
max_iterations: int = DEFAULT_MAX_ITERATIONS,
|
||||
max_tokens: int | None = None,
|
||||
response_schema: dict | None = None,
|
||||
directives: list[dict[str, Any]] | None = None,
|
||||
has_mental_models: bool = False,
|
||||
budget: str | None = None,
|
||||
) -> ReflectAgentResult:
|
||||
"""
|
||||
Execute the reflect agent loop using native tool calling.
|
||||
|
||||
The agent uses hierarchical retrieval:
|
||||
1. search_mental_models - User-curated summaries (try first)
|
||||
2. search_observations - Consolidated knowledge with freshness
|
||||
3. recall - Raw facts as ground truth
|
||||
|
||||
Args:
|
||||
llm_config: LLM provider for agent calls
|
||||
bank_id: Bank identifier
|
||||
query: Question to answer
|
||||
bank_profile: Bank profile with name and mission
|
||||
search_mental_models_fn: Tool callback for searching mental models (query, max_results) -> result
|
||||
search_observations_fn: Tool callback for searching observations (query, max_results) -> result
|
||||
recall_fn: Tool callback for recall (query, max_tokens) -> result
|
||||
expand_fn: Tool callback for expand (memory_ids, depth) -> result
|
||||
context: Optional additional context
|
||||
max_iterations: Maximum number of iterations before forcing response
|
||||
max_tokens: Maximum tokens for the final response
|
||||
response_schema: Optional JSON Schema for structured output in final response
|
||||
directives: Optional list of directive mental models to inject as hard rules
|
||||
|
||||
Returns:
|
||||
ReflectAgentResult with final answer and metadata
|
||||
"""
|
||||
reflect_id = f"{bank_id[:8]}-{int(time.time() * 1000) % 100000}"
|
||||
start_time = time.time()
|
||||
|
||||
# Build directives_applied for the trace
|
||||
directives_applied = _build_directives_applied(directives)
|
||||
|
||||
# Extract directive rules for tool schema (if any)
|
||||
directive_rules = _extract_directive_rules(directives) if directives else None
|
||||
|
||||
# Get tools for this agent (with directive compliance field if directives exist)
|
||||
tools = get_reflect_tools(directive_rules=directive_rules)
|
||||
|
||||
# Build initial messages (directives are injected into system prompt at START and END)
|
||||
system_prompt = build_system_prompt_for_tools(
|
||||
bank_profile, context, directives=directives, has_mental_models=has_mental_models, budget=budget
|
||||
)
|
||||
messages: list[dict[str, Any]] = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": query},
|
||||
]
|
||||
|
||||
# Tracking
|
||||
total_tools_called = 0
|
||||
tool_trace: list[ToolCall] = []
|
||||
tool_trace_summary: list[dict[str, Any]] = []
|
||||
llm_trace: list[dict[str, Any]] = []
|
||||
context_history: list[dict[str, Any]] = [] # For final prompt fallback
|
||||
|
||||
# Token usage tracking - accumulate across all LLM calls
|
||||
total_input_tokens = 0
|
||||
total_output_tokens = 0
|
||||
|
||||
# Track available IDs for validation (prevents hallucinated citations)
|
||||
available_memory_ids: set[str] = set()
|
||||
available_mental_model_ids: set[str] = set()
|
||||
available_observation_ids: set[str] = set()
|
||||
|
||||
def _get_llm_trace() -> list[LLMCall]:
|
||||
return [
|
||||
LLMCall(
|
||||
scope=c["scope"],
|
||||
duration_ms=c["duration_ms"],
|
||||
input_tokens=c.get("input_tokens", 0),
|
||||
output_tokens=c.get("output_tokens", 0),
|
||||
)
|
||||
for c in llm_trace
|
||||
]
|
||||
|
||||
def _get_usage() -> TokenUsageSummary:
|
||||
return TokenUsageSummary(
|
||||
input_tokens=total_input_tokens,
|
||||
output_tokens=total_output_tokens,
|
||||
total_tokens=total_input_tokens + total_output_tokens,
|
||||
)
|
||||
|
||||
def _log_completion(answer: str, iterations: int, forced: bool = False):
|
||||
elapsed_ms = int((time.time() - start_time) * 1000)
|
||||
tools_summary = (
|
||||
", ".join(
|
||||
f"{t['tool']}({t['input_summary']})={t['duration_ms']}ms/{t.get('output_chars', 0)}c"
|
||||
for t in tool_trace_summary
|
||||
)
|
||||
or "none"
|
||||
)
|
||||
llm_summary = ", ".join(f"{c['scope']}={c['duration_ms']}ms" for c in llm_trace) or "none"
|
||||
total_llm_ms = sum(c["duration_ms"] for c in llm_trace)
|
||||
total_tools_ms = sum(t["duration_ms"] for t in tool_trace_summary)
|
||||
|
||||
answer_preview = answer[:100] + "..." if len(answer) > 100 else answer
|
||||
mode = "forced" if forced else "done"
|
||||
logger.info(
|
||||
f"[REFLECT {reflect_id}] {mode} | "
|
||||
f"query='{query[:50]}...' | "
|
||||
f"iterations={iterations} | "
|
||||
f"llm=[{llm_summary}] ({total_llm_ms}ms) | "
|
||||
f"tools=[{tools_summary}] ({total_tools_ms}ms) | "
|
||||
f"answer='{answer_preview}' | "
|
||||
f"total={elapsed_ms}ms"
|
||||
)
|
||||
|
||||
for iteration in range(max_iterations):
|
||||
is_last = iteration == max_iterations - 1
|
||||
|
||||
if is_last:
|
||||
# Force text response on last iteration - no tools
|
||||
prompt = build_final_prompt(query, context_history, bank_profile, context)
|
||||
llm_start = time.time()
|
||||
response, usage = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
scope="reflect_agent_final",
|
||||
max_completion_tokens=max_tokens,
|
||||
return_usage=True,
|
||||
)
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
"duration_ms": llm_duration,
|
||||
"input_tokens": usage.input_tokens,
|
||||
"output_tokens": usage.output_tokens,
|
||||
}
|
||||
)
|
||||
answer = _clean_answer_text(response.strip())
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iteration + 1,
|
||||
tools_called=total_tools_called,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
usage=_get_usage(),
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
|
||||
# Call LLM with tools
|
||||
llm_start = time.time()
|
||||
|
||||
try:
|
||||
result = await llm_config.call_with_tools(
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
scope="reflect_agent",
|
||||
tool_choice="required" if iteration == 0 else "auto", # Force tool use on first iteration
|
||||
)
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += result.input_tokens
|
||||
total_output_tokens += result.output_tokens
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": f"agent_{iteration + 1}",
|
||||
"duration_ms": llm_duration,
|
||||
"input_tokens": result.input_tokens,
|
||||
"output_tokens": result.output_tokens,
|
||||
}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
err_duration = int((time.time() - llm_start) * 1000)
|
||||
logger.warning(f"[REFLECT {reflect_id}] LLM error on iteration {iteration + 1}: {e} ({err_duration}ms)")
|
||||
llm_trace.append({"scope": f"agent_{iteration + 1}_err", "duration_ms": err_duration})
|
||||
# Guardrail: If no evidence gathered yet, retry
|
||||
has_gathered_evidence = (
|
||||
bool(available_memory_ids) or bool(available_mental_model_ids) or bool(available_observation_ids)
|
||||
)
|
||||
if not has_gathered_evidence and iteration < max_iterations - 1:
|
||||
continue
|
||||
prompt = build_final_prompt(query, context_history, bank_profile, context)
|
||||
llm_start = time.time()
|
||||
response, usage = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
scope="reflect_agent_final",
|
||||
max_completion_tokens=max_tokens,
|
||||
return_usage=True,
|
||||
)
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
"duration_ms": llm_duration,
|
||||
"input_tokens": usage.input_tokens,
|
||||
"output_tokens": usage.output_tokens,
|
||||
}
|
||||
)
|
||||
answer = _clean_answer_text(response.strip())
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iteration + 1,
|
||||
tools_called=total_tools_called,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
usage=_get_usage(),
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
|
||||
# No tool calls - LLM wants to respond with text
|
||||
if not result.tool_calls:
|
||||
if result.content:
|
||||
answer = _clean_answer_text(result.content.strip())
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
|
||||
_log_completion(answer, iteration + 1)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iteration + 1,
|
||||
tools_called=total_tools_called,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
usage=_get_usage(),
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
# Empty response, force final
|
||||
prompt = build_final_prompt(query, context_history, bank_profile, context)
|
||||
llm_start = time.time()
|
||||
response, usage = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
scope="reflect_agent_final",
|
||||
max_completion_tokens=max_tokens,
|
||||
return_usage=True,
|
||||
)
|
||||
llm_duration = int((time.time() - llm_start) * 1000)
|
||||
total_input_tokens += usage.input_tokens
|
||||
total_output_tokens += usage.output_tokens
|
||||
llm_trace.append(
|
||||
{
|
||||
"scope": "final",
|
||||
"duration_ms": llm_duration,
|
||||
"input_tokens": usage.input_tokens,
|
||||
"output_tokens": usage.output_tokens,
|
||||
}
|
||||
)
|
||||
answer = _clean_answer_text(response.strip())
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
total_input_tokens += struct_in
|
||||
total_output_tokens += struct_out
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iteration + 1,
|
||||
tools_called=total_tools_called,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
usage=_get_usage(),
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
|
||||
# Check for done tool call (handle various LLM output formats)
|
||||
done_call = next((tc for tc in result.tool_calls if _is_done_tool(tc.name)), None)
|
||||
if done_call:
|
||||
# Guardrail: Require evidence before done
|
||||
has_gathered_evidence = (
|
||||
bool(available_memory_ids) or bool(available_mental_model_ids) or bool(available_observation_ids)
|
||||
)
|
||||
if not has_gathered_evidence and iteration < max_iterations - 1:
|
||||
# Add assistant message and fake tool result asking for evidence
|
||||
messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [_tool_call_to_dict(done_call)],
|
||||
}
|
||||
)
|
||||
messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": done_call.id,
|
||||
"name": done_call.name, # Required by Gemini
|
||||
"content": json.dumps(
|
||||
{
|
||||
"error": "You must search for information first. Use search_mental_models(), search_observations(), or recall() before providing your final answer."
|
||||
}
|
||||
),
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
# Process done tool
|
||||
return await _process_done_tool(
|
||||
done_call,
|
||||
available_memory_ids,
|
||||
available_mental_model_ids,
|
||||
available_observation_ids,
|
||||
iteration + 1,
|
||||
total_tools_called,
|
||||
tool_trace,
|
||||
_get_llm_trace(),
|
||||
_get_usage(),
|
||||
_log_completion,
|
||||
reflect_id,
|
||||
directives_applied=directives_applied,
|
||||
llm_config=llm_config,
|
||||
response_schema=response_schema,
|
||||
)
|
||||
|
||||
# Execute other tools in parallel (exclude done tool in all its format variants)
|
||||
other_tools = [tc for tc in result.tool_calls if not _is_done_tool(tc.name)]
|
||||
if other_tools:
|
||||
# Add assistant message with tool calls
|
||||
messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [_tool_call_to_dict(tc) for tc in other_tools],
|
||||
}
|
||||
)
|
||||
|
||||
# Execute tools in parallel
|
||||
tool_tasks = [
|
||||
_execute_tool_with_timing(
|
||||
tc,
|
||||
search_mental_models_fn,
|
||||
search_observations_fn,
|
||||
recall_fn,
|
||||
expand_fn,
|
||||
)
|
||||
for tc in other_tools
|
||||
]
|
||||
tool_results = await asyncio.gather(*tool_tasks, return_exceptions=True)
|
||||
total_tools_called += len(other_tools)
|
||||
|
||||
# Process results and add to messages
|
||||
for tc, result_data in zip(other_tools, tool_results):
|
||||
if isinstance(result_data, Exception):
|
||||
# Tool execution failed - send error back to LLM so it can try again
|
||||
logger.warning(f"[REFLECT {reflect_id}] Tool {tc.name} failed with exception: {result_data}")
|
||||
output = {"error": f"Tool execution failed: {result_data}"}
|
||||
duration_ms = 0
|
||||
else:
|
||||
output, duration_ms = result_data
|
||||
|
||||
# Normalize tool name for consistent tracking
|
||||
normalized_tool_name = _normalize_tool_name(tc.name)
|
||||
|
||||
# Check if tool returned an error response - log but continue (LLM will see the error)
|
||||
if isinstance(output, dict) and "error" in output:
|
||||
logger.warning(
|
||||
f"[REFLECT {reflect_id}] Tool {normalized_tool_name} returned error: {output['error']}"
|
||||
)
|
||||
|
||||
# Track available IDs from tool results (only for successful responses)
|
||||
if (
|
||||
normalized_tool_name == "search_mental_models"
|
||||
and isinstance(output, dict)
|
||||
and "mental_models" in output
|
||||
):
|
||||
for mm in output["mental_models"]:
|
||||
if "id" in mm:
|
||||
available_mental_model_ids.add(mm["id"])
|
||||
|
||||
if (
|
||||
normalized_tool_name == "search_observations"
|
||||
and isinstance(output, dict)
|
||||
and "observations" in output
|
||||
):
|
||||
for obs in output["observations"]:
|
||||
if "id" in obs:
|
||||
available_observation_ids.add(obs["id"])
|
||||
|
||||
if normalized_tool_name == "recall" and isinstance(output, dict) and "memories" in output:
|
||||
for memory in output["memories"]:
|
||||
if "id" in memory:
|
||||
available_memory_ids.add(memory["id"])
|
||||
|
||||
# Add tool result message
|
||||
messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tc.id,
|
||||
"name": tc.name, # Required by Gemini
|
||||
"content": json.dumps(output, default=str),
|
||||
}
|
||||
)
|
||||
|
||||
# Track for logging and context history
|
||||
input_dict = {"tool": tc.name, **tc.arguments}
|
||||
input_summary = _summarize_input(tc.name, tc.arguments)
|
||||
|
||||
# Extract reason from tool arguments (if provided)
|
||||
tool_reason = tc.arguments.get("reason")
|
||||
|
||||
tool_trace.append(
|
||||
ToolCall(
|
||||
tool=tc.name,
|
||||
reason=tool_reason,
|
||||
input=input_dict,
|
||||
output=output,
|
||||
duration_ms=duration_ms,
|
||||
iteration=iteration + 1,
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
output_chars = len(json.dumps(output))
|
||||
except (TypeError, ValueError):
|
||||
output_chars = len(str(output))
|
||||
|
||||
tool_trace_summary.append(
|
||||
{
|
||||
"tool": tc.name,
|
||||
"input_summary": input_summary,
|
||||
"duration_ms": duration_ms,
|
||||
"output_chars": output_chars,
|
||||
}
|
||||
)
|
||||
|
||||
# Keep context history for fallback final prompt
|
||||
context_history.append({"tool": tc.name, "input": input_dict, "output": output})
|
||||
|
||||
# Should not reach here
|
||||
answer = "I was unable to formulate a complete answer within the iteration limit."
|
||||
_log_completion(answer, max_iterations, forced=True)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
iterations=max_iterations,
|
||||
tools_called=total_tools_called,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
usage=_get_usage(),
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
|
||||
|
||||
def _tool_call_to_dict(tc: "LLMToolCall") -> dict[str, Any]:
|
||||
"""Convert LLMToolCall to OpenAI message format."""
|
||||
return {
|
||||
"id": tc.id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc.name,
|
||||
"arguments": json.dumps(tc.arguments),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
async def _process_done_tool(
|
||||
done_call: "LLMToolCall",
|
||||
available_memory_ids: set[str],
|
||||
available_mental_model_ids: set[str],
|
||||
available_observation_ids: set[str],
|
||||
iterations: int,
|
||||
total_tools_called: int,
|
||||
tool_trace: list[ToolCall],
|
||||
llm_trace: list[LLMCall],
|
||||
usage: TokenUsageSummary,
|
||||
log_completion: Callable,
|
||||
reflect_id: str,
|
||||
directives_applied: list[DirectiveInfo],
|
||||
llm_config: "LLMProvider | None" = None,
|
||||
response_schema: dict | None = None,
|
||||
) -> ReflectAgentResult:
|
||||
"""Process the done tool call and return the result."""
|
||||
args = done_call.arguments
|
||||
|
||||
# Extract and clean the answer - some LLMs leak structured output into the answer text
|
||||
raw_answer = args.get("answer", "").strip()
|
||||
answer = _clean_done_answer(raw_answer) if raw_answer else ""
|
||||
if not answer:
|
||||
answer = "No answer provided."
|
||||
|
||||
# Validate IDs (only include IDs that were actually retrieved)
|
||||
used_memory_ids = [mid for mid in args.get("memory_ids", []) if mid in available_memory_ids]
|
||||
used_mental_model_ids = [mid for mid in args.get("mental_model_ids", []) if mid in available_mental_model_ids]
|
||||
used_observation_ids = [oid for oid in args.get("observation_ids", []) if oid in available_observation_ids]
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
final_usage = usage
|
||||
if response_schema and llm_config and answer:
|
||||
structured_output, struct_in, struct_out = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
# Add structured output tokens to usage
|
||||
final_usage = TokenUsageSummary(
|
||||
input_tokens=usage.input_tokens + struct_in,
|
||||
output_tokens=usage.output_tokens + struct_out,
|
||||
total_tokens=usage.total_tokens + struct_in + struct_out,
|
||||
)
|
||||
|
||||
log_completion(answer, iterations)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iterations,
|
||||
tools_called=total_tools_called,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=llm_trace,
|
||||
usage=final_usage,
|
||||
used_memory_ids=used_memory_ids,
|
||||
used_mental_model_ids=used_mental_model_ids,
|
||||
used_observation_ids=used_observation_ids,
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
|
||||
|
||||
async def _execute_tool_with_timing(
|
||||
tc: "LLMToolCall",
|
||||
search_mental_models_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
search_observations_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
recall_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
|
||||
) -> tuple[dict[str, Any], int]:
|
||||
"""Execute a tool call and return result with timing."""
|
||||
start = time.time()
|
||||
result = await _execute_tool(
|
||||
tc.name,
|
||||
tc.arguments,
|
||||
search_mental_models_fn,
|
||||
search_observations_fn,
|
||||
recall_fn,
|
||||
expand_fn,
|
||||
)
|
||||
duration_ms = int((time.time() - start) * 1000)
|
||||
return result, duration_ms
|
||||
|
||||
|
||||
async def _execute_tool(
|
||||
tool_name: str,
|
||||
args: dict[str, Any],
|
||||
search_mental_models_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
search_observations_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
recall_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
|
||||
) -> dict[str, Any]:
|
||||
"""Execute a single tool by name."""
|
||||
# Normalize tool name for various LLM output formats
|
||||
tool_name = _normalize_tool_name(tool_name)
|
||||
|
||||
if tool_name == "search_mental_models":
|
||||
query = args.get("query")
|
||||
if not query:
|
||||
return {"error": "search_mental_models requires a query parameter"}
|
||||
max_results = int(args.get("max_results") or 5)
|
||||
return await search_mental_models_fn(query, max_results)
|
||||
|
||||
elif tool_name == "search_observations":
|
||||
query = args.get("query")
|
||||
if not query:
|
||||
return {"error": "search_observations requires a query parameter"}
|
||||
max_tokens = max(int(args.get("max_tokens") or 5000), 1000) # Default 5000, min 1000
|
||||
return await search_observations_fn(query, max_tokens)
|
||||
|
||||
elif tool_name == "recall":
|
||||
query = args.get("query")
|
||||
if not query:
|
||||
return {"error": "recall requires a query parameter"}
|
||||
max_tokens = max(int(args.get("max_tokens") or 2048), 1000) # Default 2048, min 1000
|
||||
return await recall_fn(query, max_tokens)
|
||||
|
||||
elif tool_name == "expand":
|
||||
memory_ids = args.get("memory_ids", [])
|
||||
if not memory_ids:
|
||||
return {"error": "expand requires memory_ids"}
|
||||
depth = args.get("depth", "chunk")
|
||||
return await expand_fn(memory_ids, depth)
|
||||
|
||||
else:
|
||||
return {"error": f"Unknown tool: {tool_name}"}
|
||||
|
||||
|
||||
def _summarize_input(tool_name: str, args: dict[str, Any]) -> str:
|
||||
"""Create a summary of tool input for logging, showing all params."""
|
||||
if tool_name == "search_mental_models":
|
||||
query = args.get("query", "")
|
||||
query_preview = f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
|
||||
max_results = int(args.get("max_results") or 5)
|
||||
return f"(query={query_preview}, max_results={max_results})"
|
||||
elif tool_name == "search_observations":
|
||||
query = args.get("query", "")
|
||||
query_preview = f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
|
||||
max_tokens = max(int(args.get("max_tokens") or 5000), 1000)
|
||||
return f"(query={query_preview}, max_tokens={max_tokens})"
|
||||
elif tool_name == "recall":
|
||||
query = args.get("query", "")
|
||||
query_preview = f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
|
||||
# Show actual value used (default 2048, min 1000)
|
||||
max_tokens = max(int(args.get("max_tokens") or 2048), 1000)
|
||||
return f"(query={query_preview}, max_tokens={max_tokens})"
|
||||
elif tool_name == "expand":
|
||||
memory_ids = args.get("memory_ids", [])
|
||||
depth = args.get("depth", "chunk")
|
||||
return f"(memory_ids=[{len(memory_ids)} ids], depth={depth})"
|
||||
elif tool_name == "done":
|
||||
answer = args.get("answer", "")
|
||||
answer_preview = f"'{answer[:30]}...'" if len(answer) > 30 else f"'{answer}'"
|
||||
memory_ids = args.get("memory_ids", [])
|
||||
mental_model_ids = args.get("mental_model_ids", [])
|
||||
observation_ids = args.get("observation_ids", [])
|
||||
return (
|
||||
f"(answer={answer_preview}, mem={len(memory_ids)}, mm={len(mental_model_ids)}, obs={len(observation_ids)})"
|
||||
)
|
||||
return str(args)
|
||||
@@ -0,0 +1,109 @@
|
||||
"""
|
||||
Pydantic models for the reflect agent.
|
||||
"""
|
||||
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class ObservationSection(BaseModel):
|
||||
"""A section within an observation with its supporting memories."""
|
||||
|
||||
title: str = Field(description="Section header (can be empty for intro)")
|
||||
text: str = Field(description="Section content - no headers, use lists/tables/bold")
|
||||
memory_ids: list[str] = Field(default_factory=list, description="Memory IDs supporting this section")
|
||||
|
||||
|
||||
class ReflectAction(BaseModel):
|
||||
"""Single action the reflect agent can take."""
|
||||
|
||||
tool: Literal["list_observations", "get_observation", "recall", "expand", "done"] = Field(
|
||||
description="Tool to invoke: list_observations, get_observation, recall, expand, or done"
|
||||
)
|
||||
# Tool-specific parameters
|
||||
observation_id: str | None = Field(default=None, description="Observation ID for get_observation")
|
||||
query: str | None = Field(default=None, description="Search query for recall")
|
||||
max_tokens: int | None = Field(default=None, description="Max tokens for recall results (default 2048)")
|
||||
memory_ids: list[str] | None = Field(default=None, description="Memory unit IDs for expand (batched)")
|
||||
depth: Literal["chunk", "document"] | None = Field(default=None, description="Expansion depth for expand")
|
||||
observation_sections: list[ObservationSection] | None = Field(
|
||||
default=None, description="Observation sections for done action (when output_mode=observations)"
|
||||
)
|
||||
# Plain text answer fields (for output_mode=answer)
|
||||
answer: str | None = Field(default=None, description="Well-formatted markdown answer for done action")
|
||||
answer_memory_ids: list[str] | None = Field(
|
||||
default=None, description="Memory IDs supporting the answer", alias="memory_ids"
|
||||
)
|
||||
answer_model_ids: list[str] | None = Field(
|
||||
default=None, description="Mental model IDs supporting the answer", alias="model_ids"
|
||||
)
|
||||
reasoning: str | None = Field(default=None, description="Brief reasoning for this action")
|
||||
|
||||
|
||||
class ReflectActionBatch(BaseModel):
|
||||
"""Batch of actions for parallel execution."""
|
||||
|
||||
actions: list[ReflectAction] = Field(description="List of actions to execute in parallel")
|
||||
|
||||
|
||||
class ToolCall(BaseModel):
|
||||
"""A single tool call made during reflect."""
|
||||
|
||||
tool: str = Field(description="Tool name: lookup, recall, expand")
|
||||
reason: str | None = Field(default=None, description="Agent's reasoning for making this tool call")
|
||||
input: dict = Field(description="Tool input parameters")
|
||||
output: dict = Field(description="Tool output/result")
|
||||
duration_ms: int = Field(description="Execution time in milliseconds")
|
||||
iteration: int = Field(default=0, description="Iteration number (1-based) when this tool was called")
|
||||
|
||||
|
||||
class LLMCall(BaseModel):
|
||||
"""A single LLM call made during reflect."""
|
||||
|
||||
scope: str = Field(description="Call scope: agent_1, agent_2, final, etc.")
|
||||
duration_ms: int = Field(description="Execution time in milliseconds")
|
||||
input_tokens: int = Field(default=0, description="Input tokens used")
|
||||
output_tokens: int = Field(default=0, description="Output tokens used")
|
||||
|
||||
|
||||
class DirectiveInfo(BaseModel):
|
||||
"""Information about a directive that was applied during reflect."""
|
||||
|
||||
id: str = Field(description="Directive mental model ID")
|
||||
name: str = Field(description="Directive name")
|
||||
content: str = Field(description="Directive content")
|
||||
|
||||
|
||||
class TokenUsageSummary(BaseModel):
|
||||
"""Total token usage across all LLM calls."""
|
||||
|
||||
input_tokens: int = Field(default=0, description="Total input tokens used")
|
||||
output_tokens: int = Field(default=0, description="Total output tokens used")
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output)")
|
||||
|
||||
|
||||
class ReflectAgentResult(BaseModel):
|
||||
"""Result from the reflect agent."""
|
||||
|
||||
text: str = Field(description="Final answer text")
|
||||
structured_output: dict[str, Any] | None = Field(
|
||||
default=None, description="Structured output parsed according to provided response_schema"
|
||||
)
|
||||
iterations: int = Field(default=0, description="Number of iterations taken")
|
||||
tools_called: int = Field(default=0, description="Total number of tool calls made")
|
||||
tool_trace: list[ToolCall] = Field(default_factory=list, description="Trace of all tool calls made")
|
||||
llm_trace: list[LLMCall] = Field(default_factory=list, description="Trace of all LLM calls made")
|
||||
usage: TokenUsageSummary = Field(
|
||||
default_factory=TokenUsageSummary, description="Total token usage across all LLM calls"
|
||||
)
|
||||
used_memory_ids: list[str] = Field(default_factory=list, description="Validated memory IDs actually used in answer")
|
||||
used_mental_model_ids: list[str] = Field(
|
||||
default_factory=list, description="Validated mental model IDs actually used in answer"
|
||||
)
|
||||
used_observation_ids: list[str] = Field(
|
||||
default_factory=list, description="Validated observation IDs actually used in answer"
|
||||
)
|
||||
directives_applied: list[DirectiveInfo] = Field(
|
||||
default_factory=list, description="Directive mental models that affected this reflection"
|
||||
)
|
||||
@@ -0,0 +1,186 @@
|
||||
"""
|
||||
Models and utilities for evidence-grounded observations with computed trends.
|
||||
|
||||
Observations are part of mental models and represent patterns/beliefs derived
|
||||
from memories. Each observation must be grounded in specific evidence (quotes)
|
||||
from memories, and trends are computed algorithmically from evidence timestamps.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from enum import Enum
|
||||
|
||||
from pydantic import BaseModel, Field, computed_field, field_validator
|
||||
|
||||
|
||||
class Trend(str, Enum):
|
||||
"""Computed trend for an observation based on evidence timestamps.
|
||||
|
||||
Trends indicate how an observation's evidence is distributed over time:
|
||||
- STABLE: Evidence spread across time, continues to present
|
||||
- STRENGTHENING: More/denser evidence recently than before
|
||||
- WEAKENING: Evidence mostly old, sparse recently
|
||||
- NEW: All evidence within recent window
|
||||
- STALE: No evidence in recent window (may no longer apply)
|
||||
"""
|
||||
|
||||
STABLE = "stable"
|
||||
STRENGTHENING = "strengthening"
|
||||
WEAKENING = "weakening"
|
||||
NEW = "new"
|
||||
STALE = "stale"
|
||||
|
||||
|
||||
class ObservationEvidence(BaseModel):
|
||||
"""A single piece of evidence supporting an observation.
|
||||
|
||||
Each evidence item must include an exact quote from the source memory
|
||||
to ensure observations are grounded and verifiable.
|
||||
"""
|
||||
|
||||
memory_id: str = Field(description="ID of the memory unit this evidence comes from")
|
||||
quote: str = Field(description="Exact quote from the memory supporting the observation")
|
||||
relevance: str = Field(default="", description="Brief explanation of how this quote supports the observation")
|
||||
timestamp: datetime = Field(description="When the source memory was created")
|
||||
|
||||
@field_validator("timestamp", mode="before")
|
||||
@classmethod
|
||||
def ensure_timezone_aware(cls, v: datetime | str | None) -> datetime:
|
||||
"""Ensure timestamp is always timezone-aware UTC."""
|
||||
if v is None:
|
||||
return datetime.now(timezone.utc)
|
||||
if isinstance(v, str):
|
||||
# Parse ISO format string, handling 'Z' suffix
|
||||
v = datetime.fromisoformat(v.replace("Z", "+00:00"))
|
||||
if isinstance(v, datetime):
|
||||
if v.tzinfo is None:
|
||||
return v.replace(tzinfo=timezone.utc)
|
||||
return v
|
||||
raise ValueError(f"Invalid timestamp type: {type(v)}")
|
||||
|
||||
|
||||
class Observation(BaseModel):
|
||||
"""A single observation within a mental model.
|
||||
|
||||
Observations represent patterns, preferences, beliefs, or other insights
|
||||
derived from memories. Each observation must be grounded in evidence
|
||||
with exact quotes from source memories.
|
||||
"""
|
||||
|
||||
title: str = Field(description="Short summary title for the observation (5-10 words)")
|
||||
content: str = Field(description="The observation content - detailed explanation of what we believe to be true")
|
||||
evidence: list[ObservationEvidence] = Field(default_factory=list, description="Supporting evidence with quotes")
|
||||
created_at: datetime = Field(
|
||||
default_factory=lambda: datetime.now(timezone.utc), description="When this observation was first created"
|
||||
)
|
||||
|
||||
@field_validator("created_at", mode="before")
|
||||
@classmethod
|
||||
def ensure_created_at_timezone_aware(cls, v: datetime | str | None) -> datetime:
|
||||
"""Ensure created_at is always timezone-aware UTC."""
|
||||
if v is None:
|
||||
return datetime.now(timezone.utc)
|
||||
if isinstance(v, str):
|
||||
v = datetime.fromisoformat(v.replace("Z", "+00:00"))
|
||||
if isinstance(v, datetime):
|
||||
if v.tzinfo is None:
|
||||
return v.replace(tzinfo=timezone.utc)
|
||||
return v
|
||||
raise ValueError(f"Invalid created_at type: {type(v)}")
|
||||
|
||||
@computed_field
|
||||
@property
|
||||
def trend(self) -> Trend:
|
||||
"""Compute trend from evidence timestamps."""
|
||||
return compute_trend(self.evidence)
|
||||
|
||||
@computed_field
|
||||
@property
|
||||
def evidence_span(self) -> dict[str, str | None]:
|
||||
"""Get the time span covered by evidence."""
|
||||
if not self.evidence:
|
||||
return {"from": None, "to": None}
|
||||
timestamps = [e.timestamp for e in self.evidence]
|
||||
return {
|
||||
"from": min(timestamps).isoformat(),
|
||||
"to": max(timestamps).isoformat(),
|
||||
}
|
||||
|
||||
@computed_field
|
||||
@property
|
||||
def evidence_count(self) -> int:
|
||||
"""Number of evidence items supporting this observation."""
|
||||
return len(self.evidence)
|
||||
|
||||
|
||||
def compute_trend(
|
||||
evidence: list[ObservationEvidence],
|
||||
now: datetime | None = None,
|
||||
recent_days: int = 30,
|
||||
old_days: int = 90,
|
||||
) -> Trend:
|
||||
"""Compute the trend for an observation based on evidence timestamps.
|
||||
|
||||
The trend indicates how the evidence is distributed over time:
|
||||
- STABLE: Evidence spread across time, continues to present
|
||||
- STRENGTHENING: More evidence recently than historically
|
||||
- WEAKENING: Evidence mostly old, sparse recently
|
||||
- NEW: All evidence is recent (within recent_days)
|
||||
- STALE: No evidence in recent window
|
||||
|
||||
Args:
|
||||
evidence: List of evidence items with timestamps
|
||||
now: Reference time for calculations (defaults to current UTC time)
|
||||
recent_days: Number of days to consider "recent" (default 30)
|
||||
old_days: Number of days to consider "old" (default 90)
|
||||
|
||||
Returns:
|
||||
Computed Trend enum value
|
||||
"""
|
||||
if now is None:
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
# Ensure now is timezone-aware
|
||||
if now.tzinfo is None:
|
||||
now = now.replace(tzinfo=timezone.utc)
|
||||
|
||||
if not evidence:
|
||||
return Trend.STALE
|
||||
|
||||
recent_cutoff = now - timedelta(days=recent_days)
|
||||
old_cutoff = now - timedelta(days=old_days)
|
||||
|
||||
# Normalize timestamps to UTC for comparison
|
||||
def normalize_ts(ts: datetime) -> datetime:
|
||||
if ts.tzinfo is None:
|
||||
return ts.replace(tzinfo=timezone.utc)
|
||||
return ts
|
||||
|
||||
recent = [e for e in evidence if normalize_ts(e.timestamp) > recent_cutoff]
|
||||
old = [e for e in evidence if normalize_ts(e.timestamp) < old_cutoff]
|
||||
middle = [e for e in evidence if old_cutoff <= normalize_ts(e.timestamp) <= recent_cutoff]
|
||||
|
||||
# No recent evidence = stale
|
||||
if not recent:
|
||||
return Trend.STALE
|
||||
|
||||
# All evidence is recent = new
|
||||
if not old and not middle:
|
||||
return Trend.NEW
|
||||
|
||||
# Compare density (evidence per day)
|
||||
recent_density = len(recent) / recent_days if recent_days > 0 else 0
|
||||
older_period = old_days - recent_days
|
||||
older_density = (len(old) + len(middle)) / older_period if older_period > 0 else 0
|
||||
|
||||
# Avoid division by zero
|
||||
if older_density == 0:
|
||||
return Trend.NEW
|
||||
|
||||
ratio = recent_density / older_density
|
||||
|
||||
if ratio > 1.5:
|
||||
return Trend.STRENGTHENING
|
||||
elif ratio < 0.5:
|
||||
return Trend.WEAKENING
|
||||
else:
|
||||
return Trend.STABLE
|
||||
@@ -0,0 +1,513 @@
|
||||
"""
|
||||
System prompts for the reflect agent.
|
||||
|
||||
The reflect agent uses hierarchical retrieval:
|
||||
1. search_mental_models - User-curated summaries (highest quality)
|
||||
2. search_observations - Consolidated knowledge with freshness awareness
|
||||
3. recall - Raw facts as ground truth fallback
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
|
||||
def _extract_directive_rules(directives: list[dict[str, Any]]) -> list[str]:
|
||||
"""
|
||||
Extract directive rules as a list of strings.
|
||||
|
||||
Args:
|
||||
directives: List of directives with name and content
|
||||
|
||||
Returns:
|
||||
List of directive rule strings
|
||||
"""
|
||||
rules = []
|
||||
for directive in directives:
|
||||
directive_name = directive.get("name", "")
|
||||
# New format: directives have direct content field
|
||||
content = directive.get("content", "")
|
||||
if content:
|
||||
if directive_name:
|
||||
rules.append(f"**{directive_name}**: {content}")
|
||||
else:
|
||||
rules.append(content)
|
||||
else:
|
||||
# Legacy format: check for observations
|
||||
observations = directive.get("observations", [])
|
||||
if observations:
|
||||
for obs in observations:
|
||||
# Support both Pydantic Observation objects and dicts
|
||||
if hasattr(obs, "title"):
|
||||
title = obs.title
|
||||
obs_content = obs.content
|
||||
else:
|
||||
title = obs.get("title", "")
|
||||
obs_content = obs.get("content", "")
|
||||
if title and obs_content:
|
||||
rules.append(f"**{title}**: {obs_content}")
|
||||
elif obs_content:
|
||||
rules.append(obs_content)
|
||||
elif directive_name:
|
||||
# Fallback to description
|
||||
desc = directive.get("description", "")
|
||||
if desc:
|
||||
rules.append(f"**{directive_name}**: {desc}")
|
||||
return rules
|
||||
|
||||
|
||||
def build_directives_section(directives: list[dict[str, Any]]) -> str:
|
||||
"""
|
||||
Build the directives section for the system prompt.
|
||||
|
||||
Directives are hard rules that MUST be followed in all responses.
|
||||
|
||||
Args:
|
||||
directives: List of directive mental models with observations
|
||||
"""
|
||||
if not directives:
|
||||
return ""
|
||||
|
||||
rules = _extract_directive_rules(directives)
|
||||
if not rules:
|
||||
return ""
|
||||
|
||||
parts = [
|
||||
"## DIRECTIVES (MANDATORY)",
|
||||
"These are hard rules you MUST follow in ALL responses:",
|
||||
"",
|
||||
]
|
||||
|
||||
for rule in rules:
|
||||
parts.append(f"- {rule}")
|
||||
|
||||
parts.extend(
|
||||
[
|
||||
"",
|
||||
"NEVER violate these directives, even if other context suggests otherwise.",
|
||||
"IMPORTANT: Do NOT explain or justify how you handled directives in your answer. Just follow them silently.",
|
||||
"",
|
||||
]
|
||||
)
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def build_directives_reminder(directives: list[dict[str, Any]]) -> str:
|
||||
"""
|
||||
Build a reminder section for directives to place at the end of the prompt.
|
||||
|
||||
Args:
|
||||
directives: List of directive mental models with observations
|
||||
"""
|
||||
if not directives:
|
||||
return ""
|
||||
|
||||
rules = _extract_directive_rules(directives)
|
||||
if not rules:
|
||||
return ""
|
||||
|
||||
parts = [
|
||||
"",
|
||||
"## REMINDER: MANDATORY DIRECTIVES",
|
||||
"Before responding, ensure your answer complies with ALL of these directives:",
|
||||
"",
|
||||
]
|
||||
|
||||
for i, rule in enumerate(rules, 1):
|
||||
parts.append(f"{i}. {rule}")
|
||||
|
||||
parts.append("")
|
||||
parts.append("Your response will be REJECTED if it violates any directive above.")
|
||||
parts.append("Do NOT include any commentary about how you handled directives - just follow them.")
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def build_system_prompt_for_tools(
|
||||
bank_profile: dict[str, Any],
|
||||
context: str | None = None,
|
||||
directives: list[dict[str, Any]] | None = None,
|
||||
has_mental_models: bool = False,
|
||||
budget: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Build the system prompt for tool-calling reflect agent.
|
||||
|
||||
The agent uses hierarchical retrieval:
|
||||
1. search_mental_models - User-curated summaries (try first, if available)
|
||||
2. search_observations - Consolidated knowledge with freshness
|
||||
3. recall - Raw facts as ground truth
|
||||
|
||||
Args:
|
||||
bank_profile: Bank profile with name and mission
|
||||
context: Optional additional context
|
||||
directives: Optional list of directive mental models to inject as hard rules
|
||||
has_mental_models: Whether the bank has any mental models (skip if not)
|
||||
budget: Search depth budget - "low", "mid", or "high". Controls exploration thoroughness.
|
||||
"""
|
||||
name = bank_profile.get("name", "Assistant")
|
||||
mission = bank_profile.get("mission", "")
|
||||
|
||||
parts = []
|
||||
|
||||
# Anti-hallucination rule at the very top
|
||||
parts.extend(
|
||||
[
|
||||
"CRITICAL: You MUST ONLY use information from retrieved tool results. NEVER make up names, people, events, or entities.",
|
||||
"",
|
||||
]
|
||||
)
|
||||
|
||||
# Inject directives after anti-hallucination rule
|
||||
if directives:
|
||||
parts.append(build_directives_section(directives))
|
||||
|
||||
parts.extend(
|
||||
[
|
||||
"You are a reflection agent that answers questions by reasoning over retrieved memories.",
|
||||
"",
|
||||
]
|
||||
)
|
||||
|
||||
parts.extend(
|
||||
[
|
||||
"## CRITICAL RULES",
|
||||
"- ONLY use information from tool results - no external knowledge or guessing",
|
||||
"- You SHOULD synthesize, infer, and reason from the retrieved memories",
|
||||
"- You MUST search before saying you don't have information",
|
||||
"",
|
||||
"## How to Reason",
|
||||
"- If memories mention someone did an activity, you can infer they likely enjoyed it",
|
||||
"- Synthesize a coherent narrative from related memories",
|
||||
"- Be a thoughtful interpreter, not just a literal repeater",
|
||||
"- When the exact answer isn't stated, use what IS stated to give the best answer",
|
||||
"",
|
||||
"## HIERARCHICAL RETRIEVAL STRATEGY",
|
||||
"",
|
||||
]
|
||||
)
|
||||
|
||||
# Build retrieval levels based on what's available
|
||||
if has_mental_models:
|
||||
parts.extend(
|
||||
[
|
||||
"You have access to THREE levels of knowledge. Use them in this order:",
|
||||
"",
|
||||
"### 1. MENTAL MODELS (search_mental_models) - Try First",
|
||||
"- User-curated summaries about specific topics",
|
||||
"- HIGHEST quality - manually created and maintained",
|
||||
"- If a relevant mental model exists and is FRESH, it may fully answer the question",
|
||||
"- Check `is_stale` field - if stale, also verify with lower levels",
|
||||
"",
|
||||
"### 2. OBSERVATIONS (search_observations) - Second Priority",
|
||||
"- Auto-consolidated knowledge from memories",
|
||||
"- Check `is_stale` field - if stale, ALSO use recall() to verify",
|
||||
"- Good for understanding patterns and summaries",
|
||||
"",
|
||||
"### 3. RAW FACTS (recall) - Ground Truth",
|
||||
"- Individual memories (world facts and experiences)",
|
||||
"- Use when: no mental models/observations exist, they're stale, or you need specific details",
|
||||
"- This is the source of truth that other levels are built from",
|
||||
"",
|
||||
]
|
||||
)
|
||||
else:
|
||||
parts.extend(
|
||||
[
|
||||
"You have access to TWO levels of knowledge. Use them in this order:",
|
||||
"",
|
||||
"### 1. OBSERVATIONS (search_observations) - Try First",
|
||||
"- Auto-consolidated knowledge from memories",
|
||||
"- Check `is_stale` field - if stale, ALSO use recall() to verify",
|
||||
"- Good for understanding patterns and summaries",
|
||||
"",
|
||||
"### 2. RAW FACTS (recall) - Ground Truth",
|
||||
"- Individual memories (world facts and experiences)",
|
||||
"- Use when: no observations exist, they're stale, or you need specific details",
|
||||
"- This is the source of truth that observations are built from",
|
||||
"",
|
||||
]
|
||||
)
|
||||
|
||||
parts.extend(
|
||||
[
|
||||
"## Query Strategy",
|
||||
"recall() uses semantic search. NEVER just echo the user's question - decompose it into targeted searches:",
|
||||
"",
|
||||
"BAD: User asks 'recurring lesson themes between students' → recall('recurring lesson themes between students')",
|
||||
"GOOD: Break it down into component searches:",
|
||||
" 1. recall('lessons') - find all lesson-related memories",
|
||||
" 2. recall('teaching sessions') - alternative phrasing",
|
||||
" 3. recall('student progress') - find student-related memories",
|
||||
"",
|
||||
"Think: What ENTITIES and CONCEPTS does this question involve? Search for each separately.",
|
||||
"",
|
||||
]
|
||||
)
|
||||
|
||||
# Add budget guidance
|
||||
if budget:
|
||||
budget_lower = budget.lower()
|
||||
if budget_lower == "low":
|
||||
parts.extend(
|
||||
[
|
||||
"## RESEARCH DEPTH: SHALLOW (Quick Response)",
|
||||
"- Prioritize speed over completeness",
|
||||
"- If mental models or observations provide a reasonable answer, stop there",
|
||||
"- Only dig deeper if the initial results are clearly insufficient",
|
||||
"- Prefer a quick overview rather than exhaustive details",
|
||||
"- Answer promptly with available information",
|
||||
"",
|
||||
]
|
||||
)
|
||||
elif budget_lower == "mid":
|
||||
parts.extend(
|
||||
[
|
||||
"## RESEARCH DEPTH: MODERATE (Balanced)",
|
||||
"- Balance thoroughness with efficiency",
|
||||
"- Check multiple sources when the question warrants it",
|
||||
"- Verify stale data if it's central to the answer",
|
||||
"- Don't over-explore, but ensure reasonable coverage",
|
||||
"",
|
||||
]
|
||||
)
|
||||
elif budget_lower == "high":
|
||||
parts.extend(
|
||||
[
|
||||
"## RESEARCH DEPTH: DEEP (Thorough Exploration)",
|
||||
"- Explore comprehensively before answering",
|
||||
"- Search across all available knowledge levels",
|
||||
"- Use multiple query variations to ensure coverage",
|
||||
"- Verify information across different retrieval levels",
|
||||
"- Use expand() to get full context on important memories",
|
||||
"- Take time to synthesize a complete, well-researched answer",
|
||||
"",
|
||||
]
|
||||
)
|
||||
|
||||
parts.append("## Workflow")
|
||||
|
||||
if has_mental_models:
|
||||
parts.extend(
|
||||
[
|
||||
"1. First, try search_mental_models() - check if a curated summary exists",
|
||||
"2. If no mental model or it's stale, try search_observations() for consolidated knowledge",
|
||||
"3. If observations are stale OR you need specific details, use recall() for raw facts",
|
||||
"4. Use expand() if you need more context on specific memories",
|
||||
"5. When ready, call done() with your answer and supporting IDs",
|
||||
]
|
||||
)
|
||||
else:
|
||||
parts.extend(
|
||||
[
|
||||
"1. First, try search_observations() - check for consolidated knowledge",
|
||||
"2. If observations are stale OR you need specific details, use recall() for raw facts",
|
||||
"3. Use expand() if you need more context on specific memories",
|
||||
"4. When ready, call done() with your answer and supporting IDs",
|
||||
]
|
||||
)
|
||||
|
||||
parts.extend(
|
||||
[
|
||||
"",
|
||||
"## Output Format: Well-Formatted Markdown Answer",
|
||||
"Call done() with a well-formatted markdown 'answer' field.",
|
||||
"- USE markdown formatting for structure (headers, lists, bold, italic, code blocks, tables, etc.)",
|
||||
"- CRITICAL: Add blank lines before and after block elements (tables, code blocks, lists)",
|
||||
"- Format for clarity and readability with proper spacing and hierarchy",
|
||||
"- NEVER include memory IDs, UUIDs, or 'Memory references' in the answer text",
|
||||
"- Put IDs ONLY in the memory_ids/mental_model_ids/observation_ids arrays, not in the answer",
|
||||
]
|
||||
)
|
||||
|
||||
parts.append("")
|
||||
parts.append(f"## Memory Bank: {name}")
|
||||
|
||||
if mission:
|
||||
parts.append(f"Mission: {mission}")
|
||||
|
||||
# Disposition traits
|
||||
disposition = bank_profile.get("disposition", {})
|
||||
if disposition:
|
||||
traits = []
|
||||
if "skepticism" in disposition:
|
||||
traits.append(f"skepticism={disposition['skepticism']}")
|
||||
if "literalism" in disposition:
|
||||
traits.append(f"literalism={disposition['literalism']}")
|
||||
if "empathy" in disposition:
|
||||
traits.append(f"empathy={disposition['empathy']}")
|
||||
if traits:
|
||||
parts.append(f"Disposition: {', '.join(traits)}")
|
||||
|
||||
if context:
|
||||
parts.append(f"\n## Additional Context\n{context}")
|
||||
|
||||
# Add directive reminder at the END for recency effect
|
||||
if directives:
|
||||
parts.append(build_directives_reminder(directives))
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def build_agent_prompt(
|
||||
query: str,
|
||||
context_history: list[dict],
|
||||
bank_profile: dict,
|
||||
additional_context: str | None = None,
|
||||
) -> str:
|
||||
"""Build the user prompt for the reflect agent."""
|
||||
parts = []
|
||||
|
||||
# Bank identity
|
||||
name = bank_profile.get("name", "Assistant")
|
||||
mission = bank_profile.get("mission", "")
|
||||
|
||||
parts.append(f"## Memory Bank Context\nName: {name}")
|
||||
if mission:
|
||||
parts.append(f"Mission: {mission}")
|
||||
|
||||
# Disposition traits if present
|
||||
disposition = bank_profile.get("disposition", {})
|
||||
if disposition:
|
||||
traits = []
|
||||
if "skepticism" in disposition:
|
||||
traits.append(f"skepticism={disposition['skepticism']}")
|
||||
if "literalism" in disposition:
|
||||
traits.append(f"literalism={disposition['literalism']}")
|
||||
if "empathy" in disposition:
|
||||
traits.append(f"empathy={disposition['empathy']}")
|
||||
if traits:
|
||||
parts.append(f"Disposition: {', '.join(traits)}")
|
||||
|
||||
# Additional context from caller
|
||||
if additional_context:
|
||||
parts.append(f"\n## Additional Context\n{additional_context}")
|
||||
|
||||
# Tool call history
|
||||
if context_history:
|
||||
parts.append("\n## Tool Results (synthesize and reason from this data)")
|
||||
for i, entry in enumerate(context_history, 1):
|
||||
tool = entry["tool"]
|
||||
output = entry["output"]
|
||||
# Format as proper JSON for LLM readability
|
||||
try:
|
||||
output_str = json.dumps(output, indent=2, default=str)
|
||||
except (TypeError, ValueError):
|
||||
output_str = str(output)
|
||||
parts.append(f"\n### Call {i}: {tool}\n```json\n{output_str}\n```")
|
||||
|
||||
# The question
|
||||
parts.append(f"\n## Question\n{query}")
|
||||
|
||||
# Instructions
|
||||
if context_history:
|
||||
parts.append(
|
||||
"\n## Instructions\n"
|
||||
"Based on the tool results above, either call more tools or provide your final answer. "
|
||||
"Synthesize and reason from the data - make reasonable inferences when helpful. "
|
||||
"If you have related information, use it to give the best possible answer."
|
||||
)
|
||||
else:
|
||||
parts.append(
|
||||
"\n## Instructions\n"
|
||||
"Start by searching for relevant information using the hierarchical retrieval strategy:\n"
|
||||
"1. Try search_mental_models() first for curated summaries\n"
|
||||
"2. Try search_observations() for consolidated knowledge\n"
|
||||
"3. Use recall() for specific details or to verify stale data"
|
||||
)
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def build_final_prompt(
|
||||
query: str,
|
||||
context_history: list[dict],
|
||||
bank_profile: dict,
|
||||
additional_context: str | None = None,
|
||||
) -> str:
|
||||
"""Build the final prompt when forcing a text response (no tools)."""
|
||||
parts = []
|
||||
|
||||
# Bank identity
|
||||
name = bank_profile.get("name", "Assistant")
|
||||
mission = bank_profile.get("mission", "")
|
||||
|
||||
parts.append(f"## Memory Bank Context\nName: {name}")
|
||||
if mission:
|
||||
parts.append(f"Mission: {mission}")
|
||||
|
||||
# Disposition traits if present
|
||||
disposition = bank_profile.get("disposition", {})
|
||||
if disposition:
|
||||
traits = []
|
||||
if "skepticism" in disposition:
|
||||
traits.append(f"skepticism={disposition['skepticism']}")
|
||||
if "literalism" in disposition:
|
||||
traits.append(f"literalism={disposition['literalism']}")
|
||||
if "empathy" in disposition:
|
||||
traits.append(f"empathy={disposition['empathy']}")
|
||||
if traits:
|
||||
parts.append(f"Disposition: {', '.join(traits)}")
|
||||
|
||||
# Additional context from caller
|
||||
if additional_context:
|
||||
parts.append(f"\n## Additional Context\n{additional_context}")
|
||||
|
||||
# Tool call history
|
||||
if context_history:
|
||||
parts.append("\n## Retrieved Data (synthesize and reason from this data)")
|
||||
for entry in context_history:
|
||||
tool = entry["tool"]
|
||||
output = entry["output"]
|
||||
# Format as proper JSON for LLM readability
|
||||
try:
|
||||
output_str = json.dumps(output, indent=2, default=str)
|
||||
except (TypeError, ValueError):
|
||||
output_str = str(output)
|
||||
parts.append(f"\n### From {tool}:\n```json\n{output_str}\n```")
|
||||
else:
|
||||
parts.append("\n## Retrieved Data\nNo data was retrieved.")
|
||||
|
||||
# The question
|
||||
parts.append(f"\n## Question\n{query}")
|
||||
|
||||
# Final instructions
|
||||
parts.append(
|
||||
"\n## Instructions\n"
|
||||
"Provide a thoughtful answer by synthesizing and reasoning from the retrieved data above. "
|
||||
"You can make reasonable inferences from the memories, but don't completely fabricate information. "
|
||||
"If the exact answer isn't stated, use what IS stated to give the best possible answer. "
|
||||
"Only say 'I don't have information' if the retrieved data is truly unrelated to the question.\n\n"
|
||||
"IMPORTANT: Output ONLY the final answer. Do NOT include meta-commentary like "
|
||||
'"I\'ll search..." or "Let me analyze...". Do NOT explain your reasoning process. '
|
||||
"Just provide the direct synthesized answer."
|
||||
)
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
FINAL_SYSTEM_PROMPT = """CRITICAL: You MUST ONLY use information from retrieved tool results. NEVER make up names, people, events, or entities.
|
||||
|
||||
You are a thoughtful assistant that synthesizes answers from retrieved memories.
|
||||
|
||||
Your approach:
|
||||
- Reason over the retrieved memories to answer the question
|
||||
- Make reasonable inferences when the exact answer isn't explicitly stated
|
||||
- Connect related memories to form a complete picture
|
||||
- Be helpful - if you have related information, use it to give the best possible answer
|
||||
- ONLY use information from tool results - no external knowledge or guessing
|
||||
|
||||
Only say "I don't have information" if the retrieved data is truly unrelated to the question.
|
||||
|
||||
FORMATTING: Use proper markdown formatting in your answer:
|
||||
- Headers (##, ###) for sections
|
||||
- Lists (bullet or numbered) for enumerations
|
||||
- Bold/italic for emphasis
|
||||
- Tables with proper syntax (ensure blank line before and after)
|
||||
- Code blocks where appropriate
|
||||
- CRITICAL: Always add blank lines before and after block elements (tables, code blocks, lists)
|
||||
- Proper spacing between sections
|
||||
|
||||
CRITICAL: Output ONLY the final synthesized answer. Do NOT include:
|
||||
- Meta-commentary about what you're doing ("I'll search...", "Let me analyze...")
|
||||
- Explanations of your reasoning process
|
||||
- Descriptions of your approach
|
||||
Just provide the direct answer with proper markdown formatting."""
|
||||
@@ -0,0 +1,436 @@
|
||||
"""
|
||||
Tool implementations for the reflect agent.
|
||||
|
||||
Implements hierarchical retrieval:
|
||||
1. search_mental_models - User-curated stored reflect responses (highest quality)
|
||||
2. search_observations - Consolidated knowledge with freshness
|
||||
3. recall - Raw facts as ground truth
|
||||
"""
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from asyncpg import Connection
|
||||
|
||||
from ...api.http import RequestContext
|
||||
from ..memory_engine import MemoryEngine
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Observation is considered stale if not updated in this many days
|
||||
STALE_THRESHOLD_DAYS = 7
|
||||
|
||||
|
||||
async def tool_search_mental_models(
|
||||
conn: "Connection",
|
||||
bank_id: str,
|
||||
query: str,
|
||||
query_embedding: list[float],
|
||||
max_results: int = 5,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str = "any",
|
||||
exclude_ids: list[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Search user-curated mental models by semantic similarity.
|
||||
|
||||
Mental models are high-quality, manually created summaries about specific topics.
|
||||
They should be searched FIRST as they represent the most reliable synthesized knowledge.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
bank_id: Bank identifier
|
||||
query: Search query (for logging/tracing)
|
||||
query_embedding: Pre-computed embedding for semantic search
|
||||
max_results: Maximum number of mental models to return
|
||||
tags: Optional tags to filter mental models
|
||||
tags_match: How to match tags - "any" (OR), "all" (AND)
|
||||
exclude_ids: Optional list of mental model IDs to exclude (e.g., when refreshing a mental model)
|
||||
|
||||
Returns:
|
||||
Dict with matching mental models including content and freshness info
|
||||
"""
|
||||
from ..memory_engine import fq_table
|
||||
from ..search.tags import build_tags_where_clause
|
||||
|
||||
# Build filters dynamically
|
||||
filters = ""
|
||||
params: list[Any] = [bank_id, str(query_embedding), max_results]
|
||||
next_param = 4
|
||||
|
||||
# Use the centralized tag filtering logic
|
||||
if tags:
|
||||
tag_clause, tag_params, next_param = build_tags_where_clause(tags, param_offset=next_param, match=tags_match)
|
||||
filters += f" {tag_clause}"
|
||||
params.extend(tag_params)
|
||||
|
||||
if exclude_ids:
|
||||
filters += f" AND id != ALL(${next_param}::text[])"
|
||||
params.append(exclude_ids)
|
||||
next_param += 1
|
||||
|
||||
# Search mental models by embedding similarity
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT
|
||||
id, name, content,
|
||||
tags, created_at, last_refreshed_at,
|
||||
1 - (embedding <=> $2::vector) as relevance
|
||||
FROM {fq_table("mental_models")}
|
||||
WHERE bank_id = $1 AND embedding IS NOT NULL {filters}
|
||||
ORDER BY embedding <=> $2::vector
|
||||
LIMIT $3
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
mental_models = []
|
||||
|
||||
for row in rows:
|
||||
last_refreshed_at = row["last_refreshed_at"]
|
||||
if last_refreshed_at and last_refreshed_at.tzinfo is None:
|
||||
last_refreshed_at = last_refreshed_at.replace(tzinfo=timezone.utc)
|
||||
|
||||
# Calculate freshness
|
||||
is_stale = False
|
||||
if last_refreshed_at:
|
||||
age = now - last_refreshed_at
|
||||
is_stale = age > timedelta(days=STALE_THRESHOLD_DAYS)
|
||||
|
||||
mental_models.append(
|
||||
{
|
||||
"id": str(row["id"]),
|
||||
"name": row["name"],
|
||||
"content": row["content"],
|
||||
"tags": row["tags"] or [],
|
||||
"relevance": round(row["relevance"], 4),
|
||||
"updated_at": last_refreshed_at.isoformat() if last_refreshed_at else None,
|
||||
"is_stale": is_stale,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"query": query,
|
||||
"count": len(mental_models),
|
||||
"mental_models": mental_models,
|
||||
}
|
||||
|
||||
|
||||
async def tool_search_observations(
|
||||
memory_engine: "MemoryEngine",
|
||||
bank_id: str,
|
||||
query: str,
|
||||
request_context: "RequestContext",
|
||||
max_tokens: int = 5000,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str = "any",
|
||||
last_consolidated_at: datetime | None = None,
|
||||
pending_consolidation: int = 0,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Search consolidated observations using recall with include_observations.
|
||||
|
||||
Observations are auto-generated from memories. Returns freshness info
|
||||
so the agent knows if it should also verify with recall().
|
||||
|
||||
Args:
|
||||
memory_engine: Memory engine instance
|
||||
bank_id: Bank identifier
|
||||
query: Search query
|
||||
request_context: Request context for authentication
|
||||
max_tokens: Maximum tokens for results (default 5000)
|
||||
tags: Optional tags to filter observations
|
||||
tags_match: How to match tags - "any" (OR), "all" (AND)
|
||||
last_consolidated_at: When consolidation last ran (for staleness check)
|
||||
pending_consolidation: Number of memories waiting to be consolidated
|
||||
|
||||
Returns:
|
||||
Dict with matching observations including freshness info
|
||||
"""
|
||||
from ..memory_engine import fq_table
|
||||
|
||||
# Use recall to search observations (they come back in results field when fact_type=["observation"])
|
||||
result = await memory_engine.recall_async(
|
||||
bank_id=bank_id,
|
||||
query=query,
|
||||
fact_type=["observation"], # Only retrieve observations
|
||||
max_tokens=max_tokens, # Token budget controls how many observations are returned
|
||||
enable_trace=False,
|
||||
request_context=request_context,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
_connection_budget=1,
|
||||
_quiet=True,
|
||||
)
|
||||
|
||||
observations = []
|
||||
|
||||
# When fact_type=["observation"], results come back in `results` field as MemoryFact objects
|
||||
# We need to fetch additional fields (proof_count, source_memory_ids) from the database
|
||||
if result.results:
|
||||
obs_ids = [m.id for m in result.results]
|
||||
|
||||
# Fetch proof_count and source_memory_ids for these observations
|
||||
pool = await memory_engine._get_pool()
|
||||
async with pool.acquire() as conn:
|
||||
obs_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, proof_count, source_memory_ids
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = ANY($1::uuid[])
|
||||
""",
|
||||
obs_ids,
|
||||
)
|
||||
obs_data = {str(row["id"]): row for row in obs_rows}
|
||||
|
||||
for m in result.results:
|
||||
# Get additional data from DB lookup
|
||||
extra = obs_data.get(m.id, {})
|
||||
proof_count = extra.get("proof_count", 1) if extra else 1
|
||||
source_ids = extra.get("source_memory_ids", []) if extra else []
|
||||
# Convert UUIDs to strings
|
||||
source_memory_ids = [str(sid) for sid in (source_ids or [])]
|
||||
|
||||
# Determine staleness
|
||||
is_stale = False
|
||||
staleness_reason = None
|
||||
if pending_consolidation > 0:
|
||||
is_stale = True
|
||||
staleness_reason = f"{pending_consolidation} memories pending consolidation"
|
||||
|
||||
observations.append(
|
||||
{
|
||||
"id": str(m.id),
|
||||
"text": m.text,
|
||||
"proof_count": proof_count,
|
||||
"source_memory_ids": source_memory_ids,
|
||||
"tags": m.tags or [],
|
||||
"is_stale": is_stale,
|
||||
"staleness_reason": staleness_reason,
|
||||
}
|
||||
)
|
||||
|
||||
# Return freshness info (more understandable than raw pending_consolidation count)
|
||||
if pending_consolidation == 0:
|
||||
freshness = "up_to_date"
|
||||
elif pending_consolidation < 10:
|
||||
freshness = "slightly_stale"
|
||||
else:
|
||||
freshness = "stale"
|
||||
|
||||
return {
|
||||
"query": query,
|
||||
"count": len(observations),
|
||||
"observations": observations,
|
||||
"freshness": freshness,
|
||||
}
|
||||
|
||||
|
||||
async def tool_recall(
|
||||
memory_engine: "MemoryEngine",
|
||||
bank_id: str,
|
||||
query: str,
|
||||
request_context: "RequestContext",
|
||||
max_tokens: int = 2048,
|
||||
max_results: int = 50,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str = "any",
|
||||
connection_budget: int = 1,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Search memories using TEMPR retrieval.
|
||||
|
||||
This is the ground truth - raw facts and experiences.
|
||||
Use when mental models/observations don't exist, are stale, or need verification.
|
||||
|
||||
Args:
|
||||
memory_engine: Memory engine instance
|
||||
bank_id: Bank identifier
|
||||
query: Search query
|
||||
request_context: Request context for authentication
|
||||
max_tokens: Maximum tokens for results (default 2048)
|
||||
max_results: Maximum number of results
|
||||
tags: Filter by tags (includes untagged memories)
|
||||
tags_match: How to match tags - "any" (OR), "all" (AND), or "exact"
|
||||
connection_budget: Max DB connections for this recall (default 1 for internal ops)
|
||||
|
||||
Returns:
|
||||
Dict with list of matching memories
|
||||
"""
|
||||
result = await memory_engine.recall_async(
|
||||
bank_id=bank_id,
|
||||
query=query,
|
||||
fact_type=["experience", "world"], # Exclude opinions and observations
|
||||
max_tokens=max_tokens,
|
||||
enable_trace=False,
|
||||
request_context=request_context,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
_connection_budget=connection_budget,
|
||||
_quiet=True, # Suppress logging for internal operations
|
||||
)
|
||||
|
||||
memories = []
|
||||
for m in result.results[:max_results]:
|
||||
memories.append(
|
||||
{
|
||||
"id": str(m.id),
|
||||
"text": m.text,
|
||||
"type": m.fact_type,
|
||||
"entities": m.entities or [],
|
||||
"occurred": m.occurred_start, # Already ISO format string
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"query": query,
|
||||
"count": len(memories),
|
||||
"memories": memories,
|
||||
}
|
||||
|
||||
|
||||
async def tool_expand(
|
||||
conn: "Connection",
|
||||
bank_id: str,
|
||||
memory_ids: list[str],
|
||||
depth: str,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Expand multiple memories to get chunk or document context.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
bank_id: Bank identifier
|
||||
memory_ids: List of memory unit IDs
|
||||
depth: "chunk" or "document"
|
||||
|
||||
Returns:
|
||||
Dict with results array, each containing memory, chunk, and optionally document data
|
||||
"""
|
||||
from ..memory_engine import fq_table
|
||||
|
||||
if not memory_ids:
|
||||
return {"error": "memory_ids is required and must not be empty"}
|
||||
|
||||
# Validate and convert UUIDs
|
||||
valid_uuids: list[uuid.UUID] = []
|
||||
errors: dict[str, str] = {}
|
||||
for mid in memory_ids:
|
||||
try:
|
||||
valid_uuids.append(uuid.UUID(mid))
|
||||
except ValueError:
|
||||
errors[mid] = f"Invalid memory_id format: {mid}"
|
||||
|
||||
if not valid_uuids:
|
||||
return {"error": "No valid memory IDs provided", "details": errors}
|
||||
|
||||
# Batch fetch all memory units
|
||||
memories = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, chunk_id, document_id, fact_type, context
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = ANY($1) AND bank_id = $2
|
||||
""",
|
||||
valid_uuids,
|
||||
bank_id,
|
||||
)
|
||||
memory_map = {row["id"]: row for row in memories}
|
||||
|
||||
# Collect chunk_ids and document_ids for batch fetching
|
||||
chunk_ids = [m["chunk_id"] for m in memories if m["chunk_id"]]
|
||||
doc_ids_from_chunks: set[str] = set()
|
||||
doc_ids_direct: set[str] = set()
|
||||
|
||||
# Batch fetch all chunks
|
||||
chunk_map: dict[str, Any] = {}
|
||||
if chunk_ids:
|
||||
chunks = await conn.fetch(
|
||||
f"""
|
||||
SELECT chunk_id, chunk_text, chunk_index, document_id
|
||||
FROM {fq_table("chunks")}
|
||||
WHERE chunk_id = ANY($1)
|
||||
""",
|
||||
chunk_ids,
|
||||
)
|
||||
chunk_map = {row["chunk_id"]: row for row in chunks}
|
||||
if depth == "document":
|
||||
doc_ids_from_chunks = {c["document_id"] for c in chunks if c["document_id"]}
|
||||
|
||||
# Collect direct document IDs (memories without chunks)
|
||||
if depth == "document":
|
||||
for m in memories:
|
||||
if not m["chunk_id"] and m["document_id"]:
|
||||
doc_ids_direct.add(m["document_id"])
|
||||
|
||||
# Batch fetch all documents
|
||||
doc_map: dict[str, Any] = {}
|
||||
all_doc_ids = list(doc_ids_from_chunks | doc_ids_direct)
|
||||
if all_doc_ids:
|
||||
docs = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, original_text, metadata, retain_params
|
||||
FROM {fq_table("documents")}
|
||||
WHERE id = ANY($1) AND bank_id = $2
|
||||
""",
|
||||
all_doc_ids,
|
||||
bank_id,
|
||||
)
|
||||
doc_map = {row["id"]: row for row in docs}
|
||||
|
||||
# Build results
|
||||
results: list[dict[str, Any]] = []
|
||||
for mid, mem_uuid in zip(memory_ids, valid_uuids):
|
||||
if mid in errors:
|
||||
results.append({"memory_id": mid, "error": errors[mid]})
|
||||
continue
|
||||
|
||||
memory = memory_map.get(mem_uuid)
|
||||
if not memory:
|
||||
results.append({"memory_id": mid, "error": f"Memory not found: {mid}"})
|
||||
continue
|
||||
|
||||
item: dict[str, Any] = {
|
||||
"memory_id": mid,
|
||||
"memory": {
|
||||
"id": str(memory["id"]),
|
||||
"text": memory["text"],
|
||||
"type": memory["fact_type"],
|
||||
"context": memory["context"],
|
||||
},
|
||||
}
|
||||
|
||||
# Add chunk if available
|
||||
if memory["chunk_id"] and memory["chunk_id"] in chunk_map:
|
||||
chunk = chunk_map[memory["chunk_id"]]
|
||||
item["chunk"] = {
|
||||
"id": chunk["chunk_id"],
|
||||
"text": chunk["chunk_text"],
|
||||
"index": chunk["chunk_index"],
|
||||
"document_id": chunk["document_id"],
|
||||
}
|
||||
# Add document if depth=document
|
||||
if depth == "document" and chunk["document_id"] in doc_map:
|
||||
doc = doc_map[chunk["document_id"]]
|
||||
item["document"] = {
|
||||
"id": doc["id"],
|
||||
"full_text": doc["original_text"],
|
||||
"metadata": doc["metadata"],
|
||||
"retain_params": doc["retain_params"],
|
||||
}
|
||||
elif memory["document_id"] and depth == "document" and memory["document_id"] in doc_map:
|
||||
# No chunk, but has document_id
|
||||
doc = doc_map[memory["document_id"]]
|
||||
item["document"] = {
|
||||
"id": doc["id"],
|
||||
"full_text": doc["original_text"],
|
||||
"metadata": doc["metadata"],
|
||||
"retain_params": doc["retain_params"],
|
||||
}
|
||||
|
||||
results.append(item)
|
||||
|
||||
return {"results": results, "count": len(results)}
|
||||
@@ -0,0 +1,250 @@
|
||||
"""
|
||||
Tool schema definitions for the reflect agent.
|
||||
|
||||
These are OpenAI-format tool definitions used with native tool calling.
|
||||
The reflect agent uses a hierarchical retrieval strategy:
|
||||
1. search_mental_models - User-curated stored reflect responses (highest quality, if applicable)
|
||||
2. search_observations - Consolidated knowledge with freshness awareness
|
||||
3. recall - Raw facts (world/experience) as ground truth fallback
|
||||
"""
|
||||
|
||||
# Tool definitions in OpenAI format
|
||||
|
||||
TOOL_SEARCH_MENTAL_MODELS = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "search_mental_models",
|
||||
"description": (
|
||||
"Search user-curated mental models (stored reflect responses). These are high-quality, manually created "
|
||||
"summaries about specific topics. Use FIRST when the question might be covered by an "
|
||||
"existing mental model. Returns mental models with their content and last refresh time."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"reason": {
|
||||
"type": "string",
|
||||
"description": "Brief explanation of why you're making this search (for debugging)",
|
||||
},
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "Search query to find relevant mental models",
|
||||
},
|
||||
"max_results": {
|
||||
"type": "integer",
|
||||
"description": "Maximum number of mental models to return (default 5)",
|
||||
},
|
||||
},
|
||||
"required": ["reason", "query"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
TOOL_SEARCH_OBSERVATIONS = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "search_observations",
|
||||
"description": (
|
||||
"Search consolidated observations (auto-generated knowledge). These are automatically "
|
||||
"synthesized from memories. Returns observations with freshness info (updated_at, is_stale). "
|
||||
"If an observation is STALE, you should ALSO use recall() to verify with current facts."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"reason": {
|
||||
"type": "string",
|
||||
"description": "Brief explanation of why you're making this search (for debugging)",
|
||||
},
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "Search query to find relevant observations",
|
||||
},
|
||||
"max_tokens": {
|
||||
"type": "integer",
|
||||
"description": "Maximum tokens for results (default 5000). Use higher values for broader searches.",
|
||||
},
|
||||
},
|
||||
"required": ["reason", "query"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
TOOL_RECALL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "recall",
|
||||
"description": (
|
||||
"Search raw memories (facts and experiences). This is the ground truth data. "
|
||||
"Use when: (1) no reflections/mental models exist, (2) mental models are stale, "
|
||||
"(3) you need specific details not in synthesized knowledge. "
|
||||
"Returns individual memory facts with their timestamps."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"reason": {
|
||||
"type": "string",
|
||||
"description": "Brief explanation of why you're making this search (for debugging)",
|
||||
},
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "Search query string",
|
||||
},
|
||||
"max_tokens": {
|
||||
"type": "integer",
|
||||
"description": "Optional limit on result size (default 2048). Use higher values for broader searches.",
|
||||
},
|
||||
},
|
||||
"required": ["reason", "query"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
TOOL_EXPAND = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "expand",
|
||||
"description": "Get more context for one or more memories. Memory hierarchy: memory -> chunk -> document.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"reason": {
|
||||
"type": "string",
|
||||
"description": "Brief explanation of why you need more context (for debugging)",
|
||||
},
|
||||
"memory_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of memory IDs from recall results (batch multiple for efficiency)",
|
||||
},
|
||||
"depth": {
|
||||
"type": "string",
|
||||
"enum": ["chunk", "document"],
|
||||
"description": "chunk: surrounding text chunk, document: full source document",
|
||||
},
|
||||
},
|
||||
"required": ["reason", "memory_ids", "depth"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
TOOL_DONE_ANSWER = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "done",
|
||||
"description": "Signal completion with your final answer. Use this when you have gathered enough information to answer the question.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"answer": {
|
||||
"type": "string",
|
||||
"description": "Your response as well-formatted markdown. Use headers, lists, bold/italic, and code blocks for clarity. NEVER include memory IDs, UUIDs, or 'Memory references' in this text - put IDs only in memory_ids array.",
|
||||
},
|
||||
"memory_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of memory IDs that support your answer (put IDs here, NOT in answer text)",
|
||||
},
|
||||
"mental_model_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of mental model IDs that support your answer",
|
||||
},
|
||||
"observation_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of observation IDs that support your answer",
|
||||
},
|
||||
},
|
||||
"required": ["answer"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _build_done_tool_with_directives(directive_rules: list[str]) -> dict:
|
||||
"""
|
||||
Build the done tool schema with directive compliance field.
|
||||
|
||||
When directives are present, adds a required field that forces the agent
|
||||
to confirm compliance with each directive before submitting.
|
||||
|
||||
Args:
|
||||
directive_rules: List of directive rule strings
|
||||
"""
|
||||
# Build rules list for description
|
||||
rules_list = "\n".join(f" {i + 1}. {rule}" for i, rule in enumerate(directive_rules))
|
||||
|
||||
# Build the tool with directive compliance field
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "done",
|
||||
"description": (
|
||||
"Signal completion with your final answer. IMPORTANT: You must confirm directive compliance before submitting. "
|
||||
"Your answer will be REJECTED if it violates any directive."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"answer": {
|
||||
"type": "string",
|
||||
"description": "Your response as well-formatted markdown. Use headers, lists, bold/italic, and code blocks for clarity. NEVER include memory IDs, UUIDs, or 'Memory references' in this text - put IDs only in memory_ids array.",
|
||||
},
|
||||
"memory_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of memory IDs that support your answer (put IDs here, NOT in answer text)",
|
||||
},
|
||||
"mental_model_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of mental model IDs that support your answer",
|
||||
},
|
||||
"observation_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of observation IDs that support your answer",
|
||||
},
|
||||
"directive_compliance": {
|
||||
"type": "string",
|
||||
"description": f"REQUIRED: Confirm your answer complies with ALL directives. List each directive and how your answer follows it:\n{rules_list}\n\nFormat: 'Directive 1: [how answer complies]. Directive 2: [how answer complies]...'",
|
||||
},
|
||||
},
|
||||
"required": ["answer", "directive_compliance"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def get_reflect_tools(directive_rules: list[str] | None = None) -> list[dict]:
|
||||
"""
|
||||
Get the list of tools for the reflect agent.
|
||||
|
||||
The tools support a hierarchical retrieval strategy:
|
||||
1. search_mental_models - User-curated stored reflect responses (try first)
|
||||
2. search_observations - Consolidated knowledge with freshness
|
||||
3. recall - Raw facts as ground truth
|
||||
|
||||
Args:
|
||||
directive_rules: Optional list of directive rule strings. If provided,
|
||||
the done() tool will require directive compliance confirmation.
|
||||
|
||||
Returns:
|
||||
List of tool definitions in OpenAI format
|
||||
"""
|
||||
tools = [
|
||||
TOOL_SEARCH_MENTAL_MODELS,
|
||||
TOOL_SEARCH_OBSERVATIONS,
|
||||
TOOL_RECALL,
|
||||
TOOL_EXPAND,
|
||||
]
|
||||
|
||||
# Use directive-aware done tool if directives are present
|
||||
if directive_rules:
|
||||
tools.append(_build_done_tool_with_directives(directive_rules))
|
||||
else:
|
||||
tools.append(TOOL_DONE_ANSWER)
|
||||
|
||||
return tools
|
||||
@@ -10,8 +10,63 @@ from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
# Valid fact types for recall operations (excludes 'observation' which is internal)
|
||||
VALID_RECALL_FACT_TYPES = frozenset(["world", "experience", "opinion"])
|
||||
# Valid fact types for recall operations (excludes 'opinion' which is deprecated)
|
||||
VALID_RECALL_FACT_TYPES = frozenset(["world", "experience", "observation"])
|
||||
|
||||
|
||||
class LLMToolCall(BaseModel):
|
||||
"""A tool call requested by the LLM."""
|
||||
|
||||
id: str = Field(description="Unique identifier for this tool call")
|
||||
name: str = Field(description="Name of the tool to call")
|
||||
arguments: dict[str, Any] = Field(description="Arguments to pass to the tool")
|
||||
|
||||
|
||||
class LLMToolCallResult(BaseModel):
|
||||
"""Result from an LLM call that may include tool calls."""
|
||||
|
||||
content: str | None = Field(default=None, description="Text content if any")
|
||||
tool_calls: list[LLMToolCall] = Field(default_factory=list, description="Tool calls requested by the LLM")
|
||||
finish_reason: str | None = Field(default=None, description="Reason the LLM stopped: 'stop', 'tool_calls', etc.")
|
||||
input_tokens: int = Field(default=0, description="Input tokens used in this call")
|
||||
output_tokens: int = Field(default=0, description="Output tokens used in this call")
|
||||
|
||||
|
||||
class ToolCallTrace(BaseModel):
|
||||
"""A single tool call made during reflect."""
|
||||
|
||||
tool: str = Field(description="Tool name: lookup, recall, learn, expand")
|
||||
reason: str | None = Field(default=None, description="Agent's reasoning for making this tool call")
|
||||
input: dict = Field(description="Tool input parameters")
|
||||
output: dict = Field(description="Tool output/result")
|
||||
duration_ms: int = Field(description="Execution time in milliseconds")
|
||||
iteration: int = Field(default=0, description="Iteration number (1-based) when this tool was called")
|
||||
|
||||
|
||||
class LLMCallTrace(BaseModel):
|
||||
"""A single LLM call made during reflect."""
|
||||
|
||||
scope: str = Field(description="Call scope: agent_1, agent_2, final, etc.")
|
||||
duration_ms: int = Field(description="Execution time in milliseconds")
|
||||
|
||||
|
||||
class ObservationRef(BaseModel):
|
||||
"""Reference to an observation accessed during reflect."""
|
||||
|
||||
id: str = Field(description="Observation ID")
|
||||
name: str = Field(description="Observation name")
|
||||
type: str = Field(description="Observation type: entity, concept, event")
|
||||
subtype: str = Field(description="Observation subtype: structural, emergent, learned")
|
||||
description: str = Field(description="Brief description")
|
||||
summary: str | None = Field(default=None, description="Full summary (when looked up in detail)")
|
||||
|
||||
|
||||
class DirectiveRef(BaseModel):
|
||||
"""Reference to a directive that was applied during reflect."""
|
||||
|
||||
id: str = Field(description="Directive mental model ID")
|
||||
name: str = Field(description="Directive name")
|
||||
content: str = Field(description="Directive content")
|
||||
|
||||
|
||||
class TokenUsage(BaseModel):
|
||||
@@ -114,6 +169,28 @@ class ChunkInfo(BaseModel):
|
||||
truncated: bool = Field(default=False, description="Whether the chunk was truncated due to token limits")
|
||||
|
||||
|
||||
class ObservationResult(BaseModel):
|
||||
"""An observation result from recall (consolidated knowledge synthesized from facts)."""
|
||||
|
||||
id: str = Field(description="Unique observation ID")
|
||||
text: str = Field(description="The observation text")
|
||||
proof_count: int = Field(description="Number of facts supporting this observation")
|
||||
relevance: float = Field(default=0.0, description="Relevance score to the query")
|
||||
tags: list[str] | None = Field(default=None, description="Tags for visibility scoping")
|
||||
source_memory_ids: list[str] = Field(
|
||||
default_factory=list, description="IDs of facts that contribute to this observation"
|
||||
)
|
||||
|
||||
|
||||
class MentalModelResult(BaseModel):
|
||||
"""A mental model result from recall (stored reflect response)."""
|
||||
|
||||
id: str = Field(description="Unique mental model ID")
|
||||
name: str = Field(description="Human-readable name")
|
||||
content: str = Field(description="The synthesized content")
|
||||
relevance: float = Field(default=0.0, description="Relevance score to the query")
|
||||
|
||||
|
||||
class RecallResult(BaseModel):
|
||||
"""
|
||||
Result from a recall operation.
|
||||
@@ -177,8 +254,15 @@ class ReflectResult(BaseModel):
|
||||
],
|
||||
"experience": [],
|
||||
"opinion": [],
|
||||
"mental_models": [],
|
||||
"directives": [
|
||||
{
|
||||
"id": "directive-123",
|
||||
"name": "Response Style",
|
||||
"rules": ["Always be concise"],
|
||||
}
|
||||
],
|
||||
},
|
||||
"new_opinions": ["Machine learning has great potential in healthcare"],
|
||||
"structured_output": {"summary": "ML in healthcare", "confidence": 0.9},
|
||||
"usage": {"input_tokens": 1500, "output_tokens": 500, "total_tokens": 2000},
|
||||
}
|
||||
@@ -186,10 +270,9 @@ class ReflectResult(BaseModel):
|
||||
)
|
||||
|
||||
text: str = Field(description="The formulated answer text")
|
||||
based_on: dict[str, list[MemoryFact]] = Field(
|
||||
description="Facts used to formulate the answer, organized by type (world, experience, opinion)"
|
||||
based_on: dict[str, Any] = Field(
|
||||
description="Facts used to formulate the answer, organized by type (world, experience, mental_models, directives)"
|
||||
)
|
||||
new_opinions: list[str] = Field(default_factory=list, description="List of newly formed opinions during reflection")
|
||||
structured_output: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description="Structured output parsed according to the provided response schema. Only present when response_schema was provided.",
|
||||
@@ -198,24 +281,18 @@ class ReflectResult(BaseModel):
|
||||
default=None,
|
||||
description="Token usage metrics for the LLM calls made during this reflect operation.",
|
||||
)
|
||||
|
||||
|
||||
class Opinion(BaseModel):
|
||||
"""
|
||||
An opinion with confidence score.
|
||||
|
||||
Opinions represent the bank's formed perspectives on topics,
|
||||
with a confidence level indicating strength of belief.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {"text": "Machine learning has great potential in healthcare", "confidence": 0.85}
|
||||
}
|
||||
tool_trace: list[ToolCallTrace] = Field(
|
||||
default_factory=list,
|
||||
description="Trace of tool calls made during reflection. Only present when include.tool_calls is enabled.",
|
||||
)
|
||||
llm_trace: list[LLMCallTrace] = Field(
|
||||
default_factory=list,
|
||||
description="Trace of LLM calls made during reflection. Only present when include.tool_calls is enabled.",
|
||||
)
|
||||
directives_applied: list[DirectiveRef] = Field(
|
||||
default_factory=list,
|
||||
description="Directive mental models that were applied during this reflection.",
|
||||
)
|
||||
|
||||
text: str = Field(description="The opinion text")
|
||||
confidence: float = Field(description="Confidence score between 0.0 and 1.0")
|
||||
|
||||
|
||||
class EntityObservation(BaseModel):
|
||||
@@ -261,3 +338,32 @@ class EntityState(BaseModel):
|
||||
observations: list[EntityObservation] = Field(
|
||||
default_factory=list, description="List of observations about this entity"
|
||||
)
|
||||
|
||||
|
||||
class MentalModel(BaseModel):
|
||||
"""
|
||||
A manually configured mental model for tracking specific topics/areas.
|
||||
|
||||
Mental models are user-defined focus areas that the agent should track
|
||||
and maintain summaries for, unlike auto-extracted entities.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"id": "team-dynamics",
|
||||
"name": "Team Dynamics",
|
||||
"description": "Track how the team collaborates, communication patterns, conflicts, and resolutions",
|
||||
"summary": "The team has strong collaboration...",
|
||||
"summary_updated_at": "2024-01-15T10:30:00Z",
|
||||
"created_at": "2024-01-10T08:00:00Z",
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
id: str = Field(description="Unique identifier (alphanumeric lowercase)")
|
||||
name: str = Field(description="Display name for the mental model")
|
||||
description: str = Field(description="Prompt/directions for what to track and summarize")
|
||||
summary: str | None = Field(None, description="Generated summary based on relevant facts")
|
||||
summary_updated_at: str | None = Field(None, description="ISO format date when summary was last updated")
|
||||
created_at: str = Field(description="ISO format date when the mental model was created")
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
"""
|
||||
bank profile utilities for disposition and background management.
|
||||
bank profile utilities for disposition and mission management.
|
||||
"""
|
||||
|
||||
import json
|
||||
@@ -27,19 +27,18 @@ class BankProfile(TypedDict):
|
||||
|
||||
name: str
|
||||
disposition: DispositionTraits
|
||||
background: str
|
||||
mission: str
|
||||
|
||||
|
||||
class BackgroundMergeResponse(BaseModel):
|
||||
"""LLM response for background merge with disposition inference."""
|
||||
class MissionMergeResponse(BaseModel):
|
||||
"""LLM response for mission merge."""
|
||||
|
||||
background: str = Field(description="Merged background in first person perspective")
|
||||
disposition: DispositionTraits = Field(description="Inferred disposition traits (skepticism, literalism, empathy)")
|
||||
mission: str = Field(description="Merged mission in first person perspective")
|
||||
|
||||
|
||||
async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
"""
|
||||
Get bank profile (name, disposition + background).
|
||||
Get bank profile (name, disposition + mission).
|
||||
Auto-creates bank with default values if not exists.
|
||||
|
||||
Args:
|
||||
@@ -47,13 +46,13 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
bank_id: bank IDentifier
|
||||
|
||||
Returns:
|
||||
BankProfile with name, typed DispositionTraits, and background
|
||||
BankProfile with name, typed DispositionTraits, and mission
|
||||
"""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Try to get existing bank
|
||||
row = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT name, disposition, background
|
||||
SELECT name, disposition, mission
|
||||
FROM {fq_table("banks")} WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
@@ -66,13 +65,15 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
disposition_data = json.loads(disposition_data)
|
||||
|
||||
return BankProfile(
|
||||
name=row["name"], disposition=DispositionTraits(**disposition_data), background=row["background"]
|
||||
name=row["name"],
|
||||
disposition=DispositionTraits(**disposition_data),
|
||||
mission=row["mission"] or "",
|
||||
)
|
||||
|
||||
# Bank doesn't exist, create with defaults
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("banks")} (bank_id, name, disposition, background)
|
||||
INSERT INTO {fq_table("banks")} (bank_id, name, disposition, mission)
|
||||
VALUES ($1, $2, $3::jsonb, $4)
|
||||
ON CONFLICT (bank_id) DO NOTHING
|
||||
""",
|
||||
@@ -82,7 +83,7 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
"",
|
||||
)
|
||||
|
||||
return BankProfile(name=bank_id, disposition=DispositionTraits(**DEFAULT_DISPOSITION), background="")
|
||||
return BankProfile(name=bank_id, disposition=DispositionTraits(**DEFAULT_DISPOSITION), mission="")
|
||||
|
||||
|
||||
async def update_bank_disposition(pool, bank_id: str, disposition: dict[str, int]) -> None:
|
||||
@@ -110,244 +111,121 @@ async def update_bank_disposition(pool, bank_id: str, disposition: dict[str, int
|
||||
)
|
||||
|
||||
|
||||
async def merge_bank_background(pool, llm_config, bank_id: str, new_info: str, update_disposition: bool = True) -> dict:
|
||||
async def set_bank_mission(pool, bank_id: str, mission: str) -> None:
|
||||
"""
|
||||
Merge new background information with existing background using LLM.
|
||||
Normalizes to first person ("I") and resolves conflicts.
|
||||
Optionally infers disposition traits from the merged background.
|
||||
Set bank mission (replacing any existing mission).
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
llm_config: LLM configuration for background merging
|
||||
bank_id: bank IDentifier
|
||||
new_info: New background information to add/merge
|
||||
update_disposition: If True, infer Big Five traits from background (default: True)
|
||||
mission: The mission text
|
||||
"""
|
||||
# Ensure bank exists first
|
||||
await get_bank_profile(pool, bank_id)
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("banks")}
|
||||
SET mission = $2,
|
||||
updated_at = NOW()
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
mission,
|
||||
)
|
||||
|
||||
|
||||
async def merge_bank_mission(pool, llm_config, bank_id: str, new_info: str) -> dict:
|
||||
"""
|
||||
Merge new mission information with existing mission using LLM.
|
||||
Normalizes to first person ("I") and resolves conflicts.
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
llm_config: LLM configuration for mission merging
|
||||
bank_id: bank IDentifier
|
||||
new_info: New mission information to add/merge
|
||||
|
||||
Returns:
|
||||
Dict with 'background' (str) and optionally 'disposition' (dict) keys
|
||||
Dict with 'mission' (str) key
|
||||
"""
|
||||
# Get current profile
|
||||
profile = await get_bank_profile(pool, bank_id)
|
||||
current_background = profile["background"]
|
||||
current_mission = profile["mission"]
|
||||
|
||||
# Use LLM to merge backgrounds and optionally infer disposition
|
||||
result = await _llm_merge_background(llm_config, current_background, new_info, infer_disposition=update_disposition)
|
||||
# Use LLM to merge missions
|
||||
result = await _llm_merge_mission(llm_config, current_mission, new_info)
|
||||
|
||||
merged_background = result["background"]
|
||||
inferred_disposition = result.get("disposition")
|
||||
merged_mission = result["mission"]
|
||||
|
||||
# Update in database
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
if inferred_disposition:
|
||||
# Update both background and disposition
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("banks")}
|
||||
SET background = $2,
|
||||
disposition = $3::jsonb,
|
||||
updated_at = NOW()
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
merged_background,
|
||||
json.dumps(inferred_disposition),
|
||||
)
|
||||
else:
|
||||
# Update only background
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("banks")}
|
||||
SET background = $2,
|
||||
updated_at = NOW()
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
merged_background,
|
||||
)
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("banks")}
|
||||
SET mission = $2,
|
||||
updated_at = NOW()
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
merged_mission,
|
||||
)
|
||||
|
||||
response = {"background": merged_background}
|
||||
if inferred_disposition:
|
||||
response["disposition"] = inferred_disposition
|
||||
|
||||
return response
|
||||
return {"mission": merged_mission}
|
||||
|
||||
|
||||
async def _llm_merge_background(llm_config, current: str, new_info: str, infer_disposition: bool = False) -> dict:
|
||||
async def _llm_merge_mission(llm_config, current: str, new_info: str) -> dict:
|
||||
"""
|
||||
Use LLM to intelligently merge background information.
|
||||
Optionally infer Big Five disposition traits from the merged background.
|
||||
Use LLM to intelligently merge mission information.
|
||||
|
||||
Args:
|
||||
llm_config: LLM configuration to use
|
||||
current: Current background text
|
||||
current: Current mission text
|
||||
new_info: New information to merge
|
||||
infer_disposition: If True, also infer disposition traits
|
||||
|
||||
Returns:
|
||||
Dict with 'background' (str) and optionally 'disposition' (dict) keys
|
||||
Dict with 'mission' (str) key
|
||||
"""
|
||||
if infer_disposition:
|
||||
prompt = f"""You are helping maintain a memory bank's background/profile and infer their disposition. You MUST respond with ONLY valid JSON.
|
||||
prompt = f"""You are helping maintain an agent's mission statement.
|
||||
|
||||
Current background: {current if current else "(empty)"}
|
||||
Current mission: {current if current else "(empty)"}
|
||||
|
||||
New information to add: {new_info}
|
||||
|
||||
Instructions:
|
||||
1. Merge the new information with the current background
|
||||
2. If there are conflicts (e.g., different birthplaces), the NEW information overwrites the old
|
||||
3. Keep additions that don't conflict
|
||||
4. Output in FIRST PERSON ("I") perspective
|
||||
5. Be concise - keep merged background under 500 characters
|
||||
6. Infer disposition traits from the merged background (each 1-5 integer):
|
||||
- Skepticism: 1-5 (1=trusting, takes things at face value; 5=skeptical, questions everything)
|
||||
- Literalism: 1-5 (1=flexible interpretation, reads between lines; 5=literal, exact interpretation)
|
||||
- Empathy: 1-5 (1=detached, focuses on facts; 5=empathetic, considers emotional context)
|
||||
|
||||
CRITICAL: You MUST respond with ONLY a valid JSON object. No markdown, no code blocks, no explanations. Just the JSON.
|
||||
|
||||
Format:
|
||||
{{
|
||||
"background": "the merged background text in first person",
|
||||
"disposition": {{
|
||||
"skepticism": 3,
|
||||
"literalism": 3,
|
||||
"empathy": 3
|
||||
}}
|
||||
}}
|
||||
|
||||
Trait inference examples:
|
||||
- "I'm a lawyer" → skepticism: 4, literalism: 5, empathy: 2
|
||||
- "I'm a therapist" → skepticism: 2, literalism: 2, empathy: 5
|
||||
- "I'm an engineer" → skepticism: 3, literalism: 4, empathy: 3
|
||||
- "I've been burned before by trusting people" → skepticism: 5, literalism: 3, empathy: 3
|
||||
- "I try to understand what people really mean" → skepticism: 3, literalism: 2, empathy: 4
|
||||
- "I take contracts very seriously" → skepticism: 4, literalism: 5, empathy: 2"""
|
||||
else:
|
||||
prompt = f"""You are helping maintain a memory bank's background/profile.
|
||||
|
||||
Current background: {current if current else "(empty)"}
|
||||
|
||||
New information to add: {new_info}
|
||||
|
||||
Instructions:
|
||||
1. Merge the new information with the current background
|
||||
2. If there are conflicts (e.g., different birthplaces), the NEW information overwrites the old
|
||||
1. Merge the new information with the current mission
|
||||
2. If there are conflicts, the NEW information overwrites the old
|
||||
3. Keep additions that don't conflict
|
||||
4. Output in FIRST PERSON ("I") perspective
|
||||
5. Be concise - keep it under 500 characters
|
||||
6. Return ONLY the merged background text, no explanations
|
||||
6. Return ONLY the merged mission text, no explanations
|
||||
|
||||
Merged background:"""
|
||||
Merged mission:"""
|
||||
|
||||
try:
|
||||
# Prepare messages
|
||||
messages = [{"role": "user", "content": prompt}]
|
||||
|
||||
if infer_disposition:
|
||||
# Use structured output with Pydantic model for disposition inference
|
||||
try:
|
||||
parsed = await llm_config.call(
|
||||
messages=messages,
|
||||
response_format=BackgroundMergeResponse,
|
||||
scope="bank_background",
|
||||
temperature=0.3,
|
||||
max_completion_tokens=8192,
|
||||
)
|
||||
logger.info(f"Successfully got structured response: background={parsed.background[:100]}")
|
||||
|
||||
# Convert Pydantic model to dict format
|
||||
return {"background": parsed.background, "disposition": parsed.disposition.model_dump()}
|
||||
except Exception as e:
|
||||
logger.warning(f"Structured output failed, falling back to manual parsing: {e}")
|
||||
# Fall through to manual parsing below
|
||||
|
||||
# Manual parsing fallback or non-disposition merge
|
||||
content = await llm_config.call(
|
||||
messages=messages, scope="bank_background", temperature=0.3, max_completion_tokens=8192
|
||||
messages=messages, scope="bank_mission", temperature=0.3, max_completion_tokens=8192
|
||||
)
|
||||
|
||||
logger.info(f"LLM response for background merge (first 500 chars): {content[:500]}")
|
||||
logger.info(f"LLM response for mission merge (first 500 chars): {content[:500]}")
|
||||
|
||||
if infer_disposition:
|
||||
# Parse JSON response - try multiple extraction methods
|
||||
result = None
|
||||
|
||||
# Method 1: Direct parse
|
||||
try:
|
||||
result = json.loads(content)
|
||||
logger.info("Successfully parsed JSON directly")
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# Method 2: Extract from markdown code blocks
|
||||
if result is None:
|
||||
# Remove markdown code blocks
|
||||
code_block_match = re.search(r"```(?:json)?\s*(\{.*?\})\s*```", content, re.DOTALL)
|
||||
if code_block_match:
|
||||
try:
|
||||
result = json.loads(code_block_match.group(1))
|
||||
logger.info("Successfully extracted JSON from markdown code block")
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# Method 3: Find nested JSON structure
|
||||
if result is None:
|
||||
# Look for JSON object with nested structure
|
||||
json_match = re.search(
|
||||
r'\{[^{}]*"background"[^{}]*"disposition"[^{}]*\{[^{}]*\}[^{}]*\}', content, re.DOTALL
|
||||
)
|
||||
if json_match:
|
||||
try:
|
||||
result = json.loads(json_match.group())
|
||||
logger.info("Successfully extracted JSON using nested pattern")
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# All parsing methods failed - use fallback
|
||||
if result is None:
|
||||
logger.warning(f"Failed to extract JSON from LLM response. Raw content: {content[:200]}")
|
||||
# Fallback: use new_info as background with default disposition
|
||||
return {
|
||||
"background": new_info if new_info else current if current else "",
|
||||
"disposition": DEFAULT_DISPOSITION.copy(),
|
||||
}
|
||||
|
||||
# Validate disposition values
|
||||
disposition = result.get("disposition", {})
|
||||
for key in ["skepticism", "literalism", "empathy"]:
|
||||
if key not in disposition:
|
||||
disposition[key] = 3 # Default to neutral
|
||||
else:
|
||||
# Clamp to [1, 5] and convert to int
|
||||
disposition[key] = max(1, min(5, int(disposition[key])))
|
||||
|
||||
result["disposition"] = disposition
|
||||
|
||||
# Ensure background exists
|
||||
if "background" not in result or not result["background"]:
|
||||
result["background"] = new_info if new_info else ""
|
||||
|
||||
return result
|
||||
else:
|
||||
# Just background merge
|
||||
merged = content
|
||||
if not merged or merged.lower() in ["(empty)", "none", "n/a"]:
|
||||
merged = new_info if new_info else ""
|
||||
return {"background": merged}
|
||||
merged = content.strip()
|
||||
if not merged or merged.lower() in ["(empty)", "none", "n/a"]:
|
||||
merged = new_info if new_info else ""
|
||||
return {"mission": merged}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error merging background with LLM: {e}")
|
||||
logger.error(f"Error merging mission with LLM: {e}")
|
||||
# Fallback: just append new info
|
||||
if current:
|
||||
merged = f"{current} {new_info}".strip()
|
||||
else:
|
||||
merged = new_info
|
||||
|
||||
result = {"background": merged}
|
||||
if infer_disposition:
|
||||
result["disposition"] = DEFAULT_DISPOSITION.copy()
|
||||
return result
|
||||
return {"mission": merged}
|
||||
|
||||
|
||||
async def list_banks(pool) -> list:
|
||||
@@ -358,12 +236,12 @@ async def list_banks(pool) -> list:
|
||||
pool: Database connection pool
|
||||
|
||||
Returns:
|
||||
List of dicts with bank_id, name, disposition, background, created_at, updated_at
|
||||
List of dicts with bank_id, name, disposition, mission, created_at, updated_at
|
||||
"""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT bank_id, name, disposition, background, created_at, updated_at
|
||||
SELECT bank_id, name, disposition, mission, created_at, updated_at
|
||||
FROM {fq_table("banks")}
|
||||
ORDER BY updated_at DESC
|
||||
"""
|
||||
@@ -381,7 +259,7 @@ async def list_banks(pool) -> list:
|
||||
"bank_id": row["bank_id"],
|
||||
"name": row["name"],
|
||||
"disposition": disposition_data,
|
||||
"background": row["background"],
|
||||
"mission": row["mission"] or "",
|
||||
"created_at": row["created_at"].isoformat() if row["created_at"] else None,
|
||||
"updated_at": row["updated_at"].isoformat() if row["updated_at"] else None,
|
||||
}
|
||||
|
||||
@@ -57,21 +57,25 @@ def _infer_temporal_date(fact_text: str, event_date: datetime) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _sanitize_text(text: str) -> str:
|
||||
def _sanitize_text(text: str | None) -> str | None:
|
||||
"""
|
||||
Sanitize text by removing invalid Unicode surrogate characters.
|
||||
Sanitize text by removing characters that break downstream systems.
|
||||
|
||||
Surrogate characters (U+D800 to U+DFFF) are used in UTF-16 encoding
|
||||
but cannot be encoded in UTF-8. They can appear in Python strings
|
||||
from improperly decoded data (e.g., from JavaScript or broken files).
|
||||
Removes:
|
||||
- Null bytes (\\x00): Invalid in PostgreSQL UTF-8 encoding
|
||||
- Unicode surrogates (U+D800-U+DFFF): Invalid in UTF-8, break LLM APIs
|
||||
|
||||
This function removes unpaired surrogates to prevent UnicodeEncodeError
|
||||
when the text is sent to the LLM API.
|
||||
Surrogate characters are used in UTF-16 encoding but cannot be encoded
|
||||
in UTF-8. They can appear in Python strings from improperly decoded data
|
||||
(e.g., from JavaScript or broken files). Null bytes commonly appear in
|
||||
OCR output, PDF extraction, or copy-paste from binary sources.
|
||||
"""
|
||||
if text is None:
|
||||
return None
|
||||
if not text:
|
||||
return text
|
||||
# Remove surrogate characters (U+D800 to U+DFFF) using regex
|
||||
# These are invalid in UTF-8 and cause encoding errors
|
||||
# Remove null bytes and surrogate characters
|
||||
text = text.replace("\x00", "")
|
||||
return re.sub(r"[\ud800-\udfff]", "", text)
|
||||
|
||||
|
||||
@@ -114,11 +118,8 @@ class CausalRelation(BaseModel):
|
||||
"""Causal relationship from this fact to a previous fact (stored format)."""
|
||||
|
||||
target_fact_index: int = Field(description="Index of the related fact in the facts array (0-based).")
|
||||
relation_type: Literal["caused_by", "enabled_by", "prevented_by"] = Field(
|
||||
description="How this fact relates to the target: "
|
||||
"'caused_by' = this fact was caused by the target, "
|
||||
"'enabled_by' = this fact was enabled by the target, "
|
||||
"'prevented_by' = this fact was prevented by the target"
|
||||
relation_type: Literal["caused_by"] = Field(
|
||||
description="How this fact relates to the target: 'caused_by' = this fact was caused by the target"
|
||||
)
|
||||
strength: float = Field(
|
||||
description="Strength of relationship (0.0 to 1.0)",
|
||||
@@ -141,11 +142,8 @@ class FactCausalRelation(BaseModel):
|
||||
"MUST be less than this fact's position in the list. "
|
||||
"Example: if this is fact #5, target_index can only be 0, 1, 2, 3, or 4."
|
||||
)
|
||||
relation_type: Literal["caused_by", "enabled_by", "prevented_by"] = Field(
|
||||
description="How this fact relates to the target fact: "
|
||||
"'caused_by' = this fact was caused by the target fact, "
|
||||
"'enabled_by' = this fact was enabled by the target fact, "
|
||||
"'prevented_by' = this fact was blocked/prevented by the target fact"
|
||||
relation_type: Literal["caused_by"] = Field(
|
||||
description="How this fact relates to the target fact: 'caused_by' = this fact was caused by the target fact"
|
||||
)
|
||||
strength: float = Field(
|
||||
description="Strength of relationship (0.0 to 1.0). 1.0 = strong, 0.5 = moderate",
|
||||
@@ -438,34 +436,15 @@ def _chunk_conversation(turns: list[dict], max_chars: int) -> list[str]:
|
||||
# FACT EXTRACTION PROMPTS
|
||||
# =============================================================================
|
||||
|
||||
# Concise extraction prompt (default) - selective, high-quality facts
|
||||
CONCISE_FACT_EXTRACTION_PROMPT = """Extract SIGNIFICANT facts from text. Be SELECTIVE - only extract facts worth remembering long-term.
|
||||
# Base prompt template (shared by concise and custom modes)
|
||||
# Uses {extraction_guidelines} placeholder for mode-specific instructions
|
||||
_BASE_FACT_EXTRACTION_PROMPT = """Extract SIGNIFICANT facts from text. Be SELECTIVE - only extract facts worth remembering long-term.
|
||||
|
||||
LANGUAGE RULE (CRITICAL): Output facts in the EXACT SAME language as the input text. If input is Japanese, output Japanese. If input is Chinese, output Chinese. NEVER translate to English. Preserve original language completely.
|
||||
LANGUAGE REQUIREMENT: Detect the language of the input text. All extracted facts, entity names, descriptions, and other output MUST be in the SAME language as the input. Do not translate to another language.
|
||||
|
||||
{fact_types_instruction}
|
||||
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
SELECTIVITY - CRITICAL (Reduces 90% of unnecessary output)
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
ONLY extract facts that are:
|
||||
✅ Personal info: names, relationships, roles, background
|
||||
✅ Preferences: likes, dislikes, habits, interests (e.g., "Alice likes coffee")
|
||||
✅ Significant events: milestones, decisions, achievements, changes
|
||||
✅ Plans/goals: future intentions, deadlines, commitments
|
||||
✅ Expertise: skills, knowledge, certifications, experience
|
||||
✅ Important context: projects, problems, constraints
|
||||
✅ Sensory/emotional details: feelings, sensations, perceptions that provide context
|
||||
✅ Observations: descriptions of people, places, things with specific details
|
||||
|
||||
DO NOT extract:
|
||||
❌ Generic greetings: "how are you", "hello", pleasantries without substance
|
||||
❌ Pure filler: "thanks", "sounds good", "ok", "got it", "sure"
|
||||
❌ Process chatter: "let me check", "one moment", "I'll look into it"
|
||||
❌ Repeated info: if already stated, don't extract again
|
||||
|
||||
CONSOLIDATE related statements into ONE fact when possible.
|
||||
{extraction_guidelines}
|
||||
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
FACT FORMAT - BE CONCISE
|
||||
@@ -513,7 +492,33 @@ ENTITIES
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
Include: people names, organizations, places, key objects, abstract concepts (career, friendship, etc.)
|
||||
Always include "user" when fact is about the user.
|
||||
Always include "user" when fact is about the user.{examples}"""
|
||||
|
||||
# Concise mode guidelines
|
||||
_CONCISE_GUIDELINES = """══════════════════════════════════════════════════════════════════════════
|
||||
SELECTIVITY - CRITICAL (Reduces 90% of unnecessary output)
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
ONLY extract facts that are:
|
||||
✅ Personal info: names, relationships, roles, background
|
||||
✅ Preferences: likes, dislikes, habits, interests (e.g., "Alice likes coffee")
|
||||
✅ Significant events: milestones, decisions, achievements, changes
|
||||
✅ Plans/goals: future intentions, deadlines, commitments
|
||||
✅ Expertise: skills, knowledge, certifications, experience
|
||||
✅ Important context: projects, problems, constraints
|
||||
✅ Sensory/emotional details: feelings, sensations, perceptions that provide context
|
||||
✅ Observations: descriptions of people, places, things with specific details
|
||||
|
||||
DO NOT extract:
|
||||
❌ Generic greetings: "how are you", "hello", pleasantries without substance
|
||||
❌ Pure filler: "thanks", "sounds good", "ok", "got it", "sure"
|
||||
❌ Process chatter: "let me check", "one moment", "I'll look into it"
|
||||
❌ Repeated info: if already stated, don't extract again
|
||||
|
||||
CONSOLIDATE related statements into ONE fact when possible."""
|
||||
|
||||
# Concise mode examples
|
||||
_CONCISE_EXAMPLES = """
|
||||
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
EXAMPLES
|
||||
@@ -537,7 +542,26 @@ Output: ONLY 2 facts (skip coffee preference - too trivial):
|
||||
QUALITY OVER QUANTITY
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
Ask: "Would this be useful to recall in 6 months?" If no, skip it."""
|
||||
Ask: "Would this be useful to recall in 6 months?" If no, skip it.
|
||||
|
||||
IMPORTANT: Sensory/emotional details and observations that provide meaningful context
|
||||
about experiences ARE important to remember, even if they seem small (e.g., how food
|
||||
tasted, how someone looked, how loud music was). Extract these if they characterize
|
||||
an experience or person."""
|
||||
|
||||
# Assembled concise prompt (backward compatible - exact same output as before)
|
||||
CONCISE_FACT_EXTRACTION_PROMPT = _BASE_FACT_EXTRACTION_PROMPT.format(
|
||||
fact_types_instruction="{fact_types_instruction}",
|
||||
extraction_guidelines=_CONCISE_GUIDELINES,
|
||||
examples=_CONCISE_EXAMPLES,
|
||||
)
|
||||
|
||||
# Custom prompt uses same base but without examples
|
||||
CUSTOM_FACT_EXTRACTION_PROMPT = _BASE_FACT_EXTRACTION_PROMPT.format(
|
||||
fact_types_instruction="{fact_types_instruction}",
|
||||
extraction_guidelines="{custom_instructions}",
|
||||
examples="", # No examples for custom mode
|
||||
)
|
||||
|
||||
|
||||
# Verbose extraction prompt - detailed, comprehensive facts (legacy mode)
|
||||
@@ -622,6 +646,7 @@ For EVENTS (fact_kind="event") - MUST SET BOTH occurred_start AND occurred_end:
|
||||
- Convert relative dates → absolute using Event Date as reference
|
||||
- If Event Date is "Saturday, March 15, 2020", then "yesterday" = Friday, March 14, 2020
|
||||
- Dates mentioned in text (e.g., "in March 2020") should use THAT year, not current year
|
||||
- CRITICAL: If the content mentions an absolute date (e.g., "March 15, 2024", "2024-03-15"), you MUST extract it and set occurred_start in ISO format
|
||||
- Always include the day name (Monday, Tuesday, etc.) in the 'when' field
|
||||
- Set occurred_start AND occurred_end to WHEN IT HAPPENED (not when mentioned)
|
||||
- For single-day/point events: set occurred_end = occurred_start (same timestamp)
|
||||
@@ -662,7 +687,7 @@ CAUSAL RELATIONSHIPS
|
||||
══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
Link facts with causal_relations (max 2 per fact). target_index must be < this fact's index.
|
||||
Types: "caused_by", "enabled_by", "prevented_by"
|
||||
Type: "caused_by" (this fact was caused by the target fact)
|
||||
|
||||
Example: "Lost job → couldn't pay rent → moved apartment"
|
||||
- Fact 0: Lost job, causal_relations: null
|
||||
@@ -678,7 +703,6 @@ async def _extract_facts_from_chunk(
|
||||
context: str,
|
||||
llm_config: "LLMConfig",
|
||||
agent_name: str = None,
|
||||
extract_opinions: bool = False,
|
||||
) -> tuple[list[dict[str, str]], TokenUsage]:
|
||||
"""
|
||||
Extract facts from a single chunk (internal helper for parallel processing).
|
||||
@@ -686,17 +710,15 @@ async def _extract_facts_from_chunk(
|
||||
Note: event_date parameter is kept for backward compatibility but not used in prompt.
|
||||
The LLM extracts temporal information from the context string instead.
|
||||
"""
|
||||
memory_bank_context = f"\n- Your name: {agent_name}" if agent_name and extract_opinions else ""
|
||||
import logging
|
||||
|
||||
# Determine which fact types to extract based on the flag
|
||||
from openai import BadRequestError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Determine which fact types to extract
|
||||
# Note: We use "assistant" in the prompt but convert to "bank" for storage
|
||||
if extract_opinions:
|
||||
# Opinion extraction uses a separate prompt (not this one)
|
||||
fact_types_instruction = "Extract ONLY 'opinion' type facts (formed opinions, beliefs, and perspectives). DO NOT extract 'world' or 'assistant' facts."
|
||||
else:
|
||||
fact_types_instruction = (
|
||||
"Extract ONLY 'world' and 'assistant' type facts. DO NOT extract opinions - those are extracted separately."
|
||||
)
|
||||
fact_types_instruction = "Extract ONLY 'world' and 'assistant' type facts."
|
||||
|
||||
# Check config for extraction mode and causal link extraction
|
||||
config = get_config()
|
||||
@@ -704,13 +726,27 @@ async def _extract_facts_from_chunk(
|
||||
extract_causal_links = config.retain_extract_causal_links
|
||||
|
||||
# Select base prompt based on extraction mode
|
||||
if extraction_mode == "verbose":
|
||||
if extraction_mode == "custom":
|
||||
# Custom mode: inject user-provided guidelines
|
||||
if not config.retain_custom_instructions:
|
||||
logger.warning(
|
||||
"extraction_mode='custom' but HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS not set. "
|
||||
"Falling back to 'concise' mode."
|
||||
)
|
||||
base_prompt = CONCISE_FACT_EXTRACTION_PROMPT
|
||||
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
|
||||
else:
|
||||
base_prompt = CUSTOM_FACT_EXTRACTION_PROMPT
|
||||
prompt = base_prompt.format(
|
||||
fact_types_instruction=fact_types_instruction,
|
||||
custom_instructions=config.retain_custom_instructions,
|
||||
)
|
||||
elif extraction_mode == "verbose":
|
||||
base_prompt = VERBOSE_FACT_EXTRACTION_PROMPT
|
||||
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
|
||||
else:
|
||||
base_prompt = CONCISE_FACT_EXTRACTION_PROMPT
|
||||
|
||||
# Format the prompt with fact types instruction
|
||||
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
|
||||
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
|
||||
|
||||
# Build the full prompt with or without causal relationships section
|
||||
# Select appropriate response schema based on extraction mode and causal links
|
||||
@@ -723,12 +759,6 @@ async def _extract_facts_from_chunk(
|
||||
else:
|
||||
response_schema = FactExtractionResponseNoCausal
|
||||
|
||||
import logging
|
||||
|
||||
from openai import BadRequestError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Retry logic for JSON validation errors
|
||||
max_retries = 2
|
||||
last_error = None
|
||||
@@ -739,9 +769,12 @@ async def _extract_facts_from_chunk(
|
||||
|
||||
# Build user message with metadata and chunk content in a clear format
|
||||
# Format event_date with day of week for better temporal reasoning
|
||||
# Handle both datetime objects and ISO string formats (from deserialized async tasks)
|
||||
from .orchestrator import parse_datetime_flexible
|
||||
|
||||
event_date = parse_datetime_flexible(event_date)
|
||||
event_date_formatted = event_date.strftime("%A, %B %d, %Y") # e.g., "Monday, June 10, 2024"
|
||||
user_message = f"""Extract facts from the following text chunk.
|
||||
{memory_bank_context}
|
||||
|
||||
Chunk: {chunk_index + 1}/{total_chunks}
|
||||
Event Date: {event_date_formatted} ({event_date.isoformat()})
|
||||
@@ -753,12 +786,28 @@ Text:
|
||||
usage = TokenUsage() # Track cumulative usage across retries
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
# Use retain-specific overrides if set, otherwise fall back to global LLM config
|
||||
max_retries = (
|
||||
config.retain_llm_max_retries if config.retain_llm_max_retries is not None else config.llm_max_retries
|
||||
)
|
||||
initial_backoff = (
|
||||
config.retain_llm_initial_backoff
|
||||
if config.retain_llm_initial_backoff is not None
|
||||
else config.llm_initial_backoff
|
||||
)
|
||||
max_backoff = (
|
||||
config.retain_llm_max_backoff if config.retain_llm_max_backoff is not None else config.llm_max_backoff
|
||||
)
|
||||
|
||||
extraction_response_json, call_usage = await llm_config.call(
|
||||
messages=[{"role": "system", "content": prompt}, {"role": "user", "content": user_message}],
|
||||
response_format=response_schema,
|
||||
scope="memory_extract_facts",
|
||||
temperature=0.1,
|
||||
max_completion_tokens=config.retain_max_completion_tokens,
|
||||
max_retries=max_retries,
|
||||
initial_backoff=initial_backoff,
|
||||
max_backoff=max_backoff,
|
||||
skip_validation=True, # Get raw JSON, we'll validate leniently
|
||||
return_usage=True,
|
||||
)
|
||||
@@ -823,7 +872,8 @@ Text:
|
||||
|
||||
# Critical field: fact_type
|
||||
# LLM uses "assistant" but we convert to "experience" for storage
|
||||
fact_type = llm_fact.get("fact_type")
|
||||
original_fact_type = llm_fact.get("fact_type")
|
||||
fact_type = original_fact_type
|
||||
|
||||
# Convert "assistant" → "experience" for storage
|
||||
if fact_type == "assistant":
|
||||
@@ -840,7 +890,10 @@ Text:
|
||||
else:
|
||||
# Default to 'world' if we can't determine
|
||||
fact_type = "world"
|
||||
logger.warning(f"Fact {i}: defaulting to fact_type='world'")
|
||||
logger.warning(
|
||||
f"Fact {i}: defaulting to fact_type='world' "
|
||||
f"(original fact_type={original_fact_type!r}, fact_kind={fact_kind!r})"
|
||||
)
|
||||
|
||||
# Get fact_kind for temporal handling (but don't store it)
|
||||
fact_kind = llm_fact.get("fact_kind", "conversation")
|
||||
@@ -958,6 +1011,29 @@ Text:
|
||||
|
||||
except BadRequestError as e:
|
||||
last_error = e
|
||||
error_str = str(e).lower()
|
||||
|
||||
# Check if error is related to max_tokens/completion_tokens not being supported
|
||||
if any(
|
||||
keyword in error_str
|
||||
for keyword in [
|
||||
"max_tokens",
|
||||
"max_completion_tokens",
|
||||
"maximum context",
|
||||
"token limit",
|
||||
"context length",
|
||||
]
|
||||
):
|
||||
# Provide helpful error message with configuration suggestions
|
||||
raise ValueError(
|
||||
f"Model does not support the required output token limit.\n\n"
|
||||
f"The model '{llm_config.model}' (provider: {llm_config.provider}) failed with: {e}\n\n"
|
||||
f"You have two options to fix this:\n"
|
||||
f" 1. Use a different model that supports at least {config.retain_max_completion_tokens} output tokens\n"
|
||||
f" 2. Decrease HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS to a value your model supports\n"
|
||||
f" (current value: {config.retain_max_completion_tokens}, must be > RETAIN_CHUNK_SIZE={config.retain_chunk_size})"
|
||||
) from e
|
||||
|
||||
if "json_validate_failed" in str(e):
|
||||
logger.warning(
|
||||
f" [1.3.{chunk_index + 1}] Attempt {attempt + 1}/{max_retries} failed with JSON validation error: {e}"
|
||||
@@ -980,7 +1056,6 @@ async def _extract_facts_with_auto_split(
|
||||
context: str,
|
||||
llm_config: LLMConfig,
|
||||
agent_name: str = None,
|
||||
extract_opinions: bool = False,
|
||||
) -> tuple[list[dict[str, str]], TokenUsage]:
|
||||
"""
|
||||
Extract facts from a chunk with automatic splitting if output exceeds token limits.
|
||||
@@ -996,7 +1071,6 @@ async def _extract_facts_with_auto_split(
|
||||
context: Context about the conversation/document
|
||||
llm_config: LLM configuration to use
|
||||
agent_name: Optional agent name (memory owner)
|
||||
extract_opinions: If True, extract ONLY opinions. If False, extract world and agent facts (no opinions)
|
||||
|
||||
Returns:
|
||||
Tuple of (facts list, token usage) extracted from the chunk (possibly from sub-chunks)
|
||||
@@ -1015,7 +1089,6 @@ async def _extract_facts_with_auto_split(
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions,
|
||||
)
|
||||
except OutputTooLongError:
|
||||
# Output exceeded token limits - split the chunk in half and retry
|
||||
@@ -1060,7 +1133,6 @@ async def _extract_facts_with_auto_split(
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions,
|
||||
),
|
||||
_extract_facts_with_auto_split(
|
||||
chunk=second_half,
|
||||
@@ -1070,7 +1142,6 @@ async def _extract_facts_with_auto_split(
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions,
|
||||
),
|
||||
]
|
||||
|
||||
@@ -1094,7 +1165,6 @@ async def extract_facts_from_text(
|
||||
llm_config: LLMConfig,
|
||||
agent_name: str,
|
||||
context: str = "",
|
||||
extract_opinions: bool = False,
|
||||
) -> tuple[list[Fact], list[tuple[str, int]], TokenUsage]:
|
||||
"""
|
||||
Extract semantic facts from conversational or narrative text using LLM.
|
||||
@@ -1111,7 +1181,6 @@ async def extract_facts_from_text(
|
||||
context: Context about the conversation/document
|
||||
llm_config: LLM configuration to use
|
||||
agent_name: Agent name (memory owner)
|
||||
extract_opinions: If True, extract ONLY opinions. If False, extract world and bank facts (no opinions)
|
||||
|
||||
Returns:
|
||||
Tuple of (facts, chunks, usage) where:
|
||||
@@ -1139,7 +1208,6 @@ async def extract_facts_from_text(
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions,
|
||||
)
|
||||
for i, chunk in enumerate(chunks)
|
||||
]
|
||||
@@ -1171,7 +1239,7 @@ SECONDS_PER_FACT = 10
|
||||
|
||||
|
||||
async def extract_facts_from_contents(
|
||||
contents: list[RetainContent], llm_config, agent_name: str, extract_opinions: bool = False
|
||||
contents: list[RetainContent], llm_config, agent_name: str
|
||||
) -> tuple[list[ExtractedFactType], list[ChunkMetadata], TokenUsage]:
|
||||
"""
|
||||
Extract facts from multiple content items in parallel.
|
||||
@@ -1186,7 +1254,6 @@ async def extract_facts_from_contents(
|
||||
contents: List of RetainContent objects to process
|
||||
llm_config: LLM configuration for fact extraction
|
||||
agent_name: Name of the agent (for agent-related fact detection)
|
||||
extract_opinions: If True, extract only opinions; otherwise world/bank facts
|
||||
|
||||
Returns:
|
||||
Tuple of (extracted_facts, chunks_metadata, usage)
|
||||
@@ -1205,7 +1272,6 @@ async def extract_facts_from_contents(
|
||||
context=item.context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions,
|
||||
)
|
||||
fact_extraction_tasks.append(task)
|
||||
|
||||
@@ -1310,31 +1376,26 @@ def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> list[C
|
||||
|
||||
def _add_temporal_offsets(facts: list[ExtractedFactType], contents: list[RetainContent]) -> None:
|
||||
"""
|
||||
Add time offsets to preserve fact ordering within each content.
|
||||
Add time offsets to preserve fact ordering across all contents.
|
||||
|
||||
This allows retrieval to distinguish between facts that happened earlier vs later
|
||||
in the same conversation, even when the base event_date is the same.
|
||||
This allows retrieval to distinguish between facts from different documents/conversations
|
||||
even when they have the same base event_date, and also between facts within the same
|
||||
conversation.
|
||||
|
||||
Uses absolute position across all facts to ensure unique timestamps.
|
||||
|
||||
Modifies facts in place.
|
||||
"""
|
||||
# Group facts by content_index
|
||||
current_content_idx = 0
|
||||
content_fact_start = 0
|
||||
from .orchestrator import parse_datetime_flexible
|
||||
|
||||
for i, fact in enumerate(facts):
|
||||
if fact.content_index != current_content_idx:
|
||||
# Moved to next content
|
||||
current_content_idx = fact.content_index
|
||||
content_fact_start = i
|
||||
# Use absolute position across all facts to ensure uniqueness across different contents
|
||||
offset = timedelta(seconds=i * SECONDS_PER_FACT)
|
||||
|
||||
# Calculate position within this content
|
||||
fact_position = i - content_fact_start
|
||||
offset = timedelta(seconds=fact_position * SECONDS_PER_FACT)
|
||||
|
||||
# Apply offset to all temporal fields
|
||||
# Apply offset to all temporal fields (handle both datetime objects and ISO strings)
|
||||
if fact.occurred_start:
|
||||
fact.occurred_start = fact.occurred_start + offset
|
||||
fact.occurred_start = parse_datetime_flexible(fact.occurred_start) + offset
|
||||
if fact.occurred_end:
|
||||
fact.occurred_end = fact.occurred_end + offset
|
||||
fact.occurred_end = parse_datetime_flexible(fact.occurred_end) + offset
|
||||
if fact.mentioned_at:
|
||||
fact.mentioned_at = fact.mentioned_at + offset
|
||||
fact.mentioned_at = parse_datetime_flexible(fact.mentioned_at) + offset
|
||||
|
||||
@@ -8,6 +8,7 @@ import json
|
||||
import logging
|
||||
|
||||
from ..memory_engine import fq_table
|
||||
from .fact_extraction import _sanitize_text
|
||||
from .types import ProcessedFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -41,14 +42,13 @@ async def insert_facts_batch(
|
||||
contexts = []
|
||||
fact_types = []
|
||||
confidence_scores = []
|
||||
access_counts = []
|
||||
metadata_jsons = []
|
||||
chunk_ids = []
|
||||
document_ids = []
|
||||
tags_list = []
|
||||
|
||||
for fact in facts:
|
||||
fact_texts.append(fact.fact_text)
|
||||
fact_texts.append(_sanitize_text(fact.fact_text))
|
||||
# Convert embedding to string for asyncpg vector type
|
||||
embeddings.append(str(fact.embedding))
|
||||
# event_date: Use occurred_start if available, otherwise use mentioned_at
|
||||
@@ -57,11 +57,10 @@ async def insert_facts_batch(
|
||||
occurred_starts.append(fact.occurred_start)
|
||||
occurred_ends.append(fact.occurred_end)
|
||||
mentioned_ats.append(fact.mentioned_at)
|
||||
contexts.append(fact.context)
|
||||
contexts.append(_sanitize_text(fact.context))
|
||||
fact_types.append(fact.fact_type)
|
||||
# confidence_score is only for opinion facts
|
||||
confidence_scores.append(1.0 if fact.fact_type == "opinion" else None)
|
||||
access_counts.append(0) # Initial access count
|
||||
metadata_jsons.append(json.dumps(fact.metadata))
|
||||
chunk_ids.append(fact.chunk_id)
|
||||
# Use per-fact document_id if available, otherwise fallback to batch-level document_id
|
||||
@@ -76,16 +75,16 @@ async def insert_facts_batch(
|
||||
WITH input_data AS (
|
||||
SELECT * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[], $15::jsonb[]
|
||||
$8::text[], $9::text[], $10::float[], $11::jsonb[], $12::text[], $13::text[], $14::jsonb[]
|
||||
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags_json)
|
||||
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags_json)
|
||||
)
|
||||
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags)
|
||||
context, fact_type, confidence_score, metadata, chunk_id, document_id, tags)
|
||||
SELECT
|
||||
$1,
|
||||
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id,
|
||||
context, fact_type, confidence_score, metadata, chunk_id, document_id,
|
||||
COALESCE(
|
||||
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
|
||||
'{{}}'::varchar[]
|
||||
@@ -103,7 +102,6 @@ async def insert_facts_batch(
|
||||
contexts,
|
||||
fact_types,
|
||||
confidence_scores,
|
||||
access_counts,
|
||||
metadata_jsons,
|
||||
chunk_ids,
|
||||
document_ids,
|
||||
@@ -126,7 +124,7 @@ async def ensure_bank_exists(conn, bank_id: str) -> None:
|
||||
"""
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("banks")} (bank_id, disposition, background)
|
||||
INSERT INTO {fq_table("banks")} (bank_id, disposition, mission)
|
||||
VALUES ($1, $2::jsonb, $3)
|
||||
ON CONFLICT (bank_id) DO UPDATE
|
||||
SET updated_at = NOW()
|
||||
@@ -160,7 +158,8 @@ async def handle_document_tracking(
|
||||
"""
|
||||
import hashlib
|
||||
|
||||
# Calculate content hash
|
||||
# Sanitize and calculate content hash
|
||||
combined_content = _sanitize_text(combined_content) or ""
|
||||
content_hash = hashlib.sha256(combined_content.encode()).hexdigest()
|
||||
|
||||
# Always delete old document first if it exists (cascades to units and links)
|
||||
|
||||
@@ -754,17 +754,14 @@ async def create_causal_links_batch(
|
||||
causal_relations_per_fact: List of causal relations for each fact.
|
||||
Each element is a list of dicts with:
|
||||
- target_fact_index: Index into unit_ids for the target fact
|
||||
- relation_type: "causes", "caused_by", "enables", or "prevents"
|
||||
- relation_type: "caused_by"
|
||||
- strength: Float in [0.0, 1.0] representing relationship strength
|
||||
|
||||
Returns:
|
||||
Number of causal links created
|
||||
|
||||
Causal link types:
|
||||
- "causes": This fact directly causes the target fact (forward causation)
|
||||
- "caused_by": This fact was caused by the target fact (backward causation)
|
||||
- "enables": This fact enables/allows the target fact (enablement)
|
||||
- "prevents": This fact prevents/blocks the target fact (prevention)
|
||||
Causal link type:
|
||||
- "caused_by": This fact was caused by the target fact
|
||||
"""
|
||||
if not unit_ids or not causal_relations_per_fact:
|
||||
return 0
|
||||
@@ -787,8 +784,8 @@ async def create_causal_links_batch(
|
||||
relation_type = relation["relation_type"]
|
||||
strength = relation.get("strength", 1.0)
|
||||
|
||||
# Validate relation_type - must match database constraint
|
||||
valid_types = {"causes", "caused_by", "enables", "prevents"}
|
||||
# Validate relation_type - only "caused_by" is supported (DB constraint)
|
||||
valid_types = {"caused_by"}
|
||||
if relation_type not in valid_types:
|
||||
logger.error(
|
||||
f"Invalid relation_type '{relation_type}' (type: {type(relation_type).__name__}) "
|
||||
|
||||
@@ -1,254 +0,0 @@
|
||||
"""
|
||||
Observation regeneration for retain pipeline.
|
||||
|
||||
Regenerates entity observations as part of the retain transaction.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from ...config import get_config
|
||||
from ..memory_engine import fq_table
|
||||
from ..search import observation_utils
|
||||
from . import embedding_utils
|
||||
from .types import EntityLink
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def utcnow():
|
||||
"""Get current UTC time."""
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
# Simple dataclass-like container for facts (avoid importing from memory_engine)
|
||||
class MemoryFactForObservation:
|
||||
def __init__(self, id: str, text: str, fact_type: str, context: str, occurred_start: str | None):
|
||||
self.id = id
|
||||
self.text = text
|
||||
self.fact_type = fact_type
|
||||
self.context = context
|
||||
self.occurred_start = occurred_start
|
||||
|
||||
|
||||
async def regenerate_observations_batch(
|
||||
conn, embeddings_model, llm_config, bank_id: str, entity_links: list[EntityLink], log_buffer: list[str] = None
|
||||
) -> None:
|
||||
"""
|
||||
Regenerate observations for top entities in this batch.
|
||||
|
||||
Called INSIDE the retain transaction for atomicity - if observations
|
||||
fail, the entire retain batch is rolled back.
|
||||
|
||||
Args:
|
||||
conn: Database connection (from the retain transaction)
|
||||
embeddings_model: Embeddings model for generating observation embeddings
|
||||
llm_config: LLM configuration for observation extraction
|
||||
bank_id: Bank identifier
|
||||
entity_links: Entity links from this batch
|
||||
log_buffer: Optional log buffer for timing
|
||||
"""
|
||||
config = get_config()
|
||||
TOP_N_ENTITIES = config.observation_top_entities
|
||||
MIN_FACTS_THRESHOLD = config.observation_min_facts
|
||||
|
||||
if not entity_links:
|
||||
return
|
||||
|
||||
# Count mentions per entity in this batch
|
||||
entity_mention_counts: dict[str, int] = {}
|
||||
for link in entity_links:
|
||||
if link.entity_id:
|
||||
entity_id = str(link.entity_id)
|
||||
entity_mention_counts[entity_id] = entity_mention_counts.get(entity_id, 0) + 1
|
||||
|
||||
if not entity_mention_counts:
|
||||
return
|
||||
|
||||
# Sort by mention count descending and take top N
|
||||
sorted_entities = sorted(entity_mention_counts.items(), key=lambda x: x[1], reverse=True)
|
||||
entities_to_process = [e[0] for e in sorted_entities[:TOP_N_ENTITIES]]
|
||||
|
||||
obs_start = time.time()
|
||||
|
||||
# Convert to UUIDs
|
||||
entity_uuids = [uuid.UUID(eid) if isinstance(eid, str) else eid for eid in entities_to_process]
|
||||
|
||||
# Batch query for entity names
|
||||
entity_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, canonical_name FROM {fq_table("entities")}
|
||||
WHERE id = ANY($1) AND bank_id = $2
|
||||
""",
|
||||
entity_uuids,
|
||||
bank_id,
|
||||
)
|
||||
entity_names = {row["id"]: row["canonical_name"] for row in entity_rows}
|
||||
|
||||
# Batch query for fact counts
|
||||
fact_counts = await conn.fetch(
|
||||
f"""
|
||||
SELECT ue.entity_id, COUNT(*) as cnt
|
||||
FROM {fq_table("unit_entities")} ue
|
||||
JOIN {fq_table("memory_units")} mu ON ue.unit_id = mu.id
|
||||
WHERE ue.entity_id = ANY($1) AND mu.bank_id = $2
|
||||
GROUP BY ue.entity_id
|
||||
""",
|
||||
entity_uuids,
|
||||
bank_id,
|
||||
)
|
||||
entity_fact_counts = {row["entity_id"]: row["cnt"] for row in fact_counts}
|
||||
|
||||
# Filter entities that meet the threshold
|
||||
entities_with_names = []
|
||||
for entity_id in entities_to_process:
|
||||
entity_uuid = uuid.UUID(entity_id) if isinstance(entity_id, str) else entity_id
|
||||
if entity_uuid not in entity_names:
|
||||
continue
|
||||
fact_count = entity_fact_counts.get(entity_uuid, 0)
|
||||
if fact_count >= MIN_FACTS_THRESHOLD:
|
||||
entities_with_names.append((entity_id, entity_names[entity_uuid]))
|
||||
|
||||
if not entities_with_names:
|
||||
return
|
||||
|
||||
# Process entities SEQUENTIALLY (asyncpg doesn't allow concurrent queries on same connection)
|
||||
# We must use the same connection to stay in the retain transaction
|
||||
total_observations = 0
|
||||
|
||||
for entity_id, entity_name in entities_with_names:
|
||||
try:
|
||||
obs_ids = await _regenerate_entity_observations(
|
||||
conn, embeddings_model, llm_config, bank_id, entity_id, entity_name
|
||||
)
|
||||
total_observations += len(obs_ids)
|
||||
except Exception as e:
|
||||
logger.error(f"[OBSERVATIONS] Error processing entity {entity_id}: {e}")
|
||||
|
||||
obs_time = time.time() - obs_start
|
||||
if log_buffer is not None:
|
||||
log_buffer.append(
|
||||
f"[11] Observations: {total_observations} observations for {len(entities_with_names)} entities in {obs_time:.3f}s"
|
||||
)
|
||||
|
||||
|
||||
async def _regenerate_entity_observations(
|
||||
conn, embeddings_model, llm_config, bank_id: str, entity_id: str, entity_name: str
|
||||
) -> list[str]:
|
||||
"""
|
||||
Regenerate observations for a single entity.
|
||||
|
||||
Uses the provided connection (part of retain transaction).
|
||||
|
||||
Args:
|
||||
conn: Database connection (from the retain transaction)
|
||||
embeddings_model: Embeddings model
|
||||
llm_config: LLM configuration
|
||||
bank_id: Bank identifier
|
||||
entity_id: Entity UUID
|
||||
entity_name: Canonical name of the entity
|
||||
|
||||
Returns:
|
||||
List of created observation IDs
|
||||
"""
|
||||
entity_uuid = uuid.UUID(entity_id) if isinstance(entity_id, str) else entity_id
|
||||
|
||||
# Get all facts mentioning this entity (exclude observations themselves)
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.fact_type
|
||||
FROM {fq_table("memory_units")} mu
|
||||
JOIN {fq_table("unit_entities")} ue ON mu.id = ue.unit_id
|
||||
WHERE mu.bank_id = $1
|
||||
AND ue.entity_id = $2
|
||||
AND mu.fact_type IN ('world', 'experience')
|
||||
ORDER BY mu.occurred_start DESC
|
||||
LIMIT 50
|
||||
""",
|
||||
bank_id,
|
||||
entity_uuid,
|
||||
)
|
||||
|
||||
if not rows:
|
||||
return []
|
||||
|
||||
# Convert to fact objects for observation extraction
|
||||
facts = []
|
||||
for row in rows:
|
||||
occurred_start = row["occurred_start"].isoformat() if row["occurred_start"] else None
|
||||
facts.append(
|
||||
MemoryFactForObservation(
|
||||
id=str(row["id"]),
|
||||
text=row["text"],
|
||||
fact_type=row["fact_type"],
|
||||
context=row["context"],
|
||||
occurred_start=occurred_start,
|
||||
)
|
||||
)
|
||||
|
||||
# Extract observations using LLM
|
||||
observations = await observation_utils.extract_observations_from_facts(llm_config, entity_name, facts)
|
||||
|
||||
if not observations:
|
||||
return []
|
||||
|
||||
# Delete old observations for this entity
|
||||
await conn.execute(
|
||||
f"""
|
||||
DELETE FROM {fq_table("memory_units")}
|
||||
WHERE id IN (
|
||||
SELECT mu.id
|
||||
FROM {fq_table("memory_units")} mu
|
||||
JOIN {fq_table("unit_entities")} ue ON mu.id = ue.unit_id
|
||||
WHERE mu.bank_id = $1
|
||||
AND mu.fact_type = 'observation'
|
||||
AND ue.entity_id = $2
|
||||
)
|
||||
""",
|
||||
bank_id,
|
||||
entity_uuid,
|
||||
)
|
||||
|
||||
# Generate embeddings for new observations
|
||||
embeddings = await embedding_utils.generate_embeddings_batch(embeddings_model, observations)
|
||||
|
||||
# Insert new observations
|
||||
current_time = utcnow()
|
||||
created_ids = []
|
||||
|
||||
for obs_text, embedding in zip(observations, embeddings):
|
||||
result = await conn.fetchrow(
|
||||
f"""
|
||||
INSERT INTO {fq_table("memory_units")} (
|
||||
bank_id, text, embedding, context, event_date,
|
||||
occurred_start, occurred_end, mentioned_at,
|
||||
fact_type, access_count
|
||||
)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, 'observation', 0)
|
||||
RETURNING id
|
||||
""",
|
||||
bank_id,
|
||||
obs_text,
|
||||
str(embedding),
|
||||
f"observation about {entity_name}",
|
||||
current_time,
|
||||
current_time,
|
||||
current_time,
|
||||
current_time,
|
||||
)
|
||||
obs_id = str(result["id"])
|
||||
created_ids.append(obs_id)
|
||||
|
||||
# Link observation to entity
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("unit_entities")} (unit_id, entity_id)
|
||||
VALUES ($1, $2)
|
||||
""",
|
||||
uuid.UUID(obs_id),
|
||||
entity_uuid,
|
||||
)
|
||||
|
||||
return created_ids
|
||||
@@ -8,8 +8,8 @@ import logging
|
||||
import time
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from ...config import get_config
|
||||
from ..db_utils import acquire_with_retry
|
||||
from . import bank_utils
|
||||
|
||||
@@ -19,6 +19,39 @@ def utcnow():
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
def parse_datetime_flexible(value: Any) -> datetime:
|
||||
"""
|
||||
Parse a datetime value that could be either a datetime object or an ISO string.
|
||||
|
||||
This handles datetime values from both direct Python calls and deserialized JSON
|
||||
(where datetime objects are serialized as ISO strings).
|
||||
|
||||
Args:
|
||||
value: Either a datetime object or an ISO format string
|
||||
|
||||
Returns:
|
||||
datetime object (timezone-aware)
|
||||
|
||||
Raises:
|
||||
TypeError: If value is neither datetime nor string
|
||||
ValueError: If string is not a valid ISO datetime
|
||||
"""
|
||||
if isinstance(value, datetime):
|
||||
# Ensure timezone-aware
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=UTC)
|
||||
return value
|
||||
elif isinstance(value, str):
|
||||
# Parse ISO format string (handles both 'Z' and '+00:00' timezone formats)
|
||||
dt = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
# Ensure timezone-aware
|
||||
if dt.tzinfo is None:
|
||||
return dt.replace(tzinfo=UTC)
|
||||
return dt
|
||||
else:
|
||||
raise TypeError(f"Expected datetime or string, got {type(value).__name__}")
|
||||
|
||||
|
||||
from ..response_models import TokenUsage
|
||||
from . import (
|
||||
chunk_storage,
|
||||
@@ -28,9 +61,8 @@ from . import (
|
||||
fact_extraction,
|
||||
fact_storage,
|
||||
link_creation,
|
||||
observation_regeneration,
|
||||
)
|
||||
from .types import ExtractedFact, ProcessedFact, RetainContent, RetainContentDict
|
||||
from .types import EntityLink, ExtractedFact, ProcessedFact, RetainContent, RetainContentDict
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -40,7 +72,6 @@ async def retain_batch(
|
||||
embeddings_model,
|
||||
llm_config,
|
||||
entity_resolver,
|
||||
task_backend,
|
||||
format_date_fn,
|
||||
duplicate_checker_fn,
|
||||
bank_id: str,
|
||||
@@ -59,7 +90,6 @@ async def retain_batch(
|
||||
embeddings_model: Embeddings model for generating embeddings
|
||||
llm_config: LLM configuration for fact extraction
|
||||
entity_resolver: Entity resolver for entity processing
|
||||
task_backend: Task backend for background jobs
|
||||
format_date_fn: Function to format datetime to readable string
|
||||
duplicate_checker_fn: Function to check for duplicate facts
|
||||
bank_id: Bank identifier
|
||||
@@ -93,10 +123,18 @@ async def retain_batch(
|
||||
# Merge item-level tags with document-level tags
|
||||
item_tags = item.get("tags", []) or []
|
||||
merged_tags = list(set(item_tags + (document_tags or [])))
|
||||
|
||||
# Handle event_date: parse flexibly (handles both datetime objects and ISO strings)
|
||||
event_date_value = item.get("event_date")
|
||||
if event_date_value:
|
||||
event_date_value = parse_datetime_flexible(event_date_value)
|
||||
else:
|
||||
event_date_value = utcnow()
|
||||
|
||||
content = RetainContent(
|
||||
content=item["content"],
|
||||
context=item.get("context", ""),
|
||||
event_date=item.get("event_date") or utcnow(),
|
||||
event_date=event_date_value,
|
||||
metadata=item.get("metadata", {}),
|
||||
entities=item.get("entities", []),
|
||||
tags=merged_tags,
|
||||
@@ -105,11 +143,8 @@ async def retain_batch(
|
||||
|
||||
# Step 1: Extract facts from all contents
|
||||
step_start = time.time()
|
||||
extract_opinions = fact_type_override == "opinion"
|
||||
|
||||
extracted_facts, chunks, usage = await fact_extraction.extract_facts_from_contents(
|
||||
contents, llm_config, agent_name, extract_opinions
|
||||
)
|
||||
extracted_facts, chunks, usage = await fact_extraction.extract_facts_from_contents(contents, llm_config, agent_name)
|
||||
log_buffer.append(
|
||||
f"[1] Extract facts: {len(extracted_facts)} facts, {len(chunks)} chunks from {len(contents)} contents in {time.time() - step_start:.3f}s"
|
||||
)
|
||||
@@ -123,6 +158,13 @@ async def retain_batch(
|
||||
# Handle document tracking even with no facts
|
||||
if document_id:
|
||||
combined_content = "\n".join([c.get("content", "") for c in contents_dicts])
|
||||
# Collect tags from all content items and merge with document_tags
|
||||
all_tags = set(document_tags or [])
|
||||
for item in contents_dicts:
|
||||
item_tags = item.get("tags", []) or []
|
||||
all_tags.update(item_tags)
|
||||
merged_tags = list(all_tags)
|
||||
|
||||
retain_params = {}
|
||||
if contents_dicts:
|
||||
first_item = contents_dicts[0]
|
||||
@@ -137,7 +179,7 @@ async def retain_batch(
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, merged_tags
|
||||
)
|
||||
else:
|
||||
# Check for per-item document_ids
|
||||
@@ -151,6 +193,13 @@ async def retain_batch(
|
||||
|
||||
for doc_id, doc_contents in contents_by_doc.items():
|
||||
combined_content = "\n".join([c.get("content", "") for _, c in doc_contents])
|
||||
# Collect tags from all content items for this document and merge with document_tags
|
||||
all_tags = set(document_tags or [])
|
||||
for _, item in doc_contents:
|
||||
item_tags = item.get("tags", []) or []
|
||||
all_tags.update(item_tags)
|
||||
merged_tags = list(all_tags)
|
||||
|
||||
retain_params = {}
|
||||
if doc_contents:
|
||||
first_item = doc_contents[0][1]
|
||||
@@ -165,7 +214,7 @@ async def retain_batch(
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params, merged_tags
|
||||
)
|
||||
|
||||
total_time = time.time() - start_time
|
||||
@@ -217,6 +266,13 @@ async def retain_batch(
|
||||
# Legacy: single document_id parameter
|
||||
combined_content = "\n".join([c.get("content", "") for c in contents_dicts])
|
||||
retain_params = {}
|
||||
# Collect tags from all content items and merge with document_tags
|
||||
all_tags = set(document_tags or [])
|
||||
for item in contents_dicts:
|
||||
item_tags = item.get("tags", []) or []
|
||||
all_tags.update(item_tags)
|
||||
merged_tags = list(all_tags)
|
||||
|
||||
if contents_dicts:
|
||||
first_item = contents_dicts[0]
|
||||
if first_item.get("context"):
|
||||
@@ -231,7 +287,7 @@ async def retain_batch(
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, merged_tags
|
||||
)
|
||||
document_ids_added.append(document_id)
|
||||
doc_id_mapping[None] = document_id # For backwards compatibility
|
||||
@@ -259,6 +315,13 @@ async def retain_batch(
|
||||
# Combine content for this document
|
||||
combined_content = "\n".join([c.get("content", "") for _, c in doc_contents])
|
||||
|
||||
# Collect tags from all content items for this document and merge with document_tags
|
||||
all_tags = set(document_tags or [])
|
||||
for _, item in doc_contents:
|
||||
item_tags = item.get("tags", []) or []
|
||||
all_tags.update(item_tags)
|
||||
merged_tags = list(all_tags)
|
||||
|
||||
# Extract retain params from first content item
|
||||
retain_params = {}
|
||||
if doc_contents:
|
||||
@@ -281,7 +344,7 @@ async def retain_batch(
|
||||
combined_content,
|
||||
is_first_batch,
|
||||
retain_params,
|
||||
document_tags,
|
||||
merged_tags,
|
||||
)
|
||||
document_ids_added.append(actual_doc_id)
|
||||
|
||||
@@ -408,27 +471,9 @@ async def retain_batch(
|
||||
causal_link_count = await link_creation.create_causal_links_batch(conn, unit_ids, non_duplicate_facts)
|
||||
log_buffer.append(f"[10] Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s")
|
||||
|
||||
# Regenerate observations - sync (in transaction) or async (background task)
|
||||
config = get_config()
|
||||
if config.retain_observations_async:
|
||||
# Queue for async processing after transaction commits
|
||||
entity_ids_for_async = list(set(link.entity_id for link in entity_links)) if entity_links else []
|
||||
log_buffer.append(
|
||||
f"[11] Observations: queued {len(entity_ids_for_async)} entities for async processing"
|
||||
)
|
||||
else:
|
||||
# Run synchronously inside transaction for atomicity
|
||||
await observation_regeneration.regenerate_observations_batch(
|
||||
conn, embeddings_model, llm_config, bank_id, entity_links, log_buffer
|
||||
)
|
||||
entity_ids_for_async = []
|
||||
|
||||
# Map results back to original content items
|
||||
result_unit_ids = _map_results_to_contents(contents, extracted_facts, is_duplicate_flags, unit_ids)
|
||||
|
||||
# Trigger background tasks AFTER transaction commits
|
||||
await _trigger_background_tasks(task_backend, bank_id, unit_ids, non_duplicate_facts, entity_ids_for_async)
|
||||
|
||||
# Log final summary
|
||||
total_time = time.time() - start_time
|
||||
log_buffer.append(f"{'=' * 60}")
|
||||
@@ -470,35 +515,3 @@ def _map_results_to_contents(
|
||||
result_unit_ids.append(content_unit_ids)
|
||||
|
||||
return result_unit_ids
|
||||
|
||||
|
||||
async def _trigger_background_tasks(
|
||||
task_backend,
|
||||
bank_id: str,
|
||||
unit_ids: list[str],
|
||||
facts: list[ProcessedFact],
|
||||
entity_ids_for_observations: list[str] | None = None,
|
||||
) -> None:
|
||||
"""Trigger background tasks after transaction commits."""
|
||||
# Trigger opinion reinforcement if there are entities
|
||||
fact_entities = [[e.name for e in fact.entities] for fact in facts]
|
||||
if any(fact_entities):
|
||||
await task_backend.submit_task(
|
||||
{
|
||||
"type": "reinforce_opinion",
|
||||
"bank_id": bank_id,
|
||||
"created_unit_ids": unit_ids,
|
||||
"unit_texts": [fact.fact_text for fact in facts],
|
||||
"unit_entities": fact_entities,
|
||||
}
|
||||
)
|
||||
|
||||
# Trigger observation regeneration if async mode is enabled
|
||||
if entity_ids_for_observations:
|
||||
await task_backend.submit_task(
|
||||
{
|
||||
"type": "regenerate_observations",
|
||||
"bank_id": bank_id,
|
||||
"entity_ids": entity_ids_for_observations,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -86,10 +86,10 @@ class CausalRelation:
|
||||
"""
|
||||
Causal relationship between facts.
|
||||
|
||||
Represents how one fact causes, enables, or prevents another.
|
||||
Represents how one fact was caused by another.
|
||||
"""
|
||||
|
||||
relation_type: str # "causes", "enables", "prevents", "caused_by"
|
||||
relation_type: str # "caused_by"
|
||||
target_fact_index: int # Index of the target fact in the batch
|
||||
strength: float = 1.0 # Strength of the causal relationship
|
||||
|
||||
|
||||
@@ -162,7 +162,7 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
entry_points = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
@@ -216,7 +216,7 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
neighbors = await conn.fetch(
|
||||
f"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.occurred_end,
|
||||
mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type,
|
||||
mu.mentioned_at, mu.embedding, mu.fact_type,
|
||||
mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight, ml.link_type, ml.from_unit_id
|
||||
FROM {fq_table("memory_links")} ml
|
||||
|
||||
@@ -45,7 +45,7 @@ async def _find_semantic_seeds(
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
@@ -155,7 +155,6 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
all_seeds.extend(temporal_seeds)
|
||||
|
||||
if not all_seeds:
|
||||
logger.debug("[LinkExpansion] No seeds found, returning empty results")
|
||||
return [], timings
|
||||
|
||||
seed_ids = list({s.id for s in all_seeds})
|
||||
@@ -164,36 +163,108 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
# Run entity and causal expansion sequentially on same connection
|
||||
query_start = time.time()
|
||||
|
||||
entity_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
COUNT(*)::float AS score
|
||||
FROM {fq_table("unit_entities")} seed_ue
|
||||
JOIN {fq_table("entities")} e ON seed_ue.entity_id = e.id
|
||||
JOIN {fq_table("unit_entities")} other_ue ON seed_ue.entity_id = other_ue.entity_id
|
||||
JOIN {fq_table("memory_units")} mu ON other_ue.unit_id = mu.id
|
||||
WHERE seed_ue.unit_id = ANY($1::uuid[])
|
||||
AND e.mention_count < $2
|
||||
AND mu.id != ALL($1::uuid[])
|
||||
AND mu.fact_type = $3
|
||||
GROUP BY mu.id
|
||||
ORDER BY score DESC
|
||||
LIMIT $4
|
||||
""",
|
||||
seed_ids,
|
||||
self.max_entity_frequency,
|
||||
fact_type,
|
||||
budget,
|
||||
)
|
||||
# For observations, traverse through source_memory_ids to find entity connections.
|
||||
# Observations don't have direct unit_entities - they inherit entities via their
|
||||
# source world/experience facts.
|
||||
#
|
||||
# Path: observation → source_memory_ids → world fact → entities →
|
||||
# ALL world facts with those entities → their observations (excluding seeds)
|
||||
if fact_type == "observation":
|
||||
# Debug: Check what source_memory_ids exist on seed observations
|
||||
debug_sources = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, source_memory_ids
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = ANY($1::uuid[])
|
||||
""",
|
||||
seed_ids,
|
||||
)
|
||||
source_ids_found = []
|
||||
for row in debug_sources:
|
||||
if row["source_memory_ids"]:
|
||||
source_ids_found.extend(row["source_memory_ids"])
|
||||
logger.debug(
|
||||
f"[LinkExpansion] observation graph: {len(seed_ids)} seeds, "
|
||||
f"{len(source_ids_found)} source_memory_ids found"
|
||||
)
|
||||
|
||||
entity_rows = await conn.fetch(
|
||||
f"""
|
||||
WITH seed_sources AS (
|
||||
-- Get source memory IDs from seed observations
|
||||
SELECT DISTINCT unnest(source_memory_ids) AS source_id
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = ANY($1::uuid[])
|
||||
AND source_memory_ids IS NOT NULL
|
||||
),
|
||||
source_entities AS (
|
||||
-- Get entities from those source memories (filtered by frequency)
|
||||
SELECT DISTINCT ue.entity_id
|
||||
FROM seed_sources ss
|
||||
JOIN {fq_table("unit_entities")} ue ON ss.source_id = ue.unit_id
|
||||
JOIN {fq_table("entities")} e ON ue.entity_id = e.id
|
||||
WHERE e.mention_count < $2
|
||||
),
|
||||
all_connected_sources AS (
|
||||
-- Find ALL world facts sharing those entities (don't exclude seed sources)
|
||||
-- The exclusion happens at the observation level, not the source level
|
||||
SELECT DISTINCT other_ue.unit_id AS source_id
|
||||
FROM source_entities se
|
||||
JOIN {fq_table("unit_entities")} other_ue ON se.entity_id = other_ue.entity_id
|
||||
)
|
||||
-- Find observations derived from connected source memories
|
||||
-- Only exclude the actual seed observations
|
||||
SELECT
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
COUNT(DISTINCT cs.source_id)::float AS score
|
||||
FROM all_connected_sources cs
|
||||
JOIN {fq_table("memory_units")} mu
|
||||
ON mu.source_memory_ids @> ARRAY[cs.source_id]
|
||||
WHERE mu.fact_type = 'observation'
|
||||
AND mu.id != ALL($1::uuid[])
|
||||
GROUP BY mu.id
|
||||
ORDER BY score DESC
|
||||
LIMIT $3
|
||||
""",
|
||||
seed_ids,
|
||||
self.max_entity_frequency,
|
||||
budget,
|
||||
)
|
||||
logger.debug(f"[LinkExpansion] observation graph: found {len(entity_rows)} connected observations")
|
||||
else:
|
||||
# For world/experience facts, use direct entity lookup
|
||||
entity_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
COUNT(*)::float AS score
|
||||
FROM {fq_table("unit_entities")} seed_ue
|
||||
JOIN {fq_table("entities")} e ON seed_ue.entity_id = e.id
|
||||
JOIN {fq_table("unit_entities")} other_ue ON seed_ue.entity_id = other_ue.entity_id
|
||||
JOIN {fq_table("memory_units")} mu ON other_ue.unit_id = mu.id
|
||||
WHERE seed_ue.unit_id = ANY($1::uuid[])
|
||||
AND e.mention_count < $2
|
||||
AND mu.id != ALL($1::uuid[])
|
||||
AND mu.fact_type = $3
|
||||
GROUP BY mu.id
|
||||
ORDER BY score DESC
|
||||
LIMIT $4
|
||||
""",
|
||||
seed_ids,
|
||||
self.max_entity_frequency,
|
||||
fact_type,
|
||||
budget,
|
||||
)
|
||||
|
||||
causal_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT DISTINCT ON (mu.id)
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
|
||||
mu.occurred_end, mu.mentioned_at, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight + 1.0 AS score
|
||||
FROM {fq_table("memory_links")} ml
|
||||
@@ -211,11 +282,69 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
budget,
|
||||
)
|
||||
|
||||
# Fallback: semantic/temporal/entity links from memory_links table
|
||||
# These are secondary to entity links (via unit_entities) and causal links
|
||||
# Weight is halved (0.5x) to prioritize primary link types
|
||||
# Check both directions: seeds -> others AND others -> seeds
|
||||
fallback_rows = await conn.fetch(
|
||||
f"""
|
||||
WITH outgoing AS (
|
||||
-- Links FROM seeds TO other facts
|
||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight
|
||||
FROM {fq_table("memory_links")} ml
|
||||
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
||||
WHERE ml.from_unit_id = ANY($1::uuid[])
|
||||
AND ml.link_type IN ('semantic', 'temporal', 'entity')
|
||||
AND ml.weight >= $2
|
||||
AND mu.fact_type = $3
|
||||
AND mu.id != ALL($1::uuid[])
|
||||
),
|
||||
incoming AS (
|
||||
-- Links FROM other facts TO seeds (reverse direction)
|
||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight
|
||||
FROM {fq_table("memory_links")} ml
|
||||
JOIN {fq_table("memory_units")} mu ON ml.from_unit_id = mu.id
|
||||
WHERE ml.to_unit_id = ANY($1::uuid[])
|
||||
AND ml.link_type IN ('semantic', 'temporal', 'entity')
|
||||
AND ml.weight >= $2
|
||||
AND mu.fact_type = $3
|
||||
AND mu.id != ALL($1::uuid[])
|
||||
),
|
||||
combined AS (
|
||||
SELECT * FROM outgoing
|
||||
UNION ALL
|
||||
SELECT * FROM incoming
|
||||
)
|
||||
SELECT DISTINCT ON (id)
|
||||
id, text, context, event_date, occurred_start,
|
||||
occurred_end, mentioned_at, embedding,
|
||||
fact_type, document_id, chunk_id, tags,
|
||||
(MAX(weight) * 0.5) AS score
|
||||
FROM combined
|
||||
GROUP BY id, text, context, event_date, occurred_start,
|
||||
occurred_end, mentioned_at, embedding,
|
||||
fact_type, document_id, chunk_id, tags
|
||||
ORDER BY id, score DESC
|
||||
LIMIT $4
|
||||
""",
|
||||
seed_ids,
|
||||
self.causal_weight_threshold,
|
||||
fact_type,
|
||||
budget,
|
||||
)
|
||||
|
||||
timings.edge_load_time = time.time() - query_start
|
||||
timings.db_queries = 2
|
||||
timings.edge_count = len(entity_rows) + len(causal_rows)
|
||||
timings.db_queries = 3
|
||||
timings.edge_count = len(entity_rows) + len(causal_rows) + len(fallback_rows)
|
||||
|
||||
# Merge results, taking max score per fact
|
||||
# Priority: entity links (unit_entities) > causal links > fallback links
|
||||
score_map: dict[str, float] = {}
|
||||
row_map: dict[str, dict] = {}
|
||||
|
||||
@@ -230,6 +359,12 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
if fact_id not in row_map:
|
||||
row_map[fact_id] = dict(row)
|
||||
|
||||
for row in fallback_rows:
|
||||
fact_id = str(row["id"])
|
||||
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
|
||||
if fact_id not in row_map:
|
||||
row_map[fact_id] = dict(row)
|
||||
|
||||
# Sort by score and limit
|
||||
sorted_ids = sorted(score_map.keys(), key=lambda x: score_map[x], reverse=True)[:budget]
|
||||
rows = [row_map[fact_id] for fact_id in sorted_ids]
|
||||
|
||||
@@ -449,7 +449,7 @@ async def fetch_memory_units_by_ids(
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags
|
||||
mentioned_at, embedding, fact_type, document_id, chunk_id, tags
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = ANY($1::uuid[])
|
||||
AND fact_type = $2
|
||||
|
||||
@@ -1,125 +0,0 @@
|
||||
"""
|
||||
Observation utilities for generating entity observations from facts.
|
||||
|
||||
Observations are objective facts synthesized from multiple memory facts
|
||||
about an entity, without personality influence.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..response_models import MemoryFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Observation(BaseModel):
|
||||
"""An observation about an entity."""
|
||||
|
||||
observation: str = Field(description="The observation text - a factual statement about the entity")
|
||||
|
||||
|
||||
class ObservationExtractionResponse(BaseModel):
|
||||
"""Response containing extracted observations."""
|
||||
|
||||
observations: list[Observation] = Field(default_factory=list, description="List of observations about the entity")
|
||||
|
||||
|
||||
def format_facts_for_observation_prompt(facts: list[MemoryFact]) -> str:
|
||||
"""Format facts as text for observation extraction prompt."""
|
||||
import json
|
||||
|
||||
if not facts:
|
||||
return "[]"
|
||||
formatted = []
|
||||
for fact in facts:
|
||||
fact_obj = {"text": fact.text}
|
||||
|
||||
# Add context if available
|
||||
if fact.context:
|
||||
fact_obj["context"] = fact.context
|
||||
|
||||
# Add occurred_start if available
|
||||
if fact.occurred_start:
|
||||
fact_obj["occurred_at"] = fact.occurred_start
|
||||
|
||||
formatted.append(fact_obj)
|
||||
|
||||
return json.dumps(formatted, indent=2)
|
||||
|
||||
|
||||
def build_observation_prompt(
|
||||
entity_name: str,
|
||||
facts_text: str,
|
||||
) -> str:
|
||||
"""Build the observation extraction prompt for the LLM."""
|
||||
return f"""Based on the following facts about "{entity_name}", generate a list of key observations.
|
||||
|
||||
FACTS ABOUT {entity_name.upper()}:
|
||||
{facts_text}
|
||||
|
||||
Your task: Synthesize the facts into clear, objective observations about {entity_name}.
|
||||
|
||||
GUIDELINES:
|
||||
1. Each observation should be a factual statement about {entity_name}
|
||||
2. Combine related facts into single observations where appropriate
|
||||
3. Be objective - do not add opinions, judgments, or interpretations
|
||||
4. Focus on what we KNOW about {entity_name}, not what we assume
|
||||
5. Include observations about: identity, characteristics, roles, relationships, activities
|
||||
6. Write in third person (e.g., "John is..." not "I think John is...")
|
||||
7. If there are conflicting facts, note the most recent or most supported one
|
||||
|
||||
EXAMPLES of good observations:
|
||||
- "John works at Google as a software engineer"
|
||||
- "John is detail-oriented and methodical in his approach"
|
||||
- "John collaborates frequently with Sarah on the AI project"
|
||||
- "John joined the company in 2023"
|
||||
|
||||
EXAMPLES of bad observations (avoid these):
|
||||
- "John seems like a good person" (opinion/judgment)
|
||||
- "John probably likes his job" (assumption)
|
||||
- "I believe John is reliable" (first-person opinion)
|
||||
|
||||
Generate 3-7 observations based on the available facts. If there are very few facts, generate fewer observations."""
|
||||
|
||||
|
||||
def get_observation_system_message() -> str:
|
||||
"""Get the system message for observation extraction."""
|
||||
return "You are an objective observer synthesizing facts about an entity. Generate clear, factual observations without opinions or personality influence. Be concise and accurate."
|
||||
|
||||
|
||||
async def extract_observations_from_facts(llm_config, entity_name: str, facts: list[MemoryFact]) -> list[str]:
|
||||
"""
|
||||
Extract observations from facts about an entity using LLM.
|
||||
|
||||
Args:
|
||||
llm_config: LLM configuration to use
|
||||
entity_name: Name of the entity to generate observations about
|
||||
facts: List of facts mentioning the entity
|
||||
|
||||
Returns:
|
||||
List of observation strings
|
||||
"""
|
||||
if not facts:
|
||||
return []
|
||||
|
||||
facts_text = format_facts_for_observation_prompt(facts)
|
||||
prompt = build_observation_prompt(entity_name, facts_text)
|
||||
|
||||
try:
|
||||
result = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": get_observation_system_message()},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
response_format=ObservationExtractionResponse,
|
||||
scope="memory_extract_observation",
|
||||
)
|
||||
|
||||
observations = [op.observation for op in result.observations]
|
||||
return observations
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to extract observations for {entity_name}: {str(e)}")
|
||||
return []
|
||||
@@ -116,7 +116,7 @@ async def retrieve_semantic(
|
||||
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
@@ -180,7 +180,7 @@ async def retrieve_bm25(
|
||||
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
ts_rank_cd(search_vector, to_tsquery('english', $1)) AS bm25_score
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
@@ -237,7 +237,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
WITH semantic_ranked AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity,
|
||||
NULL::float AS bm25_score,
|
||||
'semantic' AS source,
|
||||
@@ -249,7 +249,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
AND (1 - (embedding <=> $1::vector)) >= 0.3
|
||||
{tags_clause}
|
||||
)
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
similarity, bm25_score, source
|
||||
FROM semantic_ranked
|
||||
WHERE rn <= $4
|
||||
@@ -281,7 +281,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
WITH semantic_ranked AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity,
|
||||
NULL::float AS bm25_score,
|
||||
'semantic' AS source,
|
||||
@@ -294,7 +294,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
{tags_clause}
|
||||
),
|
||||
bm25_ranked AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
NULL::float AS similarity,
|
||||
ts_rank_cd(search_vector, to_tsquery('english', $5)) AS bm25_score,
|
||||
'bm25' AS source,
|
||||
@@ -306,12 +306,12 @@ async def retrieve_semantic_bm25_combined(
|
||||
{tags_clause}
|
||||
),
|
||||
semantic AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
similarity, bm25_score, source
|
||||
FROM semantic_ranked WHERE rn <= $4
|
||||
),
|
||||
bm25 AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
similarity, bm25_score, source
|
||||
FROM bm25_ranked WHERE rn <= $4
|
||||
)
|
||||
@@ -386,7 +386,7 @@ async def retrieve_temporal_combined(
|
||||
entry_points = await conn.fetch(
|
||||
f"""
|
||||
WITH ranked_entries AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity,
|
||||
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, embedding <=> $1::vector) AS rn
|
||||
FROM {fq_table("memory_units")}
|
||||
@@ -406,7 +406,7 @@ async def retrieve_temporal_combined(
|
||||
AND (1 - (embedding <=> $1::vector)) >= $6
|
||||
{tags_clause}
|
||||
)
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags, similarity
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags, similarity
|
||||
FROM ranked_entries
|
||||
WHERE rn <= 10
|
||||
""",
|
||||
@@ -486,7 +486,7 @@ async def retrieve_temporal_combined(
|
||||
|
||||
neighbors = await conn.fetch(
|
||||
f"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight, ml.link_type, ml.from_unit_id,
|
||||
1 - (mu.embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_links")} ml
|
||||
@@ -610,7 +610,7 @@ async def retrieve_temporal(
|
||||
|
||||
entry_points = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
@@ -691,7 +691,7 @@ async def retrieve_temporal(
|
||||
# Batch fetch all neighbors for this batch of nodes
|
||||
neighbors = await conn.fetch(
|
||||
f"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id,
|
||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id,
|
||||
ml.weight, ml.link_type, ml.from_unit_id,
|
||||
1 - (mu.embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_links")} ml
|
||||
@@ -1023,7 +1023,7 @@ async def _get_temporal_entry_points(
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
access_count, embedding, fact_type, document_id, chunk_id,
|
||||
embedding, fact_type, document_id, chunk_id,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
|
||||
@@ -1,159 +0,0 @@
|
||||
"""
|
||||
Scoring functions for memory search and retrieval.
|
||||
|
||||
Includes recency weighting, frequency weighting, temporal proximity,
|
||||
and similarity calculations used in memory activation and ranking.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
|
||||
"""
|
||||
Calculate cosine similarity between two vectors.
|
||||
|
||||
Args:
|
||||
vec1: First vector
|
||||
vec2: Second vector
|
||||
|
||||
Returns:
|
||||
Similarity score between 0 and 1
|
||||
"""
|
||||
if len(vec1) != len(vec2):
|
||||
raise ValueError("Vectors must have same dimension")
|
||||
|
||||
dot_product = sum(a * b for a, b in zip(vec1, vec2))
|
||||
magnitude1 = sum(a * a for a in vec1) ** 0.5
|
||||
magnitude2 = sum(b * b for b in vec2) ** 0.5
|
||||
|
||||
if magnitude1 == 0 or magnitude2 == 0:
|
||||
return 0.0
|
||||
|
||||
return dot_product / (magnitude1 * magnitude2)
|
||||
|
||||
|
||||
def calculate_recency_weight(days_since: float, half_life_days: float = 365.0) -> float:
|
||||
"""
|
||||
Calculate recency weight using logarithmic decay.
|
||||
|
||||
This provides much better differentiation over long time periods compared to
|
||||
exponential decay. Uses a log-based decay where the half-life parameter controls
|
||||
when memories reach 50% weight.
|
||||
|
||||
Examples:
|
||||
- Today (0 days): 1.0
|
||||
- 1 year (365 days): ~0.5 (with default half_life=365)
|
||||
- 2 years (730 days): ~0.33
|
||||
- 5 years (1825 days): ~0.17
|
||||
- 10 years (3650 days): ~0.09
|
||||
|
||||
This ensures that 2-year-old and 5-year-old memories have meaningfully
|
||||
different weights, unlike exponential decay which makes them both ~0.
|
||||
|
||||
Args:
|
||||
days_since: Number of days since the memory was created
|
||||
half_life_days: Number of days for weight to reach 0.5 (default: 1 year)
|
||||
|
||||
Returns:
|
||||
Weight between 0 and 1
|
||||
"""
|
||||
import math
|
||||
|
||||
# Logarithmic decay: 1 / (1 + log(1 + days_since/half_life))
|
||||
# This decays much slower than exponential, giving better long-term differentiation
|
||||
normalized_age = days_since / half_life_days
|
||||
return 1.0 / (1.0 + math.log1p(normalized_age))
|
||||
|
||||
|
||||
def calculate_frequency_weight(access_count: int, max_boost: float = 2.0) -> float:
|
||||
"""
|
||||
Calculate frequency weight based on access count.
|
||||
|
||||
Frequently accessed memories are weighted higher.
|
||||
Uses logarithmic scaling to avoid over-weighting.
|
||||
|
||||
Args:
|
||||
access_count: Number of times the memory was accessed
|
||||
max_boost: Maximum multiplier for frequently accessed memories
|
||||
|
||||
Returns:
|
||||
Weight between 1.0 and max_boost
|
||||
"""
|
||||
import math
|
||||
|
||||
if access_count <= 0:
|
||||
return 1.0
|
||||
|
||||
# Logarithmic scaling: log(access_count + 1) / log(10)
|
||||
# This gives: 0 accesses = 1.0, 9 accesses ~= 1.5, 99 accesses ~= 2.0
|
||||
normalized = math.log(access_count + 1) / math.log(10)
|
||||
return 1.0 + min(normalized, max_boost - 1.0)
|
||||
|
||||
|
||||
def calculate_temporal_anchor(occurred_start: datetime, occurred_end: datetime) -> datetime:
|
||||
"""
|
||||
Calculate a single temporal anchor point from a temporal range.
|
||||
|
||||
Used for spreading activation - we need a single representative date
|
||||
to calculate temporal proximity between facts. This simplifies the
|
||||
range-to-range distance problem.
|
||||
|
||||
Strategy: Use midpoint of the range for balanced representation.
|
||||
|
||||
Args:
|
||||
occurred_start: Start of temporal range
|
||||
occurred_end: End of temporal range
|
||||
|
||||
Returns:
|
||||
Single datetime representing the temporal anchor (midpoint)
|
||||
|
||||
Examples:
|
||||
- Point event (July 14): start=July 14, end=July 14 → anchor=July 14
|
||||
- Month range (February): start=Feb 1, end=Feb 28 → anchor=Feb 14
|
||||
- Year range (2023): start=Jan 1, end=Dec 31 → anchor=July 1
|
||||
"""
|
||||
# Calculate midpoint
|
||||
time_delta = occurred_end - occurred_start
|
||||
midpoint = occurred_start + (time_delta / 2)
|
||||
return midpoint
|
||||
|
||||
|
||||
def calculate_temporal_proximity(anchor_a: datetime, anchor_b: datetime, half_life_days: float = 30.0) -> float:
|
||||
"""
|
||||
Calculate temporal proximity between two temporal anchors.
|
||||
|
||||
Used for spreading activation to determine how "close" two facts are
|
||||
in time. Uses logarithmic decay so that temporal similarity doesn't
|
||||
drop off too quickly.
|
||||
|
||||
Args:
|
||||
anchor_a: Temporal anchor of first fact
|
||||
anchor_b: Temporal anchor of second fact
|
||||
half_life_days: Number of days for proximity to reach 0.5
|
||||
(default: 30 days = 1 month)
|
||||
|
||||
Returns:
|
||||
Proximity score in [0, 1] where:
|
||||
- 1.0 = same day
|
||||
- 0.5 = ~half_life days apart
|
||||
- 0.0 = very distant in time
|
||||
|
||||
Examples:
|
||||
- Same day: 1.0
|
||||
- 1 week apart (half_life=30): ~0.7
|
||||
- 1 month apart (half_life=30): ~0.5
|
||||
- 1 year apart (half_life=30): ~0.2
|
||||
"""
|
||||
import math
|
||||
|
||||
days_apart = abs((anchor_a - anchor_b).days)
|
||||
|
||||
if days_apart == 0:
|
||||
return 1.0
|
||||
|
||||
# Logarithmic decay: 1 / (1 + log(1 + days_apart/half_life))
|
||||
# Similar to calculate_recency_weight but for proximity between events
|
||||
normalized_distance = days_apart / half_life_days
|
||||
proximity = 1.0 / (1.0 + math.log1p(normalized_distance))
|
||||
|
||||
return proximity
|
||||
@@ -3,31 +3,13 @@ Think operation utilities for formulating answers based on agent and world facts
|
||||
"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..response_models import DispositionTraits, MemoryFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Opinion(BaseModel):
|
||||
"""An opinion formed by the bank."""
|
||||
|
||||
opinion: str = Field(description="The opinion or perspective with reasoning included")
|
||||
confidence: float = Field(description="Confidence score for this opinion (0.0 to 1.0, where 1.0 is very confident)")
|
||||
|
||||
|
||||
class OpinionExtractionResponse(BaseModel):
|
||||
"""Response containing extracted opinions."""
|
||||
|
||||
opinions: list[Opinion] = Field(
|
||||
default_factory=list, description="List of opinions formed with their supporting reasons and confidence scores"
|
||||
)
|
||||
|
||||
|
||||
def describe_trait_level(value: int) -> str:
|
||||
"""Convert trait value (1-5) to descriptive text."""
|
||||
levels = {1: "very low", 2: "low", 3: "moderate", 4: "high", 5: "very high"}
|
||||
@@ -93,17 +75,46 @@ def format_facts_for_prompt(facts: list[MemoryFact]) -> str:
|
||||
return json.dumps(formatted, indent=2)
|
||||
|
||||
|
||||
def format_entity_summaries_for_prompt(entities: dict) -> str:
|
||||
"""Format entity summaries for inclusion in the reflect prompt.
|
||||
|
||||
Args:
|
||||
entities: Dict mapping entity name to EntityState objects
|
||||
|
||||
Returns:
|
||||
Formatted string with entity summaries, or empty string if no summaries
|
||||
"""
|
||||
if not entities:
|
||||
return ""
|
||||
|
||||
summaries = []
|
||||
for name, state in entities.items():
|
||||
# Get summary from observations (summary is stored as single observation)
|
||||
if state.observations:
|
||||
summary_text = state.observations[0].text
|
||||
summaries.append(f"## {name}\n{summary_text}")
|
||||
|
||||
if not summaries:
|
||||
return ""
|
||||
|
||||
return "\n\n".join(summaries)
|
||||
|
||||
|
||||
def build_think_prompt(
|
||||
agent_facts_text: str,
|
||||
world_facts_text: str,
|
||||
opinion_facts_text: str,
|
||||
query: str,
|
||||
name: str,
|
||||
disposition: DispositionTraits,
|
||||
background: str,
|
||||
context: str | None = None,
|
||||
entity_summaries_text: str | None = None,
|
||||
) -> str:
|
||||
"""Build the think prompt for the LLM."""
|
||||
"""Build the think prompt for the LLM.
|
||||
|
||||
Note: opinion_facts_text parameter removed - opinions are now stored as mental models
|
||||
and included via entity_summaries_text.
|
||||
"""
|
||||
disposition_desc = build_disposition_description(disposition)
|
||||
|
||||
name_section = f"""
|
||||
@@ -125,6 +136,14 @@ Your background:
|
||||
ADDITIONAL CONTEXT:
|
||||
{context}
|
||||
|
||||
"""
|
||||
|
||||
entity_section = ""
|
||||
if entity_summaries_text:
|
||||
entity_section = f"""
|
||||
KEY PEOPLE, PLACES & THINGS I KNOW ABOUT:
|
||||
{entity_summaries_text}
|
||||
|
||||
"""
|
||||
|
||||
return f"""Here's what I know and have experienced:
|
||||
@@ -135,14 +154,11 @@ MY IDENTITY & EXPERIENCES:
|
||||
WHAT I KNOW ABOUT THE WORLD:
|
||||
{world_facts_text}
|
||||
|
||||
MY EXISTING OPINIONS & BELIEFS:
|
||||
{opinion_facts_text}
|
||||
|
||||
{context_section}{name_section}{disposition_desc}{background_section}
|
||||
{entity_section}{context_section}{name_section}{disposition_desc}{background_section}
|
||||
|
||||
QUESTION: {query}
|
||||
|
||||
Based on everything I know, believe, and who I am (including my name, disposition and background), here's what I genuinely think about this question. I'll draw on my experiences, knowledge, opinions, and personal traits to give you my honest perspective."""
|
||||
Based on everything I know, believe, and who I am (including my name, disposition and background), here's what I genuinely think about this question. I'll draw on my experiences, knowledge, and personal traits to give you my honest perspective."""
|
||||
|
||||
|
||||
def get_system_message(disposition: DispositionTraits) -> str:
|
||||
@@ -172,117 +188,7 @@ def get_system_message(disposition: DispositionTraits) -> str:
|
||||
" ".join(instructions) if instructions else "Balance your disposition traits when interpreting information."
|
||||
)
|
||||
|
||||
return f"You are a person with your own thoughts, experiences, opinions, and disposition. When asked a question, respond naturally like a human would - share what you genuinely think based on what you know and have experienced. {disposition_instruction} Be direct, express your views confidently, and use 'I think', 'I believe', 'in my view', etc. Respond in plain text without markdown formatting. IMPORTANT: Detect the language of the question and respond in the SAME language. Do not translate to English if the question is in another language."
|
||||
|
||||
|
||||
async def extract_opinions_from_text(llm_config, text: str, query: str) -> list[Opinion]:
|
||||
"""
|
||||
Extract opinions with reasons and confidence from text using LLM.
|
||||
|
||||
Args:
|
||||
llm_config: LLM configuration to use
|
||||
text: Text to extract opinions from
|
||||
query: The original query that prompted this response
|
||||
|
||||
Returns:
|
||||
List of Opinion objects with text and confidence
|
||||
"""
|
||||
extraction_prompt = f"""Extract any NEW opinions or perspectives from the answer below and rewrite them in FIRST-PERSON as if YOU are stating the opinion directly.
|
||||
|
||||
ORIGINAL QUESTION:
|
||||
{query}
|
||||
|
||||
ANSWER PROVIDED:
|
||||
{text}
|
||||
|
||||
Your task: Find opinions in the answer and rewrite them AS IF YOU ARE THE ONE SAYING THEM.
|
||||
|
||||
An opinion is a judgment, viewpoint, or conclusion that goes beyond just stating facts.
|
||||
|
||||
IMPORTANT: Do NOT extract statements like:
|
||||
- "I don't have enough information"
|
||||
- "The facts don't contain information about X"
|
||||
- "I cannot answer because..."
|
||||
|
||||
ONLY extract actual opinions about substantive topics.
|
||||
|
||||
CRITICAL FORMAT REQUIREMENTS:
|
||||
1. **ALWAYS start with first-person phrases**: "I think...", "I believe...", "In my view...", "I've come to believe...", "Previously I thought... but now..."
|
||||
2. **NEVER use third-person**: Do NOT say "The speaker thinks..." or "They believe..." - always use "I"
|
||||
3. Include the reasoning naturally within the statement
|
||||
4. Provide a confidence score (0.0 to 1.0)
|
||||
|
||||
CORRECT Examples (✓ FIRST-PERSON):
|
||||
- "I think Alice is more reliable because she consistently delivers on time and writes clean code"
|
||||
- "Previously I thought all engineers were equal, but now I feel that experience and track record really matter"
|
||||
- "I believe reliability is best measured by consistent output over time"
|
||||
- "I've come to believe that track records are more important than potential"
|
||||
|
||||
WRONG Examples (✗ THIRD-PERSON - DO NOT USE):
|
||||
- "The speaker thinks Alice is more reliable"
|
||||
- "They believe reliability matters"
|
||||
- "It is believed that Alice is better"
|
||||
|
||||
If no genuine opinions are expressed (e.g., the response just says "I don't know"), return an empty list."""
|
||||
|
||||
try:
|
||||
result = await llm_config.call(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are converting opinions from text into first-person statements. Always use 'I think', 'I believe', 'I feel', etc. NEVER use third-person like 'The speaker' or 'They'.",
|
||||
},
|
||||
{"role": "user", "content": extraction_prompt},
|
||||
],
|
||||
response_format=OpinionExtractionResponse,
|
||||
scope="memory_extract_opinion",
|
||||
)
|
||||
|
||||
# Format opinions with confidence score and convert to first-person
|
||||
formatted_opinions = []
|
||||
for op in result.opinions:
|
||||
# Convert third-person to first-person if needed
|
||||
opinion_text = op.opinion
|
||||
|
||||
# Replace common third-person patterns with first-person
|
||||
def singularize_verb(verb):
|
||||
if verb.endswith("es"):
|
||||
return verb[:-1] # believes -> believe
|
||||
elif verb.endswith("s"):
|
||||
return verb[:-1] # thinks -> think
|
||||
return verb
|
||||
|
||||
# Pattern: "The speaker/user [verb]..." -> "I [verb]..."
|
||||
match = re.match(
|
||||
r"^(The speaker|The user|They|It is believed) (believes?|thinks?|feels?|says|asserts?|considers?)(\s+that)?(.*)$",
|
||||
opinion_text,
|
||||
re.IGNORECASE,
|
||||
)
|
||||
if match:
|
||||
verb = singularize_verb(match.group(2))
|
||||
that_part = match.group(3) or "" # Keep " that" if present
|
||||
rest = match.group(4)
|
||||
opinion_text = f"I {verb}{that_part}{rest}"
|
||||
|
||||
# If still doesn't start with first-person, prepend "I believe that "
|
||||
first_person_starters = [
|
||||
"I think",
|
||||
"I believe",
|
||||
"I feel",
|
||||
"In my view",
|
||||
"I've come to believe",
|
||||
"Previously I",
|
||||
]
|
||||
if not any(opinion_text.startswith(starter) for starter in first_person_starters):
|
||||
opinion_text = "I believe that " + opinion_text[0].lower() + opinion_text[1:]
|
||||
|
||||
formatted_opinions.append(Opinion(opinion=opinion_text, confidence=op.confidence))
|
||||
|
||||
return formatted_opinions
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to extract opinions: {str(e)}")
|
||||
return []
|
||||
return f"You are a person with your own thoughts, experiences, opinions, and disposition. When asked a question, respond naturally like a human would - share what you genuinely think based on what you know and have experienced. {disposition_instruction} Be direct, express your views confidently, and use 'I think', 'I believe', 'in my view', etc. Respond in plain text without markdown formatting. CRITICAL: ONLY use the facts and information provided in the prompt - do not make up names, events, or information that weren't mentioned. If you don't have enough information to answer, say so. IMPORTANT: Detect the language of the question and respond in the SAME language. Do not translate to English if the question is in another language."
|
||||
|
||||
|
||||
async def reflect(
|
||||
@@ -290,7 +196,6 @@ async def reflect(
|
||||
query: str,
|
||||
experience_facts: list[str] = None,
|
||||
world_facts: list[str] = None,
|
||||
opinion_facts: list[str] = None,
|
||||
name: str = "Assistant",
|
||||
disposition: DispositionTraits = None,
|
||||
background: str = "",
|
||||
@@ -307,7 +212,6 @@ async def reflect(
|
||||
query: Question to answer
|
||||
experience_facts: List of experience/agent fact strings
|
||||
world_facts: List of world fact strings
|
||||
opinion_facts: List of opinion fact strings
|
||||
name: Name of the agent/persona
|
||||
disposition: Disposition traits (defaults to neutral)
|
||||
background: Background information
|
||||
@@ -328,18 +232,15 @@ async def reflect(
|
||||
|
||||
agent_results = to_memory_facts(experience_facts or [], "experience")
|
||||
world_results = to_memory_facts(world_facts or [], "world")
|
||||
opinion_results = to_memory_facts(opinion_facts or [], "opinion")
|
||||
|
||||
# Format facts for prompt
|
||||
agent_facts_text = format_facts_for_prompt(agent_results)
|
||||
world_facts_text = format_facts_for_prompt(world_results)
|
||||
opinion_facts_text = format_facts_for_prompt(opinion_results)
|
||||
|
||||
# Build prompt
|
||||
prompt = build_think_prompt(
|
||||
agent_facts_text=agent_facts_text,
|
||||
world_facts_text=world_facts_text,
|
||||
opinion_facts_text=opinion_facts_text,
|
||||
query=query,
|
||||
name=name,
|
||||
disposition=disposition,
|
||||
|
||||
@@ -85,7 +85,6 @@ class NodeVisit(BaseModel):
|
||||
text: str = Field(description="Memory unit text content")
|
||||
context: str = Field(description="Memory unit context")
|
||||
event_date: datetime | None = Field(default=None, description="When the memory occurred")
|
||||
access_count: int = Field(description="Number of times accessed before this search")
|
||||
|
||||
# How this node was reached
|
||||
is_entry_point: bool = Field(description="Whether this is an entry point")
|
||||
|
||||
@@ -136,7 +136,6 @@ class SearchTracer:
|
||||
text: str,
|
||||
context: str,
|
||||
event_date: datetime | None,
|
||||
access_count: int,
|
||||
is_entry_point: bool,
|
||||
parent_node_id: str | None,
|
||||
link_type: Literal["temporal", "semantic", "entity"] | None,
|
||||
@@ -155,7 +154,6 @@ class SearchTracer:
|
||||
text: Memory unit text
|
||||
context: Memory unit context
|
||||
event_date: When the memory occurred
|
||||
access_count: Access count before this search
|
||||
is_entry_point: Whether this is an entry point
|
||||
parent_node_id: Node that led here (None for entry points)
|
||||
link_type: Type of link from parent
|
||||
@@ -194,7 +192,6 @@ class SearchTracer:
|
||||
text=text,
|
||||
context=context,
|
||||
event_date=event_date,
|
||||
access_count=access_count,
|
||||
is_entry_point=is_entry_point,
|
||||
parent_node_id=parent_node_id,
|
||||
link_type=link_type,
|
||||
@@ -333,8 +330,8 @@ class SearchTracer:
|
||||
RetrievalResult(
|
||||
rank=rank,
|
||||
node_id=doc_id,
|
||||
text=data.get("text", ""),
|
||||
context=data.get("context", ""),
|
||||
text=data.get("text") or "",
|
||||
context=data.get("context") or "",
|
||||
event_date=data.get("event_date"),
|
||||
fact_type=data.get("fact_type") or fact_type,
|
||||
score=score,
|
||||
|
||||
@@ -46,7 +46,6 @@ class RetrievalResult:
|
||||
mentioned_at: datetime | None = None
|
||||
document_id: str | None = None
|
||||
chunk_id: str | None = None
|
||||
access_count: int = 0
|
||||
embedding: list[float] | None = None
|
||||
tags: list[str] | None = None # Visibility scope tags
|
||||
|
||||
@@ -71,7 +70,6 @@ class RetrievalResult:
|
||||
mentioned_at=row.get("mentioned_at"),
|
||||
document_id=row.get("document_id"),
|
||||
chunk_id=row.get("chunk_id"),
|
||||
access_count=row.get("access_count", 0),
|
||||
embedding=row.get("embedding"),
|
||||
tags=row.get("tags"),
|
||||
similarity=row.get("similarity"),
|
||||
@@ -156,7 +154,6 @@ class ScoredResult:
|
||||
"mentioned_at": self.retrieval.mentioned_at,
|
||||
"document_id": self.retrieval.document_id,
|
||||
"chunk_id": self.retrieval.chunk_id,
|
||||
"access_count": self.retrieval.access_count,
|
||||
"embedding": self.retrieval.embedding,
|
||||
"tags": self.retrieval.tags,
|
||||
"semantic_similarity": self.retrieval.similarity,
|
||||
|
||||
@@ -1,31 +1,40 @@
|
||||
"""
|
||||
Abstract task backend for running async tasks.
|
||||
Task backend for distributed task processing.
|
||||
|
||||
This provides an abstraction that can be adapted to different execution models:
|
||||
- AsyncIO queue (default implementation)
|
||||
- Pub/Sub architectures (future)
|
||||
- Message brokers (future)
|
||||
This provides an abstraction for task storage and execution:
|
||||
- BrokerTaskBackend: Uses PostgreSQL as broker (production)
|
||||
- SyncTaskBackend: Executes tasks immediately (testing/embedded)
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import asyncpg
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def fq_table(table: str, schema: str | None = None) -> str:
|
||||
"""Get fully-qualified table name with optional schema prefix."""
|
||||
if schema:
|
||||
return f'"{schema}".{table}'
|
||||
return table
|
||||
|
||||
|
||||
class TaskBackend(ABC):
|
||||
"""
|
||||
Abstract base class for task execution backends.
|
||||
|
||||
Implementations must:
|
||||
1. Store/publish task events (as serializable dicts)
|
||||
2. Execute tasks through a provided executor callback
|
||||
2. Execute tasks through a provided executor callback (optional)
|
||||
|
||||
The backend treats tasks as pure dictionaries that can be serialized
|
||||
and sent over the network. The executor (typically MemoryEngine.execute_task)
|
||||
and stored in the database. The executor (typically MemoryEngine.execute_task)
|
||||
receives the dict and routes it to the appropriate handler.
|
||||
"""
|
||||
|
||||
@@ -46,7 +55,7 @@ class TaskBackend(ABC):
|
||||
@abstractmethod
|
||||
async def initialize(self):
|
||||
"""
|
||||
Initialize the backend (e.g., start workers, connect to broker).
|
||||
Initialize the backend (e.g., connect to database).
|
||||
"""
|
||||
pass
|
||||
|
||||
@@ -63,7 +72,7 @@ class TaskBackend(ABC):
|
||||
@abstractmethod
|
||||
async def shutdown(self):
|
||||
"""
|
||||
Shutdown the backend gracefully (e.g., stop workers, close connections).
|
||||
Shutdown the backend gracefully.
|
||||
"""
|
||||
pass
|
||||
|
||||
@@ -93,9 +102,8 @@ class SyncTaskBackend(TaskBackend):
|
||||
"""
|
||||
Synchronous task backend that executes tasks immediately.
|
||||
|
||||
This is useful for embedded/CLI usage where we don't want background
|
||||
workers that prevent clean exit. Tasks are executed inline rather than
|
||||
being queued.
|
||||
This is useful for tests and embedded/CLI usage where we don't want
|
||||
background workers. Tasks are executed inline rather than being queued.
|
||||
"""
|
||||
|
||||
async def initialize(self):
|
||||
@@ -121,221 +129,138 @@ class SyncTaskBackend(TaskBackend):
|
||||
logger.debug("SyncTaskBackend shutdown")
|
||||
|
||||
|
||||
class NoopTaskBackend(TaskBackend):
|
||||
class BrokerTaskBackend(TaskBackend):
|
||||
"""
|
||||
No-op task backend that discards all tasks.
|
||||
Task backend using PostgreSQL as broker.
|
||||
|
||||
This is useful for tests where background task execution is not needed
|
||||
and would only slow down the test suite.
|
||||
submit_task() stores task_payload in async_operations table.
|
||||
Actual polling and execution is handled separately by WorkerPoller.
|
||||
|
||||
This backend is used by the API to store tasks. Workers poll
|
||||
the database separately to claim and execute tasks.
|
||||
"""
|
||||
|
||||
async def initialize(self):
|
||||
"""No-op."""
|
||||
self._initialized = True
|
||||
logger.debug("NoopTaskBackend initialized")
|
||||
|
||||
async def submit_task(self, task_dict: dict[str, Any]):
|
||||
"""Discard the task (do nothing)."""
|
||||
pass
|
||||
|
||||
async def shutdown(self):
|
||||
"""No-op."""
|
||||
self._initialized = False
|
||||
logger.debug("NoopTaskBackend shutdown")
|
||||
|
||||
|
||||
class AsyncIOQueueBackend(TaskBackend):
|
||||
"""
|
||||
Task backend implementation using asyncio queues.
|
||||
|
||||
This is the default implementation that uses in-process asyncio queues
|
||||
and a periodic consumer worker.
|
||||
"""
|
||||
|
||||
def __init__(self, batch_size: int = 10, batch_interval: float = 1.0):
|
||||
def __init__(
|
||||
self,
|
||||
pool_getter: Callable[[], "asyncpg.Pool"],
|
||||
schema: str | None = None,
|
||||
schema_getter: Callable[[], str | None] | None = None,
|
||||
):
|
||||
"""
|
||||
Initialize AsyncIO queue backend.
|
||||
Initialize the broker task backend.
|
||||
|
||||
Args:
|
||||
batch_size: Maximum number of tasks to process in one batch
|
||||
batch_interval: Maximum time (seconds) to wait before processing batch
|
||||
pool_getter: Callable that returns the asyncpg connection pool
|
||||
schema: Database schema for multi-tenant support (optional, static)
|
||||
schema_getter: Callable that returns current schema dynamically (optional).
|
||||
If set, takes precedence over static schema for submit_task.
|
||||
"""
|
||||
super().__init__()
|
||||
self._queue: asyncio.Queue | None = None
|
||||
self._worker_task: asyncio.Task | None = None
|
||||
self._shutdown_event: asyncio.Event | None = None
|
||||
self._batch_size = batch_size
|
||||
self._batch_interval = batch_interval
|
||||
self._in_flight_count = 0
|
||||
self._in_flight_lock = asyncio.Lock()
|
||||
self._pool_getter = pool_getter
|
||||
self._schema = schema
|
||||
self._schema_getter = schema_getter
|
||||
|
||||
async def initialize(self):
|
||||
"""Initialize the queue and start the worker."""
|
||||
if self._initialized:
|
||||
return
|
||||
|
||||
self._queue = asyncio.Queue()
|
||||
self._shutdown_event = asyncio.Event()
|
||||
self._worker_task = asyncio.create_task(self._worker())
|
||||
"""Initialize the backend."""
|
||||
self._initialized = True
|
||||
logger.info("AsyncIOQueueBackend initialized")
|
||||
logger.info("BrokerTaskBackend initialized")
|
||||
|
||||
async def submit_task(self, task_dict: dict[str, Any]):
|
||||
"""
|
||||
Submit a task by putting it in the queue.
|
||||
Store task payload in async_operations table.
|
||||
|
||||
The task_dict should contain an 'operation_id' if updating an existing
|
||||
operation record, otherwise a new operation will be created.
|
||||
|
||||
Args:
|
||||
task_dict: Task dictionary to execute
|
||||
task_dict: Task dictionary to store (must be JSON serializable)
|
||||
"""
|
||||
if not self._initialized:
|
||||
await self.initialize()
|
||||
|
||||
await self._queue.put(task_dict)
|
||||
pool = self._pool_getter()
|
||||
operation_id = task_dict.get("operation_id")
|
||||
task_type = task_dict.get("type", "unknown")
|
||||
bank_id = task_dict.get("bank_id")
|
||||
|
||||
# Custom encoder to handle datetime objects
|
||||
from datetime import datetime
|
||||
|
||||
def datetime_encoder(obj):
|
||||
if isinstance(obj, datetime):
|
||||
return obj.isoformat()
|
||||
raise TypeError(f"Object of type {type(obj).__name__} is not JSON serializable")
|
||||
|
||||
payload_json = json.dumps(task_dict, default=datetime_encoder)
|
||||
|
||||
schema = self._schema_getter() if self._schema_getter else self._schema
|
||||
table = fq_table("async_operations", schema)
|
||||
|
||||
if operation_id:
|
||||
# Update existing operation with task payload
|
||||
await pool.execute(
|
||||
f"""
|
||||
UPDATE {table}
|
||||
SET task_payload = $1::jsonb, updated_at = now()
|
||||
WHERE operation_id = $2
|
||||
""",
|
||||
payload_json,
|
||||
operation_id,
|
||||
)
|
||||
logger.debug(f"Updated task payload for operation {operation_id}")
|
||||
else:
|
||||
# Insert new operation (for tasks without pre-created records)
|
||||
# e.g., access_count_update tasks
|
||||
import uuid
|
||||
|
||||
new_id = uuid.uuid4()
|
||||
await pool.execute(
|
||||
f"""
|
||||
INSERT INTO {table} (operation_id, bank_id, operation_type, status, task_payload)
|
||||
VALUES ($1, $2, $3, 'pending', $4::jsonb)
|
||||
""",
|
||||
new_id,
|
||||
bank_id,
|
||||
task_type,
|
||||
payload_json,
|
||||
)
|
||||
logger.debug(f"Created new operation {new_id} for task type {task_type}")
|
||||
|
||||
async def shutdown(self):
|
||||
"""Shutdown the backend."""
|
||||
self._initialized = False
|
||||
logger.info("BrokerTaskBackend shutdown")
|
||||
|
||||
async def wait_for_pending_tasks(self, timeout: float = 120.0):
|
||||
"""
|
||||
Wait for all pending tasks in the queue and in-flight tasks to complete.
|
||||
Wait for pending tasks to be processed.
|
||||
|
||||
This is useful in tests to ensure background tasks complete before assertions.
|
||||
In the broker model, this polls the database to check if tasks
|
||||
for this process have been completed. This is useful in tests
|
||||
when worker_enabled=True (API processes its own tasks).
|
||||
|
||||
Args:
|
||||
timeout: Maximum time to wait in seconds (default 120s for long-running tasks)
|
||||
timeout: Maximum time to wait in seconds
|
||||
"""
|
||||
if not self._initialized or self._queue is None:
|
||||
return
|
||||
import asyncio
|
||||
|
||||
pool = self._pool_getter()
|
||||
schema = self._schema_getter() if self._schema_getter else self._schema
|
||||
table = fq_table("async_operations", schema)
|
||||
|
||||
# Wait for queue to be empty AND no in-flight tasks
|
||||
start_time = asyncio.get_event_loop().time()
|
||||
while asyncio.get_event_loop().time() - start_time < timeout:
|
||||
async with self._in_flight_lock:
|
||||
in_flight = self._in_flight_count
|
||||
# Check if there are any pending tasks with payloads
|
||||
count = await pool.fetchval(
|
||||
f"""
|
||||
SELECT COUNT(*) FROM {table}
|
||||
WHERE status = 'pending' AND task_payload IS NOT NULL
|
||||
"""
|
||||
)
|
||||
|
||||
if self._queue.empty() and in_flight == 0:
|
||||
# Queue is empty and no tasks in flight, we're done
|
||||
if count == 0:
|
||||
return
|
||||
|
||||
# Wait a bit before checking again
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
async def shutdown(self):
|
||||
"""Shutdown the worker and drain the queue."""
|
||||
if not self._initialized:
|
||||
return
|
||||
|
||||
logger.info("Shutting down AsyncIOQueueBackend...")
|
||||
|
||||
# Signal shutdown
|
||||
self._shutdown_event.set()
|
||||
|
||||
# Cancel worker
|
||||
if self._worker_task is not None:
|
||||
self._worker_task.cancel()
|
||||
try:
|
||||
await self._worker_task
|
||||
except asyncio.CancelledError:
|
||||
pass # Worker cancelled successfully
|
||||
|
||||
self._initialized = False
|
||||
logger.info("AsyncIOQueueBackend shutdown complete")
|
||||
|
||||
async def _execute_task_with_tracking(self, task_dict: dict[str, Any]):
|
||||
"""Execute a task and track its in-flight status."""
|
||||
async with self._in_flight_lock:
|
||||
self._in_flight_count += 1
|
||||
try:
|
||||
await self._execute_task(task_dict)
|
||||
finally:
|
||||
async with self._in_flight_lock:
|
||||
self._in_flight_count -= 1
|
||||
|
||||
async def _execute_task_no_tracking(self, task_dict: dict[str, Any]):
|
||||
"""Execute a task without in-flight tracking (tracking done at batch level)."""
|
||||
await self._execute_task(task_dict)
|
||||
|
||||
def _get_queue_stats(self) -> tuple[int, dict[str, int]]:
|
||||
"""Get current queue size and bank_id distribution."""
|
||||
queue_size = self._queue.qsize() if self._queue else 0
|
||||
bank_distribution: dict[str, int] = {}
|
||||
|
||||
if queue_size > 0 and self._queue:
|
||||
# Peek at queue items without removing them
|
||||
# Note: This is a snapshot and may not be perfectly accurate due to concurrency
|
||||
try:
|
||||
# Access internal deque for logging purposes only
|
||||
items = list(self._queue._queue) # type: ignore[attr-defined]
|
||||
for item in items:
|
||||
bank_id = item.get("bank_id", "unknown")
|
||||
bank_distribution[bank_id] = bank_distribution.get(bank_id, 0) + 1
|
||||
except Exception:
|
||||
pass # Queue access failed, return empty distribution
|
||||
|
||||
return queue_size, bank_distribution
|
||||
|
||||
async def _worker(self):
|
||||
"""
|
||||
Background worker that processes tasks in batches.
|
||||
|
||||
Collects tasks for up to batch_interval seconds or batch_size items,
|
||||
then processes them.
|
||||
"""
|
||||
while not self._shutdown_event.is_set():
|
||||
try:
|
||||
# Collect tasks for batching
|
||||
tasks = []
|
||||
deadline = asyncio.get_event_loop().time() + self._batch_interval
|
||||
|
||||
while len(tasks) < self._batch_size and asyncio.get_event_loop().time() < deadline:
|
||||
try:
|
||||
remaining_time = max(0.1, deadline - asyncio.get_event_loop().time())
|
||||
task_dict = await asyncio.wait_for(self._queue.get(), timeout=remaining_time)
|
||||
# Track task as in-flight immediately when picked up from queue
|
||||
# This prevents wait_for_pending_tasks from returning too early
|
||||
async with self._in_flight_lock:
|
||||
self._in_flight_count += 1
|
||||
tasks.append(task_dict)
|
||||
except TimeoutError:
|
||||
break
|
||||
|
||||
# Process batch
|
||||
if tasks:
|
||||
# Log batch start with queue stats
|
||||
queue_size, bank_distribution = self._get_queue_stats()
|
||||
|
||||
# Summarize batch by task type and bank
|
||||
batch_summary: dict[str, dict[str, int]] = {}
|
||||
for task_dict in tasks:
|
||||
task_type = task_dict.get("type", "unknown")
|
||||
bank_id = task_dict.get("bank_id", "unknown")
|
||||
if task_type not in batch_summary:
|
||||
batch_summary[task_type] = {}
|
||||
batch_summary[task_type][bank_id] = batch_summary[task_type].get(bank_id, 0) + 1
|
||||
|
||||
# Build log message
|
||||
batch_parts = []
|
||||
for task_type, banks in sorted(batch_summary.items()):
|
||||
bank_str = ", ".join(f"{b}:{c}" for b, c in sorted(banks.items()))
|
||||
batch_parts.append(f"{task_type}[{bank_str}]")
|
||||
batch_str = ", ".join(batch_parts)
|
||||
|
||||
if queue_size > 0:
|
||||
pending_str = ", ".join(f"{k}:{v}" for k, v in sorted(bank_distribution.items()))
|
||||
logger.info(
|
||||
f"Processing {len(tasks)} tasks: {batch_str} (pending={queue_size} [{pending_str}])"
|
||||
)
|
||||
else:
|
||||
logger.info(f"Processing {len(tasks)} tasks: {batch_str}")
|
||||
|
||||
# Execute tasks concurrently (in_flight already tracked when picked up)
|
||||
await asyncio.gather(
|
||||
*[self._execute_task_no_tracking(task_dict) for task_dict in tasks], return_exceptions=True
|
||||
)
|
||||
|
||||
# Decrement in_flight count after all tasks complete
|
||||
async with self._in_flight_lock:
|
||||
self._in_flight_count -= len(tasks)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Worker error: {e}")
|
||||
await asyncio.sleep(1) # Backoff on error
|
||||
logger.warning(f"Timeout waiting for pending tasks after {timeout}s")
|
||||
|
||||
@@ -19,7 +19,6 @@ async def extract_facts(
|
||||
context: str = "",
|
||||
llm_config: "LLMConfig" = None,
|
||||
agent_name: str = None,
|
||||
extract_opinions: bool = False,
|
||||
) -> tuple[list["Fact"], list[tuple[str, int]]]:
|
||||
"""
|
||||
Extract semantic facts from text using LLM.
|
||||
@@ -36,7 +35,6 @@ async def extract_facts(
|
||||
context: Context about the conversation/document
|
||||
llm_config: LLM configuration to use
|
||||
agent_name: Optional agent name to help identify agent-related facts
|
||||
extract_opinions: If True, extract ONLY opinions. If False, extract world and agent facts (no opinions)
|
||||
|
||||
Returns:
|
||||
Tuple of (facts, chunks) where:
|
||||
@@ -55,7 +53,6 @@ async def extract_facts(
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions,
|
||||
)
|
||||
|
||||
if not facts:
|
||||
@@ -65,154 +62,3 @@ async def extract_facts(
|
||||
return [], chunks
|
||||
|
||||
return facts, chunks
|
||||
|
||||
|
||||
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
|
||||
"""
|
||||
Calculate cosine similarity between two vectors.
|
||||
|
||||
Args:
|
||||
vec1: First vector
|
||||
vec2: Second vector
|
||||
|
||||
Returns:
|
||||
Similarity score between 0 and 1
|
||||
"""
|
||||
if len(vec1) != len(vec2):
|
||||
raise ValueError("Vectors must have same dimension")
|
||||
|
||||
dot_product = sum(a * b for a, b in zip(vec1, vec2))
|
||||
magnitude1 = sum(a * a for a in vec1) ** 0.5
|
||||
magnitude2 = sum(b * b for b in vec2) ** 0.5
|
||||
|
||||
if magnitude1 == 0 or magnitude2 == 0:
|
||||
return 0.0
|
||||
|
||||
return dot_product / (magnitude1 * magnitude2)
|
||||
|
||||
|
||||
def calculate_recency_weight(days_since: float, half_life_days: float = 365.0) -> float:
|
||||
"""
|
||||
Calculate recency weight using logarithmic decay.
|
||||
|
||||
This provides much better differentiation over long time periods compared to
|
||||
exponential decay. Uses a log-based decay where the half-life parameter controls
|
||||
when memories reach 50% weight.
|
||||
|
||||
Examples:
|
||||
- Today (0 days): 1.0
|
||||
- 1 year (365 days): ~0.5 (with default half_life=365)
|
||||
- 2 years (730 days): ~0.33
|
||||
- 5 years (1825 days): ~0.17
|
||||
- 10 years (3650 days): ~0.09
|
||||
|
||||
This ensures that 2-year-old and 5-year-old memories have meaningfully
|
||||
different weights, unlike exponential decay which makes them both ~0.
|
||||
|
||||
Args:
|
||||
days_since: Number of days since the memory was created
|
||||
half_life_days: Number of days for weight to reach 0.5 (default: 1 year)
|
||||
|
||||
Returns:
|
||||
Weight between 0 and 1
|
||||
"""
|
||||
import math
|
||||
|
||||
# Logarithmic decay: 1 / (1 + log(1 + days_since/half_life))
|
||||
# This decays much slower than exponential, giving better long-term differentiation
|
||||
normalized_age = days_since / half_life_days
|
||||
return 1.0 / (1.0 + math.log1p(normalized_age))
|
||||
|
||||
|
||||
def calculate_frequency_weight(access_count: int, max_boost: float = 2.0) -> float:
|
||||
"""
|
||||
Calculate frequency weight based on access count.
|
||||
|
||||
Frequently accessed memories are weighted higher.
|
||||
Uses logarithmic scaling to avoid over-weighting.
|
||||
|
||||
Args:
|
||||
access_count: Number of times the memory was accessed
|
||||
max_boost: Maximum multiplier for frequently accessed memories
|
||||
|
||||
Returns:
|
||||
Weight between 1.0 and max_boost
|
||||
"""
|
||||
import math
|
||||
|
||||
if access_count <= 0:
|
||||
return 1.0
|
||||
|
||||
# Logarithmic scaling: log(access_count + 1) / log(10)
|
||||
# This gives: 0 accesses = 1.0, 9 accesses ~= 1.5, 99 accesses ~= 2.0
|
||||
normalized = math.log(access_count + 1) / math.log(10)
|
||||
return 1.0 + min(normalized, max_boost - 1.0)
|
||||
|
||||
|
||||
def calculate_temporal_anchor(occurred_start: datetime, occurred_end: datetime) -> datetime:
|
||||
"""
|
||||
Calculate a single temporal anchor point from a temporal range.
|
||||
|
||||
Used for spreading activation - we need a single representative date
|
||||
to calculate temporal proximity between facts. This simplifies the
|
||||
range-to-range distance problem.
|
||||
|
||||
Strategy: Use midpoint of the range for balanced representation.
|
||||
|
||||
Args:
|
||||
occurred_start: Start of temporal range
|
||||
occurred_end: End of temporal range
|
||||
|
||||
Returns:
|
||||
Single datetime representing the temporal anchor (midpoint)
|
||||
|
||||
Examples:
|
||||
- Point event (July 14): start=July 14, end=July 14 → anchor=July 14
|
||||
- Month range (February): start=Feb 1, end=Feb 28 → anchor=Feb 14
|
||||
- Year range (2023): start=Jan 1, end=Dec 31 → anchor=July 1
|
||||
"""
|
||||
# Calculate midpoint
|
||||
time_delta = occurred_end - occurred_start
|
||||
midpoint = occurred_start + (time_delta / 2)
|
||||
return midpoint
|
||||
|
||||
|
||||
def calculate_temporal_proximity(anchor_a: datetime, anchor_b: datetime, half_life_days: float = 30.0) -> float:
|
||||
"""
|
||||
Calculate temporal proximity between two temporal anchors.
|
||||
|
||||
Used for spreading activation to determine how "close" two facts are
|
||||
in time. Uses logarithmic decay so that temporal similarity doesn't
|
||||
drop off too quickly.
|
||||
|
||||
Args:
|
||||
anchor_a: Temporal anchor of first fact
|
||||
anchor_b: Temporal anchor of second fact
|
||||
half_life_days: Number of days for proximity to reach 0.5
|
||||
(default: 30 days = 1 month)
|
||||
|
||||
Returns:
|
||||
Proximity score in [0, 1] where:
|
||||
- 1.0 = same day
|
||||
- 0.5 = ~half_life days apart
|
||||
- 0.0 = very distant in time
|
||||
|
||||
Examples:
|
||||
- Same day: 1.0
|
||||
- 1 week apart (half_life=30): ~0.7
|
||||
- 1 month apart (half_life=30): ~0.5
|
||||
- 1 year apart (half_life=30): ~0.2
|
||||
"""
|
||||
import math
|
||||
|
||||
days_apart = abs((anchor_a - anchor_b).days)
|
||||
|
||||
if days_apart == 0:
|
||||
return 1.0
|
||||
|
||||
# Logarithmic decay: 1 / (1 + log(1 + days_apart/half_life))
|
||||
# Similar to calculate_recency_weight but for proximity between events
|
||||
normalized_distance = days_apart / half_life_days
|
||||
proximity = 1.0 / (1.0 + math.log1p(normalized_distance))
|
||||
|
||||
return proximity
|
||||
|
||||
@@ -16,11 +16,21 @@ with the system (e.g., running migrations for tenant schemas).
|
||||
"""
|
||||
|
||||
from hindsight_api.extensions.base import Extension
|
||||
from hindsight_api.extensions.builtin import ApiKeyTenantExtension
|
||||
from hindsight_api.extensions.builtin import ApiKeyTenantExtension, SupabaseTenantExtension
|
||||
from hindsight_api.extensions.context import DefaultExtensionContext, ExtensionContext
|
||||
from hindsight_api.extensions.http import HttpExtension
|
||||
from hindsight_api.extensions.loader import load_extension
|
||||
from hindsight_api.extensions.mcp import MCPExtension
|
||||
from hindsight_api.extensions.operation_validator import (
|
||||
# Consolidation operation
|
||||
ConsolidateContext,
|
||||
ConsolidateResult,
|
||||
# Mental Model operations
|
||||
MentalModelGetContext,
|
||||
MentalModelGetResult,
|
||||
MentalModelRefreshContext,
|
||||
MentalModelRefreshResult,
|
||||
# Core operations
|
||||
OperationValidationError,
|
||||
OperationValidatorExtension,
|
||||
RecallContext,
|
||||
@@ -33,6 +43,7 @@ from hindsight_api.extensions.operation_validator import (
|
||||
)
|
||||
from hindsight_api.extensions.tenant import (
|
||||
AuthenticationError,
|
||||
Tenant,
|
||||
TenantContext,
|
||||
TenantExtension,
|
||||
)
|
||||
@@ -47,7 +58,9 @@ __all__ = [
|
||||
"DefaultExtensionContext",
|
||||
# HTTP Extension
|
||||
"HttpExtension",
|
||||
# Operation Validator
|
||||
# MCP Extension
|
||||
"MCPExtension",
|
||||
# Operation Validator - Core
|
||||
"OperationValidationError",
|
||||
"OperationValidatorExtension",
|
||||
"RecallContext",
|
||||
@@ -57,10 +70,20 @@ __all__ = [
|
||||
"RetainContext",
|
||||
"RetainResult",
|
||||
"ValidationResult",
|
||||
# Operation Validator - Consolidation
|
||||
"ConsolidateContext",
|
||||
"ConsolidateResult",
|
||||
# Operation Validator - Mental Model
|
||||
"MentalModelGetContext",
|
||||
"MentalModelGetResult",
|
||||
"MentalModelRefreshContext",
|
||||
"MentalModelRefreshResult",
|
||||
# Tenant/Auth
|
||||
"ApiKeyTenantExtension",
|
||||
"SupabaseTenantExtension",
|
||||
"AuthenticationError",
|
||||
"RequestContext",
|
||||
"Tenant",
|
||||
"TenantContext",
|
||||
"TenantExtension",
|
||||
]
|
||||
|
||||
@@ -6,13 +6,17 @@ They can be used directly or serve as examples for custom implementations.
|
||||
|
||||
Available built-in extensions:
|
||||
- ApiKeyTenantExtension: Simple API key validation with public schema
|
||||
- SupabaseTenantExtension: Supabase JWT validation with per-user schema isolation
|
||||
|
||||
Example usage:
|
||||
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.tenant:ApiKeyTenantExtension
|
||||
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.supabase_tenant:SupabaseTenantExtension
|
||||
"""
|
||||
|
||||
from hindsight_api.extensions.builtin.supabase_tenant import SupabaseTenantExtension
|
||||
from hindsight_api.extensions.builtin.tenant import ApiKeyTenantExtension
|
||||
|
||||
__all__ = [
|
||||
"ApiKeyTenantExtension",
|
||||
"SupabaseTenantExtension",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,433 @@
|
||||
"""
|
||||
Supabase Tenant Extension for Hindsight
|
||||
|
||||
Validates Supabase JWTs and maps authenticated users to isolated memory banks.
|
||||
Each user gets their own PostgreSQL schema based on their Supabase user ID.
|
||||
|
||||
This extension enables multi-tenant memory isolation for applications using
|
||||
Supabase Auth - each authenticated user's memories are stored in a separate
|
||||
schema, ensuring complete data isolation.
|
||||
|
||||
Features:
|
||||
- Local JWT Verification: Validates tokens locally using JWKS public keys
|
||||
(no network call per request)
|
||||
- Automatic Schema Isolation: Each user gets {prefix}_{user_id} schema
|
||||
- Zero User Management: Leverages your existing Supabase Auth setup
|
||||
- Production Ready: Includes health checks, timeouts, key rotation handling,
|
||||
and error handling
|
||||
- Built-in: Ships with Hindsight, no extra installation needed
|
||||
- Legacy Support: Falls back to /auth/v1/user endpoint for HS256 projects
|
||||
|
||||
JWT Verification Strategy:
|
||||
By default, JWTs are verified locally using public keys from the Supabase
|
||||
JWKS endpoint (/auth/v1/.well-known/jwks.json). This is the Supabase-recommended
|
||||
approach: no network call per request, fast, and secure.
|
||||
|
||||
If JWKS keys are unavailable (e.g., legacy HS256 projects), the extension
|
||||
falls back to calling /auth/v1/user per request for validation. This requires
|
||||
the service_role key to be configured.
|
||||
|
||||
Configuration via environment variables:
|
||||
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.supabase_tenant:SupabaseTenantExtension
|
||||
HINDSIGHT_API_TENANT_SUPABASE_URL=https://your-project.supabase.co
|
||||
|
||||
# Optional - only required for legacy HS256 projects or health checks
|
||||
HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY=your-service-role-key
|
||||
|
||||
# Optional
|
||||
HINDSIGHT_API_TENANT_SCHEMA_PREFIX=user # Default: "user" (creates user_<uuid> schemas)
|
||||
|
||||
Usage:
|
||||
Clients pass their Supabase JWT in the Authorization header:
|
||||
|
||||
curl -H "Authorization: Bearer <supabase_jwt>" \\
|
||||
https://your-hindsight-server/v1/default/banks/my-bank/memories/recall
|
||||
|
||||
Author: BrighterBalance (https://brighterbalance.app)
|
||||
License: MIT
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
|
||||
import httpx
|
||||
import jwt as pyjwt
|
||||
from jwt import PyJWK
|
||||
|
||||
from hindsight_api.extensions.tenant import AuthenticationError, Tenant, TenantContext, TenantExtension
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
__all__ = ["SupabaseTenantExtension"]
|
||||
|
||||
# Minimum expected JWT length (JWTs are typically 100+ characters)
|
||||
MIN_TOKEN_LENGTH = 20
|
||||
|
||||
# Timeout for Supabase API calls
|
||||
REQUEST_TIMEOUT_SECONDS = 10.0
|
||||
|
||||
# JWKS cache TTL — Supabase Edge caches JWKS for 10 minutes, so we match that
|
||||
JWKS_CACHE_TTL_SECONDS = 600
|
||||
|
||||
# Minimum interval between JWKS refreshes to avoid hammering the endpoint
|
||||
JWKS_MIN_REFRESH_INTERVAL_SECONDS = 30
|
||||
|
||||
# Algorithms supported by Supabase Auth for asymmetric JWT signing
|
||||
SUPPORTED_ALGORITHMS = ["RS256", "ES256"]
|
||||
|
||||
# Supabase user IDs are UUIDs — validate before using in schema names
|
||||
_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}$", re.IGNORECASE)
|
||||
|
||||
# Schema prefix must be a valid Postgres identifier component (letters, digits, underscores)
|
||||
_SCHEMA_PREFIX_RE = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_]*$")
|
||||
|
||||
|
||||
class SupabaseTenantExtension(TenantExtension):
|
||||
"""
|
||||
TenantExtension that validates Supabase JWTs for multi-tenant isolation.
|
||||
|
||||
Each authenticated user gets their own PostgreSQL schema, ensuring complete
|
||||
memory isolation between users. The schema name is derived from the user's
|
||||
Supabase user ID (the ``sub`` claim in the JWT).
|
||||
|
||||
JWT verification uses JWKS (local, no network call per request) when
|
||||
asymmetric keys are configured in Supabase, and falls back to the
|
||||
``/auth/v1/user`` endpoint for legacy HS256 projects.
|
||||
|
||||
Example:
|
||||
User with ID "a1b2c3d4-e5f6-7890-abcd-ef1234567890"
|
||||
gets schema "user_a1b2c3d4_e5f6_7890_abcd_ef1234567890"
|
||||
"""
|
||||
|
||||
def __init__(self, config: dict[str, str]) -> None:
|
||||
"""
|
||||
Initialize with configuration from environment variables.
|
||||
|
||||
Config keys are derived from HINDSIGHT_API_TENANT_* env vars:
|
||||
- HINDSIGHT_API_TENANT_SUPABASE_URL -> config["supabase_url"] (required)
|
||||
- HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY -> config["supabase_service_key"] (optional)
|
||||
- HINDSIGHT_API_TENANT_SCHEMA_PREFIX -> config["schema_prefix"] (optional)
|
||||
|
||||
Args:
|
||||
config: Dictionary of configuration values from environment
|
||||
|
||||
Raises:
|
||||
ValueError: If required configuration is missing
|
||||
"""
|
||||
super().__init__(config)
|
||||
|
||||
self.supabase_url = (config.get("supabase_url") or "").rstrip("/")
|
||||
self.supabase_service_key = config.get("supabase_service_key")
|
||||
self.schema_prefix = config.get("schema_prefix", "user")
|
||||
|
||||
# Track initialized schemas to avoid redundant migrations
|
||||
self._initialized_schemas: set[str] = set()
|
||||
|
||||
# Reusable HTTP client (created on startup)
|
||||
self._http_client: httpx.AsyncClient | None = None
|
||||
|
||||
# JWKS state
|
||||
self._jwks_keys: dict[str, PyJWK] = {}
|
||||
self._jwks_last_fetched: float = 0
|
||||
self._use_jwks: bool = False
|
||||
|
||||
if not self.supabase_url:
|
||||
raise ValueError(
|
||||
"HINDSIGHT_API_TENANT_SUPABASE_URL is required. "
|
||||
"Set it to your Supabase project URL (e.g., https://xxx.supabase.co)"
|
||||
)
|
||||
|
||||
if not _SCHEMA_PREFIX_RE.match(self.schema_prefix):
|
||||
raise ValueError(
|
||||
f"Invalid schema_prefix '{self.schema_prefix}'. "
|
||||
"Must be a valid Postgres identifier (letters, digits, underscores, starting with a letter or underscore)."
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Lifecycle
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def on_startup(self) -> None:
|
||||
"""
|
||||
Called when Hindsight starts.
|
||||
|
||||
Creates a reusable HTTP client, fetches JWKS for local JWT verification,
|
||||
and optionally verifies connectivity to Supabase.
|
||||
"""
|
||||
logger.info("Initializing Supabase tenant extension")
|
||||
logger.info("Supabase URL: %s", self.supabase_url)
|
||||
logger.info("Schema prefix: %s_", self.schema_prefix)
|
||||
|
||||
self._http_client = httpx.AsyncClient(timeout=REQUEST_TIMEOUT_SECONDS)
|
||||
|
||||
# Attempt to fetch JWKS for fast local JWT verification
|
||||
await self._try_init_jwks()
|
||||
|
||||
# Optional health check using service key
|
||||
if self.supabase_service_key:
|
||||
await self._health_check()
|
||||
|
||||
async def on_shutdown(self) -> None:
|
||||
"""Called when Hindsight shuts down. Closes the HTTP client."""
|
||||
logger.info("Shutting down Supabase tenant extension")
|
||||
if self._http_client:
|
||||
await self._http_client.aclose()
|
||||
self._http_client = None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# JWKS management
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _try_init_jwks(self) -> None:
|
||||
"""Fetch JWKS and decide verification mode (local JWKS vs legacy endpoint)."""
|
||||
try:
|
||||
await self._fetch_jwks()
|
||||
if self._jwks_keys:
|
||||
self._use_jwks = True
|
||||
logger.info(
|
||||
"JWKS loaded — using local JWT verification with %d key(s)",
|
||||
len(self._jwks_keys),
|
||||
)
|
||||
return
|
||||
|
||||
# JWKS endpoint returned no keys — project likely uses legacy HS256
|
||||
logger.warning(
|
||||
"JWKS endpoint returned no signing keys. "
|
||||
"Falling back to /auth/v1/user endpoint for JWT verification. "
|
||||
"For better performance, enable asymmetric JWT signing in your "
|
||||
"Supabase dashboard (Project Settings → Auth → JWT Algorithm)."
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Could not fetch JWKS (%s). Falling back to /auth/v1/user endpoint for JWT verification.",
|
||||
e,
|
||||
)
|
||||
|
||||
# Legacy mode requires service key
|
||||
if not self.supabase_service_key:
|
||||
raise ValueError(
|
||||
"HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY is required when JWKS "
|
||||
"is not available. Either enable asymmetric JWT signing in your "
|
||||
"Supabase project or provide the service_role key."
|
||||
)
|
||||
self._use_jwks = False
|
||||
|
||||
async def _fetch_jwks(self) -> None:
|
||||
"""Fetch public signing keys from the Supabase JWKS endpoint."""
|
||||
if self._http_client is None:
|
||||
raise RuntimeError("HTTP client not initialized")
|
||||
|
||||
url = f"{self.supabase_url}/auth/v1/.well-known/jwks.json"
|
||||
response = await self._http_client.get(url)
|
||||
response.raise_for_status()
|
||||
|
||||
jwks_data = response.json()
|
||||
keys: dict[str, PyJWK] = {}
|
||||
for key_data in jwks_data.get("keys", []):
|
||||
kid = key_data.get("kid")
|
||||
if kid:
|
||||
keys[kid] = PyJWK(key_data)
|
||||
|
||||
self._jwks_keys = keys
|
||||
self._jwks_last_fetched = time.monotonic()
|
||||
|
||||
async def _get_signing_key(self, token: str) -> PyJWK:
|
||||
"""
|
||||
Resolve the signing key for a token from the JWKS cache.
|
||||
|
||||
If the key ID (``kid``) is not in the cache, triggers one JWKS refresh
|
||||
to handle key rotation before raising an error.
|
||||
"""
|
||||
header = pyjwt.get_unverified_header(token)
|
||||
kid = header.get("kid")
|
||||
if not kid:
|
||||
raise AuthenticationError("Token missing key ID (kid) header")
|
||||
|
||||
# Refresh cache if stale
|
||||
now = time.monotonic()
|
||||
if now - self._jwks_last_fetched > JWKS_CACHE_TTL_SECONDS:
|
||||
logger.debug("JWKS cache expired, refreshing")
|
||||
await self._fetch_jwks()
|
||||
|
||||
if kid in self._jwks_keys:
|
||||
return self._jwks_keys[kid]
|
||||
|
||||
# Key not found — try one forced refresh to handle key rotation,
|
||||
# but only if we haven't just refreshed
|
||||
if now - self._jwks_last_fetched > JWKS_MIN_REFRESH_INTERVAL_SECONDS:
|
||||
logger.info("Signing key %s not in cache, refreshing JWKS for possible key rotation", kid)
|
||||
await self._fetch_jwks()
|
||||
if kid in self._jwks_keys:
|
||||
return self._jwks_keys[kid]
|
||||
|
||||
raise AuthenticationError("Unable to find signing key for token")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Authentication
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def authenticate(self, context: RequestContext) -> TenantContext:
|
||||
"""
|
||||
Validate a Supabase JWT and return tenant context.
|
||||
|
||||
Uses local JWKS verification when available (no network call per
|
||||
request), falling back to the ``/auth/v1/user`` endpoint for legacy
|
||||
HS256 projects.
|
||||
|
||||
Args:
|
||||
context: Request context containing the API key (JWT)
|
||||
|
||||
Returns:
|
||||
TenantContext with schema_name set to ``{prefix}_{user_uuid}``
|
||||
|
||||
Raises:
|
||||
AuthenticationError: If token is missing, invalid, or expired
|
||||
"""
|
||||
token = context.api_key
|
||||
|
||||
if not token:
|
||||
raise AuthenticationError("Missing Authorization header. Expected: Bearer <supabase_jwt>")
|
||||
|
||||
if len(token) < MIN_TOKEN_LENGTH:
|
||||
raise AuthenticationError("Invalid token format")
|
||||
|
||||
if self._http_client is None:
|
||||
raise AuthenticationError("Extension not initialized")
|
||||
|
||||
# Verify the JWT and extract user ID
|
||||
if self._use_jwks:
|
||||
user_id = await self._verify_token_jwks(token)
|
||||
else:
|
||||
user_id = await self._verify_token_legacy(token)
|
||||
|
||||
# Validate user ID format before using in schema name
|
||||
if not _UUID_RE.match(user_id):
|
||||
raise AuthenticationError("Invalid user ID format in token")
|
||||
|
||||
# Build isolated schema name — hyphens to underscores for Postgres compatibility
|
||||
safe_user_id = user_id.replace("-", "_")
|
||||
schema_name = f"{self.schema_prefix}_{safe_user_id}"
|
||||
|
||||
# Initialize schema on first access
|
||||
if schema_name not in self._initialized_schemas:
|
||||
await self._initialize_schema(schema_name)
|
||||
|
||||
return TenantContext(schema_name=schema_name)
|
||||
|
||||
async def _verify_token_jwks(self, token: str) -> str:
|
||||
"""
|
||||
Verify a JWT locally using cached JWKS public keys.
|
||||
|
||||
Validates signature, expiration, issuer, and audience. Returns the
|
||||
user ID from the ``sub`` claim.
|
||||
|
||||
Raises:
|
||||
AuthenticationError: If the token is invalid or expired.
|
||||
"""
|
||||
try:
|
||||
signing_key = await self._get_signing_key(token)
|
||||
payload = pyjwt.decode(
|
||||
token,
|
||||
signing_key.key,
|
||||
algorithms=SUPPORTED_ALGORITHMS,
|
||||
audience="authenticated",
|
||||
issuer=f"{self.supabase_url}/auth/v1",
|
||||
)
|
||||
except pyjwt.ExpiredSignatureError:
|
||||
raise AuthenticationError("Token has expired")
|
||||
except pyjwt.InvalidAudienceError:
|
||||
raise AuthenticationError("Invalid token audience")
|
||||
except pyjwt.InvalidIssuerError:
|
||||
raise AuthenticationError("Invalid token issuer")
|
||||
except pyjwt.DecodeError:
|
||||
raise AuthenticationError("Invalid token")
|
||||
except AuthenticationError:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise AuthenticationError(f"Token verification failed: {e!s}")
|
||||
|
||||
user_id = payload.get("sub")
|
||||
if not user_id:
|
||||
raise AuthenticationError("Token valid but missing subject (sub) claim")
|
||||
return user_id
|
||||
|
||||
async def _verify_token_legacy(self, token: str) -> str:
|
||||
"""
|
||||
Verify a JWT by calling the Supabase ``/auth/v1/user`` endpoint.
|
||||
|
||||
This is the fallback for projects using legacy HS256 JWT signing.
|
||||
Adds a network round-trip per request.
|
||||
|
||||
Raises:
|
||||
AuthenticationError: If the token is invalid or the request fails.
|
||||
"""
|
||||
try:
|
||||
response = await self._http_client.get(
|
||||
f"{self.supabase_url}/auth/v1/user",
|
||||
headers={
|
||||
"Authorization": f"Bearer {token}",
|
||||
"apikey": self.supabase_service_key,
|
||||
},
|
||||
)
|
||||
|
||||
if response.status_code == 401:
|
||||
raise AuthenticationError("Invalid or expired token")
|
||||
|
||||
if response.status_code != 200:
|
||||
raise AuthenticationError(f"Authentication failed: {response.status_code}")
|
||||
|
||||
user_data = response.json()
|
||||
user_id = user_data.get("id")
|
||||
|
||||
if not user_id:
|
||||
raise AuthenticationError("Token valid but no user ID found")
|
||||
|
||||
return user_id
|
||||
|
||||
except AuthenticationError:
|
||||
raise
|
||||
except httpx.TimeoutException:
|
||||
raise AuthenticationError("Authentication timeout - please retry")
|
||||
except httpx.RequestError as e:
|
||||
raise AuthenticationError(f"Connection error: {e!s}")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Schema management
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _initialize_schema(self, schema_name: str) -> None:
|
||||
"""Run migrations for a new tenant schema and cache the result."""
|
||||
logger.info("Initializing schema: %s", schema_name)
|
||||
try:
|
||||
await self.context.run_migration(schema_name)
|
||||
self._initialized_schemas.add(schema_name)
|
||||
logger.info("Schema ready: %s", schema_name)
|
||||
except Exception as e:
|
||||
logger.error("Schema initialization failed for %s: %s", schema_name, e)
|
||||
raise AuthenticationError(f"Failed to initialize tenant: {e!s}")
|
||||
|
||||
async def list_tenants(self) -> list[Tenant]:
|
||||
"""Return all tenant schemas that have been initialized."""
|
||||
return [Tenant(schema=schema) for schema in self._initialized_schemas]
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Health check
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _health_check(self) -> None:
|
||||
"""Verify connectivity to Supabase using the auth health endpoint."""
|
||||
try:
|
||||
response = await self._http_client.get(
|
||||
f"{self.supabase_url}/auth/v1/health",
|
||||
headers={"apikey": self.supabase_service_key},
|
||||
)
|
||||
if response.status_code == 200:
|
||||
logger.info("Supabase connection verified")
|
||||
else:
|
||||
logger.warning("Supabase health check returned %d", response.status_code)
|
||||
except Exception as e:
|
||||
logger.warning("Could not verify Supabase connection: %s", e)
|
||||
@@ -1,20 +1,60 @@
|
||||
"""Built-in tenant extension implementations."""
|
||||
|
||||
from hindsight_api.extensions.tenant import AuthenticationError, TenantContext, TenantExtension
|
||||
from hindsight_api.config import get_config
|
||||
from hindsight_api.extensions.tenant import AuthenticationError, Tenant, TenantContext, TenantExtension
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
|
||||
class DefaultTenantExtension(TenantExtension):
|
||||
"""
|
||||
Default single-tenant extension with no authentication.
|
||||
|
||||
This is the default extension used when no tenant extension is configured.
|
||||
It provides single-tenant behavior using the configured schema from
|
||||
HINDSIGHT_API_DATABASE_SCHEMA (defaults to 'public').
|
||||
|
||||
Features:
|
||||
- No authentication required (passes all requests)
|
||||
- Uses configured schema from environment
|
||||
- Perfect for single-tenant deployments without auth
|
||||
|
||||
Configuration:
|
||||
HINDSIGHT_API_DATABASE_SCHEMA=your-schema (optional, defaults to 'public')
|
||||
|
||||
This is automatically enabled by default. To use custom authentication,
|
||||
configure a different tenant extension:
|
||||
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.tenant:ApiKeyTenantExtension
|
||||
"""
|
||||
|
||||
def __init__(self, config: dict[str, str]):
|
||||
super().__init__(config)
|
||||
# Cache the schema at initialization for consistency
|
||||
# Support explicit schema override via config, otherwise use environment
|
||||
self._schema = config.get("schema", get_config().database_schema)
|
||||
|
||||
async def authenticate(self, context: RequestContext) -> TenantContext:
|
||||
"""Return configured schema without any authentication."""
|
||||
return TenantContext(schema_name=self._schema)
|
||||
|
||||
async def list_tenants(self) -> list[Tenant]:
|
||||
"""Return configured schema for single-tenant setup."""
|
||||
return [Tenant(schema=self._schema)]
|
||||
|
||||
|
||||
class ApiKeyTenantExtension(TenantExtension):
|
||||
"""
|
||||
Built-in tenant extension that validates API key against an environment variable.
|
||||
|
||||
This is a simple implementation that:
|
||||
1. Validates the API key matches HINDSIGHT_API_TENANT_API_KEY
|
||||
2. Returns 'public' as the schema for all authenticated requests
|
||||
2. Returns the configured schema (HINDSIGHT_API_DATABASE_SCHEMA, default 'public')
|
||||
for all authenticated requests
|
||||
|
||||
Configuration:
|
||||
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.tenant:ApiKeyTenantExtension
|
||||
HINDSIGHT_API_TENANT_API_KEY=your-secret-key
|
||||
HINDSIGHT_API_DATABASE_SCHEMA=your-schema (optional, defaults to 'public')
|
||||
HINDSIGHT_API_TENANT_MCP_AUTH_DISABLED=true (optional, disable auth for MCP endpoints)
|
||||
|
||||
For multi-tenant setups with separate schemas per tenant, implement a custom
|
||||
TenantExtension that looks up the schema based on the API key or token claims.
|
||||
@@ -25,9 +65,26 @@ class ApiKeyTenantExtension(TenantExtension):
|
||||
self.expected_api_key = config.get("api_key")
|
||||
if not self.expected_api_key:
|
||||
raise ValueError("HINDSIGHT_API_TENANT_API_KEY is required when using ApiKeyTenantExtension")
|
||||
# Allow disabling MCP auth for backwards compatibility
|
||||
self.mcp_auth_disabled = config.get("mcp_auth_disabled", "").lower() in ("true", "1", "yes")
|
||||
|
||||
async def authenticate(self, context: RequestContext) -> TenantContext:
|
||||
"""Validate API key and return public schema context."""
|
||||
"""Validate API key and return configured schema context."""
|
||||
if context.api_key != self.expected_api_key:
|
||||
raise AuthenticationError("Invalid API key")
|
||||
return TenantContext(schema_name="public")
|
||||
return TenantContext(schema_name=get_config().database_schema)
|
||||
|
||||
async def list_tenants(self) -> list[Tenant]:
|
||||
"""Return configured schema for single-tenant setup."""
|
||||
return [Tenant(schema=get_config().database_schema)]
|
||||
|
||||
async def authenticate_mcp(self, context: RequestContext) -> TenantContext:
|
||||
"""
|
||||
Authenticate MCP requests.
|
||||
|
||||
If mcp_auth_disabled is set, skip authentication for backwards compatibility.
|
||||
Otherwise, delegate to authenticate().
|
||||
"""
|
||||
if self.mcp_auth_disabled:
|
||||
return TenantContext(schema_name=get_config().database_schema)
|
||||
return await self.authenticate(context)
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
"""MCP Extension for registering additional MCP tools.
|
||||
|
||||
This extension allows external packages (like hindsight-cloud) to register
|
||||
additional MCP tools on the Hindsight MCP server.
|
||||
|
||||
Example:
|
||||
HINDSIGHT_API_MCP_EXTENSION=hindsight_cloud.extensions:CloudMCPExtension
|
||||
"""
|
||||
|
||||
import logging
|
||||
from abc import abstractmethod
|
||||
|
||||
from fastmcp import FastMCP
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
from hindsight_api.extensions.base import Extension
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MCPExtension(Extension):
|
||||
"""Base class for MCP extensions that register additional tools.
|
||||
|
||||
Subclass this to add MCP tools in extension packages.
|
||||
|
||||
Example:
|
||||
class CloudMCPExtension(MCPExtension):
|
||||
def register_tools(self, mcp: FastMCP, memory: MemoryEngine) -> None:
|
||||
@mcp.tool()
|
||||
async def my_custom_tool(query: str) -> str:
|
||||
return "result"
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def register_tools(self, mcp: FastMCP, memory: MemoryEngine) -> None:
|
||||
"""Register additional MCP tools.
|
||||
|
||||
Args:
|
||||
mcp: FastMCP server instance to register tools on
|
||||
memory: MemoryEngine instance for accessing memory operations
|
||||
"""
|
||||
pass
|
||||
@@ -1,4 +1,4 @@
|
||||
"""Operation Validator Extension for validating retain/recall/reflect operations."""
|
||||
"""Operation Validator Extension for validating retain/recall/reflect/consolidate operations."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
@@ -97,6 +97,19 @@ class ReflectContext:
|
||||
context: str | None = None
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Consolidation Pre-operation Context
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConsolidateContext:
|
||||
"""Context for a consolidation operation validation (pre-operation)."""
|
||||
|
||||
bank_id: str
|
||||
request_context: "RequestContext"
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Post-operation Contexts (includes results)
|
||||
# =============================================================================
|
||||
@@ -164,9 +177,79 @@ class ReflectResultContext:
|
||||
error: str | None = None
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Consolidation Post-operation Context
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConsolidateResult:
|
||||
"""Result context for post-consolidation hook."""
|
||||
|
||||
bank_id: str
|
||||
request_context: "RequestContext"
|
||||
# Result
|
||||
processed: int = 0
|
||||
created: int = 0
|
||||
updated: int = 0
|
||||
success: bool = True
|
||||
error: str | None = None
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Mental Model Contexts
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@dataclass
|
||||
class MentalModelGetContext:
|
||||
"""Context for a mental model GET operation validation (pre-operation)."""
|
||||
|
||||
bank_id: str
|
||||
mental_model_id: str
|
||||
request_context: "RequestContext"
|
||||
|
||||
|
||||
@dataclass
|
||||
class MentalModelRefreshContext:
|
||||
"""Context for a mental model refresh/create operation validation (pre-operation)."""
|
||||
|
||||
bank_id: str
|
||||
mental_model_id: str | None # None for create (not yet assigned)
|
||||
request_context: "RequestContext"
|
||||
|
||||
|
||||
@dataclass
|
||||
class MentalModelGetResult:
|
||||
"""Result context for post-mental-model-GET hook."""
|
||||
|
||||
bank_id: str
|
||||
mental_model_id: str
|
||||
request_context: "RequestContext"
|
||||
output_tokens: int # tokens in the returned content
|
||||
success: bool = True
|
||||
error: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class MentalModelRefreshResult:
|
||||
"""Result context for post-mental-model-refresh hook."""
|
||||
|
||||
bank_id: str
|
||||
mental_model_id: str
|
||||
request_context: "RequestContext"
|
||||
query_tokens: int # tokens in source_query
|
||||
output_tokens: int # tokens in generated content
|
||||
context_tokens: int # tokens in context (if any)
|
||||
facts_used: int # facts referenced in based_on
|
||||
mental_models_used: int # mental models referenced in based_on
|
||||
success: bool = True
|
||||
error: str | None = None
|
||||
|
||||
|
||||
class OperationValidatorExtension(Extension, ABC):
|
||||
"""
|
||||
Validates and hooks into retain/recall/reflect operations.
|
||||
Validates and hooks into retain/recall/reflect/consolidate operations.
|
||||
|
||||
This extension allows implementing custom logic such as:
|
||||
- Rate limiting (pre-operation)
|
||||
@@ -185,9 +268,13 @@ class OperationValidatorExtension(Extension, ABC):
|
||||
-> config = {"max_requests": "100"}
|
||||
|
||||
Hook execution order:
|
||||
1. validate_retain/validate_recall/validate_reflect (pre-operation)
|
||||
1. validate_* (pre-operation)
|
||||
2. [operation executes]
|
||||
3. on_retain_complete/on_recall_complete/on_reflect_complete (post-operation)
|
||||
3. on_*_complete (post-operation)
|
||||
|
||||
Supported operations:
|
||||
- retain, recall, reflect (core memory operations)
|
||||
- consolidate (mental models consolidation)
|
||||
"""
|
||||
|
||||
# =========================================================================
|
||||
@@ -325,3 +412,122 @@ class OperationValidatorExtension(Extension, ABC):
|
||||
- error: Error message (if failed)
|
||||
"""
|
||||
pass
|
||||
|
||||
# =========================================================================
|
||||
# Consolidation - Pre-operation validation hook (optional - override to implement)
|
||||
# =========================================================================
|
||||
|
||||
async def validate_consolidate(self, ctx: ConsolidateContext) -> ValidationResult:
|
||||
"""
|
||||
Validate a consolidation operation before execution.
|
||||
|
||||
Override to implement custom validation logic for consolidation.
|
||||
|
||||
Args:
|
||||
ctx: Context containing:
|
||||
- bank_id: Bank identifier
|
||||
- request_context: Request context with auth info
|
||||
|
||||
Returns:
|
||||
ValidationResult indicating whether the operation is allowed.
|
||||
"""
|
||||
return ValidationResult.accept()
|
||||
|
||||
# =========================================================================
|
||||
# Consolidation - Post-operation hook (optional - override to implement)
|
||||
# =========================================================================
|
||||
|
||||
async def on_consolidate_complete(self, result: ConsolidateResult) -> None:
|
||||
"""
|
||||
Called after a consolidation operation completes (success or failure).
|
||||
|
||||
Override to implement post-operation logic such as usage tracking or audit logging.
|
||||
|
||||
Args:
|
||||
result: Result context containing:
|
||||
- bank_id: Bank identifier
|
||||
- processed: Number of memories processed
|
||||
- created: Number of mental models created
|
||||
- updated: Number of mental models updated
|
||||
- success: Whether the operation succeeded
|
||||
- error: Error message (if failed)
|
||||
"""
|
||||
pass
|
||||
|
||||
# =========================================================================
|
||||
# Mental Model - Pre-operation validation hook (optional - override to implement)
|
||||
# =========================================================================
|
||||
|
||||
async def validate_mental_model_get(self, ctx: MentalModelGetContext) -> ValidationResult:
|
||||
"""
|
||||
Validate a mental model GET operation before execution.
|
||||
|
||||
Override to implement custom validation logic for mental model retrieval.
|
||||
|
||||
Args:
|
||||
ctx: Context containing:
|
||||
- bank_id: Bank identifier
|
||||
- mental_model_id: Mental model identifier
|
||||
- request_context: Request context with auth info
|
||||
|
||||
Returns:
|
||||
ValidationResult indicating whether the operation is allowed.
|
||||
"""
|
||||
return ValidationResult.accept()
|
||||
|
||||
async def validate_mental_model_refresh(self, ctx: MentalModelRefreshContext) -> ValidationResult:
|
||||
"""
|
||||
Validate a mental model refresh/create operation before execution.
|
||||
|
||||
Override to implement custom validation logic for mental model refresh.
|
||||
|
||||
Args:
|
||||
ctx: Context containing:
|
||||
- bank_id: Bank identifier
|
||||
- mental_model_id: Mental model identifier (None for create)
|
||||
- request_context: Request context with auth info
|
||||
|
||||
Returns:
|
||||
ValidationResult indicating whether the operation is allowed.
|
||||
"""
|
||||
return ValidationResult.accept()
|
||||
|
||||
# =========================================================================
|
||||
# Mental Model - Post-operation hooks (optional - override to implement)
|
||||
# =========================================================================
|
||||
|
||||
async def on_mental_model_get_complete(self, result: MentalModelGetResult) -> None:
|
||||
"""
|
||||
Called after a mental model GET operation completes (success or failure).
|
||||
|
||||
Override to implement post-operation logic such as tracking or audit logging.
|
||||
|
||||
Args:
|
||||
result: Result context containing:
|
||||
- bank_id: Bank identifier
|
||||
- mental_model_id: Mental model identifier
|
||||
- output_tokens: Token count of the returned content
|
||||
- success: Whether the operation succeeded
|
||||
- error: Error message (if failed)
|
||||
"""
|
||||
pass
|
||||
|
||||
async def on_mental_model_refresh_complete(self, result: MentalModelRefreshResult) -> None:
|
||||
"""
|
||||
Called after a mental model refresh operation completes (success or failure).
|
||||
|
||||
Override to implement post-operation logic such as tracking or audit logging.
|
||||
|
||||
Args:
|
||||
result: Result context containing:
|
||||
- bank_id: Bank identifier
|
||||
- mental_model_id: Mental model identifier
|
||||
- query_tokens: Tokens in source_query
|
||||
- output_tokens: Tokens in generated content
|
||||
- context_tokens: Tokens in context
|
||||
- facts_used: Number of facts referenced
|
||||
- mental_models_used: Number of mental models referenced
|
||||
- success: Whether the operation succeeded
|
||||
- error: Error message (if failed)
|
||||
"""
|
||||
pass
|
||||
|
||||
@@ -28,6 +28,18 @@ class TenantContext:
|
||||
schema_name: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class Tenant:
|
||||
"""
|
||||
Represents a tenant for worker discovery.
|
||||
|
||||
Used by list_tenants() to return tenant information including
|
||||
the PostgreSQL schema name for database operations.
|
||||
"""
|
||||
|
||||
schema: str
|
||||
|
||||
|
||||
class TenantExtension(Extension, ABC):
|
||||
"""
|
||||
Extension for multi-tenancy and API key authentication.
|
||||
@@ -61,3 +73,36 @@ class TenantExtension(Extension, ABC):
|
||||
AuthenticationError: If authentication fails.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def list_tenants(self) -> list[Tenant]:
|
||||
"""
|
||||
List all tenants that should be processed by workers.
|
||||
|
||||
This method is used by the worker to discover all tenants that need
|
||||
task polling. Workers will poll for pending tasks in each tenant's schema.
|
||||
|
||||
Returns:
|
||||
List of Tenant objects containing schema information.
|
||||
For single-tenant setups, return [Tenant(schema="public")].
|
||||
"""
|
||||
...
|
||||
|
||||
async def authenticate_mcp(self, context: RequestContext) -> TenantContext:
|
||||
"""
|
||||
Authenticate MCP requests.
|
||||
|
||||
By default, this calls authenticate(). Override this method to provide
|
||||
different authentication behavior for MCP endpoints (e.g., to disable
|
||||
auth for backwards compatibility with existing MCP servers).
|
||||
|
||||
Args:
|
||||
context: The action context containing API key and other auth data.
|
||||
|
||||
Returns:
|
||||
TenantContext with the schema_name for database operations.
|
||||
|
||||
Raises:
|
||||
AuthenticationError: If authentication fails.
|
||||
"""
|
||||
return await self.authenticate(context)
|
||||
|
||||
@@ -20,14 +20,13 @@ import warnings
|
||||
|
||||
import uvicorn
|
||||
|
||||
from . import MemoryEngine
|
||||
from . import MemoryEngine, __version__
|
||||
from .api import create_app
|
||||
from .banner import print_banner
|
||||
from .config import DEFAULT_WORKERS, ENV_WORKERS, HindsightConfig, get_config
|
||||
from .daemon import (
|
||||
DEFAULT_DAEMON_PORT,
|
||||
DEFAULT_IDLE_TIMEOUT,
|
||||
DaemonLock,
|
||||
IdleTimeoutMiddleware,
|
||||
daemonize,
|
||||
)
|
||||
@@ -136,30 +135,15 @@ def main():
|
||||
|
||||
# Daemon mode handling
|
||||
if args.daemon:
|
||||
# Use fixed daemon port
|
||||
args.port = DEFAULT_DAEMON_PORT
|
||||
# Use port from args (may be custom for profiles)
|
||||
if args.port == config.port: # No custom port specified
|
||||
args.port = DEFAULT_DAEMON_PORT
|
||||
args.host = "127.0.0.1" # Only bind to localhost for security
|
||||
|
||||
# Check if another daemon is already running
|
||||
daemon_lock = DaemonLock()
|
||||
if not daemon_lock.acquire():
|
||||
print(f"Daemon already running (PID: {daemon_lock.get_pid()})", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
# Fork into background
|
||||
# No lockfile needed - port binding prevents duplicate daemons
|
||||
daemonize()
|
||||
|
||||
# Re-acquire lock in child process
|
||||
daemon_lock = DaemonLock()
|
||||
if not daemon_lock.acquire():
|
||||
sys.exit(1)
|
||||
|
||||
# Register cleanup to release lock
|
||||
def release_lock():
|
||||
daemon_lock.release()
|
||||
|
||||
atexit.register(release_lock)
|
||||
|
||||
# Print banner (not in daemon mode)
|
||||
if not args.daemon:
|
||||
print()
|
||||
@@ -170,27 +154,56 @@ def main():
|
||||
if args.log_level != config.log_level:
|
||||
config = HindsightConfig(
|
||||
database_url=config.database_url,
|
||||
database_schema=config.database_schema,
|
||||
llm_provider=config.llm_provider,
|
||||
llm_api_key=config.llm_api_key,
|
||||
llm_model=config.llm_model,
|
||||
llm_base_url=config.llm_base_url,
|
||||
llm_max_concurrent=config.llm_max_concurrent,
|
||||
llm_max_retries=config.llm_max_retries,
|
||||
llm_initial_backoff=config.llm_initial_backoff,
|
||||
llm_max_backoff=config.llm_max_backoff,
|
||||
llm_timeout=config.llm_timeout,
|
||||
llm_vertexai_project_id=config.llm_vertexai_project_id,
|
||||
llm_vertexai_region=config.llm_vertexai_region,
|
||||
llm_vertexai_service_account_key=config.llm_vertexai_service_account_key,
|
||||
retain_llm_provider=config.retain_llm_provider,
|
||||
retain_llm_api_key=config.retain_llm_api_key,
|
||||
retain_llm_model=config.retain_llm_model,
|
||||
retain_llm_base_url=config.retain_llm_base_url,
|
||||
retain_llm_max_concurrent=config.retain_llm_max_concurrent,
|
||||
retain_llm_max_retries=config.retain_llm_max_retries,
|
||||
retain_llm_initial_backoff=config.retain_llm_initial_backoff,
|
||||
retain_llm_max_backoff=config.retain_llm_max_backoff,
|
||||
retain_llm_timeout=config.retain_llm_timeout,
|
||||
reflect_llm_provider=config.reflect_llm_provider,
|
||||
reflect_llm_api_key=config.reflect_llm_api_key,
|
||||
reflect_llm_model=config.reflect_llm_model,
|
||||
reflect_llm_base_url=config.reflect_llm_base_url,
|
||||
reflect_llm_max_concurrent=config.reflect_llm_max_concurrent,
|
||||
reflect_llm_max_retries=config.reflect_llm_max_retries,
|
||||
reflect_llm_initial_backoff=config.reflect_llm_initial_backoff,
|
||||
reflect_llm_max_backoff=config.reflect_llm_max_backoff,
|
||||
reflect_llm_timeout=config.reflect_llm_timeout,
|
||||
consolidation_llm_provider=config.consolidation_llm_provider,
|
||||
consolidation_llm_api_key=config.consolidation_llm_api_key,
|
||||
consolidation_llm_model=config.consolidation_llm_model,
|
||||
consolidation_llm_base_url=config.consolidation_llm_base_url,
|
||||
consolidation_llm_max_concurrent=config.consolidation_llm_max_concurrent,
|
||||
consolidation_llm_max_retries=config.consolidation_llm_max_retries,
|
||||
consolidation_llm_initial_backoff=config.consolidation_llm_initial_backoff,
|
||||
consolidation_llm_max_backoff=config.consolidation_llm_max_backoff,
|
||||
consolidation_llm_timeout=config.consolidation_llm_timeout,
|
||||
embeddings_provider=config.embeddings_provider,
|
||||
embeddings_local_model=config.embeddings_local_model,
|
||||
embeddings_local_force_cpu=config.embeddings_local_force_cpu,
|
||||
embeddings_tei_url=config.embeddings_tei_url,
|
||||
embeddings_openai_base_url=config.embeddings_openai_base_url,
|
||||
embeddings_cohere_base_url=config.embeddings_cohere_base_url,
|
||||
reranker_provider=config.reranker_provider,
|
||||
reranker_local_model=config.reranker_local_model,
|
||||
reranker_local_force_cpu=config.reranker_local_force_cpu,
|
||||
reranker_local_max_concurrent=config.reranker_local_max_concurrent,
|
||||
reranker_tei_url=config.reranker_tei_url,
|
||||
reranker_tei_batch_size=config.reranker_tei_batch_size,
|
||||
reranker_tei_max_concurrent=config.reranker_tei_max_concurrent,
|
||||
@@ -199,18 +212,20 @@ def main():
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
log_level=args.log_level,
|
||||
log_format=config.log_format,
|
||||
mcp_enabled=config.mcp_enabled,
|
||||
graph_retriever=config.graph_retriever,
|
||||
mpfp_top_k_neighbors=config.mpfp_top_k_neighbors,
|
||||
recall_max_concurrent=config.recall_max_concurrent,
|
||||
recall_connection_budget=config.recall_connection_budget,
|
||||
observation_min_facts=config.observation_min_facts,
|
||||
observation_top_entities=config.observation_top_entities,
|
||||
retain_max_completion_tokens=config.retain_max_completion_tokens,
|
||||
retain_chunk_size=config.retain_chunk_size,
|
||||
retain_extract_causal_links=config.retain_extract_causal_links,
|
||||
retain_extraction_mode=config.retain_extraction_mode,
|
||||
retain_observations_async=config.retain_observations_async,
|
||||
retain_custom_instructions=config.retain_custom_instructions,
|
||||
enable_observations=config.enable_observations,
|
||||
consolidation_batch_size=config.consolidation_batch_size,
|
||||
consolidation_max_tokens=config.consolidation_max_tokens,
|
||||
skip_llm_verification=config.skip_llm_verification,
|
||||
lazy_reranker=config.lazy_reranker,
|
||||
run_migrations_on_startup=config.run_migrations_on_startup,
|
||||
@@ -218,9 +233,15 @@ def main():
|
||||
db_pool_max_size=config.db_pool_max_size,
|
||||
db_command_timeout=config.db_command_timeout,
|
||||
db_acquire_timeout=config.db_acquire_timeout,
|
||||
task_backend=config.task_backend,
|
||||
task_backend_memory_batch_size=config.task_backend_memory_batch_size,
|
||||
task_backend_memory_batch_interval=config.task_backend_memory_batch_interval,
|
||||
worker_enabled=config.worker_enabled,
|
||||
worker_id=config.worker_id,
|
||||
worker_poll_interval_ms=config.worker_poll_interval_ms,
|
||||
worker_max_retries=config.worker_max_retries,
|
||||
worker_http_port=config.worker_http_port,
|
||||
worker_max_slots=config.worker_max_slots,
|
||||
worker_consolidation_max_slots=config.worker_consolidation_max_slots,
|
||||
reflect_max_iterations=config.reflect_max_iterations,
|
||||
mental_model_refresh_concurrency=config.mental_model_refresh_concurrency,
|
||||
)
|
||||
config.configure_logging()
|
||||
if not args.daemon:
|
||||
@@ -325,11 +346,13 @@ def main():
|
||||
embeddings_provider=config.embeddings_provider,
|
||||
reranker_provider=config.reranker_provider,
|
||||
mcp_enabled=config.mcp_enabled,
|
||||
version=__version__,
|
||||
)
|
||||
|
||||
# Start idle checker in daemon mode
|
||||
if idle_middleware is not None:
|
||||
# Start the idle checker in a background thread with its own event loop
|
||||
import logging
|
||||
import threading
|
||||
|
||||
def run_idle_checker():
|
||||
@@ -340,12 +363,12 @@ def main():
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
loop.run_until_complete(idle_middleware._check_idle())
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as e:
|
||||
logging.error(f"Idle checker error: {e}", exc_info=True)
|
||||
|
||||
threading.Thread(target=run_idle_checker, daemon=True).start()
|
||||
|
||||
uvicorn.run(**uvicorn_config) # type: ignore[invalid-argument-type] - dict kwargs
|
||||
uvicorn.run(**uvicorn_config)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user