Compare commits
129
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
44f7940f47 | ||
|
|
a03fd32214 | ||
|
|
c935c576f1 | ||
|
|
16b85a4faa | ||
|
|
4c792400c1 | ||
|
|
0284595909 | ||
|
|
fe4ed1db73 | ||
|
|
bac4b24e30 | ||
|
|
3290f4bfff | ||
|
|
63a65d0723 | ||
|
|
870cfccabb | ||
|
|
4476a10aa3 | ||
|
|
4f2833873c | ||
|
|
1eeced3116 | ||
|
|
55c216e069 | ||
|
|
e64d3634a9 | ||
|
|
70ce979fbe | ||
|
|
de132501c6 | ||
|
|
a75dcfebf5 | ||
|
|
20c8f8b06a | ||
|
|
f5f3fca4ad | ||
|
|
d47c8a28cc | ||
|
|
1ffc2a418c | ||
|
|
fa53917c63 | ||
|
|
59913086be | ||
|
|
7935b0accd | ||
|
|
26bf5714cd | ||
|
|
6232e690fc | ||
|
|
4135a6cee5 | ||
|
|
eb2702bcba | ||
|
|
0d0abaaa9f | ||
|
|
a6798f7e2a | ||
|
|
fb31a35a86 | ||
|
|
ba99b4422a | ||
|
|
6fe93140a7 | ||
|
|
d6ff191198 | ||
|
|
3bb6a38b5c | ||
|
|
b5df8657e8 | ||
|
|
1dacd0e904 | ||
|
|
4b82d2d7ec | ||
|
|
33fac2c5e2 | ||
|
|
49e233cdb7 | ||
|
|
e6709d541f | ||
|
|
9fd567984c | ||
|
|
c65c6a9dc0 | ||
|
|
4de0730c40 | ||
|
|
5e1f13e4f2 | ||
|
|
67c1a4295f | ||
|
|
37fc7fb8bd | ||
|
|
29a542dc23 | ||
|
|
ecc1f31996 | ||
|
|
233bd2e5d4 | ||
|
|
b3becb6e9a | ||
|
|
67b273de69 | ||
|
|
5a3090b5e5 | ||
|
|
2a00df0bc0 | ||
|
|
7715a5110e | ||
|
|
c06d9b4e4f | ||
|
|
39e3f7c528 | ||
|
|
d899d1890d | ||
|
|
70de23ed85 | ||
|
|
1984936150 | ||
|
|
4f21886a0e | ||
|
|
5e65691743 | ||
|
|
76fd052b3a | ||
|
|
6b5f593dca | ||
|
|
dd59bc8ef9 | ||
|
|
eea0f27118 | ||
|
|
964537f885 | ||
|
|
1a620697b1 | ||
|
|
ce45d301ce | ||
|
|
d49e8201b4 | ||
|
|
c8c7603580 | ||
|
|
787ed60763 | ||
|
|
6b78f7d949 | ||
|
|
54e2df0baf | ||
|
|
967e586e01 | ||
|
|
dfa7cec05b | ||
|
|
36e48a7166 | ||
|
|
786b1ecbbd | ||
|
|
f14f277692 | ||
|
|
c9f3657de6 | ||
|
|
0ae0374dc8 | ||
|
|
f7ff32d49d | ||
|
|
e06a6120a3 | ||
|
|
e599346e59 | ||
|
|
0b352d1bfa | ||
|
|
c882511f10 | ||
|
|
234d426499 | ||
|
|
e6511e7d77 | ||
|
|
904ea4de24 | ||
|
|
6168a77846 | ||
|
|
da44a5e839 | ||
|
|
32bca12c6f | ||
|
|
26850a0156 | ||
|
|
2a0c490c9e | ||
|
|
a831a7b77b | ||
|
|
d405b4feed | ||
|
|
b94b5cf26e | ||
|
|
6d820ef91b | ||
|
|
cf8882a867 | ||
|
|
490fccdc6f | ||
|
|
2948cb62d2 | ||
|
|
9053a51a88 | ||
|
|
f2c28cfd98 | ||
|
|
67fc532c43 | ||
|
|
9474f950f2 | ||
|
|
6a0c034f5d | ||
|
|
b52eb905ad | ||
|
|
1c6acc3ba0 | ||
|
|
8ecb5d3a0c | ||
|
|
ae80876671 | ||
|
|
476a62da47 | ||
|
|
5aaa769ab9 | ||
|
|
04f01ab9ab | ||
|
|
63f51385c4 | ||
|
|
e468a4e19f | ||
|
|
c0a0f447b7 | ||
|
|
84927ccc99 | ||
|
|
a6e8944ff0 | ||
|
|
f6d890f6ed | ||
|
|
1fa8d9150c | ||
|
|
656777c2be | ||
|
|
b36807ad3b | ||
|
|
11ac9cd9a5 | ||
|
|
9394cf92f2 | ||
|
|
47be07f97f | ||
|
|
bb1f9cb221 | ||
|
|
7dd68538bb |
@@ -2,11 +2,23 @@
|
||||
# Copy this file to .env and fill in your values
|
||||
|
||||
# LLM Configuration (Required)
|
||||
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio
|
||||
HINDSIGHT_API_LLM_PROVIDER=openai
|
||||
HINDSIGHT_API_LLM_API_KEY=your-api-key-here
|
||||
HINDSIGHT_API_LLM_MODEL=o3-mini
|
||||
HINDSIGHT_API_LLM_BASE_URL=https://api.openai.com/v1
|
||||
|
||||
# Example: Anthropic Claude configuration
|
||||
# HINDSIGHT_API_LLM_PROVIDER=anthropic
|
||||
# HINDSIGHT_API_LLM_API_KEY=your-anthropic-api-key
|
||||
# HINDSIGHT_API_LLM_MODEL=claude-sonnet-4-20250514
|
||||
|
||||
# Example: LM Studio local configuration (Qwen 2.5 32B recommended)
|
||||
# HINDSIGHT_API_LLM_PROVIDER=lmstudio
|
||||
# HINDSIGHT_API_LLM_API_KEY=lmstudio
|
||||
# HINDSIGHT_API_LLM_BASE_URL=http://localhost:1234/v1
|
||||
# HINDSIGHT_API_LLM_MODEL=qwen2.5-32b-instruct
|
||||
|
||||
# API Configuration (Optional)
|
||||
HINDSIGHT_API_HOST=0.0.0.0
|
||||
HINDSIGHT_API_PORT=8888
|
||||
|
||||
Executable
+27
@@ -0,0 +1,27 @@
|
||||
#!/bin/bash
|
||||
# Pre-commit hook - runs all scripts in scripts/hooks/
|
||||
|
||||
set -e
|
||||
|
||||
REPO_ROOT="$(git rev-parse --show-toplevel)"
|
||||
HOOKS_DIR="$REPO_ROOT/scripts/hooks"
|
||||
|
||||
if [ ! -d "$HOOKS_DIR" ]; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=== Running pre-commit hooks ==="
|
||||
echo ""
|
||||
|
||||
# Run all executable scripts in hooks directory
|
||||
for hook in "$HOOKS_DIR"/*.sh; do
|
||||
if [ -x "$hook" ]; then
|
||||
echo "[hook] $(basename "$hook")"
|
||||
(cd "$REPO_ROOT" && "$hook")
|
||||
fi
|
||||
done
|
||||
|
||||
echo ""
|
||||
echo "=== Pre-commit hooks completed ==="
|
||||
echo ""
|
||||
@@ -0,0 +1,71 @@
|
||||
name: Bug Report
|
||||
description: Report a bug or unexpected behavior
|
||||
labels: ["bug", "triage"]
|
||||
body:
|
||||
- type: markdown
|
||||
attributes:
|
||||
value: |
|
||||
Thanks for taking the time to report a bug! Please fill out the sections below.
|
||||
|
||||
- type: textarea
|
||||
id: description
|
||||
attributes:
|
||||
label: Bug Description
|
||||
description: A clear and concise description of the bug
|
||||
placeholder: What happened?
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: textarea
|
||||
id: reproduction
|
||||
attributes:
|
||||
label: Steps to Reproduce
|
||||
description: Steps to reproduce the behavior
|
||||
placeholder: |
|
||||
1. Configure '...'
|
||||
2. Call '...'
|
||||
3. See error
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: textarea
|
||||
id: expected
|
||||
attributes:
|
||||
label: Expected Behavior
|
||||
description: What did you expect to happen?
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: textarea
|
||||
id: actual
|
||||
attributes:
|
||||
label: Actual Behavior
|
||||
description: What actually happened?
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: input
|
||||
id: version
|
||||
attributes:
|
||||
label: Version
|
||||
description: What version are you using?
|
||||
placeholder: e.g., 0.1.0 or commit hash
|
||||
validations:
|
||||
required: false
|
||||
|
||||
- type: dropdown
|
||||
id: llm-provider
|
||||
attributes:
|
||||
label: LLM Provider
|
||||
description: Which LLM provider are you using?
|
||||
options:
|
||||
- OpenAI
|
||||
- Anthropic
|
||||
- Gemini
|
||||
- Groq
|
||||
- Ollama
|
||||
- LM Studio
|
||||
- Other
|
||||
validations:
|
||||
required: false
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
blank_issues_enabled: false
|
||||
contact_links:
|
||||
- name: Questions & Help
|
||||
url: https://github.com/vectorize-io/hindsight/discussions/categories/q-a
|
||||
about: Please ask questions and get help in Discussions instead of opening an issue.
|
||||
- name: Ideas & Feedback
|
||||
url: https://github.com/vectorize-io/hindsight/discussions/categories/ideas
|
||||
about: Share ideas or give feedback in Discussions.
|
||||
@@ -0,0 +1,82 @@
|
||||
name: Feature Request
|
||||
description: Suggest a new feature or enhancement
|
||||
labels: ["enhancement", "triage"]
|
||||
body:
|
||||
- type: markdown
|
||||
attributes:
|
||||
value: |
|
||||
Thanks for suggesting a feature! Please describe what you'd like to see added.
|
||||
|
||||
- type: textarea
|
||||
id: use-case
|
||||
attributes:
|
||||
label: Use Case
|
||||
description: Describe your specific use case. What are you building? What's your goal?
|
||||
placeholder: |
|
||||
I'm building an AI agent that needs to...
|
||||
My application handles...
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: textarea
|
||||
id: problem
|
||||
attributes:
|
||||
label: Problem Statement
|
||||
description: What problem are you facing? What's missing or difficult today?
|
||||
placeholder: Currently I have to... which causes...
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: textarea
|
||||
id: benefit
|
||||
attributes:
|
||||
label: How This Feature Would Help
|
||||
description: Explain how this feature would improve your workflow or solve your problem
|
||||
placeholder: With this feature, I would be able to...
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: textarea
|
||||
id: solution
|
||||
attributes:
|
||||
label: Proposed Solution
|
||||
description: Describe your ideal solution (optional - we may have ideas too!)
|
||||
placeholder: It would be great if Hindsight could...
|
||||
validations:
|
||||
required: false
|
||||
|
||||
- type: textarea
|
||||
id: alternatives
|
||||
attributes:
|
||||
label: Alternatives Considered
|
||||
description: Have you considered any alternative solutions or workarounds?
|
||||
validations:
|
||||
required: false
|
||||
|
||||
- type: dropdown
|
||||
id: priority
|
||||
attributes:
|
||||
label: Priority
|
||||
description: How important is this feature to you?
|
||||
options:
|
||||
- Nice to have
|
||||
- Important - affects my workflow
|
||||
- Critical - blocking my use case
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: textarea
|
||||
id: additional
|
||||
attributes:
|
||||
label: Additional Context
|
||||
description: Any other context, mockups, or examples?
|
||||
validations:
|
||||
required: false
|
||||
|
||||
- type: checkboxes
|
||||
id: checklist
|
||||
attributes:
|
||||
label: Checklist
|
||||
options:
|
||||
- label: I would be willing to contribute this feature
|
||||
required: false
|
||||
@@ -27,7 +27,9 @@ jobs:
|
||||
node-version: 20
|
||||
cache: npm
|
||||
cache-dependency-path: package-lock.json
|
||||
- uses: astral-sh/setup-uv@v4
|
||||
- run: npm ci --workspace=hindsight-docs
|
||||
- run: uv run generate-llms-full
|
||||
- run: npm run build --workspace=hindsight-docs
|
||||
- uses: actions/upload-pages-artifact@v3
|
||||
with:
|
||||
|
||||
@@ -42,6 +42,10 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/litellm
|
||||
run: uv build --out-dir dist
|
||||
|
||||
- name: Build hindsight-embed
|
||||
working-directory: ./hindsight-embed
|
||||
run: uv build --out-dir dist
|
||||
|
||||
# Publish in order (client and api first, then hindsight-all which depends on them)
|
||||
- name: Publish hindsight-client to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
@@ -67,6 +71,12 @@ jobs:
|
||||
packages-dir: ./hindsight-integrations/litellm/dist
|
||||
skip-existing: true
|
||||
|
||||
- name: Publish hindsight-embed to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: ./hindsight-embed/dist
|
||||
skip-existing: true
|
||||
|
||||
# Upload artifacts for GitHub release
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v4
|
||||
@@ -77,6 +87,7 @@ jobs:
|
||||
hindsight-api/dist/*
|
||||
hindsight/dist/*
|
||||
hindsight-integrations/litellm/dist/*
|
||||
hindsight-embed/dist/*
|
||||
retention-days: 1
|
||||
|
||||
release-typescript-client:
|
||||
@@ -102,7 +113,18 @@ jobs:
|
||||
|
||||
- name: Publish to npm
|
||||
working-directory: ./hindsight-clients/typescript
|
||||
run: npm publish --access public
|
||||
run: |
|
||||
set +e
|
||||
OUTPUT=$(npm publish --access public 2>&1)
|
||||
EXIT_CODE=$?
|
||||
echo "$OUTPUT"
|
||||
if [ $EXIT_CODE -ne 0 ]; then
|
||||
if echo "$OUTPUT" | grep -q "cannot publish over"; then
|
||||
echo "Package version already published, skipping..."
|
||||
exit 0
|
||||
fi
|
||||
exit $EXIT_CODE
|
||||
fi
|
||||
env:
|
||||
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
|
||||
|
||||
@@ -117,6 +139,65 @@ jobs:
|
||||
path: hindsight-clients/typescript/*.tgz
|
||||
retention-days: 1
|
||||
|
||||
release-control-plane:
|
||||
runs-on: ubuntu-latest
|
||||
environment: npm
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
registry-url: 'https://registry.npmjs.org'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: package-lock.json
|
||||
|
||||
- name: Install dependencies
|
||||
run: npm ci
|
||||
|
||||
- name: Build TypeScript client (dependency)
|
||||
run: npm run build --workspace=hindsight-clients/typescript
|
||||
|
||||
- name: Fix platform-specific native modules
|
||||
run: |
|
||||
# npm ci installs from lockfile which may have wrong platform binaries
|
||||
# Delete hoisted native modules and reinstall for current platform
|
||||
rm -rf node_modules/lightningcss node_modules/@tailwindcss
|
||||
npm install lightningcss @tailwindcss/postcss @tailwindcss/node
|
||||
|
||||
- name: Build
|
||||
run: npm run build --workspace=hindsight-control-plane
|
||||
|
||||
- name: Publish to npm
|
||||
working-directory: ./hindsight-control-plane
|
||||
run: |
|
||||
set +e
|
||||
OUTPUT=$(npm publish --access public 2>&1)
|
||||
EXIT_CODE=$?
|
||||
echo "$OUTPUT"
|
||||
if [ $EXIT_CODE -ne 0 ]; then
|
||||
if echo "$OUTPUT" | grep -q "cannot publish over"; then
|
||||
echo "Package version already published, skipping..."
|
||||
exit 0
|
||||
fi
|
||||
exit $EXIT_CODE
|
||||
fi
|
||||
env:
|
||||
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
|
||||
|
||||
- name: Pack for GitHub release
|
||||
working-directory: ./hindsight-control-plane
|
||||
run: npm pack
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: control-plane
|
||||
path: hindsight-control-plane/*.tgz
|
||||
retention-days: 1
|
||||
|
||||
release-rust-cli:
|
||||
runs-on: ${{ matrix.os }}
|
||||
strategy:
|
||||
@@ -181,7 +262,7 @@ jobs:
|
||||
- name: Free Disk Space
|
||||
uses: jlumbroso/free-disk-space@main
|
||||
with:
|
||||
tool-cache: false
|
||||
tool-cache: true
|
||||
android: true
|
||||
dotnet: true
|
||||
haskell: true
|
||||
@@ -206,7 +287,7 @@ jobs:
|
||||
id: get_version
|
||||
run: echo "VERSION=${GITHUB_REF#refs/tags/v}" >> $GITHUB_OUTPUT
|
||||
|
||||
- name: Extract metadata
|
||||
- name: Extract metadata for release tags
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
@@ -217,7 +298,29 @@ jobs:
|
||||
type=semver,pattern={{major}},value=${{ steps.get_version.outputs.VERSION }}
|
||||
type=raw,value=latest
|
||||
|
||||
- name: Build and push
|
||||
# TODO: Re-enable smoke test when disk space issue is resolved
|
||||
# # Step 1: Build for local testing (single platform, no push)
|
||||
# # This creates an identical image to what will be released, just for one platform
|
||||
# - name: Build image for testing
|
||||
# uses: docker/build-push-action@v6
|
||||
# with:
|
||||
# context: .
|
||||
# file: docker/standalone/Dockerfile
|
||||
# target: ${{ matrix.target }}
|
||||
# push: false
|
||||
# load: true
|
||||
# tags: ${{ matrix.image_name }}:test
|
||||
# cache-from: type=gha
|
||||
# cache-to: type=gha,mode=max
|
||||
|
||||
# # Step 2: Test the image before pushing anything
|
||||
# - name: Smoke test - verify container starts
|
||||
# env:
|
||||
# GROQ_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
# run: ./scripts/docker-smoke-test.sh "${{ matrix.image_name }}:test" "${{ matrix.target }}"
|
||||
|
||||
# Build multi-platform and push to release tags
|
||||
- name: Build and push release images
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
@@ -263,7 +366,7 @@ jobs:
|
||||
|
||||
create-github-release:
|
||||
runs-on: ubuntu-latest
|
||||
needs: [release-python-packages, release-typescript-client, release-rust-cli, release-docker-images, release-helm-chart]
|
||||
needs: [release-python-packages, release-typescript-client, release-control-plane, release-rust-cli, release-docker-images, release-helm-chart]
|
||||
permissions:
|
||||
contents: write
|
||||
|
||||
@@ -286,6 +389,12 @@ jobs:
|
||||
name: typescript-client
|
||||
path: ./artifacts/typescript-client
|
||||
|
||||
- name: Download Control Plane
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: control-plane
|
||||
path: ./artifacts/control-plane
|
||||
|
||||
- name: Download Rust CLI (Linux)
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
@@ -318,8 +427,11 @@ jobs:
|
||||
cp artifacts/python-packages/hindsight-api/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight-integrations/litellm/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight-embed/dist/* release-assets/ || true
|
||||
# TypeScript client
|
||||
cp artifacts/typescript-client/*.tgz release-assets/ || true
|
||||
# Control Plane
|
||||
cp artifacts/control-plane/*.tgz release-assets/ || true
|
||||
# Rust CLI binaries
|
||||
cp artifacts/rust-cli-linux/hindsight-linux-amd64 release-assets/ || true
|
||||
cp artifacts/rust-cli-darwin-amd64/hindsight-darwin-amd64 release-assets/ || true
|
||||
|
||||
+482
-9
@@ -20,6 +20,8 @@ jobs:
|
||||
path: hindsight-api
|
||||
- name: hindsight-client
|
||||
path: hindsight-clients/python
|
||||
- name: hindsight-embed
|
||||
path: hindsight-embed
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
@@ -38,6 +40,29 @@ jobs:
|
||||
working-directory: ./${{ matrix.path }}
|
||||
run: uv build
|
||||
|
||||
build-api-python-versions:
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
matrix:
|
||||
python-version: ['3.11', '3.12', '3.13']
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Set up Python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Build hindsight-api
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
build-typescript-client:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
@@ -57,6 +82,58 @@ jobs:
|
||||
- name: Build TypeScript client
|
||||
run: npm run build --workspace=hindsight-clients/typescript
|
||||
|
||||
build-control-plane:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: package-lock.json
|
||||
|
||||
- name: Install SDK dependencies
|
||||
run: npm ci --workspace=hindsight-clients/typescript
|
||||
|
||||
- name: Build SDK
|
||||
run: npm run build --workspace=hindsight-clients/typescript
|
||||
|
||||
# Install control plane deps and fix hoisted lightningcss binary
|
||||
# lightningcss gets hoisted to root node_modules, so we need to reinstall it there
|
||||
- name: Install Control Plane dependencies
|
||||
run: |
|
||||
npm install --workspace=hindsight-control-plane
|
||||
rm -rf node_modules/lightningcss node_modules/@tailwindcss
|
||||
npm install lightningcss @tailwindcss/postcss @tailwindcss/node
|
||||
|
||||
- name: Build Control Plane
|
||||
run: npm run build --workspace=hindsight-control-plane
|
||||
|
||||
- name: Verify standalone build
|
||||
run: |
|
||||
test -f hindsight-control-plane/standalone/server.js || exit 1
|
||||
test -d hindsight-control-plane/standalone/node_modules || exit 1
|
||||
node hindsight-control-plane/bin/cli.js --help
|
||||
|
||||
- name: Smoke test - verify server starts
|
||||
run: |
|
||||
cd hindsight-control-plane
|
||||
node bin/cli.js --port 9999 &
|
||||
SERVER_PID=$!
|
||||
sleep 5
|
||||
if curl -sf http://localhost:9999 > /dev/null 2>&1; then
|
||||
echo "Server started successfully"
|
||||
kill $SERVER_PID 2>/dev/null || true
|
||||
exit 0
|
||||
else
|
||||
echo "Server failed to respond"
|
||||
kill $SERVER_PID 2>/dev/null || true
|
||||
exit 1
|
||||
fi
|
||||
|
||||
build-docs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
@@ -76,8 +153,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
|
||||
@@ -94,10 +178,75 @@ 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
|
||||
|
||||
- name: Upload CLI artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: hindsight-cli
|
||||
path: hindsight-cli/target/release/hindsight
|
||||
retention-days: 1
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Create .env file
|
||||
run: |
|
||||
cat > .env << EOF
|
||||
HINDSIGHT_API_LLM_PROVIDER=${{ env.HINDSIGHT_API_LLM_PROVIDER }}
|
||||
HINDSIGHT_API_LLM_API_KEY=${{ env.HINDSIGHT_API_LLM_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL=${{ env.HINDSIGHT_API_LLM_MODEL }}
|
||||
EOF
|
||||
|
||||
- name: Start API server
|
||||
run: |
|
||||
./scripts/dev/start-api.sh > /tmp/api-server.log 2>&1 &
|
||||
echo "Waiting for API server to be ready..."
|
||||
for i in {1..60}; do
|
||||
if curl -sf http://localhost:8888/health > /dev/null 2>&1; then
|
||||
echo "API server is ready after ${i}s"
|
||||
break
|
||||
fi
|
||||
if [ $i -eq 60 ]; then
|
||||
echo "API server failed to start after 60s"
|
||||
cat /tmp/api-server.log
|
||||
exit 1
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
- name: Run CLI smoke test
|
||||
run: |
|
||||
HINDSIGHT_CLI=hindsight-cli/target/release/hindsight ./hindsight-cli/smoke-test.sh
|
||||
|
||||
- name: Show API server logs
|
||||
if: always()
|
||||
run: |
|
||||
echo "=== API Server Logs ==="
|
||||
cat /tmp/api-server.log || echo "No API server log found"
|
||||
|
||||
lint-helm-chart:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
@@ -130,7 +279,7 @@ jobs:
|
||||
- name: Free Disk Space
|
||||
uses: jlumbroso/free-disk-space@main
|
||||
with:
|
||||
tool-cache: false
|
||||
tool-cache: true
|
||||
android: true
|
||||
dotnet: true
|
||||
haskell: true
|
||||
@@ -148,6 +297,13 @@ jobs:
|
||||
file: docker/standalone/Dockerfile
|
||||
target: ${{ matrix.target }}
|
||||
push: false
|
||||
load: false
|
||||
|
||||
# TODO: Re-enable smoke test when disk space issue is resolved
|
||||
# - name: Smoke test - verify container starts
|
||||
# env:
|
||||
# GROQ_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
# run: ./scripts/docker-smoke-test.sh "hindsight-${{ matrix.name }}:test" "${{ matrix.target }}"
|
||||
|
||||
test-api:
|
||||
runs-on: ubuntu-latest
|
||||
@@ -157,6 +313,8 @@ jobs:
|
||||
GROQ_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
GEMINI_API_KEY: ${{ secrets.GEMINI_API_KEY }}
|
||||
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
|
||||
COHERE_API_KEY: ${{ secrets.COHERE_API_KEY }}
|
||||
HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
# Prefer CPU-only PyTorch in CI (but keep PyPI for everything else)
|
||||
@@ -182,7 +340,7 @@ jobs:
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --extra test --no-install-project --index-strategy unsafe-best-match
|
||||
run: uv sync --frozen --extra test --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v4
|
||||
@@ -243,11 +401,11 @@ jobs:
|
||||
|
||||
- name: Install client test dependencies
|
||||
working-directory: ./hindsight-clients/python
|
||||
run: uv sync --extra test --index-strategy unsafe-best-match
|
||||
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --no-install-project --index-strategy unsafe-best-match
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Create .env file
|
||||
run: |
|
||||
@@ -320,7 +478,7 @@ jobs:
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --no-install-project --index-strategy unsafe-best-match
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install TypeScript client dependencies
|
||||
working-directory: ./hindsight-clients/typescript
|
||||
@@ -408,7 +566,7 @@ jobs:
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --no-install-project --index-strategy unsafe-best-match
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Create .env file
|
||||
run: |
|
||||
@@ -445,6 +603,97 @@ jobs:
|
||||
echo "=== API Server Logs ==="
|
||||
cat /tmp/api-server.log || echo "No API server log found"
|
||||
|
||||
test-integration:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
HINDSIGHT_API_LLM_PROVIDER: groq
|
||||
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
|
||||
HINDSIGHT_API_URL: http://localhost:8888
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install integration test dependencies
|
||||
working-directory: ./hindsight-integration-tests
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
print('Downloading embedding model...')
|
||||
SentenceTransformer('BAAI/bge-small-en-v1.5')
|
||||
print('Downloading cross-encoder model...')
|
||||
CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2')
|
||||
print('Models downloaded successfully')
|
||||
"
|
||||
|
||||
- name: Create .env file
|
||||
run: |
|
||||
cat > .env << EOF
|
||||
HINDSIGHT_API_LLM_PROVIDER=${{ env.HINDSIGHT_API_LLM_PROVIDER }}
|
||||
HINDSIGHT_API_LLM_API_KEY=${{ env.HINDSIGHT_API_LLM_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL=${{ env.HINDSIGHT_API_LLM_MODEL }}
|
||||
EOF
|
||||
|
||||
- name: Start API server
|
||||
run: |
|
||||
./scripts/dev/start-api.sh > /tmp/api-server.log 2>&1 &
|
||||
echo "Waiting for API server to be ready..."
|
||||
for i in {1..60}; do
|
||||
if curl -sf http://localhost:8888/health > /dev/null 2>&1; then
|
||||
echo "API server is ready after ${i}s"
|
||||
break
|
||||
fi
|
||||
if [ $i -eq 60 ]; then
|
||||
echo "API server failed to start after 60s"
|
||||
cat /tmp/api-server.log
|
||||
exit 1
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
- name: Run integration tests
|
||||
working-directory: ./hindsight-integration-tests
|
||||
run: uv run pytest tests/ -v
|
||||
|
||||
- name: Show API server logs
|
||||
if: always()
|
||||
run: |
|
||||
echo "=== API Server Logs ==="
|
||||
cat /tmp/api-server.log || echo "No API server log found"
|
||||
|
||||
test-litellm-integration:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
@@ -468,8 +717,232 @@ jobs:
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-integrations/litellm
|
||||
run: uv sync --extra dev
|
||||
run: uv sync --frozen --extra dev
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/litellm
|
||||
run: uv run pytest tests -v
|
||||
run: uv run pytest tests -v
|
||||
|
||||
test-embed:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
HINDSIGHT_EMBED_LLM_PROVIDER: groq
|
||||
HINDSIGHT_EMBED_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_EMBED_LLM_MODEL: openai/gpt-oss-20b
|
||||
# Prefer CPU-only PyTorch in CI
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-embed
|
||||
run: uv sync --frozen --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-embed-${{ hashFiles('hindsight-embed/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-embed-
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Run smoke test
|
||||
working-directory: ./hindsight-embed
|
||||
run: ./test.sh
|
||||
|
||||
test-doc-examples:
|
||||
runs-on: ubuntu-latest
|
||||
needs: test-rust-cli
|
||||
env:
|
||||
HINDSIGHT_API_LLM_PROVIDER: groq
|
||||
HINDSIGHT_API_LLM_API_KEY: ${{ secrets.GROQ_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL: openai/gpt-oss-20b
|
||||
HINDSIGHT_API_URL: http://localhost:8888
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Download CLI artifact
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: hindsight-cli
|
||||
path: /usr/local/bin
|
||||
|
||||
- name: Make CLI executable
|
||||
run: chmod +x /usr/local/bin/hindsight
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: package-lock.json
|
||||
|
||||
- name: Build and install API
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv build
|
||||
uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install Python client dependencies
|
||||
working-directory: ./hindsight-clients/python
|
||||
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install TypeScript client
|
||||
run: |
|
||||
npm ci --workspace=hindsight-clients/typescript
|
||||
npm run build --workspace=hindsight-clients/typescript
|
||||
|
||||
- name: Create .env file
|
||||
run: |
|
||||
cat > .env << EOF
|
||||
HINDSIGHT_API_LLM_PROVIDER=${{ env.HINDSIGHT_API_LLM_PROVIDER }}
|
||||
HINDSIGHT_API_LLM_API_KEY=${{ env.HINDSIGHT_API_LLM_API_KEY }}
|
||||
HINDSIGHT_API_LLM_MODEL=${{ env.HINDSIGHT_API_LLM_MODEL }}
|
||||
EOF
|
||||
|
||||
- name: Start API server
|
||||
run: |
|
||||
./scripts/dev/start-api.sh > /tmp/api-server.log 2>&1 &
|
||||
echo "Waiting for API server to be ready..."
|
||||
for i in {1..60}; do
|
||||
if curl -sf http://localhost:8888/health > /dev/null 2>&1; then
|
||||
echo "API server is ready after ${i}s"
|
||||
break
|
||||
fi
|
||||
if [ $i -eq 60 ]; then
|
||||
echo "API server failed to start after 60s"
|
||||
cat /tmp/api-server.log
|
||||
exit 1
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
- name: Run Python doc examples
|
||||
working-directory: ./hindsight-clients/python
|
||||
run: |
|
||||
for f in ../../hindsight-docs/examples/api/*.py; do
|
||||
echo "Running $f..."
|
||||
uv run python "$f"
|
||||
done
|
||||
|
||||
- name: Run Node.js doc examples
|
||||
run: |
|
||||
for f in hindsight-docs/examples/api/*.mjs; do
|
||||
echo "Running $f..."
|
||||
node "$f"
|
||||
done
|
||||
|
||||
- name: Configure CLI
|
||||
run: hindsight configure --api-url http://localhost:8888
|
||||
|
||||
- name: Run CLI doc examples
|
||||
run: |
|
||||
for f in hindsight-docs/examples/api/*.sh; do
|
||||
echo "Running $f..."
|
||||
bash "$f"
|
||||
done
|
||||
|
||||
- name: Show API server logs
|
||||
if: always()
|
||||
run: |
|
||||
echo "=== API Server Logs ==="
|
||||
cat /tmp/api-server.log || echo "No API server log found"
|
||||
|
||||
verify-generated-files:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: package-lock.json
|
||||
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
|
||||
- name: Cache cargo
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
~/.cargo/git
|
||||
key: ${{ runner.os }}-cargo-gen-${{ hashFiles('**/Cargo.lock') }}
|
||||
|
||||
- name: Install Node dependencies
|
||||
run: npm ci
|
||||
|
||||
- name: Install Python dependencies
|
||||
run: |
|
||||
cd hindsight-dev && uv sync --frozen --index-strategy unsafe-best-match
|
||||
cd ../hindsight-api && uv sync --frozen --index-strategy unsafe-best-match
|
||||
cd ../hindsight-embed && uv sync --frozen --index-strategy unsafe-best-match
|
||||
|
||||
- name: Run generate-openapi
|
||||
run: ./scripts/generate-openapi.sh
|
||||
|
||||
- name: Run generate-clients
|
||||
run: ./scripts/generate-clients.sh
|
||||
|
||||
- name: Run lint
|
||||
run: ./scripts/hooks/lint.sh
|
||||
|
||||
- name: Verify no uncommitted changes
|
||||
run: |
|
||||
if [ -n "$(git status --porcelain)" ]; then
|
||||
echo "❌ Error: Generated files are out of sync with committed files."
|
||||
echo ""
|
||||
echo "The following files have changed after running generation scripts:"
|
||||
git status --porcelain
|
||||
echo ""
|
||||
echo "Please run the following commands locally and commit the changes:"
|
||||
echo " ./scripts/generate-openapi.sh"
|
||||
echo " ./scripts/generate-clients.sh"
|
||||
echo " ./scripts/hooks/lint.sh"
|
||||
echo ""
|
||||
git diff --stat
|
||||
exit 1
|
||||
fi
|
||||
echo "✓ All generated files are up to date"
|
||||
+17
-3
@@ -5,15 +5,18 @@ build/
|
||||
dist/
|
||||
wheels/
|
||||
*.egg-info
|
||||
|
||||
.mcp.json
|
||||
.osgrep
|
||||
# Virtual environments
|
||||
.venv
|
||||
|
||||
# Node
|
||||
node_modules/
|
||||
|
||||
# Environment variables
|
||||
# Environment variables and local config
|
||||
.env
|
||||
docker-compose.yml
|
||||
docker-compose.override.yml
|
||||
|
||||
# IDE
|
||||
.idea/
|
||||
@@ -24,6 +27,10 @@ node_modules/
|
||||
# NLTK data (will be downloaded automatically)
|
||||
nltk_data/
|
||||
|
||||
# Monitoring stack (Prometheus/Grafana binaries and data)
|
||||
.monitoring/
|
||||
.pgbouncer/
|
||||
|
||||
# Large benchmark datasets (will be downloaded automatically)
|
||||
**/longmemeval_s_cleaned.json
|
||||
|
||||
@@ -32,8 +39,15 @@ logs/
|
||||
|
||||
.DS_Store
|
||||
|
||||
# Generated docs files
|
||||
hindsight-docs/static/llms-full.txt
|
||||
|
||||
|
||||
hindsight-dev/benchmarks/locomo/results/
|
||||
hindsight-dev/benchmarks/longmemeval/results/
|
||||
hindsight-cli/target
|
||||
hindsight-clients/rust/target
|
||||
hindsight-clients/rust/target
|
||||
.claude
|
||||
whats-next.md
|
||||
TASK.md
|
||||
CHANGELOG.md
|
||||
@@ -1,151 +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-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.
|
||||
|
||||
@@ -0,0 +1,283 @@
|
||||
# CLAUDE.md
|
||||
|
||||
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
|
||||
|
||||
## Project Overview
|
||||
|
||||
Hindsight is an agent memory system that provides long-term memory for AI agents using biomimetic data structures. Memories are organized as:
|
||||
- **World facts**: General knowledge ("The sky is blue")
|
||||
- **Experience facts**: Personal experiences ("I visited Paris in 2023")
|
||||
- **Opinion facts**: Beliefs with confidence scores ("Paris is beautiful" - 0.9 confidence)
|
||||
- **Observations**: Complex mental models derived from reflection
|
||||
|
||||
## Development Commands
|
||||
|
||||
### API Server (Python/FastAPI)
|
||||
```bash
|
||||
# Start API server (loads .env automatically)
|
||||
./scripts/dev/start-api.sh
|
||||
|
||||
# Run all tests (parallelized with pytest-xdist)
|
||||
cd hindsight-api && uv run pytest tests/
|
||||
|
||||
# Run specific test file
|
||||
cd hindsight-api && uv run pytest tests/test_http_api_integration.py -v
|
||||
|
||||
# Run single test function
|
||||
cd hindsight-api && uv run pytest tests/test_retain.py::test_retain_simple -v
|
||||
|
||||
# Lint and format
|
||||
cd hindsight-api && uv run ruff check .
|
||||
cd hindsight-api && uv run ruff format .
|
||||
|
||||
# Type checking (uses ty - extremely fast type checker from Astral)
|
||||
cd hindsight-api && uv run ty check hindsight_api/
|
||||
```
|
||||
|
||||
### Control Plane (Next.js)
|
||||
```bash
|
||||
./scripts/dev/start-control-plane.sh
|
||||
# Or manually:
|
||||
cd hindsight-control-plane && npm run dev
|
||||
```
|
||||
|
||||
### Documentation Site (Docusaurus)
|
||||
```bash
|
||||
./scripts/dev/start-docs.sh
|
||||
```
|
||||
|
||||
### Generating Clients/OpenAPI
|
||||
```bash
|
||||
# Regenerate OpenAPI spec after API changes (REQUIRED after changing endpoints)
|
||||
./scripts/generate-openapi.sh
|
||||
|
||||
# Regenerate all client SDKs (Python, TypeScript, Rust)
|
||||
./scripts/generate-clients.sh
|
||||
```
|
||||
|
||||
### Benchmarks
|
||||
```bash
|
||||
./scripts/benchmarks/run-longmemeval.sh
|
||||
./scripts/benchmarks/run-locomo.sh
|
||||
./scripts/benchmarks/start-visualizer.sh # View results at localhost:8001
|
||||
```
|
||||
|
||||
## Architecture
|
||||
|
||||
### Monorepo Structure
|
||||
- **hindsight-api/**: Core FastAPI server with memory engine (Python, uv)
|
||||
- **hindsight/**: Embedded Python bundle (hindsight-all package)
|
||||
- **hindsight-control-plane/**: Admin UI (Next.js, npm)
|
||||
- **hindsight-cli/**: CLI tool (Rust, cargo, uses progenitor for API client)
|
||||
- **hindsight-clients/**: Generated SDK clients (Python, TypeScript, Rust)
|
||||
- **hindsight-docs/**: Docusaurus documentation site
|
||||
- **hindsight-integrations/**: Framework integrations (LiteLLM, OpenAI)
|
||||
- **hindsight-dev/**: Development tools and benchmarks
|
||||
|
||||
### Core Engine (hindsight-api/hindsight_api/engine/)
|
||||
- `memory_engine.py`: Main orchestrator (~170KB) for retain/recall/reflect operations
|
||||
- `llm_wrapper.py`: LLM abstraction supporting OpenAI, Anthropic, Gemini, Groq, Ollama, LM Studio
|
||||
- `embeddings.py`: Embedding generation (local sentence-transformers or TEI)
|
||||
- `cross_encoder.py`: Reranking (local or TEI)
|
||||
- `entity_resolver.py`: Entity extraction and normalization
|
||||
- `query_analyzer.py`: Query intent analysis
|
||||
|
||||
**retain/**: Memory ingestion pipeline
|
||||
- `orchestrator.py`: Coordinates the retain flow
|
||||
- `fact_extraction.py`: LLM-based fact extraction from content
|
||||
- `link_utils.py`: Entity link creation and management
|
||||
|
||||
**search/**: Multi-strategy retrieval
|
||||
- `retrieval.py`: Main retrieval orchestrator
|
||||
- `graph_retrieval.py`: Entity/relationship graph traversal
|
||||
- `mpfp_retrieval.py`: Multi-Path Fact Propagation retrieval
|
||||
- `fusion.py`: Reciprocal rank fusion for combining results
|
||||
- `reranking.py`: Cross-encoder reranking
|
||||
|
||||
### API Layer (hindsight-api/hindsight_api/api/)
|
||||
- `http.py`: FastAPI HTTP routers (~80KB) for all REST endpoints
|
||||
- `mcp.py`: Model Context Protocol server implementation
|
||||
|
||||
Main operations:
|
||||
- **Retain**: Store memories, extracts facts/entities/relationships
|
||||
- **Recall**: Retrieve memories via 4 parallel strategies (semantic, BM25, graph, temporal) + reranking
|
||||
- **Reflect**: Deep analysis forming new opinions/observations (disposition-aware)
|
||||
|
||||
### Database
|
||||
PostgreSQL with pgvector. Schema managed via Alembic migrations in `hindsight-api/hindsight_api/alembic/`. Migrations run automatically on API startup.
|
||||
|
||||
Key tables: `banks`, `memory_units`, `documents`, `entities`, `entity_links`
|
||||
|
||||
### Adding Database Migrations
|
||||
|
||||
1. **Create a new migration file** in `hindsight-api/hindsight_api/alembic/versions/`:
|
||||
- File name format: `<revision_id>_<description>.py` (e.g., `f1a2b3c4d5e6_add_new_index.py`)
|
||||
- Use a unique hex revision ID (12 chars)
|
||||
- Set `down_revision` to the previous migration's revision ID
|
||||
|
||||
2. **Migration template**:
|
||||
```python
|
||||
"""Description of the migration
|
||||
|
||||
Revision ID: f1a2b3c4d5e6
|
||||
Revises: <previous_revision_id>
|
||||
Create Date: YYYY-MM-DD
|
||||
"""
|
||||
from collections.abc import Sequence
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "f1a2b3c4d5e6"
|
||||
down_revision: str | Sequence[str] | None = "<previous_revision_id>"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
def upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
op.execute(f"CREATE INDEX ... ON {schema}table_name(...)")
|
||||
|
||||
def downgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}index_name")
|
||||
```
|
||||
|
||||
3. **Run migrations locally**:
|
||||
```bash
|
||||
# Set database URL and run migrations
|
||||
uv run hindsight-admin run-db-migration
|
||||
|
||||
# Run on a specific tenant schema
|
||||
uv run hindsight-admin run-db-migration --schema tenant_xyz
|
||||
```
|
||||
|
||||
## Key Conventions
|
||||
|
||||
### Code Quality
|
||||
**Always run the lint script after making Python or TypeScript/Node changes:**
|
||||
```bash
|
||||
./scripts/hooks/lint.sh
|
||||
```
|
||||
This runs the same checks as the pre-commit hook (Ruff for Python, ESLint/Prettier for TypeScript).
|
||||
|
||||
### Memory Banks
|
||||
- Each bank is an isolated memory store (like a "brain" for one user/agent)
|
||||
- Banks have dispositions (skepticism, literalism, empathy traits 1-5) affecting reflect
|
||||
- Banks can have background context
|
||||
- Bank isolation is strict - no cross-bank data leakage
|
||||
|
||||
### API Design
|
||||
- All endpoints operate on a single bank per request
|
||||
- Multi-bank queries are client responsibility to orchestrate
|
||||
- Disposition traits only affect reflect, not recall
|
||||
|
||||
### Control Plane API Routes
|
||||
|
||||
When adding or modifying parameters in the dataplane API (hindsight-api), you must also update the control plane routes that proxy to it:
|
||||
|
||||
1. **API Routes** (`hindsight-control-plane/src/app/api/`):
|
||||
- `recall/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/recall`
|
||||
- `reflect/route.ts` - proxies to `/v1/default/banks/{bank_id}/reflect`
|
||||
- `memories/retain/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/retain`
|
||||
- Other routes follow the same pattern
|
||||
|
||||
2. **Client types** (`hindsight-control-plane/src/lib/api.ts`):
|
||||
- Update the TypeScript type definitions for `recall()`, `reflect()`, `retain()` etc.
|
||||
|
||||
3. **Checklist when adding new API parameters**:
|
||||
- Add parameter extraction in the route handler (destructure from `body`)
|
||||
- Pass the parameter to the SDK call
|
||||
- Update the client type definition in `lib/api.ts`
|
||||
- Update any UI components that need to use the new parameter
|
||||
|
||||
### Python Style
|
||||
- Python 3.11+, type hints required
|
||||
- Async throughout (asyncpg, async FastAPI)
|
||||
- Pydantic models for request/response
|
||||
- Ruff for linting (line-length 120)
|
||||
- No Python files at project root - maintain clean directory structure
|
||||
- **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
|
||||
- Tailwind CSS with shadcn/ui components
|
||||
|
||||
### Adding New API Configuration Flags
|
||||
|
||||
When adding a new environment variable configuration:
|
||||
|
||||
1. **config.py** (`hindsight-api/hindsight_api/config.py`):
|
||||
- Add `ENV_*` constant for the environment variable name
|
||||
- Add `DEFAULT_*` constant for the default value
|
||||
- Add field to `HindsightConfig` dataclass
|
||||
- Add initialization in `from_env()` method
|
||||
|
||||
2. **main.py** (`hindsight-api/hindsight_api/main.py`):
|
||||
- Add field to the manual `HindsightConfig()` constructor call (search for "CLI override")
|
||||
|
||||
3. **Use the config** in code:
|
||||
```python
|
||||
from ...config import get_config
|
||||
config = get_config()
|
||||
value = config.your_new_field
|
||||
```
|
||||
|
||||
4. **Documentation** (`hindsight-docs/docs/developer/configuration.md`):
|
||||
- Add to appropriate section table with Variable, Description, Default
|
||||
|
||||
## Environment Setup
|
||||
|
||||
```bash
|
||||
cp .env.example .env
|
||||
# Edit .env with LLM API key
|
||||
|
||||
# Python deps
|
||||
uv sync --directory hindsight-api/
|
||||
|
||||
# Node deps (uses npm workspaces)
|
||||
npm install
|
||||
```
|
||||
|
||||
Required env vars:
|
||||
- `HINDSIGHT_API_LLM_PROVIDER`: openai, anthropic, gemini, groq, ollama, lmstudio
|
||||
- `HINDSIGHT_API_LLM_API_KEY`: Your API key
|
||||
- `HINDSIGHT_API_LLM_MODEL`: Model name (e.g., o3-mini, claude-sonnet-4-20250514)
|
||||
|
||||
Optional (uses local models by default):
|
||||
- `HINDSIGHT_API_EMBEDDINGS_PROVIDER`: local (default) or tei
|
||||
- `HINDSIGHT_API_RERANKER_PROVIDER`: local (default) or tei
|
||||
- `HINDSIGHT_API_DATABASE_URL`: External PostgreSQL (uses embedded pg0 by default)
|
||||
+30
-1
@@ -51,7 +51,36 @@ cd hindsight-api
|
||||
uv run pytest tests/
|
||||
```
|
||||
|
||||
### Code style
|
||||
### Code Style
|
||||
|
||||
We use [Ruff](https://docs.astral.sh/ruff/) for Python linting and formatting, and ESLint/Prettier for TypeScript.
|
||||
|
||||
#### Setting up git hooks (recommended)
|
||||
|
||||
Set up git hooks to automatically lint and format code before each commit:
|
||||
|
||||
```bash
|
||||
./scripts/setup-hooks.sh
|
||||
```
|
||||
|
||||
This configures git to use the hooks in `.githooks/`, which run all scripts in `scripts/hooks/` on commit. The lint hook runs in parallel:
|
||||
- **Python**: `ruff check --fix`, `ruff format`, `ty check`
|
||||
- **TypeScript**: `eslint --fix`, `prettier`
|
||||
|
||||
#### Manual linting and formatting
|
||||
|
||||
```bash
|
||||
# Run all lints (same as pre-commit)
|
||||
./scripts/hooks/lint.sh
|
||||
|
||||
# Or run individually for Python:
|
||||
cd hindsight-api
|
||||
uv run ruff check --fix . # Lint and auto-fix
|
||||
uv run ruff format . # Format code
|
||||
uv run ty check hindsight_api # Type check
|
||||
```
|
||||
|
||||
#### Style guidelines
|
||||
|
||||
- Use Python type hints
|
||||
- Follow existing code patterns
|
||||
|
||||
@@ -2,14 +2,13 @@
|
||||
|
||||

|
||||
|
||||
[Documentation](https://vectorize-io.github.io/hindsight) • [Paper](#coming-soon) • [Examples](https://github.com/vectorize-io/hindsight-cookbook)
|
||||
[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)
|
||||
|
||||
[](https://github.com/vectorize-io/hindsight/actions/workflows/release.yml)
|
||||
[](https://join.slack.com/t/hindsight-space/shared_invite/zt-3nhbm4w29-LeSJ5Ixi6j8PdiYOCPlOgg)
|
||||
[](https://opensource.org/licenses/MIT)
|
||||
[](https://pypi.org/project/hindsight-api/)
|
||||
[](https://pypi.org/project/hindsight-client/)
|
||||
[](https://www.npmjs.com/package/@vectorize-io/hindsight-client)
|
||||
[](https://join.slack.com/t/hindsight-space/shared_invite/zt-3klo21kua-VUCC_zHP5rIcXFB1_5yw6A)
|
||||

|
||||

|
||||
|
||||
|
||||
</div>
|
||||
@@ -18,7 +17,7 @@
|
||||
|
||||
## What is Hindsight?
|
||||
|
||||
Hindsight™ is an agent memory system built to create smarter agents that learn over time. It eliminates the shortcomings of alternative techniques such as RAG and knowledge graph.
|
||||
Hindsight™ is an agent memory system built to create smarter agents that learn over time. It eliminates the shortcomings of alternative techniques such as RAG and knowledge graph and delivers state-of-the-art performance on long term memory tasks.
|
||||
|
||||
Hindsight addresses common challenges that have frustrated AI engineers building agents to automate tasks and assist users with conversational interfaces. Many of these challenges stem directly from a lack of memory.
|
||||
|
||||
@@ -26,27 +25,48 @@ Hindsight addresses common challenges that have frustrated AI engineers building
|
||||
- **Hallucinations:** Long term memory can be seeded with external knowledge to ground agent behavior in reliable sources to augment training data.
|
||||
- **Cognitive Overload:** As workflows get complex, retrievals, tool calls, user messages and agent responses can grow to fill the context window leading to context rot. Short term memory optimization allows agents to reduce tokens and focus context by removing irrelevant details.
|
||||
|
||||
## How Hindsight Works
|
||||
## How is Hindsight Different From Other Memory Systems?
|
||||
|
||||

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

|
||||
|
||||
The benchmark performance data for Hindsight and GPT-4o (full context) have been reproduced by research collaborators at the Virginia Tech [Sanghani Center for Artificial Intelligence and Data Analytics](https://sanghani.cs.vt.edu/) and The Washington Post. Other scores are self-reported by software vendors.
|
||||
|
||||
A thorough examination of the techniques implemented in Hindsight and detailed breakdowns of benchmark performance are [available on arXiv](https://arxiv.org/abs/2512.12818). This research is currently being prepared for conference submission and the wider peer review process.
|
||||
|
||||
The benchmark results from this research can be inspected in our [visual benchmark explorer](https://hindsight-benchmarks.vercel.app). As additional improvements are made to Hindsight, new benchmark data will be available for review using this same tool.
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Docker (recommended)
|
||||
@@ -61,6 +81,8 @@ docker run --rm -it --pull always -p 8888:8888 -p 9999:9999 \
|
||||
ghcr.io/vectorize-io/hindsight:latest
|
||||
```
|
||||
|
||||
You can modify the LLM provider by setting `HINDSIGHT_API_LLM_PROVIDER`. Valid options are `openai`, `anthropic`, `gemini`, `groq`, `ollama`, and `lmstudio`. The documentation provides more details on [supported models](https://hindsight.vectorize.io/developer/models).
|
||||
|
||||
API: http://localhost:8888
|
||||
UI: http://localhost:9999
|
||||
|
||||
@@ -220,9 +242,13 @@ client.reflect(bank_id="my-bank", query="What should I know about Alice?")
|
||||
- [CLI](https://hindsight.vectorize.io/sdks/cli)
|
||||
|
||||
**Community:**
|
||||
- [Slack](https://join.slack.com/t/hindsight-space/shared_invite/zt-3klo21kua-VUCC_zHP5rIcXFB1_5yw6A)
|
||||
- [Slack](https://join.slack.com/t/hindsight-space/shared_invite/zt-3nhbm4w29-LeSJ5Ixi6j8PdiYOCPlOgg)
|
||||
- [GitHub Issues](https://github.com/vectorize-io/hindsight/issues)
|
||||
|
||||
---
|
||||
## Star History
|
||||
|
||||
[](https://www.star-history.com/#vectorize-io/hindsight&type=date&legend=top-left)
|
||||
---
|
||||
|
||||
## Contributing
|
||||
|
||||
@@ -2,16 +2,24 @@
|
||||
# Supports building API-only, Control Plane-only, or both
|
||||
#
|
||||
# Build args:
|
||||
# INCLUDE_API=true/false - Include API (default: true)
|
||||
# INCLUDE_CP=true/false - Include Control Plane (default: true)
|
||||
# INCLUDE_API=true/false - Include API (default: true)
|
||||
# INCLUDE_CP=true/false - Include Control Plane (default: true)
|
||||
# INCLUDE_LOCAL_MODELS=true/false - Include local ML models for embeddings/reranking (default: true)
|
||||
# Set to false when using external providers (TEI, OpenAI, Cohere)
|
||||
# PRELOAD_ML_MODELS=true/false - Pre-download ML models during build (default: true)
|
||||
# Only effective when INCLUDE_LOCAL_MODELS=true
|
||||
#
|
||||
# Examples:
|
||||
# docker build -t hindsight . # Both (standalone)
|
||||
# docker build -t hindsight-api --build-arg INCLUDE_CP=false . # API only
|
||||
# docker build -t hindsight-cp --build-arg INCLUDE_API=false . # Control Plane only
|
||||
# docker build -t hindsight . # Both (standalone)
|
||||
# docker build -t hindsight-api --build-arg INCLUDE_CP=false . # API only
|
||||
# docker build -t hindsight-cp --build-arg INCLUDE_API=false . # Control Plane only
|
||||
# docker build -t hindsight --build-arg PRELOAD_ML_MODELS=false . # Skip ML model preload
|
||||
# docker build -t hindsight --build-arg INCLUDE_LOCAL_MODELS=false . # Skip local ML deps (for external providers)
|
||||
|
||||
ARG INCLUDE_API=true
|
||||
ARG INCLUDE_CP=true
|
||||
ARG PRELOAD_ML_MODELS=true
|
||||
ARG INCLUDE_LOCAL_MODELS=true
|
||||
|
||||
# =============================================================================
|
||||
# Stage: API Builder
|
||||
@@ -19,6 +27,7 @@ ARG INCLUDE_CP=true
|
||||
FROM python:3.11-slim AS api-builder
|
||||
|
||||
ARG INCLUDE_API
|
||||
ARG INCLUDE_LOCAL_MODELS
|
||||
RUN if [ "$INCLUDE_API" != "true" ]; then echo "Skipping API build" && exit 0; fi
|
||||
|
||||
WORKDIR /app
|
||||
@@ -37,6 +46,15 @@ COPY hindsight-api/README.md ./api/
|
||||
|
||||
WORKDIR /app/api
|
||||
|
||||
# Remove local ML model dependencies if INCLUDE_LOCAL_MODELS=false
|
||||
# This creates a smaller image when using external providers (TEI, OpenAI, Cohere)
|
||||
RUN if [ "$INCLUDE_LOCAL_MODELS" != "true" ]; then \
|
||||
echo "Removing local-models dependencies (sentence-transformers, torch, transformers)..." && \
|
||||
sed -i '/"sentence-transformers/d' pyproject.toml && \
|
||||
sed -i '/"transformers/d' pyproject.toml && \
|
||||
sed -i '/"torch/d' pyproject.toml; \
|
||||
fi
|
||||
|
||||
# Sync dependencies (will create lock file if needed)
|
||||
RUN uv sync
|
||||
|
||||
@@ -60,8 +78,8 @@ WORKDIR /app
|
||||
COPY package.json package-lock.json ./
|
||||
COPY hindsight-clients/typescript/ ./hindsight-clients/typescript/
|
||||
|
||||
# Install and build SDK using workspace
|
||||
RUN npm ci -w @vectorize-io/hindsight-client
|
||||
# Install and build SDK using workspace (--ignore-scripts skips git hooks setup)
|
||||
RUN npm ci --ignore-scripts -w @vectorize-io/hindsight-client
|
||||
RUN npm run build -w @vectorize-io/hindsight-client
|
||||
|
||||
# =============================================================================
|
||||
@@ -72,30 +90,48 @@ FROM node:20-slim AS cp-builder
|
||||
ARG INCLUDE_CP
|
||||
RUN if [ "$INCLUDE_CP" != "true" ]; then echo "Skipping CP build" && exit 0; fi
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Copy built SDK
|
||||
COPY --from=sdk-builder /app/hindsight-clients/typescript /app/sdk
|
||||
# Create directory structure matching the monorepo layout
|
||||
# This is required because build:standalone script expects .next/standalone/memory-poc/hindsight-control-plane
|
||||
WORKDIR /app/memory-poc/hindsight-control-plane
|
||||
|
||||
# Install Control Plane dependencies
|
||||
# Only copy package.json (not package-lock.json) to ensure npm installs
|
||||
# correct platform-specific native bindings for lightningcss/tailwindcss
|
||||
COPY hindsight-control-plane/package.json ./
|
||||
# Remove the file: dependency on SDK (we'll copy it directly later)
|
||||
RUN sed -i '/"@vectorize-io\/hindsight-client":/d' package.json
|
||||
RUN npm install
|
||||
|
||||
# Copy Control Plane source (excluding node_modules via .dockerignore)
|
||||
COPY hindsight-control-plane/ ./
|
||||
# Remove package-lock.json to avoid conflicts with installed native bindings
|
||||
RUN rm -f package-lock.json
|
||||
# Also remove the file: dependency from package.json (restored by COPY above)
|
||||
RUN rm -f package-lock.json && sed -i '/"@vectorize-io\/hindsight-client":/d' package.json
|
||||
|
||||
# Link SDK (temporary for build)
|
||||
RUN cd /app/sdk && npm link && cd /app && npm link @vectorize-io/hindsight-client
|
||||
# Copy built SDK directly into node_modules (more reliable than npm link in Docker)
|
||||
COPY --from=sdk-builder /app/hindsight-clients/typescript ./node_modules/@vectorize-io/hindsight-client
|
||||
|
||||
# Build Control Plane
|
||||
RUN npm run build
|
||||
# Build Control Plane - run next build first, then custom standalone copy
|
||||
# (The build:standalone script expects a specific path structure that differs in Docker)
|
||||
RUN npm exec -- next build
|
||||
|
||||
# Create public directory if it doesn't exist
|
||||
RUN mkdir -p public
|
||||
# Create standalone directory structure manually
|
||||
# Note: Must exclude node_modules from find to avoid wrong server.js from next/dist/experimental/testmode/
|
||||
# Note: Must explicitly copy .next since glob * doesn't match hidden directories
|
||||
RUN STANDALONE_ROOT=$(find .next/standalone -path '*/node_modules' -prune -o -name 'server.js' -print | head -1 | xargs dirname) && \
|
||||
mkdir -p standalone && \
|
||||
cp -r "$STANDALONE_ROOT"/* standalone/ && \
|
||||
cp -r "$STANDALONE_ROOT"/.next standalone/.next && \
|
||||
# Copy node_modules if separate from app dir (monorepo structure)
|
||||
if [ -d ".next/standalone/node_modules" ] && [ "$STANDALONE_ROOT" != ".next/standalone" ]; then \
|
||||
cp -r .next/standalone/node_modules standalone/node_modules; \
|
||||
fi && \
|
||||
cp -r .next/static standalone/.next/static && \
|
||||
mkdir -p standalone/public && \
|
||||
cp -r public/* standalone/public/ 2>/dev/null || true && \
|
||||
# Verify required files exist
|
||||
test -f standalone/server.js || (echo "ERROR: server.js missing!" && exit 1) && \
|
||||
test -f standalone/.next/BUILD_ID || (echo "ERROR: BUILD_ID missing!" && exit 1)
|
||||
|
||||
# =============================================================================
|
||||
# Stage: Final Image - API Only
|
||||
@@ -104,18 +140,18 @@ FROM python:3.11-slim AS api-only
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install pg0 dependencies
|
||||
# Note: libicu version varies by Debian version - try common versions in order
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
procps \
|
||||
libxml2 \
|
||||
libssl3 \
|
||||
libgssapi-krb5-2 \
|
||||
libossp-uuid16 \
|
||||
&& apt-get install -y libicu72 || apt-get install -y libicu74 || apt-get install -y libicu* \
|
||||
&& (apt-get install -y libicu72 2>/dev/null || apt-get install -y libicu74 2>/dev/null || apt-get install -y libicu76 2>/dev/null || true) \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& pip install --no-cache-dir uv
|
||||
|
||||
# Create non-root user (PostgreSQL cannot run as root)
|
||||
RUN useradd -m -s /bin/bash hindsight
|
||||
|
||||
# Copy API with virtual environment from builder
|
||||
@@ -125,28 +161,26 @@ COPY --from=api-builder /app/api /app/api
|
||||
COPY docker/standalone/start-all.sh /app/start-all.sh
|
||||
RUN chmod +x /app/start-all.sh
|
||||
|
||||
# Create data directory for pg0 and set ownership
|
||||
RUN mkdir -p /app/data && chown -R hindsight:hindsight /app
|
||||
RUN chown -R hindsight:hindsight /app
|
||||
|
||||
# Switch to non-root user
|
||||
USER hindsight
|
||||
|
||||
# Set PATH for hindsight user
|
||||
ENV PATH="/app/api/.venv/bin:${PATH}"
|
||||
|
||||
# Pre-cache PostgreSQL binaries by starting/stopping pg0-embedded
|
||||
ENV PG0_HOME=/home/hindsight/.pg0-cache
|
||||
|
||||
ENV PG0_HOME=/home/hindsight/.pg0
|
||||
|
||||
# Pre-download ML models to avoid runtime download
|
||||
RUN /app/api/.venv/bin/python -c "\
|
||||
# Pre-download ML models to avoid runtime download (conditional)
|
||||
# Only runs if both PRELOAD_ML_MODELS=true AND INCLUDE_LOCAL_MODELS=true
|
||||
ARG PRELOAD_ML_MODELS
|
||||
ARG INCLUDE_LOCAL_MODELS
|
||||
RUN if [ "$PRELOAD_ML_MODELS" = "true" ] && [ "$INCLUDE_LOCAL_MODELS" = "true" ]; then \
|
||||
/app/api/.venv/bin/python -c "\
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder; \
|
||||
print('Downloading embedding model...'); \
|
||||
SentenceTransformer('BAAI/bge-small-en-v1.5'); \
|
||||
print('Downloading cross-encoder model...'); \
|
||||
CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2'); \
|
||||
print('Models cached successfully')"
|
||||
print('Models cached successfully')"; \
|
||||
elif [ "$INCLUDE_LOCAL_MODELS" != "true" ]; then echo "Skipping ML model preload (local-models not included)"; \
|
||||
else echo "Skipping ML model preload"; fi
|
||||
|
||||
EXPOSE 8888
|
||||
|
||||
@@ -171,9 +205,9 @@ COPY --from=sdk-builder /app/hindsight-clients/typescript /app/sdk
|
||||
|
||||
# Copy Control Plane standalone build
|
||||
WORKDIR /app/control-plane
|
||||
COPY --from=cp-builder /app/.next/standalone ./
|
||||
COPY --from=cp-builder /app/.next/static ./.next/static
|
||||
COPY --from=cp-builder /app/public ./public
|
||||
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/standalone ./
|
||||
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/.next/static ./.next/static
|
||||
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/public ./public
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
@@ -200,20 +234,21 @@ FROM python:3.11-slim AS standalone
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install Node.js, curl, uv, and pg0 dependencies
|
||||
# Install Node.js, curl, uv, and system dependencies
|
||||
# Note: libicu version varies by Debian version - try common versions in order
|
||||
RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
procps \
|
||||
libxml2 \
|
||||
libssl3 \
|
||||
libgssapi-krb5-2 \
|
||||
libossp-uuid16 \
|
||||
&& apt-get install -y libicu72 || apt-get install -y libicu74 || apt-get install -y libicu* \
|
||||
&& (apt-get install -y libicu72 2>/dev/null || apt-get install -y libicu74 2>/dev/null || apt-get install -y libicu76 2>/dev/null || true) \
|
||||
&& curl -fsSL https://deb.nodesource.com/setup_20.x | bash - \
|
||||
&& apt-get install -y nodejs \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& pip install --no-cache-dir uv
|
||||
|
||||
# Create non-root user (PostgreSQL cannot run as root)
|
||||
RUN useradd -m -s /bin/bash hindsight
|
||||
|
||||
# Copy API with virtual environment from builder
|
||||
@@ -224,9 +259,9 @@ COPY --from=sdk-builder /app/hindsight-clients/typescript /app/sdk
|
||||
|
||||
# Copy Control Plane standalone build
|
||||
WORKDIR /app/control-plane
|
||||
COPY --from=cp-builder /app/.next/standalone ./
|
||||
COPY --from=cp-builder /app/.next/static ./.next/static
|
||||
COPY --from=cp-builder /app/public ./public
|
||||
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/standalone ./
|
||||
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/.next/static ./.next/static
|
||||
COPY --from=cp-builder /app/memory-poc/hindsight-control-plane/public ./public
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
@@ -234,35 +269,26 @@ WORKDIR /app
|
||||
COPY docker/standalone/start-all.sh /app/start-all.sh
|
||||
RUN chmod +x /app/start-all.sh
|
||||
|
||||
# Create data directory for pg0 and set ownership
|
||||
RUN mkdir -p /app/data && chown -R hindsight:hindsight /app
|
||||
RUN chown -R hindsight:hindsight /app
|
||||
|
||||
# Switch to non-root user
|
||||
USER hindsight
|
||||
|
||||
# Set PATH for hindsight user
|
||||
ENV PATH="/app/api/.venv/bin:${PATH}"
|
||||
|
||||
# Pre-cache PostgreSQL binaries by starting/stopping pg0-embedded
|
||||
ENV PG0_HOME=/home/hindsight/.pg0-cache
|
||||
RUN /app/api/.venv/bin/python -c "\
|
||||
from pg0 import Pg0; \
|
||||
print('Pre-caching PostgreSQL binaries...'); \
|
||||
pg = Pg0(name='hindsight', port=5555, username='hindsight', password='hindsight', database='hindsight'); \
|
||||
pg.start(); \
|
||||
pg.stop(); \
|
||||
print('PostgreSQL pre-cached to PG0_HOME')" || echo "Pre-download skipped"
|
||||
|
||||
ENV PG0_HOME=/home/hindsight/.pg0
|
||||
|
||||
# Pre-download ML models to avoid runtime download
|
||||
RUN /app/api/.venv/bin/python -c "\
|
||||
# Pre-download ML models to avoid runtime download (conditional)
|
||||
# Only runs if both PRELOAD_ML_MODELS=true AND INCLUDE_LOCAL_MODELS=true
|
||||
ARG PRELOAD_ML_MODELS
|
||||
ARG INCLUDE_LOCAL_MODELS
|
||||
RUN if [ "$PRELOAD_ML_MODELS" = "true" ] && [ "$INCLUDE_LOCAL_MODELS" = "true" ]; then \
|
||||
/app/api/.venv/bin/python -c "\
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder; \
|
||||
print('Downloading embedding model...'); \
|
||||
SentenceTransformer('BAAI/bge-small-en-v1.5'); \
|
||||
print('Downloading cross-encoder model...'); \
|
||||
CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2'); \
|
||||
print('Models cached successfully')"
|
||||
print('Models cached successfully')"; \
|
||||
elif [ "$INCLUDE_LOCAL_MODELS" != "true" ]; then echo "Skipping ML model preload (local-models not included)"; \
|
||||
else echo "Skipping ML model preload"; fi
|
||||
|
||||
EXPOSE 8888 9999
|
||||
|
||||
|
||||
@@ -5,16 +5,70 @@ set -e
|
||||
ENABLE_API="${HINDSIGHT_ENABLE_API:-true}"
|
||||
ENABLE_CP="${HINDSIGHT_ENABLE_CP:-true}"
|
||||
|
||||
# Copy pre-cached PostgreSQL data if runtime directory is empty (first run with volume)
|
||||
if [ "$ENABLE_API" = "true" ]; then
|
||||
PG0_CACHE="/home/hindsight/.pg0-cache"
|
||||
PG0_HOME="/home/hindsight/.pg0"
|
||||
if [ -d "$PG0_CACHE" ] && [ "$(ls -A $PG0_CACHE 2>/dev/null)" ]; then
|
||||
if [ ! "$(ls -A $PG0_HOME 2>/dev/null)" ]; then
|
||||
echo "📦 Copying pre-cached PostgreSQL data..."
|
||||
cp -r "$PG0_CACHE"/* "$PG0_HOME"/ 2>/dev/null || true
|
||||
fi
|
||||
# =============================================================================
|
||||
# Dependency waiting (opt-in via HINDSIGHT_WAIT_FOR_DEPS=true)
|
||||
#
|
||||
# Problem: When running with LM Studio, the LLM may take time to load models.
|
||||
# If Hindsight starts before LM Studio is ready, it fails on LLM verification.
|
||||
# This wait loop ensures dependencies are ready before starting.
|
||||
# =============================================================================
|
||||
if [ "${HINDSIGHT_WAIT_FOR_DEPS:-false}" = "true" ]; then
|
||||
LLM_BASE_URL="${HINDSIGHT_API_LLM_BASE_URL:-http://host.docker.internal:1234/v1}"
|
||||
MAX_RETRIES="${HINDSIGHT_RETRY_MAX:-0}" # 0 = infinite
|
||||
RETRY_INTERVAL="${HINDSIGHT_RETRY_INTERVAL:-10}"
|
||||
|
||||
# Check if external database is configured (skip check for embedded pg0)
|
||||
SKIP_DB_CHECK=false
|
||||
if [ -z "${HINDSIGHT_API_DATABASE_URL}" ]; then
|
||||
SKIP_DB_CHECK=true
|
||||
else
|
||||
DB_CHECK_HOST=$(echo "$HINDSIGHT_API_DATABASE_URL" | sed -E 's|.*@([^:/]+):([0-9]+)/.*|\1 \2|')
|
||||
fi
|
||||
|
||||
check_db() {
|
||||
if $SKIP_DB_CHECK; then
|
||||
return 0
|
||||
fi
|
||||
if command -v pg_isready &> /dev/null; then
|
||||
pg_isready -h $(echo $DB_CHECK_HOST | cut -d' ' -f1) -p $(echo $DB_CHECK_HOST | cut -d' ' -f2) &>/dev/null
|
||||
else
|
||||
python3 -c "import socket; s=socket.socket(); s.settimeout(5); exit(0 if s.connect_ex(('$(echo $DB_CHECK_HOST | cut -d' ' -f1)', $(echo $DB_CHECK_HOST | cut -d' ' -f2))) == 0 else 1)" 2>/dev/null
|
||||
fi
|
||||
}
|
||||
|
||||
check_llm() {
|
||||
curl -sf "${LLM_BASE_URL}/models" --connect-timeout 5 &>/dev/null
|
||||
}
|
||||
|
||||
echo "⏳ Waiting for dependencies to be ready..."
|
||||
attempt=1
|
||||
|
||||
while true; do
|
||||
db_ok=false
|
||||
llm_ok=false
|
||||
|
||||
if check_db; then
|
||||
db_ok=true
|
||||
fi
|
||||
|
||||
if check_llm; then
|
||||
llm_ok=true
|
||||
fi
|
||||
|
||||
if $db_ok && $llm_ok; then
|
||||
echo "✅ Dependencies ready!"
|
||||
break
|
||||
fi
|
||||
|
||||
if [ "$MAX_RETRIES" -ne 0 ] && [ "$attempt" -ge "$MAX_RETRIES" ]; then
|
||||
echo "❌ Max retries ($MAX_RETRIES) reached. Dependencies not available."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo " Attempt $attempt: DB=$( $db_ok && echo 'ok' || echo 'waiting' ), LLM=$( $llm_ok && echo 'ok' || echo 'waiting' )"
|
||||
sleep "$RETRY_INTERVAL"
|
||||
((attempt++))
|
||||
done
|
||||
fi
|
||||
|
||||
# Track PIDs for wait
|
||||
@@ -23,7 +77,8 @@ PIDS=()
|
||||
# Start API if enabled
|
||||
if [ "$ENABLE_API" = "true" ]; then
|
||||
cd /app/api
|
||||
hindsight-api 2>&1 | sed -u 's/^/[api] /' &
|
||||
# Run API directly - Python's PYTHONUNBUFFERED=1 handles output buffering
|
||||
hindsight-api &
|
||||
API_PID=$!
|
||||
PIDS+=($API_PID)
|
||||
|
||||
@@ -42,7 +97,7 @@ fi
|
||||
if [ "$ENABLE_CP" = "true" ]; then
|
||||
echo "🎛️ Starting Control Plane..."
|
||||
cd /app/control-plane
|
||||
PORT=9999 node server.js 2>&1 | grep -v -E "^[[:space:]]*(▲|✓|-|$)" | sed -u 's/^/[control-plane] /' &
|
||||
PORT=9999 node server.js &
|
||||
CP_PID=$!
|
||||
PIDS+=($CP_PID)
|
||||
else
|
||||
|
||||
@@ -2,8 +2,8 @@ apiVersion: v2
|
||||
name: hindsight
|
||||
description: Hindsight helm chart
|
||||
type: application
|
||||
version: 0.1.5
|
||||
appVersion: "0.1.5"
|
||||
version: 0.3.0
|
||||
appVersion: "0.3.0"
|
||||
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
|
||||
*/}}
|
||||
@@ -110,3 +126,14 @@ API URL for control plane
|
||||
{{- define "hindsight.apiUrl" -}}
|
||||
{{- printf "http://%s-api:%d" (include "hindsight.fullname" .) (.Values.api.service.port | int) }}
|
||||
{{- end }}
|
||||
|
||||
{{/*
|
||||
Get the name of the secret to use
|
||||
*/}}
|
||||
{{- define "hindsight.secretName" -}}
|
||||
{{- if .Values.existingSecret }}
|
||||
{{- .Values.existingSecret }}
|
||||
{{- else }}
|
||||
{{- printf "%s-secret" (include "hindsight.fullname" .) }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
|
||||
@@ -15,7 +15,9 @@ spec:
|
||||
template:
|
||||
metadata:
|
||||
annotations:
|
||||
{{- if not .Values.existingSecret }}
|
||||
checksum/secret: {{ include (print $.Template.BasePath "/secret.yaml") . | sha256sum }}
|
||||
{{- end }}
|
||||
{{- with .Values.podAnnotations }}
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
@@ -37,27 +39,41 @@ spec:
|
||||
- name: http
|
||||
containerPort: {{ .Values.api.service.targetPort }}
|
||||
protocol: TCP
|
||||
{{- if .Values.existingSecret }}
|
||||
envFrom:
|
||||
- secretRef:
|
||||
name: {{ .Values.existingSecret }}
|
||||
{{- end }}
|
||||
env:
|
||||
- name: HINDSIGHT_API_DATABASE_URL
|
||||
value: {{ include "hindsight.databaseUrl" . | quote }}
|
||||
{{- /* POSTGRES_PASSWORD must be defined before DATABASE_URL for $(VAR) interpolation */}}
|
||||
{{- if not .Values.postgresql.enabled }}
|
||||
- name: POSTGRES_PASSWORD
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: {{ include "hindsight.fullname" . }}-secret
|
||||
name: {{ include "hindsight.secretName" . }}
|
||||
key: postgres-password
|
||||
{{- end }}
|
||||
- name: HINDSIGHT_API_DATABASE_URL
|
||||
value: {{ include "hindsight.databaseUrl" . | quote }}
|
||||
{{- /* Disable internal worker when dedicated workers are enabled */}}
|
||||
{{- if .Values.worker.enabled }}
|
||||
- name: HINDSIGHT_API_WORKER_ENABLED
|
||||
value: "false"
|
||||
{{- end }}
|
||||
{{- range $key, $value := .Values.api.env }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
{{- /* Only use api.secrets when not using existingSecret (for chart-managed secrets) */}}
|
||||
{{- if not .Values.existingSecret }}
|
||||
{{- range $key, $value := .Values.api.secrets }}
|
||||
- name: {{ $key }}
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: {{ include "hindsight.fullname" $ }}-secret
|
||||
name: {{ include "hindsight.secretName" $ }}
|
||||
key: {{ $key }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
livenessProbe:
|
||||
{{- toYaml .Values.api.livenessProbe | nindent 10 }}
|
||||
readinessProbe:
|
||||
|
||||
@@ -15,7 +15,9 @@ spec:
|
||||
template:
|
||||
metadata:
|
||||
annotations:
|
||||
{{- if not .Values.existingSecret }}
|
||||
checksum/secret: {{ include (print $.Template.BasePath "/secret.yaml") . | sha256sum }}
|
||||
{{- end }}
|
||||
{{- with .Values.podAnnotations }}
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
@@ -37,6 +39,11 @@ spec:
|
||||
- name: http
|
||||
containerPort: {{ .Values.controlPlane.service.targetPort }}
|
||||
protocol: TCP
|
||||
{{- if .Values.existingSecret }}
|
||||
envFrom:
|
||||
- secretRef:
|
||||
name: {{ .Values.existingSecret }}
|
||||
{{- end }}
|
||||
env:
|
||||
- name: HINDSIGHT_CP_DATAPLANE_API_URL
|
||||
value: {{ include "hindsight.apiUrl" . | quote }}
|
||||
@@ -44,13 +51,16 @@ spec:
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
{{- /* Only use controlPlane.secrets when not using existingSecret (for chart-managed secrets) */}}
|
||||
{{- if not .Values.existingSecret }}
|
||||
{{- range $key, $value := .Values.controlPlane.secrets }}
|
||||
- name: {{ $key }}
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: {{ include "hindsight.fullname" $ }}-secret
|
||||
name: {{ include "hindsight.secretName" $ }}
|
||||
key: {{ $key }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
livenessProbe:
|
||||
{{- toYaml .Values.controlPlane.livenessProbe | nindent 10 }}
|
||||
readinessProbe:
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
{{- if not .Values.existingSecret }}
|
||||
apiVersion: v1
|
||||
kind: Secret
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-secret
|
||||
name: {{ include "hindsight.secretName" . }}
|
||||
labels:
|
||||
{{- include "hindsight.labels" . | nindent 4 }}
|
||||
type: Opaque
|
||||
@@ -15,3 +16,4 @@ data:
|
||||
{{- if and (not .Values.postgresql.enabled) .Values.postgresql.external.password }}
|
||||
postgres-password: {{ .Values.postgresql.external.password | b64enc | quote }}
|
||||
{{- end }}
|
||||
{{- 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 }}"
|
||||
imagePullPolicy: {{ .Values.worker.image.pullPolicy }}
|
||||
command: ["hindsight-worker"]
|
||||
ports:
|
||||
- name: http
|
||||
containerPort: {{ .Values.worker.service.targetPort }}
|
||||
protocol: TCP
|
||||
{{- if .Values.existingSecret }}
|
||||
envFrom:
|
||||
- secretRef:
|
||||
name: {{ .Values.existingSecret }}
|
||||
{{- end }}
|
||||
env:
|
||||
{{- /* POSTGRES_PASSWORD must be defined before DATABASE_URL for $(VAR) interpolation */}}
|
||||
{{- if not .Values.postgresql.enabled }}
|
||||
- name: POSTGRES_PASSWORD
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: {{ include "hindsight.secretName" . }}
|
||||
key: postgres-password
|
||||
{{- end }}
|
||||
- name: HINDSIGHT_API_DATABASE_URL
|
||||
value: {{ include "hindsight.databaseUrl" . | quote }}
|
||||
{{- /* Worker ID uses pod name (StatefulSet provides stable names like worker-0, worker-1) */}}
|
||||
- name: HINDSIGHT_API_WORKER_ID
|
||||
valueFrom:
|
||||
fieldRef:
|
||||
fieldPath: metadata.name
|
||||
{{- /* Inherit LLM config from api.env */}}
|
||||
{{- range $key, $value := .Values.api.env }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
{{- /* Worker-specific env vars */}}
|
||||
{{- range $key, $value := .Values.worker.env }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
{{- /* Only use secrets when not using existingSecret */}}
|
||||
{{- if not .Values.existingSecret }}
|
||||
{{- /* Inherit secrets from api.secrets */}}
|
||||
{{- range $key, $value := .Values.api.secrets }}
|
||||
- name: {{ $key }}
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: {{ include "hindsight.secretName" $ }}
|
||||
key: {{ $key }}
|
||||
{{- end }}
|
||||
{{- /* Worker-specific secrets (can override api.secrets) */}}
|
||||
{{- range $key, $value := .Values.worker.secrets }}
|
||||
- name: {{ $key }}
|
||||
valueFrom:
|
||||
secretKeyRef:
|
||||
name: {{ include "hindsight.secretName" $ }}
|
||||
key: {{ $key }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
livenessProbe:
|
||||
{{- toYaml .Values.worker.livenessProbe | nindent 10 }}
|
||||
readinessProbe:
|
||||
{{- toYaml .Values.worker.readinessProbe | nindent 10 }}
|
||||
resources:
|
||||
{{- toYaml .Values.worker.resources | nindent 10 }}
|
||||
{{- with .Values.nodeSelector }}
|
||||
nodeSelector:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.affinity }}
|
||||
affinity:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.tolerations }}
|
||||
tolerations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
@@ -3,6 +3,15 @@
|
||||
# Chart version - use this to set a consistent image tag across all components
|
||||
version: "0.1.1"
|
||||
|
||||
# Use an existing secret instead of creating one from values
|
||||
# When set, all keys from this secret are injected as environment variables via envFrom
|
||||
# Required keys:
|
||||
# - postgres-password: PostgreSQL password (when postgresql.enabled=false)
|
||||
# Optional keys (any key becomes an env var):
|
||||
# - HINDSIGHT_API_LLM_API_KEY: API key for LLM provider
|
||||
# - Any other env vars you want to inject
|
||||
# existingSecret: "my-hindsight-secret"
|
||||
|
||||
# Global settings
|
||||
replicaCount: 1
|
||||
|
||||
@@ -58,6 +67,63 @@ api:
|
||||
# HINDSIGHT_API_LLM_API_KEY: "your-api-key"
|
||||
# HINDSIGHT_API_LLM_BASE_URL: "https://api.groq.com/openai/v1"
|
||||
|
||||
# Worker settings (distributed task processing)
|
||||
# When enabled, dedicated worker pods process tasks and the API's internal worker is disabled
|
||||
worker:
|
||||
enabled: false
|
||||
replicaCount: 2
|
||||
image:
|
||||
repository: ghcr.io/vectorize-io/hindsight-api
|
||||
pullPolicy: IfNotPresent
|
||||
# tag defaults to .Values.version if not specified
|
||||
|
||||
service:
|
||||
# Service for metrics scraping (headless for StatefulSet)
|
||||
port: 8889
|
||||
targetPort: 8889
|
||||
|
||||
# Resource limits and requests
|
||||
resources:
|
||||
limits:
|
||||
cpu: 2000m
|
||||
memory: 4Gi
|
||||
requests:
|
||||
cpu: 500m
|
||||
memory: 1Gi
|
||||
|
||||
# Liveness and readiness probes
|
||||
livenessProbe:
|
||||
httpGet:
|
||||
path: /health
|
||||
port: 8889
|
||||
initialDelaySeconds: 30
|
||||
periodSeconds: 10
|
||||
timeoutSeconds: 5
|
||||
failureThreshold: 3
|
||||
|
||||
readinessProbe:
|
||||
httpGet:
|
||||
path: /health
|
||||
port: 8889
|
||||
initialDelaySeconds: 10
|
||||
periodSeconds: 5
|
||||
timeoutSeconds: 3
|
||||
failureThreshold: 3
|
||||
|
||||
# Worker-specific environment variables
|
||||
env:
|
||||
# Poll interval in milliseconds (how often to check for new tasks)
|
||||
HINDSIGHT_API_WORKER_POLL_INTERVAL_MS: "500"
|
||||
# Number of tasks to claim per poll cycle
|
||||
HINDSIGHT_API_WORKER_BATCH_SIZE: "10"
|
||||
# Max retries before marking a task as failed
|
||||
HINDSIGHT_API_WORKER_MAX_RETRIES: "3"
|
||||
# HTTP port for metrics/health (matches service.targetPort)
|
||||
HINDSIGHT_API_WORKER_HTTP_PORT: "8889"
|
||||
|
||||
# Secret environment variables (inherited from api.secrets if not specified)
|
||||
secrets: {}
|
||||
|
||||
# Image settings for control plane
|
||||
controlPlane:
|
||||
enabled: true
|
||||
|
||||
+137
-1
@@ -1 +1,137 @@
|
||||
# Memory
|
||||
# Hindsight API
|
||||
|
||||
**Memory System for AI Agents** — Temporal + Semantic + Entity Memory Architecture using PostgreSQL with pgvector.
|
||||
|
||||
Hindsight gives AI agents persistent memory that works like human memory: it stores facts, tracks entities and relationships, handles temporal reasoning ("what happened last spring?"), and forms opinions based on configurable disposition traits.
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
pip install hindsight-api
|
||||
```
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Run the Server
|
||||
|
||||
```bash
|
||||
# Set your LLM provider
|
||||
export HINDSIGHT_API_LLM_PROVIDER=openai
|
||||
export HINDSIGHT_API_LLM_API_KEY=sk-xxxxxxxxxxxx
|
||||
|
||||
# Start the server (uses embedded PostgreSQL by default)
|
||||
hindsight-api
|
||||
```
|
||||
|
||||
The server starts at http://localhost:8888 with:
|
||||
- REST API for memory operations
|
||||
- MCP server at `/mcp` for tool-use integration
|
||||
|
||||
### Use the Python API
|
||||
|
||||
```python
|
||||
from hindsight_api import MemoryEngine
|
||||
|
||||
# Create and initialize the memory engine
|
||||
memory = MemoryEngine()
|
||||
await memory.initialize()
|
||||
|
||||
# Create a memory bank for your agent
|
||||
bank = await memory.create_memory_bank(
|
||||
name="my-assistant",
|
||||
background="A helpful coding assistant"
|
||||
)
|
||||
|
||||
# Store a memory
|
||||
await memory.retain(
|
||||
memory_bank_id=bank.id,
|
||||
content="The user prefers Python for data science projects"
|
||||
)
|
||||
|
||||
# Recall memories
|
||||
results = await memory.recall(
|
||||
memory_bank_id=bank.id,
|
||||
query="What programming language does the user prefer?"
|
||||
)
|
||||
|
||||
# Reflect with reasoning
|
||||
response = await memory.reflect(
|
||||
memory_bank_id=bank.id,
|
||||
query="Should I recommend Python or R for this ML project?"
|
||||
)
|
||||
```
|
||||
|
||||
## CLI Options
|
||||
|
||||
```bash
|
||||
hindsight-api --help
|
||||
|
||||
# Common options
|
||||
hindsight-api --port 9000 # Custom port (default: 8888)
|
||||
hindsight-api --host 127.0.0.1 # Bind to localhost only
|
||||
hindsight-api --workers 4 # Multiple worker processes
|
||||
hindsight-api --log-level debug # Verbose logging
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
Configure via environment variables:
|
||||
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `HINDSIGHT_API_DATABASE_URL` | PostgreSQL connection string | `pg0` (embedded) |
|
||||
| `HINDSIGHT_API_LLM_PROVIDER` | `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio` | `openai` |
|
||||
| `HINDSIGHT_API_LLM_API_KEY` | API key for LLM provider | - |
|
||||
| `HINDSIGHT_API_LLM_MODEL` | Model name | `gpt-4o-mini` |
|
||||
| `HINDSIGHT_API_HOST` | Server bind address | `0.0.0.0` |
|
||||
| `HINDSIGHT_API_PORT` | Server port | `8888` |
|
||||
|
||||
### Example with External PostgreSQL
|
||||
|
||||
```bash
|
||||
export HINDSIGHT_API_DATABASE_URL=postgresql://user:pass@localhost:5432/hindsight
|
||||
export HINDSIGHT_API_LLM_PROVIDER=groq
|
||||
export HINDSIGHT_API_LLM_API_KEY=gsk_xxxxxxxxxxxx
|
||||
|
||||
hindsight-api
|
||||
```
|
||||
|
||||
## Docker
|
||||
|
||||
```bash
|
||||
docker run --rm -it -p 8888:8888 \
|
||||
-e HINDSIGHT_API_LLM_API_KEY=$OPENAI_API_KEY \
|
||||
-v $HOME/.hindsight-docker:/home/hindsight/.pg0 \
|
||||
ghcr.io/vectorize-io/hindsight:latest
|
||||
```
|
||||
|
||||
## MCP Server
|
||||
|
||||
For local MCP integration without running the full API server:
|
||||
|
||||
```bash
|
||||
hindsight-local-mcp
|
||||
```
|
||||
|
||||
This runs a stdio-based MCP server that can be used directly with MCP-compatible clients.
|
||||
|
||||
## Key Features
|
||||
|
||||
- **Multi-Strategy Retrieval (TEMPR)** — Semantic, keyword, graph, and temporal search combined with RRF fusion
|
||||
- **Entity Graph** — Automatic entity extraction and relationship tracking
|
||||
- **Temporal Reasoning** — Native support for time-based queries
|
||||
- **Disposition Traits** — Configurable skepticism, literalism, and empathy influence opinion formation
|
||||
- **Three Memory Types** — World facts, bank actions, and formed opinions with confidence scores
|
||||
|
||||
## Documentation
|
||||
|
||||
Full documentation: [https://hindsight.vectorize.io](https://hindsight.vectorize.io)
|
||||
|
||||
- [Installation Guide](https://hindsight.vectorize.io/developer/installation)
|
||||
- [Configuration Reference](https://hindsight.vectorize.io/developer/configuration)
|
||||
- [API Reference](https://hindsight.vectorize.io/api-reference)
|
||||
- [Python SDK](https://hindsight.vectorize.io/sdks/python)
|
||||
|
||||
## License
|
||||
|
||||
Apache 2.0
|
||||
|
||||
@@ -3,26 +3,29 @@ Memory System for AI Agents.
|
||||
|
||||
Temporal + Semantic Memory Architecture using PostgreSQL with pgvector.
|
||||
"""
|
||||
|
||||
from .config import HindsightConfig, get_config
|
||||
from .engine.cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
|
||||
from .engine.embeddings import Embeddings, LocalSTEmbeddings, RemoteTEIEmbeddings
|
||||
from .engine.llm_wrapper import LLMConfig
|
||||
from .engine.memory_engine import MemoryEngine
|
||||
from .engine.search.trace import (
|
||||
SearchTrace,
|
||||
QueryInfo,
|
||||
EntryPoint,
|
||||
NodeVisit,
|
||||
WeightComponents,
|
||||
LinkInfo,
|
||||
NodeVisit,
|
||||
PruningDecision,
|
||||
SearchSummary,
|
||||
QueryInfo,
|
||||
SearchPhaseMetrics,
|
||||
SearchSummary,
|
||||
SearchTrace,
|
||||
WeightComponents,
|
||||
)
|
||||
from .engine.search.tracer import SearchTracer
|
||||
from .engine.embeddings import Embeddings, LocalSTEmbeddings, RemoteTEIEmbeddings
|
||||
from .engine.cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
|
||||
from .engine.llm_wrapper import LLMConfig
|
||||
from .config import HindsightConfig, get_config
|
||||
from .models import RequestContext
|
||||
|
||||
__all__ = [
|
||||
"MemoryEngine",
|
||||
"RequestContext",
|
||||
"HindsightConfig",
|
||||
"get_config",
|
||||
"SearchTrace",
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# Admin CLI for Hindsight
|
||||
@@ -0,0 +1,311 @@
|
||||
"""
|
||||
Hindsight Admin CLI - backup and restore operations.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import zipfile
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import asyncpg
|
||||
import typer
|
||||
|
||||
from ..config import HindsightConfig
|
||||
from ..pg0 import parse_pg0_url, resolve_database_url
|
||||
|
||||
|
||||
def _fq_table(table: str, schema: str) -> str:
|
||||
"""Get fully-qualified table name with schema prefix."""
|
||||
return f"{schema}.{table}"
|
||||
|
||||
|
||||
# Setup logging
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(message)s",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
app = typer.Typer(name="hindsight-admin", help="Hindsight administrative commands")
|
||||
|
||||
# Tables to backup/restore in dependency order
|
||||
# Import must happen in this order due to foreign key constraints
|
||||
BACKUP_TABLES = [
|
||||
"banks",
|
||||
"documents",
|
||||
"entities",
|
||||
"chunks",
|
||||
"memory_units",
|
||||
"unit_entities",
|
||||
"entity_cooccurrences",
|
||||
"memory_links",
|
||||
]
|
||||
|
||||
MANIFEST_VERSION = "1"
|
||||
|
||||
|
||||
async def _backup(database_url: str, output_path: Path, schema: str = "public") -> dict[str, Any]:
|
||||
"""Backup all tables to a zip file using binary COPY protocol."""
|
||||
conn = await asyncpg.connect(database_url)
|
||||
try:
|
||||
tables: dict[str, Any] = {}
|
||||
manifest: dict[str, Any] = {
|
||||
"version": MANIFEST_VERSION,
|
||||
"created_at": datetime.now(timezone.utc).isoformat(),
|
||||
"schema": schema,
|
||||
"tables": tables,
|
||||
}
|
||||
|
||||
# Use a transaction with REPEATABLE READ isolation to get a consistent
|
||||
# snapshot across all tables. This prevents race conditions where
|
||||
# entity_cooccurrences could reference entities created after the
|
||||
# entities table was backed up.
|
||||
async with conn.transaction(isolation="repeatable_read"):
|
||||
with zipfile.ZipFile(output_path, "w", zipfile.ZIP_DEFLATED) as zf:
|
||||
for i, table in enumerate(BACKUP_TABLES, 1):
|
||||
typer.echo(f" [{i}/{len(BACKUP_TABLES)}] Backing up {table}...", nl=False)
|
||||
|
||||
buffer = io.BytesIO()
|
||||
|
||||
# Use binary COPY for exact type preservation
|
||||
# asyncpg requires schema_name as separate parameter
|
||||
await conn.copy_from_table(table, schema_name=schema, output=buffer, format="binary")
|
||||
|
||||
data = buffer.getvalue()
|
||||
zf.writestr(f"{table}.bin", data)
|
||||
|
||||
# Get row count for manifest
|
||||
qualified_table = _fq_table(table, schema)
|
||||
row_count = await conn.fetchval(f"SELECT COUNT(*) FROM {qualified_table}")
|
||||
tables[table] = {
|
||||
"rows": row_count,
|
||||
"size_bytes": len(data),
|
||||
}
|
||||
|
||||
typer.echo(f" {row_count} rows")
|
||||
|
||||
zf.writestr("manifest.json", json.dumps(manifest, indent=2))
|
||||
|
||||
return manifest
|
||||
finally:
|
||||
await conn.close()
|
||||
|
||||
|
||||
async def _restore(database_url: str, input_path: Path, schema: str = "public") -> dict[str, Any]:
|
||||
"""Restore all tables from a zip file using binary COPY protocol."""
|
||||
conn = await asyncpg.connect(database_url)
|
||||
try:
|
||||
with zipfile.ZipFile(input_path, "r") as zf:
|
||||
# Read and validate manifest
|
||||
manifest: dict[str, Any] = json.loads(zf.read("manifest.json"))
|
||||
if manifest.get("version") != MANIFEST_VERSION:
|
||||
raise ValueError(f"Unsupported backup version: {manifest.get('version')}")
|
||||
|
||||
# Use a transaction for atomic restore - either all tables are
|
||||
# restored or none are, preventing partial/inconsistent state.
|
||||
async with conn.transaction():
|
||||
typer.echo(" Clearing existing data...")
|
||||
# Truncate tables in reverse order (respects FK constraints)
|
||||
for table in reversed(BACKUP_TABLES):
|
||||
qualified_table = _fq_table(table, schema)
|
||||
await conn.execute(f"TRUNCATE TABLE {qualified_table} CASCADE")
|
||||
|
||||
# Restore tables in forward order
|
||||
for i, table in enumerate(BACKUP_TABLES, 1):
|
||||
filename = f"{table}.bin"
|
||||
if filename not in zf.namelist():
|
||||
typer.echo(f" [{i}/{len(BACKUP_TABLES)}] {table}: skipped (not in backup)")
|
||||
continue
|
||||
|
||||
expected_rows = manifest["tables"].get(table, {}).get("rows", "?")
|
||||
typer.echo(f" [{i}/{len(BACKUP_TABLES)}] Restoring {table}... {expected_rows} rows")
|
||||
|
||||
data = zf.read(filename)
|
||||
buffer = io.BytesIO(data)
|
||||
# asyncpg requires schema_name as separate parameter
|
||||
await conn.copy_to_table(table, schema_name=schema, source=buffer, format="binary")
|
||||
|
||||
# Refresh materialized view
|
||||
typer.echo(" Refreshing materialized views...")
|
||||
await conn.execute(f"REFRESH MATERIALIZED VIEW {_fq_table('memory_units_bm25', schema)}")
|
||||
|
||||
return manifest
|
||||
finally:
|
||||
await conn.close()
|
||||
|
||||
|
||||
async def _run_backup(db_url: str, output: Path, schema: str = "public") -> dict[str, Any]:
|
||||
"""Resolve database URL and run backup."""
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
return await _backup(resolved_url, output, schema)
|
||||
|
||||
|
||||
async def _run_restore(db_url: str, input_file: Path, schema: str = "public") -> dict[str, Any]:
|
||||
"""Resolve database URL and run restore."""
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
return await _restore(resolved_url, input_file, schema)
|
||||
|
||||
|
||||
@app.command()
|
||||
def backup(
|
||||
output: Path = typer.Argument(..., help="Output file path (.zip)"),
|
||||
schema: str = typer.Option("public", "--schema", "-s", help="Database schema to backup"),
|
||||
):
|
||||
"""Backup the Hindsight database to a zip file."""
|
||||
config = HindsightConfig.from_env()
|
||||
|
||||
if not config.database_url:
|
||||
typer.echo("Error: Database URL not configured.", err=True)
|
||||
typer.echo("Set HINDSIGHT_API_DATABASE_URL environment variable.", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
if output.suffix != ".zip":
|
||||
output = output.with_suffix(".zip")
|
||||
|
||||
typer.echo(f"Backing up database (schema: {schema}) to {output}...")
|
||||
|
||||
manifest = asyncio.run(_run_backup(config.database_url, output, schema))
|
||||
|
||||
total_rows = sum(t["rows"] for t in manifest["tables"].values())
|
||||
typer.echo(f"Backed up {total_rows} rows across {len(BACKUP_TABLES)} tables")
|
||||
typer.echo(f"Backup saved to {output}")
|
||||
|
||||
|
||||
@app.command()
|
||||
def restore(
|
||||
input_file: Path = typer.Argument(..., help="Input backup file (.zip)"),
|
||||
schema: str = typer.Option("public", "--schema", "-s", help="Database schema to restore to"),
|
||||
yes: bool = typer.Option(False, "--yes", "-y", help="Skip confirmation prompt"),
|
||||
):
|
||||
"""Restore the database from a backup file. WARNING: This deletes all existing data."""
|
||||
config = HindsightConfig.from_env()
|
||||
|
||||
if not config.database_url:
|
||||
typer.echo("Error: Database URL not configured.", err=True)
|
||||
typer.echo("Set HINDSIGHT_API_DATABASE_URL environment variable.", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
if not input_file.exists():
|
||||
typer.echo(f"Error: File not found: {input_file}", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
if not yes:
|
||||
typer.confirm(
|
||||
"This will DELETE all existing data and replace it with the backup. Continue?",
|
||||
abort=True,
|
||||
)
|
||||
|
||||
typer.echo(f"Restoring database (schema: {schema}) from {input_file}...")
|
||||
|
||||
manifest = asyncio.run(_run_restore(config.database_url, input_file, schema))
|
||||
|
||||
total_rows = sum(t["rows"] for t in manifest["tables"].values())
|
||||
typer.echo(f"Restored {total_rows} rows across {len(BACKUP_TABLES)} tables")
|
||||
typer.echo("Restore complete")
|
||||
|
||||
|
||||
async def _run_migration(db_url: str, schema: str = "public") -> None:
|
||||
"""Resolve database URL and run migrations."""
|
||||
from ..migrations import run_migrations
|
||||
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
run_migrations(resolved_url, schema=schema)
|
||||
|
||||
|
||||
@app.command(name="run-db-migration")
|
||||
def run_db_migration(
|
||||
schema: str = typer.Option("public", "--schema", "-s", help="Database schema to run migrations on"),
|
||||
):
|
||||
"""Run database migrations to the latest version."""
|
||||
config = HindsightConfig.from_env()
|
||||
|
||||
if not config.database_url:
|
||||
typer.echo("Error: Database URL not configured.", err=True)
|
||||
typer.echo("Set HINDSIGHT_API_DATABASE_URL environment variable.", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
typer.echo(f"Running database migrations (schema: {schema})...")
|
||||
|
||||
asyncio.run(_run_migration(config.database_url, schema))
|
||||
|
||||
typer.echo("Database migrations completed successfully")
|
||||
|
||||
|
||||
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()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -2,20 +2,19 @@
|
||||
Alembic environment configuration for SQLAlchemy with pgvector.
|
||||
Uses synchronous psycopg2 driver for migrations to avoid pgbouncer issues.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy import pool, engine_from_config
|
||||
from sqlalchemy.engine import Connection
|
||||
|
||||
from alembic import context
|
||||
from dotenv import load_dotenv
|
||||
from sqlalchemy import engine_from_config, pool
|
||||
|
||||
# Import your models here
|
||||
from hindsight_api.models import Base
|
||||
|
||||
|
||||
# Load environment variables based on HINDSIGHT_API_DATABASE_URL env var or default to local
|
||||
def load_env():
|
||||
"""Load environment variables from .env"""
|
||||
@@ -30,6 +29,7 @@ def load_env():
|
||||
if env_file.exists():
|
||||
load_dotenv(env_file)
|
||||
|
||||
|
||||
load_env()
|
||||
|
||||
# this is the Alembic Config object, which provides
|
||||
@@ -109,6 +109,9 @@ def run_migrations_online() -> None:
|
||||
|
||||
get_database_url() # Process and set the database URL in config
|
||||
|
||||
# Check if we're targeting a specific schema (for multi-tenant isolation)
|
||||
target_schema = config.get_main_option("target_schema")
|
||||
|
||||
connectable = engine_from_config(
|
||||
config.get_section(config.config_ini_section, {}),
|
||||
prefix="sqlalchemy.",
|
||||
@@ -121,17 +124,34 @@ def run_migrations_online() -> None:
|
||||
def set_read_write_mode(dbapi_connection, connection_record):
|
||||
cursor = dbapi_connection.cursor()
|
||||
cursor.execute("SET SESSION CHARACTERISTICS AS TRANSACTION READ WRITE")
|
||||
# If targeting a specific schema, set search_path
|
||||
# Include public in search_path for access to shared extensions (pgvector)
|
||||
if target_schema:
|
||||
cursor.execute(f'CREATE SCHEMA IF NOT EXISTS "{target_schema}"')
|
||||
cursor.execute(f'SET search_path TO "{target_schema}", public')
|
||||
cursor.close()
|
||||
|
||||
with connectable.connect() as connection:
|
||||
# Also explicitly set read-write mode on this connection
|
||||
connection.execute(text("SET SESSION CHARACTERISTICS AS TRANSACTION READ WRITE"))
|
||||
|
||||
# If targeting a specific schema, set search_path
|
||||
# Include public in search_path for access to shared extensions (pgvector)
|
||||
if target_schema:
|
||||
connection.execute(text(f'CREATE SCHEMA IF NOT EXISTS "{target_schema}"'))
|
||||
connection.execute(text(f'SET search_path TO "{target_schema}", public'))
|
||||
|
||||
connection.commit() # Commit the SET command
|
||||
|
||||
context.configure(
|
||||
connection=connection,
|
||||
target_metadata=target_metadata
|
||||
)
|
||||
# Configure context with version_table_schema if using a specific schema
|
||||
context_opts = {
|
||||
"connection": connection,
|
||||
"target_metadata": target_metadata,
|
||||
}
|
||||
if target_schema:
|
||||
context_opts["version_table_schema"] = target_schema
|
||||
|
||||
context.configure(**context_opts)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
@@ -5,120 +5,150 @@ Revises:
|
||||
Create Date: 2025-11-27 11:54:19.228030
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
from alembic import op
|
||||
from pgvector.sqlalchemy import Vector
|
||||
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = '5a366d414dce'
|
||||
down_revision: Union[str, Sequence[str], None] = None
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
revision: str = "5a366d414dce"
|
||||
down_revision: str | Sequence[str] | None = None
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Upgrade schema - create all tables from scratch."""
|
||||
|
||||
# Enable required extensions
|
||||
op.execute('CREATE EXTENSION IF NOT EXISTS vector')
|
||||
op.execute("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
|
||||
# Create banks table
|
||||
op.create_table(
|
||||
'banks',
|
||||
sa.Column('bank_id', sa.Text(), nullable=False),
|
||||
sa.Column('name', sa.Text(), nullable=True),
|
||||
sa.Column('personality', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
|
||||
sa.Column('background', sa.Text(), nullable=True),
|
||||
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.PrimaryKeyConstraint('bank_id', name=op.f('pk_banks'))
|
||||
"banks",
|
||||
sa.Column("bank_id", sa.Text(), nullable=False),
|
||||
sa.Column("name", sa.Text(), nullable=True),
|
||||
sa.Column(
|
||||
"personality",
|
||||
postgresql.JSONB(astext_type=sa.Text()),
|
||||
server_default=sa.text("'{}'::jsonb"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("background", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.PrimaryKeyConstraint("bank_id", name=op.f("pk_banks")),
|
||||
)
|
||||
|
||||
# Create documents table
|
||||
op.create_table(
|
||||
'documents',
|
||||
sa.Column('id', sa.Text(), nullable=False),
|
||||
sa.Column('bank_id', sa.Text(), nullable=False),
|
||||
sa.Column('original_text', sa.Text(), nullable=True),
|
||||
sa.Column('content_hash', sa.Text(), nullable=True),
|
||||
sa.Column('metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
|
||||
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.PrimaryKeyConstraint('id', 'bank_id', name=op.f('pk_documents'))
|
||||
"documents",
|
||||
sa.Column("id", sa.Text(), nullable=False),
|
||||
sa.Column("bank_id", sa.Text(), nullable=False),
|
||||
sa.Column("original_text", sa.Text(), nullable=True),
|
||||
sa.Column("content_hash", sa.Text(), nullable=True),
|
||||
sa.Column(
|
||||
"metadata", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False
|
||||
),
|
||||
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id", "bank_id", name=op.f("pk_documents")),
|
||||
)
|
||||
op.create_index('idx_documents_bank_id', 'documents', ['bank_id'])
|
||||
op.create_index('idx_documents_content_hash', 'documents', ['content_hash'])
|
||||
op.create_index("idx_documents_bank_id", "documents", ["bank_id"])
|
||||
op.create_index("idx_documents_content_hash", "documents", ["content_hash"])
|
||||
|
||||
# Create async_operations table
|
||||
op.create_table(
|
||||
'async_operations',
|
||||
sa.Column('operation_id', postgresql.UUID(as_uuid=True), server_default=sa.text('gen_random_uuid()'), nullable=False),
|
||||
sa.Column('bank_id', sa.Text(), nullable=False),
|
||||
sa.Column('operation_type', sa.Text(), nullable=False),
|
||||
sa.Column('status', sa.Text(), server_default='pending', nullable=False),
|
||||
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.Column('completed_at', postgresql.TIMESTAMP(timezone=True), nullable=True),
|
||||
sa.Column('error_message', sa.Text(), nullable=True),
|
||||
sa.Column('result_metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
|
||||
sa.PrimaryKeyConstraint('operation_id', name=op.f('pk_async_operations')),
|
||||
sa.CheckConstraint("status IN ('pending', 'processing', 'completed', 'failed')", name='async_operations_status_check')
|
||||
"async_operations",
|
||||
sa.Column(
|
||||
"operation_id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False
|
||||
),
|
||||
sa.Column("bank_id", sa.Text(), nullable=False),
|
||||
sa.Column("operation_type", sa.Text(), nullable=False),
|
||||
sa.Column("status", sa.Text(), server_default="pending", nullable=False),
|
||||
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("completed_at", postgresql.TIMESTAMP(timezone=True), nullable=True),
|
||||
sa.Column("error_message", sa.Text(), nullable=True),
|
||||
sa.Column(
|
||||
"result_metadata",
|
||||
postgresql.JSONB(astext_type=sa.Text()),
|
||||
server_default=sa.text("'{}'::jsonb"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.PrimaryKeyConstraint("operation_id", name=op.f("pk_async_operations")),
|
||||
sa.CheckConstraint(
|
||||
"status IN ('pending', 'processing', 'completed', 'failed')", name="async_operations_status_check"
|
||||
),
|
||||
)
|
||||
op.create_index('idx_async_operations_bank_id', 'async_operations', ['bank_id'])
|
||||
op.create_index('idx_async_operations_status', 'async_operations', ['status'])
|
||||
op.create_index('idx_async_operations_bank_status', 'async_operations', ['bank_id', 'status'])
|
||||
op.create_index("idx_async_operations_bank_id", "async_operations", ["bank_id"])
|
||||
op.create_index("idx_async_operations_status", "async_operations", ["status"])
|
||||
op.create_index("idx_async_operations_bank_status", "async_operations", ["bank_id", "status"])
|
||||
|
||||
# Create entities table
|
||||
op.create_table(
|
||||
'entities',
|
||||
sa.Column('id', postgresql.UUID(as_uuid=True), server_default=sa.text('gen_random_uuid()'), nullable=False),
|
||||
sa.Column('canonical_name', sa.Text(), nullable=False),
|
||||
sa.Column('bank_id', sa.Text(), nullable=False),
|
||||
sa.Column('metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
|
||||
sa.Column('first_seen', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.Column('last_seen', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.Column('mention_count', sa.Integer(), server_default='1', nullable=False),
|
||||
sa.PrimaryKeyConstraint('id', name=op.f('pk_entities'))
|
||||
"entities",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
|
||||
sa.Column("canonical_name", sa.Text(), nullable=False),
|
||||
sa.Column("bank_id", sa.Text(), nullable=False),
|
||||
sa.Column(
|
||||
"metadata", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False
|
||||
),
|
||||
sa.Column("first_seen", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("last_seen", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("mention_count", sa.Integer(), server_default="1", nullable=False),
|
||||
sa.PrimaryKeyConstraint("id", name=op.f("pk_entities")),
|
||||
)
|
||||
op.create_index('idx_entities_bank_id', 'entities', ['bank_id'])
|
||||
op.create_index('idx_entities_canonical_name', 'entities', ['canonical_name'])
|
||||
op.create_index('idx_entities_bank_name', 'entities', ['bank_id', 'canonical_name'])
|
||||
op.create_index("idx_entities_bank_id", "entities", ["bank_id"])
|
||||
op.create_index("idx_entities_canonical_name", "entities", ["canonical_name"])
|
||||
op.create_index("idx_entities_bank_name", "entities", ["bank_id", "canonical_name"])
|
||||
# Create unique index on (bank_id, LOWER(canonical_name)) for entity resolution
|
||||
op.execute('CREATE UNIQUE INDEX idx_entities_bank_lower_name ON entities (bank_id, LOWER(canonical_name))')
|
||||
op.execute("CREATE UNIQUE INDEX idx_entities_bank_lower_name ON entities (bank_id, LOWER(canonical_name))")
|
||||
|
||||
# Create memory_units table
|
||||
op.create_table(
|
||||
'memory_units',
|
||||
sa.Column('id', postgresql.UUID(as_uuid=True), server_default=sa.text('gen_random_uuid()'), nullable=False),
|
||||
sa.Column('bank_id', sa.Text(), nullable=False),
|
||||
sa.Column('document_id', sa.Text(), nullable=True),
|
||||
sa.Column('text', sa.Text(), nullable=False),
|
||||
sa.Column('embedding', Vector(384), nullable=True),
|
||||
sa.Column('context', sa.Text(), nullable=True),
|
||||
sa.Column('event_date', postgresql.TIMESTAMP(timezone=True), nullable=False),
|
||||
sa.Column('occurred_start', postgresql.TIMESTAMP(timezone=True), nullable=True),
|
||||
sa.Column('occurred_end', postgresql.TIMESTAMP(timezone=True), nullable=True),
|
||||
sa.Column('mentioned_at', postgresql.TIMESTAMP(timezone=True), nullable=True),
|
||||
sa.Column('fact_type', sa.Text(), server_default='world', nullable=False),
|
||||
sa.Column('confidence_score', sa.Float(), nullable=True),
|
||||
sa.Column('access_count', sa.Integer(), server_default='0', nullable=False),
|
||||
sa.Column('metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
|
||||
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.ForeignKeyConstraint(['document_id', 'bank_id'], ['documents.id', 'documents.bank_id'], name='memory_units_document_fkey', ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id', name=op.f('pk_memory_units')),
|
||||
sa.CheckConstraint("fact_type IN ('world', 'bank', 'opinion', 'observation')", name='memory_units_fact_type_check'),
|
||||
sa.CheckConstraint("confidence_score IS NULL OR (confidence_score >= 0.0 AND confidence_score <= 1.0)", name='memory_units_confidence_range_check'),
|
||||
"memory_units",
|
||||
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
|
||||
sa.Column("bank_id", sa.Text(), nullable=False),
|
||||
sa.Column("document_id", sa.Text(), nullable=True),
|
||||
sa.Column("text", sa.Text(), nullable=False),
|
||||
sa.Column("embedding", Vector(384), nullable=True),
|
||||
sa.Column("context", sa.Text(), nullable=True),
|
||||
sa.Column("event_date", postgresql.TIMESTAMP(timezone=True), nullable=False),
|
||||
sa.Column("occurred_start", postgresql.TIMESTAMP(timezone=True), nullable=True),
|
||||
sa.Column("occurred_end", postgresql.TIMESTAMP(timezone=True), nullable=True),
|
||||
sa.Column("mentioned_at", postgresql.TIMESTAMP(timezone=True), nullable=True),
|
||||
sa.Column("fact_type", sa.Text(), server_default="world", nullable=False),
|
||||
sa.Column("confidence_score", sa.Float(), nullable=True),
|
||||
sa.Column("access_count", sa.Integer(), server_default="0", nullable=False),
|
||||
sa.Column(
|
||||
"metadata", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False
|
||||
),
|
||||
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.ForeignKeyConstraint(
|
||||
["document_id", "bank_id"],
|
||||
["documents.id", "documents.bank_id"],
|
||||
name="memory_units_document_fkey",
|
||||
ondelete="CASCADE",
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id", name=op.f("pk_memory_units")),
|
||||
sa.CheckConstraint(
|
||||
"fact_type IN ('world', 'bank', 'opinion', 'observation')", name="memory_units_fact_type_check"
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"confidence_score IS NULL OR (confidence_score >= 0.0 AND confidence_score <= 1.0)",
|
||||
name="memory_units_confidence_range_check",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"(fact_type = 'opinion' AND confidence_score IS NOT NULL) OR "
|
||||
"(fact_type = 'observation') OR "
|
||||
"(fact_type NOT IN ('opinion', 'observation') AND confidence_score IS NULL)",
|
||||
name='confidence_score_fact_type_check'
|
||||
)
|
||||
name="confidence_score_fact_type_check",
|
||||
),
|
||||
)
|
||||
|
||||
# Add search_vector column for full-text search
|
||||
@@ -128,18 +158,41 @@ def upgrade() -> None:
|
||||
GENERATED ALWAYS AS (to_tsvector('english', COALESCE(text, '') || ' ' || COALESCE(context, ''))) STORED
|
||||
""")
|
||||
|
||||
op.create_index('idx_memory_units_bank_id', 'memory_units', ['bank_id'])
|
||||
op.create_index('idx_memory_units_document_id', 'memory_units', ['document_id'])
|
||||
op.create_index('idx_memory_units_event_date', 'memory_units', [sa.text('event_date DESC')])
|
||||
op.create_index('idx_memory_units_bank_date', 'memory_units', ['bank_id', sa.text('event_date DESC')])
|
||||
op.create_index('idx_memory_units_access_count', 'memory_units', [sa.text('access_count DESC')])
|
||||
op.create_index('idx_memory_units_fact_type', 'memory_units', ['fact_type'])
|
||||
op.create_index('idx_memory_units_bank_fact_type', 'memory_units', ['bank_id', 'fact_type'])
|
||||
op.create_index('idx_memory_units_bank_type_date', 'memory_units', ['bank_id', 'fact_type', sa.text('event_date DESC')])
|
||||
op.create_index('idx_memory_units_opinion_confidence', 'memory_units', ['bank_id', sa.text('confidence_score DESC')], postgresql_where=sa.text("fact_type = 'opinion'"))
|
||||
op.create_index('idx_memory_units_opinion_date', 'memory_units', ['bank_id', sa.text('event_date DESC')], postgresql_where=sa.text("fact_type = 'opinion'"))
|
||||
op.create_index('idx_memory_units_observation_date', 'memory_units', ['bank_id', sa.text('event_date DESC')], postgresql_where=sa.text("fact_type = 'observation'"))
|
||||
op.create_index('idx_memory_units_embedding', 'memory_units', ['embedding'], postgresql_using='hnsw', postgresql_ops={'embedding': 'vector_cosine_ops'})
|
||||
op.create_index("idx_memory_units_bank_id", "memory_units", ["bank_id"])
|
||||
op.create_index("idx_memory_units_document_id", "memory_units", ["document_id"])
|
||||
op.create_index("idx_memory_units_event_date", "memory_units", [sa.text("event_date DESC")])
|
||||
op.create_index("idx_memory_units_bank_date", "memory_units", ["bank_id", sa.text("event_date DESC")])
|
||||
op.create_index("idx_memory_units_access_count", "memory_units", [sa.text("access_count DESC")])
|
||||
op.create_index("idx_memory_units_fact_type", "memory_units", ["fact_type"])
|
||||
op.create_index("idx_memory_units_bank_fact_type", "memory_units", ["bank_id", "fact_type"])
|
||||
op.create_index(
|
||||
"idx_memory_units_bank_type_date", "memory_units", ["bank_id", "fact_type", sa.text("event_date DESC")]
|
||||
)
|
||||
op.create_index(
|
||||
"idx_memory_units_opinion_confidence",
|
||||
"memory_units",
|
||||
["bank_id", sa.text("confidence_score DESC")],
|
||||
postgresql_where=sa.text("fact_type = 'opinion'"),
|
||||
)
|
||||
op.create_index(
|
||||
"idx_memory_units_opinion_date",
|
||||
"memory_units",
|
||||
["bank_id", sa.text("event_date DESC")],
|
||||
postgresql_where=sa.text("fact_type = 'opinion'"),
|
||||
)
|
||||
op.create_index(
|
||||
"idx_memory_units_observation_date",
|
||||
"memory_units",
|
||||
["bank_id", sa.text("event_date DESC")],
|
||||
postgresql_where=sa.text("fact_type = 'observation'"),
|
||||
)
|
||||
op.create_index(
|
||||
"idx_memory_units_embedding",
|
||||
"memory_units",
|
||||
["embedding"],
|
||||
postgresql_using="hnsw",
|
||||
postgresql_ops={"embedding": "vector_cosine_ops"},
|
||||
)
|
||||
|
||||
# Create BM25 full-text search index on search_vector
|
||||
op.execute("""
|
||||
@@ -158,116 +211,149 @@ def upgrade() -> None:
|
||||
FROM memory_units
|
||||
""")
|
||||
|
||||
op.create_index('idx_memory_units_bm25_bank', 'memory_units_bm25', ['bank_id'])
|
||||
op.create_index('idx_memory_units_bm25_text_vector', 'memory_units_bm25', ['text_vector'], postgresql_using='gin')
|
||||
op.create_index("idx_memory_units_bm25_bank", "memory_units_bm25", ["bank_id"])
|
||||
op.create_index("idx_memory_units_bm25_text_vector", "memory_units_bm25", ["text_vector"], postgresql_using="gin")
|
||||
|
||||
# Create entity_cooccurrences table
|
||||
op.create_table(
|
||||
'entity_cooccurrences',
|
||||
sa.Column('entity_id_1', postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column('entity_id_2', postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column('cooccurrence_count', sa.Integer(), server_default='1', nullable=False),
|
||||
sa.Column('last_cooccurred', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.ForeignKeyConstraint(['entity_id_1'], ['entities.id'], name=op.f('fk_entity_cooccurrences_entity_id_1_entities'), ondelete='CASCADE'),
|
||||
sa.ForeignKeyConstraint(['entity_id_2'], ['entities.id'], name=op.f('fk_entity_cooccurrences_entity_id_2_entities'), ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('entity_id_1', 'entity_id_2', name=op.f('pk_entity_cooccurrences')),
|
||||
sa.CheckConstraint('entity_id_1 < entity_id_2', name='entity_cooccurrence_order_check')
|
||||
"entity_cooccurrences",
|
||||
sa.Column("entity_id_1", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("entity_id_2", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("cooccurrence_count", sa.Integer(), server_default="1", nullable=False),
|
||||
sa.Column(
|
||||
"last_cooccurred", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False
|
||||
),
|
||||
sa.ForeignKeyConstraint(
|
||||
["entity_id_1"],
|
||||
["entities.id"],
|
||||
name=op.f("fk_entity_cooccurrences_entity_id_1_entities"),
|
||||
ondelete="CASCADE",
|
||||
),
|
||||
sa.ForeignKeyConstraint(
|
||||
["entity_id_2"],
|
||||
["entities.id"],
|
||||
name=op.f("fk_entity_cooccurrences_entity_id_2_entities"),
|
||||
ondelete="CASCADE",
|
||||
),
|
||||
sa.PrimaryKeyConstraint("entity_id_1", "entity_id_2", name=op.f("pk_entity_cooccurrences")),
|
||||
sa.CheckConstraint("entity_id_1 < entity_id_2", name="entity_cooccurrence_order_check"),
|
||||
)
|
||||
op.create_index('idx_entity_cooccurrences_entity1', 'entity_cooccurrences', ['entity_id_1'])
|
||||
op.create_index('idx_entity_cooccurrences_entity2', 'entity_cooccurrences', ['entity_id_2'])
|
||||
op.create_index('idx_entity_cooccurrences_count', 'entity_cooccurrences', [sa.text('cooccurrence_count DESC')])
|
||||
op.create_index("idx_entity_cooccurrences_entity1", "entity_cooccurrences", ["entity_id_1"])
|
||||
op.create_index("idx_entity_cooccurrences_entity2", "entity_cooccurrences", ["entity_id_2"])
|
||||
op.create_index("idx_entity_cooccurrences_count", "entity_cooccurrences", [sa.text("cooccurrence_count DESC")])
|
||||
|
||||
# Create memory_links table
|
||||
op.create_table(
|
||||
'memory_links',
|
||||
sa.Column('from_unit_id', postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column('to_unit_id', postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column('link_type', sa.Text(), nullable=False),
|
||||
sa.Column('entity_id', postgresql.UUID(as_uuid=True), nullable=True),
|
||||
sa.Column('weight', sa.Float(), server_default='1.0', nullable=False),
|
||||
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.ForeignKeyConstraint(['entity_id'], ['entities.id'], name=op.f('fk_memory_links_entity_id_entities'), ondelete='CASCADE'),
|
||||
sa.ForeignKeyConstraint(['from_unit_id'], ['memory_units.id'], name=op.f('fk_memory_links_from_unit_id_memory_units'), ondelete='CASCADE'),
|
||||
sa.ForeignKeyConstraint(['to_unit_id'], ['memory_units.id'], name=op.f('fk_memory_links_to_unit_id_memory_units'), ondelete='CASCADE'),
|
||||
sa.CheckConstraint("link_type IN ('temporal', 'semantic', 'entity', 'causes', 'caused_by', 'enables', 'prevents')", name='memory_links_link_type_check'),
|
||||
sa.CheckConstraint('weight >= 0.0 AND weight <= 1.0', name='memory_links_weight_check')
|
||||
"memory_links",
|
||||
sa.Column("from_unit_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("to_unit_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("link_type", sa.Text(), nullable=False),
|
||||
sa.Column("entity_id", postgresql.UUID(as_uuid=True), nullable=True),
|
||||
sa.Column("weight", sa.Float(), server_default="1.0", nullable=False),
|
||||
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.ForeignKeyConstraint(
|
||||
["entity_id"], ["entities.id"], name=op.f("fk_memory_links_entity_id_entities"), ondelete="CASCADE"
|
||||
),
|
||||
sa.ForeignKeyConstraint(
|
||||
["from_unit_id"],
|
||||
["memory_units.id"],
|
||||
name=op.f("fk_memory_links_from_unit_id_memory_units"),
|
||||
ondelete="CASCADE",
|
||||
),
|
||||
sa.ForeignKeyConstraint(
|
||||
["to_unit_id"],
|
||||
["memory_units.id"],
|
||||
name=op.f("fk_memory_links_to_unit_id_memory_units"),
|
||||
ondelete="CASCADE",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"link_type IN ('temporal', 'semantic', 'entity', 'causes', 'caused_by', 'enables', 'prevents')",
|
||||
name="memory_links_link_type_check",
|
||||
),
|
||||
sa.CheckConstraint("weight >= 0.0 AND weight <= 1.0", name="memory_links_weight_check"),
|
||||
)
|
||||
# Create unique constraint using COALESCE for nullable entity_id
|
||||
op.execute("CREATE UNIQUE INDEX idx_memory_links_unique ON memory_links (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid))")
|
||||
op.create_index('idx_memory_links_from_unit', 'memory_links', ['from_unit_id'])
|
||||
op.create_index('idx_memory_links_to_unit', 'memory_links', ['to_unit_id'])
|
||||
op.create_index('idx_memory_links_entity', 'memory_links', ['entity_id'])
|
||||
op.create_index('idx_memory_links_link_type', 'memory_links', ['link_type'])
|
||||
op.execute(
|
||||
"CREATE UNIQUE INDEX idx_memory_links_unique ON memory_links (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid))"
|
||||
)
|
||||
op.create_index("idx_memory_links_from_unit", "memory_links", ["from_unit_id"])
|
||||
op.create_index("idx_memory_links_to_unit", "memory_links", ["to_unit_id"])
|
||||
op.create_index("idx_memory_links_entity", "memory_links", ["entity_id"])
|
||||
op.create_index("idx_memory_links_link_type", "memory_links", ["link_type"])
|
||||
|
||||
# Create unit_entities table
|
||||
op.create_table(
|
||||
'unit_entities',
|
||||
sa.Column('unit_id', postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column('entity_id', postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['entity_id'], ['entities.id'], name=op.f('fk_unit_entities_entity_id_entities'), ondelete='CASCADE'),
|
||||
sa.ForeignKeyConstraint(['unit_id'], ['memory_units.id'], name=op.f('fk_unit_entities_unit_id_memory_units'), ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('unit_id', 'entity_id', name=op.f('pk_unit_entities'))
|
||||
"unit_entities",
|
||||
sa.Column("unit_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.Column("entity_id", postgresql.UUID(as_uuid=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(
|
||||
["entity_id"], ["entities.id"], name=op.f("fk_unit_entities_entity_id_entities"), ondelete="CASCADE"
|
||||
),
|
||||
sa.ForeignKeyConstraint(
|
||||
["unit_id"], ["memory_units.id"], name=op.f("fk_unit_entities_unit_id_memory_units"), ondelete="CASCADE"
|
||||
),
|
||||
sa.PrimaryKeyConstraint("unit_id", "entity_id", name=op.f("pk_unit_entities")),
|
||||
)
|
||||
op.create_index('idx_unit_entities_unit', 'unit_entities', ['unit_id'])
|
||||
op.create_index('idx_unit_entities_entity', 'unit_entities', ['entity_id'])
|
||||
op.create_index("idx_unit_entities_unit", "unit_entities", ["unit_id"])
|
||||
op.create_index("idx_unit_entities_entity", "unit_entities", ["entity_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Downgrade schema - drop all tables."""
|
||||
|
||||
# Drop tables in reverse dependency order
|
||||
op.drop_index('idx_unit_entities_entity', table_name='unit_entities')
|
||||
op.drop_index('idx_unit_entities_unit', table_name='unit_entities')
|
||||
op.drop_table('unit_entities')
|
||||
op.drop_index("idx_unit_entities_entity", table_name="unit_entities")
|
||||
op.drop_index("idx_unit_entities_unit", table_name="unit_entities")
|
||||
op.drop_table("unit_entities")
|
||||
|
||||
op.drop_index('idx_memory_links_link_type', table_name='memory_links')
|
||||
op.drop_index('idx_memory_links_entity', table_name='memory_links')
|
||||
op.drop_index('idx_memory_links_to_unit', table_name='memory_links')
|
||||
op.drop_index('idx_memory_links_from_unit', table_name='memory_links')
|
||||
op.execute('DROP INDEX IF EXISTS idx_memory_links_unique')
|
||||
op.drop_table('memory_links')
|
||||
op.drop_index("idx_memory_links_link_type", table_name="memory_links")
|
||||
op.drop_index("idx_memory_links_entity", table_name="memory_links")
|
||||
op.drop_index("idx_memory_links_to_unit", table_name="memory_links")
|
||||
op.drop_index("idx_memory_links_from_unit", table_name="memory_links")
|
||||
op.execute("DROP INDEX IF EXISTS idx_memory_links_unique")
|
||||
op.drop_table("memory_links")
|
||||
|
||||
op.drop_index('idx_entity_cooccurrences_count', table_name='entity_cooccurrences')
|
||||
op.drop_index('idx_entity_cooccurrences_entity2', table_name='entity_cooccurrences')
|
||||
op.drop_index('idx_entity_cooccurrences_entity1', table_name='entity_cooccurrences')
|
||||
op.drop_table('entity_cooccurrences')
|
||||
op.drop_index("idx_entity_cooccurrences_count", table_name="entity_cooccurrences")
|
||||
op.drop_index("idx_entity_cooccurrences_entity2", table_name="entity_cooccurrences")
|
||||
op.drop_index("idx_entity_cooccurrences_entity1", table_name="entity_cooccurrences")
|
||||
op.drop_table("entity_cooccurrences")
|
||||
|
||||
# Drop BM25 materialized view and index
|
||||
op.drop_index('idx_memory_units_bm25_text_vector', table_name='memory_units_bm25')
|
||||
op.drop_index('idx_memory_units_bm25_bank', table_name='memory_units_bm25')
|
||||
op.execute('DROP MATERIALIZED VIEW IF EXISTS memory_units_bm25')
|
||||
op.drop_index("idx_memory_units_bm25_text_vector", table_name="memory_units_bm25")
|
||||
op.drop_index("idx_memory_units_bm25_bank", table_name="memory_units_bm25")
|
||||
op.execute("DROP MATERIALIZED VIEW IF EXISTS memory_units_bm25")
|
||||
|
||||
op.drop_index('idx_memory_units_embedding', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_observation_date', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_opinion_date', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_opinion_confidence', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_bank_type_date', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_bank_fact_type', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_fact_type', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_access_count', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_bank_date', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_event_date', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_document_id', table_name='memory_units')
|
||||
op.drop_index('idx_memory_units_bank_id', table_name='memory_units')
|
||||
op.execute('DROP INDEX IF EXISTS idx_memory_units_text_search')
|
||||
op.drop_table('memory_units')
|
||||
op.drop_index("idx_memory_units_embedding", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_observation_date", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_opinion_date", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_opinion_confidence", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_bank_type_date", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_bank_fact_type", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_fact_type", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_access_count", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_bank_date", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_event_date", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_document_id", table_name="memory_units")
|
||||
op.drop_index("idx_memory_units_bank_id", table_name="memory_units")
|
||||
op.execute("DROP INDEX IF EXISTS idx_memory_units_text_search")
|
||||
op.drop_table("memory_units")
|
||||
|
||||
op.execute('DROP INDEX IF EXISTS idx_entities_bank_lower_name')
|
||||
op.drop_index('idx_entities_bank_name', table_name='entities')
|
||||
op.drop_index('idx_entities_canonical_name', table_name='entities')
|
||||
op.drop_index('idx_entities_bank_id', table_name='entities')
|
||||
op.drop_table('entities')
|
||||
op.execute("DROP INDEX IF EXISTS idx_entities_bank_lower_name")
|
||||
op.drop_index("idx_entities_bank_name", table_name="entities")
|
||||
op.drop_index("idx_entities_canonical_name", table_name="entities")
|
||||
op.drop_index("idx_entities_bank_id", table_name="entities")
|
||||
op.drop_table("entities")
|
||||
|
||||
op.drop_index('idx_async_operations_bank_status', table_name='async_operations')
|
||||
op.drop_index('idx_async_operations_status', table_name='async_operations')
|
||||
op.drop_index('idx_async_operations_bank_id', table_name='async_operations')
|
||||
op.drop_table('async_operations')
|
||||
op.drop_index("idx_async_operations_bank_status", table_name="async_operations")
|
||||
op.drop_index("idx_async_operations_status", table_name="async_operations")
|
||||
op.drop_index("idx_async_operations_bank_id", table_name="async_operations")
|
||||
op.drop_table("async_operations")
|
||||
|
||||
op.drop_index('idx_documents_content_hash', table_name='documents')
|
||||
op.drop_index('idx_documents_bank_id', table_name='documents')
|
||||
op.drop_table('documents')
|
||||
op.drop_index("idx_documents_content_hash", table_name="documents")
|
||||
op.drop_index("idx_documents_bank_id", table_name="documents")
|
||||
op.drop_table("documents")
|
||||
|
||||
op.drop_table('banks')
|
||||
op.drop_table("banks")
|
||||
|
||||
# Drop extensions (optional - comment out if you want to keep them)
|
||||
# op.execute('DROP EXTENSION IF EXISTS vector')
|
||||
|
||||
@@ -5,18 +5,18 @@ Revises: 5a366d414dce
|
||||
Create Date: 2025-11-28 00:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = 'b7c4d8e9f1a2'
|
||||
down_revision: Union[str, Sequence[str], None] = '5a366d414dce'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
revision: str = "b7c4d8e9f1a2"
|
||||
down_revision: str | Sequence[str] | None = "5a366d414dce"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
@@ -24,47 +24,47 @@ def upgrade() -> None:
|
||||
|
||||
# Create chunks table with single text PK (bank_id_document_id_chunk_index)
|
||||
op.create_table(
|
||||
'chunks',
|
||||
sa.Column('chunk_id', sa.Text(), nullable=False),
|
||||
sa.Column('document_id', sa.Text(), nullable=False),
|
||||
sa.Column('bank_id', sa.Text(), nullable=False),
|
||||
sa.Column('chunk_index', sa.Integer(), nullable=False),
|
||||
sa.Column('chunk_text', sa.Text(), nullable=False),
|
||||
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
|
||||
sa.ForeignKeyConstraint(['document_id', 'bank_id'], ['documents.id', 'documents.bank_id'], name='chunks_document_fkey', ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('chunk_id', name=op.f('pk_chunks'))
|
||||
"chunks",
|
||||
sa.Column("chunk_id", sa.Text(), nullable=False),
|
||||
sa.Column("document_id", sa.Text(), nullable=False),
|
||||
sa.Column("bank_id", sa.Text(), nullable=False),
|
||||
sa.Column("chunk_index", sa.Integer(), nullable=False),
|
||||
sa.Column("chunk_text", sa.Text(), nullable=False),
|
||||
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
|
||||
sa.ForeignKeyConstraint(
|
||||
["document_id", "bank_id"],
|
||||
["documents.id", "documents.bank_id"],
|
||||
name="chunks_document_fkey",
|
||||
ondelete="CASCADE",
|
||||
),
|
||||
sa.PrimaryKeyConstraint("chunk_id", name=op.f("pk_chunks")),
|
||||
)
|
||||
|
||||
# Add indexes for efficient queries
|
||||
op.create_index('idx_chunks_document_id', 'chunks', ['document_id'])
|
||||
op.create_index('idx_chunks_bank_id', 'chunks', ['bank_id'])
|
||||
op.create_index("idx_chunks_document_id", "chunks", ["document_id"])
|
||||
op.create_index("idx_chunks_bank_id", "chunks", ["bank_id"])
|
||||
|
||||
# Add chunk_id column to memory_units (nullable, as existing records won't have chunks)
|
||||
op.add_column('memory_units', sa.Column('chunk_id', sa.Text(), nullable=True))
|
||||
op.add_column("memory_units", sa.Column("chunk_id", sa.Text(), nullable=True))
|
||||
|
||||
# Add foreign key constraint to chunks table
|
||||
op.create_foreign_key(
|
||||
'memory_units_chunk_fkey',
|
||||
'memory_units',
|
||||
'chunks',
|
||||
['chunk_id'],
|
||||
['chunk_id'],
|
||||
ondelete='SET NULL'
|
||||
"memory_units_chunk_fkey", "memory_units", "chunks", ["chunk_id"], ["chunk_id"], ondelete="SET NULL"
|
||||
)
|
||||
|
||||
# Add index on chunk_id for efficient lookups
|
||||
op.create_index('idx_memory_units_chunk_id', 'memory_units', ['chunk_id'])
|
||||
op.create_index("idx_memory_units_chunk_id", "memory_units", ["chunk_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove chunks table and chunk_id from memory_units."""
|
||||
|
||||
# Drop index and foreign key from memory_units
|
||||
op.drop_index('idx_memory_units_chunk_id', table_name='memory_units')
|
||||
op.drop_constraint('memory_units_chunk_fkey', 'memory_units', type_='foreignkey')
|
||||
op.drop_column('memory_units', 'chunk_id')
|
||||
op.drop_index("idx_memory_units_chunk_id", table_name="memory_units")
|
||||
op.drop_constraint("memory_units_chunk_fkey", "memory_units", type_="foreignkey")
|
||||
op.drop_column("memory_units", "chunk_id")
|
||||
|
||||
# Drop chunks table indexes and table
|
||||
op.drop_index('idx_chunks_bank_id', table_name='chunks')
|
||||
op.drop_index('idx_chunks_document_id', table_name='chunks')
|
||||
op.drop_table('chunks')
|
||||
op.drop_index("idx_chunks_bank_id", table_name="chunks")
|
||||
op.drop_index("idx_chunks_document_id", table_name="chunks")
|
||||
op.drop_table("chunks")
|
||||
|
||||
+11
-11
@@ -5,35 +5,35 @@ Revises: b7c4d8e9f1a2
|
||||
Create Date: 2025-12-02 00:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = 'c8e5f2a3b4d1'
|
||||
down_revision: Union[str, Sequence[str], None] = 'b7c4d8e9f1a2'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
revision: str = "c8e5f2a3b4d1"
|
||||
down_revision: str | Sequence[str] | None = "b7c4d8e9f1a2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add retain_params JSONB column to documents table."""
|
||||
|
||||
# Add retain_params column to store parameters passed during retain
|
||||
op.add_column('documents', sa.Column('retain_params', postgresql.JSONB(), nullable=True))
|
||||
op.add_column("documents", sa.Column("retain_params", postgresql.JSONB(), nullable=True))
|
||||
|
||||
# Add index for efficient queries on retain_params
|
||||
op.create_index('idx_documents_retain_params', 'documents', ['retain_params'], postgresql_using='gin')
|
||||
op.create_index("idx_documents_retain_params", "documents", ["retain_params"], postgresql_using="gin")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove retain_params column from documents table."""
|
||||
|
||||
# Drop index
|
||||
op.drop_index('idx_documents_retain_params', table_name='documents')
|
||||
op.drop_index("idx_documents_retain_params", table_name="documents")
|
||||
|
||||
# Drop column
|
||||
op.drop_column('documents', 'retain_params')
|
||||
op.drop_column("documents", "retain_params")
|
||||
|
||||
+20
-15
@@ -5,44 +5,49 @@ Revises: c8e5f2a3b4d1
|
||||
Create Date: 2024-12-04 15:00:00.000000
|
||||
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'd9f6a3b4c5e2'
|
||||
down_revision = 'c8e5f2a3b4d1'
|
||||
revision = "d9f6a3b4c5e2"
|
||||
down_revision = "c8e5f2a3b4d1"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (e.g., 'tenant_x.' or '' for public)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade():
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop old check constraint FIRST (before updating data)
|
||||
op.drop_constraint('memory_units_fact_type_check', 'memory_units', type_='check')
|
||||
op.drop_constraint("memory_units_fact_type_check", "memory_units", type_="check")
|
||||
|
||||
# Update existing 'bank' values to 'experience'
|
||||
op.execute("UPDATE memory_units SET fact_type = 'experience' WHERE fact_type = 'bank'")
|
||||
op.execute(f"UPDATE {schema}memory_units SET fact_type = 'experience' WHERE fact_type = 'bank'")
|
||||
# Also update any 'interactions' values (in case of partial migration)
|
||||
op.execute("UPDATE memory_units SET fact_type = 'experience' WHERE fact_type = 'interactions'")
|
||||
op.execute(f"UPDATE {schema}memory_units SET fact_type = 'experience' WHERE fact_type = 'interactions'")
|
||||
|
||||
# Create new check constraint with 'experience' instead of 'bank'
|
||||
op.create_check_constraint(
|
||||
'memory_units_fact_type_check',
|
||||
'memory_units',
|
||||
"fact_type IN ('world', 'experience', 'opinion', 'observation')"
|
||||
"memory_units_fact_type_check", "memory_units", "fact_type IN ('world', 'experience', 'opinion', 'observation')"
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop new check constraint FIRST
|
||||
op.drop_constraint('memory_units_fact_type_check', 'memory_units', type_='check')
|
||||
op.drop_constraint("memory_units_fact_type_check", "memory_units", type_="check")
|
||||
|
||||
# Update 'experience' back to 'bank'
|
||||
op.execute("UPDATE memory_units SET fact_type = 'bank' WHERE fact_type = 'experience'")
|
||||
op.execute(f"UPDATE {schema}memory_units SET fact_type = 'bank' WHERE fact_type = 'experience'")
|
||||
|
||||
# Recreate old check constraint
|
||||
op.create_check_constraint(
|
||||
'memory_units_fact_type_check',
|
||||
'memory_units',
|
||||
"fact_type IN ('world', 'bank', 'opinion', 'observation')"
|
||||
"memory_units_fact_type_check", "memory_units", "fact_type IN ('world', 'bank', 'opinion', 'observation')"
|
||||
)
|
||||
|
||||
+72
-23
@@ -8,22 +8,49 @@ Migrate disposition traits from Big Five (openness, conscientiousness, extravers
|
||||
agreeableness, neuroticism, bias_strength with 0-1 float values) to the new 3-trait
|
||||
system (skepticism, literalism, empathy with 1-5 integer values).
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = 'e0a1b2c3d4e5'
|
||||
down_revision: Union[str, Sequence[str], None] = 'rename_personality'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
revision: str = "e0a1b2c3d4e5"
|
||||
down_revision: str | Sequence[str] | None = "rename_personality"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (e.g., 'tenant_x.' or '' for public)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def _get_target_schema() -> str:
|
||||
"""Get the target schema name (tenant schema or 'public')."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return schema if schema else "public"
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Convert Big Five disposition to 3-trait disposition."""
|
||||
conn = op.get_bind()
|
||||
schema = _get_schema_prefix()
|
||||
target_schema = _get_target_schema()
|
||||
|
||||
# Check if disposition column exists (should have been created by previous migration)
|
||||
result = conn.execute(
|
||||
sa.text("""
|
||||
SELECT column_name
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'disposition'
|
||||
"""),
|
||||
{"schema": target_schema},
|
||||
)
|
||||
if not result.fetchone():
|
||||
# Column doesn't exist yet (shouldn't happen but be safe)
|
||||
return
|
||||
|
||||
# Update all existing banks to use the new disposition format
|
||||
# Convert from old format to new format with reasonable mappings:
|
||||
@@ -31,32 +58,54 @@ def upgrade() -> None:
|
||||
# - literalism: derived from conscientiousness (detail-oriented people are more literal)
|
||||
# - empathy: derived from agreeableness + inverse of neuroticism
|
||||
# Default all to 3 (neutral) for simplicity
|
||||
conn.execute(sa.text("""
|
||||
UPDATE banks
|
||||
SET disposition = '{"skepticism": 3, "literalism": 3, "empathy": 3}'::jsonb
|
||||
conn.execute(
|
||||
sa.text(f"""
|
||||
UPDATE {schema}banks
|
||||
SET disposition = '{{"skepticism": 3, "literalism": 3, "empathy": 3}}'::jsonb
|
||||
WHERE disposition IS NOT NULL
|
||||
"""))
|
||||
""")
|
||||
)
|
||||
|
||||
# Update the default for new banks
|
||||
conn.execute(sa.text("""
|
||||
ALTER TABLE banks
|
||||
ALTER COLUMN disposition SET DEFAULT '{"skepticism": 3, "literalism": 3, "empathy": 3}'::jsonb
|
||||
"""))
|
||||
conn.execute(
|
||||
sa.text(f"""
|
||||
ALTER TABLE {schema}banks
|
||||
ALTER COLUMN disposition SET DEFAULT '{{"skepticism": 3, "literalism": 3, "empathy": 3}}'::jsonb
|
||||
""")
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Convert back to Big Five disposition."""
|
||||
conn = op.get_bind()
|
||||
schema = _get_schema_prefix()
|
||||
target_schema = _get_target_schema()
|
||||
|
||||
# Check if disposition column exists
|
||||
result = conn.execute(
|
||||
sa.text("""
|
||||
SELECT column_name
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'disposition'
|
||||
"""),
|
||||
{"schema": target_schema},
|
||||
)
|
||||
if not result.fetchone():
|
||||
return
|
||||
|
||||
# Revert to Big Five format with default values
|
||||
conn.execute(sa.text("""
|
||||
UPDATE banks
|
||||
SET disposition = '{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}'::jsonb
|
||||
conn.execute(
|
||||
sa.text(f"""
|
||||
UPDATE {schema}banks
|
||||
SET disposition = '{{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}}'::jsonb
|
||||
WHERE disposition IS NOT NULL
|
||||
"""))
|
||||
""")
|
||||
)
|
||||
|
||||
# Update the default for new banks
|
||||
conn.execute(sa.text("""
|
||||
ALTER TABLE banks
|
||||
ALTER COLUMN disposition SET DEFAULT '{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}'::jsonb
|
||||
"""))
|
||||
conn.execute(
|
||||
sa.text(f"""
|
||||
ALTER TABLE {schema}banks
|
||||
ALTER COLUMN disposition SET DEFAULT '{{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}}'::jsonb
|
||||
""")
|
||||
)
|
||||
|
||||
+44
@@ -0,0 +1,44 @@
|
||||
"""add_memory_links_from_type_weight_index
|
||||
|
||||
Revision ID: f1a2b3c4d5e6
|
||||
Revises: e0a1b2c3d4e5
|
||||
Create Date: 2025-01-12
|
||||
|
||||
Add composite index on memory_links (from_unit_id, link_type, weight DESC)
|
||||
to optimize MPFP graph traversal queries that need top-k edges per type.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "f1a2b3c4d5e6"
|
||||
down_revision: str | Sequence[str] | None = "e0a1b2c3d4e5"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (e.g., 'tenant_x.' or '' for public)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add composite index for efficient MPFP edge loading."""
|
||||
schema = _get_schema_prefix()
|
||||
# Create composite index for efficient top-k per (from_node, link_type) queries
|
||||
# This enables LATERAL joins to use index-only scans with early termination
|
||||
# Note: Not using CONCURRENTLY here as it requires running outside a transaction
|
||||
# For production with large tables, consider running this manually with CONCURRENTLY
|
||||
op.execute(
|
||||
f"CREATE INDEX IF NOT EXISTS idx_memory_links_from_type_weight "
|
||||
f"ON {schema}memory_links(from_unit_id, link_type, weight DESC)"
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove the composite index."""
|
||||
schema = _get_schema_prefix()
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_links_from_type_weight")
|
||||
@@ -0,0 +1,48 @@
|
||||
"""add_tags_column
|
||||
|
||||
Revision ID: g2a3b4c5d6e7
|
||||
Revises: f1a2b3c4d5e6
|
||||
Create Date: 2025-01-13
|
||||
|
||||
Add tags column to memory_units and documents tables for visibility scoping.
|
||||
Tags enable filtering memories by scope (e.g., user IDs, session IDs) during recall/reflect.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "g2a3b4c5d6e7"
|
||||
down_revision: str | Sequence[str] | None = "f1a2b3c4d5e6"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (e.g., 'tenant_x.' or '' for public)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add tags column to memory_units and documents tables."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Add tags column to memory_units table
|
||||
op.execute(f"ALTER TABLE {schema}memory_units ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
|
||||
|
||||
# Create GIN index for efficient array containment queries (tags && ARRAY['x'])
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_memory_units_tags ON {schema}memory_units USING GIN (tags)")
|
||||
|
||||
# Add tags column to documents table for document-level tags
|
||||
op.execute(f"ALTER TABLE {schema}documents ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove tags columns and index."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_tags")
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS tags")
|
||||
op.execute(f"ALTER TABLE {schema}documents DROP COLUMN IF EXISTS tags")
|
||||
@@ -0,0 +1,112 @@
|
||||
"""mental_models_v4
|
||||
|
||||
Revision ID: h3c4d5e6f7g8
|
||||
Revises: g2a3b4c5d6e7
|
||||
Create Date: 2026-01-08 00:00:00.000000
|
||||
|
||||
This migration implements the v4 mental models system:
|
||||
1. Deletes existing observation memory_units (observations now in mental models)
|
||||
2. Adds mission column to banks (replacing background)
|
||||
3. Creates mental_models table with final schema
|
||||
|
||||
Mental models can reference entities when an entity is "promoted" to a mental model.
|
||||
Summary content is stored as JSONB observations with per-observation fact attribution.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "h3c4d5e6f7g8"
|
||||
down_revision: str | Sequence[str] | None = "g2a3b4c5d6e7"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Apply mental models v4 changes."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Step 1: Delete observation memory_units (cascades to unit_entities links)
|
||||
# Observations are now handled through mental models, not memory_units
|
||||
op.execute(f"DELETE FROM {schema}memory_units WHERE fact_type = 'observation'")
|
||||
|
||||
# Step 2: Drop observation-specific index (if it exists)
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_observation_date")
|
||||
|
||||
# Step 3: Add mission column to banks (replacing background)
|
||||
op.execute(f"ALTER TABLE {schema}banks ADD COLUMN IF NOT EXISTS mission TEXT")
|
||||
|
||||
# Migrate: copy background to mission if background column exists
|
||||
# Use DO block to check column existence first (idempotent for re-runs)
|
||||
schema_name = context.config.get_main_option("target_schema") or "public"
|
||||
op.execute(f"""
|
||||
DO $$
|
||||
BEGIN
|
||||
IF EXISTS (
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_schema = '{schema_name}' AND table_name = 'banks' AND column_name = 'background'
|
||||
) THEN
|
||||
UPDATE {schema}banks
|
||||
SET mission = background
|
||||
WHERE mission IS NULL;
|
||||
END IF;
|
||||
END $$;
|
||||
""")
|
||||
|
||||
# Remove background column (replaced by mission)
|
||||
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS background")
|
||||
|
||||
# Step 4: Create mental_models table with final v4 schema (if not exists)
|
||||
op.execute(f"""
|
||||
CREATE TABLE IF NOT EXISTS {schema}mental_models (
|
||||
id VARCHAR(64) NOT NULL,
|
||||
bank_id VARCHAR(64) NOT NULL,
|
||||
subtype VARCHAR(32) NOT NULL,
|
||||
name VARCHAR(256) NOT NULL,
|
||||
description TEXT NOT NULL,
|
||||
entity_id UUID,
|
||||
observations JSONB DEFAULT '{{"observations": []}}'::jsonb,
|
||||
links VARCHAR[],
|
||||
tags VARCHAR[] DEFAULT '{{}}',
|
||||
last_updated TIMESTAMP WITH TIME ZONE,
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT now(),
|
||||
PRIMARY KEY (id, bank_id),
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id) ON DELETE CASCADE,
|
||||
FOREIGN KEY (entity_id) REFERENCES {schema}entities(id) ON DELETE SET NULL,
|
||||
CONSTRAINT ck_mental_models_subtype CHECK (subtype IN ('structural', 'emergent', 'pinned', 'learned'))
|
||||
)
|
||||
""")
|
||||
|
||||
# Step 5: Create indexes for efficient queries (if not exist)
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_bank_id ON {schema}mental_models(bank_id)")
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_subtype ON {schema}mental_models(bank_id, subtype)")
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_entity_id ON {schema}mental_models(entity_id)")
|
||||
# GIN index for efficient tags array filtering
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_mental_models_tags ON {schema}mental_models USING GIN(tags)")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Revert mental models v4 changes."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop mental_models table (cascades to indexes)
|
||||
op.execute(f"DROP TABLE IF EXISTS {schema}mental_models CASCADE")
|
||||
|
||||
# Add back background column to banks
|
||||
op.execute(f"ALTER TABLE {schema}banks ADD COLUMN IF NOT EXISTS background TEXT")
|
||||
|
||||
# Migrate mission back to background
|
||||
op.execute(f"UPDATE {schema}banks SET background = mission WHERE background IS NULL")
|
||||
|
||||
# Remove mission column
|
||||
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS mission")
|
||||
|
||||
# Note: Cannot restore deleted observations - they are lost on downgrade
|
||||
@@ -0,0 +1,41 @@
|
||||
"""delete_opinions
|
||||
|
||||
Revision ID: i4d5e6f7g8h9
|
||||
Revises: h3c4d5e6f7g8
|
||||
Create Date: 2026-01-15 00:00:00.000000
|
||||
|
||||
This migration removes opinion facts from memory_units.
|
||||
Opinions are no longer a separate fact type - they are now represented
|
||||
through mental model observations with confidence scores.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "i4d5e6f7g8h9"
|
||||
down_revision: str | Sequence[str] | None = "h3c4d5e6f7g8"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (required for multi-tenant support)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Delete opinion memory_units."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Delete opinion memory_units (cascades to unit_entities links)
|
||||
# Opinions are now handled through mental model observations
|
||||
op.execute(f"DELETE FROM {schema}memory_units WHERE fact_type = 'opinion'")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Cannot restore deleted opinions."""
|
||||
# Note: Cannot restore deleted opinions - they are lost on downgrade
|
||||
pass
|
||||
@@ -0,0 +1,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,
|
||||
)
|
||||
@@ -5,61 +5,81 @@ Revises: d9f6a3b4c5e2
|
||||
Create Date: 2024-12-04
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import context, op
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = 'rename_personality'
|
||||
down_revision: Union[str, Sequence[str], None] = 'd9f6a3b4c5e2'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
revision: str = "rename_personality"
|
||||
down_revision: str | Sequence[str] | None = "d9f6a3b4c5e2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_target_schema() -> str:
|
||||
"""Get the target schema name (tenant schema or 'public')."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return schema if schema else "public"
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Rename personality column to disposition in banks table (if it exists)."""
|
||||
conn = op.get_bind()
|
||||
target_schema = _get_target_schema()
|
||||
|
||||
# Check if 'personality' column exists (old database)
|
||||
result = conn.execute(sa.text("""
|
||||
result = conn.execute(
|
||||
sa.text("""
|
||||
SELECT column_name
|
||||
FROM information_schema.columns
|
||||
WHERE table_name = 'banks' AND column_name = 'personality'
|
||||
"""))
|
||||
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'personality'
|
||||
"""),
|
||||
{"schema": target_schema},
|
||||
)
|
||||
has_personality = result.fetchone() is not None
|
||||
|
||||
# Check if 'disposition' column exists (new database)
|
||||
result = conn.execute(sa.text("""
|
||||
result = conn.execute(
|
||||
sa.text("""
|
||||
SELECT column_name
|
||||
FROM information_schema.columns
|
||||
WHERE table_name = 'banks' AND column_name = 'disposition'
|
||||
"""))
|
||||
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'disposition'
|
||||
"""),
|
||||
{"schema": target_schema},
|
||||
)
|
||||
has_disposition = result.fetchone() is not None
|
||||
|
||||
if has_personality and not has_disposition:
|
||||
# Old database: rename personality -> disposition
|
||||
op.alter_column('banks', 'personality', new_column_name='disposition')
|
||||
op.alter_column("banks", "personality", new_column_name="disposition")
|
||||
elif not has_personality and not has_disposition:
|
||||
# Neither exists (shouldn't happen, but be safe): add disposition column
|
||||
op.add_column('banks', sa.Column(
|
||||
'disposition',
|
||||
postgresql.JSONB(astext_type=sa.Text()),
|
||||
server_default=sa.text("'{}'::jsonb"),
|
||||
nullable=False
|
||||
))
|
||||
op.add_column(
|
||||
"banks",
|
||||
sa.Column(
|
||||
"disposition",
|
||||
postgresql.JSONB(astext_type=sa.Text()),
|
||||
server_default=sa.text("'{}'::jsonb"),
|
||||
nullable=False,
|
||||
),
|
||||
)
|
||||
# else: disposition already exists, nothing to do
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Revert disposition column back to personality."""
|
||||
conn = op.get_bind()
|
||||
result = conn.execute(sa.text("""
|
||||
target_schema = _get_target_schema()
|
||||
result = conn.execute(
|
||||
sa.text("""
|
||||
SELECT column_name
|
||||
FROM information_schema.columns
|
||||
WHERE table_name = 'banks' AND column_name = 'disposition'
|
||||
"""))
|
||||
WHERE table_schema = :schema AND table_name = 'banks' AND column_name = 'disposition'
|
||||
"""),
|
||||
{"schema": target_schema},
|
||||
)
|
||||
if result.fetchone():
|
||||
op.alter_column('banks', 'disposition', new_column_name='personality')
|
||||
op.alter_column("banks", "disposition", new_column_name="personality")
|
||||
|
||||
@@ -3,8 +3,11 @@ Unified API module for Hindsight.
|
||||
|
||||
Provides both HTTP REST API and MCP (Model Context Protocol) server.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import FastAPI
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
@@ -17,7 +20,7 @@ def create_app(
|
||||
http_api_enabled: bool = True,
|
||||
mcp_api_enabled: bool = False,
|
||||
mcp_mount_path: str = "/mcp",
|
||||
initialize_memory: bool = True
|
||||
initialize_memory: bool = True,
|
||||
) -> FastAPI:
|
||||
"""
|
||||
Create and configure the unified Hindsight API application.
|
||||
@@ -43,49 +46,70 @@ def create_app(
|
||||
# Both HTTP and MCP
|
||||
app = create_app(memory, mcp_api_enabled=True)
|
||||
"""
|
||||
mcp_app = None
|
||||
|
||||
# Create MCP app first if enabled (we need its lifespan for chaining)
|
||||
if mcp_api_enabled:
|
||||
try:
|
||||
from .mcp import create_mcp_app
|
||||
|
||||
mcp_app = create_mcp_app(memory=memory)
|
||||
except ImportError as e:
|
||||
logger.error(f"MCP server requested but dependencies not available: {e}")
|
||||
logger.error("Install with: pip install hindsight-api[mcp]")
|
||||
raise
|
||||
|
||||
# Import and create HTTP API if enabled
|
||||
if http_api_enabled:
|
||||
from .http import create_app as create_http_app
|
||||
app = create_http_app(
|
||||
memory=memory,
|
||||
initialize_memory=initialize_memory
|
||||
)
|
||||
|
||||
app = create_http_app(memory=memory, initialize_memory=initialize_memory)
|
||||
logger.info("HTTP REST API enabled")
|
||||
else:
|
||||
# Create minimal FastAPI app
|
||||
app = FastAPI(title="Hindsight API", version="0.0.7")
|
||||
logger.info("HTTP REST API disabled")
|
||||
|
||||
# Mount MCP server if enabled
|
||||
if mcp_api_enabled:
|
||||
try:
|
||||
from .mcp import create_mcp_app
|
||||
# Mount MCP server and chain its lifespan if enabled
|
||||
if mcp_app is not None:
|
||||
# Get the MCP app's underlying Starlette app for lifespan access
|
||||
mcp_starlette_app = mcp_app.mcp_app
|
||||
|
||||
# Create MCP app with dynamic bank_id support
|
||||
# Supports: /mcp/{bank_id}/sse (bank-specific SSE endpoint)
|
||||
mcp_app = create_mcp_app(memory=memory)
|
||||
app.mount(mcp_mount_path, mcp_app)
|
||||
logger.info(f"MCP server enabled at {mcp_mount_path}/{{bank_id}}/sse")
|
||||
except ImportError as e:
|
||||
logger.error(f"MCP server requested but dependencies not available: {e}")
|
||||
logger.error("Install with: pip install hindsight-api[mcp]")
|
||||
raise
|
||||
# Store the original lifespan
|
||||
original_lifespan = app.router.lifespan_context
|
||||
|
||||
@asynccontextmanager
|
||||
async def chained_lifespan(app_instance: FastAPI):
|
||||
"""Chain the MCP lifespan with the main app lifespan."""
|
||||
# Start MCP lifespan first
|
||||
async with mcp_starlette_app.router.lifespan_context(mcp_starlette_app):
|
||||
logger.info("MCP lifespan started")
|
||||
# Then start the original app lifespan
|
||||
async with original_lifespan(app_instance):
|
||||
yield
|
||||
logger.info("MCP lifespan stopped")
|
||||
|
||||
# Replace the app's lifespan with the chained version
|
||||
app.router.lifespan_context = chained_lifespan
|
||||
|
||||
# Mount the MCP middleware
|
||||
app.mount(mcp_mount_path, mcp_app)
|
||||
logger.info(f"MCP server enabled at {mcp_mount_path}/")
|
||||
|
||||
return app
|
||||
|
||||
|
||||
# Re-export commonly used items for backwards compatibility
|
||||
from .http import (
|
||||
RecallRequest,
|
||||
RecallResult,
|
||||
RecallResponse,
|
||||
MemoryItem,
|
||||
RetainRequest,
|
||||
ReflectRequest,
|
||||
ReflectResponse,
|
||||
CreateBankRequest,
|
||||
DispositionTraits,
|
||||
MemoryItem,
|
||||
RecallRequest,
|
||||
RecallResponse,
|
||||
RecallResult,
|
||||
ReflectRequest,
|
||||
ReflectResponse,
|
||||
RetainRequest,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
|
||||
+2234
-778
File diff suppressed because it is too large
Load Diff
@@ -4,28 +4,38 @@ import json
|
||||
import logging
|
||||
import os
|
||||
from contextvars import ContextVar
|
||||
from typing import Optional
|
||||
|
||||
from fastmcp import FastMCP
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
# Configure logging from HINDSIGHT_API_LOG_LEVEL environment variable
|
||||
_log_level_str = os.environ.get("HINDSIGHT_API_LOG_LEVEL", "info").lower()
|
||||
_log_level_map = {"critical": logging.CRITICAL, "error": logging.ERROR, "warning": logging.WARNING,
|
||||
"info": logging.INFO, "debug": logging.DEBUG, "trace": logging.DEBUG}
|
||||
_log_level_map = {
|
||||
"critical": logging.CRITICAL,
|
||||
"error": logging.ERROR,
|
||||
"warning": logging.WARNING,
|
||||
"info": logging.INFO,
|
||||
"debug": logging.DEBUG,
|
||||
"trace": logging.DEBUG,
|
||||
}
|
||||
logging.basicConfig(
|
||||
level=_log_level_map.get(_log_level_str, logging.INFO),
|
||||
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s"
|
||||
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Context variable to hold the current bank_id from the URL path
|
||||
_current_bank_id: ContextVar[Optional[str]] = ContextVar("current_bank_id", default=None)
|
||||
# Default bank_id from environment variable
|
||||
DEFAULT_BANK_ID = os.environ.get("HINDSIGHT_MCP_BANK_ID", "default")
|
||||
|
||||
# Context variable to hold the current bank_id
|
||||
_current_bank_id: ContextVar[str | None] = ContextVar("current_bank_id", default=None)
|
||||
|
||||
|
||||
def get_current_bank_id() -> Optional[str]:
|
||||
"""Get the current bank_id from context (set from URL path)."""
|
||||
def get_current_bank_id() -> str | None:
|
||||
"""Get the current bank_id from context."""
|
||||
return _current_bank_id.get()
|
||||
|
||||
|
||||
@@ -37,12 +47,18 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
|
||||
memory: MemoryEngine instance (required)
|
||||
|
||||
Returns:
|
||||
Configured FastMCP server instance
|
||||
Configured FastMCP server instance with stateless_http enabled
|
||||
"""
|
||||
mcp = FastMCP("hindsight-mcp-server")
|
||||
# Use stateless_http=True for Claude Code compatibility
|
||||
mcp = FastMCP("hindsight-mcp-server", stateless_http=True)
|
||||
|
||||
@mcp.tool()
|
||||
async def retain(content: str, context: str = "general") -> str:
|
||||
async def retain(
|
||||
content: str,
|
||||
context: str = "general",
|
||||
async_processing: bool = True,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Store important information to long-term memory.
|
||||
|
||||
@@ -58,20 +74,34 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
|
||||
Args:
|
||||
content: The fact/memory to store (be specific and include relevant details)
|
||||
context: Category for the memory (e.g., 'preferences', 'work', 'hobbies', 'family'). Default: 'general'
|
||||
async_processing: If True, queue for background processing and return immediately. If False, wait for completion. Default: True
|
||||
bank_id: Optional bank to store in (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
bank_id = get_current_bank_id()
|
||||
await memory.put_batch_async(
|
||||
bank_id=bank_id,
|
||||
contents=[{"content": content, "context": context}]
|
||||
)
|
||||
return "Memory stored successfully"
|
||||
target_bank = bank_id or get_current_bank_id()
|
||||
if target_bank is None:
|
||||
return "Error: No bank_id configured"
|
||||
contents = [{"content": content, "context": context}]
|
||||
if async_processing:
|
||||
# Queue for background processing and return immediately
|
||||
result = await memory.submit_async_retain(
|
||||
bank_id=target_bank, contents=contents, request_context=RequestContext()
|
||||
)
|
||||
return f"Memory queued for background processing (operation_id: {result.get('operation_id', 'N/A')})"
|
||||
else:
|
||||
# Wait for completion
|
||||
await memory.retain_batch_async(
|
||||
bank_id=target_bank,
|
||||
contents=contents,
|
||||
request_context=RequestContext(),
|
||||
)
|
||||
return f"Memory stored successfully in bank '{target_bank}'"
|
||||
except Exception as e:
|
||||
logger.error(f"Error storing memory: {e}", exc_info=True)
|
||||
return f"Error: {str(e)}"
|
||||
|
||||
@mcp.tool()
|
||||
async def recall(query: str, max_results: int = 10) -> str:
|
||||
async def recall(query: str, max_tokens: int = 4096, bank_id: str | None = None) -> str:
|
||||
"""
|
||||
Search memories to provide personalized, context-aware responses.
|
||||
|
||||
@@ -83,45 +113,165 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
|
||||
|
||||
Args:
|
||||
query: Natural language search query (e.g., "user's food preferences", "what projects is user working on")
|
||||
max_results: Maximum number of results to return (default: 10)
|
||||
max_tokens: Maximum tokens in the response (default: 4096)
|
||||
bank_id: Optional bank to search in (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
bank_id = get_current_bank_id()
|
||||
target_bank = bank_id or get_current_bank_id()
|
||||
if target_bank is None:
|
||||
return "Error: No bank_id configured"
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
search_result = await memory.recall_async(
|
||||
bank_id=bank_id,
|
||||
|
||||
recall_result = await memory.recall_async(
|
||||
bank_id=target_bank,
|
||||
query=query,
|
||||
fact_type=list(VALID_RECALL_FACT_TYPES),
|
||||
budget=Budget.LOW
|
||||
budget=Budget.HIGH,
|
||||
max_tokens=max_tokens,
|
||||
request_context=RequestContext(),
|
||||
)
|
||||
|
||||
results = [
|
||||
{
|
||||
"id": fact.id,
|
||||
"text": fact.text,
|
||||
"type": fact.fact_type,
|
||||
"context": fact.context,
|
||||
"event_date": fact.event_date,
|
||||
}
|
||||
for fact in search_result.results[:max_results]
|
||||
]
|
||||
|
||||
return json.dumps({"results": results}, indent=2)
|
||||
# Use model's JSON serialization
|
||||
return recall_result.model_dump_json(indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"Error searching: {e}", exc_info=True)
|
||||
return json.dumps({"error": str(e), "results": []})
|
||||
return f'{{"error": "{e}", "results": []}}'
|
||||
|
||||
@mcp.tool()
|
||||
async def reflect(query: str, context: str | None = None, budget: str = "low", bank_id: str | None = None) -> str:
|
||||
"""
|
||||
Generate thoughtful analysis by synthesizing stored memories with the bank's personality.
|
||||
|
||||
WHEN TO USE THIS TOOL:
|
||||
Use reflect when you need reasoned analysis, not just fact retrieval. This tool
|
||||
thinks through the question using everything the bank knows and its personality traits.
|
||||
|
||||
EXAMPLES OF GOOD QUERIES:
|
||||
- "What patterns have emerged in how I approach debugging?"
|
||||
- "Based on my past decisions, what architectural style do I prefer?"
|
||||
- "What might be the best approach for this problem given what you know about me?"
|
||||
- "How should I prioritize these tasks based on my goals?"
|
||||
|
||||
HOW IT DIFFERS FROM RECALL:
|
||||
- recall: Returns raw facts matching your search (fast lookup)
|
||||
- reflect: Reasons across memories to form a synthesized answer (deeper analysis)
|
||||
|
||||
Use recall for "what did I say about X?" and reflect for "what should I do about X?"
|
||||
|
||||
Args:
|
||||
query: The question or topic to reflect on
|
||||
context: Optional context about why this reflection is needed
|
||||
budget: Search budget - 'low', 'mid', or 'high' (default: 'low')
|
||||
bank_id: Optional bank to reflect in (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or get_current_bank_id()
|
||||
if target_bank is None:
|
||||
return "Error: No bank_id configured"
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
|
||||
# Map string budget to enum
|
||||
budget_map = {"low": Budget.LOW, "mid": Budget.MID, "high": Budget.HIGH}
|
||||
budget_enum = budget_map.get(budget.lower(), Budget.LOW)
|
||||
|
||||
reflect_result = await memory.reflect_async(
|
||||
bank_id=target_bank,
|
||||
query=query,
|
||||
budget=budget_enum,
|
||||
context=context,
|
||||
request_context=RequestContext(),
|
||||
)
|
||||
|
||||
return reflect_result.model_dump_json(indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"Error reflecting: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}", "text": ""}}'
|
||||
|
||||
@mcp.tool()
|
||||
async def list_banks() -> str:
|
||||
"""
|
||||
List all available memory banks.
|
||||
|
||||
Use this tool to discover what memory banks exist in the system.
|
||||
Each bank is an isolated memory store (like a separate "brain").
|
||||
|
||||
Returns:
|
||||
JSON list of banks with their IDs, names, dispositions, and missions.
|
||||
"""
|
||||
try:
|
||||
banks = await memory.list_banks(request_context=RequestContext())
|
||||
return json.dumps({"banks": banks}, indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"Error listing banks: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}", "banks": []}}'
|
||||
|
||||
@mcp.tool()
|
||||
async def create_bank(bank_id: str, name: str | None = None, mission: str | None = None) -> str:
|
||||
"""
|
||||
Create a new memory bank or get an existing one.
|
||||
|
||||
Memory banks are isolated stores - each one is like a separate "brain" for a user/agent.
|
||||
Banks are auto-created with default settings if they don't exist.
|
||||
|
||||
Args:
|
||||
bank_id: Unique identifier for the bank (e.g., 'user-123', 'agent-alpha')
|
||||
name: Optional human-friendly name for the bank
|
||||
mission: Optional mission describing who the agent is and what they're trying to accomplish
|
||||
"""
|
||||
try:
|
||||
# get_bank_profile auto-creates bank if it doesn't exist
|
||||
profile = await memory.get_bank_profile(bank_id, request_context=RequestContext())
|
||||
|
||||
# Update name/mission if provided
|
||||
if name is not None or mission is not None:
|
||||
await memory.update_bank(
|
||||
bank_id,
|
||||
name=name,
|
||||
mission=mission,
|
||||
request_context=RequestContext(),
|
||||
)
|
||||
# Fetch updated profile
|
||||
profile = await memory.get_bank_profile(bank_id, request_context=RequestContext())
|
||||
|
||||
# Serialize disposition if it's a Pydantic model
|
||||
if "disposition" in profile and hasattr(profile["disposition"], "model_dump"):
|
||||
profile["disposition"] = profile["disposition"].model_dump()
|
||||
return json.dumps(profile, indent=2)
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating bank: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}"}}'
|
||||
|
||||
return mcp
|
||||
|
||||
|
||||
class MCPMiddleware:
|
||||
"""ASGI middleware that extracts bank_id from path and sets context."""
|
||||
"""ASGI middleware that extracts bank_id from header or path and sets context.
|
||||
|
||||
Bank ID can be provided via:
|
||||
1. X-Bank-Id header (recommended for Claude Code)
|
||||
2. URL path: /mcp/{bank_id}/
|
||||
3. Environment variable HINDSIGHT_MCP_BANK_ID (fallback default)
|
||||
|
||||
For Claude Code, configure with:
|
||||
claude mcp add --transport http hindsight http://localhost:8888/mcp \\
|
||||
--header "X-Bank-Id: my-bank"
|
||||
"""
|
||||
|
||||
def __init__(self, app, memory: MemoryEngine):
|
||||
self.app = app
|
||||
self.memory = memory
|
||||
self.mcp_server = create_mcp_server(memory)
|
||||
self.mcp_app = self.mcp_server.http_app()
|
||||
self.mcp_app = self.mcp_server.http_app(path="/")
|
||||
# Expose the lifespan for the parent app to chain
|
||||
self.lifespan = self.mcp_app.lifespan_handler if hasattr(self.mcp_app, "lifespan_handler") else None
|
||||
|
||||
def _get_header(self, scope: dict, name: str) -> str | None:
|
||||
"""Extract a header value from ASGI scope."""
|
||||
name_lower = name.lower().encode()
|
||||
for header_name, header_value in scope.get("headers", []):
|
||||
if header_name.lower() == name_lower:
|
||||
return header_value.decode()
|
||||
return None
|
||||
|
||||
async def __call__(self, scope, receive, send):
|
||||
if scope["type"] != "http":
|
||||
@@ -133,46 +283,50 @@ class MCPMiddleware:
|
||||
# Strip any mount prefix (e.g., /mcp) that FastAPI might not have stripped
|
||||
root_path = scope.get("root_path", "")
|
||||
if root_path and path.startswith(root_path):
|
||||
path = path[len(root_path):] or "/"
|
||||
path = path[len(root_path) :] or "/"
|
||||
|
||||
# Also handle case where mount path wasn't stripped (e.g., /mcp/...)
|
||||
if path.startswith("/mcp/"):
|
||||
path = path[4:] # Remove /mcp prefix
|
||||
elif path == "/mcp":
|
||||
path = "/"
|
||||
|
||||
# Extract bank_id from path: /{bank_id}/ or /{bank_id}
|
||||
# http_app expects requests at /
|
||||
if not path.startswith("/") or len(path) <= 1:
|
||||
# No bank_id in path - return error
|
||||
await self._send_error(send, 400, "bank_id required in path: /mcp/{bank_id}/")
|
||||
return
|
||||
# Try to get bank_id from header first (for Claude Code compatibility)
|
||||
bank_id = self._get_header(scope, "X-Bank-Id")
|
||||
|
||||
# Extract bank_id from first path segment
|
||||
parts = path[1:].split("/", 1)
|
||||
if not parts[0]:
|
||||
await self._send_error(send, 400, "bank_id required in path: /mcp/{bank_id}/")
|
||||
return
|
||||
# MCP endpoint paths that should not be treated as bank_ids
|
||||
MCP_ENDPOINTS = {"sse", "messages"}
|
||||
|
||||
bank_id = parts[0]
|
||||
new_path = "/" + parts[1] if len(parts) > 1 else "/"
|
||||
# If no header, try to extract from path: /{bank_id}/...
|
||||
new_path = path
|
||||
if not bank_id and path.startswith("/") and len(path) > 1:
|
||||
parts = path[1:].split("/", 1)
|
||||
# Don't treat MCP endpoints as bank_ids
|
||||
if parts[0] and parts[0] not in MCP_ENDPOINTS:
|
||||
# First segment looks like a bank_id
|
||||
bank_id = parts[0]
|
||||
new_path = "/" + parts[1] if len(parts) > 1 else "/"
|
||||
|
||||
# Fall back to default bank_id
|
||||
if not bank_id:
|
||||
bank_id = DEFAULT_BANK_ID
|
||||
logger.debug(f"Using default bank_id: {bank_id}")
|
||||
|
||||
# Set bank_id context
|
||||
token = _current_bank_id.set(bank_id)
|
||||
try:
|
||||
new_scope = scope.copy()
|
||||
new_scope["path"] = new_path
|
||||
# Clear root_path since we're passing directly to the app
|
||||
new_scope["root_path"] = ""
|
||||
|
||||
# Wrap send to rewrite the SSE endpoint URL to include bank_id
|
||||
# The SSE app sends "event: endpoint\ndata: /messages\n" but we need
|
||||
# the client to POST to /{bank_id}/messages instead
|
||||
# Wrap send to rewrite the SSE endpoint URL to include bank_id if using path-based routing
|
||||
async def send_wrapper(message):
|
||||
if message["type"] == "http.response.body":
|
||||
body = message.get("body", b"")
|
||||
if body and b"/messages" in body:
|
||||
# Rewrite /messages to /{bank_id}/messages in SSE endpoint event
|
||||
body = body.replace(
|
||||
b"data: /messages",
|
||||
f"data: /{bank_id}/messages".encode()
|
||||
)
|
||||
body = body.replace(b"data: /messages", f"data: /{bank_id}/messages".encode())
|
||||
message = {**message, "body": body}
|
||||
await send(message)
|
||||
|
||||
@@ -183,24 +337,29 @@ class MCPMiddleware:
|
||||
async def _send_error(self, send, status: int, message: str):
|
||||
"""Send an error response."""
|
||||
body = json.dumps({"error": message}).encode()
|
||||
await send({
|
||||
"type": "http.response.start",
|
||||
"status": status,
|
||||
"headers": [(b"content-type", b"application/json")],
|
||||
})
|
||||
await send({
|
||||
"type": "http.response.body",
|
||||
"body": body,
|
||||
})
|
||||
await send(
|
||||
{
|
||||
"type": "http.response.start",
|
||||
"status": status,
|
||||
"headers": [(b"content-type", b"application/json")],
|
||||
}
|
||||
)
|
||||
await send(
|
||||
{
|
||||
"type": "http.response.body",
|
||||
"body": body,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def create_mcp_app(memory: MemoryEngine):
|
||||
"""
|
||||
Create an ASGI app that handles MCP requests.
|
||||
|
||||
URL pattern: /mcp/{bank_id}/
|
||||
|
||||
The bank_id is extracted from the URL path and made available to tools.
|
||||
Bank ID can be provided via:
|
||||
1. X-Bank-Id header: claude mcp add --transport http hindsight http://localhost:8888/mcp --header "X-Bank-Id: my-bank"
|
||||
2. URL path: /mcp/{bank_id}/
|
||||
3. Environment variable HINDSIGHT_MCP_BANK_ID (fallback, default: "default")
|
||||
|
||||
Args:
|
||||
memory: MemoryEngine instance
|
||||
|
||||
@@ -6,7 +6,7 @@ Shows the logo and tagline with gradient colors.
|
||||
|
||||
# Gradient colors: #0074d9 -> #009296
|
||||
GRADIENT_START = (0, 116, 217) # #0074d9
|
||||
GRADIENT_END = (0, 146, 150) # #009296
|
||||
GRADIENT_END = (0, 146, 150) # #009296
|
||||
|
||||
# Pre-generated logo (generated by test-logo.py)
|
||||
LOGO = """\
|
||||
@@ -31,8 +31,8 @@ def gradient_text(text: str, start: tuple = GRADIENT_START, end: tuple = GRADIEN
|
||||
result = []
|
||||
length = len(text)
|
||||
for i, char in enumerate(text):
|
||||
if char == ' ':
|
||||
result.append(' ')
|
||||
if char == " ":
|
||||
result.append(" ")
|
||||
else:
|
||||
t = i / max(length - 1, 1)
|
||||
r, g, b = _interpolate_color(start, end, t)
|
||||
@@ -74,9 +74,16 @@ def dim(text: str) -> str:
|
||||
return f"\033[38;2;128;128;128m{text}\033[0m"
|
||||
|
||||
|
||||
def print_startup_info(host: str, port: int, database_url: str, llm_provider: str,
|
||||
llm_model: str, embeddings_provider: str, reranker_provider: str,
|
||||
mcp_enabled: bool = False):
|
||||
def print_startup_info(
|
||||
host: str,
|
||||
port: int,
|
||||
database_url: str,
|
||||
llm_provider: str,
|
||||
llm_model: str,
|
||||
embeddings_provider: str,
|
||||
reranker_provider: str,
|
||||
mcp_enabled: bool = False,
|
||||
):
|
||||
"""Print styled startup information."""
|
||||
print(color_start("Starting Hindsight API..."))
|
||||
print(f" {dim('URL:')} {color(f'http://{host}:{port}', 0.2)}")
|
||||
|
||||
@@ -3,10 +3,18 @@ Centralized configuration for Hindsight API.
|
||||
|
||||
All environment variables and their defaults are defined here.
|
||||
"""
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from dotenv import find_dotenv, load_dotenv
|
||||
|
||||
# Load .env file, searching current and parent directories (overrides existing env vars)
|
||||
load_dotenv(find_dotenv(usecwd=True), override=True)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -16,40 +24,237 @@ ENV_LLM_PROVIDER = "HINDSIGHT_API_LLM_PROVIDER"
|
||||
ENV_LLM_API_KEY = "HINDSIGHT_API_LLM_API_KEY"
|
||||
ENV_LLM_MODEL = "HINDSIGHT_API_LLM_MODEL"
|
||||
ENV_LLM_BASE_URL = "HINDSIGHT_API_LLM_BASE_URL"
|
||||
ENV_LLM_MAX_CONCURRENT = "HINDSIGHT_API_LLM_MAX_CONCURRENT"
|
||||
ENV_LLM_TIMEOUT = "HINDSIGHT_API_LLM_TIMEOUT"
|
||||
ENV_LLM_GROQ_SERVICE_TIER = "HINDSIGHT_API_LLM_GROQ_SERVICE_TIER"
|
||||
|
||||
# Per-operation LLM configuration (optional, falls back to global LLM config)
|
||||
ENV_RETAIN_LLM_PROVIDER = "HINDSIGHT_API_RETAIN_LLM_PROVIDER"
|
||||
ENV_RETAIN_LLM_API_KEY = "HINDSIGHT_API_RETAIN_LLM_API_KEY"
|
||||
ENV_RETAIN_LLM_MODEL = "HINDSIGHT_API_RETAIN_LLM_MODEL"
|
||||
ENV_RETAIN_LLM_BASE_URL = "HINDSIGHT_API_RETAIN_LLM_BASE_URL"
|
||||
|
||||
ENV_REFLECT_LLM_PROVIDER = "HINDSIGHT_API_REFLECT_LLM_PROVIDER"
|
||||
ENV_REFLECT_LLM_API_KEY = "HINDSIGHT_API_REFLECT_LLM_API_KEY"
|
||||
ENV_REFLECT_LLM_MODEL = "HINDSIGHT_API_REFLECT_LLM_MODEL"
|
||||
ENV_REFLECT_LLM_BASE_URL = "HINDSIGHT_API_REFLECT_LLM_BASE_URL"
|
||||
|
||||
ENV_EMBEDDINGS_PROVIDER = "HINDSIGHT_API_EMBEDDINGS_PROVIDER"
|
||||
ENV_EMBEDDINGS_LOCAL_MODEL = "HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL"
|
||||
ENV_EMBEDDINGS_TEI_URL = "HINDSIGHT_API_EMBEDDINGS_TEI_URL"
|
||||
ENV_EMBEDDINGS_OPENAI_API_KEY = "HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY"
|
||||
ENV_EMBEDDINGS_OPENAI_MODEL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL"
|
||||
ENV_EMBEDDINGS_OPENAI_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL"
|
||||
|
||||
ENV_COHERE_API_KEY = "HINDSIGHT_API_COHERE_API_KEY"
|
||||
ENV_EMBEDDINGS_COHERE_MODEL = "HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL"
|
||||
ENV_EMBEDDINGS_COHERE_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL"
|
||||
ENV_RERANKER_COHERE_MODEL = "HINDSIGHT_API_RERANKER_COHERE_MODEL"
|
||||
ENV_RERANKER_COHERE_BASE_URL = "HINDSIGHT_API_RERANKER_COHERE_BASE_URL"
|
||||
|
||||
# LiteLLM gateway configuration (for embeddings and reranker via LiteLLM proxy)
|
||||
ENV_LITELLM_API_BASE = "HINDSIGHT_API_LITELLM_API_BASE"
|
||||
ENV_LITELLM_API_KEY = "HINDSIGHT_API_LITELLM_API_KEY"
|
||||
ENV_EMBEDDINGS_LITELLM_MODEL = "HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL"
|
||||
ENV_RERANKER_LITELLM_MODEL = "HINDSIGHT_API_RERANKER_LITELLM_MODEL"
|
||||
|
||||
ENV_RERANKER_PROVIDER = "HINDSIGHT_API_RERANKER_PROVIDER"
|
||||
ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
|
||||
ENV_RERANKER_LOCAL_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT"
|
||||
ENV_RERANKER_TEI_URL = "HINDSIGHT_API_RERANKER_TEI_URL"
|
||||
ENV_RERANKER_TEI_BATCH_SIZE = "HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE"
|
||||
ENV_RERANKER_TEI_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT"
|
||||
ENV_RERANKER_MAX_CANDIDATES = "HINDSIGHT_API_RERANKER_MAX_CANDIDATES"
|
||||
ENV_RERANKER_FLASHRANK_MODEL = "HINDSIGHT_API_RERANKER_FLASHRANK_MODEL"
|
||||
ENV_RERANKER_FLASHRANK_CACHE_DIR = "HINDSIGHT_API_RERANKER_FLASHRANK_CACHE_DIR"
|
||||
|
||||
ENV_HOST = "HINDSIGHT_API_HOST"
|
||||
ENV_PORT = "HINDSIGHT_API_PORT"
|
||||
ENV_LOG_LEVEL = "HINDSIGHT_API_LOG_LEVEL"
|
||||
ENV_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"
|
||||
ENV_MPFP_TOP_K_NEIGHBORS = "HINDSIGHT_API_MPFP_TOP_K_NEIGHBORS"
|
||||
ENV_RECALL_MAX_CONCURRENT = "HINDSIGHT_API_RECALL_MAX_CONCURRENT"
|
||||
ENV_RECALL_CONNECTION_BUDGET = "HINDSIGHT_API_RECALL_CONNECTION_BUDGET"
|
||||
ENV_MCP_LOCAL_BANK_ID = "HINDSIGHT_API_MCP_LOCAL_BANK_ID"
|
||||
ENV_MCP_INSTRUCTIONS = "HINDSIGHT_API_MCP_INSTRUCTIONS"
|
||||
ENV_MENTAL_MODEL_REFRESH_CONCURRENCY = "HINDSIGHT_API_MENTAL_MODEL_REFRESH_CONCURRENCY"
|
||||
|
||||
# Observation thresholds
|
||||
ENV_OBSERVATION_MIN_FACTS = "HINDSIGHT_API_OBSERVATION_MIN_FACTS"
|
||||
ENV_OBSERVATION_TOP_ENTITIES = "HINDSIGHT_API_OBSERVATION_TOP_ENTITIES"
|
||||
|
||||
# Retain settings
|
||||
ENV_RETAIN_MAX_COMPLETION_TOKENS = "HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS"
|
||||
ENV_RETAIN_CHUNK_SIZE = "HINDSIGHT_API_RETAIN_CHUNK_SIZE"
|
||||
ENV_RETAIN_EXTRACT_CAUSAL_LINKS = "HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS"
|
||||
ENV_RETAIN_EXTRACTION_MODE = "HINDSIGHT_API_RETAIN_EXTRACTION_MODE"
|
||||
ENV_RETAIN_OBSERVATIONS_ASYNC = "HINDSIGHT_API_RETAIN_OBSERVATIONS_ASYNC"
|
||||
|
||||
# Optimization flags
|
||||
ENV_SKIP_LLM_VERIFICATION = "HINDSIGHT_API_SKIP_LLM_VERIFICATION"
|
||||
ENV_LAZY_RERANKER = "HINDSIGHT_API_LAZY_RERANKER"
|
||||
|
||||
# Database migrations
|
||||
ENV_RUN_MIGRATIONS_ON_STARTUP = "HINDSIGHT_API_RUN_MIGRATIONS_ON_STARTUP"
|
||||
|
||||
# Database connection pool
|
||||
ENV_DB_POOL_MIN_SIZE = "HINDSIGHT_API_DB_POOL_MIN_SIZE"
|
||||
ENV_DB_POOL_MAX_SIZE = "HINDSIGHT_API_DB_POOL_MAX_SIZE"
|
||||
ENV_DB_COMMAND_TIMEOUT = "HINDSIGHT_API_DB_COMMAND_TIMEOUT"
|
||||
ENV_DB_ACQUIRE_TIMEOUT = "HINDSIGHT_API_DB_ACQUIRE_TIMEOUT"
|
||||
|
||||
# Worker configuration (distributed task processing)
|
||||
ENV_WORKER_ENABLED = "HINDSIGHT_API_WORKER_ENABLED"
|
||||
ENV_WORKER_ID = "HINDSIGHT_API_WORKER_ID"
|
||||
ENV_WORKER_POLL_INTERVAL_MS = "HINDSIGHT_API_WORKER_POLL_INTERVAL_MS"
|
||||
ENV_WORKER_MAX_RETRIES = "HINDSIGHT_API_WORKER_MAX_RETRIES"
|
||||
ENV_WORKER_BATCH_SIZE = "HINDSIGHT_API_WORKER_BATCH_SIZE"
|
||||
ENV_WORKER_HTTP_PORT = "HINDSIGHT_API_WORKER_HTTP_PORT"
|
||||
|
||||
# Reflect agent settings
|
||||
ENV_REFLECT_MAX_ITERATIONS = "HINDSIGHT_API_REFLECT_MAX_ITERATIONS"
|
||||
|
||||
# Default values
|
||||
DEFAULT_DATABASE_URL = "pg0"
|
||||
DEFAULT_LLM_PROVIDER = "openai"
|
||||
DEFAULT_LLM_MODEL = "gpt-5-mini"
|
||||
DEFAULT_LLM_MAX_CONCURRENT = 32
|
||||
DEFAULT_LLM_TIMEOUT = 120.0 # seconds
|
||||
|
||||
DEFAULT_EMBEDDINGS_PROVIDER = "local"
|
||||
DEFAULT_EMBEDDINGS_LOCAL_MODEL = "BAAI/bge-small-en-v1.5"
|
||||
DEFAULT_EMBEDDINGS_OPENAI_MODEL = "text-embedding-3-small"
|
||||
DEFAULT_EMBEDDING_DIMENSION = 384
|
||||
|
||||
DEFAULT_RERANKER_PROVIDER = "local"
|
||||
DEFAULT_RERANKER_LOCAL_MODEL = "cross-encoder/ms-marco-MiniLM-L-6-v2"
|
||||
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT = 4 # Limit concurrent CPU-bound reranking to prevent thrashing
|
||||
DEFAULT_RERANKER_TEI_BATCH_SIZE = 128
|
||||
DEFAULT_RERANKER_TEI_MAX_CONCURRENT = 8
|
||||
DEFAULT_RERANKER_MAX_CANDIDATES = 300
|
||||
DEFAULT_RERANKER_FLASHRANK_MODEL = "ms-marco-MiniLM-L-12-v2" # Best balance of speed and quality
|
||||
DEFAULT_RERANKER_FLASHRANK_CACHE_DIR = None # Use default cache directory
|
||||
|
||||
DEFAULT_EMBEDDINGS_COHERE_MODEL = "embed-english-v3.0"
|
||||
DEFAULT_RERANKER_COHERE_MODEL = "rerank-english-v3.0"
|
||||
|
||||
# LiteLLM defaults
|
||||
DEFAULT_LITELLM_API_BASE = "http://localhost:4000"
|
||||
DEFAULT_EMBEDDINGS_LITELLM_MODEL = "text-embedding-3-small"
|
||||
DEFAULT_RERANKER_LITELLM_MODEL = "cohere/rerank-english-v3.0"
|
||||
|
||||
DEFAULT_HOST = "0.0.0.0"
|
||||
DEFAULT_PORT = 8888
|
||||
DEFAULT_LOG_LEVEL = "info"
|
||||
DEFAULT_LOG_FORMAT = "text" # Options: "text", "json"
|
||||
DEFAULT_WORKERS = 1
|
||||
DEFAULT_MCP_ENABLED = True
|
||||
DEFAULT_GRAPH_RETRIEVER = "bfs" # Options: "bfs", "mpfp"
|
||||
DEFAULT_GRAPH_RETRIEVER = "link_expansion" # Options: "link_expansion", "mpfp", "bfs"
|
||||
DEFAULT_MPFP_TOP_K_NEIGHBORS = 20 # Fan-out limit per node in MPFP graph traversal
|
||||
DEFAULT_RECALL_MAX_CONCURRENT = 32 # Max concurrent recall operations per worker
|
||||
DEFAULT_RECALL_CONNECTION_BUDGET = 4 # Max concurrent DB connections per recall operation
|
||||
DEFAULT_MCP_LOCAL_BANK_ID = "mcp"
|
||||
DEFAULT_MENTAL_MODEL_REFRESH_CONCURRENCY = 8 # Max concurrent mental model refreshes
|
||||
|
||||
# Required embedding dimension for database schema
|
||||
EMBEDDING_DIMENSION = 384
|
||||
# Observation thresholds
|
||||
DEFAULT_OBSERVATION_MIN_FACTS = 5 # Min facts required to generate entity observations
|
||||
DEFAULT_OBSERVATION_TOP_ENTITIES = 5 # Max entities to process per retain batch
|
||||
|
||||
# Retain settings
|
||||
DEFAULT_RETAIN_MAX_COMPLETION_TOKENS = 64000 # Max tokens for fact extraction LLM call
|
||||
DEFAULT_RETAIN_CHUNK_SIZE = 3000 # Max chars per chunk for fact extraction
|
||||
DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS = True # Extract causal links between facts
|
||||
DEFAULT_RETAIN_EXTRACTION_MODE = "concise" # Extraction mode: "concise" or "verbose"
|
||||
RETAIN_EXTRACTION_MODES = ("concise", "verbose") # Allowed extraction modes
|
||||
DEFAULT_RETAIN_OBSERVATIONS_ASYNC = False # Run observation generation async (after retain completes)
|
||||
|
||||
# Database migrations
|
||||
DEFAULT_RUN_MIGRATIONS_ON_STARTUP = True
|
||||
|
||||
# Database connection pool
|
||||
DEFAULT_DB_POOL_MIN_SIZE = 5
|
||||
DEFAULT_DB_POOL_MAX_SIZE = 100
|
||||
DEFAULT_DB_COMMAND_TIMEOUT = 60 # seconds
|
||||
DEFAULT_DB_ACQUIRE_TIMEOUT = 30 # seconds
|
||||
|
||||
# Worker configuration (distributed task processing)
|
||||
DEFAULT_WORKER_ENABLED = True # API runs worker by default (standalone mode)
|
||||
DEFAULT_WORKER_ID = None # Will use hostname if not specified
|
||||
DEFAULT_WORKER_POLL_INTERVAL_MS = 500 # Poll database every 500ms
|
||||
DEFAULT_WORKER_MAX_RETRIES = 3 # Max retries before marking task failed
|
||||
DEFAULT_WORKER_BATCH_SIZE = 10 # Tasks to claim per poll cycle
|
||||
DEFAULT_WORKER_HTTP_PORT = 8889 # HTTP port for worker metrics/health
|
||||
|
||||
# Reflect agent settings
|
||||
DEFAULT_REFLECT_MAX_ITERATIONS = 10 # Max tool call iterations before forcing response
|
||||
|
||||
# Default MCP tool descriptions (can be customized via env vars)
|
||||
DEFAULT_MCP_RETAIN_DESCRIPTION = """Store important information to long-term memory.
|
||||
|
||||
Use this tool PROACTIVELY whenever the user shares:
|
||||
- Personal facts, preferences, or interests
|
||||
- Important events or milestones
|
||||
- User history, experiences, or background
|
||||
- Decisions, opinions, or stated preferences
|
||||
- Goals, plans, or future intentions
|
||||
- Relationships or people mentioned
|
||||
- Work context, projects, or responsibilities"""
|
||||
|
||||
DEFAULT_MCP_RECALL_DESCRIPTION = """Search memories to provide personalized, context-aware responses.
|
||||
|
||||
Use this tool PROACTIVELY to:
|
||||
- Check user's preferences before making suggestions
|
||||
- Recall user's history to provide continuity
|
||||
- Remember user's goals and context
|
||||
- Personalize responses based on past interactions"""
|
||||
|
||||
# Default embedding dimension (used by initial migration, adjusted at runtime)
|
||||
EMBEDDING_DIMENSION = DEFAULT_EMBEDDING_DIMENSION
|
||||
|
||||
|
||||
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()
|
||||
if mode_lower not in RETAIN_EXTRACTION_MODES:
|
||||
logger.warning(
|
||||
f"Invalid extraction mode '{mode}', must be one of {RETAIN_EXTRACTION_MODES}. "
|
||||
f"Defaulting to '{DEFAULT_RETAIN_EXTRACTION_MODE}'."
|
||||
)
|
||||
return DEFAULT_RETAIN_EXTRACTION_MODE
|
||||
return mode_lower
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -59,30 +264,89 @@ class HindsightConfig:
|
||||
# Database
|
||||
database_url: str
|
||||
|
||||
# LLM
|
||||
# LLM (default, used as fallback for per-operation config)
|
||||
llm_provider: str
|
||||
llm_api_key: Optional[str]
|
||||
llm_api_key: str | None
|
||||
llm_model: str
|
||||
llm_base_url: Optional[str]
|
||||
llm_base_url: str | None
|
||||
llm_max_concurrent: int
|
||||
llm_timeout: float
|
||||
|
||||
# Per-operation LLM configuration (None = use default LLM config)
|
||||
retain_llm_provider: str | None
|
||||
retain_llm_api_key: str | None
|
||||
retain_llm_model: str | None
|
||||
retain_llm_base_url: str | None
|
||||
|
||||
reflect_llm_provider: str | None
|
||||
reflect_llm_api_key: str | None
|
||||
reflect_llm_model: str | None
|
||||
reflect_llm_base_url: str | None
|
||||
|
||||
# Embeddings
|
||||
embeddings_provider: str
|
||||
embeddings_local_model: str
|
||||
embeddings_tei_url: Optional[str]
|
||||
embeddings_tei_url: str | None
|
||||
embeddings_openai_base_url: str | None
|
||||
embeddings_cohere_base_url: str | None
|
||||
|
||||
# Reranker
|
||||
reranker_provider: str
|
||||
reranker_local_model: str
|
||||
reranker_tei_url: Optional[str]
|
||||
reranker_tei_url: str | None
|
||||
reranker_tei_batch_size: int
|
||||
reranker_tei_max_concurrent: int
|
||||
reranker_max_candidates: int
|
||||
reranker_cohere_base_url: str | None
|
||||
|
||||
# Server
|
||||
host: str
|
||||
port: int
|
||||
log_level: str
|
||||
log_format: str
|
||||
mcp_enabled: bool
|
||||
|
||||
# Recall
|
||||
graph_retriever: str
|
||||
mpfp_top_k_neighbors: int
|
||||
recall_max_concurrent: int
|
||||
recall_connection_budget: int
|
||||
mental_model_refresh_concurrency: int
|
||||
|
||||
# Observation thresholds
|
||||
observation_min_facts: int
|
||||
observation_top_entities: int
|
||||
|
||||
# Retain settings
|
||||
retain_max_completion_tokens: int
|
||||
retain_chunk_size: int
|
||||
retain_extract_causal_links: bool
|
||||
retain_extraction_mode: str
|
||||
retain_observations_async: bool
|
||||
|
||||
# Optimization flags
|
||||
skip_llm_verification: bool
|
||||
lazy_reranker: bool
|
||||
|
||||
# Database migrations
|
||||
run_migrations_on_startup: bool
|
||||
|
||||
# Database connection pool
|
||||
db_pool_min_size: int
|
||||
db_pool_max_size: int
|
||||
db_command_timeout: int
|
||||
db_acquire_timeout: int
|
||||
|
||||
# Worker configuration (distributed task processing)
|
||||
worker_enabled: bool
|
||||
worker_id: str | None
|
||||
worker_poll_interval_ms: int
|
||||
worker_max_retries: int
|
||||
worker_batch_size: int
|
||||
worker_http_port: int
|
||||
|
||||
# Reflect agent settings
|
||||
reflect_max_iterations: int
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> "HindsightConfig":
|
||||
@@ -90,31 +354,94 @@ class HindsightConfig:
|
||||
return cls(
|
||||
# Database
|
||||
database_url=os.getenv(ENV_DATABASE_URL, DEFAULT_DATABASE_URL),
|
||||
|
||||
# LLM
|
||||
llm_provider=os.getenv(ENV_LLM_PROVIDER, DEFAULT_LLM_PROVIDER),
|
||||
llm_api_key=os.getenv(ENV_LLM_API_KEY),
|
||||
llm_model=os.getenv(ENV_LLM_MODEL, DEFAULT_LLM_MODEL),
|
||||
llm_base_url=os.getenv(ENV_LLM_BASE_URL) or None,
|
||||
|
||||
llm_max_concurrent=int(os.getenv(ENV_LLM_MAX_CONCURRENT, str(DEFAULT_LLM_MAX_CONCURRENT))),
|
||||
llm_timeout=float(os.getenv(ENV_LLM_TIMEOUT, str(DEFAULT_LLM_TIMEOUT))),
|
||||
# Per-operation LLM config (None = use default)
|
||||
retain_llm_provider=os.getenv(ENV_RETAIN_LLM_PROVIDER) or None,
|
||||
retain_llm_api_key=os.getenv(ENV_RETAIN_LLM_API_KEY) or None,
|
||||
retain_llm_model=os.getenv(ENV_RETAIN_LLM_MODEL) or None,
|
||||
retain_llm_base_url=os.getenv(ENV_RETAIN_LLM_BASE_URL) or None,
|
||||
reflect_llm_provider=os.getenv(ENV_REFLECT_LLM_PROVIDER) or None,
|
||||
reflect_llm_api_key=os.getenv(ENV_REFLECT_LLM_API_KEY) or None,
|
||||
reflect_llm_model=os.getenv(ENV_REFLECT_LLM_MODEL) or None,
|
||||
reflect_llm_base_url=os.getenv(ENV_REFLECT_LLM_BASE_URL) or None,
|
||||
# Embeddings
|
||||
embeddings_provider=os.getenv(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER),
|
||||
embeddings_local_model=os.getenv(ENV_EMBEDDINGS_LOCAL_MODEL, DEFAULT_EMBEDDINGS_LOCAL_MODEL),
|
||||
embeddings_tei_url=os.getenv(ENV_EMBEDDINGS_TEI_URL),
|
||||
|
||||
embeddings_openai_base_url=os.getenv(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None,
|
||||
embeddings_cohere_base_url=os.getenv(ENV_EMBEDDINGS_COHERE_BASE_URL) or None,
|
||||
# Reranker
|
||||
reranker_provider=os.getenv(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER),
|
||||
reranker_local_model=os.getenv(ENV_RERANKER_LOCAL_MODEL, DEFAULT_RERANKER_LOCAL_MODEL),
|
||||
reranker_tei_url=os.getenv(ENV_RERANKER_TEI_URL),
|
||||
|
||||
reranker_tei_batch_size=int(os.getenv(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE))),
|
||||
reranker_tei_max_concurrent=int(
|
||||
os.getenv(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT))
|
||||
),
|
||||
reranker_max_candidates=int(os.getenv(ENV_RERANKER_MAX_CANDIDATES, str(DEFAULT_RERANKER_MAX_CANDIDATES))),
|
||||
reranker_cohere_base_url=os.getenv(ENV_RERANKER_COHERE_BASE_URL) or None,
|
||||
# Server
|
||||
host=os.getenv(ENV_HOST, DEFAULT_HOST),
|
||||
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
|
||||
log_level=os.getenv(ENV_LOG_LEVEL, DEFAULT_LOG_LEVEL),
|
||||
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),
|
||||
mpfp_top_k_neighbors=int(os.getenv(ENV_MPFP_TOP_K_NEIGHBORS, str(DEFAULT_MPFP_TOP_K_NEIGHBORS))),
|
||||
recall_max_concurrent=int(os.getenv(ENV_RECALL_MAX_CONCURRENT, str(DEFAULT_RECALL_MAX_CONCURRENT))),
|
||||
recall_connection_budget=int(
|
||||
os.getenv(ENV_RECALL_CONNECTION_BUDGET, str(DEFAULT_RECALL_CONNECTION_BUDGET))
|
||||
),
|
||||
mental_model_refresh_concurrency=int(
|
||||
os.getenv(ENV_MENTAL_MODEL_REFRESH_CONCURRENCY, str(DEFAULT_MENTAL_MODEL_REFRESH_CONCURRENCY))
|
||||
),
|
||||
# Optimization flags
|
||||
skip_llm_verification=os.getenv(ENV_SKIP_LLM_VERIFICATION, "false").lower() == "true",
|
||||
lazy_reranker=os.getenv(ENV_LAZY_RERANKER, "false").lower() == "true",
|
||||
# Observation thresholds
|
||||
observation_min_facts=int(os.getenv(ENV_OBSERVATION_MIN_FACTS, str(DEFAULT_OBSERVATION_MIN_FACTS))),
|
||||
observation_top_entities=int(
|
||||
os.getenv(ENV_OBSERVATION_TOP_ENTITIES, str(DEFAULT_OBSERVATION_TOP_ENTITIES))
|
||||
),
|
||||
# Retain settings
|
||||
retain_max_completion_tokens=int(
|
||||
os.getenv(ENV_RETAIN_MAX_COMPLETION_TOKENS, str(DEFAULT_RETAIN_MAX_COMPLETION_TOKENS))
|
||||
),
|
||||
retain_chunk_size=int(os.getenv(ENV_RETAIN_CHUNK_SIZE, str(DEFAULT_RETAIN_CHUNK_SIZE))),
|
||||
retain_extract_causal_links=os.getenv(
|
||||
ENV_RETAIN_EXTRACT_CAUSAL_LINKS, str(DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS)
|
||||
).lower()
|
||||
== "true",
|
||||
retain_extraction_mode=_validate_extraction_mode(
|
||||
os.getenv(ENV_RETAIN_EXTRACTION_MODE, DEFAULT_RETAIN_EXTRACTION_MODE)
|
||||
),
|
||||
retain_observations_async=os.getenv(
|
||||
ENV_RETAIN_OBSERVATIONS_ASYNC, str(DEFAULT_RETAIN_OBSERVATIONS_ASYNC)
|
||||
).lower()
|
||||
== "true",
|
||||
# Database migrations
|
||||
run_migrations_on_startup=os.getenv(ENV_RUN_MIGRATIONS_ON_STARTUP, "true").lower() == "true",
|
||||
# Database connection pool
|
||||
db_pool_min_size=int(os.getenv(ENV_DB_POOL_MIN_SIZE, str(DEFAULT_DB_POOL_MIN_SIZE))),
|
||||
db_pool_max_size=int(os.getenv(ENV_DB_POOL_MAX_SIZE, str(DEFAULT_DB_POOL_MAX_SIZE))),
|
||||
db_command_timeout=int(os.getenv(ENV_DB_COMMAND_TIMEOUT, str(DEFAULT_DB_COMMAND_TIMEOUT))),
|
||||
db_acquire_timeout=int(os.getenv(ENV_DB_ACQUIRE_TIMEOUT, str(DEFAULT_DB_ACQUIRE_TIMEOUT))),
|
||||
# Worker configuration
|
||||
worker_enabled=os.getenv(ENV_WORKER_ENABLED, str(DEFAULT_WORKER_ENABLED)).lower() == "true",
|
||||
worker_id=os.getenv(ENV_WORKER_ID) or DEFAULT_WORKER_ID,
|
||||
worker_poll_interval_ms=int(os.getenv(ENV_WORKER_POLL_INTERVAL_MS, str(DEFAULT_WORKER_POLL_INTERVAL_MS))),
|
||||
worker_max_retries=int(os.getenv(ENV_WORKER_MAX_RETRIES, str(DEFAULT_WORKER_MAX_RETRIES))),
|
||||
worker_batch_size=int(os.getenv(ENV_WORKER_BATCH_SIZE, str(DEFAULT_WORKER_BATCH_SIZE))),
|
||||
worker_http_port=int(os.getenv(ENV_WORKER_HTTP_PORT, str(DEFAULT_WORKER_HTTP_PORT))),
|
||||
# Reflect agent settings
|
||||
reflect_max_iterations=int(os.getenv(ENV_REFLECT_MAX_ITERATIONS, str(DEFAULT_REFLECT_MAX_ITERATIONS))),
|
||||
)
|
||||
|
||||
def get_llm_base_url(self) -> str:
|
||||
@@ -127,6 +454,8 @@ class HindsightConfig:
|
||||
return "https://api.groq.com/openai/v1"
|
||||
elif provider == "ollama":
|
||||
return "http://localhost:11434/v1"
|
||||
elif provider == "lmstudio":
|
||||
return "http://localhost:1234/v1"
|
||||
else:
|
||||
return ""
|
||||
|
||||
@@ -143,21 +472,59 @@ 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"
|
||||
)
|
||||
"""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"LLM: provider={self.llm_provider}, model={self.llm_model}")
|
||||
if self.retain_llm_provider or self.retain_llm_model:
|
||||
retain_provider = self.retain_llm_provider or self.llm_provider
|
||||
retain_model = self.retain_llm_model or self.llm_model
|
||||
logger.info(f"LLM (retain): provider={retain_provider}, model={retain_model}")
|
||||
if self.reflect_llm_provider or self.reflect_llm_model:
|
||||
reflect_provider = self.reflect_llm_provider or self.llm_provider
|
||||
reflect_model = self.reflect_llm_model or self.llm_model
|
||||
logger.info(f"LLM (reflect): provider={reflect_provider}, model={reflect_model}")
|
||||
logger.info(f"Embeddings: provider={self.embeddings_provider}")
|
||||
logger.info(f"Reranker: provider={self.reranker_provider}")
|
||||
logger.info(f"Graph retriever: {self.graph_retriever}")
|
||||
|
||||
|
||||
# Cached config instance
|
||||
_config_cache: HindsightConfig | None = None
|
||||
|
||||
|
||||
def get_config() -> HindsightConfig:
|
||||
"""Get the current configuration from environment variables."""
|
||||
return HindsightConfig.from_env()
|
||||
"""Get the cached configuration, loading from environment on first call."""
|
||||
global _config_cache
|
||||
if _config_cache is None:
|
||||
_config_cache = HindsightConfig.from_env()
|
||||
return _config_cache
|
||||
|
||||
|
||||
def clear_config_cache() -> None:
|
||||
"""Clear the config cache. Useful for testing or reloading config."""
|
||||
global _config_cache
|
||||
_config_cache = None
|
||||
|
||||
@@ -0,0 +1,204 @@
|
||||
"""
|
||||
Daemon mode support for Hindsight API.
|
||||
|
||||
Provides idle timeout and lockfile management for running as a background daemon.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import fcntl
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Default daemon configuration
|
||||
DEFAULT_DAEMON_PORT = 8889
|
||||
DEFAULT_IDLE_TIMEOUT = 0 # 0 = no auto-exit (hindsight-embed passes its own timeout)
|
||||
LOCKFILE_PATH = Path.home() / ".hindsight" / "daemon.lock"
|
||||
DAEMON_LOG_PATH = Path.home() / ".hindsight" / "daemon.log"
|
||||
|
||||
|
||||
class IdleTimeoutMiddleware:
|
||||
"""ASGI middleware that tracks activity and exits after idle timeout."""
|
||||
|
||||
def __init__(self, app, idle_timeout: int = DEFAULT_IDLE_TIMEOUT):
|
||||
self.app = app
|
||||
self.idle_timeout = idle_timeout
|
||||
self.last_activity = time.time()
|
||||
self._checker_task = None
|
||||
|
||||
async def __call__(self, scope, receive, send):
|
||||
# Update activity timestamp on each request
|
||||
self.last_activity = time.time()
|
||||
await self.app(scope, receive, send)
|
||||
|
||||
def start_idle_checker(self):
|
||||
"""Start the background task that checks for idle timeout."""
|
||||
self._checker_task = asyncio.create_task(self._check_idle())
|
||||
|
||||
async def _check_idle(self):
|
||||
"""Background task that exits the process after idle timeout."""
|
||||
# If idle_timeout is 0, don't auto-exit
|
||||
if self.idle_timeout <= 0:
|
||||
return
|
||||
|
||||
while True:
|
||||
await asyncio.sleep(30) # Check every 30 seconds
|
||||
idle_time = time.time() - self.last_activity
|
||||
if idle_time > self.idle_timeout:
|
||||
logger.info(f"Idle timeout reached ({self.idle_timeout}s), shutting down daemon")
|
||||
# Give a moment for any in-flight requests
|
||||
await asyncio.sleep(1)
|
||||
os._exit(0)
|
||||
|
||||
|
||||
class DaemonLock:
|
||||
"""
|
||||
File-based lock to prevent multiple daemon instances.
|
||||
|
||||
Uses fcntl.flock for atomic locking on Unix systems.
|
||||
"""
|
||||
|
||||
def __init__(self, lockfile: Path = LOCKFILE_PATH):
|
||||
self.lockfile = lockfile
|
||||
self._fd = None
|
||||
|
||||
def acquire(self) -> bool:
|
||||
"""
|
||||
Try to acquire the daemon lock.
|
||||
|
||||
Returns True if lock acquired, False if another daemon is running.
|
||||
"""
|
||||
self.lockfile.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
try:
|
||||
self._fd = open(self.lockfile, "w")
|
||||
fcntl.flock(self._fd.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||
# Write PID for debugging
|
||||
self._fd.write(str(os.getpid()))
|
||||
self._fd.flush()
|
||||
return True
|
||||
except (IOError, OSError):
|
||||
# Lock is held by another process
|
||||
if self._fd:
|
||||
self._fd.close()
|
||||
self._fd = None
|
||||
return False
|
||||
|
||||
def release(self):
|
||||
"""Release the daemon lock."""
|
||||
if self._fd:
|
||||
try:
|
||||
fcntl.flock(self._fd.fileno(), fcntl.LOCK_UN)
|
||||
self._fd.close()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
self._fd = None
|
||||
# Remove lockfile
|
||||
try:
|
||||
self.lockfile.unlink()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def is_locked(self) -> bool:
|
||||
"""Check if the lock is held by another process."""
|
||||
if not self.lockfile.exists():
|
||||
return False
|
||||
|
||||
try:
|
||||
fd = open(self.lockfile, "r")
|
||||
fcntl.flock(fd.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||
# We got the lock, so no one else has it
|
||||
fcntl.flock(fd.fileno(), fcntl.LOCK_UN)
|
||||
fd.close()
|
||||
return False
|
||||
except (IOError, OSError):
|
||||
return True
|
||||
|
||||
def get_pid(self) -> int | None:
|
||||
"""Get the PID of the daemon holding the lock."""
|
||||
if not self.lockfile.exists():
|
||||
return None
|
||||
try:
|
||||
with open(self.lockfile, "r") as f:
|
||||
return int(f.read().strip())
|
||||
except (ValueError, IOError):
|
||||
return None
|
||||
|
||||
|
||||
def daemonize():
|
||||
"""
|
||||
Fork the current process into a background daemon.
|
||||
|
||||
Uses double-fork technique to properly detach from terminal.
|
||||
"""
|
||||
# First fork
|
||||
pid = os.fork()
|
||||
if pid > 0:
|
||||
# Parent exits
|
||||
sys.exit(0)
|
||||
|
||||
# Create new session
|
||||
os.setsid()
|
||||
|
||||
# Second fork to prevent zombie processes
|
||||
pid = os.fork()
|
||||
if pid > 0:
|
||||
sys.exit(0)
|
||||
|
||||
# Redirect standard file descriptors to log file
|
||||
DAEMON_LOG_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
sys.stdout.flush()
|
||||
sys.stderr.flush()
|
||||
|
||||
# Redirect stdin to /dev/null
|
||||
with open("/dev/null", "r") as devnull:
|
||||
os.dup2(devnull.fileno(), sys.stdin.fileno())
|
||||
|
||||
# Redirect stdout/stderr to log file
|
||||
log_fd = open(DAEMON_LOG_PATH, "a")
|
||||
os.dup2(log_fd.fileno(), sys.stdout.fileno())
|
||||
os.dup2(log_fd.fileno(), sys.stderr.fileno())
|
||||
|
||||
|
||||
def check_daemon_running(port: int = DEFAULT_DAEMON_PORT) -> bool:
|
||||
"""Check if a daemon is running and responsive on the given port."""
|
||||
import socket
|
||||
|
||||
try:
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
sock.settimeout(1)
|
||||
result = sock.connect_ex(("127.0.0.1", port))
|
||||
sock.close()
|
||||
return result == 0
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def stop_daemon(port: int = DEFAULT_DAEMON_PORT) -> bool:
|
||||
"""Stop a running daemon by sending SIGTERM to the process."""
|
||||
lock = DaemonLock()
|
||||
pid = lock.get_pid()
|
||||
|
||||
if pid is None:
|
||||
return False
|
||||
|
||||
try:
|
||||
import signal
|
||||
|
||||
os.kill(pid, signal.SIGTERM)
|
||||
# Wait for process to exit
|
||||
for _ in range(50): # Wait up to 5 seconds
|
||||
time.sleep(0.1)
|
||||
try:
|
||||
os.kill(pid, 0) # Check if process exists
|
||||
except OSError:
|
||||
return True # Process exited
|
||||
return False
|
||||
except OSError:
|
||||
return False
|
||||
@@ -7,24 +7,30 @@ This package contains all the implementation details of the memory engine:
|
||||
- Supporting modules: embeddings, cross_encoder, entity_resolver, etc.
|
||||
"""
|
||||
|
||||
from .memory_engine import MemoryEngine
|
||||
from .cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
|
||||
from .db_utils import acquire_with_retry
|
||||
from .embeddings import Embeddings, LocalSTEmbeddings, RemoteTEIEmbeddings
|
||||
from .cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
|
||||
from .llm_wrapper import LLMConfig
|
||||
from .memory_engine import (
|
||||
MemoryEngine,
|
||||
UnqualifiedTableError,
|
||||
fq_table,
|
||||
get_current_schema,
|
||||
validate_sql_schema,
|
||||
)
|
||||
from .response_models import MemoryFact, RecallResult, ReflectResult
|
||||
from .search.trace import (
|
||||
SearchTrace,
|
||||
QueryInfo,
|
||||
EntryPoint,
|
||||
NodeVisit,
|
||||
WeightComponents,
|
||||
LinkInfo,
|
||||
NodeVisit,
|
||||
PruningDecision,
|
||||
SearchSummary,
|
||||
QueryInfo,
|
||||
SearchPhaseMetrics,
|
||||
SearchSummary,
|
||||
SearchTrace,
|
||||
WeightComponents,
|
||||
)
|
||||
from .search.tracer import SearchTracer
|
||||
from .llm_wrapper import LLMConfig
|
||||
from .response_models import RecallResult, ReflectResult, MemoryFact
|
||||
|
||||
__all__ = [
|
||||
"MemoryEngine",
|
||||
@@ -49,4 +55,9 @@ __all__ = [
|
||||
"RecallResult",
|
||||
"ReflectResult",
|
||||
"MemoryFact",
|
||||
# Schema safety utilities
|
||||
"fq_table",
|
||||
"get_current_schema",
|
||||
"validate_sql_schema",
|
||||
"UnqualifiedTableError",
|
||||
]
|
||||
|
||||
@@ -5,19 +5,40 @@ Provides an interface for reranking with different backends.
|
||||
|
||||
Configuration via environment variables - see hindsight_api.config for all env var names.
|
||||
"""
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Tuple, Optional
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import httpx
|
||||
|
||||
from ..config import (
|
||||
ENV_RERANKER_PROVIDER,
|
||||
ENV_RERANKER_LOCAL_MODEL,
|
||||
ENV_RERANKER_TEI_URL,
|
||||
DEFAULT_RERANKER_PROVIDER,
|
||||
DEFAULT_LITELLM_API_BASE,
|
||||
DEFAULT_RERANKER_COHERE_MODEL,
|
||||
DEFAULT_RERANKER_FLASHRANK_CACHE_DIR,
|
||||
DEFAULT_RERANKER_FLASHRANK_MODEL,
|
||||
DEFAULT_RERANKER_LITELLM_MODEL,
|
||||
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||
DEFAULT_RERANKER_LOCAL_MODEL,
|
||||
DEFAULT_RERANKER_PROVIDER,
|
||||
DEFAULT_RERANKER_TEI_BATCH_SIZE,
|
||||
DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
|
||||
ENV_COHERE_API_KEY,
|
||||
ENV_LITELLM_API_BASE,
|
||||
ENV_LITELLM_API_KEY,
|
||||
ENV_RERANKER_COHERE_BASE_URL,
|
||||
ENV_RERANKER_COHERE_MODEL,
|
||||
ENV_RERANKER_FLASHRANK_CACHE_DIR,
|
||||
ENV_RERANKER_FLASHRANK_MODEL,
|
||||
ENV_RERANKER_LITELLM_MODEL,
|
||||
ENV_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||
ENV_RERANKER_LOCAL_MODEL,
|
||||
ENV_RERANKER_PROVIDER,
|
||||
ENV_RERANKER_TEI_BATCH_SIZE,
|
||||
ENV_RERANKER_TEI_MAX_CONCURRENT,
|
||||
ENV_RERANKER_TEI_URL,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -47,7 +68,7 @@ class CrossEncoderModel(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs for relevance.
|
||||
|
||||
@@ -70,25 +91,34 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
- Fast inference (~80ms for 100 pairs on CPU)
|
||||
- Small model (80MB)
|
||||
- Trained for passage re-ranking
|
||||
|
||||
Uses a dedicated thread pool to limit concurrent CPU-bound work.
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: Optional[str] = None):
|
||||
# Shared executor across all instances (one model loaded anyway)
|
||||
_executor: ThreadPoolExecutor | None = None
|
||||
_max_concurrent: int = 4 # Limit concurrent CPU-bound reranking calls
|
||||
|
||||
def __init__(self, model_name: str | None = None, max_concurrent: int = 4):
|
||||
"""
|
||||
Initialize local SentenceTransformers cross-encoder.
|
||||
|
||||
Args:
|
||||
model_name: Name of the CrossEncoder model to use.
|
||||
Default: cross-encoder/ms-marco-MiniLM-L-6-v2
|
||||
max_concurrent: Maximum concurrent reranking calls (default: 2).
|
||||
Higher values may cause CPU thrashing under load.
|
||||
"""
|
||||
self.model_name = model_name or DEFAULT_RERANKER_LOCAL_MODEL
|
||||
self._model = None
|
||||
LocalSTCrossEncoder._max_concurrent = max_concurrent
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "local"
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Load the cross-encoder model."""
|
||||
"""Load the cross-encoder model and initialize the executor."""
|
||||
if self._model is not None:
|
||||
return
|
||||
|
||||
@@ -101,13 +131,55 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
)
|
||||
|
||||
logger.info(f"Reranker: initializing local provider with model {self.model_name}")
|
||||
self._model = CrossEncoder(self.model_name)
|
||||
logger.info("Reranker: local provider initialized")
|
||||
|
||||
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
|
||||
# Determine device and device_map based on hardware and installed packages.
|
||||
# When accelerate is installed but no GPU/MPS is available, transformers can
|
||||
# incorrectly use lazy loading (meta tensors) which fails on .to(device).
|
||||
# We use device_map="cpu" in that case to force direct CPU loading.
|
||||
import torch
|
||||
|
||||
try:
|
||||
import accelerate # type: ignore[import-not-found] # noqa: F401
|
||||
|
||||
accelerate_available = True
|
||||
except ImportError:
|
||||
accelerate_available = False
|
||||
|
||||
# Check for GPU (CUDA) or Apple Silicon (MPS)
|
||||
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
|
||||
device_map = None
|
||||
elif accelerate_available:
|
||||
device = "cpu"
|
||||
device_map = "cpu" # Force direct CPU loading to avoid meta tensors
|
||||
else:
|
||||
device = "cpu"
|
||||
device_map = None
|
||||
|
||||
self._model = CrossEncoder(
|
||||
self.model_name,
|
||||
device=device,
|
||||
model_kwargs={"low_cpu_mem_usage": False, "device_map": device_map},
|
||||
)
|
||||
|
||||
# Initialize shared executor (limited workers naturally limits concurrency)
|
||||
if LocalSTCrossEncoder._executor is None:
|
||||
LocalSTCrossEncoder._executor = ThreadPoolExecutor(
|
||||
max_workers=LocalSTCrossEncoder._max_concurrent,
|
||||
thread_name_prefix="reranker",
|
||||
)
|
||||
logger.info(f"Reranker: local provider initialized (max_concurrent={LocalSTCrossEncoder._max_concurrent})")
|
||||
else:
|
||||
logger.info("Reranker: local provider initialized (using existing executor)")
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs for relevance.
|
||||
|
||||
Uses a dedicated thread pool with limited workers to prevent CPU thrashing.
|
||||
|
||||
Args:
|
||||
pairs: List of (query, document) tuples to score
|
||||
|
||||
@@ -116,8 +188,14 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
"""
|
||||
if self._model is None:
|
||||
raise RuntimeError("Reranker not initialized. Call initialize() first.")
|
||||
scores = self._model.predict(pairs, show_progress_bar=False)
|
||||
return scores.tolist() if hasattr(scores, 'tolist') else list(scores)
|
||||
|
||||
# Use dedicated executor - limited workers naturally limits concurrency
|
||||
loop = asyncio.get_event_loop()
|
||||
scores = await loop.run_in_executor(
|
||||
LocalSTCrossEncoder._executor,
|
||||
lambda: self._model.predict(pairs, show_progress_bar=False),
|
||||
)
|
||||
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
|
||||
|
||||
|
||||
class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
@@ -128,13 +206,21 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
See: https://github.com/huggingface/text-embeddings-inference
|
||||
|
||||
Note: The TEI server must be running a cross-encoder/reranker model.
|
||||
|
||||
Requests are made in parallel with configurable batch size and max concurrency (backpressure).
|
||||
Uses a GLOBAL semaphore to limit concurrent requests across ALL recall operations.
|
||||
"""
|
||||
|
||||
# Global semaphore shared across all instances and calls to prevent thundering herd
|
||||
_global_semaphore: asyncio.Semaphore | None = None
|
||||
_global_max_concurrent: int = DEFAULT_RERANKER_TEI_MAX_CONCURRENT
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
timeout: float = 30.0,
|
||||
batch_size: int = 32,
|
||||
batch_size: int = DEFAULT_RERANKER_TEI_BATCH_SIZE,
|
||||
max_concurrent: int = DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
|
||||
max_retries: int = 3,
|
||||
retry_delay: float = 0.5,
|
||||
):
|
||||
@@ -144,75 +230,246 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
Args:
|
||||
base_url: Base URL of the TEI server (e.g., "http://localhost:8080")
|
||||
timeout: Request timeout in seconds (default: 30.0)
|
||||
batch_size: Maximum batch size for rerank requests (default: 32)
|
||||
batch_size: Maximum batch size for rerank requests (default: 128)
|
||||
max_concurrent: Maximum concurrent requests for backpressure (default: 8).
|
||||
This is a GLOBAL limit across all parallel recall operations.
|
||||
max_retries: Maximum number of retries for failed requests (default: 3)
|
||||
retry_delay: Initial delay between retries in seconds, doubles each retry (default: 0.5)
|
||||
"""
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.timeout = timeout
|
||||
self.batch_size = batch_size
|
||||
self.max_concurrent = max_concurrent
|
||||
self.max_retries = max_retries
|
||||
self.retry_delay = retry_delay
|
||||
self._client: Optional[httpx.Client] = None
|
||||
self._model_id: Optional[str] = None
|
||||
self._async_client: httpx.AsyncClient | None = None
|
||||
self._model_id: str | None = None
|
||||
|
||||
# Update global semaphore if max_concurrent changed
|
||||
if (
|
||||
RemoteTEICrossEncoder._global_semaphore is None
|
||||
or RemoteTEICrossEncoder._global_max_concurrent != max_concurrent
|
||||
):
|
||||
RemoteTEICrossEncoder._global_max_concurrent = max_concurrent
|
||||
RemoteTEICrossEncoder._global_semaphore = asyncio.Semaphore(max_concurrent)
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "tei"
|
||||
|
||||
def _request_with_retry(self, method: str, url: str, **kwargs) -> httpx.Response:
|
||||
"""Make an HTTP request with automatic retries on transient errors."""
|
||||
import time
|
||||
async def _async_request_with_retry(
|
||||
self,
|
||||
client: httpx.AsyncClient,
|
||||
semaphore: asyncio.Semaphore,
|
||||
method: str,
|
||||
url: str,
|
||||
**kwargs,
|
||||
) -> httpx.Response:
|
||||
"""Make an async HTTP request with automatic retries on transient errors and semaphore for backpressure."""
|
||||
last_error = None
|
||||
delay = self.retry_delay
|
||||
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
if method == "GET":
|
||||
response = self._client.get(url, **kwargs)
|
||||
else:
|
||||
response = self._client.post(url, **kwargs)
|
||||
response.raise_for_status()
|
||||
return response
|
||||
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
logger.warning(f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
|
||||
time.sleep(delay)
|
||||
delay *= 2 # Exponential backoff
|
||||
except httpx.HTTPStatusError as e:
|
||||
# Retry on 5xx server errors
|
||||
if e.response.status_code >= 500 and attempt < self.max_retries:
|
||||
async with semaphore:
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
if method == "GET":
|
||||
response = await client.get(url, **kwargs)
|
||||
else:
|
||||
response = await client.post(url, **kwargs)
|
||||
response.raise_for_status()
|
||||
return response
|
||||
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
|
||||
last_error = e
|
||||
logger.warning(f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
|
||||
time.sleep(delay)
|
||||
delay *= 2
|
||||
else:
|
||||
raise
|
||||
if attempt < self.max_retries:
|
||||
logger.warning(
|
||||
f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. "
|
||||
f"Retrying in {delay}s..."
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
delay *= 2 # Exponential backoff
|
||||
except httpx.HTTPStatusError as e:
|
||||
# Retry on 5xx server errors
|
||||
if e.response.status_code >= 500 and attempt < self.max_retries:
|
||||
last_error = e
|
||||
logger.warning(
|
||||
f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. "
|
||||
f"Retrying in {delay}s..."
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
delay *= 2
|
||||
else:
|
||||
raise
|
||||
|
||||
raise last_error
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Initialize the HTTP client and verify server connectivity."""
|
||||
if self._client is not None:
|
||||
if self._async_client is not None:
|
||||
return
|
||||
|
||||
logger.info(f"Reranker: initializing TEI provider at {self.base_url}")
|
||||
self._client = httpx.Client(timeout=self.timeout)
|
||||
logger.info(
|
||||
f"Reranker: initializing TEI provider at {self.base_url} "
|
||||
f"(batch_size={self.batch_size}, max_concurrent={self.max_concurrent})"
|
||||
)
|
||||
self._async_client = httpx.AsyncClient(timeout=self.timeout)
|
||||
|
||||
# Verify server is reachable and get model info
|
||||
# Use a temporary semaphore for initialization
|
||||
init_semaphore = asyncio.Semaphore(1)
|
||||
try:
|
||||
response = self._request_with_retry("GET", f"{self.base_url}/info")
|
||||
response = await self._async_request_with_retry(
|
||||
self._async_client, init_semaphore, "GET", f"{self.base_url}/info"
|
||||
)
|
||||
info = response.json()
|
||||
self._model_id = info.get("model_id", "unknown")
|
||||
logger.info(f"Reranker: TEI provider initialized (model: {self._model_id})")
|
||||
except httpx.HTTPError as e:
|
||||
self._async_client = None
|
||||
raise RuntimeError(f"Failed to connect to TEI server at {self.base_url}: {e}")
|
||||
|
||||
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
|
||||
async def _rerank_query_group(
|
||||
self,
|
||||
client: httpx.AsyncClient,
|
||||
semaphore: asyncio.Semaphore,
|
||||
query: str,
|
||||
texts: list[str],
|
||||
) -> list[tuple[int, float]]:
|
||||
"""Rerank a single query group and return list of (original_index, score) tuples."""
|
||||
try:
|
||||
response = await self._async_request_with_retry(
|
||||
client,
|
||||
semaphore,
|
||||
"POST",
|
||||
f"{self.base_url}/rerank",
|
||||
json={
|
||||
"query": query,
|
||||
"texts": texts,
|
||||
"return_text": False,
|
||||
},
|
||||
)
|
||||
results = response.json()
|
||||
# TEI returns results sorted by score descending, with original index
|
||||
return [(result["index"], result["score"]) for result in results]
|
||||
except httpx.HTTPError as e:
|
||||
raise RuntimeError(f"TEI rerank request failed: {e}")
|
||||
|
||||
async def _predict_async(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""Async implementation of predict that runs requests in parallel with backpressure."""
|
||||
if not pairs:
|
||||
return []
|
||||
|
||||
# Group all pairs by query
|
||||
query_groups: dict[str, list[tuple[int, str]]] = {}
|
||||
for idx, (query, text) in enumerate(pairs):
|
||||
if query not in query_groups:
|
||||
query_groups[query] = []
|
||||
query_groups[query].append((idx, text))
|
||||
|
||||
# Split each query group into batches
|
||||
tasks_info: list[tuple[str, list[int], list[str]]] = [] # (query, indices, texts)
|
||||
for query, indexed_texts in query_groups.items():
|
||||
indices = [idx for idx, _ in indexed_texts]
|
||||
texts = [text for _, text in indexed_texts]
|
||||
|
||||
# Split into batches
|
||||
for i in range(0, len(texts), self.batch_size):
|
||||
batch_indices = indices[i : i + self.batch_size]
|
||||
batch_texts = texts[i : i + self.batch_size]
|
||||
tasks_info.append((query, batch_indices, batch_texts))
|
||||
|
||||
# Run all requests in parallel with GLOBAL semaphore for backpressure
|
||||
# This ensures max_concurrent is respected across ALL parallel recall operations
|
||||
all_scores = [0.0] * len(pairs)
|
||||
semaphore = RemoteTEICrossEncoder._global_semaphore
|
||||
|
||||
tasks = [
|
||||
self._rerank_query_group(self._async_client, semaphore, query, texts) for query, _, texts in tasks_info
|
||||
]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
||||
# Map scores back to original positions
|
||||
for (_, indices, _), result_scores in zip(tasks_info, results):
|
||||
for original_idx_in_batch, score in result_scores:
|
||||
global_idx = indices[original_idx_in_batch]
|
||||
all_scores[global_idx] = score
|
||||
|
||||
return all_scores
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs using the remote TEI reranker.
|
||||
|
||||
Requests are made in parallel with configurable backpressure.
|
||||
|
||||
Args:
|
||||
pairs: List of (query, document) tuples to score
|
||||
|
||||
Returns:
|
||||
List of relevance scores
|
||||
"""
|
||||
if self._async_client is None:
|
||||
raise RuntimeError("Reranker not initialized. Call initialize() first.")
|
||||
|
||||
return await self._predict_async(pairs)
|
||||
|
||||
|
||||
class CohereCrossEncoder(CrossEncoderModel):
|
||||
"""
|
||||
Cohere cross-encoder implementation using the Cohere Rerank API.
|
||||
|
||||
Supports rerank-english-v3.0 and rerank-multilingual-v3.0 models.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str,
|
||||
model: str = DEFAULT_RERANKER_COHERE_MODEL,
|
||||
base_url: str | None = None,
|
||||
timeout: float = 60.0,
|
||||
):
|
||||
"""
|
||||
Initialize Cohere cross-encoder client.
|
||||
|
||||
Args:
|
||||
api_key: Cohere API key
|
||||
model: Cohere rerank model name (default: rerank-english-v3.0)
|
||||
base_url: Custom base URL for Cohere-compatible API (e.g., Azure-hosted endpoint)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.base_url = base_url
|
||||
self.timeout = timeout
|
||||
self._client = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "cohere"
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Initialize the Cohere client."""
|
||||
if self._client is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
import cohere
|
||||
except ImportError:
|
||||
raise ImportError("cohere is required for CohereCrossEncoder. Install it with: pip install cohere")
|
||||
|
||||
base_url_msg = f" at {self.base_url}" if self.base_url else ""
|
||||
logger.info(f"Reranker: initializing Cohere provider with model {self.model}{base_url_msg}")
|
||||
|
||||
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
|
||||
client_kwargs = {"api_key": self.api_key, "timeout": self.timeout}
|
||||
if self.base_url:
|
||||
client_kwargs["base_url"] = self.base_url
|
||||
self._client = cohere.Client(**client_kwargs)
|
||||
logger.info("Reranker: Cohere provider initialized")
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs using the Cohere Rerank API.
|
||||
|
||||
Args:
|
||||
pairs: List of (query, document) tuples to score
|
||||
|
||||
@@ -225,50 +482,312 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
if not pairs:
|
||||
return []
|
||||
|
||||
all_scores = []
|
||||
# Run sync Cohere API calls in thread pool
|
||||
loop = asyncio.get_event_loop()
|
||||
return await loop.run_in_executor(None, self._predict_sync, pairs)
|
||||
|
||||
# Process in batches
|
||||
for i in range(0, len(pairs), self.batch_size):
|
||||
batch = pairs[i:i + self.batch_size]
|
||||
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""Synchronous predict implementation for Cohere API."""
|
||||
# Group pairs by query for efficient batching
|
||||
# Cohere rerank expects one query with multiple documents
|
||||
query_groups: dict[str, list[tuple[int, str]]] = {}
|
||||
for idx, (query, text) in enumerate(pairs):
|
||||
if query not in query_groups:
|
||||
query_groups[query] = []
|
||||
query_groups[query].append((idx, text))
|
||||
|
||||
# TEI rerank endpoint expects query and texts separately
|
||||
# All pairs in a batch should have the same query for optimal performance
|
||||
# but we handle mixed queries by making separate requests per unique query
|
||||
query_groups: dict[str, list[tuple[int, str]]] = {}
|
||||
for idx, (query, text) in enumerate(batch):
|
||||
if query not in query_groups:
|
||||
query_groups[query] = []
|
||||
query_groups[query].append((idx, text))
|
||||
all_scores = [0.0] * len(pairs)
|
||||
|
||||
batch_scores = [0.0] * len(batch)
|
||||
for query, indexed_texts in query_groups.items():
|
||||
texts = [text for _, text in indexed_texts]
|
||||
indices = [idx for idx, _ in indexed_texts]
|
||||
|
||||
for query, indexed_texts in query_groups.items():
|
||||
texts = [text for _, text in indexed_texts]
|
||||
indices = [idx for idx, _ in indexed_texts]
|
||||
response = self._client.rerank(
|
||||
query=query,
|
||||
documents=texts,
|
||||
model=self.model,
|
||||
return_documents=False,
|
||||
)
|
||||
|
||||
try:
|
||||
response = self._request_with_retry(
|
||||
"POST",
|
||||
f"{self.base_url}/rerank",
|
||||
json={
|
||||
"query": query,
|
||||
"texts": texts,
|
||||
"return_text": False,
|
||||
},
|
||||
)
|
||||
results = response.json()
|
||||
# Map scores back to original positions
|
||||
for result in response.results:
|
||||
original_idx = result.index
|
||||
score = result.relevance_score
|
||||
all_scores[indices[original_idx]] = score
|
||||
|
||||
# TEI returns results sorted by score descending, with original index
|
||||
for result in results:
|
||||
original_idx = result["index"]
|
||||
score = result["score"]
|
||||
# Map back to batch position
|
||||
batch_scores[indices[original_idx]] = score
|
||||
return all_scores
|
||||
|
||||
except httpx.HTTPError as e:
|
||||
raise RuntimeError(f"TEI rerank request failed: {e}")
|
||||
|
||||
all_scores.extend(batch_scores)
|
||||
class RRFPassthroughCrossEncoder(CrossEncoderModel):
|
||||
"""
|
||||
Passthrough cross-encoder that preserves RRF scores without neural reranking.
|
||||
|
||||
This is useful for:
|
||||
- Testing retrieval quality without reranking overhead
|
||||
- Deployments where reranking latency is unacceptable
|
||||
- Debugging to isolate retrieval vs reranking issues
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize RRF passthrough cross-encoder."""
|
||||
pass
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "rrf"
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""No initialization needed."""
|
||||
logger.info("Reranker: RRF passthrough provider initialized (neural reranking disabled)")
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Return neutral scores - actual ranking uses RRF scores from retrieval.
|
||||
|
||||
Args:
|
||||
pairs: List of (query, document) tuples (ignored)
|
||||
|
||||
Returns:
|
||||
List of 0.5 scores (neutral, lets RRF scores dominate)
|
||||
"""
|
||||
# Return neutral scores so RRF ranking is preserved
|
||||
return [0.5] * len(pairs)
|
||||
|
||||
|
||||
class FlashRankCrossEncoder(CrossEncoderModel):
|
||||
"""
|
||||
FlashRank cross-encoder implementation.
|
||||
|
||||
FlashRank is an ultra-lite reranking library that runs on CPU without
|
||||
requiring PyTorch or Transformers. It's ideal for serverless deployments
|
||||
with minimal cold-start overhead.
|
||||
|
||||
Available models:
|
||||
- ms-marco-TinyBERT-L-2-v2: Fastest, ~4MB
|
||||
- ms-marco-MiniLM-L-12-v2: Best quality, ~34MB (default)
|
||||
- rank-T5-flan: Best zero-shot, ~110MB
|
||||
- ms-marco-MultiBERT-L-12: Multi-lingual, ~150MB
|
||||
"""
|
||||
|
||||
# Shared executor for CPU-bound reranking
|
||||
_executor: ThreadPoolExecutor | None = None
|
||||
_max_concurrent: int = 4
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str | None = None,
|
||||
cache_dir: str | None = None,
|
||||
max_length: int = 512,
|
||||
max_concurrent: int = 4,
|
||||
):
|
||||
"""
|
||||
Initialize FlashRank cross-encoder.
|
||||
|
||||
Args:
|
||||
model_name: FlashRank model name. Default: ms-marco-MiniLM-L-12-v2
|
||||
cache_dir: Directory to cache downloaded models. Default: system cache
|
||||
max_length: Maximum sequence length for reranking. Default: 512
|
||||
max_concurrent: Maximum concurrent reranking calls. Default: 4
|
||||
"""
|
||||
self.model_name = model_name or DEFAULT_RERANKER_FLASHRANK_MODEL
|
||||
self.cache_dir = cache_dir or DEFAULT_RERANKER_FLASHRANK_CACHE_DIR
|
||||
self.max_length = max_length
|
||||
self._ranker = None
|
||||
FlashRankCrossEncoder._max_concurrent = max_concurrent
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "flashrank"
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Load the FlashRank model."""
|
||||
if self._ranker is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
from flashrank import Ranker # type: ignore[import-untyped]
|
||||
except ImportError:
|
||||
raise ImportError("flashrank is required for FlashRankCrossEncoder. Install it with: pip install flashrank")
|
||||
|
||||
logger.info(f"Reranker: initializing FlashRank provider with model {self.model_name}")
|
||||
|
||||
# Initialize ranker with optional cache directory
|
||||
ranker_kwargs = {"model_name": self.model_name, "max_length": self.max_length}
|
||||
if self.cache_dir:
|
||||
ranker_kwargs["cache_dir"] = self.cache_dir
|
||||
|
||||
self._ranker = Ranker(**ranker_kwargs)
|
||||
|
||||
# Initialize shared executor
|
||||
if FlashRankCrossEncoder._executor is None:
|
||||
FlashRankCrossEncoder._executor = ThreadPoolExecutor(
|
||||
max_workers=FlashRankCrossEncoder._max_concurrent,
|
||||
thread_name_prefix="flashrank",
|
||||
)
|
||||
logger.info(
|
||||
f"Reranker: FlashRank provider initialized (max_concurrent={FlashRankCrossEncoder._max_concurrent})"
|
||||
)
|
||||
else:
|
||||
logger.info("Reranker: FlashRank provider initialized (using existing executor)")
|
||||
|
||||
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""Synchronous predict - processes each query group."""
|
||||
from flashrank import RerankRequest # type: ignore[import-untyped]
|
||||
|
||||
if not pairs:
|
||||
return []
|
||||
|
||||
# Group pairs by query
|
||||
query_groups: dict[str, list[tuple[int, str]]] = {}
|
||||
for idx, (query, text) in enumerate(pairs):
|
||||
if query not in query_groups:
|
||||
query_groups[query] = []
|
||||
query_groups[query].append((idx, text))
|
||||
|
||||
all_scores = [0.0] * len(pairs)
|
||||
|
||||
for query, indexed_texts in query_groups.items():
|
||||
# Build passages list for FlashRank
|
||||
passages = [{"id": i, "text": text} for i, (_, text) in enumerate(indexed_texts)]
|
||||
global_indices = [idx for idx, _ in indexed_texts]
|
||||
|
||||
# Create rerank request
|
||||
request = RerankRequest(query=query, passages=passages)
|
||||
results = self._ranker.rerank(request)
|
||||
|
||||
# Map scores back to original positions
|
||||
for result in results:
|
||||
local_idx = result["id"]
|
||||
score = result["score"]
|
||||
global_idx = global_indices[local_idx]
|
||||
all_scores[global_idx] = score
|
||||
|
||||
return all_scores
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs using FlashRank.
|
||||
|
||||
Args:
|
||||
pairs: List of (query, document) tuples to score
|
||||
|
||||
Returns:
|
||||
List of relevance scores (higher = more relevant)
|
||||
"""
|
||||
if self._ranker is None:
|
||||
raise RuntimeError("Reranker not initialized. Call initialize() first.")
|
||||
|
||||
# Run in thread pool to avoid blocking event loop
|
||||
loop = asyncio.get_event_loop()
|
||||
return await loop.run_in_executor(FlashRankCrossEncoder._executor, self._predict_sync, pairs)
|
||||
|
||||
|
||||
class LiteLLMCrossEncoder(CrossEncoderModel):
|
||||
"""
|
||||
LiteLLM cross-encoder implementation using LiteLLM proxy's /rerank endpoint.
|
||||
|
||||
LiteLLM provides a unified interface for multiple reranking providers via
|
||||
the Cohere-compatible /rerank endpoint.
|
||||
See: https://docs.litellm.ai/docs/rerank
|
||||
|
||||
Supported providers via LiteLLM:
|
||||
- Cohere (rerank-english-v3.0, etc.) - prefix with cohere/
|
||||
- Together AI - prefix with together_ai/
|
||||
- Azure AI - prefix with azure_ai/
|
||||
- Jina AI - prefix with jina_ai/
|
||||
- AWS Bedrock - prefix with bedrock/
|
||||
- Voyage AI - prefix with voyage/
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_base: str = DEFAULT_LITELLM_API_BASE,
|
||||
api_key: str | None = None,
|
||||
model: str = DEFAULT_RERANKER_LITELLM_MODEL,
|
||||
timeout: float = 60.0,
|
||||
):
|
||||
"""
|
||||
Initialize LiteLLM cross-encoder client.
|
||||
|
||||
Args:
|
||||
api_base: Base URL of the LiteLLM proxy (default: http://localhost:4000)
|
||||
api_key: API key for the LiteLLM proxy (optional, depends on proxy config)
|
||||
model: Reranking model name (default: cohere/rerank-english-v3.0)
|
||||
Use provider prefix (e.g., cohere/, together_ai/, voyage/)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
"""
|
||||
self.api_base = api_base.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self._async_client: httpx.AsyncClient | None = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "litellm"
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Initialize the async HTTP client."""
|
||||
if self._async_client is not None:
|
||||
return
|
||||
|
||||
logger.info(f"Reranker: initializing LiteLLM provider at {self.api_base} with model {self.model}")
|
||||
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
|
||||
self._async_client = httpx.AsyncClient(timeout=self.timeout, headers=headers)
|
||||
logger.info("Reranker: LiteLLM provider initialized")
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs using the LiteLLM proxy's /rerank endpoint.
|
||||
|
||||
Args:
|
||||
pairs: List of (query, document) tuples to score
|
||||
|
||||
Returns:
|
||||
List of relevance scores
|
||||
"""
|
||||
if self._async_client is None:
|
||||
raise RuntimeError("Reranker not initialized. Call initialize() first.")
|
||||
|
||||
if not pairs:
|
||||
return []
|
||||
|
||||
# Group pairs by query (LiteLLM rerank expects one query with multiple documents)
|
||||
query_groups: dict[str, list[tuple[int, str]]] = {}
|
||||
for idx, (query, text) in enumerate(pairs):
|
||||
if query not in query_groups:
|
||||
query_groups[query] = []
|
||||
query_groups[query].append((idx, text))
|
||||
|
||||
all_scores = [0.0] * len(pairs)
|
||||
|
||||
for query, indexed_texts in query_groups.items():
|
||||
texts = [text for _, text in indexed_texts]
|
||||
indices = [idx for idx, _ in indexed_texts]
|
||||
|
||||
# LiteLLM /rerank follows Cohere API format
|
||||
response = await self._async_client.post(
|
||||
f"{self.api_base}/rerank",
|
||||
json={
|
||||
"model": self.model,
|
||||
"query": query,
|
||||
"documents": texts,
|
||||
"top_n": len(texts), # Return all scores
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
# Map scores back to original positions
|
||||
# Response format: {"results": [{"index": 0, "relevance_score": 0.9}, ...]}
|
||||
for item in result.get("results", []):
|
||||
original_idx = item["index"]
|
||||
score = item.get("relevance_score", item.get("score", 0.0))
|
||||
all_scores[indices[original_idx]] = score
|
||||
|
||||
return all_scores
|
||||
|
||||
@@ -287,15 +806,36 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
if provider == "tei":
|
||||
url = os.environ.get(ENV_RERANKER_TEI_URL)
|
||||
if not url:
|
||||
raise ValueError(
|
||||
f"{ENV_RERANKER_TEI_URL} is required when {ENV_RERANKER_PROVIDER} is 'tei'"
|
||||
)
|
||||
return RemoteTEICrossEncoder(base_url=url)
|
||||
raise ValueError(f"{ENV_RERANKER_TEI_URL} is required when {ENV_RERANKER_PROVIDER} is 'tei'")
|
||||
batch_size = int(os.environ.get(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE)))
|
||||
max_concurrent = int(os.environ.get(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT)))
|
||||
return RemoteTEICrossEncoder(base_url=url, batch_size=batch_size, max_concurrent=max_concurrent)
|
||||
elif provider == "local":
|
||||
model = os.environ.get(ENV_RERANKER_LOCAL_MODEL)
|
||||
model_name = model or DEFAULT_RERANKER_LOCAL_MODEL
|
||||
return LocalSTCrossEncoder(model_name=model_name)
|
||||
max_concurrent = int(
|
||||
os.environ.get(ENV_RERANKER_LOCAL_MAX_CONCURRENT, str(DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT))
|
||||
)
|
||||
return LocalSTCrossEncoder(model_name=model_name, max_concurrent=max_concurrent)
|
||||
elif provider == "cohere":
|
||||
api_key = os.environ.get(ENV_COHERE_API_KEY)
|
||||
if not api_key:
|
||||
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_RERANKER_PROVIDER} is 'cohere'")
|
||||
model = os.environ.get(ENV_RERANKER_COHERE_MODEL, DEFAULT_RERANKER_COHERE_MODEL)
|
||||
base_url = os.environ.get(ENV_RERANKER_COHERE_BASE_URL) or None
|
||||
return CohereCrossEncoder(api_key=api_key, model=model, base_url=base_url)
|
||||
elif provider == "flashrank":
|
||||
model = os.environ.get(ENV_RERANKER_FLASHRANK_MODEL, DEFAULT_RERANKER_FLASHRANK_MODEL)
|
||||
cache_dir = os.environ.get(ENV_RERANKER_FLASHRANK_CACHE_DIR, DEFAULT_RERANKER_FLASHRANK_CACHE_DIR)
|
||||
return FlashRankCrossEncoder(model_name=model, cache_dir=cache_dir)
|
||||
elif provider == "litellm":
|
||||
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
|
||||
api_key = os.environ.get(ENV_LITELLM_API_KEY)
|
||||
model = os.environ.get(ENV_RERANKER_LITELLM_MODEL, DEFAULT_RERANKER_LITELLM_MODEL)
|
||||
return LiteLLMCrossEncoder(api_base=api_base, api_key=api_key, model=model)
|
||||
elif provider == "rrf":
|
||||
return RRFPassthroughCrossEncoder()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei'"
|
||||
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'flashrank', 'litellm', 'rrf'"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,284 @@
|
||||
"""
|
||||
Database connection budget management.
|
||||
|
||||
Limits concurrent database connections per operation to prevent
|
||||
a single operation (e.g., recall with parallel queries) from
|
||||
exhausting the connection pool.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, AsyncIterator
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import asyncpg
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class OperationBudget:
|
||||
"""
|
||||
Tracks connection budget for a single operation.
|
||||
|
||||
Each operation gets a semaphore limiting its concurrent connections.
|
||||
"""
|
||||
|
||||
operation_id: str
|
||||
max_connections: int
|
||||
semaphore: asyncio.Semaphore = field(init=False)
|
||||
active_count: int = field(default=0, init=False)
|
||||
|
||||
def __post_init__(self):
|
||||
self.semaphore = asyncio.Semaphore(self.max_connections)
|
||||
|
||||
|
||||
class ConnectionBudgetManager:
|
||||
"""
|
||||
Manages per-operation connection budgets.
|
||||
|
||||
Usage:
|
||||
manager = ConnectionBudgetManager(default_budget=4)
|
||||
|
||||
# Start an operation
|
||||
async with manager.operation(max_connections=2) as op:
|
||||
# Acquire connections within the budget
|
||||
async with op.acquire(pool) as conn:
|
||||
await conn.fetch(...)
|
||||
|
||||
# Multiple connections respect the budget
|
||||
async with op.acquire(pool) as conn1, op.acquire(pool) as conn2:
|
||||
# At most 2 concurrent connections for this operation
|
||||
...
|
||||
"""
|
||||
|
||||
def __init__(self, default_budget: int = 4):
|
||||
"""
|
||||
Initialize the budget manager.
|
||||
|
||||
Args:
|
||||
default_budget: Default max connections per operation
|
||||
"""
|
||||
self.default_budget = default_budget
|
||||
self._operations: dict[str, OperationBudget] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
@asynccontextmanager
|
||||
async def operation(
|
||||
self,
|
||||
max_connections: int | None = None,
|
||||
operation_id: str | None = None,
|
||||
) -> AsyncIterator["BudgetedOperation"]:
|
||||
"""
|
||||
Create a budgeted operation context.
|
||||
|
||||
Args:
|
||||
max_connections: Max concurrent connections for this operation.
|
||||
Defaults to manager's default_budget.
|
||||
operation_id: Optional custom operation ID. Auto-generated if not provided.
|
||||
|
||||
Yields:
|
||||
BudgetedOperation context for acquiring connections
|
||||
"""
|
||||
op_id = operation_id or f"op-{uuid.uuid4().hex[:12]}"
|
||||
budget = max_connections or self.default_budget
|
||||
|
||||
async with self._lock:
|
||||
if op_id in self._operations:
|
||||
raise ValueError(f"Operation {op_id} already exists")
|
||||
self._operations[op_id] = OperationBudget(op_id, budget)
|
||||
|
||||
try:
|
||||
yield BudgetedOperation(self, op_id)
|
||||
finally:
|
||||
async with self._lock:
|
||||
self._operations.pop(op_id, None)
|
||||
|
||||
def _get_budget(self, operation_id: str) -> OperationBudget:
|
||||
"""Get budget for an operation (internal use)."""
|
||||
budget = self._operations.get(operation_id)
|
||||
if not budget:
|
||||
raise ValueError(f"Operation {operation_id} not found")
|
||||
return budget
|
||||
|
||||
|
||||
class BudgetedOperation:
|
||||
"""
|
||||
A single operation with connection budget.
|
||||
|
||||
Provides methods to acquire connections within the budget.
|
||||
"""
|
||||
|
||||
def __init__(self, manager: ConnectionBudgetManager, operation_id: str):
|
||||
self._manager = manager
|
||||
self.operation_id = operation_id
|
||||
|
||||
@property
|
||||
def budget(self) -> OperationBudget:
|
||||
"""Get the budget for this operation."""
|
||||
return self._manager._get_budget(self.operation_id)
|
||||
|
||||
@asynccontextmanager
|
||||
async def acquire(self, pool: "asyncpg.Pool") -> AsyncIterator["asyncpg.Connection"]:
|
||||
"""
|
||||
Acquire a connection within the operation's budget.
|
||||
|
||||
Blocks if the operation has reached its connection limit.
|
||||
|
||||
Args:
|
||||
pool: asyncpg connection pool
|
||||
|
||||
Yields:
|
||||
Database connection
|
||||
"""
|
||||
budget = self.budget
|
||||
async with budget.semaphore:
|
||||
budget.active_count += 1
|
||||
conn = await pool.acquire()
|
||||
try:
|
||||
yield conn
|
||||
finally:
|
||||
budget.active_count -= 1
|
||||
await pool.release(conn)
|
||||
|
||||
def wrap_pool(self, pool: "asyncpg.Pool") -> "BudgetedPool":
|
||||
"""
|
||||
Wrap a pool with this operation's budget.
|
||||
|
||||
The returned BudgetedPool can be passed to functions expecting a pool,
|
||||
and all acquire() calls will be limited by this operation's budget.
|
||||
|
||||
Args:
|
||||
pool: asyncpg connection pool to wrap
|
||||
|
||||
Returns:
|
||||
BudgetedPool that limits connections to this operation's budget
|
||||
"""
|
||||
return BudgetedPool(pool, self)
|
||||
|
||||
async def acquire_many(
|
||||
self,
|
||||
pool: "asyncpg.Pool",
|
||||
count: int,
|
||||
) -> AsyncIterator[list["asyncpg.Connection"]]:
|
||||
"""
|
||||
Acquire multiple connections within the budget.
|
||||
|
||||
Note: This acquires connections sequentially to respect the budget.
|
||||
For parallel acquisition, use multiple acquire() calls with asyncio.gather().
|
||||
|
||||
Args:
|
||||
pool: asyncpg connection pool
|
||||
count: Number of connections to acquire
|
||||
|
||||
Yields:
|
||||
List of database connections
|
||||
"""
|
||||
connections = []
|
||||
try:
|
||||
for _ in range(count):
|
||||
conn = await pool.acquire()
|
||||
connections.append(conn)
|
||||
yield connections
|
||||
finally:
|
||||
for conn in connections:
|
||||
await pool.release(conn)
|
||||
|
||||
|
||||
# Global default manager instance
|
||||
_default_manager: ConnectionBudgetManager | None = None
|
||||
|
||||
|
||||
def get_budget_manager(default_budget: int = 4) -> ConnectionBudgetManager:
|
||||
"""
|
||||
Get or create the global budget manager.
|
||||
|
||||
Args:
|
||||
default_budget: Default max connections per operation
|
||||
|
||||
Returns:
|
||||
Global ConnectionBudgetManager instance
|
||||
"""
|
||||
global _default_manager
|
||||
if _default_manager is None:
|
||||
_default_manager = ConnectionBudgetManager(default_budget=default_budget)
|
||||
return _default_manager
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def budgeted_operation(
|
||||
max_connections: int | None = None,
|
||||
operation_id: str | None = None,
|
||||
default_budget: int = 4,
|
||||
) -> AsyncIterator[BudgetedOperation]:
|
||||
"""
|
||||
Convenience function to create a budgeted operation.
|
||||
|
||||
Args:
|
||||
max_connections: Max concurrent connections for this operation
|
||||
operation_id: Optional custom operation ID
|
||||
default_budget: Default budget if manager not yet created
|
||||
|
||||
Yields:
|
||||
BudgetedOperation context
|
||||
|
||||
Example:
|
||||
async with budgeted_operation(max_connections=2) as op:
|
||||
async with op.acquire(pool) as conn:
|
||||
await conn.fetch(...)
|
||||
"""
|
||||
manager = get_budget_manager(default_budget)
|
||||
async with manager.operation(max_connections, operation_id) as op:
|
||||
yield op
|
||||
|
||||
|
||||
class BudgetedPool:
|
||||
"""
|
||||
A pool wrapper that limits concurrent connection acquisitions.
|
||||
|
||||
This can be passed to functions expecting a pool, and acquire()
|
||||
calls will be limited by the budget semaphore.
|
||||
|
||||
Usage:
|
||||
async with budgeted_operation(max_connections=4) as op:
|
||||
budgeted_pool = op.wrap_pool(pool)
|
||||
# Pass budgeted_pool to functions that expect a pool
|
||||
await some_function(budgeted_pool, ...)
|
||||
"""
|
||||
|
||||
def __init__(self, pool: "asyncpg.Pool", operation: BudgetedOperation):
|
||||
self._pool = pool
|
||||
self._operation = operation
|
||||
|
||||
async def acquire(self) -> "asyncpg.Connection":
|
||||
"""
|
||||
Acquire a connection within the budget.
|
||||
|
||||
Note: Caller must release the connection when done.
|
||||
Prefer using as context manager via acquire_with_retry or op.acquire().
|
||||
"""
|
||||
budget = self._operation.budget
|
||||
await budget.semaphore.acquire()
|
||||
budget.active_count += 1
|
||||
try:
|
||||
return await self._pool.acquire()
|
||||
except Exception:
|
||||
budget.active_count -= 1
|
||||
budget.semaphore.release()
|
||||
raise
|
||||
|
||||
async def release(self, conn: "asyncpg.Connection") -> None:
|
||||
"""Release a connection back to the pool."""
|
||||
budget = self._operation.budget
|
||||
try:
|
||||
await self._pool.release(conn)
|
||||
finally:
|
||||
budget.active_count -= 1
|
||||
budget.semaphore.release()
|
||||
|
||||
def __getattr__(self, name):
|
||||
"""Proxy other attributes to the underlying pool."""
|
||||
return getattr(self._pool, name)
|
||||
@@ -1,9 +1,11 @@
|
||||
"""
|
||||
Database utility functions for connection management with retry logic.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import asyncpg
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -54,16 +56,14 @@ async def retry_with_backoff(
|
||||
except retryable_exceptions as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
delay = min(base_delay * (2 ** attempt), max_delay)
|
||||
delay = min(base_delay * (2**attempt), max_delay)
|
||||
logger.warning(
|
||||
f"Database operation failed (attempt {attempt + 1}/{max_retries + 1}): {e}. "
|
||||
f"Retrying in {delay:.1f}s..."
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
else:
|
||||
logger.error(
|
||||
f"Database operation failed after {max_retries + 1} attempts: {e}"
|
||||
)
|
||||
logger.error(f"Database operation failed after {max_retries + 1} attempts: {e}")
|
||||
raise last_exception
|
||||
|
||||
|
||||
@@ -83,10 +83,22 @@ async def acquire_with_retry(pool: asyncpg.Pool, max_retries: int = DEFAULT_MAX_
|
||||
Yields:
|
||||
An asyncpg connection
|
||||
"""
|
||||
import time
|
||||
|
||||
start = time.time()
|
||||
|
||||
async def acquire():
|
||||
return await pool.acquire()
|
||||
|
||||
conn = await retry_with_backoff(acquire, max_retries=max_retries)
|
||||
acquire_time = time.time() - start
|
||||
|
||||
# Log slow connection acquisitions (indicates pool contention)
|
||||
if acquire_time > 0.05: # 50ms threshold
|
||||
pool_size = pool.get_size()
|
||||
pool_free = pool.get_idle_size()
|
||||
logger.warning(f"[DB POOL] Slow acquire: {acquire_time:.3f}s | size={pool_size}, idle={pool_free}")
|
||||
|
||||
try:
|
||||
yield conn
|
||||
finally:
|
||||
|
||||
@@ -3,25 +3,38 @@ Embeddings abstraction for the memory system.
|
||||
|
||||
Provides an interface for generating embeddings with different backends.
|
||||
|
||||
IMPORTANT: All embeddings must produce 384-dimensional vectors to match
|
||||
the database schema (pgvector column defined as vector(384)).
|
||||
The embedding dimension is auto-detected from the model at initialization.
|
||||
The database schema is automatically adjusted to match the model's dimension.
|
||||
|
||||
Configuration via environment variables - see hindsight_api.config for all env var names.
|
||||
"""
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Optional
|
||||
|
||||
import logging
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import httpx
|
||||
|
||||
from ..config import (
|
||||
ENV_EMBEDDINGS_PROVIDER,
|
||||
ENV_EMBEDDINGS_LOCAL_MODEL,
|
||||
ENV_EMBEDDINGS_TEI_URL,
|
||||
DEFAULT_EMBEDDINGS_PROVIDER,
|
||||
DEFAULT_EMBEDDINGS_COHERE_MODEL,
|
||||
DEFAULT_EMBEDDINGS_LITELLM_MODEL,
|
||||
DEFAULT_EMBEDDINGS_LOCAL_MODEL,
|
||||
EMBEDDING_DIMENSION,
|
||||
DEFAULT_EMBEDDINGS_OPENAI_MODEL,
|
||||
DEFAULT_EMBEDDINGS_PROVIDER,
|
||||
DEFAULT_LITELLM_API_BASE,
|
||||
ENV_COHERE_API_KEY,
|
||||
ENV_EMBEDDINGS_COHERE_BASE_URL,
|
||||
ENV_EMBEDDINGS_COHERE_MODEL,
|
||||
ENV_EMBEDDINGS_LITELLM_MODEL,
|
||||
ENV_EMBEDDINGS_LOCAL_MODEL,
|
||||
ENV_EMBEDDINGS_OPENAI_API_KEY,
|
||||
ENV_EMBEDDINGS_OPENAI_BASE_URL,
|
||||
ENV_EMBEDDINGS_OPENAI_MODEL,
|
||||
ENV_EMBEDDINGS_PROVIDER,
|
||||
ENV_EMBEDDINGS_TEI_URL,
|
||||
ENV_LITELLM_API_BASE,
|
||||
ENV_LITELLM_API_KEY,
|
||||
ENV_LLM_API_KEY,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -31,8 +44,8 @@ class Embeddings(ABC):
|
||||
"""
|
||||
Abstract base class for embedding generation.
|
||||
|
||||
All implementations MUST generate 384-dimensional embeddings to match
|
||||
the database schema.
|
||||
The embedding dimension is determined by the model and detected at initialization.
|
||||
The database schema is automatically adjusted to match the model's dimension.
|
||||
"""
|
||||
|
||||
@property
|
||||
@@ -41,6 +54,12 @@ class Embeddings(ABC):
|
||||
"""Return a human-readable name for this provider (e.g., 'local', 'tei')."""
|
||||
pass
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def dimension(self) -> int:
|
||||
"""Return the embedding dimension produced by this model."""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def initialize(self) -> None:
|
||||
"""
|
||||
@@ -52,15 +71,15 @@ class Embeddings(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def encode(self, texts: List[str]) -> List[List[float]]:
|
||||
def encode(self, texts: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generate 384-dimensional embeddings for a list of texts.
|
||||
Generate embeddings for a list of texts.
|
||||
|
||||
Args:
|
||||
texts: List of text strings to encode
|
||||
|
||||
Returns:
|
||||
List of 384-dimensional embedding vectors (each is a list of floats)
|
||||
List of embedding vectors (each is a list of floats)
|
||||
"""
|
||||
pass
|
||||
|
||||
@@ -70,27 +89,31 @@ class LocalSTEmbeddings(Embeddings):
|
||||
Local embeddings implementation using SentenceTransformers.
|
||||
|
||||
Call initialize() during startup to load the model and avoid cold starts.
|
||||
|
||||
Default model is BAAI/bge-small-en-v1.5 which produces 384-dimensional
|
||||
embeddings matching the database schema.
|
||||
The embedding dimension is auto-detected from the model.
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: Optional[str] = None):
|
||||
def __init__(self, model_name: str | None = None):
|
||||
"""
|
||||
Initialize local SentenceTransformers embeddings.
|
||||
|
||||
Args:
|
||||
model_name: Name of the SentenceTransformer model to use.
|
||||
Must produce 384-dimensional embeddings.
|
||||
Default: BAAI/bge-small-en-v1.5
|
||||
"""
|
||||
self.model_name = model_name or DEFAULT_EMBEDDINGS_LOCAL_MODEL
|
||||
self._model = None
|
||||
self._dimension: int | None = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "local"
|
||||
|
||||
@property
|
||||
def dimension(self) -> int:
|
||||
if self._dimension is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
return self._dimension
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Load the embedding model."""
|
||||
if self._model is not None:
|
||||
@@ -105,33 +128,51 @@ class LocalSTEmbeddings(Embeddings):
|
||||
)
|
||||
|
||||
logger.info(f"Embeddings: initializing local provider with model {self.model_name}")
|
||||
# Disable lazy loading (meta tensors) which causes issues with newer transformers/accelerate
|
||||
# Setting low_cpu_mem_usage=False and device_map=None ensures tensors are fully materialized
|
||||
|
||||
# Determine device and device_map based on hardware and installed packages.
|
||||
# When accelerate is installed but no GPU/MPS is available, transformers can
|
||||
# incorrectly use lazy loading (meta tensors) which fails on .to(device).
|
||||
# We use device_map="cpu" in that case to force direct CPU loading.
|
||||
import torch
|
||||
|
||||
try:
|
||||
import accelerate # type: ignore[import-not-found] # noqa: F401
|
||||
|
||||
accelerate_available = True
|
||||
except ImportError:
|
||||
accelerate_available = False
|
||||
|
||||
# Check for GPU (CUDA) or Apple Silicon (MPS)
|
||||
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
|
||||
device_map = None
|
||||
elif accelerate_available:
|
||||
device = "cpu"
|
||||
device_map = "cpu" # Force direct CPU loading to avoid meta tensors
|
||||
else:
|
||||
device = "cpu"
|
||||
device_map = None
|
||||
|
||||
self._model = SentenceTransformer(
|
||||
self.model_name,
|
||||
model_kwargs={"low_cpu_mem_usage": False, "device_map": None},
|
||||
device=device,
|
||||
model_kwargs={"low_cpu_mem_usage": False, "device_map": device_map},
|
||||
)
|
||||
|
||||
# Validate dimension matches database schema
|
||||
model_dim = self._model.get_sentence_embedding_dimension()
|
||||
if model_dim != EMBEDDING_DIMENSION:
|
||||
raise ValueError(
|
||||
f"Model {self.model_name} produces {model_dim}-dimensional embeddings, "
|
||||
f"but database schema requires {EMBEDDING_DIMENSION} dimensions. "
|
||||
f"Use a model that produces {EMBEDDING_DIMENSION}-dimensional embeddings."
|
||||
)
|
||||
self._dimension = self._model.get_sentence_embedding_dimension()
|
||||
logger.info(f"Embeddings: local provider initialized (dim: {self._dimension})")
|
||||
|
||||
logger.info(f"Embeddings: local provider initialized (dim: {model_dim})")
|
||||
|
||||
def encode(self, texts: List[str]) -> List[List[float]]:
|
||||
def encode(self, texts: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generate 384-dimensional embeddings for a list of texts.
|
||||
Generate embeddings for a list of texts.
|
||||
|
||||
Args:
|
||||
texts: List of text strings to encode
|
||||
|
||||
Returns:
|
||||
List of 384-dimensional embedding vectors
|
||||
List of embedding vectors
|
||||
"""
|
||||
if self._model is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
@@ -146,7 +187,7 @@ class RemoteTEIEmbeddings(Embeddings):
|
||||
TEI provides a high-performance inference server for embedding models.
|
||||
See: https://github.com/huggingface/text-embeddings-inference
|
||||
|
||||
The server should be running a model that produces 384-dimensional embeddings.
|
||||
The embedding dimension is auto-detected from the server at initialization.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -172,16 +213,24 @@ class RemoteTEIEmbeddings(Embeddings):
|
||||
self.batch_size = batch_size
|
||||
self.max_retries = max_retries
|
||||
self.retry_delay = retry_delay
|
||||
self._client: Optional[httpx.Client] = None
|
||||
self._model_id: Optional[str] = None
|
||||
self._client: httpx.Client | None = None
|
||||
self._model_id: str | None = None
|
||||
self._dimension: int | None = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "tei"
|
||||
|
||||
@property
|
||||
def dimension(self) -> int:
|
||||
if self._dimension is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
return self._dimension
|
||||
|
||||
def _request_with_retry(self, method: str, url: str, **kwargs) -> httpx.Response:
|
||||
"""Make an HTTP request with automatic retries on transient errors."""
|
||||
import time
|
||||
|
||||
last_error = None
|
||||
delay = self.retry_delay
|
||||
|
||||
@@ -196,14 +245,18 @@ class RemoteTEIEmbeddings(Embeddings):
|
||||
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
logger.warning(f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
|
||||
logger.warning(
|
||||
f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
|
||||
)
|
||||
time.sleep(delay)
|
||||
delay *= 2 # Exponential backoff
|
||||
except httpx.HTTPStatusError as e:
|
||||
# Retry on 5xx server errors
|
||||
if e.response.status_code >= 500 and attempt < self.max_retries:
|
||||
last_error = e
|
||||
logger.warning(f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
|
||||
logger.warning(
|
||||
f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
|
||||
)
|
||||
time.sleep(delay)
|
||||
delay *= 2
|
||||
else:
|
||||
@@ -224,11 +277,28 @@ class RemoteTEIEmbeddings(Embeddings):
|
||||
response = self._request_with_retry("GET", f"{self.base_url}/info")
|
||||
info = response.json()
|
||||
self._model_id = info.get("model_id", "unknown")
|
||||
logger.info(f"Embeddings: TEI provider initialized (model: {self._model_id})")
|
||||
|
||||
# Get dimension from server info or by doing a test embedding
|
||||
if "max_input_length" in info and "model_dtype" in info:
|
||||
# Try to get dimension from info endpoint (some TEI versions expose it)
|
||||
# If not available, do a test embedding
|
||||
pass
|
||||
|
||||
# Do a test embedding to detect dimension
|
||||
test_response = self._request_with_retry(
|
||||
"POST",
|
||||
f"{self.base_url}/embed",
|
||||
json={"inputs": ["test"]},
|
||||
)
|
||||
test_embeddings = test_response.json()
|
||||
if test_embeddings and len(test_embeddings) > 0:
|
||||
self._dimension = len(test_embeddings[0])
|
||||
|
||||
logger.info(f"Embeddings: TEI provider initialized (model: {self._model_id}, dim: {self._dimension})")
|
||||
except httpx.HTTPError as e:
|
||||
raise RuntimeError(f"Failed to connect to TEI server at {self.base_url}: {e}")
|
||||
|
||||
def encode(self, texts: List[str]) -> List[List[float]]:
|
||||
def encode(self, texts: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generate embeddings using the remote TEI server.
|
||||
|
||||
@@ -248,7 +318,7 @@ class RemoteTEIEmbeddings(Embeddings):
|
||||
|
||||
# Process in batches
|
||||
for i in range(0, len(texts), self.batch_size):
|
||||
batch = texts[i:i + self.batch_size]
|
||||
batch = texts[i : i + self.batch_size]
|
||||
|
||||
try:
|
||||
response = self._request_with_retry(
|
||||
@@ -264,6 +334,369 @@ class RemoteTEIEmbeddings(Embeddings):
|
||||
return all_embeddings
|
||||
|
||||
|
||||
class OpenAIEmbeddings(Embeddings):
|
||||
"""
|
||||
OpenAI embeddings implementation using the OpenAI API.
|
||||
|
||||
Supports text-embedding-3-small (1536 dims), text-embedding-3-large (3072 dims),
|
||||
and text-embedding-ada-002 (1536 dims, legacy).
|
||||
|
||||
The embedding dimension is auto-detected from the model at initialization.
|
||||
"""
|
||||
|
||||
# Known dimensions for OpenAI embedding models
|
||||
MODEL_DIMENSIONS = {
|
||||
"text-embedding-3-small": 1536,
|
||||
"text-embedding-3-large": 3072,
|
||||
"text-embedding-ada-002": 1536,
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str,
|
||||
model: str = DEFAULT_EMBEDDINGS_OPENAI_MODEL,
|
||||
base_url: str | None = None,
|
||||
batch_size: int = 100,
|
||||
max_retries: int = 3,
|
||||
):
|
||||
"""
|
||||
Initialize OpenAI embeddings client.
|
||||
|
||||
Args:
|
||||
api_key: OpenAI API key
|
||||
model: OpenAI embedding model name (default: text-embedding-3-small)
|
||||
base_url: Custom base URL for OpenAI-compatible API (e.g., Azure OpenAI endpoint)
|
||||
batch_size: Maximum batch size for embedding requests (default: 100)
|
||||
max_retries: Maximum number of retries for failed requests (default: 3)
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.base_url = base_url
|
||||
self.batch_size = batch_size
|
||||
self.max_retries = max_retries
|
||||
self._client = None
|
||||
self._dimension: int | None = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "openai"
|
||||
|
||||
@property
|
||||
def dimension(self) -> int:
|
||||
if self._dimension is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
return self._dimension
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Initialize the OpenAI client and detect dimension."""
|
||||
if self._client is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
from openai import OpenAI
|
||||
except ImportError:
|
||||
raise ImportError("openai is required for OpenAIEmbeddings. Install it with: pip install openai")
|
||||
|
||||
base_url_msg = f" at {self.base_url}" if self.base_url else ""
|
||||
logger.info(f"Embeddings: initializing OpenAI provider with model {self.model}{base_url_msg}")
|
||||
|
||||
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
|
||||
client_kwargs = {"api_key": self.api_key, "max_retries": self.max_retries}
|
||||
if self.base_url:
|
||||
client_kwargs["base_url"] = self.base_url
|
||||
self._client = OpenAI(**client_kwargs)
|
||||
|
||||
# Try to get dimension from known models, otherwise do a test embedding
|
||||
if self.model in self.MODEL_DIMENSIONS:
|
||||
self._dimension = self.MODEL_DIMENSIONS[self.model]
|
||||
else:
|
||||
# Do a test embedding to detect dimension
|
||||
response = self._client.embeddings.create(
|
||||
model=self.model,
|
||||
input=["test"],
|
||||
)
|
||||
if response.data:
|
||||
self._dimension = len(response.data[0].embedding)
|
||||
|
||||
logger.info(f"Embeddings: OpenAI provider initialized (model: {self.model}, dim: {self._dimension})")
|
||||
|
||||
def encode(self, texts: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generate embeddings using the OpenAI API.
|
||||
|
||||
Args:
|
||||
texts: List of text strings to encode
|
||||
|
||||
Returns:
|
||||
List of embedding vectors
|
||||
"""
|
||||
if self._client is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
all_embeddings = []
|
||||
|
||||
# Process in batches
|
||||
for i in range(0, len(texts), self.batch_size):
|
||||
batch = texts[i : i + self.batch_size]
|
||||
|
||||
response = self._client.embeddings.create(
|
||||
model=self.model,
|
||||
input=batch,
|
||||
)
|
||||
|
||||
# Sort by index to ensure correct order
|
||||
batch_embeddings = sorted(response.data, key=lambda x: x.index)
|
||||
all_embeddings.extend([e.embedding for e in batch_embeddings])
|
||||
|
||||
return all_embeddings
|
||||
|
||||
|
||||
class CohereEmbeddings(Embeddings):
|
||||
"""
|
||||
Cohere embeddings implementation using the Cohere API.
|
||||
|
||||
Supports embed-english-v3.0 (1024 dims) and embed-multilingual-v3.0 (1024 dims).
|
||||
|
||||
The embedding dimension is auto-detected from the model at initialization.
|
||||
"""
|
||||
|
||||
# Known dimensions for Cohere embedding models
|
||||
MODEL_DIMENSIONS = {
|
||||
"embed-english-v3.0": 1024,
|
||||
"embed-multilingual-v3.0": 1024,
|
||||
"embed-english-light-v3.0": 384,
|
||||
"embed-multilingual-light-v3.0": 384,
|
||||
"embed-english-v2.0": 4096,
|
||||
"embed-multilingual-v2.0": 768,
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str,
|
||||
model: str = DEFAULT_EMBEDDINGS_COHERE_MODEL,
|
||||
base_url: str | None = None,
|
||||
batch_size: int = 96,
|
||||
timeout: float = 60.0,
|
||||
input_type: str = "search_document",
|
||||
):
|
||||
"""
|
||||
Initialize Cohere embeddings client.
|
||||
|
||||
Args:
|
||||
api_key: Cohere API key
|
||||
model: Cohere embedding model name (default: embed-english-v3.0)
|
||||
base_url: Custom base URL for Cohere-compatible API (e.g., Azure-hosted endpoint)
|
||||
batch_size: Maximum batch size for embedding requests (default: 96, Cohere's limit)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
input_type: Input type for embeddings (default: search_document).
|
||||
Options: search_document, search_query, classification, clustering
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.base_url = base_url
|
||||
self.batch_size = batch_size
|
||||
self.timeout = timeout
|
||||
self.input_type = input_type
|
||||
self._client = None
|
||||
self._dimension: int | None = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "cohere"
|
||||
|
||||
@property
|
||||
def dimension(self) -> int:
|
||||
if self._dimension is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
return self._dimension
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Initialize the Cohere client and detect dimension."""
|
||||
if self._client is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
import cohere
|
||||
except ImportError:
|
||||
raise ImportError("cohere is required for CohereEmbeddings. Install it with: pip install cohere")
|
||||
|
||||
base_url_msg = f" at {self.base_url}" if self.base_url else ""
|
||||
logger.info(f"Embeddings: initializing Cohere provider with model {self.model}{base_url_msg}")
|
||||
|
||||
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
|
||||
client_kwargs = {"api_key": self.api_key, "timeout": self.timeout}
|
||||
if self.base_url:
|
||||
client_kwargs["base_url"] = self.base_url
|
||||
self._client = cohere.Client(**client_kwargs)
|
||||
|
||||
# Try to get dimension from known models, otherwise do a test embedding
|
||||
if self.model in self.MODEL_DIMENSIONS:
|
||||
self._dimension = self.MODEL_DIMENSIONS[self.model]
|
||||
else:
|
||||
# Do a test embedding to detect dimension
|
||||
response = self._client.embed(
|
||||
texts=["test"],
|
||||
model=self.model,
|
||||
input_type=self.input_type,
|
||||
)
|
||||
if response.embeddings:
|
||||
self._dimension = len(response.embeddings[0])
|
||||
|
||||
logger.info(f"Embeddings: Cohere provider initialized (model: {self.model}, dim: {self._dimension})")
|
||||
|
||||
def encode(self, texts: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generate embeddings using the Cohere API.
|
||||
|
||||
Args:
|
||||
texts: List of text strings to encode
|
||||
|
||||
Returns:
|
||||
List of embedding vectors
|
||||
"""
|
||||
if self._client is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
all_embeddings = []
|
||||
|
||||
# Process in batches
|
||||
for i in range(0, len(texts), self.batch_size):
|
||||
batch = texts[i : i + self.batch_size]
|
||||
|
||||
response = self._client.embed(
|
||||
texts=batch,
|
||||
model=self.model,
|
||||
input_type=self.input_type,
|
||||
)
|
||||
|
||||
all_embeddings.extend(response.embeddings)
|
||||
|
||||
return all_embeddings
|
||||
|
||||
|
||||
class LiteLLMEmbeddings(Embeddings):
|
||||
"""
|
||||
LiteLLM embeddings implementation using LiteLLM proxy's /embeddings endpoint.
|
||||
|
||||
LiteLLM provides a unified interface for multiple embedding providers.
|
||||
The proxy exposes an OpenAI-compatible /embeddings endpoint.
|
||||
See: https://docs.litellm.ai/docs/embedding/supported_embedding
|
||||
|
||||
Supported providers via LiteLLM:
|
||||
- OpenAI (text-embedding-3-small, text-embedding-ada-002, etc.)
|
||||
- Cohere (embed-english-v3.0, etc.) - prefix with cohere/
|
||||
- Vertex AI (textembedding-gecko, etc.) - prefix with vertex_ai/
|
||||
- HuggingFace, Mistral, Voyage AI, etc.
|
||||
|
||||
The embedding dimension is auto-detected from the model at initialization.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_base: str = DEFAULT_LITELLM_API_BASE,
|
||||
api_key: str | None = None,
|
||||
model: str = DEFAULT_EMBEDDINGS_LITELLM_MODEL,
|
||||
batch_size: int = 100,
|
||||
timeout: float = 60.0,
|
||||
):
|
||||
"""
|
||||
Initialize LiteLLM embeddings client.
|
||||
|
||||
Args:
|
||||
api_base: Base URL of the LiteLLM proxy (default: http://localhost:4000)
|
||||
api_key: API key for the LiteLLM proxy (optional, depends on proxy config)
|
||||
model: Embedding model name (default: text-embedding-3-small)
|
||||
Use provider prefix for non-OpenAI models (e.g., cohere/embed-english-v3.0)
|
||||
batch_size: Maximum batch size for embedding requests (default: 100)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
"""
|
||||
self.api_base = api_base.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.batch_size = batch_size
|
||||
self.timeout = timeout
|
||||
self._client: httpx.Client | None = None
|
||||
self._dimension: int | None = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "litellm"
|
||||
|
||||
@property
|
||||
def dimension(self) -> int:
|
||||
if self._dimension is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
return self._dimension
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Initialize the HTTP client and detect embedding dimension."""
|
||||
if self._client is not None:
|
||||
return
|
||||
|
||||
logger.info(f"Embeddings: initializing LiteLLM provider at {self.api_base} with model {self.model}")
|
||||
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
|
||||
self._client = httpx.Client(timeout=self.timeout, headers=headers)
|
||||
|
||||
# Do a test embedding to detect dimension
|
||||
try:
|
||||
response = self._client.post(
|
||||
f"{self.api_base}/embeddings",
|
||||
json={"model": self.model, "input": ["test"]},
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
if result.get("data") and len(result["data"]) > 0:
|
||||
self._dimension = len(result["data"][0]["embedding"])
|
||||
logger.info(f"Embeddings: LiteLLM provider initialized (model: {self.model}, dim: {self._dimension})")
|
||||
except httpx.HTTPError as e:
|
||||
raise RuntimeError(f"Failed to connect to LiteLLM proxy at {self.api_base}: {e}")
|
||||
|
||||
def encode(self, texts: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generate embeddings using the LiteLLM proxy.
|
||||
|
||||
Args:
|
||||
texts: List of text strings to encode
|
||||
|
||||
Returns:
|
||||
List of embedding vectors
|
||||
"""
|
||||
if self._client is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
all_embeddings = []
|
||||
|
||||
# Process in batches
|
||||
for i in range(0, len(texts), self.batch_size):
|
||||
batch = texts[i : i + self.batch_size]
|
||||
|
||||
response = self._client.post(
|
||||
f"{self.api_base}/embeddings",
|
||||
json={"model": self.model, "input": batch},
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
# Sort by index to ensure correct order
|
||||
batch_embeddings = sorted(result["data"], key=lambda x: x["index"])
|
||||
all_embeddings.extend([e["embedding"] for e in batch_embeddings])
|
||||
|
||||
return all_embeddings
|
||||
|
||||
|
||||
def create_embeddings_from_env() -> Embeddings:
|
||||
"""
|
||||
Create an Embeddings instance based on environment variables.
|
||||
@@ -278,15 +711,36 @@ def create_embeddings_from_env() -> Embeddings:
|
||||
if provider == "tei":
|
||||
url = os.environ.get(ENV_EMBEDDINGS_TEI_URL)
|
||||
if not url:
|
||||
raise ValueError(
|
||||
f"{ENV_EMBEDDINGS_TEI_URL} is required when {ENV_EMBEDDINGS_PROVIDER} is 'tei'"
|
||||
)
|
||||
raise ValueError(f"{ENV_EMBEDDINGS_TEI_URL} is required when {ENV_EMBEDDINGS_PROVIDER} is 'tei'")
|
||||
return RemoteTEIEmbeddings(base_url=url)
|
||||
elif provider == "local":
|
||||
model = os.environ.get(ENV_EMBEDDINGS_LOCAL_MODEL)
|
||||
model_name = model or DEFAULT_EMBEDDINGS_LOCAL_MODEL
|
||||
return LocalSTEmbeddings(model_name=model_name)
|
||||
elif provider == "openai":
|
||||
# Use dedicated embeddings API key, or fall back to LLM API key
|
||||
api_key = os.environ.get(ENV_EMBEDDINGS_OPENAI_API_KEY) or os.environ.get(ENV_LLM_API_KEY)
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
f"{ENV_EMBEDDINGS_OPENAI_API_KEY} or {ENV_LLM_API_KEY} is required "
|
||||
f"when {ENV_EMBEDDINGS_PROVIDER} is 'openai'"
|
||||
)
|
||||
model = os.environ.get(ENV_EMBEDDINGS_OPENAI_MODEL, DEFAULT_EMBEDDINGS_OPENAI_MODEL)
|
||||
base_url = os.environ.get(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None
|
||||
return OpenAIEmbeddings(api_key=api_key, model=model, base_url=base_url)
|
||||
elif provider == "cohere":
|
||||
api_key = os.environ.get(ENV_COHERE_API_KEY)
|
||||
if not api_key:
|
||||
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_EMBEDDINGS_PROVIDER} is 'cohere'")
|
||||
model = os.environ.get(ENV_EMBEDDINGS_COHERE_MODEL, DEFAULT_EMBEDDINGS_COHERE_MODEL)
|
||||
base_url = os.environ.get(ENV_EMBEDDINGS_COHERE_BASE_URL) or None
|
||||
return CohereEmbeddings(api_key=api_key, model=model, base_url=base_url)
|
||||
elif provider == "litellm":
|
||||
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
|
||||
api_key = os.environ.get(ENV_LITELLM_API_KEY)
|
||||
model = os.environ.get(ENV_EMBEDDINGS_LITELLM_MODEL, DEFAULT_EMBEDDINGS_LITELLM_MODEL)
|
||||
return LiteLLMEmbeddings(api_base=api_base, api_key=api_key, model=model)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei'"
|
||||
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei', 'openai', 'cohere', 'litellm'"
|
||||
)
|
||||
|
||||
@@ -4,12 +4,14 @@ Entity extraction and resolution for memory system.
|
||||
Uses spaCy for entity extraction and implements resolution logic
|
||||
to disambiguate entities across memory units.
|
||||
"""
|
||||
import asyncpg
|
||||
from typing import List, Dict, Optional, Set, Any
|
||||
from difflib import SequenceMatcher
|
||||
from datetime import datetime, timezone
|
||||
from .db_utils import acquire_with_retry
|
||||
|
||||
from datetime import UTC, datetime
|
||||
from difflib import SequenceMatcher
|
||||
|
||||
import asyncpg
|
||||
|
||||
from .db_utils import acquire_with_retry
|
||||
from .memory_engine import fq_table
|
||||
|
||||
# Load spaCy model (singleton)
|
||||
_nlp = None
|
||||
@@ -32,11 +34,11 @@ class EntityResolver:
|
||||
async def resolve_entities_batch(
|
||||
self,
|
||||
bank_id: str,
|
||||
entities_data: List[Dict],
|
||||
entities_data: list[dict],
|
||||
context: str,
|
||||
unit_event_date,
|
||||
conn=None,
|
||||
) -> List[str]:
|
||||
) -> list[str]:
|
||||
"""
|
||||
Resolve multiple entities in batch (MUCH faster than sequential).
|
||||
|
||||
@@ -62,36 +64,38 @@ class EntityResolver:
|
||||
else:
|
||||
return await self._resolve_entities_batch_impl(conn, bank_id, entities_data, context, unit_event_date)
|
||||
|
||||
async def _resolve_entities_batch_impl(self, conn, bank_id: str, entities_data: List[Dict], context: str, unit_event_date) -> List[str]:
|
||||
async def _resolve_entities_batch_impl(
|
||||
self, conn, bank_id: str, entities_data: list[dict], context: str, unit_event_date
|
||||
) -> list[str]:
|
||||
# Query ALL candidates for this bank
|
||||
all_entities = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT canonical_name, id, metadata, last_seen, mention_count
|
||||
FROM entities
|
||||
FROM {fq_table("entities")}
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id
|
||||
bank_id,
|
||||
)
|
||||
|
||||
# Build entity ID to name mapping for co-occurrence lookups
|
||||
entity_id_to_name = {row['id']: row['canonical_name'].lower() for row in all_entities}
|
||||
entity_id_to_name = {row["id"]: row["canonical_name"].lower() for row in all_entities}
|
||||
|
||||
# Query ALL co-occurrences for this bank's entities in one query
|
||||
# This builds a map of entity_id -> set of co-occurring entity names
|
||||
all_cooccurrences = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT ec.entity_id_1, ec.entity_id_2, ec.cooccurrence_count
|
||||
FROM entity_cooccurrences ec
|
||||
WHERE ec.entity_id_1 IN (SELECT id FROM entities WHERE bank_id = $1)
|
||||
OR ec.entity_id_2 IN (SELECT id FROM entities WHERE bank_id = $1)
|
||||
FROM {fq_table("entity_cooccurrences")} ec
|
||||
WHERE ec.entity_id_1 IN (SELECT id FROM {fq_table("entities")} WHERE bank_id = $1)
|
||||
OR ec.entity_id_2 IN (SELECT id FROM {fq_table("entities")} WHERE bank_id = $1)
|
||||
""",
|
||||
bank_id
|
||||
bank_id,
|
||||
)
|
||||
|
||||
# Build co-occurrence map: entity_id -> set of co-occurring entity names (lowercase)
|
||||
cooccurrence_map: Dict[str, Set[str]] = {}
|
||||
cooccurrence_map: dict[str, set[str]] = {}
|
||||
for row in all_cooccurrences:
|
||||
eid1, eid2 = row['entity_id_1'], row['entity_id_2']
|
||||
eid1, eid2 = row["entity_id_1"], row["entity_id_2"]
|
||||
# Add both directions
|
||||
if eid1 not in cooccurrence_map:
|
||||
cooccurrence_map[eid1] = set()
|
||||
@@ -105,22 +109,24 @@ class EntityResolver:
|
||||
|
||||
# Build candidate map for each entity text
|
||||
all_candidates = {} # Maps entity_text -> list of candidates
|
||||
entity_texts = list(set(e['text'] for e in entities_data))
|
||||
entity_texts = list(set(e["text"] for e in entities_data))
|
||||
|
||||
for entity_text in entity_texts:
|
||||
matching = []
|
||||
entity_text_lower = entity_text.lower()
|
||||
for row in all_entities:
|
||||
canonical_name = row['canonical_name']
|
||||
ent_id = row['id']
|
||||
metadata = row['metadata']
|
||||
last_seen = row['last_seen']
|
||||
mention_count = row['mention_count']
|
||||
canonical_name = row["canonical_name"]
|
||||
ent_id = row["id"]
|
||||
metadata = row["metadata"]
|
||||
last_seen = row["last_seen"]
|
||||
mention_count = row["mention_count"]
|
||||
canonical_lower = canonical_name.lower()
|
||||
# Match if exact or substring match
|
||||
if (entity_text_lower == canonical_lower or
|
||||
entity_text_lower in canonical_lower or
|
||||
canonical_lower in entity_text_lower):
|
||||
if (
|
||||
entity_text_lower == canonical_lower
|
||||
or entity_text_lower in canonical_lower
|
||||
or canonical_lower in entity_text_lower
|
||||
):
|
||||
matching.append((ent_id, canonical_name, metadata, last_seen, mention_count))
|
||||
all_candidates[entity_text] = matching
|
||||
|
||||
@@ -130,10 +136,10 @@ class EntityResolver:
|
||||
entities_to_create = [] # (idx, entity_data, event_date)
|
||||
|
||||
for idx, entity_data in enumerate(entities_data):
|
||||
entity_text = entity_data['text']
|
||||
nearby_entities = entity_data.get('nearby_entities', [])
|
||||
entity_text = entity_data["text"]
|
||||
nearby_entities = entity_data.get("nearby_entities", [])
|
||||
# Use per-entity date if available, otherwise fall back to batch-level date
|
||||
entity_event_date = entity_data.get('event_date', unit_event_date)
|
||||
entity_event_date = entity_data.get("event_date", unit_event_date)
|
||||
|
||||
candidates = all_candidates.get(entity_text, [])
|
||||
|
||||
@@ -146,17 +152,13 @@ class EntityResolver:
|
||||
best_candidate = None
|
||||
best_score = 0.0
|
||||
|
||||
nearby_entity_set = {e['text'].lower() for e in nearby_entities if e['text'] != entity_text}
|
||||
nearby_entity_set = {e["text"].lower() for e in nearby_entities if e["text"] != entity_text}
|
||||
|
||||
for candidate_id, canonical_name, metadata, last_seen, mention_count in candidates:
|
||||
score = 0.0
|
||||
|
||||
# 1. Name similarity (0-0.5)
|
||||
name_similarity = SequenceMatcher(
|
||||
None,
|
||||
entity_text.lower(),
|
||||
canonical_name.lower()
|
||||
).ratio()
|
||||
name_similarity = SequenceMatcher(None, entity_text.lower(), canonical_name.lower()).ratio()
|
||||
score += name_similarity * 0.5
|
||||
|
||||
# 2. Co-occurring entities (0-0.3)
|
||||
@@ -169,8 +171,10 @@ class EntityResolver:
|
||||
# 3. Temporal proximity (0-0.2)
|
||||
if last_seen and entity_event_date:
|
||||
# Normalize timezone awareness for comparison
|
||||
event_date_utc = entity_event_date if entity_event_date.tzinfo else entity_event_date.replace(tzinfo=timezone.utc)
|
||||
last_seen_utc = last_seen if last_seen.tzinfo else last_seen.replace(tzinfo=timezone.utc)
|
||||
event_date_utc = (
|
||||
entity_event_date if entity_event_date.tzinfo else entity_event_date.replace(tzinfo=UTC)
|
||||
)
|
||||
last_seen_utc = last_seen if last_seen.tzinfo else last_seen.replace(tzinfo=UTC)
|
||||
days_diff = abs((event_date_utc - last_seen_utc).total_seconds() / 86400)
|
||||
if days_diff < 7:
|
||||
temporal_score = max(0, 1.0 - (days_diff / 7))
|
||||
@@ -192,23 +196,23 @@ class EntityResolver:
|
||||
# Batch update existing entities
|
||||
if entities_to_update:
|
||||
await conn.executemany(
|
||||
"""
|
||||
UPDATE entities SET
|
||||
f"""
|
||||
UPDATE {fq_table("entities")} SET
|
||||
mention_count = mention_count + 1,
|
||||
last_seen = $2
|
||||
WHERE id = $1::uuid
|
||||
""",
|
||||
entities_to_update
|
||||
entities_to_update,
|
||||
)
|
||||
|
||||
# Batch create new entities using COPY + INSERT for maximum speed
|
||||
# This handles duplicates via ON CONFLICT and returns all IDs
|
||||
if entities_to_create:
|
||||
# Group entities by canonical name (lowercase) to handle duplicates within batch
|
||||
# For duplicates, we only insert once and reuse the ID
|
||||
# For duplicates, we only insert once and reuse the ID, but track the count
|
||||
unique_entities = {} # lowercase_name -> (entity_data, event_date, [indices])
|
||||
for idx, entity_data, event_date in entities_to_create:
|
||||
name_lower = entity_data['text'].lower()
|
||||
name_lower = entity_data["text"].lower()
|
||||
if name_lower not in unique_entities:
|
||||
unique_entities[name_lower] = (entity_data, event_date, [idx])
|
||||
else:
|
||||
@@ -219,34 +223,37 @@ class EntityResolver:
|
||||
# Use a single query with unnest for speed
|
||||
entity_names = []
|
||||
entity_dates = []
|
||||
entity_counts = [] # Track how many times each entity appears in this batch
|
||||
indices_map = [] # Maps result index -> list of original indices
|
||||
|
||||
for name_lower, (entity_data, event_date, indices) in unique_entities.items():
|
||||
entity_names.append(entity_data['text'])
|
||||
entity_names.append(entity_data["text"])
|
||||
entity_dates.append(event_date)
|
||||
entity_counts.append(len(indices)) # Count of occurrences in this batch
|
||||
indices_map.append(indices)
|
||||
|
||||
# Batch INSERT ... ON CONFLICT with RETURNING
|
||||
# This is much faster than individual inserts
|
||||
# Uses the batch count for mention_count instead of always 1
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
INSERT INTO entities (bank_id, canonical_name, first_seen, last_seen, mention_count)
|
||||
SELECT $1, name, event_date, event_date, 1
|
||||
FROM unnest($2::text[], $3::timestamptz[]) AS t(name, event_date)
|
||||
f"""
|
||||
INSERT INTO {fq_table("entities")} (bank_id, canonical_name, first_seen, last_seen, mention_count)
|
||||
SELECT $1, name, event_date, event_date, cnt
|
||||
FROM unnest($2::text[], $3::timestamptz[], $4::int[]) AS t(name, event_date, cnt)
|
||||
ON CONFLICT (bank_id, LOWER(canonical_name))
|
||||
DO UPDATE SET
|
||||
mention_count = entities.mention_count + 1,
|
||||
mention_count = {fq_table("entities")}.mention_count + EXCLUDED.mention_count,
|
||||
last_seen = EXCLUDED.last_seen
|
||||
RETURNING id
|
||||
""",
|
||||
bank_id,
|
||||
entity_names,
|
||||
entity_dates
|
||||
entity_dates,
|
||||
entity_counts,
|
||||
)
|
||||
|
||||
# Map returned IDs back to original indices
|
||||
for result_idx, row in enumerate(rows):
|
||||
entity_id = row['id']
|
||||
entity_id = row["id"]
|
||||
for original_idx in indices_map[result_idx]:
|
||||
entity_ids[original_idx] = entity_id
|
||||
|
||||
@@ -257,7 +264,7 @@ class EntityResolver:
|
||||
bank_id: str,
|
||||
entity_text: str,
|
||||
context: str,
|
||||
nearby_entities: List[Dict],
|
||||
nearby_entities: list[dict],
|
||||
unit_event_date,
|
||||
) -> str:
|
||||
"""
|
||||
@@ -276,9 +283,9 @@ class EntityResolver:
|
||||
async with acquire_with_retry(self.pool) as conn:
|
||||
# Find candidate entities with similar name
|
||||
candidates = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT id, canonical_name, metadata, last_seen
|
||||
FROM entities
|
||||
FROM {fq_table("entities")}
|
||||
WHERE bank_id = $1
|
||||
AND (
|
||||
canonical_name ILIKE $2
|
||||
@@ -287,14 +294,14 @@ class EntityResolver:
|
||||
)
|
||||
ORDER BY mention_count DESC
|
||||
""",
|
||||
bank_id, entity_text, f"%{entity_text}%"
|
||||
bank_id,
|
||||
entity_text,
|
||||
f"%{entity_text}%",
|
||||
)
|
||||
|
||||
if not candidates:
|
||||
# New entity - create it
|
||||
return await self._create_entity(
|
||||
conn, bank_id, entity_text, unit_event_date
|
||||
)
|
||||
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
|
||||
|
||||
# Score candidates based on:
|
||||
# 1. Name similarity
|
||||
@@ -306,31 +313,27 @@ class EntityResolver:
|
||||
best_score = 0.0
|
||||
best_name_similarity = 0.0
|
||||
|
||||
nearby_entity_set = {e['text'].lower() for e in nearby_entities if e['text'] != entity_text}
|
||||
nearby_entity_set = {e["text"].lower() for e in nearby_entities if e["text"] != entity_text}
|
||||
|
||||
for row in candidates:
|
||||
candidate_id = row['id']
|
||||
canonical_name = row['canonical_name']
|
||||
metadata = row['metadata']
|
||||
last_seen = row['last_seen']
|
||||
candidate_id = row["id"]
|
||||
canonical_name = row["canonical_name"]
|
||||
metadata = row["metadata"]
|
||||
last_seen = row["last_seen"]
|
||||
score = 0.0
|
||||
|
||||
# 1. Name similarity (0-1)
|
||||
name_similarity = SequenceMatcher(
|
||||
None,
|
||||
entity_text.lower(),
|
||||
canonical_name.lower()
|
||||
).ratio()
|
||||
name_similarity = SequenceMatcher(None, entity_text.lower(), canonical_name.lower()).ratio()
|
||||
score += name_similarity * 0.5
|
||||
|
||||
# 2. Co-occurring entities (0-0.5)
|
||||
# Get entities that co-occurred with this candidate before
|
||||
# Use the materialized co-occurrence cache for fast lookup
|
||||
co_entity_rows = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT e.canonical_name, ec.cooccurrence_count
|
||||
FROM entity_cooccurrences ec
|
||||
JOIN entities e ON (
|
||||
FROM {fq_table("entity_cooccurrences")} ec
|
||||
JOIN {fq_table("entities")} e ON (
|
||||
CASE
|
||||
WHEN ec.entity_id_1 = $1 THEN ec.entity_id_2
|
||||
WHEN ec.entity_id_2 = $1 THEN ec.entity_id_1
|
||||
@@ -338,9 +341,9 @@ class EntityResolver:
|
||||
)
|
||||
WHERE ec.entity_id_1 = $1 OR ec.entity_id_2 = $1
|
||||
""",
|
||||
candidate_id
|
||||
candidate_id,
|
||||
)
|
||||
co_entities = {r['canonical_name'].lower() for r in co_entity_rows}
|
||||
co_entities = {r["canonical_name"].lower() for r in co_entity_rows}
|
||||
|
||||
# Check overlap with nearby entities
|
||||
overlap = len(nearby_entity_set & co_entities)
|
||||
@@ -366,20 +369,19 @@ class EntityResolver:
|
||||
if best_score > threshold:
|
||||
# Update entity
|
||||
await conn.execute(
|
||||
"""
|
||||
UPDATE entities
|
||||
f"""
|
||||
UPDATE {fq_table("entities")}
|
||||
SET mention_count = mention_count + 1,
|
||||
last_seen = $1
|
||||
WHERE id = $2
|
||||
""",
|
||||
unit_event_date, best_candidate
|
||||
unit_event_date,
|
||||
best_candidate,
|
||||
)
|
||||
return best_candidate
|
||||
else:
|
||||
# Not confident - create new entity
|
||||
return await self._create_entity(
|
||||
conn, bank_id, entity_text, unit_event_date
|
||||
)
|
||||
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
|
||||
|
||||
async def _create_entity(
|
||||
self,
|
||||
@@ -404,16 +406,19 @@ class EntityResolver:
|
||||
Entity ID
|
||||
"""
|
||||
entity_id = await conn.fetchval(
|
||||
"""
|
||||
INSERT INTO entities (bank_id, canonical_name, first_seen, last_seen, mention_count)
|
||||
f"""
|
||||
INSERT INTO {fq_table("entities")} (bank_id, canonical_name, first_seen, last_seen, mention_count)
|
||||
VALUES ($1, $2, $3, $4, 1)
|
||||
ON CONFLICT (bank_id, LOWER(canonical_name))
|
||||
DO UPDATE SET
|
||||
mention_count = entities.mention_count + 1,
|
||||
mention_count = {fq_table("entities")}.mention_count + 1,
|
||||
last_seen = EXCLUDED.last_seen
|
||||
RETURNING id
|
||||
""",
|
||||
bank_id, entity_text, event_date, event_date
|
||||
bank_id,
|
||||
entity_text,
|
||||
event_date,
|
||||
event_date,
|
||||
)
|
||||
return entity_id
|
||||
|
||||
@@ -429,25 +434,27 @@ class EntityResolver:
|
||||
async with acquire_with_retry(self.pool) as conn:
|
||||
# Insert unit-entity link
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO unit_entities (unit_id, entity_id)
|
||||
f"""
|
||||
INSERT INTO {fq_table("unit_entities")} (unit_id, entity_id)
|
||||
VALUES ($1, $2)
|
||||
ON CONFLICT DO NOTHING
|
||||
""",
|
||||
unit_id, entity_id
|
||||
unit_id,
|
||||
entity_id,
|
||||
)
|
||||
|
||||
# Update co-occurrence cache: find other entities in this unit
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT entity_id
|
||||
FROM unit_entities
|
||||
FROM {fq_table("unit_entities")}
|
||||
WHERE unit_id = $1 AND entity_id != $2
|
||||
""",
|
||||
unit_id, entity_id
|
||||
unit_id,
|
||||
entity_id,
|
||||
)
|
||||
|
||||
other_entities = [row['entity_id'] for row in rows]
|
||||
other_entities = [row["entity_id"] for row in rows]
|
||||
|
||||
# Update co-occurrences for each pair
|
||||
for other_entity_id in other_entities:
|
||||
@@ -469,18 +476,19 @@ class EntityResolver:
|
||||
entity_id_1, entity_id_2 = entity_id_2, entity_id_1
|
||||
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO entity_cooccurrences (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
|
||||
f"""
|
||||
INSERT INTO {fq_table("entity_cooccurrences")} (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
|
||||
VALUES ($1, $2, 1, NOW())
|
||||
ON CONFLICT (entity_id_1, entity_id_2)
|
||||
DO UPDATE SET
|
||||
cooccurrence_count = entity_cooccurrences.cooccurrence_count + 1,
|
||||
cooccurrence_count = {fq_table("entity_cooccurrences")}.cooccurrence_count + 1,
|
||||
last_cooccurred = NOW()
|
||||
""",
|
||||
entity_id_1, entity_id_2
|
||||
entity_id_1,
|
||||
entity_id_2,
|
||||
)
|
||||
|
||||
async def link_units_to_entities_batch(self, unit_entity_pairs: List[tuple[str, str]], conn=None):
|
||||
async def link_units_to_entities_batch(self, unit_entity_pairs: list[tuple[str, str]], conn=None):
|
||||
"""
|
||||
Link multiple memory units to entities in batch (MUCH faster than sequential).
|
||||
|
||||
@@ -499,15 +507,15 @@ class EntityResolver:
|
||||
else:
|
||||
return await self._link_units_to_entities_batch_impl(conn, unit_entity_pairs)
|
||||
|
||||
async def _link_units_to_entities_batch_impl(self, conn, unit_entity_pairs: List[tuple[str, str]]):
|
||||
async def _link_units_to_entities_batch_impl(self, conn, unit_entity_pairs: list[tuple[str, str]]):
|
||||
# Batch insert all unit-entity links
|
||||
await conn.executemany(
|
||||
"""
|
||||
INSERT INTO unit_entities (unit_id, entity_id)
|
||||
f"""
|
||||
INSERT INTO {fq_table("unit_entities")} (unit_id, entity_id)
|
||||
VALUES ($1, $2)
|
||||
ON CONFLICT DO NOTHING
|
||||
""",
|
||||
unit_entity_pairs
|
||||
unit_entity_pairs,
|
||||
)
|
||||
|
||||
# Build map of unit -> entities for co-occurrence calculation
|
||||
@@ -524,7 +532,7 @@ class EntityResolver:
|
||||
entity_list = list(entity_ids) # Convert set to list for iteration
|
||||
# For each pair of entities in this unit, create co-occurrence
|
||||
for i, entity_id_1 in enumerate(entity_list):
|
||||
for entity_id_2 in entity_list[i+1:]:
|
||||
for entity_id_2 in entity_list[i + 1 :]:
|
||||
# Skip if same entity (shouldn't happen with set, but be safe)
|
||||
if entity_id_1 == entity_id_2:
|
||||
continue
|
||||
@@ -535,20 +543,20 @@ class EntityResolver:
|
||||
|
||||
# Batch update co-occurrences
|
||||
if cooccurrence_pairs:
|
||||
now = datetime.now(timezone.utc)
|
||||
now = datetime.now(UTC)
|
||||
await conn.executemany(
|
||||
"""
|
||||
INSERT INTO entity_cooccurrences (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
|
||||
f"""
|
||||
INSERT INTO {fq_table("entity_cooccurrences")} (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
|
||||
VALUES ($1, $2, $3, $4)
|
||||
ON CONFLICT (entity_id_1, entity_id_2)
|
||||
DO UPDATE SET
|
||||
cooccurrence_count = entity_cooccurrences.cooccurrence_count + 1,
|
||||
cooccurrence_count = {fq_table("entity_cooccurrences")}.cooccurrence_count + 1,
|
||||
last_cooccurred = EXCLUDED.last_cooccurred
|
||||
""",
|
||||
[(e1, e2, 1, now) for e1, e2 in cooccurrence_pairs]
|
||||
[(e1, e2, 1, now) for e1, e2 in cooccurrence_pairs],
|
||||
)
|
||||
|
||||
async def get_units_by_entity(self, entity_id: str, limit: int = 100) -> List[str]:
|
||||
async def get_units_by_entity(self, entity_id: str, limit: int = 100) -> list[str]:
|
||||
"""
|
||||
Get all units that mention an entity.
|
||||
|
||||
@@ -561,22 +569,23 @@ class EntityResolver:
|
||||
"""
|
||||
async with acquire_with_retry(self.pool) as conn:
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT unit_id
|
||||
FROM unit_entities
|
||||
FROM {fq_table("unit_entities")}
|
||||
WHERE entity_id = $1
|
||||
ORDER BY unit_id
|
||||
LIMIT $2
|
||||
""",
|
||||
entity_id, limit
|
||||
entity_id,
|
||||
limit,
|
||||
)
|
||||
return [row['unit_id'] for row in rows]
|
||||
return [row["unit_id"] for row in rows]
|
||||
|
||||
async def get_entity_by_text(
|
||||
self,
|
||||
bank_id: str,
|
||||
entity_text: str,
|
||||
) -> Optional[str]:
|
||||
) -> str | None:
|
||||
"""
|
||||
Find an entity by text (for query resolution).
|
||||
|
||||
@@ -589,14 +598,15 @@ class EntityResolver:
|
||||
"""
|
||||
async with acquire_with_retry(self.pool) as conn:
|
||||
row = await conn.fetchrow(
|
||||
"""
|
||||
SELECT id FROM entities
|
||||
f"""
|
||||
SELECT id FROM {fq_table("entities")}
|
||||
WHERE bank_id = $1
|
||||
AND canonical_name ILIKE $2
|
||||
ORDER BY mention_count DESC
|
||||
LIMIT 1
|
||||
""",
|
||||
bank_id, entity_text
|
||||
bank_id,
|
||||
entity_text,
|
||||
)
|
||||
|
||||
return row['id'] if row else None
|
||||
return row["id"] if row else None
|
||||
|
||||
@@ -0,0 +1,619 @@
|
||||
"""Abstract interface for MemoryEngine public methods.
|
||||
|
||||
This module defines the public API that HTTP endpoints and extensions should use
|
||||
to interact with the memory system. All methods require a RequestContext for
|
||||
authentication when a TenantExtension is configured.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
from hindsight_api.engine.response_models import RecallResult, ReflectResult
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
|
||||
class MemoryEngineInterface(ABC):
|
||||
"""
|
||||
Abstract interface for the Memory Engine.
|
||||
|
||||
This defines the public API that should be used by HTTP endpoints and extensions.
|
||||
All methods require a RequestContext for authentication.
|
||||
"""
|
||||
|
||||
# =========================================================================
|
||||
# Health & Status
|
||||
# =========================================================================
|
||||
|
||||
@abstractmethod
|
||||
async def health_check(self) -> dict:
|
||||
"""
|
||||
Check the health of the memory system.
|
||||
|
||||
Returns:
|
||||
Dict with 'status' key ('healthy' or 'unhealthy') and additional info.
|
||||
"""
|
||||
...
|
||||
|
||||
# =========================================================================
|
||||
# Core Memory Operations
|
||||
# =========================================================================
|
||||
|
||||
@abstractmethod
|
||||
async def retain_batch_async(
|
||||
self,
|
||||
bank_id: str,
|
||||
contents: list[dict[str, Any]],
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Retain a batch of memory items.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
contents: List of content dicts with 'content', optional 'event_date',
|
||||
'context', 'metadata', 'document_id'.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with processing results.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def recall_async(
|
||||
self,
|
||||
bank_id: str,
|
||||
query: str,
|
||||
*,
|
||||
budget: "Budget | None" = None,
|
||||
max_tokens: int = 4096,
|
||||
enable_trace: bool = False,
|
||||
fact_type: list[str] | None = None,
|
||||
question_date: datetime | None = None,
|
||||
include_entities: bool = False,
|
||||
max_entity_tokens: int = 500,
|
||||
include_chunks: bool = False,
|
||||
max_chunk_tokens: int = 8192,
|
||||
request_context: "RequestContext",
|
||||
) -> "RecallResult":
|
||||
"""
|
||||
Recall memories relevant to a query.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
query: The search query.
|
||||
budget: Search budget (LOW, MID, HIGH).
|
||||
max_tokens: Maximum tokens in response.
|
||||
enable_trace: Include trace information.
|
||||
fact_type: Filter by fact types.
|
||||
question_date: Context date for temporal relevance.
|
||||
include_entities: Include entity observations.
|
||||
max_entity_tokens: Max tokens for entity observations.
|
||||
include_chunks: Include raw chunks.
|
||||
max_chunk_tokens: Max tokens for chunks.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
RecallResult with matching memories.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def reflect_async(
|
||||
self,
|
||||
bank_id: str,
|
||||
query: str,
|
||||
*,
|
||||
budget: "Budget | None" = None,
|
||||
context: str | None = None,
|
||||
max_tokens: int = 4096,
|
||||
response_schema: dict | None = None,
|
||||
request_context: "RequestContext",
|
||||
) -> "ReflectResult":
|
||||
"""
|
||||
Reflect on a query and generate a thoughtful response.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
query: The question to reflect on.
|
||||
budget: Search budget for retrieving context.
|
||||
context: Additional context for the reflection.
|
||||
max_tokens: Maximum tokens for the response.
|
||||
response_schema: Optional JSON Schema for structured output.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
ReflectResult with generated response and supporting facts.
|
||||
"""
|
||||
...
|
||||
|
||||
# =========================================================================
|
||||
# Bank Management
|
||||
# =========================================================================
|
||||
|
||||
@abstractmethod
|
||||
async def list_banks(
|
||||
self,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
List all memory banks.
|
||||
|
||||
Args:
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
List of bank info dicts.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def get_bank_profile(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Get bank profile including disposition and mission.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Bank profile dict with bank_id, name, disposition, and mission.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def update_bank_disposition(
|
||||
self,
|
||||
bank_id: str,
|
||||
disposition: dict[str, int],
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> None:
|
||||
"""
|
||||
Update bank disposition traits.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
disposition: Dict with trait values.
|
||||
request_context: Request context for authentication.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def merge_bank_mission(
|
||||
self,
|
||||
bank_id: str,
|
||||
new_info: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Merge new mission information into bank profile.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
new_info: New mission information to merge.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Updated mission info.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def set_bank_mission(
|
||||
self,
|
||||
bank_id: str,
|
||||
mission: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Set the bank's mission (replaces existing).
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
mission: The mission text.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with bank_id and mission.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete_bank(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
fact_type: str | None = None,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, int]:
|
||||
"""
|
||||
Delete a bank or its memories.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
fact_type: If specified, only delete memories of this type.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with deletion counts.
|
||||
"""
|
||||
...
|
||||
|
||||
# =========================================================================
|
||||
# Memory Units
|
||||
# =========================================================================
|
||||
|
||||
@abstractmethod
|
||||
async def list_memory_units(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
fact_type: str | None = None,
|
||||
search_query: str | None = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
List memory units with pagination.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
fact_type: Filter by fact type.
|
||||
search_query: Full-text search query.
|
||||
limit: Maximum results.
|
||||
offset: Pagination offset.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with 'items', 'total', 'limit', 'offset'.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete_memory_unit(
|
||||
self,
|
||||
unit_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Delete a specific memory unit.
|
||||
|
||||
Args:
|
||||
unit_id: The memory unit ID.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Deletion result.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def get_graph_data(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
fact_type: str | None = None,
|
||||
limit: int = 1000,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Get graph data for visualization.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
fact_type: Filter by fact type.
|
||||
limit: Maximum number of items to return (default: 1000).
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with nodes, edges, table_rows, total_units, limit.
|
||||
"""
|
||||
...
|
||||
|
||||
# =========================================================================
|
||||
# Documents
|
||||
# =========================================================================
|
||||
|
||||
@abstractmethod
|
||||
async def list_documents(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
search_query: str | None = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
List documents with pagination.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
search_query: Search query.
|
||||
limit: Maximum results.
|
||||
offset: Pagination offset.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with 'items', 'total', 'limit', 'offset'.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def get_document(
|
||||
self,
|
||||
document_id: str,
|
||||
bank_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any] | None:
|
||||
"""
|
||||
Get a specific document.
|
||||
|
||||
Args:
|
||||
document_id: The document ID.
|
||||
bank_id: The memory bank ID.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Document dict or None if not found.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete_document(
|
||||
self,
|
||||
document_id: str,
|
||||
bank_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, int]:
|
||||
"""
|
||||
Delete a document and its memory units.
|
||||
|
||||
Args:
|
||||
document_id: The document ID.
|
||||
bank_id: The memory bank ID.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with deletion counts.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def get_chunk(
|
||||
self,
|
||||
chunk_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any] | None:
|
||||
"""
|
||||
Get a specific chunk.
|
||||
|
||||
Args:
|
||||
chunk_id: The chunk ID.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Chunk dict or None if not found.
|
||||
"""
|
||||
...
|
||||
|
||||
# =========================================================================
|
||||
# Entities
|
||||
# =========================================================================
|
||||
|
||||
@abstractmethod
|
||||
async def list_entities(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
List entities for a bank with pagination.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
limit: Maximum results.
|
||||
offset: Offset for pagination.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with items, total, limit, offset.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def get_entity_observations(
|
||||
self,
|
||||
bank_id: str,
|
||||
entity_id: str,
|
||||
*,
|
||||
limit: int = 10,
|
||||
request_context: "RequestContext",
|
||||
) -> list[Any]:
|
||||
"""
|
||||
Get observations for an entity.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
entity_id: The entity ID.
|
||||
limit: Maximum observations.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
List of EntityObservation objects.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def regenerate_entity_observations(
|
||||
self,
|
||||
bank_id: str,
|
||||
entity_id: str,
|
||||
entity_name: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> None:
|
||||
"""
|
||||
Regenerate observations for an entity.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
entity_id: The entity ID.
|
||||
entity_name: The entity's canonical name.
|
||||
request_context: Request context for authentication.
|
||||
"""
|
||||
...
|
||||
|
||||
# =========================================================================
|
||||
# Statistics & Operations
|
||||
# =========================================================================
|
||||
|
||||
@abstractmethod
|
||||
async def get_bank_stats(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Get statistics about memory nodes and links for a bank.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with node_counts, link_counts, link_counts_by_fact_type,
|
||||
link_breakdown, and operations stats.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def get_entity(
|
||||
self,
|
||||
bank_id: str,
|
||||
entity_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any] | None:
|
||||
"""
|
||||
Get entity details including metadata and observations.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
entity_id: The entity ID.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Entity dict with id, canonical_name, mention_count, first_seen,
|
||||
last_seen, metadata, and observations. None if not found.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def list_operations(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
List async operations for a bank.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with 'total' (int) and 'operations' (list of operation dicts).
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def cancel_operation(
|
||||
self,
|
||||
bank_id: str,
|
||||
operation_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Cancel a pending async operation.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
operation_id: The operation ID to cancel.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with success status and message.
|
||||
|
||||
Raises:
|
||||
ValueError: If operation not found.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def update_bank(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
name: str | None = None,
|
||||
mission: str | None = None,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Update bank name and/or mission.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
name: New bank name (optional).
|
||||
mission: New mission text (optional, replaces existing).
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Updated bank profile dict.
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def submit_async_retain(
|
||||
self,
|
||||
bank_id: str,
|
||||
contents: list[dict[str, Any]],
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Submit a batch retain operation to run asynchronously.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID.
|
||||
contents: List of content dicts to retain.
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with operation_id and items_count.
|
||||
"""
|
||||
...
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,18 @@
|
||||
"""
|
||||
Mental models module for Hindsight.
|
||||
|
||||
Mental models are synthesized summaries that represent understanding. They come
|
||||
in different subtypes based on how they were created:
|
||||
|
||||
- Structural: Derived from the bank's mission (e.g., "Be a PM for engineering team")
|
||||
These are created upfront based on what any agent with this role would need.
|
||||
|
||||
- Emergent: Discovered from data patterns (named entities, temporal clusters, etc.)
|
||||
These surface organically as facts are retained.
|
||||
|
||||
- Pinned: User-defined models that persist across refreshes.
|
||||
"""
|
||||
|
||||
from .models import MentalModel, MentalModelSubtype
|
||||
|
||||
__all__ = ["MentalModel", "MentalModelSubtype"]
|
||||
@@ -0,0 +1,311 @@
|
||||
"""
|
||||
Emergent mental model detection and promotion.
|
||||
|
||||
Emergent models are discovered from data patterns:
|
||||
- Named entity extraction (people, projects, systems)
|
||||
- Temporal clustering (events with multiple references)
|
||||
- Causal patterns ("Because X, we do Y")
|
||||
- Behavioral anchors ("After X, we started Y")
|
||||
- Reference frequency (anything mentioned repeatedly)
|
||||
|
||||
When a pattern is detected, it goes through a mission filter to check relevance,
|
||||
and if relevant, is promoted to a mental model.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from .models import EmergentCandidate
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..llm_wrapper import LLMConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MissionFilterCandidate(BaseModel):
|
||||
"""Result of mission filtering for a single candidate."""
|
||||
|
||||
name: str
|
||||
promote: bool = Field(description="True if this is a specific named entity worth tracking")
|
||||
reason: str = Field(description="Brief explanation for the decision")
|
||||
|
||||
|
||||
class MissionFilterResponse(BaseModel):
|
||||
"""Response from LLM for mission filtering."""
|
||||
|
||||
candidates: list[MissionFilterCandidate] = Field(description="Filtering decision for each candidate")
|
||||
|
||||
|
||||
def build_mission_filter_prompt(mission: str, candidates: list[EmergentCandidate]) -> str:
|
||||
"""Build the prompt for filtering candidates by mission relevance."""
|
||||
candidate_list = "\n".join(
|
||||
[f"- {c.name} (mentions: {c.mention_count}, method: {c.detection_method})" for c in candidates]
|
||||
)
|
||||
|
||||
return f"""Filter these detected entities. For each one, decide: promote=true or promote=false.
|
||||
|
||||
MISSION: {mission}
|
||||
|
||||
DETECTED ENTITIES:
|
||||
{candidate_list}
|
||||
|
||||
=== DECISION RULES ===
|
||||
|
||||
Set promote=true ONLY for specific, named entities:
|
||||
- Person names: "John", "Maria", "Alice Chen", "Dr. Smith"
|
||||
- Named organizations: "Google", "Acme Corp", "Frontend Team"
|
||||
- Named places: "Central Park Zoo", "NYC Office", "Building A"
|
||||
- Named projects: "Project Phoenix", "Auth Service v2"
|
||||
|
||||
Set promote=false for EVERYTHING ELSE, including:
|
||||
- Common English words: user, support, help, family, kids, parents, friends, people, team, photo, nature, park, office, home, work, school, joy, love, hope, fear, anger, gratitude, kindness, passion, motivation, inspiration, encouragement, positivity, energy, community, connection, commitment, collaboration, growth, impact, difference, success, progress, change, education, volunteering, veterans, homeless, shelter, meeting, project, system, process, event
|
||||
- Generic categories (even capitalized): Users, Customers, Team, Family, Kids, Veterans, Community
|
||||
- Abstract concepts: motivation, inspiration, gratitude, commitment, resilience
|
||||
|
||||
THE TEST: Is this a specific name you'd find in a contact list or org chart?
|
||||
- "John" → YES (promote=true)
|
||||
- "kids" → NO (promote=false)
|
||||
- "community" → NO (promote=false)
|
||||
- "Maria" → YES (promote=true)
|
||||
- "park" → NO (promote=false)
|
||||
|
||||
When in doubt, set promote=false."""
|
||||
|
||||
|
||||
def get_mission_filter_system_message() -> str:
|
||||
"""System message for mission filtering."""
|
||||
return """You filter entities for promotion. Output JSON with 'candidates' array.
|
||||
|
||||
Rules:
|
||||
- promote=true ONLY for specific names (people, organizations, named places/projects)
|
||||
- promote=false for common words, generic categories, abstract concepts
|
||||
|
||||
Examples:
|
||||
- "John" → promote=true (person name)
|
||||
- "kids" → promote=false (generic category)
|
||||
- "community" → promote=false (abstract concept)
|
||||
- "Google" → promote=true (organization name)
|
||||
- "motivation" → promote=false (abstract concept)
|
||||
|
||||
When in doubt, promote=false. Most entities should be rejected."""
|
||||
|
||||
|
||||
async def filter_candidates_by_mission(
|
||||
llm_config: "LLMConfig",
|
||||
mission: str,
|
||||
candidates: list[EmergentCandidate],
|
||||
) -> list[EmergentCandidate]:
|
||||
"""
|
||||
Filter emergent candidates to keep only specific, named entities.
|
||||
|
||||
Args:
|
||||
llm_config: LLM configuration
|
||||
mission: The bank's mission (used for context)
|
||||
candidates: List of detected candidates
|
||||
|
||||
Returns:
|
||||
Filtered list of candidates that are specific named entities
|
||||
"""
|
||||
if not candidates:
|
||||
return []
|
||||
|
||||
if not mission:
|
||||
# No mission = no filtering, keep all candidates
|
||||
logger.debug("[EMERGENT] No mission set, skipping filter")
|
||||
return candidates
|
||||
|
||||
prompt = build_mission_filter_prompt(mission, candidates)
|
||||
|
||||
try:
|
||||
result = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": get_mission_filter_system_message()},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
response_format=MissionFilterResponse,
|
||||
scope="mental_model_mission_filter",
|
||||
)
|
||||
|
||||
# Build name -> promote map
|
||||
promote_map = {c.name: c.promote for c in result.candidates}
|
||||
|
||||
# Filter candidates
|
||||
filtered = []
|
||||
for candidate in candidates:
|
||||
if candidate.name in promote_map:
|
||||
if promote_map[candidate.name]:
|
||||
filtered.append(candidate)
|
||||
logger.debug(f"[EMERGENT] Promoting '{candidate.name}'")
|
||||
else:
|
||||
logger.debug(f"[EMERGENT] Rejecting '{candidate.name}'")
|
||||
else:
|
||||
# Candidate not in response - reject by default
|
||||
logger.debug(f"[EMERGENT] '{candidate.name}' not in response, rejecting")
|
||||
|
||||
logger.info(f"[EMERGENT] Mission filter: {len(filtered)}/{len(candidates)} candidates promoted")
|
||||
return filtered
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"[EMERGENT] Mission filter failed, rejecting all candidates: {e}")
|
||||
return []
|
||||
|
||||
|
||||
async def evaluate_emergent_models(
|
||||
llm_config: "LLMConfig",
|
||||
models: list[dict],
|
||||
) -> list[str]:
|
||||
"""
|
||||
Evaluate existing emergent models to check if they should be kept.
|
||||
|
||||
This re-evaluates emergent models using the same filtering criteria
|
||||
as new candidates. Models that are generic/abstract will be removed.
|
||||
|
||||
Args:
|
||||
llm_config: LLM configuration
|
||||
models: List of existing emergent model dicts with 'name', 'id'
|
||||
|
||||
Returns:
|
||||
List of model IDs that should be REMOVED (no longer valid)
|
||||
"""
|
||||
if not models:
|
||||
return []
|
||||
|
||||
# Convert existing models to candidates for evaluation
|
||||
candidates = [
|
||||
EmergentCandidate(
|
||||
name=m["name"],
|
||||
detection_method="existing_emergent_model",
|
||||
mention_count=0,
|
||||
)
|
||||
for m in models
|
||||
]
|
||||
|
||||
# Build a simple prompt for re-evaluation
|
||||
names_list = "\n".join([f"- {m['name']}" for m in models])
|
||||
prompt = f"""Re-evaluate these existing mental models. For each one, decide: promote=true (keep) or promote=false (remove).
|
||||
|
||||
EXISTING MODELS:
|
||||
{names_list}
|
||||
|
||||
=== DECISION RULES ===
|
||||
|
||||
Set promote=true ONLY for specific, named entities:
|
||||
- Person names: "John", "Maria", "Alice Chen", "Dr. Smith"
|
||||
- Named organizations: "Google", "Acme Corp", "Frontend Team"
|
||||
- Named places: "Central Park Zoo", "NYC Office", "Building A"
|
||||
- Named projects: "Project Phoenix", "Auth Service v2"
|
||||
|
||||
Set promote=false for EVERYTHING ELSE, including:
|
||||
- Common English words: user, support, help, family, kids, parents, friends, people, team, photo, nature, park, office, home, work, school, joy, love, hope, fear, anger, gratitude, kindness, passion, motivation, inspiration, encouragement, positivity, energy, community, connection, commitment, collaboration, growth, impact, difference, success, progress, change, education, volunteering, veterans, homeless, shelter, meeting, project, system, process, event
|
||||
- Generic categories (even capitalized): Users, Customers, Team, Family, Kids, Veterans, Community
|
||||
- Abstract concepts: motivation, inspiration, gratitude, commitment, resilience
|
||||
|
||||
THE TEST: Is this a specific name you'd find in a contact list or org chart?
|
||||
- "John" → YES (promote=true)
|
||||
- "kids" → NO (promote=false)
|
||||
- "community" → NO (promote=false)
|
||||
|
||||
When in doubt, set promote=false."""
|
||||
|
||||
try:
|
||||
result = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": get_mission_filter_system_message()},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
response_format=MissionFilterResponse,
|
||||
scope="mental_model_emergent_evaluation",
|
||||
)
|
||||
|
||||
# Build name -> promote map
|
||||
promote_map = {c.name: c.promote for c in result.candidates}
|
||||
|
||||
# Find models to remove
|
||||
models_to_remove = []
|
||||
for model in models:
|
||||
name = model["name"]
|
||||
if name in promote_map:
|
||||
if not promote_map[name]:
|
||||
models_to_remove.append(model["id"])
|
||||
else:
|
||||
logger.debug(f"[EMERGENT] Keeping '{name}'")
|
||||
else:
|
||||
# Model not in response - remove to be safe
|
||||
logger.info(f"[EMERGENT] '{name}' not in evaluation response, marking for removal")
|
||||
models_to_remove.append(model["id"])
|
||||
|
||||
logger.info(f"[EMERGENT] Evaluation: {len(models_to_remove)}/{len(models)} emergent models marked for removal")
|
||||
return models_to_remove
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"[EMERGENT] Evaluation failed, keeping all models: {e}")
|
||||
return []
|
||||
|
||||
|
||||
async def detect_entity_candidates(
|
||||
pool,
|
||||
bank_id: str,
|
||||
min_mentions: int = 5,
|
||||
top_percent: int = 20,
|
||||
) -> list[EmergentCandidate]:
|
||||
"""
|
||||
Detect entities that are candidates for promotion to mental models.
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
bank_id: Bank identifier
|
||||
min_mentions: Minimum mention count to consider
|
||||
top_percent: Only consider top X% by mention count
|
||||
|
||||
Returns:
|
||||
List of entity candidates
|
||||
"""
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
|
||||
candidates = []
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Get entities that meet criteria and don't already have mental models
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
WITH ranked AS (
|
||||
SELECT
|
||||
e.id,
|
||||
e.canonical_name,
|
||||
e.mention_count,
|
||||
PERCENT_RANK() OVER (ORDER BY e.mention_count DESC) as rank_pct
|
||||
FROM {fq_table("entities")} e
|
||||
LEFT JOIN {fq_table("mental_models")} mm
|
||||
ON mm.entity_id = e.id AND mm.bank_id = e.bank_id
|
||||
WHERE e.bank_id = $1
|
||||
AND e.mention_count >= $2
|
||||
AND mm.id IS NULL -- Not already a mental model
|
||||
)
|
||||
SELECT id, canonical_name, mention_count
|
||||
FROM ranked
|
||||
WHERE rank_pct <= $3
|
||||
ORDER BY mention_count DESC
|
||||
LIMIT 50
|
||||
""",
|
||||
bank_id,
|
||||
min_mentions,
|
||||
top_percent / 100.0,
|
||||
)
|
||||
|
||||
for row in rows:
|
||||
candidates.append(
|
||||
EmergentCandidate(
|
||||
name=row["canonical_name"],
|
||||
detection_method="named_entity_extraction",
|
||||
mention_count=row["mention_count"],
|
||||
entity_id=str(row["id"]),
|
||||
relevance_score=0.0,
|
||||
)
|
||||
)
|
||||
|
||||
logger.debug(f"[EMERGENT] Detected {len(candidates)} entity candidates")
|
||||
return candidates
|
||||
@@ -0,0 +1,98 @@
|
||||
"""
|
||||
Pydantic models for mental models.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class MentalModelSubtype(str, Enum):
|
||||
"""Subtype of mental model - how it was created."""
|
||||
|
||||
STRUCTURAL = "structural" # Derived from mission, created upfront
|
||||
EMERGENT = "emergent" # Discovered from data patterns
|
||||
LEARNED = "learned" # Formed through reflection
|
||||
PINNED = "pinned" # User-defined topic, observations LLM-generated
|
||||
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"
|
||||
)
|
||||
|
||||
|
||||
class StructuralModelTemplate(BaseModel):
|
||||
"""
|
||||
A template for a structural mental model.
|
||||
|
||||
Generated by LLM based on the bank's mission. Represents what any agent
|
||||
with this role would need to track.
|
||||
"""
|
||||
|
||||
id: str = Field(default="", description="Existing model ID to keep, or empty for new models")
|
||||
name: str = Field(description="Human-readable name")
|
||||
description: str = Field(description="What this model should track")
|
||||
initial_probes: list[str] = Field(default_factory=list, description="Initial search queries to populate this model")
|
||||
|
||||
|
||||
class StructuralModelDerivationResponse(BaseModel):
|
||||
"""Response from LLM for structural model derivation."""
|
||||
|
||||
templates: list[StructuralModelTemplate] = Field(description="Structural model templates derived from the mission")
|
||||
|
||||
|
||||
class EmergentCandidate(BaseModel):
|
||||
"""
|
||||
A candidate for promotion to emergent mental model.
|
||||
|
||||
Detected through pattern analysis of facts.
|
||||
"""
|
||||
|
||||
name: str = Field(description="Name of the detected pattern/entity")
|
||||
detection_method: str = Field(description="How this candidate was detected")
|
||||
mention_count: int = Field(default=0, description="How many times referenced")
|
||||
entity_id: str | None = Field(default=None, description="Entity ID if detected as entity")
|
||||
relevance_score: float = Field(default=0.0, description="Score from mission filter (0-1)")
|
||||
|
||||
|
||||
class ResearchResult(BaseModel):
|
||||
"""
|
||||
Result from the research endpoint.
|
||||
|
||||
Contains the answer along with the mental models and facts used.
|
||||
"""
|
||||
|
||||
answer: str = Field(description="The synthesized answer")
|
||||
mental_models_used: list[str] = Field(default_factory=list, description="IDs of mental models that contributed")
|
||||
facts_used: list[str] = Field(default_factory=list, description="Fact IDs that contributed")
|
||||
question_type: str | None = Field(default=None, description="Detected question type (WHO, WHAT, HOW, etc.)")
|
||||
@@ -0,0 +1,228 @@
|
||||
"""
|
||||
Structural mental model derivation from bank mission.
|
||||
|
||||
Structural models are derived from the bank's mission - they represent what
|
||||
any agent with this role would need to track. For example:
|
||||
|
||||
Mission: "Be a PM for engineering team"
|
||||
Structural models:
|
||||
- Team Structure (who's on the team, roles)
|
||||
- Project Overview (current projects, status)
|
||||
- Processes (how releases work, how decisions are made)
|
||||
- Key Systems (what we own, dependencies)
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from .models import StructuralModelTemplate
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..llm_wrapper import LLMConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class StructuralDerivationResponse(BaseModel):
|
||||
"""Response from LLM for structural model derivation."""
|
||||
|
||||
templates: list[StructuralModelTemplate] = Field(description="Structural model templates derived from the mission")
|
||||
|
||||
|
||||
class StructuralRelevanceResult(BaseModel):
|
||||
"""Result of evaluating a structural model's relevance to the mission."""
|
||||
|
||||
name: str
|
||||
relevant: bool
|
||||
reason: str
|
||||
|
||||
|
||||
class StructuralRelevanceResponse(BaseModel):
|
||||
"""Response from LLM for structural model relevance evaluation."""
|
||||
|
||||
models: list[StructuralRelevanceResult] = Field(description="Relevance evaluation for each model")
|
||||
|
||||
|
||||
def build_structural_derivation_prompt(mission: str, existing_models: list[dict] | None = None) -> str:
|
||||
"""Build the prompt for deriving structural models from a mission."""
|
||||
existing_section = ""
|
||||
if existing_models:
|
||||
model_list = "\n".join([f"- id='{m['id']}' name='{m['name']}': {m['description']}" for m in existing_models])
|
||||
existing_section = f"""
|
||||
EXISTING STRUCTURAL MODELS:
|
||||
{model_list}
|
||||
|
||||
IMPORTANT: If keeping an existing model, you MUST return its EXACT 'id' value.
|
||||
Models not included in your output will be REMOVED.
|
||||
"""
|
||||
|
||||
return f"""Given this agent mission, identify the KEY THINGS to track to achieve it.
|
||||
|
||||
MISSION: {mission}
|
||||
{existing_section}
|
||||
IMPORTANT CONSTRAINTS:
|
||||
- Return 0-3 structural models MAXIMUM (less is better!)
|
||||
- Only include models for SPECIFIC, CONCRETE things the agent needs to track
|
||||
- Each model must be DIRECTLY tied to achieving the mission
|
||||
- If the mission is simple, return 0 models (empty array is fine)
|
||||
- If existing models are provided and you want to keep one, use its EXACT id
|
||||
- Do NOT create near-duplicates (e.g., don't create "topic-map" if "topic-connections" exists)
|
||||
|
||||
GOOD examples (specific, actionable):
|
||||
- Mission: "Be a PM for engineering team" → "Team Members" (track who's on the team)
|
||||
- Mission: "Track customer feedback" → "Customer Issues" (track specific complaints/requests)
|
||||
- Mission: "Manage project X" → "Project X Milestones" (track progress)
|
||||
|
||||
BAD examples (too generic, don't create these):
|
||||
- "Processes", "Workflows", "Key Systems", "Important Events"
|
||||
- "Communication", "Collaboration", "Progress", "Status"
|
||||
- Generic role-based models not tied to the specific mission
|
||||
|
||||
For each model:
|
||||
1. id: Use EXACT existing id if keeping a model, or leave empty for new models
|
||||
2. name: Short, specific name (e.g., "Team Members", "Sprint Goals")
|
||||
3. description: One line describing what to track
|
||||
4. initial_probes: 2-3 search queries to find relevant information
|
||||
|
||||
Return ONLY the models that should exist. Existing models not in your output will be deleted."""
|
||||
|
||||
|
||||
def get_structural_derivation_system_message() -> str:
|
||||
"""System message for structural model derivation."""
|
||||
return """You identify the key things to track for a mission. Be VERY selective.
|
||||
|
||||
Rules:
|
||||
- Maximum 3 models (prefer fewer)
|
||||
- Only SPECIFIC, CONCRETE things - not generic categories
|
||||
- Each must DIRECTLY help achieve the mission
|
||||
- Empty array is valid if no models are truly needed
|
||||
- If existing models are shown and you want to keep one, return its EXACT id
|
||||
- Never create duplicates - if a similar model exists, keep the existing one
|
||||
|
||||
Output JSON with 'templates' array (can be empty)."""
|
||||
|
||||
|
||||
def _normalize_id(text: str) -> str:
|
||||
"""Normalize a string to a canonical form for comparison.
|
||||
|
||||
Removes common suffixes, pluralization, and normalizes separators.
|
||||
"""
|
||||
# Lowercase and normalize separators
|
||||
normalized = text.lower().replace(" ", "-").replace("_", "-")
|
||||
|
||||
# Remove common suffixes that indicate the same concept
|
||||
suffixes_to_remove = ["-map", "-list", "-overview", "-tracker", "-s"]
|
||||
for suffix in suffixes_to_remove:
|
||||
if normalized.endswith(suffix) and len(normalized) > len(suffix):
|
||||
normalized = normalized[: -len(suffix)]
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def _find_similar_existing_id(new_id: str, existing_models: list[dict]) -> str | None:
|
||||
"""Find an existing model ID that is similar to the new ID.
|
||||
|
||||
Returns the existing ID if a similar one is found, None otherwise.
|
||||
"""
|
||||
if not existing_models:
|
||||
return None
|
||||
|
||||
new_normalized = _normalize_id(new_id)
|
||||
|
||||
for model in existing_models:
|
||||
existing_id = model.get("id", "")
|
||||
existing_normalized = _normalize_id(existing_id)
|
||||
|
||||
# Check if one is a prefix of the other (normalized)
|
||||
if new_normalized.startswith(existing_normalized) or existing_normalized.startswith(new_normalized):
|
||||
return existing_id
|
||||
|
||||
# Check if they're the same when normalized
|
||||
if new_normalized == existing_normalized:
|
||||
return existing_id
|
||||
|
||||
return None
|
||||
|
||||
|
||||
async def derive_structural_models(
|
||||
llm_config: "LLMConfig",
|
||||
mission: str,
|
||||
existing_models: list[dict] | None = None,
|
||||
) -> tuple[list[StructuralModelTemplate], list[str]]:
|
||||
"""
|
||||
Derive structural model templates from a bank's mission.
|
||||
|
||||
This combines derivation and evaluation in one call. The LLM sees existing
|
||||
models and decides which to keep. Any existing model not in the output
|
||||
will be marked for removal.
|
||||
|
||||
Args:
|
||||
llm_config: LLM configuration for calling the model
|
||||
mission: The bank's mission (e.g., "Be a PM for engineering team")
|
||||
existing_models: Optional list of existing model dicts with 'name', 'description', 'id'
|
||||
|
||||
Returns:
|
||||
Tuple of (templates to create/keep, IDs of existing models to remove)
|
||||
|
||||
Raises:
|
||||
Exception: If LLM call fails
|
||||
"""
|
||||
prompt = build_structural_derivation_prompt(mission, existing_models)
|
||||
|
||||
result = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": get_structural_derivation_system_message()},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
response_format=StructuralDerivationResponse,
|
||||
scope="mental_model_structural_derivation",
|
||||
)
|
||||
|
||||
templates = result.templates
|
||||
logger.info(f"[STRUCTURAL] LLM returned {len(templates)} structural models")
|
||||
|
||||
# Build set of existing IDs for quick lookup
|
||||
existing_ids = {m["id"] for m in existing_models} if existing_models else set()
|
||||
|
||||
# Process templates: validate IDs, deduplicate, assign stable IDs
|
||||
processed_templates: list[StructuralModelTemplate] = []
|
||||
kept_existing_ids: set[str] = set()
|
||||
|
||||
for template in templates:
|
||||
# If LLM returned an ID, check if it's a valid existing ID
|
||||
if template.id and template.id in existing_ids:
|
||||
# LLM is keeping an existing model
|
||||
kept_existing_ids.add(template.id)
|
||||
processed_templates.append(template)
|
||||
logger.info(f"[STRUCTURAL] Keeping existing model: {template.id}")
|
||||
else:
|
||||
# New model or LLM didn't return a valid ID
|
||||
# Generate ID from name
|
||||
generated_id = template.name.lower().replace(" ", "-").replace("_", "-")
|
||||
|
||||
# Check for similar existing models to prevent near-duplicates
|
||||
similar_id = _find_similar_existing_id(generated_id, existing_models)
|
||||
if similar_id and similar_id not in kept_existing_ids:
|
||||
# Use the existing similar model instead of creating a new one
|
||||
logger.info(f"[STRUCTURAL] Detected near-duplicate: '{generated_id}' matches existing '{similar_id}'")
|
||||
template.id = similar_id
|
||||
kept_existing_ids.add(similar_id)
|
||||
else:
|
||||
template.id = generated_id
|
||||
|
||||
processed_templates.append(template)
|
||||
|
||||
# Find existing models to remove (not kept in LLM output)
|
||||
models_to_remove = []
|
||||
if existing_models:
|
||||
for model in existing_models:
|
||||
if model["id"] not in kept_existing_ids:
|
||||
logger.info(f"[STRUCTURAL] Marking '{model['name']}' (id={model['id']}) for removal")
|
||||
models_to_remove.append(model["id"])
|
||||
|
||||
if models_to_remove:
|
||||
logger.info(f"[STRUCTURAL] {len(models_to_remove)} existing models will be removed")
|
||||
|
||||
return processed_templates, models_to_remove
|
||||
@@ -4,11 +4,12 @@ Query analysis abstraction for the memory system.
|
||||
Provides an interface for analyzing natural language queries to extract
|
||||
structured information like temporal constraints.
|
||||
"""
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
import logging
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -20,6 +21,7 @@ class TemporalConstraint(BaseModel):
|
||||
|
||||
Represents a time range with start and end dates.
|
||||
"""
|
||||
|
||||
start_date: datetime = Field(description="Start of the time range (inclusive)")
|
||||
end_date: datetime = Field(description="End of the time range (inclusive)")
|
||||
|
||||
@@ -33,9 +35,9 @@ class QueryAnalysis(BaseModel):
|
||||
|
||||
Contains extracted structured information like temporal constraints.
|
||||
"""
|
||||
temporal_constraint: Optional[TemporalConstraint] = Field(
|
||||
default=None,
|
||||
description="Extracted temporal constraint, if any"
|
||||
|
||||
temporal_constraint: TemporalConstraint | None = Field(
|
||||
default=None, description="Extracted temporal constraint, if any"
|
||||
)
|
||||
|
||||
|
||||
@@ -58,9 +60,7 @@ class QueryAnalyzer(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def analyze(
|
||||
self, query: str, reference_date: Optional[datetime] = None
|
||||
) -> QueryAnalysis:
|
||||
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
|
||||
"""
|
||||
Analyze a natural language query.
|
||||
|
||||
@@ -84,7 +84,7 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
|
||||
Performance:
|
||||
- ~10-50ms per query
|
||||
- No model loading required
|
||||
- No model loading required (lazy import on first use)
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
@@ -95,11 +95,10 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
"""Load dateparser (lazy import)."""
|
||||
if self._search_dates is None:
|
||||
from dateparser.search import search_dates
|
||||
|
||||
self._search_dates = search_dates
|
||||
|
||||
def analyze(
|
||||
self, query: str, reference_date: Optional[datetime] = None
|
||||
) -> QueryAnalysis:
|
||||
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
|
||||
"""
|
||||
Analyze query using dateparser.
|
||||
|
||||
@@ -113,8 +112,6 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
Returns:
|
||||
QueryAnalysis with temporal_constraint if found
|
||||
"""
|
||||
self.load()
|
||||
|
||||
if reference_date is None:
|
||||
reference_date = datetime.now()
|
||||
|
||||
@@ -124,11 +121,14 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
if period_result is not None:
|
||||
return QueryAnalysis(temporal_constraint=period_result)
|
||||
|
||||
# Lazy load dateparser (only imports on first call, then cached)
|
||||
self.load()
|
||||
|
||||
# Use dateparser's search_dates to find temporal expressions
|
||||
settings = {
|
||||
'RELATIVE_BASE': reference_date,
|
||||
'PREFER_DATES_FROM': 'past',
|
||||
'RETURN_AS_TIMEZONE_AWARE': False,
|
||||
"RELATIVE_BASE": reference_date,
|
||||
"PREFER_DATES_FROM": "past",
|
||||
"RETURN_AS_TIMEZONE_AWARE": False,
|
||||
}
|
||||
|
||||
results = self._search_dates(query, settings=settings)
|
||||
@@ -137,11 +137,8 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
return QueryAnalysis(temporal_constraint=None)
|
||||
|
||||
# Filter out false positives (common words parsed as dates)
|
||||
false_positives = {'do', 'may', 'march', 'will', 'can', 'sat', 'sun', 'mon', 'tue', 'wed', 'thu', 'fri'}
|
||||
valid_results = [
|
||||
(text, date) for text, date in results
|
||||
if text.lower() not in false_positives or len(text) > 3
|
||||
]
|
||||
false_positives = {"do", "may", "march", "will", "can", "sat", "sun", "mon", "tue", "wed", "thu", "fri"}
|
||||
valid_results = [(text, date) for text, date in results if text.lower() not in false_positives or len(text) > 3]
|
||||
|
||||
if not valid_results:
|
||||
return QueryAnalysis(temporal_constraint=None)
|
||||
@@ -153,84 +150,94 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
start_date = parsed_date.replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
end_date = parsed_date.replace(hour=23, minute=59, second=59, microsecond=999999)
|
||||
|
||||
return QueryAnalysis(
|
||||
temporal_constraint=TemporalConstraint(
|
||||
start_date=start_date,
|
||||
end_date=end_date
|
||||
)
|
||||
)
|
||||
return QueryAnalysis(temporal_constraint=TemporalConstraint(start_date=start_date, end_date=end_date))
|
||||
|
||||
def _extract_period(
|
||||
self, query: str, reference_date: datetime
|
||||
) -> Optional[TemporalConstraint]:
|
||||
def _extract_period(self, query: str, reference_date: datetime) -> TemporalConstraint | None:
|
||||
"""
|
||||
Extract period-based temporal expressions (week, month, year, weekend).
|
||||
|
||||
These need special handling as they represent date ranges, not single dates.
|
||||
Supports multiple languages.
|
||||
"""
|
||||
|
||||
def constraint(start: datetime, end: datetime) -> TemporalConstraint:
|
||||
return TemporalConstraint(
|
||||
start_date=start.replace(hour=0, minute=0, second=0, microsecond=0),
|
||||
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999)
|
||||
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999),
|
||||
)
|
||||
|
||||
# Yesterday patterns (English, Spanish, Italian, French, German)
|
||||
if re.search(r'\b(yesterday|ayer|ieri|hier|gestern)\b', query, re.IGNORECASE):
|
||||
if re.search(r"\b(yesterday|ayer|ieri|hier|gestern)\b", query, re.IGNORECASE):
|
||||
d = reference_date - timedelta(days=1)
|
||||
return constraint(d, d)
|
||||
|
||||
# Today patterns
|
||||
if re.search(r'\b(today|hoy|oggi|aujourd\'?hui|heute)\b', query, re.IGNORECASE):
|
||||
if re.search(r"\b(today|hoy|oggi|aujourd\'?hui|heute)\b", query, re.IGNORECASE):
|
||||
return constraint(reference_date, reference_date)
|
||||
|
||||
# "a couple of days ago" / "a few days ago" patterns
|
||||
# These are imprecise so we create a range
|
||||
if re.search(r'\b(a\s+)?couple\s+(of\s+)?days?\s+ago\b', query, re.IGNORECASE):
|
||||
if re.search(r"\b(a\s+)?couple\s+(of\s+)?days?\s+ago\b", query, re.IGNORECASE):
|
||||
# "a couple of days" = approximately 2 days, give range of 1-3 days
|
||||
return constraint(reference_date - timedelta(days=3), reference_date - timedelta(days=1))
|
||||
|
||||
if re.search(r'\b(a\s+)?few\s+days?\s+ago\b', query, re.IGNORECASE):
|
||||
if re.search(r"\b(a\s+)?few\s+days?\s+ago\b", query, re.IGNORECASE):
|
||||
# "a few days" = approximately 3-4 days, give range of 2-5 days
|
||||
return constraint(reference_date - timedelta(days=5), reference_date - timedelta(days=2))
|
||||
|
||||
# "a couple of weeks ago" / "a few weeks ago" patterns
|
||||
if re.search(r'\b(a\s+)?couple\s+(of\s+)?weeks?\s+ago\b', query, re.IGNORECASE):
|
||||
if re.search(r"\b(a\s+)?couple\s+(of\s+)?weeks?\s+ago\b", query, re.IGNORECASE):
|
||||
# "a couple of weeks" = approximately 2 weeks, give range of 1-3 weeks
|
||||
return constraint(reference_date - timedelta(weeks=3), reference_date - timedelta(weeks=1))
|
||||
|
||||
if re.search(r'\b(a\s+)?few\s+weeks?\s+ago\b', query, re.IGNORECASE):
|
||||
if re.search(r"\b(a\s+)?few\s+weeks?\s+ago\b", query, re.IGNORECASE):
|
||||
# "a few weeks" = approximately 3-4 weeks, give range of 2-5 weeks
|
||||
return constraint(reference_date - timedelta(weeks=5), reference_date - timedelta(weeks=2))
|
||||
|
||||
# "a couple of months ago" / "a few months ago" patterns
|
||||
if re.search(r'\b(a\s+)?couple\s+(of\s+)?months?\s+ago\b', query, re.IGNORECASE):
|
||||
if re.search(r"\b(a\s+)?couple\s+(of\s+)?months?\s+ago\b", query, re.IGNORECASE):
|
||||
# "a couple of months" = approximately 2 months, give range of 1-3 months
|
||||
return constraint(reference_date - timedelta(days=90), reference_date - timedelta(days=30))
|
||||
|
||||
if re.search(r'\b(a\s+)?few\s+months?\s+ago\b', query, re.IGNORECASE):
|
||||
if re.search(r"\b(a\s+)?few\s+months?\s+ago\b", query, re.IGNORECASE):
|
||||
# "a few months" = approximately 3-4 months, give range of 2-5 months
|
||||
return constraint(reference_date - timedelta(days=150), reference_date - timedelta(days=60))
|
||||
|
||||
# Last week patterns (English, Spanish, Italian, French, German)
|
||||
if re.search(r'\b(last\s+week|la\s+semana\s+pasada|la\s+settimana\s+scorsa|la\s+semaine\s+derni[eè]re|letzte\s+woche)\b', query, re.IGNORECASE):
|
||||
if re.search(
|
||||
r"\b(last\s+week|la\s+semana\s+pasada|la\s+settimana\s+scorsa|la\s+semaine\s+derni[eè]re|letzte\s+woche)\b",
|
||||
query,
|
||||
re.IGNORECASE,
|
||||
):
|
||||
start = reference_date - timedelta(days=reference_date.weekday() + 7)
|
||||
return constraint(start, start + timedelta(days=6))
|
||||
|
||||
# Last month patterns
|
||||
if re.search(r'\b(last\s+month|el\s+mes\s+pasado|il\s+mese\s+scorso|le\s+mois\s+dernier|letzten?\s+monat)\b', query, re.IGNORECASE):
|
||||
if re.search(
|
||||
r"\b(last\s+month|el\s+mes\s+pasado|il\s+mese\s+scorso|le\s+mois\s+dernier|letzten?\s+monat)\b",
|
||||
query,
|
||||
re.IGNORECASE,
|
||||
):
|
||||
first = reference_date.replace(day=1)
|
||||
end = first - timedelta(days=1)
|
||||
start = end.replace(day=1)
|
||||
return constraint(start, end)
|
||||
|
||||
# Last year patterns
|
||||
if re.search(r'\b(last\s+year|el\s+a[ñn]o\s+pasado|l\'anno\s+scorso|l\'ann[ée]e\s+derni[eè]re|letztes?\s+jahr)\b', query, re.IGNORECASE):
|
||||
if re.search(
|
||||
r"\b(last\s+year|el\s+a[ñn]o\s+pasado|l\'anno\s+scorso|l\'ann[ée]e\s+derni[eè]re|letztes?\s+jahr)\b",
|
||||
query,
|
||||
re.IGNORECASE,
|
||||
):
|
||||
year = reference_date.year - 1
|
||||
return constraint(datetime(year, 1, 1), datetime(year, 12, 31))
|
||||
|
||||
# Last weekend patterns
|
||||
if re.search(r'\b(last\s+weekend|el\s+fin\s+de\s+semana\s+pasado|lo\s+scorso\s+fine\s+settimana|le\s+week-?end\s+dernier|letztes?\s+wochenende)\b', query, re.IGNORECASE):
|
||||
if re.search(
|
||||
r"\b(last\s+weekend|el\s+fin\s+de\s+semana\s+pasado|lo\s+scorso\s+fine\s+settimana|le\s+week-?end\s+dernier|letztes?\s+wochenende)\b",
|
||||
query,
|
||||
re.IGNORECASE,
|
||||
):
|
||||
days_since_sat = (reference_date.weekday() + 2) % 7
|
||||
if days_since_sat == 0:
|
||||
days_since_sat = 7
|
||||
@@ -239,22 +246,22 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
|
||||
# Month + Year patterns (e.g., "June 2024", "junio 2024", "giugno 2024")
|
||||
month_patterns = {
|
||||
'january|enero|gennaio|janvier|januar': 1,
|
||||
'february|febrero|febbraio|f[ée]vrier|februar': 2,
|
||||
'march|marzo|mars|m[äa]rz': 3,
|
||||
'april|abril|aprile|avril': 4,
|
||||
'may|mayo|maggio|mai': 5,
|
||||
'june|junio|giugno|juin|juni': 6,
|
||||
'july|julio|luglio|juillet|juli': 7,
|
||||
'august|agosto|ao[uû]t': 8,
|
||||
'september|septiembre|settembre|septembre': 9,
|
||||
'october|octubre|ottobre|octobre|oktober': 10,
|
||||
'november|noviembre|novembre': 11,
|
||||
'december|diciembre|dicembre|d[ée]cembre|dezember': 12,
|
||||
"january|enero|gennaio|janvier|januar": 1,
|
||||
"february|febrero|febbraio|f[ée]vrier|februar": 2,
|
||||
"march|marzo|mars|m[äa]rz": 3,
|
||||
"april|abril|aprile|avril": 4,
|
||||
"may|mayo|maggio|mai": 5,
|
||||
"june|junio|giugno|juin|juni": 6,
|
||||
"july|julio|luglio|juillet|juli": 7,
|
||||
"august|agosto|ao[uû]t": 8,
|
||||
"september|septiembre|settembre|septembre": 9,
|
||||
"october|octubre|ottobre|octobre|oktober": 10,
|
||||
"november|noviembre|novembre": 11,
|
||||
"december|diciembre|dicembre|d[ée]cembre|dezember": 12,
|
||||
}
|
||||
|
||||
for pattern, month_num in month_patterns.items():
|
||||
match = re.search(rf'\b({pattern})\s+(\d{{4}})\b', query, re.IGNORECASE)
|
||||
match = re.search(rf"\b({pattern})\s+(\d{{4}})\b", query, re.IGNORECASE)
|
||||
if match:
|
||||
year = int(match.group(2))
|
||||
start = datetime(year, month_num, 1)
|
||||
@@ -279,11 +286,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
|
||||
- Model size: ~80M params (~300MB download)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str = "google/flan-t5-small",
|
||||
device: str = "cpu"
|
||||
):
|
||||
def __init__(self, model_name: str = "google/flan-t5-small", device: str = "cpu"):
|
||||
"""
|
||||
Initialize T5 query analyzer.
|
||||
|
||||
@@ -304,11 +307,10 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
|
||||
return
|
||||
|
||||
try:
|
||||
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"transformers is required for TransformerQueryAnalyzer. "
|
||||
"Install it with: pip install transformers"
|
||||
"transformers is required for TransformerQueryAnalyzer. Install it with: pip install transformers"
|
||||
)
|
||||
|
||||
logger.info(f"Loading query analyzer model: {self.model_name}...")
|
||||
@@ -322,9 +324,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
|
||||
"""Lazy load the T5 model for temporal extraction (calls load())."""
|
||||
self.load()
|
||||
|
||||
def _extract_with_rules(
|
||||
self, query: str, reference_date: datetime
|
||||
) -> Optional[TemporalConstraint]:
|
||||
def _extract_with_rules(self, query: str, reference_date: datetime) -> TemporalConstraint | None:
|
||||
"""
|
||||
Extract temporal expressions using rule-based patterns.
|
||||
|
||||
@@ -332,6 +332,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
|
||||
patterns that need model-based extraction.
|
||||
"""
|
||||
import re
|
||||
|
||||
query_lower = query.lower()
|
||||
|
||||
def get_last_weekday(weekday: int) -> datetime:
|
||||
@@ -343,50 +344,60 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
|
||||
def constraint(start: datetime, end: datetime) -> TemporalConstraint:
|
||||
return TemporalConstraint(
|
||||
start_date=start.replace(hour=0, minute=0, second=0, microsecond=0),
|
||||
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999)
|
||||
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999),
|
||||
)
|
||||
|
||||
# Yesterday
|
||||
if re.search(r'\byesterday\b', query_lower):
|
||||
if re.search(r"\byesterday\b", query_lower):
|
||||
d = reference_date - timedelta(days=1)
|
||||
return constraint(d, d)
|
||||
|
||||
# Last week
|
||||
if re.search(r'\blast\s+week\b', query_lower):
|
||||
if re.search(r"\blast\s+week\b", query_lower):
|
||||
start = reference_date - timedelta(days=reference_date.weekday() + 7)
|
||||
return constraint(start, start + timedelta(days=6))
|
||||
|
||||
# Last month
|
||||
if re.search(r'\blast\s+month\b', query_lower):
|
||||
if re.search(r"\blast\s+month\b", query_lower):
|
||||
first = reference_date.replace(day=1)
|
||||
end = first - timedelta(days=1)
|
||||
start = end.replace(day=1)
|
||||
return constraint(start, end)
|
||||
|
||||
# Last year
|
||||
if re.search(r'\blast\s+year\b', query_lower):
|
||||
if re.search(r"\blast\s+year\b", query_lower):
|
||||
y = reference_date.year - 1
|
||||
return constraint(datetime(y, 1, 1), datetime(y, 12, 31))
|
||||
|
||||
# Last weekend
|
||||
if re.search(r'\blast\s+weekend\b', query_lower):
|
||||
if re.search(r"\blast\s+weekend\b", query_lower):
|
||||
sat = get_last_weekday(5)
|
||||
return constraint(sat, sat + timedelta(days=1))
|
||||
|
||||
# Last <weekday>
|
||||
weekdays = {'monday': 0, 'tuesday': 1, 'wednesday': 2, 'thursday': 3,
|
||||
'friday': 4, 'saturday': 5, 'sunday': 6}
|
||||
weekdays = {"monday": 0, "tuesday": 1, "wednesday": 2, "thursday": 3, "friday": 4, "saturday": 5, "sunday": 6}
|
||||
for name, num in weekdays.items():
|
||||
if re.search(rf'\blast\s+{name}\b', query_lower):
|
||||
if re.search(rf"\blast\s+{name}\b", query_lower):
|
||||
d = get_last_weekday(num)
|
||||
return constraint(d, d)
|
||||
|
||||
# Month + Year: "June 2024", "in March 2023"
|
||||
months = {'january': 1, 'february': 2, 'march': 3, 'april': 4, 'may': 5,
|
||||
'june': 6, 'july': 7, 'august': 8, 'september': 9, 'october': 10,
|
||||
'november': 11, 'december': 12}
|
||||
months = {
|
||||
"january": 1,
|
||||
"february": 2,
|
||||
"march": 3,
|
||||
"april": 4,
|
||||
"may": 5,
|
||||
"june": 6,
|
||||
"july": 7,
|
||||
"august": 8,
|
||||
"september": 9,
|
||||
"october": 10,
|
||||
"november": 11,
|
||||
"december": 12,
|
||||
}
|
||||
for name, num in months.items():
|
||||
match = re.search(rf'\b{name}\s+(\d{{4}})\b', query_lower)
|
||||
match = re.search(rf"\b{name}\s+(\d{{4}})\b", query_lower)
|
||||
if match:
|
||||
year = int(match.group(1))
|
||||
if num == 12:
|
||||
@@ -397,9 +408,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
|
||||
|
||||
return None
|
||||
|
||||
def analyze(
|
||||
self, query: str, reference_date: Optional[datetime] = None
|
||||
) -> QueryAnalysis:
|
||||
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
|
||||
"""
|
||||
Analyze query for temporal expressions.
|
||||
|
||||
@@ -435,11 +444,11 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
|
||||
last_saturday = get_last_weekday(5)
|
||||
|
||||
# Build prompt for T5
|
||||
prompt = f"""Today is {reference_date.strftime('%Y-%m-%d')}. Extract date range or "none".
|
||||
prompt = f"""Today is {reference_date.strftime("%Y-%m-%d")}. Extract date range or "none".
|
||||
|
||||
June 2024 = 2024-06-01 to 2024-06-30
|
||||
yesterday = {yesterday.strftime('%Y-%m-%d')} to {yesterday.strftime('%Y-%m-%d')}
|
||||
last Saturday = {last_saturday.strftime('%Y-%m-%d')} to {last_saturday.strftime('%Y-%m-%d')}
|
||||
yesterday = {yesterday.strftime("%Y-%m-%d")} to {yesterday.strftime("%Y-%m-%d")}
|
||||
last Saturday = {last_saturday.strftime("%Y-%m-%d")} to {last_saturday.strftime("%Y-%m-%d")}
|
||||
what is the weather = none
|
||||
{query} ="""
|
||||
|
||||
@@ -448,13 +457,7 @@ what is the weather = none
|
||||
inputs = {k: v.to(self.device) for k, v in inputs.items()}
|
||||
|
||||
with self._no_grad():
|
||||
outputs = self._model.generate(
|
||||
**inputs,
|
||||
max_new_tokens=30,
|
||||
num_beams=3,
|
||||
do_sample=False,
|
||||
temperature=1.0
|
||||
)
|
||||
outputs = self._model.generate(**inputs, max_new_tokens=30, num_beams=3, do_sample=False, temperature=1.0)
|
||||
|
||||
result = self._tokenizer.decode(outputs[0], skip_special_tokens=True).strip()
|
||||
|
||||
@@ -466,14 +469,14 @@ what is the weather = none
|
||||
"""Get torch.no_grad context manager."""
|
||||
try:
|
||||
import torch
|
||||
|
||||
return torch.no_grad()
|
||||
except ImportError:
|
||||
from contextlib import nullcontext
|
||||
|
||||
return nullcontext()
|
||||
|
||||
def _parse_generated_output(
|
||||
self, result: str, reference_date: datetime
|
||||
) -> Optional[TemporalConstraint]:
|
||||
def _parse_generated_output(self, result: str, reference_date: datetime) -> TemporalConstraint | None:
|
||||
"""
|
||||
Parse T5 generated output into TemporalConstraint.
|
||||
|
||||
@@ -492,7 +495,8 @@ what is the weather = none
|
||||
try:
|
||||
# Parse "YYYY-MM-DD to YYYY-MM-DD"
|
||||
import re
|
||||
pattern = r'(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})'
|
||||
|
||||
pattern = r"(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})"
|
||||
match = re.search(pattern, result, re.IGNORECASE)
|
||||
|
||||
if match:
|
||||
@@ -513,7 +517,7 @@ what is the weather = none
|
||||
|
||||
return TemporalConstraint(start_date=start_date, end_date=end_date)
|
||||
|
||||
except (ValueError, AttributeError) as e:
|
||||
except (ValueError, AttributeError):
|
||||
return None
|
||||
|
||||
return None
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
"""
|
||||
Reflect agent module for agentic reflection with tools.
|
||||
|
||||
The reflect agent uses an iterative loop with tools to:
|
||||
1. Lookup mental models (existing knowledge)
|
||||
2. Recall facts (semantic + temporal search)
|
||||
3. Learn new insights (create/update mental models)
|
||||
4. Expand memories (get chunk/document context)
|
||||
"""
|
||||
|
||||
from .agent import ReflectAgentResult, run_reflect_agent
|
||||
from .models import MentalModelInput, ReflectAction, ReflectActionBatch
|
||||
|
||||
__all__ = [
|
||||
"run_reflect_agent",
|
||||
"ReflectAgentResult",
|
||||
"ReflectAction",
|
||||
"ReflectActionBatch",
|
||||
"MentalModelInput",
|
||||
]
|
||||
@@ -0,0 +1,723 @@
|
||||
"""
|
||||
Reflect agent - agentic loop for reflection with native tool calling.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable
|
||||
|
||||
from .models import DirectiveInfo, LLMCall, MentalModelInput, ReflectAgentResult, 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."""
|
||||
if not directives:
|
||||
return []
|
||||
|
||||
result = []
|
||||
for directive in directives:
|
||||
directive_id = directive.get("id", "")
|
||||
directive_name = directive.get("name", "")
|
||||
observations = directive.get("observations", [])
|
||||
|
||||
rules = []
|
||||
for obs in observations:
|
||||
# Support both Pydantic Observation objects and dicts
|
||||
if hasattr(obs, "content"):
|
||||
rules.append(obs.content)
|
||||
elif isinstance(obs, dict) and obs.get("content"):
|
||||
rules.append(obs["content"])
|
||||
|
||||
result.append(DirectiveInfo(id=directive_id, name=directive_name, rules=rules))
|
||||
|
||||
return result
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..llm_wrapper import LLMProvider
|
||||
from ..response_models import LLMToolCall
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_MAX_ITERATIONS = 10
|
||||
|
||||
|
||||
async def _generate_structured_output(
|
||||
answer: str,
|
||||
response_schema: dict,
|
||||
llm_config: "LLMProvider",
|
||||
reflect_id: str,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Generate structured output from an answer using the provided JSON schema.
|
||||
|
||||
Args:
|
||||
answer: The text answer to extract structured data from
|
||||
response_schema: JSON Schema for the expected output structure
|
||||
llm_config: LLM provider for making the extraction call
|
||||
reflect_id: Reflect ID for logging
|
||||
|
||||
Returns:
|
||||
Structured output dict if successful, None otherwise
|
||||
"""
|
||||
try:
|
||||
from typing import Any as TypingAny
|
||||
|
||||
from pydantic import create_model
|
||||
|
||||
def _json_schema_type_to_python(field_schema: dict) -> type:
|
||||
"""Map JSON schema type to Python type for better LLM guidance."""
|
||||
json_type = field_schema.get("type", "string")
|
||||
if json_type == "array":
|
||||
return list
|
||||
elif json_type == "object":
|
||||
return dict
|
||||
elif json_type == "integer":
|
||||
return int
|
||||
elif json_type == "number":
|
||||
return float
|
||||
elif json_type == "boolean":
|
||||
return bool
|
||||
else:
|
||||
return str
|
||||
|
||||
# Build fields from JSON schema properties
|
||||
schema_props = response_schema.get("properties", {})
|
||||
required_fields = set(response_schema.get("required", []))
|
||||
fields: dict[str, TypingAny] = {}
|
||||
for field_name, field_schema in schema_props.items():
|
||||
field_type = _json_schema_type_to_python(field_schema)
|
||||
default = ... if field_name in required_fields else None
|
||||
fields[field_name] = (field_type, default)
|
||||
|
||||
if not fields:
|
||||
return None
|
||||
|
||||
DynamicModel = create_model("StructuredResponse", **fields)
|
||||
|
||||
# Include the full schema in the prompt for better LLM guidance
|
||||
schema_str = json.dumps(response_schema, indent=2)
|
||||
|
||||
# Call LLM with the answer to extract structured data
|
||||
structured_prompt = f"""Based on this answer, extract the information into the requested structured format.
|
||||
|
||||
Answer: {answer}
|
||||
|
||||
JSON Schema to follow:
|
||||
```json
|
||||
{schema_str}
|
||||
```
|
||||
|
||||
Return ONLY a valid JSON object that matches this exact schema. Pay special attention to field types:
|
||||
- "type": "array" means the value must be a JSON array/list, NOT a string
|
||||
- "type": "string" means the value must be a string
|
||||
- "type": "object" means the value must be a JSON object
|
||||
|
||||
Do not include any explanation, only the JSON object."""
|
||||
|
||||
structured_result = await llm_config.call(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": "Extract structured data from the given answer. Return only valid JSON matching the provided schema exactly.",
|
||||
},
|
||||
{"role": "user", "content": structured_prompt},
|
||||
],
|
||||
response_format=DynamicModel,
|
||||
scope="reflect_structured",
|
||||
skip_validation=True, # We'll handle the dict ourselves
|
||||
)
|
||||
|
||||
# Convert to dict
|
||||
if hasattr(structured_result, "model_dump"):
|
||||
structured_output = structured_result.model_dump()
|
||||
elif isinstance(structured_result, dict):
|
||||
structured_output = structured_result
|
||||
else:
|
||||
# Try to parse as JSON
|
||||
structured_output = json.loads(str(structured_result))
|
||||
|
||||
logger.info(f"[REFLECT {reflect_id}] Generated structured output with {len(structured_output)} fields")
|
||||
return structured_output
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"[REFLECT {reflect_id}] Failed to generate structured output: {e}")
|
||||
return None
|
||||
|
||||
|
||||
async def run_reflect_agent(
|
||||
llm_config: "LLMProvider",
|
||||
bank_id: str,
|
||||
query: str,
|
||||
bank_profile: dict[str, Any],
|
||||
lookup_fn: Callable[[str | None], Awaitable[dict[str, Any]]],
|
||||
recall_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
|
||||
learn_fn: Callable[[MentalModelInput], Awaitable[dict[str, Any]]] | None = None,
|
||||
context: str | None = None,
|
||||
max_iterations: int = DEFAULT_MAX_ITERATIONS,
|
||||
max_tokens: int | None = None,
|
||||
response_schema: dict | None = None,
|
||||
directives: list[dict[str, Any]] | None = None,
|
||||
) -> ReflectAgentResult:
|
||||
"""
|
||||
Execute the reflect agent loop using native tool calling.
|
||||
|
||||
The agent iteratively calls tools to gather information and learn,
|
||||
then provides a final answer via the done() tool.
|
||||
|
||||
Args:
|
||||
llm_config: LLM provider for agent calls
|
||||
bank_id: Bank identifier
|
||||
query: Question to answer
|
||||
bank_profile: Bank profile with name and mission
|
||||
lookup_fn: Tool callback for lookup (model_id) -> result
|
||||
recall_fn: Tool callback for recall (query, max_tokens) -> result
|
||||
expand_fn: Tool callback for expand (memory_id, depth) -> result
|
||||
learn_fn: Optional tool callback for learn (MentalModelInput) -> result.
|
||||
If None, learn tool is disabled.
|
||||
context: Optional additional context
|
||||
max_iterations: Maximum number of iterations before forcing response
|
||||
max_tokens: Maximum tokens for the final response
|
||||
response_schema: Optional JSON Schema for structured output in final response
|
||||
directives: Optional list of directive mental models to inject as hard rules
|
||||
|
||||
Returns:
|
||||
ReflectAgentResult with final answer and metadata
|
||||
"""
|
||||
enable_learn = learn_fn is not None
|
||||
reflect_id = f"{bank_id[:8]}-{int(time.time() * 1000) % 100000}"
|
||||
start_time = time.time()
|
||||
|
||||
# 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(enable_learn=enable_learn, 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)
|
||||
messages: list[dict[str, Any]] = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": query},
|
||||
]
|
||||
|
||||
# Tracking
|
||||
mental_models_created: list[str] = []
|
||||
total_tools_called = 0
|
||||
tool_trace: list[ToolCall] = []
|
||||
tool_trace_summary: list[dict[str, Any]] = []
|
||||
llm_trace: list[dict[str, Any]] = []
|
||||
context_history: list[dict[str, Any]] = [] # For final prompt fallback
|
||||
|
||||
# Track available IDs for validation (prevents hallucinated citations)
|
||||
available_memory_ids: set[str] = set()
|
||||
available_model_ids: set[str] = set()
|
||||
|
||||
# Pre-fetch mental models so the agent always starts with this knowledge
|
||||
prefetch_start = time.time()
|
||||
models_result = await lookup_fn(None) # List all mental models
|
||||
prefetch_duration = int((time.time() - prefetch_start) * 1000)
|
||||
|
||||
# Track available model IDs
|
||||
if isinstance(models_result, dict) and "models" in models_result:
|
||||
for model in models_result["models"]:
|
||||
if "id" in model:
|
||||
available_model_ids.add(model["id"])
|
||||
|
||||
# Add to context history for the agent
|
||||
context_history.append({"tool": "list_mental_models", "output": models_result})
|
||||
|
||||
# Add to tool trace
|
||||
tool_trace.append(
|
||||
ToolCall(
|
||||
tool="list_mental_models",
|
||||
input={"tool": "list_mental_models"},
|
||||
output=models_result,
|
||||
duration_ms=prefetch_duration,
|
||||
iteration=0,
|
||||
)
|
||||
)
|
||||
tool_trace_summary.append(
|
||||
{
|
||||
"tool": "list_mental_models",
|
||||
"input_summary": "(prefetch)",
|
||||
"duration_ms": prefetch_duration,
|
||||
"output_chars": len(json.dumps(models_result, default=str)),
|
||||
}
|
||||
)
|
||||
total_tools_called += 1
|
||||
|
||||
# Include in the user message so the agent sees it
|
||||
models_info = json.dumps(models_result, indent=2, default=str)
|
||||
messages[1]["content"] = f"{query}\n\n## Available Mental Models (pre-fetched)\n```json\n{models_info}\n```"
|
||||
|
||||
def _get_llm_trace() -> list[LLMCall]:
|
||||
return [LLMCall(scope=c["scope"], duration_ms=c["duration_ms"]) for c in llm_trace]
|
||||
|
||||
def _log_completion(answer: str, iterations: int, forced: bool = False):
|
||||
elapsed_ms = int((time.time() - start_time) * 1000)
|
||||
tools_summary = (
|
||||
", ".join(
|
||||
f"{t['tool']}({t['input_summary']})={t['duration_ms']}ms/{t.get('output_chars', 0)}c"
|
||||
for t in tool_trace_summary
|
||||
)
|
||||
or "none"
|
||||
)
|
||||
llm_summary = ", ".join(f"{c['scope']}={c['duration_ms']}ms" for c in llm_trace) or "none"
|
||||
total_llm_ms = sum(c["duration_ms"] for c in llm_trace)
|
||||
total_tools_ms = sum(t["duration_ms"] for t in tool_trace_summary)
|
||||
|
||||
answer_preview = answer[:100] + "..." if len(answer) > 100 else answer
|
||||
mode = "forced" if forced else "done"
|
||||
logger.info(
|
||||
f"[REFLECT {reflect_id}] {mode} | "
|
||||
f"query='{query[:50]}...' | "
|
||||
f"iterations={iterations} | "
|
||||
f"llm=[{llm_summary}] ({total_llm_ms}ms) | "
|
||||
f"tools=[{tools_summary}] ({total_tools_ms}ms) | "
|
||||
f"answer='{answer_preview}' | "
|
||||
f"total={elapsed_ms}ms"
|
||||
)
|
||||
|
||||
for iteration in range(max_iterations):
|
||||
is_last = iteration == max_iterations - 1
|
||||
|
||||
if is_last:
|
||||
# Force text response on last iteration - no tools
|
||||
prompt = build_final_prompt(query, context_history, bank_profile, context)
|
||||
llm_start = time.time()
|
||||
response = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
scope="reflect_agent_final",
|
||||
max_completion_tokens=max_tokens,
|
||||
)
|
||||
llm_trace.append({"scope": "final", "duration_ms": int((time.time() - llm_start) * 1000)})
|
||||
answer = response.strip()
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iteration + 1,
|
||||
tools_called=total_tools_called,
|
||||
mental_models_created=mental_models_created,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
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)
|
||||
llm_trace.append({"scope": f"agent_{iteration + 1}", "duration_ms": llm_duration})
|
||||
|
||||
except Exception:
|
||||
llm_trace.append(
|
||||
{"scope": f"agent_{iteration + 1}_err", "duration_ms": int((time.time() - llm_start) * 1000)}
|
||||
)
|
||||
# Guardrail: If no evidence gathered yet, retry
|
||||
has_gathered_evidence = bool(available_memory_ids) or bool(available_model_ids)
|
||||
if not has_gathered_evidence and iteration < max_iterations - 1:
|
||||
continue
|
||||
prompt = build_final_prompt(query, context_history, bank_profile, context)
|
||||
llm_start = time.time()
|
||||
response = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
scope="reflect_agent_final",
|
||||
max_completion_tokens=max_tokens,
|
||||
)
|
||||
llm_trace.append({"scope": "final", "duration_ms": int((time.time() - llm_start) * 1000)})
|
||||
answer = response.strip()
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iteration + 1,
|
||||
tools_called=total_tools_called,
|
||||
mental_models_created=mental_models_created,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
|
||||
# No tool calls - LLM wants to respond with text
|
||||
if not result.tool_calls:
|
||||
if result.content:
|
||||
answer = result.content.strip()
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output = await _generate_structured_output(
|
||||
answer, response_schema, llm_config, reflect_id
|
||||
)
|
||||
|
||||
_log_completion(answer, iteration + 1)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iteration + 1,
|
||||
tools_called=total_tools_called,
|
||||
mental_models_created=mental_models_created,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
# Empty response, force final
|
||||
prompt = build_final_prompt(query, context_history, bank_profile, context)
|
||||
llm_start = time.time()
|
||||
response = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": FINAL_SYSTEM_PROMPT},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
scope="reflect_agent_final",
|
||||
max_completion_tokens=max_tokens,
|
||||
)
|
||||
llm_trace.append({"scope": "final", "duration_ms": int((time.time() - llm_start) * 1000)})
|
||||
answer = response.strip()
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and answer:
|
||||
structured_output = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
|
||||
_log_completion(answer, iteration + 1, forced=True)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iteration + 1,
|
||||
tools_called=total_tools_called,
|
||||
mental_models_created=mental_models_created,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
|
||||
# Check for done tool call (handle both 'done' and 'functions.done')
|
||||
done_call = next((tc for tc in result.tool_calls if tc.name == "done" or tc.name == "functions.done"), None)
|
||||
if done_call:
|
||||
# Guardrail: Require evidence before done
|
||||
has_gathered_evidence = bool(available_memory_ids) or bool(available_model_ids)
|
||||
if not has_gathered_evidence and iteration < max_iterations - 1:
|
||||
# Add assistant message and fake tool result asking for evidence
|
||||
messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [_tool_call_to_dict(done_call)],
|
||||
}
|
||||
)
|
||||
messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": done_call.id,
|
||||
"content": json.dumps(
|
||||
{
|
||||
"error": "You must call recall() or list_mental_models() to gather evidence before providing your final answer."
|
||||
}
|
||||
),
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
# Process done tool
|
||||
return await _process_done_tool(
|
||||
done_call,
|
||||
available_memory_ids,
|
||||
available_model_ids,
|
||||
iteration + 1,
|
||||
total_tools_called,
|
||||
mental_models_created,
|
||||
tool_trace,
|
||||
_get_llm_trace(),
|
||||
_log_completion,
|
||||
reflect_id,
|
||||
directives_applied=directives_applied,
|
||||
llm_config=llm_config,
|
||||
response_schema=response_schema,
|
||||
)
|
||||
|
||||
# Execute other tools in parallel (exclude done and functions.done)
|
||||
other_tools = [tc for tc in result.tool_calls if tc.name not in ("done", "functions.done")]
|
||||
if other_tools:
|
||||
# Add assistant message with tool calls
|
||||
messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [_tool_call_to_dict(tc) for tc in other_tools],
|
||||
}
|
||||
)
|
||||
|
||||
# Execute tools in parallel
|
||||
tool_tasks = [
|
||||
_execute_tool_with_timing(tc, lookup_fn, recall_fn, expand_fn, learn_fn) for tc in other_tools
|
||||
]
|
||||
tool_results = await asyncio.gather(*tool_tasks, return_exceptions=True)
|
||||
total_tools_called += len(other_tools)
|
||||
|
||||
# Process results and add to messages
|
||||
for tc, result_data in zip(other_tools, tool_results):
|
||||
if isinstance(result_data, Exception):
|
||||
# Tool execution failed - log and raise to fail the request
|
||||
logger.error(f"[REFLECT {reflect_id}] Tool {tc.name} failed with exception: {result_data}")
|
||||
raise RuntimeError(f"Reflect tool '{tc.name}' failed: {result_data}")
|
||||
|
||||
output, duration_ms = result_data
|
||||
|
||||
# Check if tool returned an error response
|
||||
if isinstance(output, dict) and "error" in output:
|
||||
logger.error(f"[REFLECT {reflect_id}] Tool {tc.name} returned error: {output['error']}")
|
||||
raise RuntimeError(f"Reflect tool '{tc.name}' error: {output['error']}")
|
||||
|
||||
# Track created mental models
|
||||
if tc.name == "learn" and isinstance(output, dict) and "model_id" in output:
|
||||
mental_models_created.append(output["model_id"])
|
||||
|
||||
# Track available memory IDs from recall
|
||||
if tc.name == "recall" and isinstance(output, dict) and "memories" in output:
|
||||
for memory in output["memories"]:
|
||||
if "id" in memory:
|
||||
available_memory_ids.add(memory["id"])
|
||||
|
||||
# Track available model IDs
|
||||
if tc.name in ("list_mental_models", "get_mental_model") and isinstance(output, dict):
|
||||
if output.get("found") and "model" in output:
|
||||
model_id = output["model"].get("id")
|
||||
if model_id:
|
||||
available_model_ids.add(model_id)
|
||||
elif "models" in output:
|
||||
for model in output["models"]:
|
||||
if "id" in model:
|
||||
available_model_ids.add(model["id"])
|
||||
|
||||
# Add tool result message
|
||||
messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tc.id,
|
||||
"content": json.dumps(output, default=str),
|
||||
}
|
||||
)
|
||||
|
||||
# Track for logging and context history
|
||||
input_dict = {"tool": tc.name, **tc.arguments}
|
||||
input_summary = _summarize_input(tc.name, tc.arguments)
|
||||
|
||||
tool_trace.append(
|
||||
ToolCall(
|
||||
tool=tc.name, input=input_dict, output=output, duration_ms=duration_ms, iteration=iteration + 1
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
output_chars = len(json.dumps(output))
|
||||
except (TypeError, ValueError):
|
||||
output_chars = len(str(output))
|
||||
|
||||
tool_trace_summary.append(
|
||||
{
|
||||
"tool": tc.name,
|
||||
"input_summary": input_summary,
|
||||
"duration_ms": duration_ms,
|
||||
"output_chars": output_chars,
|
||||
}
|
||||
)
|
||||
|
||||
# Keep context history for fallback final prompt
|
||||
context_history.append({"tool": tc.name, "input": input_dict, "output": output})
|
||||
|
||||
# Should not reach here
|
||||
answer = "I was unable to formulate a complete answer within the iteration limit."
|
||||
_log_completion(answer, max_iterations, forced=True)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
iterations=max_iterations,
|
||||
tools_called=total_tools_called,
|
||||
mental_models_created=mental_models_created,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=_get_llm_trace(),
|
||||
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_model_ids: set[str],
|
||||
iterations: int,
|
||||
total_tools_called: int,
|
||||
mental_models_created: list[str],
|
||||
tool_trace: list[ToolCall],
|
||||
llm_trace: list[LLMCall],
|
||||
log_completion: Callable,
|
||||
reflect_id: str,
|
||||
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
|
||||
|
||||
answer = args.get("answer", "").strip()
|
||||
if not answer:
|
||||
answer = "No answer provided."
|
||||
|
||||
# Validate IDs
|
||||
used_memory_ids = [mid for mid in args.get("memory_ids", []) if mid in available_memory_ids]
|
||||
used_model_ids = [mid for mid in args.get("model_ids", []) if mid in available_model_ids]
|
||||
|
||||
# Generate structured output if schema provided
|
||||
structured_output = None
|
||||
if response_schema and llm_config and answer:
|
||||
structured_output = await _generate_structured_output(answer, response_schema, llm_config, reflect_id)
|
||||
|
||||
log_completion(answer, iterations)
|
||||
return ReflectAgentResult(
|
||||
text=answer,
|
||||
structured_output=structured_output,
|
||||
iterations=iterations,
|
||||
tools_called=total_tools_called,
|
||||
mental_models_created=mental_models_created,
|
||||
tool_trace=tool_trace,
|
||||
llm_trace=llm_trace,
|
||||
used_memory_ids=used_memory_ids,
|
||||
used_model_ids=used_model_ids,
|
||||
directives_applied=directives_applied,
|
||||
)
|
||||
|
||||
|
||||
async def _execute_tool_with_timing(
|
||||
tc: "LLMToolCall",
|
||||
lookup_fn: Callable[[str | None], Awaitable[dict[str, Any]]],
|
||||
recall_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
|
||||
learn_fn: Callable[[MentalModelInput], Awaitable[dict[str, Any]]] | None = None,
|
||||
) -> tuple[dict[str, Any], int]:
|
||||
"""Execute a tool call and return result with timing."""
|
||||
start = time.time()
|
||||
result = await _execute_tool(tc.name, tc.arguments, lookup_fn, recall_fn, expand_fn, learn_fn)
|
||||
duration_ms = int((time.time() - start) * 1000)
|
||||
return result, duration_ms
|
||||
|
||||
|
||||
async def _execute_tool(
|
||||
tool_name: str,
|
||||
args: dict[str, Any],
|
||||
lookup_fn: Callable[[str | None], Awaitable[dict[str, Any]]],
|
||||
recall_fn: Callable[[str, int], Awaitable[dict[str, Any]]],
|
||||
expand_fn: Callable[[list[str], str], Awaitable[dict[str, Any]]],
|
||||
learn_fn: Callable[[MentalModelInput], Awaitable[dict[str, Any]]] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Execute a single tool by name."""
|
||||
# Normalize tool name - some LLMs return 'functions.done' instead of 'done'
|
||||
if tool_name.startswith("functions."):
|
||||
tool_name = tool_name[len("functions.") :]
|
||||
|
||||
if tool_name == "list_mental_models":
|
||||
return await lookup_fn(None)
|
||||
|
||||
elif tool_name == "get_mental_model":
|
||||
model_id = args.get("model_id")
|
||||
if not model_id:
|
||||
return {"error": "get_mental_model requires model_id"}
|
||||
return await lookup_fn(model_id)
|
||||
|
||||
elif tool_name == "recall":
|
||||
query = args.get("query")
|
||||
if not query:
|
||||
return {"error": "recall requires a query parameter"}
|
||||
max_tokens = max(args.get("max_tokens") or 2048, 1000) # Default 2048, min 1000
|
||||
return await recall_fn(query, max_tokens)
|
||||
|
||||
elif tool_name == "learn":
|
||||
if learn_fn is None:
|
||||
return {"error": "learn tool is not available"}
|
||||
name = args.get("name")
|
||||
description = args.get("description")
|
||||
if not name or not description:
|
||||
return {"error": "learn requires name and description"}
|
||||
return await learn_fn(MentalModelInput(name=name, description=description))
|
||||
|
||||
elif tool_name == "expand":
|
||||
memory_ids = args.get("memory_ids", [])
|
||||
if not memory_ids:
|
||||
return {"error": "expand requires memory_ids"}
|
||||
depth = args.get("depth", "chunk")
|
||||
return await expand_fn(memory_ids, depth)
|
||||
|
||||
else:
|
||||
return {"error": f"Unknown tool: {tool_name}"}
|
||||
|
||||
|
||||
def _summarize_input(tool_name: str, args: dict[str, Any]) -> str:
|
||||
"""Create a summary of tool input for logging, showing all params."""
|
||||
if tool_name == "list_mental_models":
|
||||
return "()"
|
||||
elif tool_name == "get_mental_model":
|
||||
return f"(model_id={args.get('model_id', '?')})"
|
||||
elif tool_name == "recall":
|
||||
query = args.get("query", "")
|
||||
query_preview = f"'{query[:30]}...'" if len(query) > 30 else f"'{query}'"
|
||||
# Show actual value used (default 2048, min 1000)
|
||||
max_tokens = max(args.get("max_tokens") or 2048, 1000)
|
||||
return f"(query={query_preview}, max_tokens={max_tokens})"
|
||||
elif tool_name == "learn":
|
||||
name = args.get("name", "?")
|
||||
desc = args.get("description", "")
|
||||
desc_preview = f"'{desc[:20]}...'" if len(desc) > 20 else f"'{desc}'"
|
||||
return f"(name='{name}', description={desc_preview})"
|
||||
elif tool_name == "expand":
|
||||
memory_ids = args.get("memory_ids", [])
|
||||
depth = args.get("depth", "chunk")
|
||||
return f"(memory_ids=[{len(memory_ids)} ids], depth={depth})"
|
||||
elif tool_name == "done":
|
||||
answer = args.get("answer", "")
|
||||
answer_preview = f"'{answer[:30]}...'" if len(answer) > 30 else f"'{answer}'"
|
||||
memory_ids = args.get("memory_ids", [])
|
||||
model_ids = args.get("model_ids", [])
|
||||
return f"(answer={answer_preview}, memory_ids={len(memory_ids)}, model_ids={len(model_ids)})"
|
||||
return str(args)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,114 @@
|
||||
"""
|
||||
Pydantic models for the reflect agent.
|
||||
"""
|
||||
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class MentalModelObservation(BaseModel):
|
||||
"""An observation within a mental model with its supporting memories."""
|
||||
|
||||
title: str = Field(description="Observation header (can be empty for intro)")
|
||||
text: str = Field(description="Observation content - no headers, use lists/tables/bold")
|
||||
memory_ids: list[str] = Field(default_factory=list, description="Memory IDs supporting this observation")
|
||||
|
||||
|
||||
class MentalModelInput(BaseModel):
|
||||
"""Input for the learn tool to create a mental model placeholder.
|
||||
|
||||
The agent only specifies name and description - the actual content/observations
|
||||
are generated during refresh, similar to pinned models.
|
||||
"""
|
||||
|
||||
name: str = Field(description="Human-readable name for the mental model")
|
||||
description: str = Field(description="What to track - used as prompt for content generation during refresh")
|
||||
entity_id: str | None = Field(default=None, description="Optional link to existing entity ID")
|
||||
|
||||
|
||||
class AnswerSection(BaseModel):
|
||||
"""A section of the answer with its supporting evidence (DEPRECATED)."""
|
||||
|
||||
title: str = Field(description="Section header/title")
|
||||
text: str = Field(description="Section content")
|
||||
memory_ids: list[str] = Field(default_factory=list, description="Memory IDs supporting this section")
|
||||
model_ids: list[str] = Field(default_factory=list, description="Mental model IDs supporting this section")
|
||||
|
||||
|
||||
class ReflectAction(BaseModel):
|
||||
"""Single action the reflect agent can take."""
|
||||
|
||||
tool: Literal["list_mental_models", "get_mental_model", "recall", "learn", "expand", "done"] = Field(
|
||||
description="Tool to invoke: list_mental_models, get_mental_model, recall, learn, expand, or done"
|
||||
)
|
||||
# Tool-specific parameters
|
||||
model_id: str | None = Field(default=None, description="Mental model ID for get_mental_model")
|
||||
query: str | None = Field(default=None, description="Search query for recall")
|
||||
max_tokens: int | None = Field(default=None, description="Max tokens for recall results (default 2048)")
|
||||
mental_model: MentalModelInput | None = Field(default=None, description="Mental model to create/update for learn")
|
||||
memory_ids: list[str] | None = Field(default=None, description="Memory unit IDs for expand (batched)")
|
||||
depth: Literal["chunk", "document"] | None = Field(default=None, description="Expansion depth for expand")
|
||||
sections: list[AnswerSection] | None = Field(default=None, description="DEPRECATED: Use answer field instead")
|
||||
observations: list[MentalModelObservation] | None = Field(
|
||||
default=None, description="Observations for done action (when output_mode=observations)"
|
||||
)
|
||||
# Plain text answer fields (for output_mode=answer)
|
||||
answer: str | None = Field(default=None, description="Plain text answer for done action (no markdown)")
|
||||
answer_memory_ids: list[str] | None = Field(
|
||||
default=None, description="Memory IDs supporting the answer", alias="memory_ids"
|
||||
)
|
||||
answer_model_ids: list[str] | None = Field(
|
||||
default=None, description="Mental model IDs supporting the answer", alias="model_ids"
|
||||
)
|
||||
reasoning: str | None = Field(default=None, description="Brief reasoning for this action")
|
||||
|
||||
|
||||
class ReflectActionBatch(BaseModel):
|
||||
"""Batch of actions for parallel execution."""
|
||||
|
||||
actions: list[ReflectAction] = Field(description="List of actions to execute in parallel")
|
||||
|
||||
|
||||
class ToolCall(BaseModel):
|
||||
"""A single tool call made during reflect."""
|
||||
|
||||
tool: str = Field(description="Tool name: lookup, recall, learn, expand")
|
||||
input: dict = Field(description="Tool input parameters")
|
||||
output: dict = Field(description="Tool output/result")
|
||||
duration_ms: int = Field(description="Execution time in milliseconds")
|
||||
iteration: int = Field(default=0, description="Iteration number (1-based) when this tool was called")
|
||||
|
||||
|
||||
class LLMCall(BaseModel):
|
||||
"""A single LLM call made during reflect."""
|
||||
|
||||
scope: str = Field(description="Call scope: agent_1, agent_2, final, etc.")
|
||||
duration_ms: int = Field(description="Execution time in milliseconds")
|
||||
|
||||
|
||||
class 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")
|
||||
rules: list[str] = Field(default_factory=list, description="Directive rules/observations that were applied")
|
||||
|
||||
|
||||
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")
|
||||
mental_models_created: list[str] = Field(default_factory=list, description="IDs of mental models created/updated")
|
||||
tool_trace: list[ToolCall] = Field(default_factory=list, description="Trace of all tool calls made")
|
||||
llm_trace: list[LLMCall] = Field(default_factory=list, description="Trace of all LLM calls made")
|
||||
used_memory_ids: list[str] = Field(default_factory=list, description="Validated memory IDs actually used in answer")
|
||||
used_model_ids: list[str] = Field(default_factory=list, description="Validated model IDs actually used in answer")
|
||||
directives_applied: list[DirectiveInfo] = Field(
|
||||
default_factory=list, description="Directive mental models that affected this reflection"
|
||||
)
|
||||
@@ -0,0 +1,248 @@
|
||||
"""
|
||||
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
|
||||
|
||||
|
||||
class CandidateObservation(BaseModel):
|
||||
"""A candidate observation generated during the seed phase.
|
||||
|
||||
Candidates are preliminary observations that need evidence validation
|
||||
before becoming full observations.
|
||||
"""
|
||||
|
||||
content: str = Field(description="The proposed observation content")
|
||||
seed_memory_ids: list[str] = Field(default_factory=list, description="Memory IDs that inspired this candidate")
|
||||
|
||||
|
||||
class CandidateWithEvidence(BaseModel):
|
||||
"""A candidate observation with gathered supporting and contradicting evidence."""
|
||||
|
||||
candidate: CandidateObservation
|
||||
supporting_memories: list[dict] = Field(default_factory=list, description="Memories that support this observation")
|
||||
contradicting_memories: list[dict] = Field(
|
||||
default_factory=list, description="Memories that contradict this observation"
|
||||
)
|
||||
|
||||
|
||||
class MentalModelSnapshot(BaseModel):
|
||||
"""A versioned snapshot of a mental model's observations.
|
||||
|
||||
Used for tracking changes over time and enabling diff views.
|
||||
"""
|
||||
|
||||
version: int = Field(description="Version number (1-indexed)")
|
||||
observations: list[Observation] = Field(default_factory=list, description="Observations at this version")
|
||||
created_at: datetime = Field(
|
||||
default_factory=lambda: datetime.now(timezone.utc), description="When this version was created"
|
||||
)
|
||||
reflect_summary: str | None = Field(default=None, description="Summary of changes in this version")
|
||||
|
||||
|
||||
def verify_evidence_quotes(
|
||||
observation: Observation,
|
||||
memories: dict[str, str],
|
||||
) -> tuple[bool, list[str]]:
|
||||
"""Verify that all evidence quotes exist in the referenced memories.
|
||||
|
||||
Args:
|
||||
observation: The observation to verify
|
||||
memories: Dict mapping memory_id to memory content
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, list of error messages)
|
||||
"""
|
||||
errors = []
|
||||
|
||||
for evidence in observation.evidence:
|
||||
memory_content = memories.get(evidence.memory_id)
|
||||
if memory_content is None:
|
||||
errors.append(f"Memory {evidence.memory_id} not found")
|
||||
continue
|
||||
|
||||
if evidence.quote not in memory_content:
|
||||
errors.append(f"Quote not found in memory {evidence.memory_id}: '{evidence.quote[:50]}...'")
|
||||
|
||||
return len(errors) == 0, errors
|
||||
@@ -0,0 +1,762 @@
|
||||
"""
|
||||
System prompts for the reflect agent.
|
||||
"""
|
||||
|
||||
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 directive mental models with observations
|
||||
|
||||
Returns:
|
||||
List of directive rule strings
|
||||
"""
|
||||
rules = []
|
||||
for directive in directives:
|
||||
directive_name = directive.get("name", "")
|
||||
observations = directive.get("observations", [])
|
||||
if observations:
|
||||
for obs in observations:
|
||||
# Support both Pydantic Observation objects and dicts
|
||||
if hasattr(obs, "title"):
|
||||
title = obs.title
|
||||
content = obs.content
|
||||
else:
|
||||
title = obs.get("title", "")
|
||||
content = obs.get("content", "")
|
||||
if title and content:
|
||||
rules.append(f"**{title}**: {content}")
|
||||
elif content:
|
||||
rules.append(content)
|
||||
elif directive_name:
|
||||
# Fallback to description if no observations
|
||||
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,
|
||||
) -> str:
|
||||
"""
|
||||
Build the system prompt for tool-calling reflect agent.
|
||||
|
||||
This is a simplified prompt since tools are defined separately via the tools parameter.
|
||||
|
||||
Args:
|
||||
bank_profile: Bank profile with name and mission
|
||||
context: Optional additional context
|
||||
directives: Optional list of directive mental models to inject as hard rules
|
||||
"""
|
||||
name = bank_profile.get("name", "Assistant")
|
||||
mission = bank_profile.get("mission", "")
|
||||
|
||||
no_info_rule = (
|
||||
"- Only say 'I don't have information' AFTER trying list_mental_models AND recall with no relevant results"
|
||||
)
|
||||
|
||||
parts = []
|
||||
|
||||
# Inject directives at the VERY START for maximum prominence
|
||||
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",
|
||||
"- You must NEVER fabricate information that has no basis in retrieved data",
|
||||
"- You SHOULD synthesize, infer, and reason from the retrieved memories",
|
||||
"- You MUST call recall() before saying you don't have information",
|
||||
no_info_rule,
|
||||
"",
|
||||
"## How to Reason",
|
||||
"- If memories mention someone did an activity, you can infer they likely enjoyed it",
|
||||
"- Synthesize a coherent narrative from related memories",
|
||||
"- Be a thoughtful interpreter, not just a literal repeater",
|
||||
"- When the exact answer isn't stated, use what IS stated to give the best answer",
|
||||
"",
|
||||
"## Query Strategy (IMPORTANT)",
|
||||
"recall() uses semantic search. NEVER just echo the user's question - decompose it into targeted searches:",
|
||||
"",
|
||||
"BAD: User asks 'recurring lesson themes between students' → recall('recurring lesson themes between students')",
|
||||
"GOOD: Break it down into component searches:",
|
||||
" 1. recall('lessons') - find all lesson-related memories",
|
||||
" 2. recall('teaching sessions') - alternative phrasing",
|
||||
" 3. recall('student progress') - find student-related memories",
|
||||
" 4. recall('topics taught') - find subject matter",
|
||||
"",
|
||||
"Think: What ENTITIES and CONCEPTS does this question involve? Search for each separately.",
|
||||
"- Questions about patterns → search for the individual instances first",
|
||||
"- Questions comparing things → search for each thing separately",
|
||||
"- Questions about relationships → search for each party involved",
|
||||
"",
|
||||
"## Workflow",
|
||||
]
|
||||
)
|
||||
|
||||
# Answer mode: include mental model lookup in workflow
|
||||
parts.extend(
|
||||
[
|
||||
"1. Review the pre-fetched mental models for relevant synthesized knowledge",
|
||||
"2. If relevant, call get_mental_model(model_id) for full observations",
|
||||
"3. DECOMPOSE the question into component searches (see Query Strategy above)",
|
||||
" - Identify entities and concepts in the question",
|
||||
" - Search for each separately with targeted queries",
|
||||
"4. Run multiple recall() calls - don't just echo the user's question",
|
||||
"5. Use expand() if you need more context on specific memories",
|
||||
"6. BEFORE answering: Check if any person/project/concept from the memories deserves a mental model - use learn() if so",
|
||||
"7. When ready, call done() with your answer and supporting memory_ids",
|
||||
"",
|
||||
"## When to Use learn() - IMPORTANT",
|
||||
"ACTIVELY look for opportunities to use learn() when you discover:",
|
||||
"- A person mentioned in 2+ memories who has no mental model yet",
|
||||
"- A project or concept the user asks about that has no mental model",
|
||||
"- A pattern or topic worth tracking for future questions",
|
||||
"",
|
||||
"DO NOT wait to be asked - proactively create models when you see the need.",
|
||||
"Example: learn(name='Project Alpha', description='Track goals, status, and key decisions for Project Alpha')",
|
||||
"",
|
||||
"## Output Format: Plain Text Answer",
|
||||
"Call done() with a plain text 'answer' field.",
|
||||
"- Do NOT use markdown formatting",
|
||||
"- NEVER include memory IDs, UUIDs, or 'Memory references' in the answer text",
|
||||
"- Put memory IDs ONLY in the memory_ids array parameter, not in the answer",
|
||||
]
|
||||
)
|
||||
|
||||
parts.append("")
|
||||
parts.append(f"## Memory Bank: {name}")
|
||||
|
||||
if mission:
|
||||
parts.append(f"Mission: {mission}")
|
||||
|
||||
# Disposition traits
|
||||
disposition = bank_profile.get("disposition", {})
|
||||
if disposition:
|
||||
traits = []
|
||||
if "skepticism" in disposition:
|
||||
traits.append(f"skepticism={disposition['skepticism']}")
|
||||
if "literalism" in disposition:
|
||||
traits.append(f"literalism={disposition['literalism']}")
|
||||
if "empathy" in disposition:
|
||||
traits.append(f"empathy={disposition['empathy']}")
|
||||
if traits:
|
||||
parts.append(f"Disposition: {', '.join(traits)}")
|
||||
|
||||
if context:
|
||||
parts.append(f"\n## Additional Context\n{context}")
|
||||
|
||||
# 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 calling list_mental_models() to see available mental models - they contain pre-synthesized knowledge. "
|
||||
"If a relevant model exists, use get_mental_model(model_id) to get its observations. "
|
||||
"Then use recall(query) for specific details not covered by mental models."
|
||||
)
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def build_final_prompt(
|
||||
query: str,
|
||||
context_history: list[dict],
|
||||
bank_profile: dict,
|
||||
additional_context: str | None = None,
|
||||
) -> str:
|
||||
"""Build the final prompt when forcing a text response (no tools)."""
|
||||
parts = []
|
||||
|
||||
# Bank identity
|
||||
name = bank_profile.get("name", "Assistant")
|
||||
mission = bank_profile.get("mission", "")
|
||||
|
||||
parts.append(f"## Memory Bank Context\nName: {name}")
|
||||
if mission:
|
||||
parts.append(f"Mission: {mission}")
|
||||
|
||||
# Disposition traits if present
|
||||
disposition = bank_profile.get("disposition", {})
|
||||
if disposition:
|
||||
traits = []
|
||||
if "skepticism" in disposition:
|
||||
traits.append(f"skepticism={disposition['skepticism']}")
|
||||
if "literalism" in disposition:
|
||||
traits.append(f"literalism={disposition['literalism']}")
|
||||
if "empathy" in disposition:
|
||||
traits.append(f"empathy={disposition['empathy']}")
|
||||
if traits:
|
||||
parts.append(f"Disposition: {', '.join(traits)}")
|
||||
|
||||
# Additional context from caller
|
||||
if additional_context:
|
||||
parts.append(f"\n## Additional Context\n{additional_context}")
|
||||
|
||||
# Tool call history
|
||||
if context_history:
|
||||
parts.append("\n## Retrieved Data (synthesize and reason from this data)")
|
||||
for entry in context_history:
|
||||
tool = entry["tool"]
|
||||
output = entry["output"]
|
||||
# Format as proper JSON for LLM readability
|
||||
try:
|
||||
output_str = json.dumps(output, indent=2, default=str)
|
||||
except (TypeError, ValueError):
|
||||
output_str = str(output)
|
||||
parts.append(f"\n### From {tool}:\n```json\n{output_str}\n```")
|
||||
else:
|
||||
parts.append("\n## Retrieved Data\nNo data was retrieved.")
|
||||
|
||||
# The question
|
||||
parts.append(f"\n## Question\n{query}")
|
||||
|
||||
# Final instructions
|
||||
parts.append(
|
||||
"\n## Instructions\n"
|
||||
"Provide a thoughtful answer by synthesizing and reasoning from the retrieved data above. "
|
||||
"You can make reasonable inferences from the memories, but don't completely fabricate information."
|
||||
"If the exact answer isn't stated, use what IS stated to give the best possible answer. "
|
||||
"Only say 'I don't have information' if the retrieved data is truly unrelated to the question."
|
||||
)
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
FINAL_SYSTEM_PROMPT = """You are a thoughtful assistant that synthesizes answers from retrieved memories.
|
||||
|
||||
Your approach:
|
||||
- Reason over the retrieved memories to answer the question
|
||||
- Make reasonable inferences when the exact answer isn't explicitly stated
|
||||
- Connect related memories to form a complete picture
|
||||
- Be helpful - if you have related information, use it to give the best possible answer
|
||||
|
||||
Only say "I don't have information" if the retrieved data is truly unrelated to the question.
|
||||
Do NOT fabricate information that has no basis in the retrieved data."""
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# 4-Phase Mental Model Reflect Prompts
|
||||
# =============================================================================
|
||||
|
||||
SEED_PHASE_SYSTEM_PROMPT = """You are analyzing memories to discover NEW patterns and generate candidate observations.
|
||||
|
||||
Your task is to identify potential observations (beliefs, preferences, patterns, behaviors) that could be part of a mental model about this person/topic.
|
||||
|
||||
## Important: Avoid Redundancy
|
||||
If existing observations are provided, DO NOT generate candidates that are essentially the same.
|
||||
Focus on discovering NEW patterns not already covered by existing observations.
|
||||
|
||||
## Rules
|
||||
- Generate 5-15 candidate observations for NEW patterns only
|
||||
- Each candidate should be specific and testable (can be supported or contradicted by evidence)
|
||||
- Note which memory IDs inspired each candidate (these are seeds, not final evidence)
|
||||
- Focus on patterns that appear MULTIPLE TIMES across many memories - the more the better
|
||||
- The best candidates are ones you can find 10, 20, or even 50+ supporting memories for
|
||||
- Skip patterns that are already covered by existing observations
|
||||
|
||||
## Output Format
|
||||
Return a JSON array of candidate observations:
|
||||
```json
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": "The specific observation/belief/pattern - be detailed and specific",
|
||||
"seed_memory_ids": ["memory_id_1", "memory_id_2", "memory_id_3"]
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
Focus on patterns that appear multiple times or have strong signals. Don't generate obvious or trivial observations.
|
||||
Prefer candidates with MORE seed memories - they're more likely to be real patterns.
|
||||
Return an empty candidates array if no genuinely new patterns are found."""
|
||||
|
||||
|
||||
def build_seed_phase_prompt(
|
||||
memories: list[dict],
|
||||
topic: str | None = None,
|
||||
existing_observations: list[dict] | None = None,
|
||||
) -> str:
|
||||
"""Build the user prompt for the seed phase.
|
||||
|
||||
Args:
|
||||
memories: List of memories to analyze
|
||||
topic: Optional topic focus for the mental model
|
||||
existing_observations: Optional list of existing observations to avoid rediscovering
|
||||
"""
|
||||
parts = []
|
||||
|
||||
if topic:
|
||||
parts.append(f"## Topic Focus\n{topic}\n")
|
||||
|
||||
# Include existing observations so we don't rediscover them
|
||||
if existing_observations:
|
||||
parts.append("## Existing Observations (DO NOT regenerate these)")
|
||||
parts.append("These patterns are already tracked. Focus on discovering NEW patterns:\n")
|
||||
for i, obs in enumerate(existing_observations, 1):
|
||||
title = obs.get("title", "")
|
||||
content = obs.get("content", "")
|
||||
parts.append(f"{i}. **{title}**: {content}\n")
|
||||
parts.append("")
|
||||
|
||||
parts.append("## Memories to Analyze")
|
||||
parts.append("Review these memories and identify patterns, preferences, beliefs, and behaviors:\n")
|
||||
|
||||
for mem in memories:
|
||||
mem_id = mem.get("id", "unknown")
|
||||
content = mem.get("content", mem.get("text", ""))
|
||||
timestamp = mem.get("timestamp", mem.get("created_at", ""))
|
||||
parts.append(f"[{mem_id}] ({timestamp}): {content}\n")
|
||||
|
||||
parts.append("\n## Instructions")
|
||||
if existing_observations:
|
||||
parts.append("Generate candidate observations for NEW patterns not already covered above.")
|
||||
parts.append("If all patterns are already covered by existing observations, return an empty candidates array.")
|
||||
else:
|
||||
parts.append("Generate candidate observations based on patterns you see in these memories.")
|
||||
parts.append("Look for: recurring themes, stated preferences, behavioral patterns, beliefs, values, goals.")
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
VALIDATE_PHASE_SYSTEM_PROMPT = """You are validating candidate observations against evidence.
|
||||
|
||||
For each candidate, you have:
|
||||
- Supporting memories (evidence FOR the observation)
|
||||
- Contradicting memories (evidence AGAINST the observation)
|
||||
|
||||
## Your Task
|
||||
1. Evaluate each candidate based on the evidence
|
||||
2. For valid candidates, extract EXACT QUOTES from supporting memories
|
||||
3. Discard candidates with insufficient or contradicting evidence
|
||||
4. Merge similar candidates into single, refined observations
|
||||
|
||||
## Rules for Quotes
|
||||
- Quotes must be EXACT text from the memory, not paraphrased
|
||||
- Each quote should directly support the observation
|
||||
- The MORE evidence quotes, the BETTER - don't limit yourself, include ALL relevant quotes (10, 20, 50+)
|
||||
- Observations with only 1-2 quotes are weak and should be discarded unless the evidence is exceptionally strong
|
||||
- Stronger observations have more supporting evidence - aim for comprehensive coverage
|
||||
|
||||
## Output Format
|
||||
Return validated observations with evidence:
|
||||
```json
|
||||
{
|
||||
"observations": [
|
||||
{
|
||||
"title": "Short descriptive title (3-8 words) - like a headline",
|
||||
"content": "The full observation content - detailed explanation of the pattern/belief",
|
||||
"evidence": [
|
||||
{
|
||||
"memory_id": "exact_memory_id",
|
||||
"quote": "Exact quote from the memory text",
|
||||
"relevance": "Brief explanation of how this supports the observation",
|
||||
"timestamp": "2024-01-15T10:00:00Z"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"discarded": [
|
||||
{
|
||||
"content": "The discarded candidate",
|
||||
"reason": "Why it was discarded (insufficient evidence, contradicted, etc.)"
|
||||
}
|
||||
],
|
||||
"merged": [
|
||||
{
|
||||
"from": ["candidate 1 content", "candidate 2 content"],
|
||||
"into": "The merged observation content"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
## Title Guidelines
|
||||
- Title should be a SHORT label (like "Prefers morning meetings" or "Coffee enthusiast")
|
||||
- NOT a truncated version of the content
|
||||
- Think of it as a category/tag for the observation
|
||||
|
||||
Be rigorous: only keep observations with clear, verifiable evidence from multiple memories."""
|
||||
|
||||
|
||||
def build_validate_phase_prompt(candidates_with_evidence: list[dict]) -> str:
|
||||
"""Build the user prompt for the validate phase."""
|
||||
parts = ["## Candidates to Validate\n"]
|
||||
|
||||
for i, item in enumerate(candidates_with_evidence, 1):
|
||||
candidate = item.get("candidate", {})
|
||||
supporting = item.get("supporting_memories", [])
|
||||
contradicting = item.get("contradicting_memories", [])
|
||||
|
||||
parts.append(f"### Candidate {i}: {candidate.get('content', '')}")
|
||||
|
||||
if supporting:
|
||||
parts.append("\n**Supporting Evidence:**")
|
||||
for mem in supporting:
|
||||
mem_id = mem.get("id", "unknown")
|
||||
content = mem.get("content", mem.get("text", ""))
|
||||
timestamp = mem.get("timestamp", mem.get("created_at", ""))
|
||||
parts.append(f"- [{mem_id}] ({timestamp}): {content}")
|
||||
|
||||
if contradicting:
|
||||
parts.append("\n**Contradicting Evidence:**")
|
||||
for mem in contradicting:
|
||||
mem_id = mem.get("id", "unknown")
|
||||
content = mem.get("content", mem.get("text", ""))
|
||||
timestamp = mem.get("timestamp", mem.get("created_at", ""))
|
||||
parts.append(f"- [{mem_id}] ({timestamp}): {content}")
|
||||
|
||||
if not supporting and not contradicting:
|
||||
parts.append("\n*No additional evidence found*")
|
||||
|
||||
parts.append("")
|
||||
|
||||
parts.append("## Instructions")
|
||||
parts.append("1. Evaluate each candidate based on its evidence")
|
||||
parts.append("2. Keep candidates with strong supporting evidence")
|
||||
parts.append("3. Discard candidates with no evidence or strong contradictions")
|
||||
parts.append("4. Merge similar candidates")
|
||||
parts.append("5. Extract EXACT quotes (copy-paste from memory text) for evidence")
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
COMPARE_PHASE_SYSTEM_PROMPT = """You are merging new observations with an existing mental model.
|
||||
|
||||
You have:
|
||||
- EXISTING observations (from the current mental model)
|
||||
- NEW observations (from this reflect cycle)
|
||||
|
||||
## Your Task
|
||||
Produce the final, complete mental model by:
|
||||
1. Keeping existing observations that are still valid
|
||||
2. Updating existing observations with new evidence (ADD new evidence to existing)
|
||||
3. Adding new observations that don't overlap with existing
|
||||
4. Removing existing observations that are contradicted by new evidence
|
||||
5. Merging overlapping observations
|
||||
|
||||
## Rules
|
||||
- The final model should have no contradictions
|
||||
- Each observation must have evidence with exact quotes
|
||||
- COMBINE evidence from both existing and new observations
|
||||
- If an existing observation has new supporting evidence, ADD ALL the new evidence to it
|
||||
- Include ALL relevant evidence - the more quotes the better (10, 20, 50+ is great)
|
||||
- Observations with more evidence are more reliable - don't limit the number of quotes
|
||||
|
||||
## Output Format
|
||||
Return the complete, final mental model:
|
||||
```json
|
||||
{
|
||||
"observations": [
|
||||
{
|
||||
"title": "Short descriptive title (3-8 words)",
|
||||
"content": "Full observation content - detailed explanation",
|
||||
"evidence": [
|
||||
{
|
||||
"memory_id": "id",
|
||||
"quote": "exact quote",
|
||||
"relevance": "explanation",
|
||||
"timestamp": "ISO timestamp"
|
||||
}
|
||||
],
|
||||
"created_at": "ISO timestamp of when observation was first created"
|
||||
}
|
||||
],
|
||||
"changes": {
|
||||
"kept": ["Observation that was kept unchanged"],
|
||||
"updated": [{"from": "old content", "to": "new content", "reason": "why"}],
|
||||
"added": ["New observation that was added"],
|
||||
"removed": [{"content": "removed observation", "reason": "why removed"}],
|
||||
"merged": [{"from": ["obs1", "obs2"], "into": "merged observation"}]
|
||||
}
|
||||
}
|
||||
```"""
|
||||
|
||||
|
||||
def build_compare_phase_prompt(
|
||||
existing_observations: list[dict],
|
||||
new_observations: list[dict],
|
||||
) -> str:
|
||||
"""Build the user prompt for the compare phase."""
|
||||
parts = []
|
||||
|
||||
parts.append("## Existing Mental Model Observations")
|
||||
if existing_observations:
|
||||
for i, obs in enumerate(existing_observations, 1):
|
||||
title = obs.get("title", "")
|
||||
content = obs.get("content", obs.get("text", ""))
|
||||
evidence = obs.get("evidence", [])
|
||||
parts.append(f"\n### Existing {i}: {title}")
|
||||
parts.append(f"Content: {content}")
|
||||
if evidence:
|
||||
parts.append(f"Evidence ({len(evidence)} items):")
|
||||
for ev in evidence[:5]: # Show max 5 evidence items
|
||||
parts.append(f' - [{ev.get("memory_id", "?")}]: "{ev.get("quote", "")}"')
|
||||
if len(evidence) > 5:
|
||||
parts.append(f" ... and {len(evidence) - 5} more")
|
||||
else:
|
||||
parts.append("*No existing observations*")
|
||||
|
||||
parts.append("\n## New Observations from This Reflect")
|
||||
if new_observations:
|
||||
for i, obs in enumerate(new_observations, 1):
|
||||
title = obs.get("title", "")
|
||||
content = obs.get("content", "")
|
||||
evidence = obs.get("evidence", [])
|
||||
parts.append(f"\n### New {i}: {title}")
|
||||
parts.append(f"Content: {content}")
|
||||
if evidence:
|
||||
parts.append(f"Evidence ({len(evidence)} items):")
|
||||
for ev in evidence:
|
||||
parts.append(f' - [{ev.get("memory_id", "?")}]: "{ev.get("quote", "")}"')
|
||||
else:
|
||||
parts.append("*No new observations*")
|
||||
|
||||
parts.append("\n## Instructions")
|
||||
parts.append("Merge these into a coherent, non-contradictory mental model.")
|
||||
parts.append("Preserve all valid evidence. Remove stale or contradicted observations.")
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# UPDATE EXISTING Phase Prompts (for diff-based refresh)
|
||||
# =============================================================================
|
||||
|
||||
UPDATE_EXISTING_SYSTEM_PROMPT = """You are updating existing observations with newly found evidence.
|
||||
|
||||
For each existing observation, you have been given:
|
||||
- The original observation (title, content, existing evidence)
|
||||
- Newly found supporting memories
|
||||
- Newly found contradicting memories
|
||||
|
||||
## Your Task
|
||||
1. Extract EXACT QUOTES from new supporting memories to add to the observation
|
||||
2. Flag observations with strong contradicting evidence for potential removal
|
||||
3. Keep existing evidence intact - only ADD new evidence
|
||||
|
||||
## Rules for Quotes
|
||||
- Quotes must be EXACT text from the memory, not paraphrased
|
||||
- Each quote should directly support the observation
|
||||
- Include ALL relevant quotes from the new memories
|
||||
|
||||
## Output Format
|
||||
Return updated observations with new evidence:
|
||||
```json
|
||||
{
|
||||
"updated_observations": [
|
||||
{
|
||||
"title": "Original title",
|
||||
"content": "Original content",
|
||||
"existing_evidence_count": 5,
|
||||
"new_evidence": [
|
||||
{
|
||||
"memory_id": "exact_memory_id",
|
||||
"quote": "Exact quote from the memory text",
|
||||
"relevance": "Brief explanation of how this supports the observation",
|
||||
"timestamp": "2024-01-15T10:00:00Z"
|
||||
}
|
||||
],
|
||||
"has_contradiction": false,
|
||||
"contradiction_note": null
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
If an observation has strong contradicting evidence, set has_contradiction=true and explain in contradiction_note."""
|
||||
|
||||
|
||||
def build_update_existing_prompt(observations_with_evidence: list[dict]) -> str:
|
||||
"""Build the user prompt for the update existing phase.
|
||||
|
||||
Args:
|
||||
observations_with_evidence: List of existing observations with new evidence found
|
||||
"""
|
||||
parts = ["## Existing Observations to Update\n"]
|
||||
|
||||
for i, item in enumerate(observations_with_evidence, 1):
|
||||
obs = item.get("observation", {})
|
||||
supporting = item.get("supporting_memories", [])
|
||||
contradicting = item.get("contradicting_memories", [])
|
||||
|
||||
title = obs.get("title", "")
|
||||
content = obs.get("content", "")
|
||||
existing_evidence = obs.get("evidence", [])
|
||||
|
||||
parts.append(f"### Observation {i}: {title}")
|
||||
parts.append(f"Content: {content}")
|
||||
parts.append(f"Existing evidence count: {len(existing_evidence)}")
|
||||
|
||||
if supporting:
|
||||
parts.append("\n**New Supporting Memories:**")
|
||||
for mem in supporting:
|
||||
mem_id = mem.get("id", "unknown")
|
||||
mem_content = mem.get("content", mem.get("text", ""))
|
||||
timestamp = mem.get("timestamp", mem.get("created_at", ""))
|
||||
parts.append(f"- [{mem_id}] ({timestamp}): {mem_content}")
|
||||
|
||||
if contradicting:
|
||||
parts.append("\n**New Contradicting Memories:**")
|
||||
for mem in contradicting:
|
||||
mem_id = mem.get("id", "unknown")
|
||||
mem_content = mem.get("content", mem.get("text", ""))
|
||||
timestamp = mem.get("timestamp", mem.get("created_at", ""))
|
||||
parts.append(f"- [{mem_id}] ({timestamp}): {mem_content}")
|
||||
|
||||
if not supporting and not contradicting:
|
||||
parts.append("\n*No new evidence found*")
|
||||
|
||||
parts.append("")
|
||||
|
||||
parts.append("## Instructions")
|
||||
parts.append("1. Extract EXACT quotes from new supporting memories")
|
||||
parts.append("2. Flag observations with strong contradictions")
|
||||
parts.append("3. Return the updated observations with new evidence added")
|
||||
|
||||
return "\n".join(parts)
|
||||
@@ -0,0 +1,450 @@
|
||||
"""
|
||||
Tool implementations for the reflect agent.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from .models import MentalModelInput
|
||||
from .observations import Observation, ObservationEvidence, Trend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from asyncpg import Connection
|
||||
|
||||
from ...api.http import RequestContext
|
||||
from ..memory_engine import MemoryEngine
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def generate_model_id(name: str) -> str:
|
||||
"""Generate a stable ID from mental model name."""
|
||||
# Normalize: lowercase, replace spaces/special chars with hyphens
|
||||
normalized = re.sub(r"[^a-z0-9]+", "-", name.lower()).strip("-")
|
||||
# Truncate to reasonable length
|
||||
return normalized[:50]
|
||||
|
||||
|
||||
def _parse_observations(observations_raw: list) -> list[Observation]:
|
||||
"""Parse raw observation dicts into typed Observation models."""
|
||||
observations: list[Observation] = []
|
||||
for obs in observations_raw:
|
||||
if not isinstance(obs, dict):
|
||||
continue
|
||||
|
||||
try:
|
||||
parsed = Observation(
|
||||
title=obs.get("title", ""),
|
||||
content=obs.get("content", ""),
|
||||
evidence=[
|
||||
ObservationEvidence(
|
||||
memory_id=ev.get("memory_id", ""),
|
||||
quote=ev.get("quote", ""),
|
||||
relevance=ev.get("relevance", ""),
|
||||
timestamp=ev.get("timestamp"),
|
||||
)
|
||||
for ev in obs.get("evidence", [])
|
||||
if isinstance(ev, dict)
|
||||
],
|
||||
created_at=obs.get("created_at"),
|
||||
)
|
||||
observations.append(parsed)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to parse observation: {e}")
|
||||
continue
|
||||
|
||||
return observations
|
||||
|
||||
|
||||
async def tool_lookup(
|
||||
conn: "Connection",
|
||||
bank_id: str,
|
||||
model_id: str | None = None,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str = "any",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
List or get mental models.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
bank_id: Bank identifier
|
||||
model_id: Optional specific model ID to get (if None, lists all)
|
||||
tags: Optional tags to filter models (when listing)
|
||||
tags_match: How to match tags - "any" (OR), "all" (AND)
|
||||
|
||||
Returns:
|
||||
Dict with either a list of models or a single model's details
|
||||
"""
|
||||
if model_id:
|
||||
# Get specific mental model with full details including observations
|
||||
row = await conn.fetchrow(
|
||||
"""
|
||||
SELECT id, subtype, name, description, observations, entity_id, last_updated
|
||||
FROM mental_models
|
||||
WHERE id = $1 AND bank_id = $2
|
||||
""",
|
||||
model_id,
|
||||
bank_id,
|
||||
)
|
||||
if row:
|
||||
# Parse observations JSON
|
||||
obs_data = row["observations"] or {"observations": []}
|
||||
if isinstance(obs_data, str):
|
||||
import json
|
||||
|
||||
obs_data = json.loads(obs_data)
|
||||
observations_raw = obs_data.get("observations", []) if isinstance(obs_data, dict) else obs_data
|
||||
|
||||
# Parse observations into typed models
|
||||
observations = _parse_observations(observations_raw)
|
||||
|
||||
return {
|
||||
"found": True,
|
||||
"model": {
|
||||
"id": row["id"],
|
||||
"subtype": row["subtype"],
|
||||
"name": row["name"],
|
||||
"description": row["description"],
|
||||
"observations": observations,
|
||||
"entity_id": str(row["entity_id"]) if row["entity_id"] else None,
|
||||
"last_updated": row["last_updated"].isoformat() if row["last_updated"] else None,
|
||||
},
|
||||
}
|
||||
return {"found": False, "model_id": model_id}
|
||||
else:
|
||||
# List mental models (compact: id, name, description only)
|
||||
# Full observations are retrieved via get_mental_model(model_id)
|
||||
# NOTE: Directives (subtype='directive') are excluded from listing -
|
||||
# they are injected into the system prompt, not discoverable via tools
|
||||
# Filter by tags if provided
|
||||
if tags:
|
||||
if tags_match == "all":
|
||||
# All tags must match
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT id, subtype, name, description
|
||||
FROM mental_models
|
||||
WHERE bank_id = $1 AND tags @> $2::varchar[] AND subtype != 'directive'
|
||||
ORDER BY last_updated DESC NULLS LAST, created_at DESC
|
||||
""",
|
||||
bank_id,
|
||||
tags,
|
||||
)
|
||||
else:
|
||||
# Any tag matches (OR) - default
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT id, subtype, name, description
|
||||
FROM mental_models
|
||||
WHERE bank_id = $1 AND tags && $2::varchar[] AND subtype != 'directive'
|
||||
ORDER BY last_updated DESC NULLS LAST, created_at DESC
|
||||
""",
|
||||
bank_id,
|
||||
tags,
|
||||
)
|
||||
else:
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT id, subtype, name, description
|
||||
FROM mental_models
|
||||
WHERE bank_id = $1 AND subtype != 'directive'
|
||||
ORDER BY last_updated DESC NULLS LAST, created_at DESC
|
||||
""",
|
||||
bank_id,
|
||||
)
|
||||
|
||||
return {
|
||||
"count": len(rows),
|
||||
"models": [
|
||||
{
|
||||
"id": row["id"],
|
||||
"subtype": row["subtype"],
|
||||
"name": row["name"],
|
||||
"description": row["description"],
|
||||
}
|
||||
for row in rows
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
async def tool_recall(
|
||||
memory_engine: "MemoryEngine",
|
||||
bank_id: str,
|
||||
query: str,
|
||||
request_context: "RequestContext",
|
||||
max_tokens: int = 2048,
|
||||
max_results: int = 50,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str = "any",
|
||||
connection_budget: int = 1,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Search memories using TEMPR retrieval.
|
||||
|
||||
Args:
|
||||
memory_engine: Memory engine instance
|
||||
bank_id: Bank identifier
|
||||
query: Search query
|
||||
request_context: Request context for authentication
|
||||
max_tokens: Maximum tokens for results (default 2048)
|
||||
max_results: Maximum number of results
|
||||
tags: Filter by tags (includes untagged memories)
|
||||
tags_match: How to match tags - "any" (OR), "all" (AND), or "exact"
|
||||
connection_budget: Max DB connections for this recall (default 1 for internal ops)
|
||||
|
||||
Returns:
|
||||
Dict with list of matching memories
|
||||
"""
|
||||
result = await memory_engine.recall_async(
|
||||
bank_id=bank_id,
|
||||
query=query,
|
||||
fact_type=["experience", "world"], # Exclude opinions
|
||||
max_tokens=max_tokens,
|
||||
enable_trace=False,
|
||||
request_context=request_context,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
_connection_budget=connection_budget,
|
||||
)
|
||||
|
||||
memories = []
|
||||
for m in result.results[:max_results]:
|
||||
memories.append(
|
||||
{
|
||||
"id": str(m.id),
|
||||
"text": m.text,
|
||||
"type": m.fact_type,
|
||||
"entities": m.entities or [],
|
||||
"occurred": m.occurred_start, # Already ISO format string
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"query": query,
|
||||
"count": len(memories),
|
||||
"memories": memories,
|
||||
}
|
||||
|
||||
|
||||
async def tool_learn(
|
||||
conn: "Connection",
|
||||
bank_id: str,
|
||||
input: MentalModelInput,
|
||||
tags: list[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Create a mental model placeholder with subtype='learned'.
|
||||
|
||||
The agent only specifies name and description - actual observations are generated
|
||||
in the background via refresh, similar to pinned models.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
bank_id: Bank identifier
|
||||
input: Mental model input data (name, description, optional entity_id)
|
||||
tags: Tags to apply to new mental models (from reflect context)
|
||||
|
||||
Returns:
|
||||
Dict with created model info including model_id for background generation
|
||||
"""
|
||||
model_id = generate_model_id(input.name)
|
||||
|
||||
# Parse entity_id if provided
|
||||
entity_uuid = None
|
||||
if input.entity_id:
|
||||
try:
|
||||
entity_uuid = uuid.UUID(input.entity_id)
|
||||
except ValueError:
|
||||
logger.warning(f"Invalid entity_id format: {input.entity_id}")
|
||||
|
||||
# Check if model exists
|
||||
existing = await conn.fetchrow(
|
||||
"SELECT id FROM mental_models WHERE id = $1 AND bank_id = $2",
|
||||
model_id,
|
||||
bank_id,
|
||||
)
|
||||
|
||||
if existing:
|
||||
# Update description only - observations will be regenerated
|
||||
await conn.execute(
|
||||
"""
|
||||
UPDATE mental_models SET
|
||||
description = $3,
|
||||
entity_id = $4
|
||||
WHERE id = $1 AND bank_id = $2
|
||||
""",
|
||||
model_id,
|
||||
bank_id,
|
||||
input.description,
|
||||
entity_uuid,
|
||||
)
|
||||
status = "updated"
|
||||
else:
|
||||
# Insert new model placeholder - observations will be generated in background
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO mental_models (id, bank_id, subtype, name, description, observations, entity_id, tags, created_at)
|
||||
VALUES ($1, $2, 'learned', $3, $4, '{}'::jsonb, $5, $6, NOW())
|
||||
""",
|
||||
model_id,
|
||||
bank_id,
|
||||
input.name,
|
||||
input.description,
|
||||
entity_uuid,
|
||||
tags or [],
|
||||
)
|
||||
status = "created"
|
||||
|
||||
logger.info(f"[REFLECT] Mental model '{model_id}' {status} in bank {bank_id} - pending background generation")
|
||||
|
||||
return {
|
||||
"status": status,
|
||||
"model_id": model_id,
|
||||
"name": input.name,
|
||||
"pending_generation": True,
|
||||
}
|
||||
|
||||
|
||||
async def tool_expand(
|
||||
conn: "Connection",
|
||||
bank_id: str,
|
||||
memory_ids: list[str],
|
||||
depth: str,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Expand multiple memories to get chunk or document context.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
bank_id: Bank identifier
|
||||
memory_ids: List of memory unit IDs
|
||||
depth: "chunk" or "document"
|
||||
|
||||
Returns:
|
||||
Dict with results array, each containing memory, chunk, and optionally document data
|
||||
"""
|
||||
if not memory_ids:
|
||||
return {"error": "memory_ids is required and must not be empty"}
|
||||
|
||||
# Validate and convert UUIDs
|
||||
valid_uuids: list[uuid.UUID] = []
|
||||
errors: dict[str, str] = {}
|
||||
for mid in memory_ids:
|
||||
try:
|
||||
valid_uuids.append(uuid.UUID(mid))
|
||||
except ValueError:
|
||||
errors[mid] = f"Invalid memory_id format: {mid}"
|
||||
|
||||
if not valid_uuids:
|
||||
return {"error": "No valid memory IDs provided", "details": errors}
|
||||
|
||||
# Batch fetch all memory units
|
||||
memories = await conn.fetch(
|
||||
"""
|
||||
SELECT id, text, chunk_id, document_id, fact_type, context
|
||||
FROM memory_units
|
||||
WHERE id = ANY($1) AND bank_id = $2
|
||||
""",
|
||||
valid_uuids,
|
||||
bank_id,
|
||||
)
|
||||
memory_map = {row["id"]: row for row in memories}
|
||||
|
||||
# Collect chunk_ids and document_ids for batch fetching
|
||||
chunk_ids = [m["chunk_id"] for m in memories if m["chunk_id"]]
|
||||
doc_ids_from_chunks: set[str] = set()
|
||||
doc_ids_direct: set[str] = set()
|
||||
|
||||
# Batch fetch all chunks
|
||||
chunk_map: dict[str, Any] = {}
|
||||
if chunk_ids:
|
||||
chunks = await conn.fetch(
|
||||
"""
|
||||
SELECT chunk_id, chunk_text, chunk_index, document_id
|
||||
FROM chunks
|
||||
WHERE chunk_id = ANY($1)
|
||||
""",
|
||||
chunk_ids,
|
||||
)
|
||||
chunk_map = {row["chunk_id"]: row for row in chunks}
|
||||
if depth == "document":
|
||||
doc_ids_from_chunks = {c["document_id"] for c in chunks if c["document_id"]}
|
||||
|
||||
# Collect direct document IDs (memories without chunks)
|
||||
if depth == "document":
|
||||
for m in memories:
|
||||
if not m["chunk_id"] and m["document_id"]:
|
||||
doc_ids_direct.add(m["document_id"])
|
||||
|
||||
# Batch fetch all documents
|
||||
doc_map: dict[str, Any] = {}
|
||||
all_doc_ids = list(doc_ids_from_chunks | doc_ids_direct)
|
||||
if all_doc_ids:
|
||||
docs = await conn.fetch(
|
||||
"""
|
||||
SELECT id, original_text, metadata, retain_params
|
||||
FROM documents
|
||||
WHERE id = ANY($1) AND bank_id = $2
|
||||
""",
|
||||
all_doc_ids,
|
||||
bank_id,
|
||||
)
|
||||
doc_map = {row["id"]: row for row in docs}
|
||||
|
||||
# Build results
|
||||
results: list[dict[str, Any]] = []
|
||||
for mid, mem_uuid in zip(memory_ids, valid_uuids):
|
||||
if mid in errors:
|
||||
results.append({"memory_id": mid, "error": errors[mid]})
|
||||
continue
|
||||
|
||||
memory = memory_map.get(mem_uuid)
|
||||
if not memory:
|
||||
results.append({"memory_id": mid, "error": f"Memory not found: {mid}"})
|
||||
continue
|
||||
|
||||
item: dict[str, Any] = {
|
||||
"memory_id": mid,
|
||||
"memory": {
|
||||
"id": str(memory["id"]),
|
||||
"text": memory["text"],
|
||||
"type": memory["fact_type"],
|
||||
"context": memory["context"],
|
||||
},
|
||||
}
|
||||
|
||||
# Add chunk if available
|
||||
if memory["chunk_id"] and memory["chunk_id"] in chunk_map:
|
||||
chunk = chunk_map[memory["chunk_id"]]
|
||||
item["chunk"] = {
|
||||
"id": chunk["chunk_id"],
|
||||
"text": chunk["chunk_text"],
|
||||
"index": chunk["chunk_index"],
|
||||
"document_id": chunk["document_id"],
|
||||
}
|
||||
# Add document if depth=document
|
||||
if depth == "document" and chunk["document_id"] in doc_map:
|
||||
doc = doc_map[chunk["document_id"]]
|
||||
item["document"] = {
|
||||
"id": doc["id"],
|
||||
"full_text": doc["original_text"],
|
||||
"metadata": doc["metadata"],
|
||||
"retain_params": doc["retain_params"],
|
||||
}
|
||||
elif memory["document_id"] and depth == "document" and memory["document_id"] in doc_map:
|
||||
# No chunk, but has document_id
|
||||
doc = doc_map[memory["document_id"]]
|
||||
item["document"] = {
|
||||
"id": doc["id"],
|
||||
"full_text": doc["original_text"],
|
||||
"metadata": doc["metadata"],
|
||||
"retain_params": doc["retain_params"],
|
||||
}
|
||||
|
||||
results.append(item)
|
||||
|
||||
return {"results": results, "count": len(results)}
|
||||
@@ -0,0 +1,218 @@
|
||||
"""
|
||||
Tool schema definitions for the reflect agent.
|
||||
|
||||
These are OpenAI-format tool definitions used with native tool calling.
|
||||
"""
|
||||
|
||||
# Tool definitions in OpenAI format
|
||||
TOOL_LIST_MENTAL_MODELS = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "list_mental_models",
|
||||
"description": "List all available mental models - your synthesized knowledge about entities, concepts, and events. Returns an array of models with id, name, and description.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"required": [],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
TOOL_GET_MENTAL_MODEL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_mental_model",
|
||||
"description": "Get full details of a specific mental model including all observations and memory references.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"model_id": {
|
||||
"type": "string",
|
||||
"description": "ID of the mental model (from list_mental_models results)",
|
||||
},
|
||||
},
|
||||
"required": ["model_id"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
TOOL_RECALL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "recall",
|
||||
"description": "Search memories using semantic + temporal retrieval. Returns relevant memories from experience and world knowledge, each with an 'id' you can reference.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "Search query string",
|
||||
},
|
||||
"max_tokens": {
|
||||
"type": "integer",
|
||||
"description": "Optional limit on result size (default 2048). Use higher values for broader searches.",
|
||||
},
|
||||
},
|
||||
"required": ["query"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
TOOL_LEARN = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "learn",
|
||||
"description": "Create a new mental model to track an important recurring topic. Use when you discover a person, project, concept, or pattern that appears frequently and would benefit from synthesized knowledge. The model content will be generated automatically.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "Human-readable name (e.g., 'Project Alpha', 'John Smith', 'Product Strategy')",
|
||||
},
|
||||
"description": {
|
||||
"type": "string",
|
||||
"description": "What to track and synthesize (e.g., 'Track goals, milestones, blockers, and key decisions for Project Alpha')",
|
||||
},
|
||||
},
|
||||
"required": ["name", "description"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
TOOL_EXPAND = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "expand",
|
||||
"description": "Get more context for one or more memories. Memory hierarchy: memory -> chunk -> document.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"memory_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of memory IDs from recall results (batch multiple for efficiency)",
|
||||
},
|
||||
"depth": {
|
||||
"type": "string",
|
||||
"enum": ["chunk", "document"],
|
||||
"description": "chunk: surrounding text chunk, document: full source document",
|
||||
},
|
||||
},
|
||||
"required": ["memory_ids", "depth"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
TOOL_DONE_ANSWER = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "done",
|
||||
"description": "Signal completion with your final answer. Use this when you have gathered enough information to answer the question.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"answer": {
|
||||
"type": "string",
|
||||
"description": "Your response as plain text. Do NOT use markdown formatting. NEVER include memory IDs, UUIDs, or 'Memory references' in this text - put IDs only in memory_ids array.",
|
||||
},
|
||||
"memory_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of memory IDs that support your answer (put IDs here, NOT in answer text)",
|
||||
},
|
||||
"model_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of mental model IDs that support your answer",
|
||||
},
|
||||
},
|
||||
"required": ["answer"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
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
|
||||
"""
|
||||
from typing import Any, cast
|
||||
|
||||
# Build rules list for description
|
||||
rules_list = "\n".join(f" {i + 1}. {rule}" for i, rule in enumerate(directive_rules))
|
||||
|
||||
# Build the tool with directive compliance field
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "done",
|
||||
"description": (
|
||||
"Signal completion with your final answer. IMPORTANT: You must confirm directive compliance before submitting. "
|
||||
"Your answer will be REJECTED if it violates any directive."
|
||||
),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"answer": {
|
||||
"type": "string",
|
||||
"description": "Your response as plain text. Do NOT use markdown formatting. NEVER include memory IDs, UUIDs, or 'Memory references' in this text - put IDs only in memory_ids array.",
|
||||
},
|
||||
"memory_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of memory IDs that support your answer (put IDs here, NOT in answer text)",
|
||||
},
|
||||
"model_ids": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
"description": "Array of mental model 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(enable_learn: bool = True, directive_rules: list[str] | None = None) -> list[dict]:
|
||||
"""
|
||||
Get the list of tools for the reflect agent.
|
||||
|
||||
Args:
|
||||
enable_learn: Whether to include the learn tool
|
||||
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 = []
|
||||
|
||||
# Include mental model tools for lookup
|
||||
tools.append(TOOL_LIST_MENTAL_MODELS)
|
||||
tools.append(TOOL_GET_MENTAL_MODEL)
|
||||
tools.append(TOOL_RECALL)
|
||||
|
||||
if enable_learn:
|
||||
tools.append(TOOL_LEARN)
|
||||
|
||||
tools.append(TOOL_EXPAND)
|
||||
|
||||
# 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
|
||||
@@ -6,12 +6,95 @@ API response models should be kept separate and convert from these core models t
|
||||
API stability even if internal models change.
|
||||
"""
|
||||
|
||||
from typing import Optional, List, Dict, Any
|
||||
from pydantic import BaseModel, Field, ConfigDict
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
# Valid fact types for recall operations (excludes 'observation' which is internal, and 'opinion' which is deprecated)
|
||||
VALID_RECALL_FACT_TYPES = frozenset(["world", "experience"])
|
||||
|
||||
|
||||
# Valid fact types for recall operations (excludes 'observation' which is internal)
|
||||
VALID_RECALL_FACT_TYPES = frozenset(["world", "experience", "opinion"])
|
||||
class LLMToolCall(BaseModel):
|
||||
"""A tool call requested by the LLM."""
|
||||
|
||||
id: str = Field(description="Unique identifier for this tool call")
|
||||
name: str = Field(description="Name of the tool to call")
|
||||
arguments: dict[str, Any] = Field(description="Arguments to pass to the tool")
|
||||
|
||||
|
||||
class LLMToolCallResult(BaseModel):
|
||||
"""Result from an LLM call that may include tool calls."""
|
||||
|
||||
content: str | None = Field(default=None, description="Text content if any")
|
||||
tool_calls: list[LLMToolCall] = Field(default_factory=list, description="Tool calls requested by the LLM")
|
||||
finish_reason: str | None = Field(default=None, description="Reason the LLM stopped: 'stop', 'tool_calls', etc.")
|
||||
|
||||
|
||||
class ToolCallTrace(BaseModel):
|
||||
"""A single tool call made during reflect."""
|
||||
|
||||
tool: str = Field(description="Tool name: lookup, recall, learn, expand")
|
||||
input: dict = Field(description="Tool input parameters")
|
||||
output: dict = Field(description="Tool output/result")
|
||||
duration_ms: int = Field(description="Execution time in milliseconds")
|
||||
iteration: int = Field(default=0, description="Iteration number (1-based) when this tool was called")
|
||||
|
||||
|
||||
class LLMCallTrace(BaseModel):
|
||||
"""A single LLM call made during reflect."""
|
||||
|
||||
scope: str = Field(description="Call scope: agent_1, agent_2, final, etc.")
|
||||
duration_ms: int = Field(description="Execution time in milliseconds")
|
||||
|
||||
|
||||
class MentalModelRef(BaseModel):
|
||||
"""Reference to a mental model accessed during reflect."""
|
||||
|
||||
id: str = Field(description="Mental model ID")
|
||||
name: str = Field(description="Mental model name")
|
||||
type: str = Field(description="Mental model type: entity, concept, event")
|
||||
subtype: str = Field(description="Mental model subtype: structural, emergent, learned")
|
||||
description: str = Field(description="Brief description")
|
||||
summary: str | None = Field(default=None, description="Full summary (when looked up in detail)")
|
||||
|
||||
|
||||
class 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")
|
||||
rules: list[str] = Field(default_factory=list, description="Directive rules/observations that were applied")
|
||||
|
||||
|
||||
class TokenUsage(BaseModel):
|
||||
"""
|
||||
Token usage metrics for LLM calls.
|
||||
|
||||
Tracks input/output tokens for a single request to enable
|
||||
per-request cost tracking and monitoring.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"input_tokens": 1500,
|
||||
"output_tokens": 500,
|
||||
"total_tokens": 2000,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
input_tokens: int = Field(default=0, description="Number of input/prompt tokens consumed")
|
||||
output_tokens: int = Field(default=0, description="Number of output/completion tokens generated")
|
||||
total_tokens: int = Field(default=0, description="Total tokens (input + output)")
|
||||
|
||||
def __add__(self, other: "TokenUsage") -> "TokenUsage":
|
||||
"""Allow aggregating token usage from multiple calls."""
|
||||
return TokenUsage(
|
||||
input_tokens=self.input_tokens + other.input_tokens,
|
||||
output_tokens=self.output_tokens + other.output_tokens,
|
||||
total_tokens=self.total_tokens + other.total_tokens,
|
||||
)
|
||||
|
||||
|
||||
class DispositionTraits(BaseModel):
|
||||
@@ -23,17 +106,12 @@ class DispositionTraits(BaseModel):
|
||||
- literalism: 1=flexible interpretation, 5=literal interpretation (how strictly to interpret information)
|
||||
- empathy: 1=detached, 5=empathetic (how much to consider emotional context)
|
||||
"""
|
||||
|
||||
skepticism: int = Field(ge=1, le=5, description="How skeptical vs trusting (1=trusting, 5=skeptical)")
|
||||
literalism: int = Field(ge=1, le=5, description="How literally to interpret information (1=flexible, 5=literal)")
|
||||
empathy: int = Field(ge=1, le=5, description="How much to consider emotional context (1=detached, 5=empathetic)")
|
||||
|
||||
model_config = ConfigDict(json_schema_extra={
|
||||
"example": {
|
||||
"skepticism": 3,
|
||||
"literalism": 3,
|
||||
"empathy": 3
|
||||
}
|
||||
})
|
||||
model_config = ConfigDict(json_schema_extra={"example": {"skepticism": 3, "literalism": 3, "empathy": 3}})
|
||||
|
||||
|
||||
class MemoryFact(BaseModel):
|
||||
@@ -43,38 +121,46 @@ class MemoryFact(BaseModel):
|
||||
This represents a unit of information stored in the memory system,
|
||||
including both the content and metadata.
|
||||
"""
|
||||
model_config = ConfigDict(json_schema_extra={
|
||||
"example": {
|
||||
"id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"text": "Alice works at Google on the AI team",
|
||||
"fact_type": "world",
|
||||
"entities": ["Alice", "Google"],
|
||||
"context": "work info",
|
||||
"occurred_start": "2024-01-15T10:30:00Z",
|
||||
"occurred_end": "2024-01-15T10:30:00Z",
|
||||
"mentioned_at": "2024-01-15T10:30:00Z",
|
||||
"document_id": "session_abc123",
|
||||
"metadata": {"source": "slack"},
|
||||
"chunk_id": "bank123_session_abc123_0",
|
||||
"activation": 0.95
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"text": "Alice works at Google on the AI team",
|
||||
"fact_type": "world",
|
||||
"entities": ["Alice", "Google"],
|
||||
"context": "work info",
|
||||
"occurred_start": "2024-01-15T10:30:00Z",
|
||||
"occurred_end": "2024-01-15T10:30:00Z",
|
||||
"mentioned_at": "2024-01-15T10:30:00Z",
|
||||
"document_id": "session_abc123",
|
||||
"metadata": {"source": "slack"},
|
||||
"chunk_id": "bank123_session_abc123_0",
|
||||
"activation": 0.95,
|
||||
"tags": ["user_a", "session_123"],
|
||||
}
|
||||
}
|
||||
})
|
||||
)
|
||||
|
||||
id: str = Field(description="Unique identifier for the memory fact")
|
||||
text: str = Field(description="The actual text content of the memory")
|
||||
fact_type: str = Field(description="Type of fact: 'world', 'experience', 'opinion', or 'observation'")
|
||||
entities: Optional[List[str]] = Field(None, description="Entity names mentioned in this fact")
|
||||
context: Optional[str] = Field(None, description="Additional context for the memory")
|
||||
occurred_start: Optional[str] = Field(None, description="ISO format date when the event started occurring")
|
||||
occurred_end: Optional[str] = Field(None, description="ISO format date when the event ended occurring")
|
||||
mentioned_at: Optional[str] = Field(None, description="ISO format date when the fact was mentioned/learned")
|
||||
document_id: Optional[str] = Field(None, description="ID of the document this memory belongs to")
|
||||
metadata: Optional[Dict[str, str]] = Field(None, description="User-defined metadata")
|
||||
chunk_id: Optional[str] = Field(None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)")
|
||||
entities: list[str] | None = Field(None, description="Entity names mentioned in this fact")
|
||||
context: str | None = Field(None, description="Additional context for the memory")
|
||||
occurred_start: str | None = Field(None, description="ISO format date when the event started occurring")
|
||||
occurred_end: str | None = Field(None, description="ISO format date when the event ended occurring")
|
||||
mentioned_at: str | None = Field(None, description="ISO format date when the fact was mentioned/learned")
|
||||
document_id: str | None = Field(None, description="ID of the document this memory belongs to")
|
||||
metadata: dict[str, str] | None = Field(None, description="User-defined metadata")
|
||||
chunk_id: str | None = Field(
|
||||
None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)"
|
||||
)
|
||||
tags: list[str] | None = Field(None, description="Visibility scope tags associated with this fact")
|
||||
|
||||
|
||||
class ChunkInfo(BaseModel):
|
||||
"""Information about a chunk."""
|
||||
|
||||
chunk_text: str = Field(description="The raw chunk text")
|
||||
chunk_index: int = Field(description="Index of the chunk within the document")
|
||||
truncated: bool = Field(default=False, description="Whether the chunk was truncated due to token limits")
|
||||
@@ -87,35 +173,33 @@ class RecallResult(BaseModel):
|
||||
Contains a list of matching memory facts and optional trace information
|
||||
for debugging and transparency.
|
||||
"""
|
||||
model_config = ConfigDict(json_schema_extra={
|
||||
"example": {
|
||||
"results": [
|
||||
{
|
||||
"id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"text": "Alice works at Google on the AI team",
|
||||
"fact_type": "world",
|
||||
"context": "work info",
|
||||
"occurred_start": "2024-01-15T10:30:00Z",
|
||||
"occurred_end": "2024-01-15T10:30:00Z",
|
||||
"activation": 0.95
|
||||
}
|
||||
],
|
||||
"trace": {
|
||||
"query": "What did Alice say about machine learning?",
|
||||
"num_results": 1
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"results": [
|
||||
{
|
||||
"id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"text": "Alice works at Google on the AI team",
|
||||
"fact_type": "world",
|
||||
"context": "work info",
|
||||
"occurred_start": "2024-01-15T10:30:00Z",
|
||||
"occurred_end": "2024-01-15T10:30:00Z",
|
||||
"activation": 0.95,
|
||||
}
|
||||
],
|
||||
"trace": {"query": "What did Alice say about machine learning?", "num_results": 1},
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
results: List[MemoryFact] = Field(description="List of memory facts matching the query")
|
||||
trace: Optional[Dict[str, Any]] = Field(None, description="Trace information for debugging")
|
||||
entities: Optional[Dict[str, "EntityState"]] = Field(
|
||||
None,
|
||||
description="Entity states for entities mentioned in results (keyed by canonical name)"
|
||||
)
|
||||
chunks: Optional[Dict[str, ChunkInfo]] = Field(
|
||||
None,
|
||||
description="Chunks for facts, keyed by '{document_id}_{chunk_index}'"
|
||||
|
||||
results: list[MemoryFact] = Field(description="List of memory facts matching the query")
|
||||
trace: dict[str, Any] | None = Field(None, description="Trace information for debugging")
|
||||
entities: dict[str, "EntityState"] | None = Field(
|
||||
None, description="Entity states for entities mentioned in results (keyed by canonical name)"
|
||||
)
|
||||
chunks: dict[str, ChunkInfo] | None = Field(
|
||||
None, description="Chunks for facts, keyed by '{document_id}_{chunk_index}'"
|
||||
)
|
||||
|
||||
|
||||
@@ -124,38 +208,63 @@ class ReflectResult(BaseModel):
|
||||
Result from a reflect operation.
|
||||
|
||||
Contains the formulated answer, the facts it was based on (organized by type),
|
||||
and any new opinions that were formed during the reflection process.
|
||||
any new opinions that were formed during the reflection process, and optionally
|
||||
structured output if a response schema was provided.
|
||||
"""
|
||||
model_config = ConfigDict(json_schema_extra={
|
||||
"example": {
|
||||
"text": "Based on my knowledge, machine learning is being actively used in healthcare...",
|
||||
"based_on": {
|
||||
"world": [
|
||||
{
|
||||
"id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"text": "Machine learning is used in medical diagnosis",
|
||||
"fact_type": "world",
|
||||
"context": "healthcare",
|
||||
"occurred_start": "2024-01-15T10:30:00Z",
|
||||
"occurred_end": "2024-01-15T10:30:00Z"
|
||||
}
|
||||
],
|
||||
"experience": [],
|
||||
"opinion": []
|
||||
},
|
||||
"new_opinions": [
|
||||
"Machine learning has great potential in healthcare"
|
||||
]
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"text": "Based on my knowledge, machine learning is being actively used in healthcare...",
|
||||
"based_on": {
|
||||
"world": [
|
||||
{
|
||||
"id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"text": "Machine learning is used in medical diagnosis",
|
||||
"fact_type": "world",
|
||||
"context": "healthcare",
|
||||
"occurred_start": "2024-01-15T10:30:00Z",
|
||||
"occurred_end": "2024-01-15T10:30:00Z",
|
||||
}
|
||||
],
|
||||
"experience": [],
|
||||
"opinion": [],
|
||||
},
|
||||
"new_opinions": ["Machine learning has great potential in healthcare"],
|
||||
"structured_output": {"summary": "ML in healthcare", "confidence": 0.9},
|
||||
"usage": {"input_tokens": 1500, "output_tokens": 500, "total_tokens": 2000},
|
||||
}
|
||||
}
|
||||
})
|
||||
)
|
||||
|
||||
text: str = Field(description="The formulated answer text")
|
||||
based_on: Dict[str, List[MemoryFact]] = Field(
|
||||
based_on: dict[str, list[MemoryFact]] = Field(
|
||||
description="Facts used to formulate the answer, organized by type (world, experience, opinion)"
|
||||
)
|
||||
new_opinions: List[str] = Field(
|
||||
new_opinions: list[str] = Field(default_factory=list, description="List of newly formed opinions during reflection")
|
||||
structured_output: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description="Structured output parsed according to the provided response schema. Only present when response_schema was provided.",
|
||||
)
|
||||
usage: TokenUsage | None = Field(
|
||||
default=None,
|
||||
description="Token usage metrics for the LLM calls made during this reflect operation.",
|
||||
)
|
||||
tool_trace: list[ToolCallTrace] = Field(
|
||||
default_factory=list,
|
||||
description="List of newly formed opinions during reflection"
|
||||
description="Trace of tool calls made during reflection. Only present when include.tool_calls is enabled.",
|
||||
)
|
||||
llm_trace: list[LLMCallTrace] = Field(
|
||||
default_factory=list,
|
||||
description="Trace of LLM calls made during reflection. Only present when include.tool_calls is enabled.",
|
||||
)
|
||||
mental_models: list[MentalModelRef] = Field(
|
||||
default_factory=list,
|
||||
description="Mental models accessed during reflection, including directives (subtype='directive').",
|
||||
)
|
||||
directives_applied: list[DirectiveRef] = Field(
|
||||
default_factory=list,
|
||||
description="Directive mental models that were applied during this reflection.",
|
||||
)
|
||||
|
||||
|
||||
@@ -166,12 +275,12 @@ class Opinion(BaseModel):
|
||||
Opinions represent the bank's formed perspectives on topics,
|
||||
with a confidence level indicating strength of belief.
|
||||
"""
|
||||
model_config = ConfigDict(json_schema_extra={
|
||||
"example": {
|
||||
"text": "Machine learning has great potential in healthcare",
|
||||
"confidence": 0.85
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {"text": "Machine learning has great potential in healthcare", "confidence": 0.85}
|
||||
}
|
||||
})
|
||||
)
|
||||
|
||||
text: str = Field(description="The opinion text")
|
||||
confidence: float = Field(description="Confidence score between 0.0 and 1.0")
|
||||
@@ -184,15 +293,15 @@ class EntityObservation(BaseModel):
|
||||
Observations are objective facts synthesized from multiple memory facts
|
||||
about an entity, without personality influence.
|
||||
"""
|
||||
model_config = ConfigDict(json_schema_extra={
|
||||
"example": {
|
||||
"text": "John is detail-oriented and works at Google",
|
||||
"mentioned_at": "2024-01-15T10:30:00Z"
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {"text": "John is detail-oriented and works at Google", "mentioned_at": "2024-01-15T10:30:00Z"}
|
||||
}
|
||||
})
|
||||
)
|
||||
|
||||
text: str = Field(description="The observation text")
|
||||
mentioned_at: Optional[str] = Field(None, description="ISO format date when this observation was created")
|
||||
mentioned_at: str | None = Field(None, description="ISO format date when this observation was created")
|
||||
|
||||
|
||||
class EntityState(BaseModel):
|
||||
@@ -201,20 +310,51 @@ class EntityState(BaseModel):
|
||||
|
||||
Contains observations synthesized from facts about the entity.
|
||||
"""
|
||||
model_config = ConfigDict(json_schema_extra={
|
||||
"example": {
|
||||
"entity_id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"canonical_name": "John",
|
||||
"observations": [
|
||||
{"text": "John is detail-oriented", "mentioned_at": "2024-01-15T10:30:00Z"},
|
||||
{"text": "John works at Google on the AI team", "mentioned_at": "2024-01-14T09:00:00Z"}
|
||||
]
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"entity_id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"canonical_name": "John",
|
||||
"observations": [
|
||||
{"text": "John is detail-oriented", "mentioned_at": "2024-01-15T10:30:00Z"},
|
||||
{"text": "John works at Google on the AI team", "mentioned_at": "2024-01-14T09:00:00Z"},
|
||||
],
|
||||
}
|
||||
}
|
||||
})
|
||||
)
|
||||
|
||||
entity_id: str = Field(description="Unique identifier for the entity")
|
||||
canonical_name: str = Field(description="Canonical name of the entity")
|
||||
observations: List[EntityObservation] = Field(
|
||||
default_factory=list,
|
||||
description="List of observations about this entity"
|
||||
observations: list[EntityObservation] = Field(
|
||||
default_factory=list, description="List of observations about this entity"
|
||||
)
|
||||
|
||||
|
||||
class MentalModel(BaseModel):
|
||||
"""
|
||||
A manually configured mental model for tracking specific topics/areas.
|
||||
|
||||
Mental models are user-defined focus areas that the agent should track
|
||||
and maintain summaries for, unlike auto-extracted entities.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"id": "team-dynamics",
|
||||
"name": "Team Dynamics",
|
||||
"description": "Track how the team collaborates, communication patterns, conflicts, and resolutions",
|
||||
"summary": "The team has strong collaboration...",
|
||||
"summary_updated_at": "2024-01-15T10:30:00Z",
|
||||
"created_at": "2024-01-10T08:00:00Z",
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
id: str = Field(description="Unique identifier (alphanumeric lowercase)")
|
||||
name: str = Field(description="Display name for the mental model")
|
||||
description: str = Field(description="Prompt/directions for what to track and summarize")
|
||||
summary: str | None = Field(None, description="Generated summary based on relevant facts")
|
||||
summary_updated_at: str | None = Field(None, description="ISO format date when summary was last updated")
|
||||
created_at: str = Field(description="ISO format date when the mental model was created")
|
||||
|
||||
@@ -12,23 +12,16 @@ This package contains modular components for the retain operation:
|
||||
- fact_storage: Handle fact insertion into database
|
||||
"""
|
||||
|
||||
from .types import (
|
||||
RetainContent,
|
||||
ExtractedFact,
|
||||
ProcessedFact,
|
||||
ChunkMetadata,
|
||||
EntityRef,
|
||||
CausalRelation,
|
||||
RetainBatch
|
||||
from . import (
|
||||
chunk_storage,
|
||||
deduplication,
|
||||
embedding_processing,
|
||||
entity_processing,
|
||||
fact_extraction,
|
||||
fact_storage,
|
||||
link_creation,
|
||||
)
|
||||
|
||||
from . import fact_extraction
|
||||
from . import embedding_processing
|
||||
from . import deduplication
|
||||
from . import entity_processing
|
||||
from . import link_creation
|
||||
from . import chunk_storage
|
||||
from . import fact_storage
|
||||
from .types import CausalRelation, ChunkMetadata, EntityRef, ExtractedFact, ProcessedFact, RetainBatch, RetainContent
|
||||
|
||||
__all__ = [
|
||||
# Types
|
||||
|
||||
@@ -1,13 +1,16 @@
|
||||
"""
|
||||
bank profile utilities for disposition and background management.
|
||||
bank profile utilities for disposition and mission management.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import Dict, Optional, TypedDict
|
||||
from typing import TypedDict
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from ..response_models import DispositionTraits
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -21,20 +24,21 @@ DEFAULT_DISPOSITION = {
|
||||
|
||||
class BankProfile(TypedDict):
|
||||
"""Type for bank profile data."""
|
||||
|
||||
name: str
|
||||
disposition: DispositionTraits
|
||||
background: str
|
||||
mission: str
|
||||
|
||||
|
||||
class BackgroundMergeResponse(BaseModel):
|
||||
"""LLM response for background merge with disposition inference."""
|
||||
background: str = Field(description="Merged background in first person perspective")
|
||||
disposition: DispositionTraits = Field(description="Inferred disposition traits (skepticism, literalism, empathy)")
|
||||
class MissionMergeResponse(BaseModel):
|
||||
"""LLM response for mission merge."""
|
||||
|
||||
mission: str = Field(description="Merged mission in first person perspective")
|
||||
|
||||
|
||||
async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
"""
|
||||
Get bank profile (name, disposition + background).
|
||||
Get bank profile (name, disposition + mission).
|
||||
Auto-creates bank with default values if not exists.
|
||||
|
||||
Args:
|
||||
@@ -42,16 +46,16 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
bank_id: bank IDentifier
|
||||
|
||||
Returns:
|
||||
BankProfile with name, typed DispositionTraits, and background
|
||||
BankProfile with name, typed DispositionTraits, and mission
|
||||
"""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Try to get existing bank
|
||||
row = await conn.fetchrow(
|
||||
"""
|
||||
SELECT name, disposition, background
|
||||
FROM banks WHERE bank_id = $1
|
||||
f"""
|
||||
SELECT name, disposition, mission
|
||||
FROM {fq_table("banks")} WHERE bank_id = $1
|
||||
""",
|
||||
bank_id
|
||||
bank_id,
|
||||
)
|
||||
|
||||
if row:
|
||||
@@ -63,34 +67,26 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
return BankProfile(
|
||||
name=row["name"],
|
||||
disposition=DispositionTraits(**disposition_data),
|
||||
background=row["background"]
|
||||
mission=row["mission"] or "",
|
||||
)
|
||||
|
||||
# Bank doesn't exist, create with defaults
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO banks (bank_id, name, disposition, background)
|
||||
f"""
|
||||
INSERT INTO {fq_table("banks")} (bank_id, name, disposition, mission)
|
||||
VALUES ($1, $2, $3::jsonb, $4)
|
||||
ON CONFLICT (bank_id) DO NOTHING
|
||||
""",
|
||||
bank_id,
|
||||
bank_id, # Default name is the bank_id
|
||||
json.dumps(DEFAULT_DISPOSITION),
|
||||
""
|
||||
"",
|
||||
)
|
||||
|
||||
return BankProfile(
|
||||
name=bank_id,
|
||||
disposition=DispositionTraits(**DEFAULT_DISPOSITION),
|
||||
background=""
|
||||
)
|
||||
return BankProfile(name=bank_id, disposition=DispositionTraits(**DEFAULT_DISPOSITION), mission="")
|
||||
|
||||
|
||||
async def update_bank_disposition(
|
||||
pool,
|
||||
bank_id: str,
|
||||
disposition: Dict[str, int]
|
||||
) -> None:
|
||||
async def update_bank_disposition(pool, bank_id: str, disposition: dict[str, int]) -> None:
|
||||
"""
|
||||
Update bank disposition traits.
|
||||
|
||||
@@ -104,275 +100,132 @@ async def update_bank_disposition(
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
await conn.execute(
|
||||
"""
|
||||
UPDATE banks
|
||||
f"""
|
||||
UPDATE {fq_table("banks")}
|
||||
SET disposition = $2::jsonb,
|
||||
updated_at = NOW()
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
json.dumps(disposition)
|
||||
json.dumps(disposition),
|
||||
)
|
||||
|
||||
|
||||
async def merge_bank_background(
|
||||
pool,
|
||||
llm_config,
|
||||
bank_id: str,
|
||||
new_info: str,
|
||||
update_disposition: bool = True
|
||||
) -> dict:
|
||||
async def set_bank_mission(pool, bank_id: str, mission: str) -> None:
|
||||
"""
|
||||
Merge new background information with existing background using LLM.
|
||||
Normalizes to first person ("I") and resolves conflicts.
|
||||
Optionally infers disposition traits from the merged background.
|
||||
Set bank mission (replacing any existing mission).
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
llm_config: LLM configuration for background merging
|
||||
bank_id: bank IDentifier
|
||||
new_info: New background information to add/merge
|
||||
update_disposition: If True, infer Big Five traits from background (default: True)
|
||||
mission: The mission text
|
||||
"""
|
||||
# Ensure bank exists first
|
||||
await get_bank_profile(pool, bank_id)
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("banks")}
|
||||
SET mission = $2,
|
||||
updated_at = NOW()
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
mission,
|
||||
)
|
||||
|
||||
|
||||
async def merge_bank_mission(pool, llm_config, bank_id: str, new_info: str) -> dict:
|
||||
"""
|
||||
Merge new mission information with existing mission using LLM.
|
||||
Normalizes to first person ("I") and resolves conflicts.
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
llm_config: LLM configuration for mission merging
|
||||
bank_id: bank IDentifier
|
||||
new_info: New mission information to add/merge
|
||||
|
||||
Returns:
|
||||
Dict with 'background' (str) and optionally 'disposition' (dict) keys
|
||||
Dict with 'mission' (str) key
|
||||
"""
|
||||
# Get current profile
|
||||
profile = await get_bank_profile(pool, bank_id)
|
||||
current_background = profile["background"]
|
||||
current_mission = profile["mission"]
|
||||
|
||||
# Use LLM to merge backgrounds and optionally infer disposition
|
||||
result = await _llm_merge_background(
|
||||
llm_config,
|
||||
current_background,
|
||||
new_info,
|
||||
infer_disposition=update_disposition
|
||||
)
|
||||
# Use LLM to merge missions
|
||||
result = await _llm_merge_mission(llm_config, current_mission, new_info)
|
||||
|
||||
merged_background = result["background"]
|
||||
inferred_disposition = result.get("disposition")
|
||||
merged_mission = result["mission"]
|
||||
|
||||
# Update in database
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
if inferred_disposition:
|
||||
# Update both background and disposition
|
||||
await conn.execute(
|
||||
"""
|
||||
UPDATE banks
|
||||
SET background = $2,
|
||||
disposition = $3::jsonb,
|
||||
updated_at = NOW()
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
merged_background,
|
||||
json.dumps(inferred_disposition)
|
||||
)
|
||||
else:
|
||||
# Update only background
|
||||
await conn.execute(
|
||||
"""
|
||||
UPDATE banks
|
||||
SET background = $2,
|
||||
updated_at = NOW()
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
merged_background
|
||||
)
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("banks")}
|
||||
SET mission = $2,
|
||||
updated_at = NOW()
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
merged_mission,
|
||||
)
|
||||
|
||||
response = {"background": merged_background}
|
||||
if inferred_disposition:
|
||||
response["disposition"] = inferred_disposition
|
||||
|
||||
return response
|
||||
return {"mission": merged_mission}
|
||||
|
||||
|
||||
async def _llm_merge_background(
|
||||
llm_config,
|
||||
current: str,
|
||||
new_info: str,
|
||||
infer_disposition: bool = False
|
||||
) -> dict:
|
||||
async def _llm_merge_mission(llm_config, current: str, new_info: str) -> dict:
|
||||
"""
|
||||
Use LLM to intelligently merge background information.
|
||||
Optionally infer Big Five disposition traits from the merged background.
|
||||
Use LLM to intelligently merge mission information.
|
||||
|
||||
Args:
|
||||
llm_config: LLM configuration to use
|
||||
current: Current background text
|
||||
current: Current mission text
|
||||
new_info: New information to merge
|
||||
infer_disposition: If True, also infer disposition traits
|
||||
|
||||
Returns:
|
||||
Dict with 'background' (str) and optionally 'disposition' (dict) keys
|
||||
Dict with 'mission' (str) key
|
||||
"""
|
||||
if infer_disposition:
|
||||
prompt = f"""You are helping maintain a memory bank's background/profile and infer their disposition. You MUST respond with ONLY valid JSON.
|
||||
prompt = f"""You are helping maintain an agent's mission statement.
|
||||
|
||||
Current background: {current if current else "(empty)"}
|
||||
Current mission: {current if current else "(empty)"}
|
||||
|
||||
New information to add: {new_info}
|
||||
|
||||
Instructions:
|
||||
1. Merge the new information with the current background
|
||||
2. If there are conflicts (e.g., different birthplaces), the NEW information overwrites the old
|
||||
3. Keep additions that don't conflict
|
||||
4. Output in FIRST PERSON ("I") perspective
|
||||
5. Be concise - keep merged background under 500 characters
|
||||
6. Infer disposition traits from the merged background (each 1-5 integer):
|
||||
- Skepticism: 1-5 (1=trusting, takes things at face value; 5=skeptical, questions everything)
|
||||
- Literalism: 1-5 (1=flexible interpretation, reads between lines; 5=literal, exact interpretation)
|
||||
- Empathy: 1-5 (1=detached, focuses on facts; 5=empathetic, considers emotional context)
|
||||
|
||||
CRITICAL: You MUST respond with ONLY a valid JSON object. No markdown, no code blocks, no explanations. Just the JSON.
|
||||
|
||||
Format:
|
||||
{{
|
||||
"background": "the merged background text in first person",
|
||||
"disposition": {{
|
||||
"skepticism": 3,
|
||||
"literalism": 3,
|
||||
"empathy": 3
|
||||
}}
|
||||
}}
|
||||
|
||||
Trait inference examples:
|
||||
- "I'm a lawyer" → skepticism: 4, literalism: 5, empathy: 2
|
||||
- "I'm a therapist" → skepticism: 2, literalism: 2, empathy: 5
|
||||
- "I'm an engineer" → skepticism: 3, literalism: 4, empathy: 3
|
||||
- "I've been burned before by trusting people" → skepticism: 5, literalism: 3, empathy: 3
|
||||
- "I try to understand what people really mean" → skepticism: 3, literalism: 2, empathy: 4
|
||||
- "I take contracts very seriously" → skepticism: 4, literalism: 5, empathy: 2"""
|
||||
else:
|
||||
prompt = f"""You are helping maintain a memory bank's background/profile.
|
||||
|
||||
Current background: {current if current else "(empty)"}
|
||||
|
||||
New information to add: {new_info}
|
||||
|
||||
Instructions:
|
||||
1. Merge the new information with the current background
|
||||
2. If there are conflicts (e.g., different birthplaces), the NEW information overwrites the old
|
||||
1. Merge the new information with the current mission
|
||||
2. If there are conflicts, the NEW information overwrites the old
|
||||
3. Keep additions that don't conflict
|
||||
4. Output in FIRST PERSON ("I") perspective
|
||||
5. Be concise - keep it under 500 characters
|
||||
6. Return ONLY the merged background text, no explanations
|
||||
6. Return ONLY the merged mission text, no explanations
|
||||
|
||||
Merged background:"""
|
||||
Merged mission:"""
|
||||
|
||||
try:
|
||||
# Prepare messages
|
||||
messages = [{"role": "user", "content": prompt}]
|
||||
|
||||
if infer_disposition:
|
||||
# Use structured output with Pydantic model for disposition inference
|
||||
try:
|
||||
parsed = await llm_config.call(
|
||||
messages=messages,
|
||||
response_format=BackgroundMergeResponse,
|
||||
scope="bank_background",
|
||||
temperature=0.3,
|
||||
max_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:
|
||||
@@ -383,13 +236,13 @@ async def list_banks(pool) -> list:
|
||||
pool: Database connection pool
|
||||
|
||||
Returns:
|
||||
List of dicts with bank_id, name, disposition, background, created_at, updated_at
|
||||
List of dicts with bank_id, name, disposition, mission, created_at, updated_at
|
||||
"""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT bank_id, name, disposition, background, created_at, updated_at
|
||||
FROM banks
|
||||
f"""
|
||||
SELECT bank_id, name, disposition, mission, created_at, updated_at
|
||||
FROM {fq_table("banks")}
|
||||
ORDER BY updated_at DESC
|
||||
"""
|
||||
)
|
||||
@@ -401,13 +254,15 @@ async def list_banks(pool) -> list:
|
||||
if isinstance(disposition_data, str):
|
||||
disposition_data = json.loads(disposition_data)
|
||||
|
||||
result.append({
|
||||
"bank_id": row["bank_id"],
|
||||
"name": row["name"],
|
||||
"disposition": disposition_data,
|
||||
"background": row["background"],
|
||||
"created_at": row["created_at"].isoformat() if row["created_at"] else None,
|
||||
"updated_at": row["updated_at"].isoformat() if row["updated_at"] else None,
|
||||
})
|
||||
result.append(
|
||||
{
|
||||
"bank_id": row["bank_id"],
|
||||
"name": row["name"],
|
||||
"disposition": disposition_data,
|
||||
"mission": row["mission"] or "",
|
||||
"created_at": row["created_at"].isoformat() if row["created_at"] else None,
|
||||
"updated_at": row["updated_at"].isoformat() if row["updated_at"] else None,
|
||||
}
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
@@ -3,20 +3,16 @@ Chunk storage for retain pipeline.
|
||||
|
||||
Handles storage of document chunks in the database.
|
||||
"""
|
||||
import logging
|
||||
from typing import List, Dict, Optional
|
||||
|
||||
import logging
|
||||
|
||||
from ..memory_engine import fq_table
|
||||
from .types import ChunkMetadata
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def store_chunks_batch(
|
||||
conn,
|
||||
bank_id: str,
|
||||
document_id: str,
|
||||
chunks: List[ChunkMetadata]
|
||||
) -> Dict[int, str]:
|
||||
async def store_chunks_batch(conn, bank_id: str, document_id: str, chunks: list[ChunkMetadata]) -> dict[int, str]:
|
||||
"""
|
||||
Store document chunks in the database.
|
||||
|
||||
@@ -47,24 +43,21 @@ async def store_chunks_batch(
|
||||
|
||||
# Batch insert all chunks
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO chunks (chunk_id, document_id, bank_id, chunk_text, chunk_index)
|
||||
f"""
|
||||
INSERT INTO {fq_table("chunks")} (chunk_id, document_id, bank_id, chunk_text, chunk_index)
|
||||
SELECT * FROM unnest($1::text[], $2::text[], $3::text[], $4::text[], $5::integer[])
|
||||
""",
|
||||
chunk_ids,
|
||||
[document_id] * len(chunk_texts),
|
||||
[bank_id] * len(chunk_texts),
|
||||
chunk_texts,
|
||||
chunk_indices
|
||||
chunk_indices,
|
||||
)
|
||||
|
||||
return chunk_id_map
|
||||
|
||||
|
||||
def map_facts_to_chunks(
|
||||
facts_chunk_indices: List[int],
|
||||
chunk_id_map: Dict[int, str]
|
||||
) -> List[Optional[str]]:
|
||||
def map_facts_to_chunks(facts_chunk_indices: list[int], chunk_id_map: dict[int, str]) -> list[str | None]:
|
||||
"""
|
||||
Map fact chunk indices to chunk IDs.
|
||||
|
||||
|
||||
@@ -3,22 +3,17 @@ Deduplication logic for retain pipeline.
|
||||
|
||||
Checks for duplicate facts using semantic similarity and temporal proximity.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import List
|
||||
from collections import defaultdict
|
||||
from datetime import UTC
|
||||
|
||||
from .types import ProcessedFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def check_duplicates_batch(
|
||||
conn,
|
||||
bank_id: str,
|
||||
facts: List[ProcessedFact],
|
||||
duplicate_checker_fn
|
||||
) -> List[bool]:
|
||||
async def check_duplicates_batch(conn, bank_id: str, facts: list[ProcessedFact], duplicate_checker_fn) -> list[bool]:
|
||||
"""
|
||||
Check which facts are duplicates using batched time-window queries.
|
||||
|
||||
@@ -47,16 +42,12 @@ async def check_duplicates_batch(
|
||||
|
||||
# Defensive: if both are None (shouldn't happen), use now()
|
||||
if fact_date is None:
|
||||
from datetime import datetime, timezone
|
||||
fact_date = datetime.now(timezone.utc)
|
||||
from datetime import datetime
|
||||
|
||||
fact_date = datetime.now(UTC)
|
||||
|
||||
# Round to 12-hour bucket to group similar times
|
||||
bucket_key = fact_date.replace(
|
||||
hour=(fact_date.hour // 12) * 12,
|
||||
minute=0,
|
||||
second=0,
|
||||
microsecond=0
|
||||
)
|
||||
bucket_key = fact_date.replace(hour=(fact_date.hour // 12) * 12, minute=0, second=0, microsecond=0)
|
||||
time_buckets[bucket_key].append((idx, fact))
|
||||
|
||||
# Process each bucket in batch
|
||||
@@ -68,14 +59,7 @@ async def check_duplicates_batch(
|
||||
embeddings = [item[1].embedding for item in bucket_items]
|
||||
|
||||
# Check duplicates for this time bucket
|
||||
dup_flags = await duplicate_checker_fn(
|
||||
conn,
|
||||
bank_id,
|
||||
texts,
|
||||
embeddings,
|
||||
bucket_date,
|
||||
time_window_hours=24
|
||||
)
|
||||
dup_flags = await duplicate_checker_fn(conn, bank_id, texts, embeddings, bucket_date, time_window_hours=24)
|
||||
|
||||
# Map results back to original indices
|
||||
for idx, is_dup in zip(indices, dup_flags):
|
||||
@@ -84,10 +68,7 @@ async def check_duplicates_batch(
|
||||
return all_is_duplicate
|
||||
|
||||
|
||||
def filter_duplicates(
|
||||
facts: List[ProcessedFact],
|
||||
is_duplicate_flags: List[bool]
|
||||
) -> List[ProcessedFact]:
|
||||
def filter_duplicates(facts: list[ProcessedFact], is_duplicate_flags: list[bool]) -> list[ProcessedFact]:
|
||||
"""
|
||||
Filter out duplicate facts based on duplicate flags.
|
||||
|
||||
|
||||
@@ -3,9 +3,8 @@ Embedding processing for retain pipeline.
|
||||
|
||||
Handles augmenting fact texts with temporal information and generating embeddings.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import List
|
||||
from datetime import datetime
|
||||
|
||||
from . import embedding_utils
|
||||
from .types import ExtractedFact
|
||||
@@ -13,7 +12,7 @@ from .types import ExtractedFact
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def augment_texts_with_dates(facts: List[ExtractedFact], format_date_fn) -> List[str]:
|
||||
def augment_texts_with_dates(facts: list[ExtractedFact], format_date_fn) -> list[str]:
|
||||
"""
|
||||
Augment fact texts with readable dates for better temporal matching.
|
||||
|
||||
@@ -37,10 +36,7 @@ def augment_texts_with_dates(facts: List[ExtractedFact], format_date_fn) -> List
|
||||
return augmented_texts
|
||||
|
||||
|
||||
async def generate_embeddings_batch(
|
||||
embeddings_model,
|
||||
texts: List[str]
|
||||
) -> List[List[float]]:
|
||||
async def generate_embeddings_batch(embeddings_model, texts: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generate embeddings for a batch of texts.
|
||||
|
||||
@@ -54,9 +50,6 @@ async def generate_embeddings_batch(
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
embeddings = await embedding_utils.generate_embeddings_batch(
|
||||
embeddings_model,
|
||||
texts
|
||||
)
|
||||
embeddings = await embedding_utils.generate_embeddings_batch(embeddings_model, texts)
|
||||
|
||||
return embeddings
|
||||
|
||||
@@ -4,12 +4,11 @@ Embedding generation utilities for memory units.
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import List
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def generate_embedding(embeddings_backend, text: str) -> List[float]:
|
||||
def generate_embedding(embeddings_backend, text: str) -> list[float]:
|
||||
"""
|
||||
Generate embedding for text using the provided embeddings backend.
|
||||
|
||||
@@ -27,7 +26,7 @@ def generate_embedding(embeddings_backend, text: str) -> List[float]:
|
||||
raise Exception(f"Failed to generate embedding: {str(e)}")
|
||||
|
||||
|
||||
async def generate_embeddings_batch(embeddings_backend, texts: List[str]) -> List[List[float]]:
|
||||
async def generate_embeddings_batch(embeddings_backend, texts: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generate embeddings for multiple texts using the provided embeddings backend.
|
||||
|
||||
@@ -47,7 +46,7 @@ async def generate_embeddings_batch(embeddings_backend, texts: List[str]) -> Lis
|
||||
embeddings = await loop.run_in_executor(
|
||||
None, # Use default thread pool
|
||||
embeddings_backend.encode,
|
||||
texts
|
||||
texts,
|
||||
)
|
||||
return embeddings
|
||||
except Exception as e:
|
||||
|
||||
@@ -3,12 +3,11 @@ Entity processing for retain pipeline.
|
||||
|
||||
Handles entity extraction, resolution, and link creation for stored facts.
|
||||
"""
|
||||
import logging
|
||||
from typing import List, Tuple, Dict, Any
|
||||
from uuid import UUID
|
||||
|
||||
from .types import ProcessedFact, EntityRef, EntityLink
|
||||
import logging
|
||||
|
||||
from . import link_utils
|
||||
from .types import EntityLink, ProcessedFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -17,18 +16,20 @@ async def process_entities_batch(
|
||||
entity_resolver,
|
||||
conn,
|
||||
bank_id: str,
|
||||
unit_ids: List[str],
|
||||
facts: List[ProcessedFact],
|
||||
log_buffer: List[str] = None
|
||||
) -> List[EntityLink]:
|
||||
unit_ids: list[str],
|
||||
facts: list[ProcessedFact],
|
||||
log_buffer: list[str] = None,
|
||||
user_entities_per_content: dict[int, list[dict]] = None,
|
||||
) -> list[EntityLink]:
|
||||
"""
|
||||
Process entities for all facts and create entity links.
|
||||
|
||||
This function:
|
||||
1. Extracts entity mentions from fact texts
|
||||
2. Resolves entity names to canonical entities
|
||||
3. Creates entity records in the database
|
||||
4. Returns entity links ready for insertion
|
||||
2. Merges user-provided entities with LLM-extracted entities
|
||||
3. Resolves entity names to canonical entities
|
||||
4. Creates entity records in the database
|
||||
5. Returns entity links ready for insertion
|
||||
|
||||
Args:
|
||||
entity_resolver: EntityResolver instance for entity resolution
|
||||
@@ -37,6 +38,7 @@ async def process_entities_batch(
|
||||
unit_ids: List of unit IDs (same length as facts)
|
||||
facts: List of ProcessedFact objects
|
||||
log_buffer: Optional buffer for detailed logging
|
||||
user_entities_per_content: Dict mapping content_index to list of user-provided entities
|
||||
|
||||
Returns:
|
||||
List of EntityLink objects for batch insertion
|
||||
@@ -47,15 +49,35 @@ async def process_entities_batch(
|
||||
if len(unit_ids) != len(facts):
|
||||
raise ValueError(f"Mismatch between unit_ids ({len(unit_ids)}) and facts ({len(facts)})")
|
||||
|
||||
user_entities_per_content = user_entities_per_content or {}
|
||||
|
||||
# Extract data for link_utils function
|
||||
fact_texts = [fact.fact_text for fact in facts]
|
||||
# Use occurred_start if available, otherwise use mentioned_at for entity timestamps
|
||||
fact_dates = [fact.occurred_start if fact.occurred_start is not None else fact.mentioned_at for fact in facts]
|
||||
# Convert EntityRef objects to dict format expected by link_utils
|
||||
entities_per_fact = [
|
||||
[{'text': entity.name, 'type': 'CONCEPT'} for entity in (fact.entities or [])]
|
||||
for fact in facts
|
||||
]
|
||||
|
||||
# Convert EntityRef objects to dict format and merge with user-provided entities
|
||||
entities_per_fact = []
|
||||
for fact in facts:
|
||||
# Start with LLM-extracted entities
|
||||
llm_entities = [{"text": entity.name, "type": "CONCEPT"} for entity in (fact.entities or [])]
|
||||
|
||||
# Get user entities for this content (use content_index from fact)
|
||||
user_entities = user_entities_per_content.get(fact.content_index, [])
|
||||
|
||||
# Merge with case-insensitive deduplication
|
||||
seen_texts = {e["text"].lower() for e in llm_entities}
|
||||
for user_entity in user_entities:
|
||||
if user_entity["text"].lower() not in seen_texts:
|
||||
llm_entities.append(
|
||||
{
|
||||
"text": user_entity["text"],
|
||||
"type": user_entity.get("type", "CONCEPT"),
|
||||
}
|
||||
)
|
||||
seen_texts.add(user_entity["text"].lower())
|
||||
|
||||
entities_per_fact.append(llm_entities)
|
||||
|
||||
# Use existing link_utils function for entity processing
|
||||
entity_links = await link_utils.extract_entities_batch_optimized(
|
||||
@@ -67,16 +89,13 @@ async def process_entities_batch(
|
||||
"", # context (not used in current implementation)
|
||||
fact_dates,
|
||||
entities_per_fact,
|
||||
log_buffer # Pass log_buffer for detailed logging
|
||||
log_buffer, # Pass log_buffer for detailed logging
|
||||
)
|
||||
|
||||
return entity_links
|
||||
|
||||
|
||||
async def insert_entity_links_batch(
|
||||
conn,
|
||||
entity_links: List[EntityLink]
|
||||
) -> None:
|
||||
async def insert_entity_links_batch(conn, entity_links: list[EntityLink]) -> None:
|
||||
"""
|
||||
Insert entity links in batch.
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -3,22 +3,19 @@ Fact storage for retain pipeline.
|
||||
|
||||
Handles insertion of facts into the database.
|
||||
"""
|
||||
import logging
|
||||
import json
|
||||
from typing import List, Optional
|
||||
from uuid import UUID
|
||||
|
||||
import json
|
||||
import logging
|
||||
|
||||
from ..memory_engine import fq_table
|
||||
from .types import ProcessedFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def insert_facts_batch(
|
||||
conn,
|
||||
bank_id: str,
|
||||
facts: List[ProcessedFact],
|
||||
document_id: Optional[str] = None
|
||||
) -> List[str]:
|
||||
conn, bank_id: str, facts: list[ProcessedFact], document_id: str | None = None
|
||||
) -> list[str]:
|
||||
"""
|
||||
Insert facts into the database in batch.
|
||||
|
||||
@@ -44,10 +41,10 @@ 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)
|
||||
@@ -62,22 +59,36 @@ async def insert_facts_batch(
|
||||
contexts.append(fact.context)
|
||||
fact_types.append(fact.fact_type)
|
||||
# confidence_score is only for opinion facts
|
||||
confidence_scores.append(1.0 if fact.fact_type == 'opinion' else None)
|
||||
access_counts.append(0) # Initial access count
|
||||
confidence_scores.append(1.0 if fact.fact_type == "opinion" else None)
|
||||
metadata_jsons.append(json.dumps(fact.metadata))
|
||||
chunk_ids.append(fact.chunk_id)
|
||||
# Use per-fact document_id if available, otherwise fallback to batch-level document_id
|
||||
document_ids.append(fact.document_id if fact.document_id else document_id)
|
||||
# Convert tags to JSON string for proper batch insertion (PostgreSQL unnest doesn't handle 2D arrays well)
|
||||
tags_list.append(json.dumps(fact.tags if fact.tags else []))
|
||||
|
||||
# Batch insert all facts
|
||||
# Note: tags are passed as JSON strings and converted back to varchar[] via jsonb_array_elements_text + array_agg
|
||||
results = await conn.fetch(
|
||||
"""
|
||||
INSERT INTO memory_units (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id)
|
||||
SELECT $1, * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[]
|
||||
f"""
|
||||
WITH input_data AS (
|
||||
SELECT * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::float[], $11::jsonb[], $12::text[], $13::text[], $14::jsonb[]
|
||||
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
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, metadata, chunk_id, document_id, tags)
|
||||
SELECT
|
||||
$1,
|
||||
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
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[]
|
||||
)
|
||||
FROM input_data
|
||||
RETURNING id
|
||||
""",
|
||||
bank_id,
|
||||
@@ -90,13 +101,13 @@ async def insert_facts_batch(
|
||||
contexts,
|
||||
fact_types,
|
||||
confidence_scores,
|
||||
access_counts,
|
||||
metadata_jsons,
|
||||
chunk_ids,
|
||||
document_ids
|
||||
document_ids,
|
||||
tags_list,
|
||||
)
|
||||
|
||||
unit_ids = [str(row['id']) for row in results]
|
||||
unit_ids = [str(row["id"]) for row in results]
|
||||
return unit_ids
|
||||
|
||||
|
||||
@@ -111,15 +122,15 @@ async def ensure_bank_exists(conn, bank_id: str) -> None:
|
||||
bank_id: Bank identifier
|
||||
"""
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO banks (bank_id, disposition, background)
|
||||
f"""
|
||||
INSERT INTO {fq_table("banks")} (bank_id, disposition, mission)
|
||||
VALUES ($1, $2::jsonb, $3)
|
||||
ON CONFLICT (bank_id) DO UPDATE
|
||||
SET updated_at = NOW()
|
||||
""",
|
||||
bank_id,
|
||||
'{"skepticism": 3, "literalism": 3, "empathy": 3}',
|
||||
""
|
||||
"",
|
||||
)
|
||||
|
||||
|
||||
@@ -129,7 +140,8 @@ async def handle_document_tracking(
|
||||
document_id: str,
|
||||
combined_content: str,
|
||||
is_first_batch: bool,
|
||||
retain_params: Optional[dict] = None
|
||||
retain_params: dict | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Handle document tracking in the database.
|
||||
@@ -141,6 +153,7 @@ async def handle_document_tracking(
|
||||
combined_content: Combined content text from all content items
|
||||
is_first_batch: Whether this is the first batch (for chunked operations)
|
||||
retain_params: Optional parameters passed during retain (context, event_date, etc.)
|
||||
document_tags: Optional list of tags to associate with the document
|
||||
"""
|
||||
import hashlib
|
||||
|
||||
@@ -151,20 +164,20 @@ async def handle_document_tracking(
|
||||
# Only delete on the first batch to avoid deleting data we just inserted
|
||||
if is_first_batch:
|
||||
await conn.fetchval(
|
||||
"DELETE FROM documents WHERE id = $1 AND bank_id = $2 RETURNING id",
|
||||
document_id, bank_id
|
||||
f"DELETE FROM {fq_table('documents')} WHERE id = $1 AND bank_id = $2 RETURNING id", document_id, bank_id
|
||||
)
|
||||
|
||||
# Insert document (or update if exists from concurrent operations)
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO documents (id, bank_id, original_text, content_hash, metadata, retain_params)
|
||||
VALUES ($1, $2, $3, $4, $5, $6)
|
||||
f"""
|
||||
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params, tags)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7)
|
||||
ON CONFLICT (id, bank_id) DO UPDATE
|
||||
SET original_text = EXCLUDED.original_text,
|
||||
content_hash = EXCLUDED.content_hash,
|
||||
metadata = EXCLUDED.metadata,
|
||||
retain_params = EXCLUDED.retain_params,
|
||||
tags = EXCLUDED.tags,
|
||||
updated_at = NOW()
|
||||
""",
|
||||
document_id,
|
||||
@@ -172,5 +185,6 @@ async def handle_document_tracking(
|
||||
combined_content,
|
||||
content_hash,
|
||||
json.dumps({}), # Empty metadata dict
|
||||
json.dumps(retain_params) if retain_params else None
|
||||
json.dumps(retain_params) if retain_params else None,
|
||||
document_tags or [],
|
||||
)
|
||||
|
||||
@@ -3,20 +3,16 @@ Link creation for retain pipeline.
|
||||
|
||||
Handles creation of temporal, semantic, and causal links between facts.
|
||||
"""
|
||||
import logging
|
||||
from typing import List
|
||||
|
||||
from .types import ProcessedFact, CausalRelation
|
||||
import logging
|
||||
|
||||
from . import link_utils
|
||||
from .types import ProcessedFact
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def create_temporal_links_batch(
|
||||
conn,
|
||||
bank_id: str,
|
||||
unit_ids: List[str]
|
||||
) -> int:
|
||||
async def create_temporal_links_batch(conn, bank_id: str, unit_ids: list[str]) -> int:
|
||||
"""
|
||||
Create temporal links between facts.
|
||||
|
||||
@@ -33,20 +29,10 @@ async def create_temporal_links_batch(
|
||||
if not unit_ids:
|
||||
return 0
|
||||
|
||||
return await link_utils.create_temporal_links_batch_per_fact(
|
||||
conn,
|
||||
bank_id,
|
||||
unit_ids,
|
||||
log_buffer=[]
|
||||
)
|
||||
return await link_utils.create_temporal_links_batch_per_fact(conn, bank_id, unit_ids, log_buffer=[])
|
||||
|
||||
|
||||
async def create_semantic_links_batch(
|
||||
conn,
|
||||
bank_id: str,
|
||||
unit_ids: List[str],
|
||||
embeddings: List[List[float]]
|
||||
) -> int:
|
||||
async def create_semantic_links_batch(conn, bank_id: str, unit_ids: list[str], embeddings: list[list[float]]) -> int:
|
||||
"""
|
||||
Create semantic links between facts.
|
||||
|
||||
@@ -67,20 +53,10 @@ async def create_semantic_links_batch(
|
||||
if len(unit_ids) != len(embeddings):
|
||||
raise ValueError(f"Mismatch between unit_ids ({len(unit_ids)}) and embeddings ({len(embeddings)})")
|
||||
|
||||
return await link_utils.create_semantic_links_batch(
|
||||
conn,
|
||||
bank_id,
|
||||
unit_ids,
|
||||
embeddings,
|
||||
log_buffer=[]
|
||||
)
|
||||
return await link_utils.create_semantic_links_batch(conn, bank_id, unit_ids, embeddings, log_buffer=[])
|
||||
|
||||
|
||||
async def create_causal_links_batch(
|
||||
conn,
|
||||
unit_ids: List[str],
|
||||
facts: List[ProcessedFact]
|
||||
) -> int:
|
||||
async def create_causal_links_batch(conn, unit_ids: list[str], facts: list[ProcessedFact]) -> int:
|
||||
"""
|
||||
Create causal links between facts.
|
||||
|
||||
@@ -108,9 +84,9 @@ async def create_causal_links_batch(
|
||||
# Convert CausalRelation objects to dicts
|
||||
relations_dicts = [
|
||||
{
|
||||
'relation_type': rel.relation_type,
|
||||
'target_fact_index': rel.target_fact_index,
|
||||
'strength': rel.strength
|
||||
"relation_type": rel.relation_type,
|
||||
"target_fact_index": rel.target_fact_index,
|
||||
"strength": rel.strength,
|
||||
}
|
||||
for rel in fact.causal_relations
|
||||
]
|
||||
@@ -118,10 +94,6 @@ async def create_causal_links_batch(
|
||||
else:
|
||||
causal_relations_per_fact.append([])
|
||||
|
||||
link_count = await link_utils.create_causal_links_batch(
|
||||
conn,
|
||||
unit_ids,
|
||||
causal_relations_per_fact
|
||||
)
|
||||
link_count = await link_utils.create_causal_links_batch(conn, unit_ids, causal_relations_per_fact)
|
||||
|
||||
return link_count
|
||||
|
||||
@@ -2,12 +2,12 @@
|
||||
Link creation utilities for temporal, semantic, and entity links.
|
||||
"""
|
||||
|
||||
import time
|
||||
import logging
|
||||
from typing import List
|
||||
from datetime import timedelta, datetime, timezone
|
||||
import time
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from uuid import UUID
|
||||
|
||||
from ..memory_engine import fq_table
|
||||
from .types import EntityLink
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -19,7 +19,7 @@ def _normalize_datetime(dt):
|
||||
return None
|
||||
if dt.tzinfo is None:
|
||||
# Naive datetime - assume UTC
|
||||
return dt.replace(tzinfo=timezone.utc)
|
||||
return dt.replace(tzinfo=UTC)
|
||||
return dt
|
||||
|
||||
|
||||
@@ -54,24 +54,26 @@ def compute_temporal_links(
|
||||
try:
|
||||
time_lower = unit_event_date_norm - timedelta(hours=time_window_hours)
|
||||
except OverflowError:
|
||||
time_lower = datetime.min.replace(tzinfo=timezone.utc)
|
||||
time_lower = datetime.min.replace(tzinfo=UTC)
|
||||
try:
|
||||
time_upper = unit_event_date_norm + timedelta(hours=time_window_hours)
|
||||
except OverflowError:
|
||||
time_upper = datetime.max.replace(tzinfo=timezone.utc)
|
||||
time_upper = datetime.max.replace(tzinfo=UTC)
|
||||
|
||||
# Filter candidates within this unit's time window
|
||||
matching_neighbors = [
|
||||
(row['id'], row['event_date'])
|
||||
(row["id"], row["event_date"])
|
||||
for row in candidates
|
||||
if time_lower <= _normalize_datetime(row['event_date']) <= time_upper
|
||||
if time_lower <= _normalize_datetime(row["event_date"]) <= time_upper
|
||||
][:10] # Limit to top 10
|
||||
|
||||
for recent_id, recent_event_date in matching_neighbors:
|
||||
# Calculate temporal proximity weight
|
||||
time_diff_hours = abs((unit_event_date_norm - _normalize_datetime(recent_event_date)).total_seconds() / 3600)
|
||||
time_diff_hours = abs(
|
||||
(unit_event_date_norm - _normalize_datetime(recent_event_date)).total_seconds() / 3600
|
||||
)
|
||||
weight = max(0.3, 1.0 - (time_diff_hours / time_window_hours))
|
||||
links.append((unit_id, str(recent_id), 'temporal', weight, None))
|
||||
links.append((unit_id, str(recent_id), "temporal", weight, None))
|
||||
|
||||
return links
|
||||
|
||||
@@ -99,17 +101,17 @@ def compute_temporal_query_bounds(
|
||||
try:
|
||||
min_date = min(all_dates) - timedelta(hours=time_window_hours)
|
||||
except OverflowError:
|
||||
min_date = datetime.min.replace(tzinfo=timezone.utc)
|
||||
min_date = datetime.min.replace(tzinfo=UTC)
|
||||
|
||||
try:
|
||||
max_date = max(all_dates) + timedelta(hours=time_window_hours)
|
||||
except OverflowError:
|
||||
max_date = datetime.max.replace(tzinfo=timezone.utc)
|
||||
max_date = datetime.max.replace(tzinfo=UTC)
|
||||
|
||||
return min_date, max_date
|
||||
|
||||
|
||||
def _log(log_buffer, message, level='info'):
|
||||
def _log(log_buffer, message, level="info"):
|
||||
"""Helper to log to buffer if available, otherwise use logger.
|
||||
|
||||
Args:
|
||||
@@ -117,7 +119,7 @@ def _log(log_buffer, message, level='info'):
|
||||
message: The log message
|
||||
level: 'info', 'debug', 'warning', or 'error'. Debug messages are not added to buffer.
|
||||
"""
|
||||
if level == 'debug':
|
||||
if level == "debug":
|
||||
# Debug messages only go to logger, not to buffer
|
||||
logger.debug(message)
|
||||
return
|
||||
@@ -125,23 +127,23 @@ def _log(log_buffer, message, level='info'):
|
||||
if log_buffer is not None:
|
||||
log_buffer.append(message)
|
||||
else:
|
||||
if level == 'info':
|
||||
if level == "info":
|
||||
logger.info(message)
|
||||
else:
|
||||
logger.log(logging.WARNING if level == 'warning' else logging.ERROR, message)
|
||||
logger.log(logging.WARNING if level == "warning" else logging.ERROR, message)
|
||||
|
||||
|
||||
async def extract_entities_batch_optimized(
|
||||
entity_resolver,
|
||||
conn,
|
||||
bank_id: str,
|
||||
unit_ids: List[str],
|
||||
sentences: List[str],
|
||||
unit_ids: list[str],
|
||||
sentences: list[str],
|
||||
context: str,
|
||||
fact_dates: List,
|
||||
llm_entities: List[List[dict]],
|
||||
log_buffer: List[str] = None,
|
||||
) -> List[tuple]:
|
||||
fact_dates: list,
|
||||
llm_entities: list[list[dict]],
|
||||
log_buffer: list[str] = None,
|
||||
) -> list[tuple]:
|
||||
"""
|
||||
Process LLM-extracted entities for ALL facts in batch.
|
||||
|
||||
@@ -171,15 +173,19 @@ async def extract_entities_batch_optimized(
|
||||
formatted_entities = []
|
||||
for ent in entity_list:
|
||||
# Handle both Entity objects and dicts
|
||||
if hasattr(ent, 'text'):
|
||||
if hasattr(ent, "text"):
|
||||
# Entity objects only have 'text', default type to 'CONCEPT'
|
||||
formatted_entities.append({'text': ent.text, 'type': 'CONCEPT'})
|
||||
formatted_entities.append({"text": ent.text, "type": "CONCEPT"})
|
||||
elif isinstance(ent, dict):
|
||||
formatted_entities.append({'text': ent.get('text', ''), 'type': ent.get('type', 'CONCEPT')})
|
||||
formatted_entities.append({"text": ent.get("text", ""), "type": ent.get("type", "CONCEPT")})
|
||||
all_entities.append(formatted_entities)
|
||||
|
||||
total_entities = sum(len(ents) for ents in all_entities)
|
||||
_log(log_buffer, f" [6.1] Process LLM entities: {total_entities} entities from {len(sentences)} facts in {time.time() - substep_start:.3f}s", level='debug')
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [6.1] Process LLM entities: {total_entities} entities from {len(sentences)} facts in {time.time() - substep_start:.3f}s",
|
||||
level="debug",
|
||||
)
|
||||
|
||||
# Step 2: Resolve entities in BATCH (much faster!)
|
||||
substep_start = time.time()
|
||||
@@ -195,13 +201,19 @@ async def extract_entities_batch_optimized(
|
||||
continue
|
||||
|
||||
for local_idx, entity in enumerate(entities):
|
||||
all_entities_flat.append({
|
||||
'text': entity['text'],
|
||||
'type': entity['type'],
|
||||
'nearby_entities': entities,
|
||||
})
|
||||
all_entities_flat.append(
|
||||
{
|
||||
"text": entity["text"],
|
||||
"type": entity["type"],
|
||||
"nearby_entities": entities,
|
||||
}
|
||||
)
|
||||
entity_to_unit.append((unit_id, local_idx, fact_date))
|
||||
_log(log_buffer, f" [6.2.1] Prepare entities: {len(all_entities_flat)} entities in {time.time() - substep_6_2_1_start:.3f}s", level='debug')
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [6.2.1] Prepare entities: {len(all_entities_flat)} entities in {time.time() - substep_6_2_1_start:.3f}s",
|
||||
level="debug",
|
||||
)
|
||||
|
||||
# Resolve ALL entities in one batch call
|
||||
if all_entities_flat:
|
||||
@@ -210,7 +222,7 @@ async def extract_entities_batch_optimized(
|
||||
|
||||
# Add per-entity dates to entity data for batch resolution
|
||||
for idx, (unit_id, local_idx, fact_date) in enumerate(entity_to_unit):
|
||||
all_entities_flat[idx]['event_date'] = fact_date
|
||||
all_entities_flat[idx]["event_date"] = fact_date
|
||||
|
||||
# Resolve ALL entities in ONE batch call (much faster than sequential buckets)
|
||||
# INSERT ... ON CONFLICT handles any race conditions at the DB level
|
||||
@@ -219,10 +231,14 @@ async def extract_entities_batch_optimized(
|
||||
entities_data=all_entities_flat,
|
||||
context=context,
|
||||
unit_event_date=None, # Not used when per-entity dates provided
|
||||
conn=conn # Use main transaction connection
|
||||
conn=conn, # Use main transaction connection
|
||||
)
|
||||
|
||||
_log(log_buffer, f" [6.2.2] Resolve entities: {len(all_entities_flat)} entities in single batch in {time.time() - substep_6_2_2_start:.3f}s", level='debug')
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [6.2.2] Resolve entities: {len(all_entities_flat)} entities in single batch in {time.time() - substep_6_2_2_start:.3f}s",
|
||||
level="debug",
|
||||
)
|
||||
|
||||
# [6.2.3] Create unit-entity links in BATCH
|
||||
substep_6_2_3_start = time.time()
|
||||
@@ -239,12 +255,24 @@ async def extract_entities_batch_optimized(
|
||||
|
||||
# Batch insert all unit-entity links (MUCH faster!)
|
||||
await entity_resolver.link_units_to_entities_batch(unit_entity_pairs, conn=conn)
|
||||
_log(log_buffer, f" [6.2.3] Create unit-entity links (batched): {len(unit_entity_pairs)} links in {time.time() - substep_6_2_3_start:.3f}s", level='debug')
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [6.2.3] Create unit-entity links (batched): {len(unit_entity_pairs)} links in {time.time() - substep_6_2_3_start:.3f}s",
|
||||
level="debug",
|
||||
)
|
||||
|
||||
_log(log_buffer, f" [6.2] Entity resolution (batched): {len(all_entities_flat)} entities resolved in {time.time() - step_6_2_start:.3f}s", level='debug')
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [6.2] Entity resolution (batched): {len(all_entities_flat)} entities resolved in {time.time() - step_6_2_start:.3f}s",
|
||||
level="debug",
|
||||
)
|
||||
else:
|
||||
unit_to_entity_ids = {}
|
||||
_log(log_buffer, f" [6.2] Entity resolution (batched): 0 entities in {time.time() - step_6_2_start:.3f}s", level='debug')
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [6.2] Entity resolution (batched): 0 entities in {time.time() - step_6_2_start:.3f}s",
|
||||
level="debug",
|
||||
)
|
||||
|
||||
# Step 3: Create entity links between units that share entities
|
||||
substep_start = time.time()
|
||||
@@ -253,39 +281,44 @@ async def extract_entities_batch_optimized(
|
||||
for entity_ids in unit_to_entity_ids.values():
|
||||
all_entity_ids.update(entity_ids)
|
||||
|
||||
_log(log_buffer, f" [6.3] Creating entity links for {len(all_entity_ids)} unique entities...", level='debug')
|
||||
_log(log_buffer, f" [6.3] Creating entity links for {len(all_entity_ids)} unique entities...", level="debug")
|
||||
|
||||
# Find all units that reference these entities (ONE batched query)
|
||||
entity_to_units = {}
|
||||
if all_entity_ids:
|
||||
query_start = time.time()
|
||||
import uuid
|
||||
|
||||
entity_id_list = [uuid.UUID(eid) if isinstance(eid, str) else eid for eid in all_entity_ids]
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT entity_id, unit_id
|
||||
FROM unit_entities
|
||||
FROM {fq_table("unit_entities")}
|
||||
WHERE entity_id = ANY($1::uuid[])
|
||||
""",
|
||||
entity_id_list
|
||||
entity_id_list,
|
||||
)
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [6.3.1] Query unit_entities: {len(rows)} rows in {time.time() - query_start:.3f}s",
|
||||
level="debug",
|
||||
)
|
||||
_log(log_buffer, f" [6.3.1] Query unit_entities: {len(rows)} rows in {time.time() - query_start:.3f}s", level='debug')
|
||||
|
||||
# Group by entity_id
|
||||
group_start = time.time()
|
||||
for row in rows:
|
||||
entity_id = row['entity_id']
|
||||
entity_id = row["entity_id"]
|
||||
if entity_id not in entity_to_units:
|
||||
entity_to_units[entity_id] = []
|
||||
entity_to_units[entity_id].append(row['unit_id'])
|
||||
_log(log_buffer, f" [6.3.2] Group by entity_id: {time.time() - group_start:.3f}s", level='debug')
|
||||
entity_to_units[entity_id].append(row["unit_id"])
|
||||
_log(log_buffer, f" [6.3.2] Group by entity_id: {time.time() - group_start:.3f}s", level="debug")
|
||||
|
||||
# Create bidirectional links between units that share entities
|
||||
# OPTIMIZATION: Limit links per entity to avoid N² explosion
|
||||
# Only link each new unit to the most recent MAX_LINKS_PER_ENTITY units
|
||||
MAX_LINKS_PER_ENTITY = 50 # Limit to prevent explosion when entity appears in many facts
|
||||
link_gen_start = time.time()
|
||||
links: List[EntityLink] = []
|
||||
links: list[EntityLink] = []
|
||||
new_unit_set = set(unit_ids) # Units from this batch
|
||||
|
||||
def to_uuid(val) -> UUID:
|
||||
@@ -299,27 +332,52 @@ async def extract_entities_batch_optimized(
|
||||
|
||||
# Link new units to each other (within batch) - also limited
|
||||
# For very common entities, limit within-batch links too
|
||||
new_units_to_link = new_units[-MAX_LINKS_PER_ENTITY:] if len(new_units) > MAX_LINKS_PER_ENTITY else new_units
|
||||
new_units_to_link = (
|
||||
new_units[-MAX_LINKS_PER_ENTITY:] if len(new_units) > MAX_LINKS_PER_ENTITY else new_units
|
||||
)
|
||||
for i, unit_id_1 in enumerate(new_units_to_link):
|
||||
for unit_id_2 in new_units_to_link[i+1:]:
|
||||
links.append(EntityLink(from_unit_id=to_uuid(unit_id_1), to_unit_id=to_uuid(unit_id_2), entity_id=entity_uuid))
|
||||
links.append(EntityLink(from_unit_id=to_uuid(unit_id_2), to_unit_id=to_uuid(unit_id_1), entity_id=entity_uuid))
|
||||
for unit_id_2 in new_units_to_link[i + 1 :]:
|
||||
links.append(
|
||||
EntityLink(
|
||||
from_unit_id=to_uuid(unit_id_1), to_unit_id=to_uuid(unit_id_2), entity_id=entity_uuid
|
||||
)
|
||||
)
|
||||
links.append(
|
||||
EntityLink(
|
||||
from_unit_id=to_uuid(unit_id_2), to_unit_id=to_uuid(unit_id_1), entity_id=entity_uuid
|
||||
)
|
||||
)
|
||||
|
||||
# Link new units to LIMITED existing units (most recent)
|
||||
existing_to_link = existing_units[-MAX_LINKS_PER_ENTITY:] # Take most recent
|
||||
for new_unit in new_units:
|
||||
for existing_unit in existing_to_link:
|
||||
links.append(EntityLink(from_unit_id=to_uuid(new_unit), to_unit_id=to_uuid(existing_unit), entity_id=entity_uuid))
|
||||
links.append(EntityLink(from_unit_id=to_uuid(existing_unit), to_unit_id=to_uuid(new_unit), entity_id=entity_uuid))
|
||||
links.append(
|
||||
EntityLink(
|
||||
from_unit_id=to_uuid(new_unit), to_unit_id=to_uuid(existing_unit), entity_id=entity_uuid
|
||||
)
|
||||
)
|
||||
links.append(
|
||||
EntityLink(
|
||||
from_unit_id=to_uuid(existing_unit), to_unit_id=to_uuid(new_unit), entity_id=entity_uuid
|
||||
)
|
||||
)
|
||||
|
||||
_log(log_buffer, f" [6.3.3] Generate {len(links)} links: {time.time() - link_gen_start:.3f}s", level='debug')
|
||||
_log(log_buffer, f" [6.3] Entity link creation: {len(links)} links for {len(all_entity_ids)} unique entities in {time.time() - substep_start:.3f}s", level='debug')
|
||||
_log(
|
||||
log_buffer, f" [6.3.3] Generate {len(links)} links: {time.time() - link_gen_start:.3f}s", level="debug"
|
||||
)
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [6.3] Entity link creation: {len(links)} links for {len(all_entity_ids)} unique entities in {time.time() - substep_start:.3f}s",
|
||||
level="debug",
|
||||
)
|
||||
|
||||
return links
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to extract entities in batch: {str(e)}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
raise
|
||||
|
||||
@@ -327,9 +385,9 @@ async def extract_entities_batch_optimized(
|
||||
async def create_temporal_links_batch_per_fact(
|
||||
conn,
|
||||
bank_id: str,
|
||||
unit_ids: List[str],
|
||||
unit_ids: list[str],
|
||||
time_window_hours: int = 24,
|
||||
log_buffer: List[str] = None,
|
||||
log_buffer: list[str] = None,
|
||||
) -> int:
|
||||
"""
|
||||
Create temporal links for multiple units, each with their own event_date.
|
||||
@@ -356,15 +414,18 @@ async def create_temporal_links_batch_per_fact(
|
||||
# Get the event_date for each new unit
|
||||
fetch_dates_start = time_mod.time()
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT id, event_date
|
||||
FROM memory_units
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id::text = ANY($1)
|
||||
""",
|
||||
unit_ids
|
||||
unit_ids,
|
||||
)
|
||||
new_units = {str(row["id"]): row["event_date"] for row in rows}
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [7.1] Fetch event_dates for {len(unit_ids)} units: {time_mod.time() - fetch_dates_start:.3f}s",
|
||||
)
|
||||
new_units = {str(row['id']): row['event_date'] for row in rows}
|
||||
_log(log_buffer, f" [7.1] Fetch event_dates for {len(unit_ids)} units: {time_mod.time() - fetch_dates_start:.3f}s")
|
||||
|
||||
# Fetch ALL potential temporal neighbors in ONE query (much faster!)
|
||||
# Get time range across all units with overflow protection
|
||||
@@ -372,9 +433,9 @@ async def create_temporal_links_batch_per_fact(
|
||||
|
||||
fetch_neighbors_start = time_mod.time()
|
||||
all_candidates = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT id, event_date
|
||||
FROM memory_units
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $1
|
||||
AND event_date BETWEEN $2 AND $3
|
||||
AND id::text != ALL($4)
|
||||
@@ -383,9 +444,12 @@ async def create_temporal_links_batch_per_fact(
|
||||
bank_id,
|
||||
min_date,
|
||||
max_date,
|
||||
unit_ids
|
||||
unit_ids,
|
||||
)
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [7.2] Fetch {len(all_candidates)} candidate neighbors (1 query): {time_mod.time() - fetch_neighbors_start:.3f}s",
|
||||
)
|
||||
_log(log_buffer, f" [7.2] Fetch {len(all_candidates)} candidate neighbors (1 query): {time_mod.time() - fetch_neighbors_start:.3f}s")
|
||||
|
||||
# Filter and create links in memory (much faster than N queries)
|
||||
link_gen_start = time_mod.time()
|
||||
@@ -408,21 +472,25 @@ async def create_temporal_links_batch_per_fact(
|
||||
if time_diff_hours <= time_window_hours:
|
||||
weight = max(0.3, 1.0 - (time_diff_hours / time_window_hours))
|
||||
# Create bidirectional links
|
||||
links.append((unit_id, other_id, 'temporal', weight, None))
|
||||
links.append((other_id, unit_id, 'temporal', weight, None))
|
||||
links.append((unit_id, other_id, "temporal", weight, None))
|
||||
links.append((other_id, unit_id, "temporal", weight, None))
|
||||
|
||||
_log(log_buffer, f" [7.3] Generate {len(links)} temporal links: {time_mod.time() - link_gen_start:.3f}s")
|
||||
|
||||
if links:
|
||||
insert_start = time_mod.time()
|
||||
await conn.executemany(
|
||||
"""
|
||||
INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id)
|
||||
VALUES ($1, $2, $3, $4, $5)
|
||||
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
|
||||
""",
|
||||
links
|
||||
)
|
||||
# Batch inserts to avoid timeout on large batches
|
||||
BATCH_SIZE = 1000
|
||||
for batch_start in range(0, len(links), BATCH_SIZE):
|
||||
batch = links[batch_start : batch_start + BATCH_SIZE]
|
||||
await conn.executemany(
|
||||
f"""
|
||||
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
|
||||
VALUES ($1, $2, $3, $4, $5)
|
||||
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
|
||||
""",
|
||||
batch,
|
||||
)
|
||||
_log(log_buffer, f" [7.4] Insert {len(links)} temporal links: {time_mod.time() - insert_start:.3f}s")
|
||||
|
||||
return len(links)
|
||||
@@ -430,6 +498,7 @@ async def create_temporal_links_batch_per_fact(
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create temporal links: {str(e)}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
raise
|
||||
|
||||
@@ -437,11 +506,11 @@ async def create_temporal_links_batch_per_fact(
|
||||
async def create_semantic_links_batch(
|
||||
conn,
|
||||
bank_id: str,
|
||||
unit_ids: List[str],
|
||||
embeddings: List[List[float]],
|
||||
unit_ids: list[str],
|
||||
embeddings: list[list[float]],
|
||||
top_k: int = 5,
|
||||
threshold: float = 0.7,
|
||||
log_buffer: List[str] = None,
|
||||
log_buffer: list[str] = None,
|
||||
) -> int:
|
||||
"""
|
||||
Create semantic links for multiple units efficiently.
|
||||
@@ -465,22 +534,26 @@ async def create_semantic_links_batch(
|
||||
|
||||
try:
|
||||
import time as time_mod
|
||||
|
||||
import numpy as np
|
||||
|
||||
# Fetch ALL existing units with embeddings in ONE query
|
||||
fetch_start = time_mod.time()
|
||||
all_existing = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT id, embedding
|
||||
FROM memory_units
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $1
|
||||
AND embedding IS NOT NULL
|
||||
AND id::text != ALL($2)
|
||||
""",
|
||||
bank_id,
|
||||
unit_ids
|
||||
unit_ids,
|
||||
)
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [8.1] Fetch {len(all_existing)} existing embeddings (1 query): {time_mod.time() - fetch_start:.3f}s",
|
||||
)
|
||||
_log(log_buffer, f" [8.1] Fetch {len(all_existing)} existing embeddings (1 query): {time_mod.time() - fetch_start:.3f}s")
|
||||
|
||||
# Convert to numpy for vectorized similarity computation
|
||||
compute_start = time_mod.time()
|
||||
@@ -488,15 +561,16 @@ async def create_semantic_links_batch(
|
||||
|
||||
if all_existing:
|
||||
# Convert existing embeddings to numpy array
|
||||
existing_ids = [str(row['id']) for row in all_existing]
|
||||
existing_ids = [str(row["id"]) for row in all_existing]
|
||||
# Stack embeddings as 2D array: (num_embeddings, embedding_dim)
|
||||
embedding_arrays = []
|
||||
for row in all_existing:
|
||||
raw_emb = row['embedding']
|
||||
raw_emb = row["embedding"]
|
||||
# Handle different pgvector formats
|
||||
if isinstance(raw_emb, str):
|
||||
# Parse string format: "[1.0, 2.0, ...]"
|
||||
import json
|
||||
|
||||
emb = np.array(json.loads(raw_emb), dtype=np.float32)
|
||||
elif isinstance(raw_emb, (list, tuple)):
|
||||
emb = np.array(raw_emb, dtype=np.float32)
|
||||
@@ -537,7 +611,7 @@ async def create_semantic_links_batch(
|
||||
similar_id = existing_ids[idx]
|
||||
# Clamp to [0, 1] to handle floating point precision issues
|
||||
similarity = float(min(1.0, max(0.0, similarities[idx])))
|
||||
all_links.append((unit_id, similar_id, 'semantic', similarity, None))
|
||||
all_links.append((unit_id, similar_id, "semantic", similarity, None))
|
||||
|
||||
# Also compute similarities WITHIN the new batch (new units to each other)
|
||||
# Apply the same top_k limit per unit as we do for existing units
|
||||
@@ -565,32 +639,42 @@ async def create_semantic_links_batch(
|
||||
other_id = unit_ids[other_idx]
|
||||
# Clamp to [0, 1] to handle floating point precision issues
|
||||
similarity = float(min(1.0, max(0.0, similarities[local_idx])))
|
||||
all_links.append((unit_id, other_id, 'semantic', similarity, None))
|
||||
all_links.append((unit_id, other_id, "semantic", similarity, None))
|
||||
|
||||
_log(log_buffer, f" [8.2] Compute similarities & generate {len(all_links)} semantic links: {time_mod.time() - compute_start:.3f}s")
|
||||
_log(
|
||||
log_buffer,
|
||||
f" [8.2] Compute similarities & generate {len(all_links)} semantic links: {time_mod.time() - compute_start:.3f}s",
|
||||
)
|
||||
|
||||
if all_links:
|
||||
insert_start = time_mod.time()
|
||||
await conn.executemany(
|
||||
"""
|
||||
INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id)
|
||||
VALUES ($1, $2, $3, $4, $5)
|
||||
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
|
||||
""",
|
||||
all_links
|
||||
# Batch inserts to avoid timeout on large batches
|
||||
BATCH_SIZE = 1000
|
||||
for batch_start in range(0, len(all_links), BATCH_SIZE):
|
||||
batch = all_links[batch_start : batch_start + BATCH_SIZE]
|
||||
await conn.executemany(
|
||||
f"""
|
||||
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
|
||||
VALUES ($1, $2, $3, $4, $5)
|
||||
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
|
||||
""",
|
||||
batch,
|
||||
)
|
||||
_log(
|
||||
log_buffer, f" [8.3] Insert {len(all_links)} semantic links: {time_mod.time() - insert_start:.3f}s"
|
||||
)
|
||||
_log(log_buffer, f" [8.3] Insert {len(all_links)} semantic links: {time_mod.time() - insert_start:.3f}s")
|
||||
|
||||
return len(all_links)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create semantic links: {str(e)}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
raise
|
||||
|
||||
|
||||
async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: int = 50000):
|
||||
async def insert_entity_links_batch(conn, links: list[EntityLink], chunk_size: int = 50000):
|
||||
"""
|
||||
Insert all entity links using COPY to temp table + INSERT for maximum speed.
|
||||
|
||||
@@ -606,7 +690,6 @@ async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: i
|
||||
if not links:
|
||||
return
|
||||
|
||||
import uuid as uuid_mod
|
||||
import time as time_mod
|
||||
|
||||
total_start = time_mod.time()
|
||||
@@ -633,28 +716,22 @@ async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: i
|
||||
convert_start = time_mod.time()
|
||||
records = []
|
||||
for link in links:
|
||||
records.append((
|
||||
link.from_unit_id,
|
||||
link.to_unit_id,
|
||||
link.link_type,
|
||||
link.weight,
|
||||
link.entity_id
|
||||
))
|
||||
records.append((link.from_unit_id, link.to_unit_id, link.link_type, link.weight, link.entity_id))
|
||||
logger.debug(f" [9.3] Convert {len(records)} records: {time_mod.time() - convert_start:.3f}s")
|
||||
|
||||
# Bulk load using COPY (fastest method)
|
||||
copy_start = time_mod.time()
|
||||
await conn.copy_records_to_table(
|
||||
'_temp_entity_links',
|
||||
"_temp_entity_links",
|
||||
records=records,
|
||||
columns=['from_unit_id', 'to_unit_id', 'link_type', 'weight', 'entity_id']
|
||||
columns=["from_unit_id", "to_unit_id", "link_type", "weight", "entity_id"],
|
||||
)
|
||||
logger.debug(f" [9.4] COPY {len(records)} records to temp table: {time_mod.time() - copy_start:.3f}s")
|
||||
|
||||
# Insert from temp table with ON CONFLICT (single query for all rows)
|
||||
insert_start = time_mod.time()
|
||||
await conn.execute("""
|
||||
INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id)
|
||||
await conn.execute(f"""
|
||||
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
|
||||
SELECT from_unit_id, to_unit_id, link_type, weight, entity_id
|
||||
FROM _temp_entity_links
|
||||
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
|
||||
@@ -665,8 +742,8 @@ async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: i
|
||||
|
||||
async def create_causal_links_batch(
|
||||
conn,
|
||||
unit_ids: List[str],
|
||||
causal_relations_per_fact: List[List[dict]],
|
||||
unit_ids: list[str],
|
||||
causal_relations_per_fact: list[list[dict]],
|
||||
) -> int:
|
||||
"""
|
||||
Create causal links between facts based on LLM-extracted causal relationships.
|
||||
@@ -694,6 +771,7 @@ async def create_causal_links_batch(
|
||||
|
||||
try:
|
||||
import time as time_mod
|
||||
|
||||
create_start = time_mod.time()
|
||||
|
||||
# Build links list
|
||||
@@ -705,12 +783,12 @@ async def create_causal_links_batch(
|
||||
from_unit_id = unit_ids[fact_idx]
|
||||
|
||||
for relation in causal_relations:
|
||||
target_idx = relation['target_fact_index']
|
||||
relation_type = relation['relation_type']
|
||||
strength = relation.get('strength', 1.0)
|
||||
target_idx = relation["target_fact_index"]
|
||||
relation_type = relation["relation_type"]
|
||||
strength = relation.get("strength", 1.0)
|
||||
|
||||
# Validate relation_type - must match database constraint
|
||||
valid_types = {'causes', 'caused_by', 'enables', 'prevents'}
|
||||
valid_types = {"causes", "caused_by", "enables", "prevents"}
|
||||
if relation_type not in valid_types:
|
||||
logger.error(
|
||||
f"Invalid relation_type '{relation_type}' (type: {type(relation_type).__name__}) "
|
||||
@@ -735,24 +813,25 @@ async def create_causal_links_batch(
|
||||
# weight is the strength of the relationship
|
||||
links.append((from_unit_id, to_unit_id, relation_type, strength, None))
|
||||
|
||||
|
||||
if links:
|
||||
insert_start = time_mod.time()
|
||||
try:
|
||||
await conn.executemany(
|
||||
"""
|
||||
INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id)
|
||||
f"""
|
||||
INSERT INTO {fq_table("memory_links")} (from_unit_id, to_unit_id, link_type, weight, entity_id)
|
||||
VALUES ($1, $2, $3, $4, $5)
|
||||
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
|
||||
""",
|
||||
links
|
||||
links,
|
||||
)
|
||||
except Exception as db_error:
|
||||
# Log the actual data being inserted for debugging
|
||||
logger.error(f"Database insert failed for causal links. Error: {db_error}")
|
||||
logger.error(f"Attempted to insert {len(links)} links. First few:")
|
||||
for i, link in enumerate(links[:3]):
|
||||
logger.error(f" Link {i}: from={link[0]}, to={link[1]}, type='{link[2]}' (repr={repr(link[2])}), weight={link[3]}, entity={link[4]}")
|
||||
logger.error(
|
||||
f" Link {i}: from={link[0]}, to={link[1]}, type='{link[2]}' (repr={repr(link[2])}), weight={link[3]}, entity={link[4]}"
|
||||
)
|
||||
raise
|
||||
|
||||
return len(links)
|
||||
@@ -760,5 +839,6 @@ async def create_causal_links_batch(
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create causal links: {str(e)}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
raise
|
||||
|
||||
@@ -1,264 +0,0 @@
|
||||
"""
|
||||
Observation regeneration for retain pipeline.
|
||||
|
||||
Regenerates entity observations as part of the retain transaction.
|
||||
"""
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import List, Dict, Optional
|
||||
|
||||
from ..search import observation_utils
|
||||
from . import embedding_utils
|
||||
from ..db_utils import acquire_with_retry
|
||||
from .types import EntityLink
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def utcnow():
|
||||
"""Get current UTC time."""
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
# Simple dataclass-like container for facts (avoid importing from memory_engine)
|
||||
class MemoryFactForObservation:
|
||||
def __init__(self, id: str, text: str, fact_type: str, context: str, occurred_start: Optional[str]):
|
||||
self.id = id
|
||||
self.text = text
|
||||
self.fact_type = fact_type
|
||||
self.context = context
|
||||
self.occurred_start = occurred_start
|
||||
|
||||
|
||||
async def regenerate_observations_batch(
|
||||
conn,
|
||||
embeddings_model,
|
||||
llm_config,
|
||||
bank_id: str,
|
||||
entity_links: List[EntityLink],
|
||||
log_buffer: List[str] = None
|
||||
) -> None:
|
||||
"""
|
||||
Regenerate observations for top entities in this batch.
|
||||
|
||||
Called INSIDE the retain transaction for atomicity - if observations
|
||||
fail, the entire retain batch is rolled back.
|
||||
|
||||
Args:
|
||||
conn: Database connection (from the retain transaction)
|
||||
embeddings_model: Embeddings model for generating observation embeddings
|
||||
llm_config: LLM configuration for observation extraction
|
||||
bank_id: Bank identifier
|
||||
entity_links: Entity links from this batch
|
||||
log_buffer: Optional log buffer for timing
|
||||
"""
|
||||
TOP_N_ENTITIES = 5
|
||||
MIN_FACTS_THRESHOLD = 5
|
||||
|
||||
if not entity_links:
|
||||
return
|
||||
|
||||
# Count mentions per entity in this batch
|
||||
entity_mention_counts: Dict[str, int] = {}
|
||||
for link in entity_links:
|
||||
if link.entity_id:
|
||||
entity_id = str(link.entity_id)
|
||||
entity_mention_counts[entity_id] = entity_mention_counts.get(entity_id, 0) + 1
|
||||
|
||||
if not entity_mention_counts:
|
||||
return
|
||||
|
||||
# Sort by mention count descending and take top N
|
||||
sorted_entities = sorted(
|
||||
entity_mention_counts.items(),
|
||||
key=lambda x: x[1],
|
||||
reverse=True
|
||||
)
|
||||
entities_to_process = [e[0] for e in sorted_entities[:TOP_N_ENTITIES]]
|
||||
|
||||
obs_start = time.time()
|
||||
|
||||
# Convert to UUIDs
|
||||
entity_uuids = [uuid.UUID(eid) if isinstance(eid, str) else eid for eid in entities_to_process]
|
||||
|
||||
# Batch query for entity names
|
||||
entity_rows = await conn.fetch(
|
||||
"""
|
||||
SELECT id, canonical_name FROM entities
|
||||
WHERE id = ANY($1) AND bank_id = $2
|
||||
""",
|
||||
entity_uuids, bank_id
|
||||
)
|
||||
entity_names = {row['id']: row['canonical_name'] for row in entity_rows}
|
||||
|
||||
# Batch query for fact counts
|
||||
fact_counts = await conn.fetch(
|
||||
"""
|
||||
SELECT ue.entity_id, COUNT(*) as cnt
|
||||
FROM unit_entities ue
|
||||
JOIN memory_units mu ON ue.unit_id = mu.id
|
||||
WHERE ue.entity_id = ANY($1) AND mu.bank_id = $2
|
||||
GROUP BY ue.entity_id
|
||||
""",
|
||||
entity_uuids, bank_id
|
||||
)
|
||||
entity_fact_counts = {row['entity_id']: row['cnt'] for row in fact_counts}
|
||||
|
||||
# Filter entities that meet the threshold
|
||||
entities_with_names = []
|
||||
for entity_id in entities_to_process:
|
||||
entity_uuid = uuid.UUID(entity_id) if isinstance(entity_id, str) else entity_id
|
||||
if entity_uuid not in entity_names:
|
||||
continue
|
||||
fact_count = entity_fact_counts.get(entity_uuid, 0)
|
||||
if fact_count >= MIN_FACTS_THRESHOLD:
|
||||
entities_with_names.append((entity_id, entity_names[entity_uuid]))
|
||||
|
||||
if not entities_with_names:
|
||||
return
|
||||
|
||||
# Process entities SEQUENTIALLY (asyncpg doesn't allow concurrent queries on same connection)
|
||||
# We must use the same connection to stay in the retain transaction
|
||||
total_observations = 0
|
||||
|
||||
for entity_id, entity_name in entities_with_names:
|
||||
try:
|
||||
obs_ids = await _regenerate_entity_observations(
|
||||
conn, embeddings_model, llm_config,
|
||||
bank_id, entity_id, entity_name
|
||||
)
|
||||
total_observations += len(obs_ids)
|
||||
except Exception as e:
|
||||
logger.error(f"[OBSERVATIONS] Error processing entity {entity_id}: {e}")
|
||||
|
||||
obs_time = time.time() - obs_start
|
||||
if log_buffer is not None:
|
||||
log_buffer.append(f"[11] Observations: {total_observations} observations for {len(entities_with_names)} entities in {obs_time:.3f}s")
|
||||
|
||||
|
||||
async def _regenerate_entity_observations(
|
||||
conn,
|
||||
embeddings_model,
|
||||
llm_config,
|
||||
bank_id: str,
|
||||
entity_id: str,
|
||||
entity_name: str
|
||||
) -> List[str]:
|
||||
"""
|
||||
Regenerate observations for a single entity.
|
||||
|
||||
Uses the provided connection (part of retain transaction).
|
||||
|
||||
Args:
|
||||
conn: Database connection (from the retain transaction)
|
||||
embeddings_model: Embeddings model
|
||||
llm_config: LLM configuration
|
||||
bank_id: Bank identifier
|
||||
entity_id: Entity UUID
|
||||
entity_name: Canonical name of the entity
|
||||
|
||||
Returns:
|
||||
List of created observation IDs
|
||||
"""
|
||||
entity_uuid = uuid.UUID(entity_id) if isinstance(entity_id, str) else entity_id
|
||||
|
||||
# Get all facts mentioning this entity (exclude observations themselves)
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.fact_type
|
||||
FROM memory_units mu
|
||||
JOIN unit_entities ue ON mu.id = ue.unit_id
|
||||
WHERE mu.bank_id = $1
|
||||
AND ue.entity_id = $2
|
||||
AND mu.fact_type IN ('world', 'experience')
|
||||
ORDER BY mu.occurred_start DESC
|
||||
LIMIT 50
|
||||
""",
|
||||
bank_id, entity_uuid
|
||||
)
|
||||
|
||||
if not rows:
|
||||
return []
|
||||
|
||||
# Convert to fact objects for observation extraction
|
||||
facts = []
|
||||
for row in rows:
|
||||
occurred_start = row['occurred_start'].isoformat() if row['occurred_start'] else None
|
||||
facts.append(MemoryFactForObservation(
|
||||
id=str(row['id']),
|
||||
text=row['text'],
|
||||
fact_type=row['fact_type'],
|
||||
context=row['context'],
|
||||
occurred_start=occurred_start
|
||||
))
|
||||
|
||||
# Extract observations using LLM
|
||||
observations = await observation_utils.extract_observations_from_facts(
|
||||
llm_config,
|
||||
entity_name,
|
||||
facts
|
||||
)
|
||||
|
||||
if not observations:
|
||||
return []
|
||||
|
||||
# Delete old observations for this entity
|
||||
await conn.execute(
|
||||
"""
|
||||
DELETE FROM memory_units
|
||||
WHERE id IN (
|
||||
SELECT mu.id
|
||||
FROM memory_units mu
|
||||
JOIN unit_entities ue ON mu.id = ue.unit_id
|
||||
WHERE mu.bank_id = $1
|
||||
AND mu.fact_type = 'observation'
|
||||
AND ue.entity_id = $2
|
||||
)
|
||||
""",
|
||||
bank_id, entity_uuid
|
||||
)
|
||||
|
||||
# Generate embeddings for new observations
|
||||
embeddings = await embedding_utils.generate_embeddings_batch(
|
||||
embeddings_model, observations
|
||||
)
|
||||
|
||||
# Insert new observations
|
||||
current_time = utcnow()
|
||||
created_ids = []
|
||||
|
||||
for obs_text, embedding in zip(observations, embeddings):
|
||||
result = await conn.fetchrow(
|
||||
"""
|
||||
INSERT INTO memory_units (
|
||||
bank_id, text, embedding, context, event_date,
|
||||
occurred_start, occurred_end, mentioned_at,
|
||||
fact_type, access_count
|
||||
)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, 'observation', 0)
|
||||
RETURNING id
|
||||
""",
|
||||
bank_id,
|
||||
obs_text,
|
||||
str(embedding),
|
||||
f"observation about {entity_name}",
|
||||
current_time,
|
||||
current_time,
|
||||
current_time,
|
||||
current_time
|
||||
)
|
||||
obs_id = str(result['id'])
|
||||
created_ids.append(obs_id)
|
||||
|
||||
# Link observation to entity
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO unit_entities (unit_id, entity_id)
|
||||
VALUES ($1, $2)
|
||||
""",
|
||||
uuid.UUID(obs_id), entity_uuid
|
||||
)
|
||||
|
||||
return created_ids
|
||||
@@ -3,31 +3,32 @@ Main orchestrator for the retain pipeline.
|
||||
|
||||
Coordinates all retain pipeline modules to store memories efficiently.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import List, Dict, Any, Optional
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from . import bank_utils
|
||||
from ..db_utils import acquire_with_retry
|
||||
from . import bank_utils
|
||||
|
||||
|
||||
def utcnow():
|
||||
"""Get current UTC time."""
|
||||
return datetime.now(timezone.utc)
|
||||
return datetime.now(UTC)
|
||||
|
||||
from .types import RetainContent, ExtractedFact, ProcessedFact, EntityLink
|
||||
|
||||
from ..response_models import TokenUsage
|
||||
from . import (
|
||||
fact_extraction,
|
||||
embedding_processing,
|
||||
deduplication,
|
||||
chunk_storage,
|
||||
fact_storage,
|
||||
deduplication,
|
||||
embedding_processing,
|
||||
entity_processing,
|
||||
fact_extraction,
|
||||
fact_storage,
|
||||
link_creation,
|
||||
observation_regeneration
|
||||
)
|
||||
from .types import EntityLink, ExtractedFact, ProcessedFact, RetainContent, RetainContentDict
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -37,16 +38,16 @@ async def retain_batch(
|
||||
embeddings_model,
|
||||
llm_config,
|
||||
entity_resolver,
|
||||
task_backend,
|
||||
format_date_fn,
|
||||
duplicate_checker_fn,
|
||||
bank_id: str,
|
||||
contents_dicts: List[Dict[str, Any]],
|
||||
document_id: Optional[str] = None,
|
||||
contents_dicts: list[RetainContentDict],
|
||||
document_id: str | None = None,
|
||||
is_first_batch: bool = True,
|
||||
fact_type_override: Optional[str] = None,
|
||||
confidence_score: Optional[float] = None,
|
||||
) -> List[List[str]]:
|
||||
fact_type_override: str | None = None,
|
||||
confidence_score: float | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
) -> tuple[list[list[str]], TokenUsage]:
|
||||
"""
|
||||
Process a batch of content through the retain pipeline.
|
||||
|
||||
@@ -55,7 +56,6 @@ async def retain_batch(
|
||||
embeddings_model: Embeddings model for generating embeddings
|
||||
llm_config: LLM configuration for fact extraction
|
||||
entity_resolver: Entity resolver for entity processing
|
||||
task_backend: Task backend for background jobs
|
||||
format_date_fn: Function to format datetime to readable string
|
||||
duplicate_checker_fn: Function to check for duplicate facts
|
||||
bank_id: Bank identifier
|
||||
@@ -64,19 +64,20 @@ async def retain_batch(
|
||||
is_first_batch: Whether this is the first batch
|
||||
fact_type_override: Override fact type for all facts
|
||||
confidence_score: Confidence score for opinions
|
||||
document_tags: Tags applied to all items in this batch
|
||||
|
||||
Returns:
|
||||
List of unit ID lists (one list per content item)
|
||||
Tuple of (unit ID lists, token usage for fact extraction)
|
||||
"""
|
||||
start_time = time.time()
|
||||
total_chars = sum(len(item.get("content", "")) for item in contents_dicts)
|
||||
|
||||
# Buffer all logs
|
||||
log_buffer = []
|
||||
log_buffer.append(f"{'='*60}")
|
||||
log_buffer.append(f"{'=' * 60}")
|
||||
log_buffer.append(f"RETAIN_BATCH START: {bank_id}")
|
||||
log_buffer.append(f"Batch size: {len(contents_dicts)} content items, {total_chars:,} chars")
|
||||
log_buffer.append(f"{'='*60}")
|
||||
log_buffer.append(f"{'=' * 60}")
|
||||
|
||||
# Get bank profile
|
||||
profile = await bank_utils.get_bank_profile(pool, bank_id)
|
||||
@@ -85,28 +86,89 @@ async def retain_batch(
|
||||
# Convert dicts to RetainContent objects
|
||||
contents = []
|
||||
for item in contents_dicts:
|
||||
# Merge item-level tags with document-level tags
|
||||
item_tags = item.get("tags", []) or []
|
||||
merged_tags = list(set(item_tags + (document_tags or [])))
|
||||
content = RetainContent(
|
||||
content=item["content"],
|
||||
context=item.get("context", ""),
|
||||
event_date=item.get("event_date") or utcnow(),
|
||||
metadata=item.get("metadata", {})
|
||||
metadata=item.get("metadata", {}),
|
||||
entities=item.get("entities", []),
|
||||
tags=merged_tags,
|
||||
)
|
||||
contents.append(content)
|
||||
|
||||
# Step 1: Extract facts from all contents
|
||||
step_start = time.time()
|
||||
extract_opinions = (fact_type_override == 'opinion')
|
||||
extract_opinions = fact_type_override == "opinion"
|
||||
|
||||
extracted_facts, chunks = await fact_extraction.extract_facts_from_contents(
|
||||
contents,
|
||||
llm_config,
|
||||
agent_name,
|
||||
extract_opinions
|
||||
extracted_facts, chunks, usage = await fact_extraction.extract_facts_from_contents(
|
||||
contents, llm_config, agent_name, extract_opinions
|
||||
)
|
||||
log_buffer.append(
|
||||
f"[1] Extract facts: {len(extracted_facts)} facts, {len(chunks)} chunks from {len(contents)} contents in {time.time() - step_start:.3f}s"
|
||||
)
|
||||
log_buffer.append(f"[1] Extract facts: {len(extracted_facts)} facts, {len(chunks)} chunks from {len(contents)} contents in {time.time() - step_start:.3f}s")
|
||||
|
||||
if not extracted_facts:
|
||||
return [[] for _ in contents]
|
||||
# Still need to create document if document_id was provided
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
await fact_storage.ensure_bank_exists(conn, bank_id)
|
||||
|
||||
# Handle document tracking even with no facts
|
||||
if document_id:
|
||||
combined_content = "\n".join([c.get("content", "") for c in contents_dicts])
|
||||
retain_params = {}
|
||||
if contents_dicts:
|
||||
first_item = contents_dicts[0]
|
||||
if first_item.get("context"):
|
||||
retain_params["context"] = first_item["context"]
|
||||
if first_item.get("event_date"):
|
||||
retain_params["event_date"] = (
|
||||
first_item["event_date"].isoformat()
|
||||
if hasattr(first_item["event_date"], "isoformat")
|
||||
else str(first_item["event_date"])
|
||||
)
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
)
|
||||
else:
|
||||
# Check for per-item document_ids
|
||||
from collections import defaultdict
|
||||
|
||||
contents_by_doc = defaultdict(list)
|
||||
for idx, content_dict in enumerate(contents_dicts):
|
||||
doc_id = content_dict.get("document_id")
|
||||
if doc_id:
|
||||
contents_by_doc[doc_id].append((idx, content_dict))
|
||||
|
||||
for doc_id, doc_contents in contents_by_doc.items():
|
||||
combined_content = "\n".join([c.get("content", "") for _, c in doc_contents])
|
||||
retain_params = {}
|
||||
if doc_contents:
|
||||
first_item = doc_contents[0][1]
|
||||
if first_item.get("context"):
|
||||
retain_params["context"] = first_item["context"]
|
||||
if first_item.get("event_date"):
|
||||
retain_params["event_date"] = (
|
||||
first_item["event_date"].isoformat()
|
||||
if hasattr(first_item["event_date"], "isoformat")
|
||||
else str(first_item["event_date"])
|
||||
)
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
)
|
||||
|
||||
total_time = time.time() - start_time
|
||||
logger.info(
|
||||
f"RETAIN_BATCH COMPLETE: 0 facts extracted from {len(contents)} contents in {total_time:.3f}s (document tracked, no facts)"
|
||||
)
|
||||
return [[] for _ in contents], usage
|
||||
|
||||
# Apply fact_type_override if provided
|
||||
if fact_type_override:
|
||||
@@ -130,6 +192,7 @@ async def retain_batch(
|
||||
|
||||
# Group contents by document_id for document tracking and chunk storage
|
||||
from collections import defaultdict
|
||||
|
||||
contents_by_doc = defaultdict(list)
|
||||
for idx, content_dict in enumerate(contents_dicts):
|
||||
doc_id = content_dict.get("document_id")
|
||||
@@ -155,12 +218,16 @@ async def retain_batch(
|
||||
if first_item.get("context"):
|
||||
retain_params["context"] = first_item["context"]
|
||||
if first_item.get("event_date"):
|
||||
retain_params["event_date"] = first_item["event_date"].isoformat() if hasattr(first_item["event_date"], "isoformat") else str(first_item["event_date"])
|
||||
retain_params["event_date"] = (
|
||||
first_item["event_date"].isoformat()
|
||||
if hasattr(first_item["event_date"], "isoformat")
|
||||
else str(first_item["event_date"])
|
||||
)
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
)
|
||||
document_ids_added.append(document_id)
|
||||
doc_id_mapping[None] = document_id # For backwards compatibility
|
||||
@@ -195,17 +262,29 @@ async def retain_batch(
|
||||
if first_item.get("context"):
|
||||
retain_params["context"] = first_item["context"]
|
||||
if first_item.get("event_date"):
|
||||
retain_params["event_date"] = first_item["event_date"].isoformat() if hasattr(first_item["event_date"], "isoformat") else str(first_item["event_date"])
|
||||
retain_params["event_date"] = (
|
||||
first_item["event_date"].isoformat()
|
||||
if hasattr(first_item["event_date"], "isoformat")
|
||||
else str(first_item["event_date"])
|
||||
)
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, actual_doc_id, combined_content, is_first_batch, retain_params
|
||||
conn,
|
||||
bank_id,
|
||||
actual_doc_id,
|
||||
combined_content,
|
||||
is_first_batch,
|
||||
retain_params,
|
||||
document_tags,
|
||||
)
|
||||
document_ids_added.append(actual_doc_id)
|
||||
|
||||
if document_ids_added:
|
||||
log_buffer.append(f"[2.5] Document tracking: {len(document_ids_added)} documents in {time.time() - step_start:.3f}s")
|
||||
log_buffer.append(
|
||||
f"[2.5] Document tracking: {len(document_ids_added)} documents in {time.time() - step_start:.3f}s"
|
||||
)
|
||||
|
||||
# Store chunks and map to facts for all documents
|
||||
step_start = time.time()
|
||||
@@ -230,7 +309,9 @@ async def retain_batch(
|
||||
for chunk_idx, chunk_id in chunk_id_map.items():
|
||||
chunk_id_map_by_doc[(doc_id, chunk_idx)] = chunk_id
|
||||
|
||||
log_buffer.append(f"[3] Store chunks: {len(chunks)} chunks for {len(chunks_by_doc)} documents in {time.time() - step_start:.3f}s")
|
||||
log_buffer.append(
|
||||
f"[3] Store chunks: {len(chunks)} chunks for {len(chunks_by_doc)} documents in {time.time() - step_start:.3f}s"
|
||||
)
|
||||
|
||||
# Map chunk_ids and document_ids to facts
|
||||
for fact, processed_fact in zip(extracted_facts, processed_facts):
|
||||
@@ -265,13 +346,15 @@ async def retain_batch(
|
||||
is_duplicate_flags = await deduplication.check_duplicates_batch(
|
||||
conn, bank_id, processed_facts, duplicate_checker_fn
|
||||
)
|
||||
log_buffer.append(f"[4] Deduplication: {sum(is_duplicate_flags)} duplicates in {time.time() - step_start:.3f}s")
|
||||
log_buffer.append(
|
||||
f"[4] Deduplication: {sum(is_duplicate_flags)} duplicates in {time.time() - step_start:.3f}s"
|
||||
)
|
||||
|
||||
# Filter out duplicates
|
||||
non_duplicate_facts = deduplication.filter_duplicates(processed_facts, is_duplicate_flags)
|
||||
|
||||
if not non_duplicate_facts:
|
||||
return [[] for _ in contents]
|
||||
return [[] for _ in contents], usage
|
||||
|
||||
# Insert facts (document_id is now stored per-fact)
|
||||
step_start = time.time()
|
||||
@@ -280,8 +363,18 @@ async def retain_batch(
|
||||
|
||||
# Process entities
|
||||
step_start = time.time()
|
||||
# Build map of content_index -> user entities for merging
|
||||
user_entities_per_content = {
|
||||
idx: content.entities for idx, content in enumerate(contents) if content.entities
|
||||
}
|
||||
entity_links = await entity_processing.process_entities_batch(
|
||||
entity_resolver, conn, bank_id, unit_ids, non_duplicate_facts, log_buffer
|
||||
entity_resolver,
|
||||
conn,
|
||||
bank_id,
|
||||
unit_ids,
|
||||
non_duplicate_facts,
|
||||
log_buffer,
|
||||
user_entities_per_content=user_entities_per_content,
|
||||
)
|
||||
log_buffer.append(f"[6] Process entities: {len(entity_links)} links in {time.time() - step_start:.3f}s")
|
||||
|
||||
@@ -293,62 +386,46 @@ async def retain_batch(
|
||||
# Create semantic links
|
||||
step_start = time.time()
|
||||
embeddings_for_links = [fact.embedding for fact in non_duplicate_facts]
|
||||
semantic_link_count = await link_creation.create_semantic_links_batch(conn, bank_id, unit_ids, embeddings_for_links)
|
||||
semantic_link_count = await link_creation.create_semantic_links_batch(
|
||||
conn, bank_id, unit_ids, embeddings_for_links
|
||||
)
|
||||
log_buffer.append(f"[8] Semantic links: {semantic_link_count} links in {time.time() - step_start:.3f}s")
|
||||
|
||||
# Insert entity links
|
||||
step_start = time.time()
|
||||
if entity_links:
|
||||
await entity_processing.insert_entity_links_batch(conn, entity_links)
|
||||
log_buffer.append(f"[9] Entity links: {len(entity_links) if entity_links else 0} links in {time.time() - step_start:.3f}s")
|
||||
log_buffer.append(
|
||||
f"[9] Entity links: {len(entity_links) if entity_links else 0} links in {time.time() - step_start:.3f}s"
|
||||
)
|
||||
|
||||
# Create causal links
|
||||
step_start = time.time()
|
||||
causal_link_count = await link_creation.create_causal_links_batch(conn, unit_ids, non_duplicate_facts)
|
||||
log_buffer.append(f"[10] Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s")
|
||||
|
||||
# Regenerate observations INSIDE transaction for atomicity
|
||||
await observation_regeneration.regenerate_observations_batch(
|
||||
conn,
|
||||
embeddings_model,
|
||||
llm_config,
|
||||
bank_id,
|
||||
entity_links,
|
||||
log_buffer
|
||||
)
|
||||
|
||||
# Map results back to original content items
|
||||
result_unit_ids = _map_results_to_contents(
|
||||
contents, extracted_facts, is_duplicate_flags, unit_ids
|
||||
)
|
||||
|
||||
# Trigger background tasks AFTER transaction commits (opinion reinforcement only)
|
||||
await _trigger_background_tasks(
|
||||
task_backend,
|
||||
bank_id,
|
||||
unit_ids,
|
||||
non_duplicate_facts
|
||||
)
|
||||
result_unit_ids = _map_results_to_contents(contents, extracted_facts, is_duplicate_flags, unit_ids)
|
||||
|
||||
# Log final summary
|
||||
total_time = time.time() - start_time
|
||||
log_buffer.append(f"{'='*60}")
|
||||
log_buffer.append(f"{'=' * 60}")
|
||||
log_buffer.append(f"RETAIN_BATCH COMPLETE: {len(unit_ids)} units in {total_time:.3f}s")
|
||||
if document_ids_added:
|
||||
log_buffer.append(f"Documents: {', '.join(document_ids_added)}")
|
||||
log_buffer.append(f"{'='*60}")
|
||||
log_buffer.append(f"{'=' * 60}")
|
||||
|
||||
logger.info("\n" + "\n".join(log_buffer) + "\n")
|
||||
|
||||
return result_unit_ids
|
||||
return result_unit_ids, usage
|
||||
|
||||
|
||||
def _map_results_to_contents(
|
||||
contents: List[RetainContent],
|
||||
extracted_facts: List[ExtractedFact],
|
||||
is_duplicate_flags: List[bool],
|
||||
unit_ids: List[str]
|
||||
) -> List[List[str]]:
|
||||
contents: list[RetainContent],
|
||||
extracted_facts: list[ExtractedFact],
|
||||
is_duplicate_flags: list[bool],
|
||||
unit_ids: list[str],
|
||||
) -> list[list[str]]:
|
||||
"""
|
||||
Map created unit IDs back to original content items.
|
||||
|
||||
@@ -371,22 +448,3 @@ def _map_results_to_contents(
|
||||
result_unit_ids.append(content_unit_ids)
|
||||
|
||||
return result_unit_ids
|
||||
|
||||
|
||||
async def _trigger_background_tasks(
|
||||
task_backend,
|
||||
bank_id: str,
|
||||
unit_ids: List[str],
|
||||
facts: List[ProcessedFact],
|
||||
) -> None:
|
||||
"""Trigger opinion reinforcement as background task (after transaction commits)."""
|
||||
# Trigger opinion reinforcement if there are entities
|
||||
fact_entities = [[e.name for e in fact.entities] for fact in facts]
|
||||
if any(fact_entities):
|
||||
await task_backend.submit_task({
|
||||
'type': 'reinforce_opinion',
|
||||
'bank_id': bank_id,
|
||||
'created_unit_ids': unit_ids,
|
||||
'unit_texts': [fact.fact_text for fact in facts],
|
||||
'unit_entities': fact_entities
|
||||
})
|
||||
|
||||
@@ -6,11 +6,38 @@ from content input to fact storage.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional, Dict, Any
|
||||
from datetime import datetime
|
||||
from datetime import UTC, datetime
|
||||
from typing import TypedDict
|
||||
from uuid import UUID
|
||||
|
||||
|
||||
class RetainContentDict(TypedDict, total=False):
|
||||
"""Type definition for content items in retain_batch_async.
|
||||
|
||||
Fields:
|
||||
content: Text content to store (required)
|
||||
context: Context about the content (optional)
|
||||
event_date: When the content occurred (optional, defaults to now)
|
||||
metadata: Custom key-value metadata (optional)
|
||||
document_id: Document ID for this content item (optional)
|
||||
entities: User-provided entities to merge with extracted entities (optional)
|
||||
tags: Visibility scope tags for this content item (optional)
|
||||
"""
|
||||
|
||||
content: str # Required
|
||||
context: str
|
||||
event_date: datetime
|
||||
metadata: dict[str, str]
|
||||
document_id: str
|
||||
entities: list[dict[str, str]] # [{"text": "...", "type": "..."}]
|
||||
tags: list[str] # Visibility scope tags
|
||||
|
||||
|
||||
def _now_utc() -> datetime:
|
||||
"""Factory function for default event_date."""
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RetainContent:
|
||||
"""
|
||||
@@ -18,16 +45,13 @@ class RetainContent:
|
||||
|
||||
Represents a single piece of content to extract facts from.
|
||||
"""
|
||||
|
||||
content: str
|
||||
context: str = ""
|
||||
event_date: Optional[datetime] = None
|
||||
metadata: Dict[str, str] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
"""Ensure event_date is set."""
|
||||
if self.event_date is None:
|
||||
from datetime import datetime, timezone
|
||||
self.event_date = datetime.now(timezone.utc)
|
||||
event_date: datetime = field(default_factory=_now_utc)
|
||||
metadata: dict[str, str] = field(default_factory=dict)
|
||||
entities: list[dict[str, str]] = field(default_factory=list) # User-provided entities
|
||||
tags: list[str] = field(default_factory=list) # Visibility scope tags
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -37,6 +61,7 @@ class ChunkMetadata:
|
||||
|
||||
Used to track which facts were extracted from which chunks.
|
||||
"""
|
||||
|
||||
chunk_text: str
|
||||
fact_count: int
|
||||
content_index: int # Index of the source content
|
||||
@@ -50,9 +75,10 @@ class EntityRef:
|
||||
|
||||
Entities are extracted by the LLM during fact extraction.
|
||||
"""
|
||||
|
||||
name: str
|
||||
canonical_name: Optional[str] = None # Resolved canonical name
|
||||
entity_id: Optional[UUID] = None # Resolved entity ID
|
||||
canonical_name: str | None = None # Resolved canonical name
|
||||
entity_id: UUID | None = None # Resolved entity ID
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -62,6 +88,7 @@ class CausalRelation:
|
||||
|
||||
Represents how one fact causes, enables, or prevents another.
|
||||
"""
|
||||
|
||||
relation_type: str # "causes", "enables", "prevents", "caused_by"
|
||||
target_fact_index: int # Index of the target fact in the batch
|
||||
strength: float = 1.0 # Strength of the causal relationship
|
||||
@@ -74,20 +101,22 @@ class ExtractedFact:
|
||||
|
||||
This is the raw output from fact extraction before processing.
|
||||
"""
|
||||
|
||||
fact_text: str
|
||||
fact_type: str # "world", "experience", "opinion", "observation"
|
||||
entities: List[str] = field(default_factory=list)
|
||||
occurred_start: Optional[datetime] = None
|
||||
occurred_end: Optional[datetime] = None
|
||||
where: Optional[str] = None # WHERE the fact occurred or is about
|
||||
causal_relations: List[CausalRelation] = field(default_factory=list)
|
||||
entities: list[str] = field(default_factory=list)
|
||||
occurred_start: datetime | None = None
|
||||
occurred_end: datetime | None = None
|
||||
where: str | None = None # WHERE the fact occurred or is about
|
||||
causal_relations: list[CausalRelation] = field(default_factory=list)
|
||||
|
||||
# Context from the content item
|
||||
content_index: int = 0 # Which content this fact came from
|
||||
chunk_index: int = 0 # Which chunk this fact came from
|
||||
context: str = ""
|
||||
mentioned_at: Optional[datetime] = None
|
||||
metadata: Dict[str, str] = field(default_factory=dict)
|
||||
mentioned_at: datetime | None = None
|
||||
metadata: dict[str, str] = field(default_factory=dict)
|
||||
tags: list[str] = field(default_factory=list) # Visibility scope tags
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -97,37 +126,44 @@ class ProcessedFact:
|
||||
|
||||
Includes resolved entities, embeddings, and all necessary fields.
|
||||
"""
|
||||
|
||||
# Core fact data
|
||||
fact_text: str
|
||||
fact_type: str
|
||||
embedding: List[float]
|
||||
embedding: list[float]
|
||||
|
||||
# Temporal data
|
||||
occurred_start: Optional[datetime]
|
||||
occurred_end: Optional[datetime]
|
||||
occurred_start: datetime | None
|
||||
occurred_end: datetime | None
|
||||
mentioned_at: datetime
|
||||
|
||||
# Context and metadata
|
||||
context: str
|
||||
metadata: Dict[str, str]
|
||||
metadata: dict[str, str]
|
||||
|
||||
# Location data
|
||||
where: Optional[str] = None
|
||||
where: str | None = None
|
||||
|
||||
# Entities
|
||||
entities: List[EntityRef] = field(default_factory=list)
|
||||
entities: list[EntityRef] = field(default_factory=list)
|
||||
|
||||
# Causal relations
|
||||
causal_relations: List[CausalRelation] = field(default_factory=list)
|
||||
causal_relations: list[CausalRelation] = field(default_factory=list)
|
||||
|
||||
# Chunk reference
|
||||
chunk_id: Optional[str] = None
|
||||
chunk_id: str | None = None
|
||||
|
||||
# Document reference (denormalized for query performance)
|
||||
document_id: Optional[str] = None
|
||||
document_id: str | None = None
|
||||
|
||||
# DB fields (set after insertion)
|
||||
unit_id: Optional[UUID] = None
|
||||
unit_id: UUID | None = None
|
||||
|
||||
# Track which content this fact came from (for user entity merging)
|
||||
content_index: int = 0
|
||||
|
||||
# Visibility scope tags
|
||||
tags: list[str] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def is_duplicate(self) -> bool:
|
||||
@@ -136,10 +172,8 @@ class ProcessedFact:
|
||||
|
||||
@staticmethod
|
||||
def from_extracted_fact(
|
||||
extracted_fact: 'ExtractedFact',
|
||||
embedding: List[float],
|
||||
chunk_id: Optional[str] = None
|
||||
) -> 'ProcessedFact':
|
||||
extracted_fact: "ExtractedFact", embedding: list[float], chunk_id: str | None = None
|
||||
) -> "ProcessedFact":
|
||||
"""
|
||||
Create ProcessedFact from ExtractedFact.
|
||||
|
||||
@@ -151,12 +185,12 @@ class ProcessedFact:
|
||||
Returns:
|
||||
ProcessedFact ready for storage
|
||||
"""
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
|
||||
# Use occurred dates only if explicitly provided by LLM
|
||||
occurred_start = extracted_fact.occurred_start
|
||||
occurred_end = extracted_fact.occurred_end
|
||||
mentioned_at = extracted_fact.mentioned_at or datetime.now(timezone.utc)
|
||||
mentioned_at = extracted_fact.mentioned_at or datetime.now(UTC)
|
||||
|
||||
# Convert entity strings to EntityRef objects
|
||||
entities = [EntityRef(name=name) for name in extracted_fact.entities]
|
||||
@@ -172,7 +206,9 @@ class ProcessedFact:
|
||||
metadata=extracted_fact.metadata,
|
||||
entities=entities,
|
||||
causal_relations=extracted_fact.causal_relations,
|
||||
chunk_id=chunk_id
|
||||
chunk_id=chunk_id,
|
||||
content_index=extracted_fact.content_index,
|
||||
tags=extracted_fact.tags,
|
||||
)
|
||||
|
||||
|
||||
@@ -183,10 +219,11 @@ class EntityLink:
|
||||
|
||||
Used for entity-based graph connections in the memory graph.
|
||||
"""
|
||||
|
||||
from_unit_id: UUID
|
||||
to_unit_id: UUID
|
||||
entity_id: UUID
|
||||
link_type: str = 'entity'
|
||||
link_type: str = "entity"
|
||||
weight: float = 1.0
|
||||
|
||||
|
||||
@@ -197,24 +234,26 @@ class RetainBatch:
|
||||
|
||||
Tracks all facts, chunks, and metadata for a batch operation.
|
||||
"""
|
||||
|
||||
bank_id: str
|
||||
contents: List[RetainContent]
|
||||
document_id: Optional[str] = None
|
||||
fact_type_override: Optional[str] = None
|
||||
confidence_score: Optional[float] = None
|
||||
contents: list[RetainContent]
|
||||
document_id: str | None = None
|
||||
fact_type_override: str | None = None
|
||||
confidence_score: float | None = None
|
||||
document_tags: list[str] = field(default_factory=list) # Tags applied to all items
|
||||
|
||||
# Extracted data (populated during processing)
|
||||
extracted_facts: List[ExtractedFact] = field(default_factory=list)
|
||||
processed_facts: List[ProcessedFact] = field(default_factory=list)
|
||||
chunks: List[ChunkMetadata] = field(default_factory=list)
|
||||
extracted_facts: list[ExtractedFact] = field(default_factory=list)
|
||||
processed_facts: list[ProcessedFact] = field(default_factory=list)
|
||||
chunks: list[ChunkMetadata] = field(default_factory=list)
|
||||
|
||||
# Results (populated after storage)
|
||||
unit_ids_by_content: List[List[str]] = field(default_factory=list)
|
||||
unit_ids_by_content: list[list[str]] = field(default_factory=list)
|
||||
|
||||
def get_facts_for_content(self, content_index: int) -> List[ExtractedFact]:
|
||||
def get_facts_for_content(self, content_index: int) -> list[ExtractedFact]:
|
||||
"""Get all extracted facts for a specific content item."""
|
||||
return [f for f in self.extracted_facts if f.content_index == content_index]
|
||||
|
||||
def get_chunks_for_content(self, content_index: int) -> List[ChunkMetadata]:
|
||||
def get_chunks_for_content(self, content_index: int) -> list[ChunkMetadata]:
|
||||
"""Get all chunks for a specific content item."""
|
||||
return [c for c in self.chunks if c.content_index == content_index]
|
||||
|
||||
@@ -7,15 +7,15 @@ Provides modular search architecture:
|
||||
- Reranking: Pluggable strategies (heuristic, cross-encoder)
|
||||
"""
|
||||
|
||||
from .retrieval import (
|
||||
retrieve_parallel,
|
||||
get_default_graph_retriever,
|
||||
set_default_graph_retriever,
|
||||
ParallelRetrievalResult,
|
||||
)
|
||||
from .graph_retrieval import GraphRetriever, BFSGraphRetriever
|
||||
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
|
||||
from .mpfp_retrieval import MPFPGraphRetriever
|
||||
from .reranking import CrossEncoderReranker
|
||||
from .retrieval import (
|
||||
ParallelRetrievalResult,
|
||||
get_default_graph_retriever,
|
||||
retrieve_parallel,
|
||||
set_default_graph_retriever,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"retrieve_parallel",
|
||||
|
||||
@@ -2,15 +2,12 @@
|
||||
Helper functions for hybrid search (semantic + BM25 + graph).
|
||||
"""
|
||||
|
||||
from typing import List, Dict, Any, Tuple
|
||||
import asyncio
|
||||
from .types import RetrievalResult, MergedCandidate
|
||||
from typing import Any
|
||||
|
||||
from .types import MergedCandidate, RetrievalResult
|
||||
|
||||
|
||||
def reciprocal_rank_fusion(
|
||||
result_lists: List[List[RetrievalResult]],
|
||||
k: int = 60
|
||||
) -> List[MergedCandidate]:
|
||||
def reciprocal_rank_fusion(result_lists: list[list[RetrievalResult]], k: int = 60) -> list[MergedCandidate]:
|
||||
"""
|
||||
Merge multiple ranked result lists using Reciprocal Rank Fusion.
|
||||
|
||||
@@ -73,20 +70,14 @@ def reciprocal_rank_fusion(
|
||||
sorted(rrf_scores.items(), key=lambda x: x[1], reverse=True), start=1
|
||||
):
|
||||
merged_candidate = MergedCandidate(
|
||||
retrieval=all_retrievals[doc_id],
|
||||
rrf_score=rrf_score,
|
||||
rrf_rank=rrf_rank,
|
||||
source_ranks=source_ranks[doc_id]
|
||||
retrieval=all_retrievals[doc_id], rrf_score=rrf_score, rrf_rank=rrf_rank, source_ranks=source_ranks[doc_id]
|
||||
)
|
||||
merged_results.append(merged_candidate)
|
||||
|
||||
return merged_results
|
||||
|
||||
|
||||
def normalize_scores_on_deltas(
|
||||
results: List[Dict[str, Any]],
|
||||
score_keys: List[str]
|
||||
) -> List[Dict[str, Any]]:
|
||||
def normalize_scores_on_deltas(results: list[dict[str, Any]], score_keys: list[str]) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Normalize scores based on deltas (min-max normalization within result set).
|
||||
|
||||
|
||||
@@ -6,13 +6,13 @@ allowing different algorithms (BFS spreading activation, PPR, etc.) to be
|
||||
swapped without changing the rest of the recall pipeline.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Optional
|
||||
from datetime import datetime
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from .types import RetrievalResult
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from .tags import TagsMatch, filter_results_by_tags
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -40,10 +40,13 @@ class GraphRetriever(ABC):
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
query_text: Optional[str] = None,
|
||||
semantic_seeds: Optional[List[RetrievalResult]] = None,
|
||||
temporal_seeds: Optional[List[RetrievalResult]] = None,
|
||||
) -> List[RetrievalResult]:
|
||||
query_text: str | None = None,
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # TypedAdjacency, optional pre-loaded graph
|
||||
tags: list[str] | None = None, # Visibility scope tags for filtering
|
||||
tags_match: TagsMatch = "any", # How to match tags: 'any' (OR) or 'all' (AND)
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve relevant facts via graph traversal.
|
||||
|
||||
@@ -56,9 +59,11 @@ class GraphRetriever(ABC):
|
||||
query_text: Original query text (optional, for some strategies)
|
||||
semantic_seeds: Pre-computed semantic entry points (from semantic retrieval)
|
||||
temporal_seeds: Pre-computed temporal entry points (from temporal retrieval)
|
||||
adjacency: Pre-loaded typed adjacency graph (optional, for MPFP)
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
List of RetrievalResult objects with activation scores set
|
||||
Tuple of (List of RetrievalResult with activation scores, optional timing info)
|
||||
"""
|
||||
pass
|
||||
|
||||
@@ -109,10 +114,13 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
query_text: Optional[str] = None,
|
||||
semantic_seeds: Optional[List[RetrievalResult]] = None,
|
||||
temporal_seeds: Optional[List[RetrievalResult]] = None,
|
||||
) -> List[RetrievalResult]:
|
||||
query_text: str | None = None,
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # Not used by BFS
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve facts using BFS spreading activation.
|
||||
|
||||
@@ -123,13 +131,14 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
4. Return visited nodes up to budget
|
||||
|
||||
Note: BFS finds its own entry points via embedding search.
|
||||
The semantic_seeds and temporal_seeds parameters are accepted
|
||||
The semantic_seeds, temporal_seeds, and adjacency parameters are accepted
|
||||
for interface compatibility but not used.
|
||||
"""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
return await self._retrieve_with_conn(
|
||||
conn, query_embedding_str, bank_id, fact_type, budget
|
||||
results = await self._retrieve_with_conn(
|
||||
conn, query_embedding_str, bank_id, fact_type, budget, tags=tags, tags_match=tags_match
|
||||
)
|
||||
return results, None
|
||||
|
||||
async def _retrieve_with_conn(
|
||||
self,
|
||||
@@ -138,37 +147,50 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
) -> List[RetrievalResult]:
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> list[RetrievalResult]:
|
||||
"""Internal implementation with connection."""
|
||||
from .tags import build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_embedding_str, bank_id, fact_type, self.entry_point_threshold, self.entry_point_limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
# Step 1: Find entry points
|
||||
entry_points = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM memory_units
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= $4
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str, bank_id, fact_type,
|
||||
self.entry_point_threshold, self.entry_point_limit
|
||||
*params,
|
||||
)
|
||||
|
||||
if not entry_points:
|
||||
logger.debug(
|
||||
f"[BFS] No entry points found for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
return []
|
||||
|
||||
logger.debug(
|
||||
f"[BFS] Found {len(entry_points)} entry points for fact_type={fact_type} "
|
||||
f"(tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
# Step 2: BFS spreading activation
|
||||
visited = set()
|
||||
results = []
|
||||
queue = [
|
||||
(RetrievalResult.from_db_row(dict(r)), r["similarity"])
|
||||
for r in entry_points
|
||||
]
|
||||
queue = [(RetrievalResult.from_db_row(dict(r)), r["similarity"]) for r in entry_points]
|
||||
budget_remaining = budget
|
||||
|
||||
while queue and budget_remaining > 0:
|
||||
@@ -192,20 +214,23 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
if batch_nodes and budget_remaining > 0:
|
||||
max_neighbors = len(batch_nodes) * 20
|
||||
neighbors = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.occurred_end,
|
||||
mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type,
|
||||
mu.document_id, mu.chunk_id,
|
||||
mu.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 memory_links ml
|
||||
JOIN memory_units mu ON ml.to_unit_id = mu.id
|
||||
FROM {fq_table("memory_links")} ml
|
||||
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
||||
WHERE ml.from_unit_id = ANY($1::uuid[])
|
||||
AND ml.weight >= $2
|
||||
AND mu.fact_type = $3
|
||||
ORDER BY ml.weight DESC
|
||||
LIMIT $4
|
||||
""",
|
||||
batch_nodes, self.min_activation, fact_type, max_neighbors
|
||||
batch_nodes,
|
||||
self.min_activation,
|
||||
fact_type,
|
||||
max_neighbors,
|
||||
)
|
||||
|
||||
for n in neighbors:
|
||||
@@ -232,4 +257,8 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
neighbor_result = RetrievalResult.from_db_row(dict(n))
|
||||
queue.append((neighbor_result, new_activation))
|
||||
|
||||
# Apply tags filtering (BFS may traverse into memories that don't match tags criteria)
|
||||
if tags:
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
return results
|
||||
|
||||
@@ -0,0 +1,256 @@
|
||||
"""
|
||||
Link Expansion graph retrieval.
|
||||
|
||||
A simple, fast graph retrieval that expands from seeds via:
|
||||
1. Entity links: Find facts sharing entities with seeds (filtered by entity frequency)
|
||||
2. Causal links: Find facts causally linked to seeds (top-k by weight)
|
||||
|
||||
Characteristics:
|
||||
- 2-3 DB queries (seed finding + parallel entity/causal expansion)
|
||||
- Sublinear: only touches connected facts via indexes
|
||||
- No iteration, no propagation, no normalization
|
||||
- Target: <100ms
|
||||
"""
|
||||
|
||||
import logging
|
||||
import time
|
||||
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from .graph_retrieval import GraphRetriever
|
||||
from .tags import TagsMatch, filter_results_by_tags
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def _find_semantic_seeds(
|
||||
conn,
|
||||
query_embedding_str: str,
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
limit: int = 20,
|
||||
threshold: float = 0.3,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> list[RetrievalResult]:
|
||||
"""Find semantic seeds via embedding search."""
|
||||
from .tags import build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= $4
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
|
||||
|
||||
|
||||
class LinkExpansionRetriever(GraphRetriever):
|
||||
"""
|
||||
Graph retrieval via direct link expansion from seeds.
|
||||
|
||||
Expands through entity co-occurrence and causal links in a single query.
|
||||
Fast and simple alternative to MPFP.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_entity_frequency: int = 500,
|
||||
causal_weight_threshold: float = 0.3,
|
||||
causal_limit_per_seed: int = 10,
|
||||
):
|
||||
"""
|
||||
Initialize link expansion retriever.
|
||||
|
||||
Args:
|
||||
max_entity_frequency: Skip entities appearing in more than this many facts
|
||||
causal_weight_threshold: Minimum weight for causal links
|
||||
causal_limit_per_seed: Max causal links to follow per seed
|
||||
"""
|
||||
self.max_entity_frequency = max_entity_frequency
|
||||
self.causal_weight_threshold = causal_weight_threshold
|
||||
self.causal_limit_per_seed = causal_limit_per_seed
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return "link_expansion"
|
||||
|
||||
async def retrieve(
|
||||
self,
|
||||
pool,
|
||||
query_embedding_str: str,
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
query_text: str | None = None,
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve facts by expanding links from seeds.
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
query_embedding_str: Query embedding (unused, kept for interface)
|
||||
bank_id: Memory bank ID
|
||||
fact_type: Fact type to filter
|
||||
budget: Maximum results to return
|
||||
query_text: Original query text (unused)
|
||||
semantic_seeds: Pre-computed semantic entry points
|
||||
temporal_seeds: Pre-computed temporal entry points
|
||||
adjacency: Unused, kept for interface compatibility
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
Tuple of (results, timings)
|
||||
"""
|
||||
start_time = time.time()
|
||||
timings = MPFPTimings(fact_type=fact_type)
|
||||
|
||||
# Use single connection for all queries to reduce pool pressure
|
||||
# (queries are fast ~50ms each, connection acquisition is the bottleneck)
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Find seeds if not provided
|
||||
if semantic_seeds:
|
||||
all_seeds = list(semantic_seeds)
|
||||
else:
|
||||
seeds_start = time.time()
|
||||
all_seeds = await _find_semantic_seeds(
|
||||
conn,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit=20,
|
||||
threshold=0.3,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
timings.seeds_time = time.time() - seeds_start
|
||||
logger.debug(
|
||||
f"[LinkExpansion] Found {len(all_seeds)} semantic seeds for fact_type={fact_type} "
|
||||
f"(tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
# Add temporal seeds if provided
|
||||
if temporal_seeds:
|
||||
all_seeds.extend(temporal_seeds)
|
||||
|
||||
if not all_seeds:
|
||||
logger.debug("[LinkExpansion] No seeds found, returning empty results")
|
||||
return [], timings
|
||||
|
||||
seed_ids = list({s.id for s in all_seeds})
|
||||
timings.pattern_count = len(seed_ids)
|
||||
|
||||
# Run entity and causal expansion sequentially on same connection
|
||||
query_start = time.time()
|
||||
|
||||
entity_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
COUNT(*)::float AS score
|
||||
FROM {fq_table("unit_entities")} seed_ue
|
||||
JOIN {fq_table("entities")} e ON seed_ue.entity_id = e.id
|
||||
JOIN {fq_table("unit_entities")} other_ue ON seed_ue.entity_id = other_ue.entity_id
|
||||
JOIN {fq_table("memory_units")} mu ON other_ue.unit_id = mu.id
|
||||
WHERE seed_ue.unit_id = ANY($1::uuid[])
|
||||
AND e.mention_count < $2
|
||||
AND mu.id != ALL($1::uuid[])
|
||||
AND mu.fact_type = $3
|
||||
GROUP BY mu.id
|
||||
ORDER BY score DESC
|
||||
LIMIT $4
|
||||
""",
|
||||
seed_ids,
|
||||
self.max_entity_frequency,
|
||||
fact_type,
|
||||
budget,
|
||||
)
|
||||
|
||||
causal_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT DISTINCT ON (mu.id)
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight + 1.0 AS score
|
||||
FROM {fq_table("memory_links")} ml
|
||||
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
||||
WHERE ml.from_unit_id = ANY($1::uuid[])
|
||||
AND ml.link_type IN ('causes', 'caused_by', 'enables', 'prevents')
|
||||
AND ml.weight >= $2
|
||||
AND mu.fact_type = $3
|
||||
ORDER BY mu.id, ml.weight DESC
|
||||
LIMIT $4
|
||||
""",
|
||||
seed_ids,
|
||||
self.causal_weight_threshold,
|
||||
fact_type,
|
||||
budget,
|
||||
)
|
||||
|
||||
timings.edge_load_time = time.time() - query_start
|
||||
timings.db_queries = 2
|
||||
timings.edge_count = len(entity_rows) + len(causal_rows)
|
||||
|
||||
# Merge results, taking max score per fact
|
||||
score_map: dict[str, float] = {}
|
||||
row_map: dict[str, dict] = {}
|
||||
|
||||
for row in entity_rows:
|
||||
fact_id = str(row["id"])
|
||||
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
|
||||
row_map[fact_id] = dict(row)
|
||||
|
||||
for row in causal_rows:
|
||||
fact_id = str(row["id"])
|
||||
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
|
||||
if fact_id not in row_map:
|
||||
row_map[fact_id] = dict(row)
|
||||
|
||||
# Sort by score and limit
|
||||
sorted_ids = sorted(score_map.keys(), key=lambda x: score_map[x], reverse=True)[:budget]
|
||||
rows = [row_map[fact_id] for fact_id in sorted_ids]
|
||||
|
||||
# Convert to results
|
||||
results = []
|
||||
for row in rows:
|
||||
result = RetrievalResult.from_db_row(dict(row))
|
||||
result.activation = row["score"]
|
||||
results.append(result)
|
||||
|
||||
# Apply tags filtering (graph expansion may reach untagged memories)
|
||||
if tags:
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
timings.result_count = len(results)
|
||||
timings.traverse = time.time() - start_time
|
||||
|
||||
logger.debug(
|
||||
f"LinkExpansion: {len(results)} results from {len(seed_ids)} seeds "
|
||||
f"in {timings.traverse * 1000:.1f}ms (query: {timings.edge_load_time * 1000:.1f}ms)"
|
||||
)
|
||||
|
||||
return results, timings
|
||||
@@ -9,6 +9,7 @@ propagation from Approximate PPR.
|
||||
|
||||
Key properties:
|
||||
- Sublinear in graph size (threshold pruning bounds active nodes)
|
||||
- Lazy edge loading: only loads edges for frontier nodes, not entire graph
|
||||
- Predefined patterns capture different retrieval intents
|
||||
- All patterns run in parallel, results fused via RRF
|
||||
- No LLM in the loop during traversal
|
||||
@@ -16,13 +17,14 @@ Key properties:
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from .types import RetrievalResult
|
||||
from .graph_retrieval import GraphRetriever
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from .graph_retrieval import GraphRetriever
|
||||
from .tags import TagsMatch
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -31,29 +33,43 @@ logger = logging.getLogger(__name__)
|
||||
# Data Classes
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class EdgeTarget:
|
||||
"""A neighbor node with its edge weight."""
|
||||
|
||||
node_id: str
|
||||
weight: float
|
||||
|
||||
|
||||
@dataclass
|
||||
class TypedAdjacency:
|
||||
"""Adjacency lists split by edge type."""
|
||||
# edge_type -> from_node_id -> list of (to_node_id, weight)
|
||||
graphs: Dict[str, Dict[str, List[EdgeTarget]]] = field(default_factory=dict)
|
||||
class EdgeCache:
|
||||
"""
|
||||
Cache for lazily-loaded edges.
|
||||
|
||||
def get_neighbors(self, edge_type: str, node_id: str) -> List[EdgeTarget]:
|
||||
Grows per-hop as edges are loaded for frontier nodes.
|
||||
Shared across patterns to avoid redundant loads.
|
||||
Loads ALL edge types at once to minimize DB queries.
|
||||
Thread-safe via asyncio lock to prevent redundant concurrent loads.
|
||||
"""
|
||||
|
||||
# edge_type -> from_node_id -> list of EdgeTarget
|
||||
graphs: dict[str, dict[str, list[EdgeTarget]]] = field(default_factory=dict)
|
||||
# Track which nodes have been fully loaded (all edge types)
|
||||
_fully_loaded: set[str] = field(default_factory=set)
|
||||
# Timing stats
|
||||
db_queries: int = 0
|
||||
edge_load_time: float = 0.0
|
||||
# Detailed hop timing for debugging
|
||||
hop_details: list[dict] = field(default_factory=list)
|
||||
# Lock to prevent redundant concurrent loads
|
||||
_lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
||||
|
||||
def get_neighbors(self, edge_type: str, node_id: str) -> list[EdgeTarget]:
|
||||
"""Get neighbors for a node via a specific edge type."""
|
||||
return self.graphs.get(edge_type, {}).get(node_id, [])
|
||||
|
||||
def get_normalized_neighbors(
|
||||
self,
|
||||
edge_type: str,
|
||||
node_id: str,
|
||||
top_k: int
|
||||
) -> List[EdgeTarget]:
|
||||
def get_normalized_neighbors(self, edge_type: str, node_id: str, top_k: int) -> list[EdgeTarget]:
|
||||
"""Get top-k neighbors with weights normalized to sum to 1."""
|
||||
neighbors = self.get_neighbors(edge_type, node_id)[:top_k]
|
||||
if not neighbors:
|
||||
@@ -63,123 +79,329 @@ class TypedAdjacency:
|
||||
if total == 0:
|
||||
return []
|
||||
|
||||
return [
|
||||
EdgeTarget(node_id=n.node_id, weight=n.weight / total)
|
||||
for n in neighbors
|
||||
]
|
||||
return [EdgeTarget(node_id=n.node_id, weight=n.weight / total) for n in neighbors]
|
||||
|
||||
def is_fully_loaded(self, node_id: str) -> bool:
|
||||
"""Check if all edges for this node have been loaded."""
|
||||
return node_id in self._fully_loaded
|
||||
|
||||
def get_uncached(self, node_ids: list[str]) -> list[str]:
|
||||
"""Get node IDs that haven't been fully loaded yet."""
|
||||
return [n for n in node_ids if not self.is_fully_loaded(n)]
|
||||
|
||||
def add_all_edges(self, edges_by_type: dict[str, dict[str, list[EdgeTarget]]], all_queried: list[str]):
|
||||
"""
|
||||
Add loaded edges to the cache (all edge types at once).
|
||||
|
||||
Args:
|
||||
edges_by_type: Dict mapping edge_type -> from_node_id -> list of EdgeTarget
|
||||
all_queried: All node IDs that were queried (marks them as fully loaded)
|
||||
"""
|
||||
for edge_type, edges in edges_by_type.items():
|
||||
if edge_type not in self.graphs:
|
||||
self.graphs[edge_type] = {}
|
||||
for node_id, neighbors in edges.items():
|
||||
self.graphs[edge_type][node_id] = neighbors
|
||||
|
||||
# Mark all queried nodes as fully loaded (even if they have no edges)
|
||||
self._fully_loaded.update(all_queried)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PatternResult:
|
||||
"""Result from a single pattern traversal."""
|
||||
pattern: List[str]
|
||||
scores: Dict[str, float] # node_id -> accumulated mass
|
||||
|
||||
pattern: list[str]
|
||||
scores: dict[str, float] # node_id -> accumulated mass
|
||||
|
||||
|
||||
@dataclass
|
||||
class MPFPConfig:
|
||||
"""Configuration for MPFP algorithm."""
|
||||
alpha: float = 0.15 # teleport/keep probability
|
||||
threshold: float = 1e-6 # mass pruning threshold (lower = explore more)
|
||||
top_k_neighbors: int = 20 # fan-out limit per node
|
||||
|
||||
alpha: float = 0.15 # teleport/keep probability
|
||||
threshold: float = 1e-6 # mass pruning threshold (lower = explore more)
|
||||
top_k_neighbors: int = 20 # fan-out limit per node
|
||||
|
||||
# Patterns from semantic seeds
|
||||
patterns_semantic: List[List[str]] = field(default_factory=lambda: [
|
||||
['semantic', 'semantic'], # topic expansion
|
||||
['entity', 'temporal'], # entity timeline
|
||||
['semantic', 'causes'], # reasoning chains (forward)
|
||||
['semantic', 'caused_by'], # reasoning chains (backward)
|
||||
['entity', 'semantic'], # entity context
|
||||
])
|
||||
patterns_semantic: list[list[str]] = field(
|
||||
default_factory=lambda: [
|
||||
["semantic", "semantic"], # topic expansion
|
||||
["entity", "temporal"], # entity timeline
|
||||
["semantic", "causes"], # reasoning chains (forward)
|
||||
["semantic", "caused_by"], # reasoning chains (backward)
|
||||
["entity", "semantic"], # entity context
|
||||
]
|
||||
)
|
||||
|
||||
# Patterns from temporal seeds
|
||||
patterns_temporal: List[List[str]] = field(default_factory=lambda: [
|
||||
['temporal', 'semantic'], # what was happening then
|
||||
['temporal', 'entity'], # who was involved then
|
||||
])
|
||||
patterns_temporal: list[list[str]] = field(
|
||||
default_factory=lambda: [
|
||||
["temporal", "semantic"], # what was happening then
|
||||
["temporal", "entity"], # who was involved then
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SeedNode:
|
||||
"""An entry point node with its initial score."""
|
||||
|
||||
node_id: str
|
||||
score: float # initial mass (e.g., similarity score)
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Core Algorithm
|
||||
# Lazy Edge Loading
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
def mpfp_traverse(
|
||||
seeds: List[SeedNode],
|
||||
pattern: List[str],
|
||||
adjacency: TypedAdjacency,
|
||||
config: MPFPConfig,
|
||||
) -> PatternResult:
|
||||
|
||||
async def load_all_edges_for_frontier(
|
||||
pool,
|
||||
node_ids: list[str],
|
||||
top_k_per_type: int = 20,
|
||||
) -> dict[str, dict[str, list[EdgeTarget]]]:
|
||||
"""
|
||||
Forward Push traversal following a meta-path pattern.
|
||||
Load top-k edges per (node, edge_type) for frontier nodes.
|
||||
|
||||
Uses a LATERAL join to efficiently fetch only the top-k edges per type,
|
||||
avoiding loading hundreds of entity edges when only 20 are needed.
|
||||
|
||||
Requires composite index: (from_unit_id, link_type, weight DESC)
|
||||
|
||||
Args:
|
||||
seeds: Entry point nodes with initial scores
|
||||
pattern: Sequence of edge types to follow
|
||||
adjacency: Typed adjacency structure
|
||||
config: Algorithm parameters
|
||||
pool: Database connection pool
|
||||
node_ids: Frontier node IDs to load edges for
|
||||
top_k_per_type: Max edges to load per (node, link_type) pair
|
||||
|
||||
Returns:
|
||||
PatternResult with accumulated scores per node
|
||||
Dict mapping edge_type -> from_node_id -> list of EdgeTarget
|
||||
"""
|
||||
if not node_ids:
|
||||
return {}
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Use LATERAL join to get top-k per (from_node, link_type)
|
||||
# This leverages the composite index for efficient early termination
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
WITH frontier(node_id) AS (SELECT unnest($1::uuid[]))
|
||||
SELECT f.node_id as from_unit_id, lt.link_type, edges.to_unit_id, edges.weight
|
||||
FROM frontier f
|
||||
CROSS JOIN (VALUES ('semantic'), ('temporal'), ('entity'), ('causes'), ('caused_by')) AS lt(link_type)
|
||||
CROSS JOIN LATERAL (
|
||||
SELECT ml.to_unit_id, ml.weight
|
||||
FROM {fq_table("memory_links")} ml
|
||||
WHERE ml.from_unit_id = f.node_id
|
||||
AND ml.link_type = lt.link_type
|
||||
AND ml.weight >= 0.1
|
||||
ORDER BY ml.weight DESC
|
||||
LIMIT $2
|
||||
) edges
|
||||
""",
|
||||
node_ids,
|
||||
top_k_per_type,
|
||||
)
|
||||
|
||||
# Group by edge_type -> from_node -> neighbors
|
||||
result: dict[str, dict[str, list[EdgeTarget]]] = defaultdict(lambda: defaultdict(list))
|
||||
for row in rows:
|
||||
edge_type = row["link_type"]
|
||||
from_id = str(row["from_unit_id"])
|
||||
to_id = str(row["to_unit_id"])
|
||||
weight = row["weight"]
|
||||
result[edge_type][from_id].append(EdgeTarget(node_id=to_id, weight=weight))
|
||||
|
||||
# Convert nested defaultdicts to regular dicts
|
||||
return {edge_type: dict(edges) for edge_type, edges in result.items()}
|
||||
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Core Algorithm (Async with Lazy Loading)
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class PatternState:
|
||||
"""State for a pattern traversal between hops."""
|
||||
|
||||
pattern: list[str]
|
||||
hop_index: int
|
||||
scores: dict[str, float]
|
||||
frontier: dict[str, float]
|
||||
|
||||
|
||||
def _init_pattern_state(seeds: list[SeedNode], pattern: list[str]) -> PatternState:
|
||||
"""Initialize pattern state from seeds."""
|
||||
if not seeds:
|
||||
return PatternState(pattern=pattern, hop_index=0, scores={}, frontier={})
|
||||
|
||||
total_seed_score = sum(s.score for s in seeds)
|
||||
if total_seed_score == 0:
|
||||
total_seed_score = len(seeds)
|
||||
|
||||
frontier = {s.node_id: s.score / total_seed_score for s in seeds}
|
||||
return PatternState(pattern=pattern, hop_index=0, scores={}, frontier=frontier)
|
||||
|
||||
|
||||
def _execute_hop(state: PatternState, cache: EdgeCache, config: MPFPConfig) -> set[str]:
|
||||
"""
|
||||
Execute ONE hop of traversal, return frontier nodes for next hop.
|
||||
|
||||
This is a pure function that uses cached edges (no DB access).
|
||||
Returns set of uncached nodes needed for next hop.
|
||||
"""
|
||||
if state.hop_index >= len(state.pattern):
|
||||
return set()
|
||||
|
||||
edge_type = state.pattern[state.hop_index]
|
||||
|
||||
# Collect active nodes above threshold
|
||||
active_nodes = [node_id for node_id, mass in state.frontier.items() if mass >= config.threshold]
|
||||
if not active_nodes:
|
||||
state.frontier = {}
|
||||
return set()
|
||||
|
||||
# Propagate mass using cached edges
|
||||
next_frontier: dict[str, float] = {}
|
||||
uncached_for_next: set[str] = set()
|
||||
|
||||
for node_id, mass in state.frontier.items():
|
||||
if mass < config.threshold:
|
||||
continue
|
||||
|
||||
# Keep α portion for this node
|
||||
state.scores[node_id] = state.scores.get(node_id, 0) + config.alpha * mass
|
||||
|
||||
# Push (1-α) to neighbors
|
||||
push_mass = (1 - config.alpha) * mass
|
||||
neighbors = cache.get_normalized_neighbors(edge_type, node_id, config.top_k_neighbors)
|
||||
|
||||
for neighbor in neighbors:
|
||||
next_frontier[neighbor.node_id] = next_frontier.get(neighbor.node_id, 0) + push_mass * neighbor.weight
|
||||
# Track if we'll need edges for this node in the next hop
|
||||
if not cache.is_fully_loaded(neighbor.node_id):
|
||||
uncached_for_next.add(neighbor.node_id)
|
||||
|
||||
state.frontier = next_frontier
|
||||
state.hop_index += 1
|
||||
|
||||
return uncached_for_next
|
||||
|
||||
|
||||
def _finalize_pattern(state: PatternState, config: MPFPConfig) -> PatternResult:
|
||||
"""Finalize pattern by adding remaining frontier mass to scores."""
|
||||
for node_id, mass in state.frontier.items():
|
||||
if mass >= config.threshold:
|
||||
state.scores[node_id] = state.scores.get(node_id, 0) + mass
|
||||
|
||||
return PatternResult(pattern=state.pattern, scores=state.scores)
|
||||
|
||||
|
||||
async def mpfp_traverse_hop_synchronized(
|
||||
pool,
|
||||
pattern_jobs: list[tuple[list[SeedNode], list[str]]],
|
||||
config: MPFPConfig,
|
||||
cache: EdgeCache,
|
||||
) -> list[PatternResult]:
|
||||
"""
|
||||
Execute ALL patterns with hop-synchronized edge loading.
|
||||
|
||||
Instead of running each pattern independently (causing multiple DB queries),
|
||||
this function:
|
||||
1. Runs hop 1 for ALL patterns (using pre-warmed seed edges)
|
||||
2. Collects ALL unique hop-2 frontier nodes across patterns
|
||||
3. Pre-warms hop-2 edges in ONE query
|
||||
4. Runs hop 2 for ALL patterns
|
||||
|
||||
This reduces DB queries from O(patterns * hops) to O(hops).
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
pattern_jobs: List of (seeds, pattern) tuples
|
||||
config: Algorithm parameters
|
||||
cache: Shared edge cache (should be pre-warmed with seed edges)
|
||||
|
||||
Returns:
|
||||
List of PatternResult for each pattern
|
||||
"""
|
||||
import time
|
||||
|
||||
# Initialize all pattern states
|
||||
states = [_init_pattern_state(seeds, pattern) for seeds, pattern in pattern_jobs]
|
||||
|
||||
# Determine max hops (all patterns should be same length, but be safe)
|
||||
max_hops = max((len(p) for _, p in pattern_jobs), default=0)
|
||||
|
||||
# Detailed timing for debugging
|
||||
hop_times: list[dict] = []
|
||||
|
||||
# Execute hop-by-hop across ALL patterns
|
||||
for hop in range(max_hops):
|
||||
hop_start = time.time()
|
||||
hop_timing = {"hop": hop, "patterns_executed": 0, "uncached_count": 0, "load_time": 0.0}
|
||||
|
||||
# Execute this hop for all patterns, collect uncached nodes for next hop
|
||||
all_uncached: set[str] = set()
|
||||
exec_start = time.time()
|
||||
for state in states:
|
||||
if state.hop_index < len(state.pattern):
|
||||
uncached = _execute_hop(state, cache, config)
|
||||
all_uncached.update(uncached)
|
||||
hop_timing["patterns_executed"] += 1
|
||||
hop_timing["exec_time"] = time.time() - exec_start
|
||||
|
||||
# Pre-warm edges for ALL uncached nodes before next hop
|
||||
hop_timing["uncached_count"] = len(all_uncached)
|
||||
if all_uncached:
|
||||
uncached_list = list(all_uncached - cache._fully_loaded)
|
||||
hop_timing["uncached_after_filter"] = len(uncached_list)
|
||||
if uncached_list:
|
||||
load_start = time.time()
|
||||
edges_by_type = await load_all_edges_for_frontier(pool, uncached_list, config.top_k_neighbors)
|
||||
hop_timing["load_time"] = time.time() - load_start
|
||||
cache.edge_load_time += hop_timing["load_time"]
|
||||
cache.db_queries += 1
|
||||
cache.add_all_edges(edges_by_type, uncached_list)
|
||||
hop_timing["edges_loaded"] = sum(
|
||||
len(neighbors) for edges in edges_by_type.values() for neighbors in edges.values()
|
||||
)
|
||||
|
||||
hop_timing["total_time"] = time.time() - hop_start
|
||||
hop_times.append(hop_timing)
|
||||
|
||||
# Store hop timing details in cache for logging
|
||||
cache.hop_details = hop_times
|
||||
|
||||
# Finalize all patterns
|
||||
return [_finalize_pattern(state, config) for state in states]
|
||||
|
||||
|
||||
async def mpfp_traverse_async(
|
||||
pool,
|
||||
seeds: list[SeedNode],
|
||||
pattern: list[str],
|
||||
config: MPFPConfig,
|
||||
cache: EdgeCache,
|
||||
) -> PatternResult:
|
||||
"""
|
||||
Async Forward Push traversal with lazy edge loading.
|
||||
|
||||
NOTE: For better performance with multiple patterns, use mpfp_traverse_hop_synchronized().
|
||||
This function is kept for single-pattern use cases.
|
||||
"""
|
||||
if not seeds:
|
||||
return PatternResult(pattern=pattern, scores={})
|
||||
|
||||
scores: Dict[str, float] = {}
|
||||
|
||||
# Initialize frontier with seed masses (normalized)
|
||||
total_seed_score = sum(s.score for s in seeds)
|
||||
if total_seed_score == 0:
|
||||
total_seed_score = len(seeds) # fallback to uniform
|
||||
|
||||
frontier: Dict[str, float] = {
|
||||
s.node_id: s.score / total_seed_score for s in seeds
|
||||
}
|
||||
|
||||
# Follow pattern hop by hop
|
||||
for edge_type in pattern:
|
||||
next_frontier: Dict[str, float] = {}
|
||||
|
||||
for node_id, mass in frontier.items():
|
||||
if mass < config.threshold:
|
||||
continue
|
||||
|
||||
# Keep α portion for this node
|
||||
scores[node_id] = scores.get(node_id, 0) + config.alpha * mass
|
||||
|
||||
# Push (1-α) to neighbors
|
||||
push_mass = (1 - config.alpha) * mass
|
||||
neighbors = adjacency.get_normalized_neighbors(
|
||||
edge_type, node_id, config.top_k_neighbors
|
||||
)
|
||||
|
||||
for neighbor in neighbors:
|
||||
next_frontier[neighbor.node_id] = (
|
||||
next_frontier.get(neighbor.node_id, 0) +
|
||||
push_mass * neighbor.weight
|
||||
)
|
||||
|
||||
frontier = next_frontier
|
||||
|
||||
# Final frontier nodes get their remaining mass
|
||||
for node_id, mass in frontier.items():
|
||||
if mass >= config.threshold:
|
||||
scores[node_id] = scores.get(node_id, 0) + mass
|
||||
|
||||
return PatternResult(pattern=pattern, scores=scores)
|
||||
results = await mpfp_traverse_hop_synchronized(pool, [(seeds, pattern)], config, cache)
|
||||
return results[0] if results else PatternResult(pattern=pattern, scores={})
|
||||
|
||||
|
||||
def rrf_fusion(
|
||||
results: List[PatternResult],
|
||||
results: list[PatternResult],
|
||||
k: int = 60,
|
||||
top_k: int = 50,
|
||||
) -> List[Tuple[str, float]]:
|
||||
) -> list[tuple[str, float]]:
|
||||
"""
|
||||
Reciprocal Rank Fusion to combine pattern results.
|
||||
|
||||
@@ -191,28 +413,20 @@ def rrf_fusion(
|
||||
Returns:
|
||||
List of (node_id, fused_score) tuples, sorted by score descending
|
||||
"""
|
||||
fused: Dict[str, float] = {}
|
||||
fused: dict[str, float] = {}
|
||||
|
||||
for result in results:
|
||||
if not result.scores:
|
||||
continue
|
||||
|
||||
# Rank nodes by their score in this pattern
|
||||
ranked = sorted(
|
||||
result.scores.keys(),
|
||||
key=lambda n: result.scores[n],
|
||||
reverse=True
|
||||
)
|
||||
ranked = sorted(result.scores.keys(), key=lambda n: result.scores[n], reverse=True)
|
||||
|
||||
for rank, node_id in enumerate(ranked):
|
||||
fused[node_id] = fused.get(node_id, 0) + 1.0 / (k + rank + 1)
|
||||
|
||||
# Sort by fused score and return top-k
|
||||
sorted_results = sorted(
|
||||
fused.items(),
|
||||
key=lambda x: x[1],
|
||||
reverse=True
|
||||
)
|
||||
sorted_results = sorted(fused.items(), key=lambda x: x[1], reverse=True)
|
||||
|
||||
return sorted_results[:top_k]
|
||||
|
||||
@@ -221,62 +435,27 @@ def rrf_fusion(
|
||||
# Database Loading
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
async def load_typed_adjacency(pool, bank_id: str) -> TypedAdjacency:
|
||||
"""
|
||||
Load all edges for a bank, split by edge type.
|
||||
|
||||
Single query, then organize in-memory for fast traversal.
|
||||
"""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT ml.from_unit_id, ml.to_unit_id, ml.link_type, ml.weight
|
||||
FROM memory_links ml
|
||||
JOIN memory_units mu ON ml.from_unit_id = mu.id
|
||||
WHERE mu.bank_id = $1
|
||||
AND ml.weight >= 0.1
|
||||
ORDER BY ml.from_unit_id, ml.weight DESC
|
||||
""",
|
||||
bank_id
|
||||
)
|
||||
|
||||
graphs: Dict[str, Dict[str, List[EdgeTarget]]] = defaultdict(
|
||||
lambda: defaultdict(list)
|
||||
)
|
||||
|
||||
for row in rows:
|
||||
from_id = str(row['from_unit_id'])
|
||||
to_id = str(row['to_unit_id'])
|
||||
link_type = row['link_type']
|
||||
weight = row['weight']
|
||||
|
||||
graphs[link_type][from_id].append(
|
||||
EdgeTarget(node_id=to_id, weight=weight)
|
||||
)
|
||||
|
||||
return TypedAdjacency(graphs=dict(graphs))
|
||||
|
||||
|
||||
async def fetch_memory_units_by_ids(
|
||||
pool,
|
||||
node_ids: List[str],
|
||||
node_ids: list[str],
|
||||
fact_type: str,
|
||||
) -> List[RetrievalResult]:
|
||||
) -> list[RetrievalResult]:
|
||||
"""Fetch full memory unit details for a list of node IDs."""
|
||||
if not node_ids:
|
||||
return []
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id
|
||||
FROM memory_units
|
||||
mentioned_at, embedding, fact_type, document_id, chunk_id, tags
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = ANY($1::uuid[])
|
||||
AND fact_type = $2
|
||||
""",
|
||||
node_ids,
|
||||
fact_type
|
||||
fact_type,
|
||||
)
|
||||
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
|
||||
@@ -286,23 +465,29 @@ async def fetch_memory_units_by_ids(
|
||||
# Graph Retriever Implementation
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MPFPGraphRetriever(GraphRetriever):
|
||||
"""
|
||||
Graph retrieval using Meta-Path Forward Push.
|
||||
Graph retrieval using Meta-Path Forward Push with lazy edge loading.
|
||||
|
||||
Runs predefined patterns in parallel from semantic and temporal seeds,
|
||||
then fuses results via RRF.
|
||||
loading edges on-demand per hop instead of loading entire graph upfront.
|
||||
"""
|
||||
|
||||
def __init__(self, config: Optional[MPFPConfig] = None):
|
||||
def __init__(self, config: MPFPConfig | None = None):
|
||||
"""
|
||||
Initialize MPFP retriever.
|
||||
|
||||
Args:
|
||||
config: Algorithm configuration (uses defaults if None)
|
||||
"""
|
||||
self.config = config or MPFPConfig()
|
||||
self._adjacency_cache: Dict[str, TypedAdjacency] = {}
|
||||
if config is None:
|
||||
# Read top_k_neighbors from global config
|
||||
from ...config import get_config
|
||||
|
||||
global_config = get_config()
|
||||
config = MPFPConfig(top_k_neighbors=global_config.mpfp_top_k_neighbors)
|
||||
self.config = config
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
@@ -315,12 +500,15 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
query_text: Optional[str] = None,
|
||||
semantic_seeds: Optional[List[RetrievalResult]] = None,
|
||||
temporal_seeds: Optional[List[RetrievalResult]] = None,
|
||||
) -> List[RetrievalResult]:
|
||||
query_text: str | None = None,
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # Ignored - kept for interface compatibility
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve facts using MPFP algorithm.
|
||||
Retrieve facts using MPFP algorithm with lazy edge loading.
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
@@ -331,69 +519,104 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
query_text: Original query text (optional)
|
||||
semantic_seeds: Pre-computed semantic entry points
|
||||
temporal_seeds: Pre-computed temporal entry points
|
||||
adjacency: Ignored (kept for interface compatibility)
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
List of RetrievalResult with activation scores
|
||||
Tuple of (List of RetrievalResult with activation scores, MPFPTimings)
|
||||
"""
|
||||
# Load typed adjacency (could cache per bank_id with TTL)
|
||||
adjacency = await load_typed_adjacency(pool, bank_id)
|
||||
import time
|
||||
|
||||
timings = MPFPTimings(fact_type=fact_type)
|
||||
|
||||
# Convert seeds to SeedNode format
|
||||
semantic_seed_nodes = self._convert_seeds(semantic_seeds, 'similarity')
|
||||
temporal_seed_nodes = self._convert_seeds(temporal_seeds, 'temporal_score')
|
||||
semantic_seed_nodes = self._convert_seeds(semantic_seeds, "similarity")
|
||||
temporal_seed_nodes = self._convert_seeds(temporal_seeds, "temporal_score")
|
||||
|
||||
# If no semantic seeds provided, fall back to finding our own
|
||||
if not semantic_seed_nodes:
|
||||
seeds_start = time.time()
|
||||
semantic_seed_nodes = await self._find_semantic_seeds(
|
||||
pool, query_embedding_str, bank_id, fact_type
|
||||
pool, query_embedding_str, bank_id, fact_type, tags=tags, tags_match=tags_match
|
||||
)
|
||||
timings.seeds_time = time.time() - seeds_start
|
||||
logger.debug(
|
||||
f"[MPFP] Found {len(semantic_seed_nodes)} semantic seeds for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
# Run all patterns in parallel
|
||||
tasks = []
|
||||
# Collect all pattern jobs
|
||||
pattern_jobs = []
|
||||
|
||||
# Patterns from semantic seeds
|
||||
for pattern in self.config.patterns_semantic:
|
||||
if semantic_seed_nodes:
|
||||
tasks.append(
|
||||
asyncio.to_thread(
|
||||
mpfp_traverse,
|
||||
semantic_seed_nodes,
|
||||
pattern,
|
||||
adjacency,
|
||||
self.config,
|
||||
)
|
||||
)
|
||||
pattern_jobs.append((semantic_seed_nodes, pattern))
|
||||
|
||||
# Patterns from temporal seeds
|
||||
for pattern in self.config.patterns_temporal:
|
||||
if temporal_seed_nodes:
|
||||
tasks.append(
|
||||
asyncio.to_thread(
|
||||
mpfp_traverse,
|
||||
temporal_seed_nodes,
|
||||
pattern,
|
||||
adjacency,
|
||||
self.config,
|
||||
)
|
||||
)
|
||||
pattern_jobs.append((temporal_seed_nodes, pattern))
|
||||
|
||||
if not tasks:
|
||||
return []
|
||||
if not pattern_jobs:
|
||||
logger.debug(
|
||||
f"[MPFP] No pattern jobs (semantic_seeds={len(semantic_seed_nodes)}, temporal_seeds={len(temporal_seed_nodes)})"
|
||||
)
|
||||
return [], timings
|
||||
|
||||
# Gather pattern results
|
||||
pattern_results = await asyncio.gather(*tasks)
|
||||
timings.pattern_count = len(pattern_jobs)
|
||||
|
||||
# Shared edge cache across all patterns
|
||||
cache = EdgeCache()
|
||||
|
||||
# Pre-warm cache with ALL seed node edges BEFORE running patterns
|
||||
# This prevents redundant DB queries at hop 1
|
||||
all_seed_ids = list({s.node_id for seeds, _ in pattern_jobs for s in seeds})
|
||||
if all_seed_ids:
|
||||
import time as time_module
|
||||
|
||||
prewarm_start = time_module.time()
|
||||
edges_by_type = await load_all_edges_for_frontier(pool, all_seed_ids, self.config.top_k_neighbors)
|
||||
cache.edge_load_time += time_module.time() - prewarm_start
|
||||
cache.db_queries += 1
|
||||
cache.add_all_edges(edges_by_type, all_seed_ids)
|
||||
|
||||
# Run all patterns with HOP-SYNCHRONIZED edge loading
|
||||
# This batches hop-2 edge loads across ALL patterns into ONE query
|
||||
# Reduces DB queries from O(patterns * hops) to O(hops)
|
||||
step_start = time.time()
|
||||
pattern_results = await mpfp_traverse_hop_synchronized(pool, pattern_jobs, self.config, cache)
|
||||
timings.traverse = time.time() - step_start
|
||||
|
||||
# Record edge loading stats from cache
|
||||
timings.edge_count = sum(len(neighbors) for g in cache.graphs.values() for neighbors in g.values())
|
||||
timings.db_queries = cache.db_queries
|
||||
timings.edge_load_time = cache.edge_load_time
|
||||
timings.hop_details = cache.hop_details
|
||||
|
||||
# Fuse results
|
||||
step_start = time.time()
|
||||
fused = rrf_fusion(pattern_results, top_k=budget)
|
||||
timings.fusion = time.time() - step_start
|
||||
|
||||
if not fused:
|
||||
return []
|
||||
logger.debug(f"[MPFP] No fused results after RRF fusion (pattern_count={len(pattern_results)})")
|
||||
return [], timings
|
||||
|
||||
# Get top result IDs (don't exclude seeds - they may be highly relevant)
|
||||
# Get top result IDs
|
||||
result_ids = [node_id for node_id, score in fused][:budget]
|
||||
|
||||
# Fetch full details
|
||||
step_start = time.time()
|
||||
results = await fetch_memory_units_by_ids(pool, result_ids, fact_type)
|
||||
timings.fetch = time.time() - step_start
|
||||
|
||||
# Filter results by tags (graph traversal may have picked up unfiltered memories)
|
||||
if tags:
|
||||
from .tags import filter_results_by_tags
|
||||
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
timings.result_count = len(results)
|
||||
|
||||
# Add activation scores from fusion
|
||||
score_map = {node_id: score for node_id, score in fused}
|
||||
@@ -403,13 +626,13 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
# Sort by activation
|
||||
results.sort(key=lambda r: r.activation or 0, reverse=True)
|
||||
|
||||
return results
|
||||
return results, timings
|
||||
|
||||
def _convert_seeds(
|
||||
self,
|
||||
seeds: Optional[List[RetrievalResult]],
|
||||
seeds: list[RetrievalResult] | None,
|
||||
score_attr: str,
|
||||
) -> List[SeedNode]:
|
||||
) -> list[SeedNode]:
|
||||
"""Convert RetrievalResult seeds to SeedNode format."""
|
||||
if not seeds:
|
||||
return []
|
||||
@@ -431,24 +654,31 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
fact_type: str,
|
||||
limit: int = 20,
|
||||
threshold: float = 0.3,
|
||||
) -> List[SeedNode]:
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> list[SeedNode]:
|
||||
"""Fallback: find semantic seeds via embedding search."""
|
||||
from .tags import build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
f"""
|
||||
SELECT id, 1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM memory_units
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= $4
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str, bank_id, fact_type, threshold, limit
|
||||
*params,
|
||||
)
|
||||
|
||||
return [
|
||||
SeedNode(node_id=str(r['id']), score=r['similarity'])
|
||||
for r in rows
|
||||
]
|
||||
return [SeedNode(node_id=str(r["id"]), score=r["similarity"]) for r in rows]
|
||||
|
||||
@@ -1,132 +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 typing import List, Dict, Any
|
||||
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 []
|
||||
@@ -2,7 +2,6 @@
|
||||
Cross-encoder neural reranking for search results.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from .types import MergedCandidate, ScoredResult
|
||||
|
||||
|
||||
@@ -24,14 +23,28 @@ class CrossEncoderReranker:
|
||||
"""
|
||||
if cross_encoder is None:
|
||||
from hindsight_api.engine.cross_encoder import create_cross_encoder_from_env
|
||||
|
||||
cross_encoder = create_cross_encoder_from_env()
|
||||
self.cross_encoder = cross_encoder
|
||||
self._initialized = False
|
||||
|
||||
def rerank(
|
||||
self,
|
||||
query: str,
|
||||
candidates: List[MergedCandidate]
|
||||
) -> List[ScoredResult]:
|
||||
async def ensure_initialized(self):
|
||||
"""Ensure the cross-encoder model is initialized (for lazy initialization)."""
|
||||
if self._initialized:
|
||||
return
|
||||
|
||||
import asyncio
|
||||
|
||||
cross_encoder = self.cross_encoder
|
||||
# For local providers, run in thread pool to avoid blocking event loop
|
||||
if cross_encoder.provider_name == "local":
|
||||
loop = asyncio.get_event_loop()
|
||||
await loop.run_in_executor(None, lambda: asyncio.run(cross_encoder.initialize()))
|
||||
else:
|
||||
await cross_encoder.initialize()
|
||||
self._initialized = True
|
||||
|
||||
async def rerank(self, query: str, candidates: list[MergedCandidate]) -> list[ScoredResult]:
|
||||
"""
|
||||
Rerank candidates using cross-encoder scores.
|
||||
|
||||
@@ -72,11 +85,12 @@ class CrossEncoderReranker:
|
||||
pairs.append([query, doc_text])
|
||||
|
||||
# Get cross-encoder scores
|
||||
scores = self.cross_encoder.predict(pairs)
|
||||
scores = await self.cross_encoder.predict(pairs)
|
||||
|
||||
# Normalize scores using sigmoid to [0, 1] range
|
||||
# Cross-encoder returns logits which can be negative
|
||||
import numpy as np
|
||||
|
||||
def sigmoid(x):
|
||||
return 1 / (1 + np.exp(-x))
|
||||
|
||||
@@ -89,7 +103,7 @@ class CrossEncoderReranker:
|
||||
candidate=candidate,
|
||||
cross_encoder_score=float(raw_score),
|
||||
cross_encoder_score_normalized=float(norm_score),
|
||||
weight=float(norm_score) # Initial weight is just cross-encoder score
|
||||
weight=float(norm_score), # Initial weight is just cross-encoder score
|
||||
)
|
||||
scored_results.append(scored_result)
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -4,11 +4,11 @@ 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
|
||||
from typing import List
|
||||
|
||||
|
||||
def cosine_similarity(vec1: List[float], vec2: List[float]) -> float:
|
||||
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
|
||||
"""
|
||||
Calculate cosine similarity between two vectors.
|
||||
|
||||
@@ -58,36 +58,13 @@ def calculate_recency_weight(days_since: float, half_life_days: float = 365.0) -
|
||||
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.
|
||||
@@ -116,11 +93,7 @@ def calculate_temporal_anchor(occurred_start: datetime, occurred_end: datetime)
|
||||
return midpoint
|
||||
|
||||
|
||||
def calculate_temporal_proximity(
|
||||
anchor_a: datetime,
|
||||
anchor_b: datetime,
|
||||
half_life_days: float = 30.0
|
||||
) -> float:
|
||||
def calculate_temporal_proximity(anchor_a: datetime, anchor_b: datetime, half_life_days: float = 30.0) -> float:
|
||||
"""
|
||||
Calculate temporal proximity between two temporal anchors.
|
||||
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
"""
|
||||
Tags filtering utilities for retrieval.
|
||||
|
||||
Provides SQL building functions for filtering memories by tags.
|
||||
Supports four matching modes via TagsMatch enum:
|
||||
- "any": OR matching, includes untagged memories (default, backward compatible)
|
||||
- "all": AND matching, includes untagged memories
|
||||
- "any_strict": OR matching, excludes untagged memories
|
||||
- "all_strict": AND matching, excludes untagged memories
|
||||
|
||||
OR matching (any/any_strict): Memory matches if ANY of its tags overlap with request tags
|
||||
AND matching (all/all_strict): Memory matches if ALL request tags are present in its tags
|
||||
"""
|
||||
|
||||
from typing import Literal
|
||||
|
||||
TagsMatch = Literal["any", "all", "any_strict", "all_strict"]
|
||||
|
||||
|
||||
def _parse_tags_match(match: TagsMatch) -> tuple[str, bool]:
|
||||
"""
|
||||
Parse TagsMatch into operator and include_untagged flag.
|
||||
|
||||
Returns:
|
||||
Tuple of (operator, include_untagged)
|
||||
- operator: "&&" for any/any_strict, "@>" for all/all_strict
|
||||
- include_untagged: True for any/all, False for any_strict/all_strict
|
||||
"""
|
||||
if match == "any":
|
||||
return "&&", True
|
||||
elif match == "all":
|
||||
return "@>", True
|
||||
elif match == "any_strict":
|
||||
return "&&", False
|
||||
elif match == "all_strict":
|
||||
return "@>", False
|
||||
else:
|
||||
# Default to "any" behavior
|
||||
return "&&", True
|
||||
|
||||
|
||||
def build_tags_where_clause(
|
||||
tags: list[str] | None,
|
||||
param_offset: int = 1,
|
||||
table_alias: str = "",
|
||||
match: TagsMatch = "any",
|
||||
) -> tuple[str, list, int]:
|
||||
"""
|
||||
Build a SQL WHERE clause for filtering by tags.
|
||||
|
||||
Supports four matching modes:
|
||||
- "any" (default): OR matching, includes untagged memories
|
||||
- "all": AND matching, includes untagged memories
|
||||
- "any_strict": OR matching, excludes untagged memories
|
||||
- "all_strict": AND matching, excludes untagged memories
|
||||
|
||||
Args:
|
||||
tags: List of tags to filter by. If None or empty, returns empty clause (no filtering).
|
||||
param_offset: Starting parameter number for SQL placeholders (default 1).
|
||||
table_alias: Optional table alias prefix (e.g., "mu." for "memory_units mu").
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
Tuple of (sql_clause, params, next_param_offset):
|
||||
- sql_clause: SQL WHERE clause string
|
||||
- params: List of parameter values to bind
|
||||
- next_param_offset: Next available parameter number
|
||||
|
||||
Example:
|
||||
>>> clause, params, next_offset = build_tags_where_clause(['user_a'], 3, 'mu.', 'any_strict')
|
||||
>>> print(clause) # "AND mu.tags IS NOT NULL AND mu.tags != '{}' AND mu.tags && $3"
|
||||
"""
|
||||
if not tags:
|
||||
return "", [], param_offset
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
operator, include_untagged = _parse_tags_match(match)
|
||||
|
||||
if include_untagged:
|
||||
# Include untagged memories (NULL or empty array) OR matching tags
|
||||
clause = f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_offset})"
|
||||
else:
|
||||
# Strict: only memories with matching tags (exclude NULL and empty)
|
||||
clause = f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_offset}"
|
||||
|
||||
return clause, [tags], param_offset + 1
|
||||
|
||||
|
||||
def build_tags_where_clause_simple(
|
||||
tags: list[str] | None,
|
||||
param_num: int,
|
||||
table_alias: str = "",
|
||||
match: TagsMatch = "any",
|
||||
) -> str:
|
||||
"""
|
||||
Build a simple SQL WHERE clause for tags filtering.
|
||||
|
||||
This is a convenience version that returns just the clause string,
|
||||
assuming the caller will add the tags array to their params list.
|
||||
|
||||
Args:
|
||||
tags: List of tags to filter by. If None or empty, returns empty string.
|
||||
param_num: Parameter number to use in the clause.
|
||||
table_alias: Optional table alias prefix.
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
SQL clause string or empty string.
|
||||
"""
|
||||
if not tags:
|
||||
return ""
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
operator, include_untagged = _parse_tags_match(match)
|
||||
|
||||
if include_untagged:
|
||||
# Include untagged memories (NULL or empty array) OR matching tags
|
||||
return f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_num})"
|
||||
else:
|
||||
# Strict: only memories with matching tags (exclude NULL and empty)
|
||||
return f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_num}"
|
||||
|
||||
|
||||
def filter_results_by_tags(
|
||||
results: list,
|
||||
tags: list[str] | None,
|
||||
match: TagsMatch = "any",
|
||||
) -> list:
|
||||
"""
|
||||
Filter retrieval results by tags in Python (for post-processing).
|
||||
|
||||
Used when SQL filtering isn't possible (e.g., graph traversal results).
|
||||
|
||||
Args:
|
||||
results: List of RetrievalResult objects with a 'tags' attribute.
|
||||
tags: List of tags to filter by. If None or empty, returns all results.
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
Filtered list of results.
|
||||
"""
|
||||
if not tags:
|
||||
return results
|
||||
|
||||
_, include_untagged = _parse_tags_match(match)
|
||||
is_any_match = match in ("any", "any_strict")
|
||||
|
||||
tags_set = set(tags)
|
||||
filtered = []
|
||||
|
||||
for result in results:
|
||||
result_tags = getattr(result, "tags", None)
|
||||
|
||||
# Check if untagged
|
||||
is_untagged = result_tags is None or len(result_tags) == 0
|
||||
|
||||
if is_untagged:
|
||||
if include_untagged:
|
||||
filtered.append(result)
|
||||
# else: skip untagged
|
||||
else:
|
||||
result_tags_set = set(result_tags)
|
||||
if is_any_match:
|
||||
# Any overlap
|
||||
if result_tags_set & tags_set:
|
||||
filtered.append(result)
|
||||
else:
|
||||
# All tags must be present
|
||||
if tags_set <= result_tags_set:
|
||||
filtered.append(result)
|
||||
|
||||
return filtered
|
||||
@@ -4,16 +4,16 @@ Temporal extraction for time-aware search queries.
|
||||
Handles natural language temporal expressions using transformer-based query analysis.
|
||||
"""
|
||||
|
||||
from typing import Optional, Tuple
|
||||
from datetime import datetime
|
||||
import logging
|
||||
from hindsight_api.engine.query_analyzer import QueryAnalyzer, DateparserQueryAnalyzer
|
||||
from datetime import datetime
|
||||
|
||||
from hindsight_api.engine.query_analyzer import DateparserQueryAnalyzer, QueryAnalyzer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Global default analyzer instance
|
||||
# Can be overridden by passing a custom analyzer to extract_temporal_constraint
|
||||
_default_analyzer: Optional[QueryAnalyzer] = None
|
||||
_default_analyzer: QueryAnalyzer | None = None
|
||||
|
||||
|
||||
def get_default_analyzer() -> QueryAnalyzer:
|
||||
@@ -33,9 +33,9 @@ def get_default_analyzer() -> QueryAnalyzer:
|
||||
|
||||
def extract_temporal_constraint(
|
||||
query: str,
|
||||
reference_date: Optional[datetime] = None,
|
||||
analyzer: Optional[QueryAnalyzer] = None,
|
||||
) -> Optional[Tuple[datetime, datetime]]:
|
||||
reference_date: datetime | None = None,
|
||||
analyzer: QueryAnalyzer | None = None,
|
||||
) -> tuple[datetime, datetime] | None:
|
||||
"""
|
||||
Extract temporal constraint from query.
|
||||
|
||||
@@ -55,10 +55,7 @@ def extract_temporal_constraint(
|
||||
analysis = analyzer.analyze(query, reference_date)
|
||||
|
||||
if analysis.temporal_constraint:
|
||||
result = (
|
||||
analysis.temporal_constraint.start_date,
|
||||
analysis.temporal_constraint.end_date
|
||||
)
|
||||
result = (analysis.temporal_constraint.start_date, analysis.temporal_constraint.end_date)
|
||||
return result
|
||||
|
||||
return None
|
||||
|
||||
@@ -2,41 +2,17 @@
|
||||
Think operation utilities for formulating answers based on agent and world facts.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import re
|
||||
from datetime import datetime, timezone
|
||||
from typing import Dict, List, Any
|
||||
from pydantic import BaseModel, Field
|
||||
from datetime import datetime
|
||||
|
||||
from ..response_models import ReflectResult, MemoryFact, DispositionTraits
|
||||
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"
|
||||
}
|
||||
levels = {1: "very low", 2: "low", 3: "moderate", 4: "high", 5: "very high"}
|
||||
return levels.get(value, "moderate")
|
||||
|
||||
|
||||
@@ -47,7 +23,7 @@ def build_disposition_description(disposition: DispositionTraits) -> str:
|
||||
2: "You tend to trust information but may question obvious inconsistencies.",
|
||||
3: "You have a balanced approach to information, neither too trusting nor too skeptical.",
|
||||
4: "You are somewhat skeptical and often question the reliability of information.",
|
||||
5: "You are highly skeptical and critically examine all information for accuracy and hidden motives."
|
||||
5: "You are highly skeptical and critically examine all information for accuracy and hidden motives.",
|
||||
}
|
||||
|
||||
literalism_desc = {
|
||||
@@ -55,7 +31,7 @@ def build_disposition_description(disposition: DispositionTraits) -> str:
|
||||
2: "You tend to consider context and implied meaning alongside literal statements.",
|
||||
3: "You balance literal interpretation with contextual understanding.",
|
||||
4: "You prefer to interpret information more literally and precisely.",
|
||||
5: "You interpret information very literally and focus on exact wording and commitments."
|
||||
5: "You interpret information very literally and focus on exact wording and commitments.",
|
||||
}
|
||||
|
||||
empathy_desc = {
|
||||
@@ -63,7 +39,7 @@ def build_disposition_description(disposition: DispositionTraits) -> str:
|
||||
2: "You consider facts first but acknowledge emotional factors exist.",
|
||||
3: "You balance factual analysis with emotional understanding.",
|
||||
4: "You give significant weight to emotional context and human factors.",
|
||||
5: "You strongly consider the emotional state and circumstances of others when forming memories."
|
||||
5: "You strongly consider the emotional state and circumstances of others when forming memories.",
|
||||
}
|
||||
|
||||
return f"""Your disposition traits:
|
||||
@@ -72,7 +48,7 @@ def build_disposition_description(disposition: DispositionTraits) -> str:
|
||||
- Empathy ({describe_trait_level(disposition.empathy)}): {empathy_desc.get(disposition.empathy, empathy_desc[3])}"""
|
||||
|
||||
|
||||
def format_facts_for_prompt(facts: List[MemoryFact]) -> str:
|
||||
def format_facts_for_prompt(facts: list[MemoryFact]) -> str:
|
||||
"""Format facts as JSON for LLM prompt."""
|
||||
import json
|
||||
|
||||
@@ -80,9 +56,7 @@ def format_facts_for_prompt(facts: List[MemoryFact]) -> str:
|
||||
return "[]"
|
||||
formatted = []
|
||||
for fact in facts:
|
||||
fact_obj = {
|
||||
"text": fact.text
|
||||
}
|
||||
fact_obj = {"text": fact.text}
|
||||
|
||||
# Add context if available
|
||||
if fact.context:
|
||||
@@ -94,24 +68,53 @@ def format_facts_for_prompt(facts: List[MemoryFact]) -> str:
|
||||
if isinstance(occurred_start, str):
|
||||
fact_obj["occurred_start"] = occurred_start
|
||||
elif isinstance(occurred_start, datetime):
|
||||
fact_obj["occurred_start"] = occurred_start.strftime('%Y-%m-%d %H:%M:%S')
|
||||
fact_obj["occurred_start"] = occurred_start.strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
formatted.append(fact_obj)
|
||||
|
||||
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,
|
||||
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"""
|
||||
@@ -133,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:
|
||||
@@ -143,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:
|
||||
@@ -176,109 +184,78 @@ def get_system_message(disposition: DispositionTraits) -> str:
|
||||
elif disposition.empathy <= 2:
|
||||
instructions.append("Focus on facts and outcomes rather than emotional context.")
|
||||
|
||||
disposition_instruction = " ".join(instructions) if instructions else "Balance your disposition traits when interpreting information."
|
||||
disposition_instruction = (
|
||||
" ".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."
|
||||
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(
|
||||
async def reflect(
|
||||
llm_config,
|
||||
text: str,
|
||||
query: str
|
||||
) -> List[Opinion]:
|
||||
query: str,
|
||||
experience_facts: list[str] = None,
|
||||
world_facts: list[str] = None,
|
||||
name: str = "Assistant",
|
||||
disposition: DispositionTraits = None,
|
||||
background: str = "",
|
||||
context: str = None,
|
||||
) -> str:
|
||||
"""
|
||||
Extract opinions with reasons and confidence from text using LLM.
|
||||
Standalone reflect function for generating answers based on facts.
|
||||
|
||||
This is a static version of the reflect operation that can be called
|
||||
without a MemoryEngine instance, useful for testing.
|
||||
|
||||
Args:
|
||||
llm_config: LLM configuration to use
|
||||
text: Text to extract opinions from
|
||||
query: The original query that prompted this response
|
||||
llm_config: LLM provider instance
|
||||
query: Question to answer
|
||||
experience_facts: List of experience/agent fact strings
|
||||
world_facts: List of world fact strings
|
||||
name: Name of the agent/persona
|
||||
disposition: Disposition traits (defaults to neutral)
|
||||
background: Background information
|
||||
context: Additional context for the prompt
|
||||
|
||||
Returns:
|
||||
List of Opinion objects with text and confidence
|
||||
Generated answer text
|
||||
"""
|
||||
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.
|
||||
# Default disposition if not provided
|
||||
if disposition is None:
|
||||
disposition = DispositionTraits(skepticism=3, literalism=3, empathy=3)
|
||||
|
||||
ORIGINAL QUESTION:
|
||||
{query}
|
||||
# Convert string lists to MemoryFact format for formatting
|
||||
def to_memory_facts(facts: list[str], fact_type: str) -> list[MemoryFact]:
|
||||
if not facts:
|
||||
return []
|
||||
return [MemoryFact(id=f"test-{i}", text=f, fact_type=fact_type) for i, f in enumerate(facts)]
|
||||
|
||||
ANSWER PROVIDED:
|
||||
{text}
|
||||
agent_results = to_memory_facts(experience_facts or [], "experience")
|
||||
world_results = to_memory_facts(world_facts or [], "world")
|
||||
|
||||
Your task: Find opinions in the answer and rewrite them AS IF YOU ARE THE ONE SAYING THEM.
|
||||
# Format facts for prompt
|
||||
agent_facts_text = format_facts_for_prompt(agent_results)
|
||||
world_facts_text = format_facts_for_prompt(world_results)
|
||||
|
||||
An opinion is a judgment, viewpoint, or conclusion that goes beyond just stating facts.
|
||||
# Build prompt
|
||||
prompt = build_think_prompt(
|
||||
agent_facts_text=agent_facts_text,
|
||||
world_facts_text=world_facts_text,
|
||||
query=query,
|
||||
name=name,
|
||||
disposition=disposition,
|
||||
background=background,
|
||||
context=context,
|
||||
)
|
||||
|
||||
IMPORTANT: Do NOT extract statements like:
|
||||
- "I don't have enough information"
|
||||
- "The facts don't contain information about X"
|
||||
- "I cannot answer because..."
|
||||
system_message = get_system_message(disposition)
|
||||
|
||||
ONLY extract actual opinions about substantive topics.
|
||||
# Call LLM
|
||||
answer_text = await llm_config.call(
|
||||
messages=[{"role": "system", "content": system_message}, {"role": "user", "content": prompt}],
|
||||
scope="memory_think",
|
||||
temperature=0.9,
|
||||
max_completion_tokens=1000,
|
||||
)
|
||||
|
||||
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 answer_text.strip()
|
||||
|
||||
@@ -4,22 +4,38 @@ Search trace models for debugging and visualization.
|
||||
These Pydantic models define the structure of search traces, capturing
|
||||
every step of the spreading activation search process for analysis.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import List, Optional, Dict, Any, Literal
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class TemporalConstraint(BaseModel):
|
||||
"""Detected temporal constraint from query analysis."""
|
||||
|
||||
start: datetime | None = Field(default=None, description="Start of temporal range")
|
||||
end: datetime | None = Field(default=None, description="End of temporal range")
|
||||
|
||||
|
||||
class QueryInfo(BaseModel):
|
||||
"""Information about the search query."""
|
||||
|
||||
query_text: str = Field(description="Original query text")
|
||||
query_embedding: List[float] = Field(description="Generated query embedding vector")
|
||||
query_embedding: list[float] = Field(description="Generated query embedding vector")
|
||||
timestamp: datetime = Field(description="When the query was executed")
|
||||
budget: int = Field(description="Maximum nodes to explore")
|
||||
max_tokens: int = Field(description="Maximum tokens to return in results")
|
||||
tags: list[str] | None = Field(default=None, description="Tags filter applied to recall")
|
||||
tags_match: str | None = Field(default=None, description="Tags matching mode: any, all, any_strict, all_strict")
|
||||
temporal_constraint: TemporalConstraint | None = Field(
|
||||
default=None, description="Detected temporal range from query"
|
||||
)
|
||||
|
||||
|
||||
class EntryPoint(BaseModel):
|
||||
"""An entry point node selected for search."""
|
||||
|
||||
node_id: str = Field(description="Memory unit ID")
|
||||
text: str = Field(description="Memory unit text content")
|
||||
similarity_score: float = Field(description="Cosine similarity to query", ge=0.0, le=1.0)
|
||||
@@ -28,6 +44,7 @@ class EntryPoint(BaseModel):
|
||||
|
||||
class WeightComponents(BaseModel):
|
||||
"""Breakdown of weight calculation components."""
|
||||
|
||||
activation: float = Field(description="Activation from spreading (can exceed 1.0 through accumulation)", ge=0.0)
|
||||
semantic_similarity: float = Field(description="Semantic similarity to query", ge=0.0, le=1.0)
|
||||
recency: float = Field(description="Recency weight", ge=0.0, le=1.0)
|
||||
@@ -43,99 +60,119 @@ class WeightComponents(BaseModel):
|
||||
|
||||
class LinkInfo(BaseModel):
|
||||
"""Information about a link to a neighbor."""
|
||||
|
||||
to_node_id: str = Field(description="Target node ID")
|
||||
link_type: Literal["temporal", "semantic", "entity"] = Field(description="Type of link")
|
||||
link_weight: float = Field(description="Weight of the link (can exceed 1.0 when aggregating multiple connections)", ge=0.0)
|
||||
entity_id: Optional[str] = Field(default=None, description="Entity ID if link_type is 'entity'")
|
||||
new_activation: Optional[float] = Field(default=None, description="Activation that would be passed to neighbor (None for supplementary links)")
|
||||
link_weight: float = Field(
|
||||
description="Weight of the link (can exceed 1.0 when aggregating multiple connections)", ge=0.0
|
||||
)
|
||||
entity_id: str | None = Field(default=None, description="Entity ID if link_type is 'entity'")
|
||||
new_activation: float | None = Field(
|
||||
default=None, description="Activation that would be passed to neighbor (None for supplementary links)"
|
||||
)
|
||||
followed: bool = Field(description="Whether this link was followed (or pruned)")
|
||||
prune_reason: Optional[str] = Field(default=None, description="Why link was not followed (if not followed)")
|
||||
is_supplementary: bool = Field(default=False, description="Whether this is a supplementary link (multiple connections to same node)")
|
||||
prune_reason: str | None = Field(default=None, description="Why link was not followed (if not followed)")
|
||||
is_supplementary: bool = Field(
|
||||
default=False, description="Whether this is a supplementary link (multiple connections to same node)"
|
||||
)
|
||||
|
||||
|
||||
class NodeVisit(BaseModel):
|
||||
"""Information about visiting a node during search."""
|
||||
|
||||
step: int = Field(description="Step number in search (1-based)")
|
||||
node_id: str = Field(description="Memory unit ID")
|
||||
text: str = Field(description="Memory unit text content")
|
||||
context: str = Field(description="Memory unit context")
|
||||
event_date: Optional[datetime] = Field(default=None, description="When the memory occurred")
|
||||
access_count: int = Field(description="Number of times accessed before this search")
|
||||
event_date: datetime | None = Field(default=None, description="When the memory occurred")
|
||||
|
||||
# How this node was reached
|
||||
is_entry_point: bool = Field(description="Whether this is an entry point")
|
||||
parent_node_id: Optional[str] = Field(default=None, description="Node that led to this one")
|
||||
link_type: Optional[Literal["temporal", "semantic", "entity"]] = Field(default=None, description="Type of link from parent")
|
||||
link_weight: Optional[float] = Field(default=None, description="Weight of link from parent")
|
||||
parent_node_id: str | None = Field(default=None, description="Node that led to this one")
|
||||
link_type: Literal["temporal", "semantic", "entity"] | None = Field(
|
||||
default=None, description="Type of link from parent"
|
||||
)
|
||||
link_weight: float | None = Field(default=None, description="Weight of link from parent")
|
||||
|
||||
# Weights
|
||||
weights: WeightComponents = Field(description="Weight calculation breakdown")
|
||||
|
||||
# Neighbors discovered from this node
|
||||
neighbors_explored: List[LinkInfo] = Field(default_factory=list, description="Links explored from this node")
|
||||
neighbors_explored: list[LinkInfo] = Field(default_factory=list, description="Links explored from this node")
|
||||
|
||||
# Ranking
|
||||
final_rank: Optional[int] = Field(default=None, description="Final rank in results (1-based, None if not in top-k)")
|
||||
final_rank: int | None = Field(default=None, description="Final rank in results (1-based, None if not in top-k)")
|
||||
|
||||
|
||||
class PruningDecision(BaseModel):
|
||||
"""Records when a node was considered but not visited."""
|
||||
|
||||
node_id: str = Field(description="Node that was pruned")
|
||||
reason: Literal["already_visited", "activation_too_low", "budget_exhausted"] = Field(description="Why it was pruned")
|
||||
reason: Literal["already_visited", "activation_too_low", "budget_exhausted"] = Field(
|
||||
description="Why it was pruned"
|
||||
)
|
||||
activation: float = Field(description="Activation value when pruned")
|
||||
would_have_been_step: int = Field(description="What step it would have been if visited")
|
||||
|
||||
|
||||
class SearchPhaseMetrics(BaseModel):
|
||||
"""Performance metrics for a search phase."""
|
||||
|
||||
phase_name: str = Field(description="Name of the phase")
|
||||
duration_seconds: float = Field(description="Time taken in seconds")
|
||||
details: Dict[str, Any] = Field(default_factory=dict, description="Additional phase-specific metrics")
|
||||
details: dict[str, Any] = Field(default_factory=dict, description="Additional phase-specific metrics")
|
||||
|
||||
|
||||
class RetrievalResult(BaseModel):
|
||||
"""A single result from a retrieval method."""
|
||||
|
||||
rank: int = Field(description="Rank in this retrieval method (1-based)")
|
||||
node_id: str = Field(description="Memory unit ID")
|
||||
text: str = Field(description="Memory unit text content")
|
||||
context: str = Field(default="", description="Memory unit context")
|
||||
event_date: Optional[datetime] = Field(default=None, description="When the memory occurred")
|
||||
fact_type: Optional[str] = Field(default=None, description="Fact type (world, experience, opinion)")
|
||||
event_date: datetime | None = Field(default=None, description="When the memory occurred")
|
||||
fact_type: str | None = Field(default=None, description="Fact type (world, experience, opinion)")
|
||||
score: float = Field(description="Score from this retrieval method")
|
||||
score_name: str = Field(description="Name of the score (e.g., 'similarity', 'bm25_score', 'activation')")
|
||||
|
||||
|
||||
class RetrievalMethodResults(BaseModel):
|
||||
"""Results from a single retrieval method."""
|
||||
|
||||
method_name: Literal["semantic", "bm25", "graph", "temporal"] = Field(description="Name of retrieval method")
|
||||
fact_type: Optional[str] = Field(default=None, description="Fact type this retrieval was for (world, experience, opinion)")
|
||||
results: List[RetrievalResult] = Field(description="Retrieved results with ranks")
|
||||
fact_type: str | None = Field(
|
||||
default=None, description="Fact type this retrieval was for (world, experience, opinion)"
|
||||
)
|
||||
results: list[RetrievalResult] = Field(description="Retrieved results with ranks")
|
||||
duration_seconds: float = Field(description="Time taken for this retrieval")
|
||||
metadata: Dict[str, Any] = Field(default_factory=dict, description="Method-specific metadata")
|
||||
metadata: dict[str, Any] = Field(default_factory=dict, description="Method-specific metadata")
|
||||
|
||||
|
||||
class RRFMergeResult(BaseModel):
|
||||
"""A result after RRF merging."""
|
||||
|
||||
node_id: str = Field(description="Memory unit ID")
|
||||
text: str = Field(description="Memory unit text content")
|
||||
rrf_score: float = Field(description="Reciprocal Rank Fusion score")
|
||||
source_ranks: Dict[str, int] = Field(description="Rank in each source that contributed (method_name -> rank)")
|
||||
source_ranks: dict[str, int] = Field(description="Rank in each source that contributed (method_name -> rank)")
|
||||
final_rrf_rank: int = Field(description="Rank after RRF merge (1-based)")
|
||||
|
||||
|
||||
class RerankedResult(BaseModel):
|
||||
"""A result after reranking."""
|
||||
|
||||
node_id: str = Field(description="Memory unit ID")
|
||||
text: str = Field(description="Memory unit text content")
|
||||
rerank_score: float = Field(description="Final reranking score")
|
||||
rerank_rank: int = Field(description="Rank after reranking (1-based)")
|
||||
rrf_rank: int = Field(description="Original RRF rank before reranking")
|
||||
rank_change: int = Field(description="Change in rank (positive = moved up)")
|
||||
score_components: Dict[str, float] = Field(default_factory=dict, description="Score breakdown")
|
||||
score_components: dict[str, float] = Field(default_factory=dict, description="Score breakdown")
|
||||
|
||||
|
||||
class SearchSummary(BaseModel):
|
||||
"""Summary statistics about the search."""
|
||||
|
||||
total_nodes_visited: int = Field(description="Total nodes visited")
|
||||
total_nodes_pruned: int = Field(description="Total nodes pruned")
|
||||
entry_points_found: int = Field(description="Number of entry points")
|
||||
@@ -150,33 +187,36 @@ class SearchSummary(BaseModel):
|
||||
entity_links_followed: int = Field(default=0, description="Entity links followed")
|
||||
|
||||
# Phase timings
|
||||
phase_metrics: List[SearchPhaseMetrics] = Field(default_factory=list, description="Metrics for each phase")
|
||||
phase_metrics: list[SearchPhaseMetrics] = Field(default_factory=list, description="Metrics for each phase")
|
||||
|
||||
|
||||
class SearchTrace(BaseModel):
|
||||
"""Complete trace of a search operation."""
|
||||
|
||||
query: QueryInfo = Field(description="Query information")
|
||||
|
||||
# New 4-way retrieval architecture
|
||||
retrieval_results: List[RetrievalMethodResults] = Field(default_factory=list, description="Results from each retrieval method")
|
||||
rrf_merged: List[RRFMergeResult] = Field(default_factory=list, description="Results after RRF merging")
|
||||
reranked: List[RerankedResult] = Field(default_factory=list, description="Results after reranking")
|
||||
retrieval_results: list[RetrievalMethodResults] = Field(
|
||||
default_factory=list, description="Results from each retrieval method"
|
||||
)
|
||||
rrf_merged: list[RRFMergeResult] = Field(default_factory=list, description="Results after RRF merging")
|
||||
reranked: list[RerankedResult] = Field(default_factory=list, description="Results after reranking")
|
||||
|
||||
# Legacy fields (kept for backward compatibility with graph/temporal visualizations)
|
||||
entry_points: List[EntryPoint] = Field(default_factory=list, description="Entry points selected for search (legacy)")
|
||||
visits: List[NodeVisit] = Field(default_factory=list, description="All nodes visited during search (legacy, for graph viz)")
|
||||
pruned: List[PruningDecision] = Field(default_factory=list, description="Nodes that were pruned (legacy)")
|
||||
entry_points: list[EntryPoint] = Field(
|
||||
default_factory=list, description="Entry points selected for search (legacy)"
|
||||
)
|
||||
visits: list[NodeVisit] = Field(
|
||||
default_factory=list, description="All nodes visited during search (legacy, for graph viz)"
|
||||
)
|
||||
pruned: list[PruningDecision] = Field(default_factory=list, description="Nodes that were pruned (legacy)")
|
||||
|
||||
summary: SearchSummary = Field(description="Summary statistics")
|
||||
|
||||
# Final results (for comparison with visits)
|
||||
final_results: List[Dict[str, Any]] = Field(description="Final ranked results returned to user")
|
||||
final_results: list[dict[str, Any]] = Field(description="Final ranked results returned to user")
|
||||
|
||||
model_config = {
|
||||
"json_encoders": {
|
||||
datetime: lambda v: v.isoformat()
|
||||
}
|
||||
}
|
||||
model_config = {"json_encoders": {datetime: lambda v: v.isoformat()}}
|
||||
|
||||
def to_json(self, **kwargs) -> str:
|
||||
"""Export trace as JSON string."""
|
||||
@@ -186,14 +226,14 @@ class SearchTrace(BaseModel):
|
||||
"""Export trace as dictionary."""
|
||||
return self.model_dump()
|
||||
|
||||
def get_visit_by_node_id(self, node_id: str) -> Optional[NodeVisit]:
|
||||
def get_visit_by_node_id(self, node_id: str) -> NodeVisit | None:
|
||||
"""Find a visit by node ID."""
|
||||
for visit in self.visits:
|
||||
if visit.node_id == node_id:
|
||||
return visit
|
||||
return None
|
||||
|
||||
def get_search_path_to_node(self, node_id: str) -> List[NodeVisit]:
|
||||
def get_search_path_to_node(self, node_id: str) -> list[NodeVisit]:
|
||||
"""Get the path from entry point to a specific node."""
|
||||
path = []
|
||||
current_visit = self.get_visit_by_node_id(node_id)
|
||||
@@ -207,10 +247,10 @@ class SearchTrace(BaseModel):
|
||||
|
||||
return path
|
||||
|
||||
def get_nodes_by_link_type(self, link_type: Literal["temporal", "semantic", "entity"]) -> List[NodeVisit]:
|
||||
def get_nodes_by_link_type(self, link_type: Literal["temporal", "semantic", "entity"]) -> list[NodeVisit]:
|
||||
"""Get all nodes reached via a specific link type."""
|
||||
return [v for v in self.visits if v.link_type == link_type]
|
||||
|
||||
def get_entry_point_nodes(self) -> List[NodeVisit]:
|
||||
def get_entry_point_nodes(self) -> list[NodeVisit]:
|
||||
"""Get all entry point visits."""
|
||||
return [v for v in self.visits if v.is_entry_point]
|
||||
|
||||
@@ -4,24 +4,26 @@ Search tracer for collecting detailed search execution traces.
|
||||
The SearchTracer collects comprehensive information about each step
|
||||
of the spreading activation search process for debugging and visualization.
|
||||
"""
|
||||
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from typing import List, Optional, Dict, Any, Literal
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any, Literal
|
||||
|
||||
from .trace import (
|
||||
SearchTrace,
|
||||
QueryInfo,
|
||||
EntryPoint,
|
||||
NodeVisit,
|
||||
WeightComponents,
|
||||
LinkInfo,
|
||||
NodeVisit,
|
||||
PruningDecision,
|
||||
SearchSummary,
|
||||
SearchPhaseMetrics,
|
||||
RetrievalResult,
|
||||
RetrievalMethodResults,
|
||||
RRFMergeResult,
|
||||
QueryInfo,
|
||||
RerankedResult,
|
||||
RetrievalMethodResults,
|
||||
RetrievalResult,
|
||||
RRFMergeResult,
|
||||
SearchPhaseMetrics,
|
||||
SearchSummary,
|
||||
SearchTrace,
|
||||
TemporalConstraint,
|
||||
WeightComponents,
|
||||
)
|
||||
|
||||
|
||||
@@ -44,7 +46,14 @@ class SearchTracer:
|
||||
json_output = trace.to_json()
|
||||
"""
|
||||
|
||||
def __init__(self, query: str, budget: int, max_tokens: int):
|
||||
def __init__(
|
||||
self,
|
||||
query: str,
|
||||
budget: int,
|
||||
max_tokens: int,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str | None = None,
|
||||
):
|
||||
"""
|
||||
Initialize tracer.
|
||||
|
||||
@@ -52,23 +61,30 @@ class SearchTracer:
|
||||
query: Search query text
|
||||
budget: Maximum nodes to explore
|
||||
max_tokens: Maximum tokens to return in results
|
||||
tags: Tags filter applied to recall
|
||||
tags_match: Tags matching mode (any, all, any_strict, all_strict)
|
||||
"""
|
||||
self.query_text = query
|
||||
self.budget = budget
|
||||
self.max_tokens = max_tokens
|
||||
self.tags = tags
|
||||
self.tags_match = tags_match
|
||||
|
||||
# Trace data
|
||||
self.query_embedding: Optional[List[float]] = None
|
||||
self.start_time: Optional[float] = None
|
||||
self.entry_points: List[EntryPoint] = []
|
||||
self.visits: List[NodeVisit] = []
|
||||
self.pruned: List[PruningDecision] = []
|
||||
self.phase_metrics: List[SearchPhaseMetrics] = []
|
||||
self.query_embedding: list[float] | None = None
|
||||
self.start_time: float | None = None
|
||||
self.entry_points: list[EntryPoint] = []
|
||||
self.visits: list[NodeVisit] = []
|
||||
self.pruned: list[PruningDecision] = []
|
||||
self.phase_metrics: list[SearchPhaseMetrics] = []
|
||||
|
||||
# Temporal constraint detected from query
|
||||
self.temporal_constraint: TemporalConstraint | None = None
|
||||
|
||||
# New 4-way retrieval tracking
|
||||
self.retrieval_results: List[RetrievalMethodResults] = []
|
||||
self.rrf_merged: List[RRFMergeResult] = []
|
||||
self.reranked: List[RerankedResult] = []
|
||||
self.retrieval_results: list[RetrievalMethodResults] = []
|
||||
self.rrf_merged: list[RRFMergeResult] = []
|
||||
self.reranked: list[RerankedResult] = []
|
||||
|
||||
# Tracking state
|
||||
self.current_step = 0
|
||||
@@ -83,10 +99,15 @@ class SearchTracer:
|
||||
"""Start timing the search."""
|
||||
self.start_time = time.time()
|
||||
|
||||
def record_query_embedding(self, embedding: List[float]):
|
||||
def record_query_embedding(self, embedding: list[float]):
|
||||
"""Record the query embedding."""
|
||||
self.query_embedding = embedding
|
||||
|
||||
def record_temporal_constraint(self, start: datetime | None, end: datetime | None):
|
||||
"""Record the detected temporal constraint from query analysis."""
|
||||
if start is not None or end is not None:
|
||||
self.temporal_constraint = TemporalConstraint(start=start, end=end)
|
||||
|
||||
def add_entry_point(self, node_id: str, text: str, similarity: float, rank: int):
|
||||
"""
|
||||
Record an entry point.
|
||||
@@ -114,12 +135,11 @@ class SearchTracer:
|
||||
node_id: str,
|
||||
text: str,
|
||||
context: str,
|
||||
event_date: datetime,
|
||||
access_count: int,
|
||||
event_date: datetime | None,
|
||||
is_entry_point: bool,
|
||||
parent_node_id: Optional[str],
|
||||
link_type: Optional[Literal["temporal", "semantic", "entity"]],
|
||||
link_weight: Optional[float],
|
||||
parent_node_id: str | None,
|
||||
link_type: Literal["temporal", "semantic", "entity"] | None,
|
||||
link_weight: float | None,
|
||||
activation: float,
|
||||
semantic_similarity: float,
|
||||
recency: float,
|
||||
@@ -134,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
|
||||
@@ -173,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,
|
||||
@@ -199,10 +217,10 @@ class SearchTracer:
|
||||
to_node_id: str,
|
||||
link_type: Literal["temporal", "semantic", "entity"],
|
||||
link_weight: float,
|
||||
entity_id: Optional[str],
|
||||
new_activation: Optional[float],
|
||||
entity_id: str | None,
|
||||
new_activation: float | None,
|
||||
followed: bool,
|
||||
prune_reason: Optional[str] = None,
|
||||
prune_reason: str | None = None,
|
||||
is_supplementary: bool = False,
|
||||
):
|
||||
"""
|
||||
@@ -266,7 +284,7 @@ class SearchTracer:
|
||||
)
|
||||
)
|
||||
|
||||
def add_phase_metric(self, phase_name: str, duration_seconds: float, details: Optional[Dict[str, Any]] = None):
|
||||
def add_phase_metric(self, phase_name: str, duration_seconds: float, details: dict[str, Any] | None = None):
|
||||
"""
|
||||
Record metrics for a search phase.
|
||||
|
||||
@@ -286,11 +304,11 @@ class SearchTracer:
|
||||
def add_retrieval_results(
|
||||
self,
|
||||
method_name: Literal["semantic", "bm25", "graph", "temporal"],
|
||||
results: List[tuple], # List of (doc_id, data) tuples
|
||||
results: list[tuple], # List of (doc_id, data) tuples
|
||||
duration_seconds: float,
|
||||
score_field: str, # e.g., "similarity", "bm25_score"
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
fact_type: Optional[str] = None
|
||||
metadata: dict[str, Any] | None = None,
|
||||
fact_type: str | None = None,
|
||||
):
|
||||
"""
|
||||
Record results from a single retrieval method.
|
||||
@@ -331,7 +349,7 @@ class SearchTracer:
|
||||
)
|
||||
)
|
||||
|
||||
def add_rrf_merged(self, merged_results: List[tuple]):
|
||||
def add_rrf_merged(self, merged_results: list[tuple]):
|
||||
"""
|
||||
Record RRF merged results.
|
||||
|
||||
@@ -350,7 +368,7 @@ class SearchTracer:
|
||||
)
|
||||
)
|
||||
|
||||
def add_reranked(self, reranked_results: List[Dict[str, Any]], rrf_merged: List):
|
||||
def add_reranked(self, reranked_results: list[dict[str, Any]], rrf_merged: list):
|
||||
"""
|
||||
Record reranked results.
|
||||
|
||||
@@ -373,7 +391,15 @@ class SearchTracer:
|
||||
# Keys from ScoredResult.to_dict(): cross_encoder_score, cross_encoder_score_normalized,
|
||||
# rrf_normalized, temporal, recency, combined_score, weight
|
||||
score_components = {}
|
||||
for key in ["cross_encoder_score", "cross_encoder_score_normalized", "rrf_score", "rrf_normalized", "temporal", "recency", "combined_score"]:
|
||||
for key in [
|
||||
"cross_encoder_score",
|
||||
"cross_encoder_score_normalized",
|
||||
"rrf_score",
|
||||
"rrf_normalized",
|
||||
"temporal",
|
||||
"recency",
|
||||
"combined_score",
|
||||
]:
|
||||
if key in result and result[key] is not None:
|
||||
score_components[key] = result[key]
|
||||
|
||||
@@ -389,7 +415,7 @@ class SearchTracer:
|
||||
)
|
||||
)
|
||||
|
||||
def finalize(self, final_results: List[Dict[str, Any]]) -> SearchTrace:
|
||||
def finalize(self, final_results: list[dict[str, Any]]) -> SearchTrace:
|
||||
"""
|
||||
Finalize the trace and return the complete SearchTrace object.
|
||||
|
||||
@@ -416,9 +442,12 @@ class SearchTracer:
|
||||
query_info = QueryInfo(
|
||||
query_text=self.query_text,
|
||||
query_embedding=self.query_embedding or [],
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
timestamp=datetime.now(UTC),
|
||||
budget=self.budget,
|
||||
max_tokens=self.max_tokens,
|
||||
tags=self.tags,
|
||||
tags_match=self.tags_match,
|
||||
temporal_constraint=self.temporal_constraint,
|
||||
)
|
||||
|
||||
# Create summary
|
||||
|
||||
@@ -6,8 +6,26 @@ providing type safety and making data flow explicit.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class MPFPTimings:
|
||||
"""Timing breakdown for a single MPFP retrieval call."""
|
||||
|
||||
fact_type: str
|
||||
edge_count: int = 0 # Total edges loaded
|
||||
db_queries: int = 0 # Number of DB queries for edge loading
|
||||
edge_load_time: float = 0.0 # Time spent loading edges from DB
|
||||
traverse: float = 0.0 # Total traversal time (includes edge loading)
|
||||
pattern_count: int = 0 # Number of patterns executed
|
||||
fusion: float = 0.0 # Time for RRF fusion
|
||||
fetch: float = 0.0 # Time to fetch memory unit details
|
||||
seeds_time: float = 0.0 # Time to find semantic seeds (if fallback used)
|
||||
result_count: int = 0 # Number of results returned
|
||||
# Detailed per-hop timing: list of {hop, exec_time, uncached, load_time, edges_loaded, total_time}
|
||||
hop_details: list[dict] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -17,28 +35,29 @@ class RetrievalResult:
|
||||
|
||||
This represents a raw result from the database query, before merging or reranking.
|
||||
"""
|
||||
|
||||
id: str
|
||||
text: str
|
||||
fact_type: str
|
||||
context: Optional[str] = None
|
||||
event_date: Optional[datetime] = None
|
||||
occurred_start: Optional[datetime] = None
|
||||
occurred_end: Optional[datetime] = None
|
||||
mentioned_at: Optional[datetime] = None
|
||||
document_id: Optional[str] = None
|
||||
chunk_id: Optional[str] = None
|
||||
access_count: int = 0
|
||||
embedding: Optional[List[float]] = None
|
||||
context: str | None = None
|
||||
event_date: datetime | None = None
|
||||
occurred_start: datetime | None = None
|
||||
occurred_end: datetime | None = None
|
||||
mentioned_at: datetime | None = None
|
||||
document_id: str | None = None
|
||||
chunk_id: str | None = None
|
||||
embedding: list[float] | None = None
|
||||
tags: list[str] | None = None # Visibility scope tags
|
||||
|
||||
# Retrieval-specific scores (only one will be set depending on retrieval method)
|
||||
similarity: Optional[float] = None # Semantic retrieval
|
||||
bm25_score: Optional[float] = None # BM25 retrieval
|
||||
activation: Optional[float] = None # Graph retrieval (spreading activation)
|
||||
temporal_score: Optional[float] = None # Temporal retrieval
|
||||
temporal_proximity: Optional[float] = None # Temporal retrieval
|
||||
similarity: float | None = None # Semantic retrieval
|
||||
bm25_score: float | None = None # BM25 retrieval
|
||||
activation: float | None = None # Graph retrieval (spreading activation)
|
||||
temporal_score: float | None = None # Temporal retrieval
|
||||
temporal_proximity: float | None = None # Temporal retrieval
|
||||
|
||||
@classmethod
|
||||
def from_db_row(cls, row: Dict[str, Any]) -> "RetrievalResult":
|
||||
def from_db_row(cls, row: dict[str, Any]) -> "RetrievalResult":
|
||||
"""Create from a database row (asyncpg Record converted to dict)."""
|
||||
return cls(
|
||||
id=str(row["id"]),
|
||||
@@ -51,8 +70,8 @@ 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"),
|
||||
bm25_score=row.get("bm25_score"),
|
||||
activation=row.get("activation"),
|
||||
@@ -68,13 +87,14 @@ class MergedCandidate:
|
||||
|
||||
Contains the original retrieval data plus RRF metadata.
|
||||
"""
|
||||
|
||||
# Original retrieval data
|
||||
retrieval: RetrievalResult
|
||||
|
||||
# RRF metadata
|
||||
rrf_score: float
|
||||
rrf_rank: int = 0
|
||||
source_ranks: Dict[str, int] = field(default_factory=dict) # method_name -> rank
|
||||
source_ranks: dict[str, int] = field(default_factory=dict) # method_name -> rank
|
||||
|
||||
@property
|
||||
def id(self) -> str:
|
||||
@@ -89,6 +109,7 @@ class ScoredResult:
|
||||
|
||||
Contains all retrieval/merge data plus reranking scores and combined score.
|
||||
"""
|
||||
|
||||
# Original merged candidate
|
||||
candidate: MergedCandidate
|
||||
|
||||
@@ -115,7 +136,7 @@ class ScoredResult:
|
||||
"""Convenience property to access retrieval data."""
|
||||
return self.candidate.retrieval
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
"""
|
||||
Convert to dict for backwards compatibility.
|
||||
|
||||
@@ -133,8 +154,8 @@ 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,
|
||||
"bm25_score": self.retrieval.bm25_score,
|
||||
}
|
||||
|
||||
@@ -1,38 +1,49 @@
|
||||
"""
|
||||
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)
|
||||
"""
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, Optional, Callable, Awaitable
|
||||
import asyncio
|
||||
|
||||
import json
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Awaitable, Callable
|
||||
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.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize the task backend."""
|
||||
self._executor: Optional[Callable[[Dict[str, Any]], Awaitable[None]]] = None
|
||||
self._executor: Callable[[dict[str, Any]], Awaitable[None]] | None = None
|
||||
self._initialized = False
|
||||
|
||||
def set_executor(self, executor: Callable[[Dict[str, Any]], Awaitable[None]]):
|
||||
def set_executor(self, executor: Callable[[dict[str, Any]], Awaitable[None]]):
|
||||
"""
|
||||
Set the executor callback for processing tasks.
|
||||
|
||||
@@ -44,12 +55,12 @@ 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
|
||||
|
||||
@abstractmethod
|
||||
async def submit_task(self, task_dict: Dict[str, Any]):
|
||||
async def submit_task(self, task_dict: dict[str, Any]):
|
||||
"""
|
||||
Submit a task for execution.
|
||||
|
||||
@@ -61,11 +72,11 @@ class TaskBackend(ABC):
|
||||
@abstractmethod
|
||||
async def shutdown(self):
|
||||
"""
|
||||
Shutdown the backend gracefully (e.g., stop workers, close connections).
|
||||
Shutdown the backend gracefully.
|
||||
"""
|
||||
pass
|
||||
|
||||
async def _execute_task(self, task_dict: Dict[str, Any]):
|
||||
async def _execute_task(self, task_dict: dict[str, Any]):
|
||||
"""
|
||||
Execute a task through the registered executor.
|
||||
|
||||
@@ -73,60 +84,36 @@ class TaskBackend(ABC):
|
||||
task_dict: Task dictionary to execute
|
||||
"""
|
||||
if self._executor is None:
|
||||
task_type = task_dict.get('type', 'unknown')
|
||||
task_type = task_dict.get("type", "unknown")
|
||||
logger.warning(f"No executor registered, skipping task {task_type}")
|
||||
return
|
||||
|
||||
try:
|
||||
await self._executor(task_dict)
|
||||
except Exception as e:
|
||||
task_type = task_dict.get('type', 'unknown')
|
||||
task_type = task_dict.get("type", "unknown")
|
||||
logger.error(f"Error executing task {task_type}: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
|
||||
class AsyncIOQueueBackend(TaskBackend):
|
||||
class SyncTaskBackend(TaskBackend):
|
||||
"""
|
||||
Task backend implementation using asyncio queues.
|
||||
Synchronous task backend that executes tasks immediately.
|
||||
|
||||
This is the default implementation that uses in-process asyncio queues
|
||||
and a periodic consumer worker.
|
||||
This is useful for tests and embedded/CLI usage where we don't want
|
||||
background workers. Tasks are executed inline rather than being queued.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
batch_size: int = 100,
|
||||
batch_interval: float = 1.0
|
||||
):
|
||||
"""
|
||||
Initialize AsyncIO queue backend.
|
||||
|
||||
Args:
|
||||
batch_size: Maximum number of tasks to process in one batch
|
||||
batch_interval: Maximum time (seconds) to wait before processing batch
|
||||
"""
|
||||
super().__init__()
|
||||
self._queue: Optional[asyncio.Queue] = None
|
||||
self._worker_task: Optional[asyncio.Task] = None
|
||||
self._shutdown_event: Optional[asyncio.Event] = None
|
||||
self._batch_size = batch_size
|
||||
self._batch_interval = batch_interval
|
||||
|
||||
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())
|
||||
"""No-op for sync backend."""
|
||||
self._initialized = True
|
||||
logger.info("AsyncIOQueueBackend initialized")
|
||||
logger.debug("SyncTaskBackend initialized")
|
||||
|
||||
async def submit_task(self, task_dict: Dict[str, Any]):
|
||||
async def submit_task(self, task_dict: dict[str, Any]):
|
||||
"""
|
||||
Submit a task by putting it in the queue.
|
||||
Execute the task immediately (synchronously).
|
||||
|
||||
Args:
|
||||
task_dict: Task dictionary to execute
|
||||
@@ -134,90 +121,131 @@ class AsyncIOQueueBackend(TaskBackend):
|
||||
if not self._initialized:
|
||||
await self.initialize()
|
||||
|
||||
await self._queue.put(task_dict)
|
||||
task_type = task_dict.get('type', 'unknown')
|
||||
task_id = task_dict.get('id')
|
||||
await self._execute_task(task_dict)
|
||||
|
||||
async def wait_for_pending_tasks(self, timeout: float = 5.0):
|
||||
async def shutdown(self):
|
||||
"""No-op for sync backend."""
|
||||
self._initialized = False
|
||||
logger.debug("SyncTaskBackend shutdown")
|
||||
|
||||
|
||||
class BrokerTaskBackend(TaskBackend):
|
||||
"""
|
||||
Task backend using PostgreSQL as broker.
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pool_getter: Callable[[], "asyncpg.Pool"],
|
||||
schema: str | None = None,
|
||||
):
|
||||
"""
|
||||
Wait for all pending tasks in the queue to be processed.
|
||||
Initialize the broker task backend.
|
||||
|
||||
This is useful in tests to ensure background tasks complete before assertions.
|
||||
Args:
|
||||
pool_getter: Callable that returns the asyncpg connection pool
|
||||
schema: Database schema for multi-tenant support (optional)
|
||||
"""
|
||||
super().__init__()
|
||||
self._pool_getter = pool_getter
|
||||
self._schema = schema
|
||||
|
||||
async def initialize(self):
|
||||
"""Initialize the backend."""
|
||||
self._initialized = True
|
||||
logger.info("BrokerTaskBackend initialized")
|
||||
|
||||
async def submit_task(self, task_dict: dict[str, Any]):
|
||||
"""
|
||||
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 store (must be JSON serializable)
|
||||
"""
|
||||
if not self._initialized:
|
||||
await self.initialize()
|
||||
|
||||
pool = self._pool_getter()
|
||||
operation_id = task_dict.get("operation_id")
|
||||
task_type = task_dict.get("type", "unknown")
|
||||
bank_id = task_dict.get("bank_id")
|
||||
payload_json = json.dumps(task_dict)
|
||||
|
||||
table = fq_table("async_operations", self._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 pending tasks to be processed.
|
||||
|
||||
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
|
||||
"""
|
||||
if not self._initialized or self._queue is None:
|
||||
return
|
||||
import asyncio
|
||||
|
||||
pool = self._pool_getter()
|
||||
table = fq_table("async_operations", self._schema)
|
||||
|
||||
# Wait for queue to be empty and give worker time to process
|
||||
start_time = asyncio.get_event_loop().time()
|
||||
while asyncio.get_event_loop().time() - start_time < timeout:
|
||||
if self._queue.empty():
|
||||
# Queue is empty, give worker a bit more time to finish any in-flight task
|
||||
await asyncio.sleep(0.3)
|
||||
# Check again - if still empty, we're done
|
||||
if self._queue.empty():
|
||||
return
|
||||
else:
|
||||
# Queue not empty, wait a bit
|
||||
await asyncio.sleep(0.1)
|
||||
# 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
|
||||
"""
|
||||
)
|
||||
|
||||
async def shutdown(self):
|
||||
"""Shutdown the worker and drain the queue."""
|
||||
if not self._initialized:
|
||||
return
|
||||
if count == 0:
|
||||
return
|
||||
|
||||
logger.info("Shutting down AsyncIOQueueBackend...")
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
# 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 _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
|
||||
)
|
||||
tasks.append(task_dict)
|
||||
except asyncio.TimeoutError:
|
||||
break
|
||||
|
||||
# Process batch
|
||||
if tasks:
|
||||
# Execute tasks concurrently
|
||||
await asyncio.gather(
|
||||
*[self._execute_task(task_dict) for task_dict in tasks],
|
||||
return_exceptions=True
|
||||
)
|
||||
|
||||
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")
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user