Compare commits
32
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a104e616d3 | ||
|
|
2b24ff85d5 | ||
|
|
29b7955a55 | ||
|
|
0a8ac1d050 | ||
|
|
1949c28ad8 | ||
|
|
0e73bbf3a0 | ||
|
|
1c6acc3ba0 | ||
|
|
7982c48d3d | ||
|
|
8ecb5d3a0c | ||
|
|
3d5d360806 | ||
|
|
9de2e90c46 | ||
|
|
b0dbbc51bb | ||
|
|
3b479ebee0 | ||
|
|
95127e13be | ||
|
|
ae80876671 | ||
|
|
476a62da47 | ||
|
|
5aaa769ab9 | ||
|
|
04f01ab9ab | ||
|
|
63f51385c4 | ||
|
|
e468a4e19f | ||
|
|
c0a0f447b7 | ||
|
|
84927ccc99 | ||
|
|
a6e8944ff0 | ||
|
|
f6d890f6ed | ||
|
|
1fa8d9150c | ||
|
|
656777c2be | ||
|
|
b36807ad3b | ||
|
|
11ac9cd9a5 | ||
|
|
9394cf92f2 | ||
|
|
47be07f97f | ||
|
|
bb1f9cb221 | ||
|
|
7dd68538bb |
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 ""
|
||||
@@ -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:
|
||||
|
||||
@@ -117,6 +117,47 @@ 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: Build
|
||||
run: npm run build --workspace=hindsight-control-plane
|
||||
|
||||
- name: Publish to npm
|
||||
working-directory: ./hindsight-control-plane
|
||||
run: npm publish --access public
|
||||
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 +222,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 +247,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 +258,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: .
|
||||
@@ -227,6 +290,8 @@ jobs:
|
||||
platforms: linux/amd64,linux/arm64
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
|
||||
release-helm-chart:
|
||||
runs-on: ubuntu-latest
|
||||
@@ -263,7 +328,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 +351,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:
|
||||
@@ -320,6 +391,8 @@ jobs:
|
||||
cp artifacts/python-packages/hindsight-integrations/litellm/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
|
||||
|
||||
+161
-2
@@ -38,6 +38,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 +80,41 @@ 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
|
||||
node hindsight-control-plane/bin/cli.js --help
|
||||
|
||||
build-docs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
@@ -130,7 +188,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 +206,15 @@ jobs:
|
||||
file: docker/standalone/Dockerfile
|
||||
target: ${{ matrix.target }}
|
||||
push: false
|
||||
load: false
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
|
||||
# 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
|
||||
@@ -472,4 +539,96 @@ jobs:
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-integrations/litellm
|
||||
run: uv run pytest tests -v
|
||||
run: uv run pytest tests -v
|
||||
|
||||
test-doc-examples:
|
||||
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: 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 --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install Python client dependencies
|
||||
working-directory: ./hindsight-clients/python
|
||||
run: uv sync --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: Show API server logs
|
||||
if: always()
|
||||
run: |
|
||||
echo "=== API Server Logs ==="
|
||||
cat /tmp/api-server.log || echo "No API server log found"
|
||||
@@ -32,6 +32,9 @@ logs/
|
||||
|
||||
.DS_Store
|
||||
|
||||
# Generated docs files
|
||||
hindsight-docs/static/llms-full.txt
|
||||
|
||||
|
||||
hindsight-dev/benchmarks/locomo/results/
|
||||
hindsight-dev/benchmarks/longmemeval/results/
|
||||
|
||||
@@ -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://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)
|
||||
[](https://opensource.org/licenses/MIT)
|
||||

|
||||

|
||||
|
||||
|
||||
</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)
|
||||
@@ -223,6 +243,10 @@ client.reflect(bank_id="my-bank", query="What should I know about Alice?")
|
||||
- [Slack](https://join.slack.com/t/hindsight-space/shared_invite/zt-3klo21kua-VUCC_zHP5rIcXFB1_5yw6A)
|
||||
- [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
|
||||
|
||||
@@ -60,8 +60,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 +72,39 @@ 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
|
||||
# Next.js standalone output structure varies, so we find server.js and work from there
|
||||
RUN mkdir -p standalone/.next && \
|
||||
STANDALONE_ROOT=$(dirname $(find .next/standalone -name "server.js" | head -1)) && \
|
||||
cp -r "$STANDALONE_ROOT"/* standalone/ && \
|
||||
cp -r .next/static standalone/.next/static && \
|
||||
mkdir -p standalone/public && \
|
||||
cp -r public/* standalone/public/ 2>/dev/null || true
|
||||
|
||||
# =============================================================================
|
||||
# Stage: Final Image - API Only
|
||||
@@ -104,14 +113,16 @@ FROM python:3.11-slim AS api-only
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install pg0 dependencies
|
||||
# Install pg0 dependencies (procps provides 'kill' command needed by pg0)
|
||||
# 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
|
||||
|
||||
@@ -171,9 +182,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,14 +211,16 @@ FROM python:3.11-slim AS standalone
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install Node.js, curl, uv, and pg0 dependencies
|
||||
# Install Node.js, curl, uv, and pg0 dependencies (procps provides 'kill' command needed by pg0)
|
||||
# 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/* \
|
||||
@@ -224,9 +237,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
|
||||
|
||||
|
||||
@@ -23,7 +23,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 +43,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.1.8
|
||||
appVersion: "0.1.8"
|
||||
keywords:
|
||||
- ai
|
||||
- memory
|
||||
|
||||
+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`, `groq`, `gemini`, `ollama` | `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,23 +3,24 @@ 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
|
||||
|
||||
__all__ = [
|
||||
"MemoryEngine",
|
||||
|
||||
@@ -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
|
||||
@@ -128,10 +128,7 @@ def run_migrations_online() -> None:
|
||||
connection.execute(text("SET SESSION CHARACTERISTICS AS TRANSACTION READ WRITE"))
|
||||
connection.commit() # Commit the SET command
|
||||
|
||||
context.configure(
|
||||
connection=connection,
|
||||
target_metadata=target_metadata
|
||||
)
|
||||
context.configure(connection=connection, target_metadata=target_metadata)
|
||||
|
||||
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")
|
||||
|
||||
+7
-12
@@ -5,20 +5,19 @@ Revises: c8e5f2a3b4d1
|
||||
Create Date: 2024-12-04 15:00:00.000000
|
||||
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'd9f6a3b4c5e2'
|
||||
down_revision = 'c8e5f2a3b4d1'
|
||||
revision = "d9f6a3b4c5e2"
|
||||
down_revision = "c8e5f2a3b4d1"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
# Drop old check constraint FIRST (before updating data)
|
||||
op.drop_constraint('memory_units_fact_type_check', 'memory_units', type_='check')
|
||||
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'")
|
||||
@@ -27,22 +26,18 @@ def upgrade():
|
||||
|
||||
# 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():
|
||||
# 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'")
|
||||
|
||||
# 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')"
|
||||
)
|
||||
|
||||
+23
-15
@@ -8,17 +8,17 @@ 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 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 upgrade() -> None:
|
||||
@@ -31,17 +31,21 @@ 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("""
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE banks
|
||||
SET disposition = '{"skepticism": 3, "literalism": 3, "empathy": 3}'::jsonb
|
||||
WHERE disposition IS NOT NULL
|
||||
"""))
|
||||
""")
|
||||
)
|
||||
|
||||
# Update the default for new banks
|
||||
conn.execute(sa.text("""
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
ALTER TABLE banks
|
||||
ALTER COLUMN disposition SET DEFAULT '{"skepticism": 3, "literalism": 3, "empathy": 3}'::jsonb
|
||||
"""))
|
||||
""")
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
@@ -49,14 +53,18 @@ def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# Revert to Big Five format with default values
|
||||
conn.execute(sa.text("""
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE banks
|
||||
SET disposition = '{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}'::jsonb
|
||||
WHERE disposition IS NOT NULL
|
||||
"""))
|
||||
""")
|
||||
)
|
||||
|
||||
# Update the default for new banks
|
||||
conn.execute(sa.text("""
|
||||
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
|
||||
"""))
|
||||
""")
|
||||
)
|
||||
|
||||
@@ -5,18 +5,18 @@ 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 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 upgrade() -> None:
|
||||
@@ -24,42 +24,51 @@ def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# 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'
|
||||
"""))
|
||||
""")
|
||||
)
|
||||
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'
|
||||
"""))
|
||||
""")
|
||||
)
|
||||
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("""
|
||||
result = conn.execute(
|
||||
sa.text("""
|
||||
SELECT column_name
|
||||
FROM information_schema.columns
|
||||
WHERE table_name = 'banks' AND column_name = 'disposition'
|
||||
"""))
|
||||
""")
|
||||
)
|
||||
if result.fetchone():
|
||||
op.alter_column('banks', 'disposition', new_column_name='personality')
|
||||
op.alter_column("banks", "disposition", new_column_name="personality")
|
||||
|
||||
@@ -3,8 +3,10 @@ Unified API module for Hindsight.
|
||||
|
||||
Provides both HTTP REST API and MCP (Model Context Protocol) server.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import FastAPI
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
@@ -17,7 +19,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.
|
||||
@@ -47,10 +49,8 @@ def create_app(
|
||||
# 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
|
||||
@@ -77,15 +77,15 @@ def create_app(
|
||||
|
||||
# Re-export commonly used items for backwards compatibility
|
||||
from .http import (
|
||||
RecallRequest,
|
||||
RecallResult,
|
||||
RecallResponse,
|
||||
MemoryItem,
|
||||
RetainRequest,
|
||||
ReflectRequest,
|
||||
ReflectResponse,
|
||||
CreateBankRequest,
|
||||
DispositionTraits,
|
||||
MemoryItem,
|
||||
RecallRequest,
|
||||
RecallResponse,
|
||||
RecallResult,
|
||||
ReflectRequest,
|
||||
ReflectResponse,
|
||||
RetainRequest,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -4,27 +4,33 @@ 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
|
||||
|
||||
# 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)
|
||||
_current_bank_id: ContextVar[str | None] = ContextVar("current_bank_id", default=None)
|
||||
|
||||
|
||||
def get_current_bank_id() -> Optional[str]:
|
||||
def get_current_bank_id() -> str | None:
|
||||
"""Get the current bank_id from context (set from URL path)."""
|
||||
return _current_bank_id.get()
|
||||
|
||||
@@ -61,10 +67,7 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
|
||||
"""
|
||||
try:
|
||||
bank_id = get_current_bank_id()
|
||||
await memory.put_batch_async(
|
||||
bank_id=bank_id,
|
||||
contents=[{"content": content, "context": context}]
|
||||
)
|
||||
await memory.retain_batch_async(bank_id=bank_id, contents=[{"content": content, "context": context}])
|
||||
return "Memory stored successfully"
|
||||
except Exception as e:
|
||||
logger.error(f"Error storing memory: {e}", exc_info=True)
|
||||
@@ -88,11 +91,9 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
|
||||
try:
|
||||
bank_id = get_current_bank_id()
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
|
||||
search_result = await memory.recall_async(
|
||||
bank_id=bank_id,
|
||||
query=query,
|
||||
fact_type=list(VALID_RECALL_FACT_TYPES),
|
||||
budget=Budget.LOW
|
||||
bank_id=bank_id, query=query, fact_type=list(VALID_RECALL_FACT_TYPES), budget=Budget.LOW
|
||||
)
|
||||
|
||||
results = [
|
||||
@@ -133,7 +134,7 @@ 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/"):
|
||||
@@ -169,10 +170,7 @@ class MCPMiddleware:
|
||||
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,15 +181,19 @@ 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):
|
||||
|
||||
@@ -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,10 @@ Centralized configuration for Hindsight API.
|
||||
|
||||
All environment variables and their defaults are defined here.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -30,6 +30,8 @@ ENV_PORT = "HINDSIGHT_API_PORT"
|
||||
ENV_LOG_LEVEL = "HINDSIGHT_API_LOG_LEVEL"
|
||||
ENV_MCP_ENABLED = "HINDSIGHT_API_MCP_ENABLED"
|
||||
ENV_GRAPH_RETRIEVER = "HINDSIGHT_API_GRAPH_RETRIEVER"
|
||||
ENV_MCP_LOCAL_BANK_ID = "HINDSIGHT_API_MCP_LOCAL_BANK_ID"
|
||||
ENV_MCP_INSTRUCTIONS = "HINDSIGHT_API_MCP_INSTRUCTIONS"
|
||||
|
||||
# Default values
|
||||
DEFAULT_DATABASE_URL = "pg0"
|
||||
@@ -47,6 +49,27 @@ DEFAULT_PORT = 8888
|
||||
DEFAULT_LOG_LEVEL = "info"
|
||||
DEFAULT_MCP_ENABLED = True
|
||||
DEFAULT_GRAPH_RETRIEVER = "bfs" # Options: "bfs", "mpfp"
|
||||
DEFAULT_MCP_LOCAL_BANK_ID = "mcp"
|
||||
|
||||
# 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"""
|
||||
|
||||
# Required embedding dimension for database schema
|
||||
EMBEDDING_DIMENSION = 384
|
||||
@@ -61,19 +84,19 @@ class HindsightConfig:
|
||||
|
||||
# LLM
|
||||
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
|
||||
|
||||
# Embeddings
|
||||
embeddings_provider: str
|
||||
embeddings_local_model: str
|
||||
embeddings_tei_url: Optional[str]
|
||||
embeddings_tei_url: str | None
|
||||
|
||||
# Reranker
|
||||
reranker_provider: str
|
||||
reranker_local_model: str
|
||||
reranker_tei_url: Optional[str]
|
||||
reranker_tei_url: str | None
|
||||
|
||||
# Server
|
||||
host: str
|
||||
@@ -90,29 +113,24 @@ 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,
|
||||
|
||||
# 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),
|
||||
|
||||
# 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),
|
||||
|
||||
# Server
|
||||
host=os.getenv(ENV_HOST, DEFAULT_HOST),
|
||||
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
|
||||
log_level=os.getenv(ENV_LOG_LEVEL, DEFAULT_LOG_LEVEL),
|
||||
mcp_enabled=os.getenv(ENV_MCP_ENABLED, str(DEFAULT_MCP_ENABLED)).lower() == "true",
|
||||
|
||||
# Recall
|
||||
graph_retriever=os.getenv(ENV_GRAPH_RETRIEVER, DEFAULT_GRAPH_RETRIEVER),
|
||||
)
|
||||
@@ -146,7 +164,8 @@ class HindsightConfig:
|
||||
"""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"
|
||||
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
|
||||
force=True, # Override any existing configuration
|
||||
)
|
||||
|
||||
def log_config(self) -> None:
|
||||
|
||||
@@ -7,24 +7,24 @@ 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
|
||||
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",
|
||||
|
||||
@@ -5,19 +5,19 @@ 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 logging
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import httpx
|
||||
|
||||
from ..config import (
|
||||
ENV_RERANKER_PROVIDER,
|
||||
ENV_RERANKER_LOCAL_MODEL,
|
||||
ENV_RERANKER_TEI_URL,
|
||||
DEFAULT_RERANKER_PROVIDER,
|
||||
DEFAULT_RERANKER_LOCAL_MODEL,
|
||||
DEFAULT_RERANKER_PROVIDER,
|
||||
ENV_RERANKER_LOCAL_MODEL,
|
||||
ENV_RERANKER_PROVIDER,
|
||||
ENV_RERANKER_TEI_URL,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -47,7 +47,7 @@ class CrossEncoderModel(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
|
||||
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs for relevance.
|
||||
|
||||
@@ -72,7 +72,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
- Trained for passage re-ranking
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: Optional[str] = None):
|
||||
def __init__(self, model_name: str | None = None):
|
||||
"""
|
||||
Initialize local SentenceTransformers cross-encoder.
|
||||
|
||||
@@ -104,7 +104,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
self._model = CrossEncoder(self.model_name)
|
||||
logger.info("Reranker: local provider initialized")
|
||||
|
||||
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
|
||||
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs for relevance.
|
||||
|
||||
@@ -117,7 +117,7 @@ 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)
|
||||
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
|
||||
|
||||
|
||||
class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
@@ -153,8 +153,8 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
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
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
@@ -163,6 +163,7 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
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
|
||||
|
||||
@@ -177,14 +178,18 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
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:
|
||||
@@ -209,7 +214,7 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
except httpx.HTTPError as e:
|
||||
raise RuntimeError(f"Failed to connect to TEI server at {self.base_url}: {e}")
|
||||
|
||||
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
|
||||
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs using the remote TEI reranker.
|
||||
|
||||
@@ -229,7 +234,7 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||
|
||||
# Process in batches
|
||||
for i in range(0, len(pairs), self.batch_size):
|
||||
batch = pairs[i:i + self.batch_size]
|
||||
batch = pairs[i : i + self.batch_size]
|
||||
|
||||
# TEI rerank endpoint expects query and texts separately
|
||||
# All pairs in a batch should have the same query for optimal performance
|
||||
@@ -287,15 +292,11 @@ 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'"
|
||||
)
|
||||
raise ValueError(f"{ENV_RERANKER_TEI_URL} is required when {ENV_RERANKER_PROVIDER} is 'tei'")
|
||||
return RemoteTEICrossEncoder(base_url=url)
|
||||
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)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei'"
|
||||
)
|
||||
raise ValueError(f"Unknown reranker provider: {provider}. Supported: 'local', 'tei'")
|
||||
|
||||
@@ -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,6 +83,7 @@ async def acquire_with_retry(pool: asyncpg.Pool, max_retries: int = DEFAULT_MAX_
|
||||
Yields:
|
||||
An asyncpg connection
|
||||
"""
|
||||
|
||||
async def acquire():
|
||||
return await pool.acquire()
|
||||
|
||||
|
||||
@@ -8,20 +8,20 @@ the database schema (pgvector column defined as vector(384)).
|
||||
|
||||
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_LOCAL_MODEL,
|
||||
DEFAULT_EMBEDDINGS_PROVIDER,
|
||||
EMBEDDING_DIMENSION,
|
||||
ENV_EMBEDDINGS_LOCAL_MODEL,
|
||||
ENV_EMBEDDINGS_PROVIDER,
|
||||
ENV_EMBEDDINGS_TEI_URL,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -52,7 +52,7 @@ 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.
|
||||
|
||||
@@ -75,7 +75,7 @@ class LocalSTEmbeddings(Embeddings):
|
||||
embeddings matching the database schema.
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: Optional[str] = None):
|
||||
def __init__(self, model_name: str | None = None):
|
||||
"""
|
||||
Initialize local SentenceTransformers embeddings.
|
||||
|
||||
@@ -123,7 +123,7 @@ class LocalSTEmbeddings(Embeddings):
|
||||
|
||||
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.
|
||||
|
||||
@@ -172,8 +172,8 @@ 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
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
@@ -182,6 +182,7 @@ class RemoteTEIEmbeddings(Embeddings):
|
||||
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 +197,18 @@ class RemoteTEIEmbeddings(Embeddings):
|
||||
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
logger.warning(f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
|
||||
logger.warning(
|
||||
f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
|
||||
)
|
||||
time.sleep(delay)
|
||||
delay *= 2 # Exponential backoff
|
||||
except httpx.HTTPStatusError as e:
|
||||
# Retry on 5xx server errors
|
||||
if e.response.status_code >= 500 and attempt < self.max_retries:
|
||||
last_error = e
|
||||
logger.warning(f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
|
||||
logger.warning(
|
||||
f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
|
||||
)
|
||||
time.sleep(delay)
|
||||
delay *= 2
|
||||
else:
|
||||
@@ -228,7 +233,7 @@ class RemoteTEIEmbeddings(Embeddings):
|
||||
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 +253,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(
|
||||
@@ -278,15 +283,11 @@ 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)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei'"
|
||||
)
|
||||
raise ValueError(f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei'")
|
||||
|
||||
@@ -4,12 +4,13 @@ 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
|
||||
|
||||
# Load spaCy model (singleton)
|
||||
_nlp = None
|
||||
@@ -32,11 +33,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,7 +63,9 @@ 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(
|
||||
"""
|
||||
@@ -70,11 +73,11 @@ class EntityResolver:
|
||||
FROM 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
|
||||
@@ -85,13 +88,13 @@ class EntityResolver:
|
||||
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)
|
||||
""",
|
||||
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 +108,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 +135,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 +151,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 +170,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))
|
||||
@@ -198,7 +201,7 @@ class EntityResolver:
|
||||
last_seen = $2
|
||||
WHERE id = $1::uuid
|
||||
""",
|
||||
entities_to_update
|
||||
entities_to_update,
|
||||
)
|
||||
|
||||
# Batch create new entities using COPY + INSERT for maximum speed
|
||||
@@ -208,7 +211,7 @@ class EntityResolver:
|
||||
# For duplicates, we only insert once and reuse the ID
|
||||
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:
|
||||
@@ -222,7 +225,7 @@ class EntityResolver:
|
||||
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)
|
||||
indices_map.append(indices)
|
||||
|
||||
@@ -241,12 +244,12 @@ class EntityResolver:
|
||||
""",
|
||||
bank_id,
|
||||
entity_names,
|
||||
entity_dates
|
||||
entity_dates,
|
||||
)
|
||||
|
||||
# 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 +260,7 @@ class EntityResolver:
|
||||
bank_id: str,
|
||||
entity_text: str,
|
||||
context: str,
|
||||
nearby_entities: List[Dict],
|
||||
nearby_entities: list[dict],
|
||||
unit_event_date,
|
||||
) -> str:
|
||||
"""
|
||||
@@ -287,14 +290,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,21 +309,17 @@ 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)
|
||||
@@ -338,9 +337,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)
|
||||
@@ -372,14 +371,13 @@ class EntityResolver:
|
||||
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,
|
||||
@@ -413,7 +411,10 @@ class EntityResolver:
|
||||
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
|
||||
|
||||
@@ -434,7 +435,8 @@ class EntityResolver:
|
||||
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
|
||||
@@ -444,10 +446,11 @@ class EntityResolver:
|
||||
FROM 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:
|
||||
@@ -477,10 +480,11 @@ class EntityResolver:
|
||||
cooccurrence_count = 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,7 +503,7 @@ 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(
|
||||
"""
|
||||
@@ -507,7 +511,7 @@ class EntityResolver:
|
||||
VALUES ($1, $2)
|
||||
ON CONFLICT DO NOTHING
|
||||
""",
|
||||
unit_entity_pairs
|
||||
unit_entity_pairs,
|
||||
)
|
||||
|
||||
# Build map of unit -> entities for co-occurrence calculation
|
||||
@@ -524,7 +528,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,7 +539,7 @@ 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)
|
||||
@@ -545,10 +549,10 @@ class EntityResolver:
|
||||
cooccurrence_count = 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.
|
||||
|
||||
@@ -568,15 +572,16 @@ class EntityResolver:
|
||||
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).
|
||||
|
||||
@@ -596,7 +601,8 @@ class EntityResolver:
|
||||
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
|
||||
|
||||
@@ -1,15 +1,17 @@
|
||||
"""
|
||||
LLM wrapper for unified configuration across providers.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import asyncio
|
||||
from typing import Optional, Any, Dict, List
|
||||
from openai import AsyncOpenAI, RateLimitError, APIError, APIStatusError, APIConnectionError, LengthFinishReasonError
|
||||
from typing import Any
|
||||
|
||||
from google import genai
|
||||
from google.genai import types as genai_types
|
||||
from google.genai import errors as genai_errors
|
||||
import logging
|
||||
from google.genai import types as genai_types
|
||||
from openai import APIConnectionError, APIStatusError, AsyncOpenAI, LengthFinishReasonError
|
||||
|
||||
# Seed applied to every Groq request for deterministic behavior.
|
||||
DEFAULT_LLM_SEED = 4242
|
||||
@@ -31,6 +33,7 @@ class OutputTooLongError(Exception):
|
||||
to allow callers to handle output length issues without depending on
|
||||
provider-specific implementations.
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
@@ -68,9 +71,7 @@ class LLMProvider:
|
||||
# Validate provider
|
||||
valid_providers = ["openai", "groq", "ollama", "gemini"]
|
||||
if self.provider not in valid_providers:
|
||||
raise ValueError(
|
||||
f"Invalid LLM provider: {self.provider}. Must be one of: {', '.join(valid_providers)}"
|
||||
)
|
||||
raise ValueError(f"Invalid LLM provider: {self.provider}. Must be one of: {', '.join(valid_providers)}")
|
||||
|
||||
# Set default base URLs
|
||||
if not self.base_url:
|
||||
@@ -106,7 +107,9 @@ class LLMProvider:
|
||||
RuntimeError: If the connection test fails.
|
||||
"""
|
||||
try:
|
||||
logger.info(f"Verifying LLM: provider={self.provider}, model={self.model}, base_url={self.base_url or 'default'}...")
|
||||
logger.info(
|
||||
f"Verifying LLM: provider={self.provider}, model={self.model}, base_url={self.base_url or 'default'}..."
|
||||
)
|
||||
await self.call(
|
||||
messages=[{"role": "user", "content": "Say 'ok'"}],
|
||||
max_completion_tokens=10,
|
||||
@@ -117,16 +120,14 @@ class LLMProvider:
|
||||
# If we get here without exception, the connection is working
|
||||
logger.info(f"LLM verified: {self.provider}/{self.model}")
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"LLM connection verification failed for {self.provider}/{self.model}: {e}"
|
||||
) from e
|
||||
raise RuntimeError(f"LLM connection verification failed for {self.provider}/{self.model}: {e}") from e
|
||||
|
||||
async def call(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
response_format: Optional[Any] = None,
|
||||
max_completion_tokens: Optional[int] = None,
|
||||
temperature: Optional[float] = None,
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any | None = None,
|
||||
max_completion_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
scope: str = "memory",
|
||||
max_retries: int = 10,
|
||||
initial_backoff: float = 1.0,
|
||||
@@ -161,8 +162,7 @@ class LLMProvider:
|
||||
# Handle Gemini provider separately
|
||||
if self.provider == "gemini":
|
||||
return await self._call_gemini(
|
||||
messages, response_format, max_retries, initial_backoff,
|
||||
max_backoff, skip_validation, start_time
|
||||
messages, response_format, max_retries, initial_backoff, max_backoff, skip_validation, start_time
|
||||
)
|
||||
|
||||
call_params = {
|
||||
@@ -172,7 +172,7 @@ class LLMProvider:
|
||||
|
||||
# Check if model supports reasoning parameter (o1, o3, gpt-5 families)
|
||||
model_lower = self.model.lower()
|
||||
is_reasoning_model = any(x in model_lower for x in ["gpt-5", "o1", "o3"])
|
||||
is_reasoning_model = any(x in model_lower for x in ["gpt-5", "o1", "o3", "deepseek"])
|
||||
|
||||
# For GPT-4 and GPT-4.1 models, cap max_completion_tokens to 32000
|
||||
# For GPT-4o models, cap to 16384
|
||||
@@ -194,7 +194,7 @@ class LLMProvider:
|
||||
call_params["temperature"] = temperature
|
||||
|
||||
# Set reasoning_effort for reasoning models (OpenAI gpt-5, o1, o3)
|
||||
if is_reasoning_model and self.provider == "openai":
|
||||
if is_reasoning_model:
|
||||
call_params["reasoning_effort"] = self.reasoning_effort
|
||||
|
||||
# Provider-specific parameters
|
||||
@@ -203,7 +203,6 @@ class LLMProvider:
|
||||
extra_body = {"service_tier": "auto"}
|
||||
# Only add reasoning parameters for reasoning models
|
||||
if is_reasoning_model:
|
||||
extra_body["reasoning_effort"] = self.reasoning_effort
|
||||
extra_body["include_reasoning"] = False
|
||||
call_params["extra_body"] = extra_body
|
||||
|
||||
@@ -213,16 +212,18 @@ class LLMProvider:
|
||||
try:
|
||||
if response_format is not None:
|
||||
# Add schema to system message for JSON mode
|
||||
if hasattr(response_format, 'model_json_schema'):
|
||||
if hasattr(response_format, "model_json_schema"):
|
||||
schema = response_format.model_json_schema()
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2)}"
|
||||
|
||||
if call_params['messages'] and call_params['messages'][0].get('role') == 'system':
|
||||
call_params['messages'][0]['content'] += schema_msg
|
||||
elif call_params['messages']:
|
||||
call_params['messages'][0]['content'] = schema_msg + "\n\n" + call_params['messages'][0]['content']
|
||||
if call_params["messages"] and call_params["messages"][0].get("role") == "system":
|
||||
call_params["messages"][0]["content"] += schema_msg
|
||||
elif call_params["messages"]:
|
||||
call_params["messages"][0]["content"] = (
|
||||
schema_msg + "\n\n" + call_params["messages"][0]["content"]
|
||||
)
|
||||
|
||||
call_params['response_format'] = {"type": "json_object"}
|
||||
call_params["response_format"] = {"type": "json_object"}
|
||||
response = await self._client.chat.completions.create(**call_params)
|
||||
|
||||
content = response.choices[0].message.content
|
||||
@@ -242,8 +243,8 @@ class LLMProvider:
|
||||
if duration > 10.0:
|
||||
ratio = max(1, usage.completion_tokens) / usage.prompt_tokens
|
||||
cached_tokens = 0
|
||||
if hasattr(usage, 'prompt_tokens_details') and usage.prompt_tokens_details:
|
||||
cached_tokens = getattr(usage.prompt_tokens_details, 'cached_tokens', 0) or 0
|
||||
if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details:
|
||||
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
|
||||
cache_info = f", cached_tokens={cached_tokens}" if cached_tokens > 0 else ""
|
||||
logger.info(
|
||||
f"slow llm call: model={self.provider}/{self.model}, "
|
||||
@@ -256,15 +257,19 @@ class LLMProvider:
|
||||
except LengthFinishReasonError as e:
|
||||
logger.warning(f"LLM output exceeded token limits: {str(e)}")
|
||||
raise OutputTooLongError(
|
||||
f"LLM output exceeded token limits. Input may need to be split into smaller chunks."
|
||||
"LLM output exceeded token limits. Input may need to be split into smaller chunks."
|
||||
) from e
|
||||
|
||||
except APIConnectionError as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
status_code = getattr(e, 'status_code', None) or getattr(getattr(e, 'response', None), 'status_code', None)
|
||||
logger.warning(f"Connection error, retrying... (attempt {attempt + 1}/{max_retries + 1}) - status_code={status_code}, message={e}")
|
||||
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
|
||||
status_code = getattr(e, "status_code", None) or getattr(
|
||||
getattr(e, "response", None), "status_code", None
|
||||
)
|
||||
logger.warning(
|
||||
f"Connection error, retrying... (attempt {attempt + 1}/{max_retries + 1}) - status_code={status_code}, message={e}"
|
||||
)
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
@@ -279,7 +284,7 @@ class LLMProvider:
|
||||
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
jitter = backoff * 0.2 * (2 * (time.time() % 1) - 1)
|
||||
sleep_time = backoff + jitter
|
||||
await asyncio.sleep(sleep_time)
|
||||
@@ -293,12 +298,12 @@ class LLMProvider:
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError(f"LLM call failed after all retries with no exception captured")
|
||||
raise RuntimeError("LLM call failed after all retries with no exception captured")
|
||||
|
||||
async def _call_gemini(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
response_format: Optional[Any],
|
||||
messages: list[dict[str, str]],
|
||||
response_format: Any | None,
|
||||
max_retries: int,
|
||||
initial_backoff: float,
|
||||
max_backoff: float,
|
||||
@@ -313,27 +318,21 @@ class LLMProvider:
|
||||
gemini_contents = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get('role', 'user')
|
||||
content = msg.get('content', '')
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
if role == 'system':
|
||||
if role == "system":
|
||||
if system_instruction:
|
||||
system_instruction += "\n\n" + content
|
||||
else:
|
||||
system_instruction = content
|
||||
elif role == 'assistant':
|
||||
gemini_contents.append(genai_types.Content(
|
||||
role="model",
|
||||
parts=[genai_types.Part(text=content)]
|
||||
))
|
||||
elif role == "assistant":
|
||||
gemini_contents.append(genai_types.Content(role="model", parts=[genai_types.Part(text=content)]))
|
||||
else:
|
||||
gemini_contents.append(genai_types.Content(
|
||||
role="user",
|
||||
parts=[genai_types.Part(text=content)]
|
||||
))
|
||||
gemini_contents.append(genai_types.Content(role="user", parts=[genai_types.Part(text=content)]))
|
||||
|
||||
# Add JSON schema instruction if response_format is provided
|
||||
if response_format is not None and hasattr(response_format, 'model_json_schema'):
|
||||
if response_format is not None and hasattr(response_format, "model_json_schema"):
|
||||
schema = response_format.model_json_schema()
|
||||
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2)}"
|
||||
if system_instruction:
|
||||
@@ -344,10 +343,10 @@ class LLMProvider:
|
||||
# Build generation config
|
||||
config_kwargs = {}
|
||||
if system_instruction:
|
||||
config_kwargs['system_instruction'] = system_instruction
|
||||
config_kwargs["system_instruction"] = system_instruction
|
||||
if response_format is not None:
|
||||
config_kwargs['response_mime_type'] = 'application/json'
|
||||
config_kwargs['response_schema'] = response_format
|
||||
config_kwargs["response_mime_type"] = "application/json"
|
||||
config_kwargs["response_schema"] = response_format
|
||||
|
||||
generation_config = genai_types.GenerateContentConfig(**config_kwargs) if config_kwargs else None
|
||||
|
||||
@@ -366,14 +365,14 @@ class LLMProvider:
|
||||
# Handle empty response
|
||||
if content is None:
|
||||
block_reason = None
|
||||
if hasattr(response, 'candidates') and response.candidates:
|
||||
if hasattr(response, "candidates") and response.candidates:
|
||||
candidate = response.candidates[0]
|
||||
if hasattr(candidate, 'finish_reason'):
|
||||
if hasattr(candidate, "finish_reason"):
|
||||
block_reason = candidate.finish_reason
|
||||
|
||||
if attempt < max_retries:
|
||||
logger.warning(f"Gemini returned empty response (reason: {block_reason}), retrying...")
|
||||
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
@@ -390,7 +389,7 @@ class LLMProvider:
|
||||
|
||||
# Log slow calls
|
||||
duration = time.time() - start_time
|
||||
if duration > 10.0 and hasattr(response, 'usage_metadata') and response.usage_metadata:
|
||||
if duration > 10.0 and hasattr(response, "usage_metadata") and response.usage_metadata:
|
||||
usage = response.usage_metadata
|
||||
logger.info(
|
||||
f"slow llm call: model={self.provider}/{self.model}, "
|
||||
@@ -403,8 +402,8 @@ class LLMProvider:
|
||||
except json.JSONDecodeError as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
logger.warning(f"Gemini returned invalid JSON, retrying...")
|
||||
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
|
||||
logger.warning("Gemini returned invalid JSON, retrying...")
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
else:
|
||||
@@ -421,7 +420,7 @@ class LLMProvider:
|
||||
if e.code in (400, 429, 500, 502, 503, 504) or (e.code and e.code >= 500):
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
|
||||
backoff = min(initial_backoff * (2**attempt), max_backoff)
|
||||
jitter = backoff * 0.2 * (2 * (time.time() % 1) - 1)
|
||||
await asyncio.sleep(backoff + jitter)
|
||||
else:
|
||||
@@ -437,7 +436,7 @@ class LLMProvider:
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise RuntimeError(f"Gemini call failed after all retries")
|
||||
raise RuntimeError("Gemini call failed after all retries")
|
||||
|
||||
@classmethod
|
||||
def for_memory(cls) -> "LLMProvider":
|
||||
@@ -447,13 +446,7 @@ class LLMProvider:
|
||||
base_url = os.getenv("HINDSIGHT_API_LLM_BASE_URL", "")
|
||||
model = os.getenv("HINDSIGHT_API_LLM_MODEL", "openai/gpt-oss-120b")
|
||||
|
||||
return cls(
|
||||
provider=provider,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
reasoning_effort="low"
|
||||
)
|
||||
return cls(provider=provider, api_key=api_key, base_url=base_url, model=model, reasoning_effort="low")
|
||||
|
||||
@classmethod
|
||||
def for_answer_generation(cls) -> "LLMProvider":
|
||||
@@ -463,13 +456,7 @@ class LLMProvider:
|
||||
base_url = os.getenv("HINDSIGHT_API_ANSWER_LLM_BASE_URL", os.getenv("HINDSIGHT_API_LLM_BASE_URL", ""))
|
||||
model = os.getenv("HINDSIGHT_API_ANSWER_LLM_MODEL", os.getenv("HINDSIGHT_API_LLM_MODEL", "openai/gpt-oss-120b"))
|
||||
|
||||
return cls(
|
||||
provider=provider,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
reasoning_effort="high"
|
||||
)
|
||||
return cls(provider=provider, api_key=api_key, base_url=base_url, model=model, reasoning_effort="high")
|
||||
|
||||
@classmethod
|
||||
def for_judge(cls) -> "LLMProvider":
|
||||
@@ -479,13 +466,7 @@ class LLMProvider:
|
||||
base_url = os.getenv("HINDSIGHT_API_JUDGE_LLM_BASE_URL", os.getenv("HINDSIGHT_API_LLM_BASE_URL", ""))
|
||||
model = os.getenv("HINDSIGHT_API_JUDGE_LLM_MODEL", os.getenv("HINDSIGHT_API_LLM_MODEL", "openai/gpt-oss-120b"))
|
||||
|
||||
return cls(
|
||||
provider=provider,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
reasoning_effort="high"
|
||||
)
|
||||
return cls(provider=provider, api_key=api_key, base_url=base_url, model=model, reasoning_effort="high")
|
||||
|
||||
|
||||
# Backwards compatibility alias
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -126,9 +125,9 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
|
||||
|
||||
# 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 +136,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 +149,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 +245,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 +285,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 +306,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 +323,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 +331,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 +343,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 +407,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 +443,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 +456,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 +468,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 +494,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 +516,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
|
||||
|
||||
@@ -6,9 +6,9 @@ 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)
|
||||
VALID_RECALL_FACT_TYPES = frozenset(["world", "experience", "opinion"])
|
||||
@@ -23,17 +23,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 +38,44 @@ 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,
|
||||
}
|
||||
}
|
||||
})
|
||||
)
|
||||
|
||||
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)"
|
||||
)
|
||||
|
||||
|
||||
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 +88,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}'"
|
||||
)
|
||||
|
||||
|
||||
@@ -126,37 +125,35 @@ class ReflectResult(BaseModel):
|
||||
Contains the formulated answer, the facts it was based on (organized by type),
|
||||
and any new opinions that were formed during the reflection process.
|
||||
"""
|
||||
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"],
|
||||
}
|
||||
}
|
||||
})
|
||||
)
|
||||
|
||||
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(
|
||||
default_factory=list,
|
||||
description="List of newly formed opinions during reflection"
|
||||
)
|
||||
new_opinions: list[str] = Field(default_factory=list, description="List of newly formed opinions during reflection")
|
||||
|
||||
|
||||
class Opinion(BaseModel):
|
||||
@@ -166,12 +163,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 +181,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 +198,22 @@ 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"
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -5,8 +5,10 @@ bank profile utilities for disposition and background 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 ..response_models import DispositionTraits
|
||||
|
||||
@@ -21,6 +23,7 @@ DEFAULT_DISPOSITION = {
|
||||
|
||||
class BankProfile(TypedDict):
|
||||
"""Type for bank profile data."""
|
||||
|
||||
name: str
|
||||
disposition: DispositionTraits
|
||||
background: str
|
||||
@@ -28,6 +31,7 @@ class BankProfile(TypedDict):
|
||||
|
||||
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)")
|
||||
|
||||
@@ -51,7 +55,7 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
SELECT name, disposition, background
|
||||
FROM banks WHERE bank_id = $1
|
||||
""",
|
||||
bank_id
|
||||
bank_id,
|
||||
)
|
||||
|
||||
if row:
|
||||
@@ -61,9 +65,7 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
disposition_data = json.loads(disposition_data)
|
||||
|
||||
return BankProfile(
|
||||
name=row["name"],
|
||||
disposition=DispositionTraits(**disposition_data),
|
||||
background=row["background"]
|
||||
name=row["name"], disposition=DispositionTraits(**disposition_data), background=row["background"]
|
||||
)
|
||||
|
||||
# Bank doesn't exist, create with defaults
|
||||
@@ -76,21 +78,13 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
|
||||
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), background="")
|
||||
|
||||
|
||||
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.
|
||||
|
||||
@@ -111,17 +105,11 @@ async def update_bank_disposition(
|
||||
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 merge_bank_background(pool, llm_config, bank_id: str, new_info: str, update_disposition: bool = True) -> dict:
|
||||
"""
|
||||
Merge new background information with existing background using LLM.
|
||||
Normalizes to first person ("I") and resolves conflicts.
|
||||
@@ -142,12 +130,7 @@ async def merge_bank_background(
|
||||
current_background = profile["background"]
|
||||
|
||||
# Use LLM to merge backgrounds and optionally infer disposition
|
||||
result = await _llm_merge_background(
|
||||
llm_config,
|
||||
current_background,
|
||||
new_info,
|
||||
infer_disposition=update_disposition
|
||||
)
|
||||
result = await _llm_merge_background(llm_config, current_background, new_info, infer_disposition=update_disposition)
|
||||
|
||||
merged_background = result["background"]
|
||||
inferred_disposition = result.get("disposition")
|
||||
@@ -166,7 +149,7 @@ async def merge_bank_background(
|
||||
""",
|
||||
bank_id,
|
||||
merged_background,
|
||||
json.dumps(inferred_disposition)
|
||||
json.dumps(inferred_disposition),
|
||||
)
|
||||
else:
|
||||
# Update only background
|
||||
@@ -178,7 +161,7 @@ async def merge_bank_background(
|
||||
WHERE bank_id = $1
|
||||
""",
|
||||
bank_id,
|
||||
merged_background
|
||||
merged_background,
|
||||
)
|
||||
|
||||
response = {"background": merged_background}
|
||||
@@ -188,12 +171,7 @@ async def merge_bank_background(
|
||||
return response
|
||||
|
||||
|
||||
async def _llm_merge_background(
|
||||
llm_config,
|
||||
current: str,
|
||||
new_info: str,
|
||||
infer_disposition: bool = False
|
||||
) -> dict:
|
||||
async def _llm_merge_background(llm_config, current: str, new_info: str, infer_disposition: bool = False) -> dict:
|
||||
"""
|
||||
Use LLM to intelligently merge background information.
|
||||
Optionally infer Big Five disposition traits from the merged background.
|
||||
@@ -273,25 +251,19 @@ Merged background:"""
|
||||
response_format=BackgroundMergeResponse,
|
||||
scope="bank_background",
|
||||
temperature=0.3,
|
||||
max_completion_tokens=8192
|
||||
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()
|
||||
}
|
||||
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_background", temperature=0.3, max_completion_tokens=8192
|
||||
)
|
||||
|
||||
logger.info(f"LLM response for background merge (first 500 chars): {content[:500]}")
|
||||
@@ -310,7 +282,7 @@ Merged background:"""
|
||||
# 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)
|
||||
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))
|
||||
@@ -321,7 +293,9 @@ Merged background:"""
|
||||
# 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)
|
||||
json_match = re.search(
|
||||
r'\{[^{}]*"background"[^{}]*"disposition"[^{}]*\{[^{}]*\}[^{}]*\}', content, re.DOTALL
|
||||
)
|
||||
if json_match:
|
||||
try:
|
||||
result = json.loads(json_match.group())
|
||||
@@ -335,7 +309,7 @@ Merged background:"""
|
||||
# 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()
|
||||
"disposition": DEFAULT_DISPOSITION.copy(),
|
||||
}
|
||||
|
||||
# Validate disposition values
|
||||
@@ -401,13 +375,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,
|
||||
"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,
|
||||
}
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
@@ -3,20 +3,15 @@ Chunk storage for retain pipeline.
|
||||
|
||||
Handles storage of document chunks in the database.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import List, Dict, Optional
|
||||
|
||||
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.
|
||||
|
||||
@@ -55,16 +50,13 @@ async def store_chunks_batch(
|
||||
[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,24 +3,18 @@ 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__)
|
||||
|
||||
|
||||
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]:
|
||||
entity_resolver, conn, bank_id: str, unit_ids: list[str], facts: list[ProcessedFact], log_buffer: list[str] = None
|
||||
) -> list[EntityLink]:
|
||||
"""
|
||||
Process entities for all facts and create entity links.
|
||||
|
||||
@@ -53,8 +47,7 @@ async def process_entities_batch(
|
||||
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
|
||||
[{"text": entity.name, "type": "CONCEPT"} for entity in (fact.entities or [])] for fact in facts
|
||||
]
|
||||
|
||||
# Use existing link_utils function for entity processing
|
||||
@@ -67,16 +60,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.
|
||||
|
||||
|
||||
@@ -4,16 +4,17 @@ Fact extraction from text using LLM.
|
||||
Extracts semantic facts, entities, and temporal information from text.
|
||||
Uses the LLMConfig wrapper for all LLM calls.
|
||||
"""
|
||||
import logging
|
||||
import os
|
||||
import json
|
||||
import re
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from datetime import datetime, timedelta
|
||||
from typing import List, Dict, Optional, Literal
|
||||
from openai import AsyncOpenAI
|
||||
from pydantic import BaseModel, Field, field_validator, ConfigDict
|
||||
from ..llm_wrapper import OutputTooLongError, LLMConfig
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
|
||||
from ..llm_wrapper import LLMConfig, OutputTooLongError
|
||||
|
||||
|
||||
def _sanitize_text(text: str) -> str:
|
||||
@@ -31,11 +32,12 @@ def _sanitize_text(text: str) -> str:
|
||||
return text
|
||||
# Remove surrogate characters (U+D800 to U+DFFF) using regex
|
||||
# These are invalid in UTF-8 and cause encoding errors
|
||||
return re.sub(r'[\ud800-\udfff]', '', text)
|
||||
return re.sub(r"[\ud800-\udfff]", "", text)
|
||||
|
||||
|
||||
class Entity(BaseModel):
|
||||
"""An entity extracted from text."""
|
||||
|
||||
text: str = Field(
|
||||
description="The specific, named entity as it appears in the fact. Must be a proper noun or specific identifier."
|
||||
)
|
||||
@@ -48,42 +50,46 @@ class Fact(BaseModel):
|
||||
This is what fact_extraction returns and what the rest of the pipeline expects.
|
||||
Combined fact text format: "what | when | where | who | why"
|
||||
"""
|
||||
|
||||
# Required fields
|
||||
fact: str = Field(description="Combined fact text: what | when | where | who | why")
|
||||
fact_type: Literal["world", "experience", "opinion"] = Field(description="Perspective: world/experience/opinion")
|
||||
|
||||
# Optional temporal fields
|
||||
occurred_start: Optional[str] = None
|
||||
occurred_end: Optional[str] = None
|
||||
mentioned_at: Optional[str] = None
|
||||
occurred_start: str | None = None
|
||||
occurred_end: str | None = None
|
||||
mentioned_at: str | None = None
|
||||
|
||||
# Optional location field
|
||||
where: Optional[str] = Field(None, description="WHERE the fact occurred or is about (specific location, place, or area)")
|
||||
where: str | None = Field(
|
||||
None, description="WHERE the fact occurred or is about (specific location, place, or area)"
|
||||
)
|
||||
|
||||
# Optional structured data
|
||||
entities: Optional[List[Entity]] = None
|
||||
causal_relations: Optional[List['CausalRelation']] = None
|
||||
entities: list[Entity] | None = None
|
||||
causal_relations: list["CausalRelation"] | None = None
|
||||
|
||||
|
||||
class CausalRelation(BaseModel):
|
||||
"""Causal relationship between facts."""
|
||||
|
||||
target_fact_index: int = Field(
|
||||
description="Index of the related fact in the facts array (0-based). "
|
||||
"This creates a directed causal link to another fact in the extraction."
|
||||
"This creates a directed causal link to another fact in the extraction."
|
||||
)
|
||||
relation_type: Literal["causes", "caused_by", "enables", "prevents"] = Field(
|
||||
description="Type of causal relationship: "
|
||||
"'causes' = this fact directly causes the target fact, "
|
||||
"'caused_by' = this fact was caused by the target fact, "
|
||||
"'enables' = this fact enables/allows the target fact, "
|
||||
"'prevents' = this fact prevents/blocks the target fact"
|
||||
"'causes' = this fact directly causes the target fact, "
|
||||
"'caused_by' = this fact was caused by the target fact, "
|
||||
"'enables' = this fact enables/allows the target fact, "
|
||||
"'prevents' = this fact prevents/blocks the target fact"
|
||||
)
|
||||
strength: float = Field(
|
||||
description="Strength of causal relationship (0.0 to 1.0). "
|
||||
"1.0 = direct/strong causation, 0.5 = moderate, 0.3 = weak/indirect",
|
||||
"1.0 = direct/strong causation, 0.5 = moderate, 0.3 = weak/indirect",
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
default=1.0
|
||||
default=1.0,
|
||||
)
|
||||
|
||||
|
||||
@@ -92,9 +98,7 @@ class ExtractedFact(BaseModel):
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_mode="validation",
|
||||
json_schema_extra={
|
||||
"required": ["what", "when", "where", "who", "why", "fact_type"]
|
||||
}
|
||||
json_schema_extra={"required": ["what", "when", "where", "who", "why", "fact_type"]},
|
||||
)
|
||||
|
||||
# ==========================================================================
|
||||
@@ -103,43 +107,43 @@ class ExtractedFact(BaseModel):
|
||||
|
||||
what: str = Field(
|
||||
description="WHAT happened - COMPLETE, DETAILED description with ALL specifics. "
|
||||
"NEVER summarize or omit details. Include: exact actions, objects, quantities, specifics. "
|
||||
"BE VERBOSE - capture every detail that was mentioned. "
|
||||
"Example: 'Emily got married to Sarah at a rooftop garden ceremony with 50 guests attending and a live jazz band playing' "
|
||||
"NOT: 'A wedding happened' or 'Emily got married'"
|
||||
"NEVER summarize or omit details. Include: exact actions, objects, quantities, specifics. "
|
||||
"BE VERBOSE - capture every detail that was mentioned. "
|
||||
"Example: 'Emily got married to Sarah at a rooftop garden ceremony with 50 guests attending and a live jazz band playing' "
|
||||
"NOT: 'A wedding happened' or 'Emily got married'"
|
||||
)
|
||||
|
||||
when: str = Field(
|
||||
description="WHEN it happened - ALWAYS include temporal information if mentioned. "
|
||||
"Include: specific dates, times, durations, relative time references. "
|
||||
"Examples: 'on June 15th, 2024 at 3pm', 'last weekend', 'for the past 3 years', 'every morning at 6am'. "
|
||||
"Write 'N/A' ONLY if absolutely no temporal context exists. Prefer converting to absolute dates when possible."
|
||||
"Include: specific dates, times, durations, relative time references. "
|
||||
"Examples: 'on June 15th, 2024 at 3pm', 'last weekend', 'for the past 3 years', 'every morning at 6am'. "
|
||||
"Write 'N/A' ONLY if absolutely no temporal context exists. Prefer converting to absolute dates when possible."
|
||||
)
|
||||
|
||||
where: str = Field(
|
||||
description="WHERE it happened or is about - SPECIFIC locations, places, areas, regions if applicable. "
|
||||
"Include: cities, neighborhoods, venues, buildings, countries, specific addresses when mentioned. "
|
||||
"Examples: 'downtown San Francisco at a rooftop garden venue', 'at the user's home in Brooklyn', 'online via Zoom', 'Paris, France'. "
|
||||
"Write 'N/A' ONLY if absolutely no location context exists or if the fact is completely location-agnostic."
|
||||
"Include: cities, neighborhoods, venues, buildings, countries, specific addresses when mentioned. "
|
||||
"Examples: 'downtown San Francisco at a rooftop garden venue', 'at the user's home in Brooklyn', 'online via Zoom', 'Paris, France'. "
|
||||
"Write 'N/A' ONLY if absolutely no location context exists or if the fact is completely location-agnostic."
|
||||
)
|
||||
|
||||
who: str = Field(
|
||||
description="WHO is involved - ALL people/entities with FULL context and relationships. "
|
||||
"Include: names, roles, relationships to user, background details. "
|
||||
"Resolve coreferences (if 'my roommate' is later named 'Emily', write 'Emily, the user's college roommate'). "
|
||||
"BE DETAILED about relationships and roles. "
|
||||
"Example: 'Emily (user's college roommate from Stanford, now works at Google), Sarah (Emily's partner of 5 years, software engineer)' "
|
||||
"NOT: 'my friend' or 'Emily and Sarah'"
|
||||
"Include: names, roles, relationships to user, background details. "
|
||||
"Resolve coreferences (if 'my roommate' is later named 'Emily', write 'Emily, the user's college roommate'). "
|
||||
"BE DETAILED about relationships and roles. "
|
||||
"Example: 'Emily (user's college roommate from Stanford, now works at Google), Sarah (Emily's partner of 5 years, software engineer)' "
|
||||
"NOT: 'my friend' or 'Emily and Sarah'"
|
||||
)
|
||||
|
||||
why: str = Field(
|
||||
description="WHY it matters - ALL emotional, contextual, and motivational details. "
|
||||
"Include EVERYTHING: feelings, preferences, motivations, observations, context, background, significance. "
|
||||
"BE VERBOSE - capture all the nuance and meaning. "
|
||||
"FOR ASSISTANT FACTS: MUST include what the user asked/requested that led to this interaction! "
|
||||
"Example (world): 'The user felt thrilled and inspired, has always dreamed of an outdoor ceremony, mentioned wanting a similar garden venue, was particularly moved by the intimate atmosphere and personal vows' "
|
||||
"Example (assistant): 'User asked how to fix slow API performance with 1000+ concurrent users, expected 70-80% reduction in database load' "
|
||||
"NOT: 'User liked it' or 'To help user'"
|
||||
"Include EVERYTHING: feelings, preferences, motivations, observations, context, background, significance. "
|
||||
"BE VERBOSE - capture all the nuance and meaning. "
|
||||
"FOR ASSISTANT FACTS: MUST include what the user asked/requested that led to this interaction! "
|
||||
"Example (world): 'The user felt thrilled and inspired, has always dreamed of an outdoor ceremony, mentioned wanting a similar garden venue, was particularly moved by the intimate atmosphere and personal vows' "
|
||||
"Example (assistant): 'User asked how to fix slow API performance with 1000+ concurrent users, expected 70-80% reduction in database load' "
|
||||
"NOT: 'User liked it' or 'To help user'"
|
||||
)
|
||||
|
||||
# ==========================================================================
|
||||
@@ -148,17 +152,17 @@ class ExtractedFact(BaseModel):
|
||||
|
||||
fact_kind: str = Field(
|
||||
default="conversation",
|
||||
description="'event' = specific datable occurrence (set occurred dates), 'conversation' = general info (no occurred dates)"
|
||||
description="'event' = specific datable occurrence (set occurred dates), 'conversation' = general info (no occurred dates)",
|
||||
)
|
||||
|
||||
# Temporal fields - optional
|
||||
occurred_start: Optional[str] = Field(
|
||||
occurred_start: str | None = Field(
|
||||
default=None,
|
||||
description="WHEN the event happened (ISO timestamp). Only for fact_kind='event'. Leave null for conversations."
|
||||
description="WHEN the event happened (ISO timestamp). Only for fact_kind='event'. Leave null for conversations.",
|
||||
)
|
||||
occurred_end: Optional[str] = Field(
|
||||
occurred_end: str | None = Field(
|
||||
default=None,
|
||||
description="WHEN the event ended (ISO timestamp). Only for events with duration. Leave null for conversations."
|
||||
description="WHEN the event ended (ISO timestamp). Only for events with duration. Leave null for conversations.",
|
||||
)
|
||||
|
||||
# Classification (CRITICAL - required)
|
||||
@@ -168,16 +172,15 @@ class ExtractedFact(BaseModel):
|
||||
)
|
||||
|
||||
# Entities - extracted from fact content
|
||||
entities: Optional[List[Entity]] = Field(
|
||||
entities: list[Entity] | None = Field(
|
||||
default=None,
|
||||
description="Named entities, objects, AND abstract concepts from the fact. Include: people names, organizations, places, significant objects (e.g., 'coffee maker', 'car'), AND abstract concepts/themes (e.g., 'friendship', 'career growth', 'loss', 'celebration'). Extract anything that could help link related facts together."
|
||||
description="Named entities, objects, AND abstract concepts from the fact. Include: people names, organizations, places, significant objects (e.g., 'coffee maker', 'car'), AND abstract concepts/themes (e.g., 'friendship', 'career growth', 'loss', 'celebration'). Extract anything that could help link related facts together.",
|
||||
)
|
||||
causal_relations: Optional[List[CausalRelation]] = Field(
|
||||
default=None,
|
||||
description="Causal links to other facts. Can be null."
|
||||
causal_relations: list[CausalRelation] | None = Field(
|
||||
default=None, description="Causal links to other facts. Can be null."
|
||||
)
|
||||
|
||||
@field_validator('entities', mode='before')
|
||||
@field_validator("entities", mode="before")
|
||||
@classmethod
|
||||
def ensure_entities_list(cls, v):
|
||||
"""Ensure entities is always a list (convert None to empty list)."""
|
||||
@@ -185,7 +188,7 @@ class ExtractedFact(BaseModel):
|
||||
return []
|
||||
return v
|
||||
|
||||
@field_validator('causal_relations', mode='before')
|
||||
@field_validator("causal_relations", mode="before")
|
||||
@classmethod
|
||||
def ensure_causal_relations_list(cls, v):
|
||||
"""Ensure causal_relations is always a list (convert None to empty list)."""
|
||||
@@ -198,11 +201,11 @@ class ExtractedFact(BaseModel):
|
||||
parts = [self.what]
|
||||
|
||||
# Add 'who' if not N/A
|
||||
if self.who and self.who.upper() != 'N/A':
|
||||
if self.who and self.who.upper() != "N/A":
|
||||
parts.append(f"Involving: {self.who}")
|
||||
|
||||
# Add 'why' if not N/A
|
||||
if self.why and self.why.upper() != 'N/A':
|
||||
if self.why and self.why.upper() != "N/A":
|
||||
parts.append(self.why)
|
||||
|
||||
if len(parts) == 1:
|
||||
@@ -213,12 +216,11 @@ class ExtractedFact(BaseModel):
|
||||
|
||||
class FactExtractionResponse(BaseModel):
|
||||
"""Response containing all extracted facts."""
|
||||
facts: List[ExtractedFact] = Field(
|
||||
description="List of extracted factual statements"
|
||||
)
|
||||
|
||||
facts: list[ExtractedFact] = Field(description="List of extracted factual statements")
|
||||
|
||||
|
||||
def chunk_text(text: str, max_chars: int) -> List[str]:
|
||||
def chunk_text(text: str, max_chars: int) -> list[str]:
|
||||
"""
|
||||
Split text into chunks, preserving conversation structure when possible.
|
||||
|
||||
@@ -232,7 +234,6 @@ def chunk_text(text: str, max_chars: int) -> List[str]:
|
||||
Returns:
|
||||
List of text chunks, roughly under max_chars
|
||||
"""
|
||||
import json
|
||||
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
||||
|
||||
# If text is small enough, return as-is
|
||||
@@ -256,21 +257,21 @@ def chunk_text(text: str, max_chars: int) -> List[str]:
|
||||
is_separator_regex=False,
|
||||
separators=[
|
||||
"\n\n", # Paragraph breaks
|
||||
"\n", # Line breaks
|
||||
". ", # Sentence endings
|
||||
"! ", # Exclamations
|
||||
"? ", # Questions
|
||||
"; ", # Semicolons
|
||||
", ", # Commas
|
||||
" ", # Words
|
||||
"", # Characters (last resort)
|
||||
"\n", # Line breaks
|
||||
". ", # Sentence endings
|
||||
"! ", # Exclamations
|
||||
"? ", # Questions
|
||||
"; ", # Semicolons
|
||||
", ", # Commas
|
||||
" ", # Words
|
||||
"", # Characters (last resort)
|
||||
],
|
||||
)
|
||||
|
||||
return splitter.split_text(text)
|
||||
|
||||
|
||||
def _chunk_conversation(turns: List[dict], max_chars: int) -> List[str]:
|
||||
def _chunk_conversation(turns: list[dict], max_chars: int) -> list[str]:
|
||||
"""
|
||||
Chunk a conversation array at turn boundaries, preserving complete turns.
|
||||
|
||||
@@ -281,7 +282,6 @@ def _chunk_conversation(turns: List[dict], max_chars: int) -> List[str]:
|
||||
Returns:
|
||||
List of JSON-serialized chunks, each containing complete turns
|
||||
"""
|
||||
import json
|
||||
|
||||
chunks = []
|
||||
current_chunk = []
|
||||
@@ -315,10 +315,10 @@ async def _extract_facts_from_chunk(
|
||||
total_chunks: int,
|
||||
event_date: datetime,
|
||||
context: str,
|
||||
llm_config: 'LLMConfig',
|
||||
llm_config: "LLMConfig",
|
||||
agent_name: str = None,
|
||||
extract_opinions: bool = False
|
||||
) -> List[Dict[str, str]]:
|
||||
extract_opinions: bool = False,
|
||||
) -> list[dict[str, str]]:
|
||||
"""
|
||||
Extract facts from a single chunk (internal helper for parallel processing).
|
||||
|
||||
@@ -333,7 +333,9 @@ async def _extract_facts_from_chunk(
|
||||
# Opinion extraction uses a separate prompt (not this one)
|
||||
fact_types_instruction = "Extract ONLY 'opinion' type facts (formed opinions, beliefs, and perspectives). DO NOT extract 'world' or 'assistant' facts."
|
||||
else:
|
||||
fact_types_instruction = "Extract ONLY 'world' and 'assistant' type facts. DO NOT extract opinions - those are extracted separately."
|
||||
fact_types_instruction = (
|
||||
"Extract ONLY 'world' and 'assistant' type facts. DO NOT extract opinions - those are extracted separately."
|
||||
)
|
||||
|
||||
prompt = f"""Extract facts from text into structured format with FOUR required dimensions - BE EXTREMELY DETAILED.
|
||||
|
||||
@@ -534,10 +536,8 @@ WHAT TO EXTRACT vs SKIP
|
||||
✅ EXTRACT: User preferences (ALWAYS as separate facts!), feelings, plans, events, relationships, achievements
|
||||
❌ SKIP: Greetings, filler ("thanks", "cool"), purely structural statements"""
|
||||
|
||||
|
||||
|
||||
|
||||
import logging
|
||||
|
||||
from openai import BadRequestError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -548,11 +548,11 @@ WHAT TO EXTRACT vs SKIP
|
||||
|
||||
# Sanitize input text to prevent Unicode encoding errors (e.g., unpaired surrogates)
|
||||
sanitized_chunk = _sanitize_text(chunk)
|
||||
sanitized_context = _sanitize_text(context) if context else 'none'
|
||||
sanitized_context = _sanitize_text(context) if context else "none"
|
||||
|
||||
# Build user message with metadata and chunk content in a clear format
|
||||
# Format event_date with day of week for better temporal reasoning
|
||||
event_date_formatted = event_date.strftime('%A, %B %d, %Y') # e.g., "Monday, June 10, 2024"
|
||||
event_date_formatted = event_date.strftime("%A, %B %d, %Y") # e.g., "Monday, June 10, 2024"
|
||||
user_message = f"""Extract facts from the following text chunk.
|
||||
{memory_bank_context}
|
||||
|
||||
@@ -566,16 +566,7 @@ Text:
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
extraction_response_json = await llm_config.call(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": prompt
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": user_message
|
||||
}
|
||||
],
|
||||
messages=[{"role": "system", "content": prompt}, {"role": "user", "content": user_message}],
|
||||
response_format=FactExtractionResponse,
|
||||
scope="memory_extract_facts",
|
||||
temperature=0.1,
|
||||
@@ -601,7 +592,7 @@ Text:
|
||||
)
|
||||
return []
|
||||
|
||||
raw_facts = extraction_response_json.get('facts', [])
|
||||
raw_facts = extraction_response_json.get("facts", [])
|
||||
if not raw_facts:
|
||||
logger.debug(
|
||||
f"LLM response missing 'facts' field or returned empty list. "
|
||||
@@ -622,48 +613,48 @@ Text:
|
||||
# Helper to get non-empty value
|
||||
def get_value(field_name):
|
||||
value = llm_fact.get(field_name)
|
||||
if value and value != '' and value != [] and value != {} and str(value).upper() != 'N/A':
|
||||
if value and value != "" and value != [] and value != {} and str(value).upper() != "N/A":
|
||||
return value
|
||||
return None
|
||||
|
||||
# NEW FORMAT: what, when, who, why (all required)
|
||||
what = get_value('what')
|
||||
when = get_value('when')
|
||||
who = get_value('who')
|
||||
why = get_value('why')
|
||||
what = get_value("what")
|
||||
when = get_value("when")
|
||||
who = get_value("who")
|
||||
why = get_value("why")
|
||||
|
||||
# Fallback to old format if new fields not present
|
||||
if not what:
|
||||
what = get_value('factual_core')
|
||||
what = get_value("factual_core")
|
||||
if not what:
|
||||
logger.warning(f"Skipping fact {i}: missing 'what' field")
|
||||
continue
|
||||
|
||||
# Critical field: fact_type
|
||||
# LLM uses "assistant" but we convert to "experience" for storage
|
||||
fact_type = llm_fact.get('fact_type')
|
||||
fact_type = llm_fact.get("fact_type")
|
||||
|
||||
# Convert "assistant" → "experience" for storage
|
||||
if fact_type == 'assistant':
|
||||
fact_type = 'experience'
|
||||
if fact_type == "assistant":
|
||||
fact_type = "experience"
|
||||
|
||||
# Validate fact_type (after conversion)
|
||||
if fact_type not in ['world', 'experience', 'opinion']:
|
||||
if fact_type not in ["world", "experience", "opinion"]:
|
||||
# Try to fix common mistakes - check if they swapped fact_type and fact_kind
|
||||
fact_kind = llm_fact.get('fact_kind')
|
||||
if fact_kind == 'assistant':
|
||||
fact_type = 'experience'
|
||||
elif fact_kind in ['world', 'experience', 'opinion']:
|
||||
fact_kind = llm_fact.get("fact_kind")
|
||||
if fact_kind == "assistant":
|
||||
fact_type = "experience"
|
||||
elif fact_kind in ["world", "experience", "opinion"]:
|
||||
fact_type = fact_kind
|
||||
else:
|
||||
# Default to 'world' if we can't determine
|
||||
fact_type = 'world'
|
||||
fact_type = "world"
|
||||
logger.warning(f"Fact {i}: defaulting to fact_type='world'")
|
||||
|
||||
# Get fact_kind for temporal handling (but don't store it)
|
||||
fact_kind = llm_fact.get('fact_kind', 'conversation')
|
||||
if fact_kind not in ['conversation', 'event', 'other']:
|
||||
fact_kind = 'conversation'
|
||||
fact_kind = llm_fact.get("fact_kind", "conversation")
|
||||
if fact_kind not in ["conversation", "event", "other"]:
|
||||
fact_kind = "conversation"
|
||||
|
||||
# Build combined fact text from the 4 dimensions: what | when | who | why
|
||||
fact_data = {}
|
||||
@@ -682,20 +673,20 @@ Text:
|
||||
|
||||
# Add temporal fields
|
||||
# For events: occurred_start/occurred_end (when the event happened)
|
||||
if fact_kind == 'event':
|
||||
occurred_start = get_value('occurred_start')
|
||||
occurred_end = get_value('occurred_end')
|
||||
if fact_kind == "event":
|
||||
occurred_start = get_value("occurred_start")
|
||||
occurred_end = get_value("occurred_end")
|
||||
if occurred_start:
|
||||
fact_data['occurred_start'] = occurred_start
|
||||
fact_data["occurred_start"] = occurred_start
|
||||
# For point events: if occurred_end not set, default to occurred_start
|
||||
if occurred_end:
|
||||
fact_data['occurred_end'] = occurred_end
|
||||
fact_data["occurred_end"] = occurred_end
|
||||
else:
|
||||
fact_data['occurred_end'] = occurred_start
|
||||
fact_data["occurred_end"] = occurred_start
|
||||
|
||||
# Add entities if present (validate as Entity objects)
|
||||
# LLM sometimes returns strings instead of {"text": "..."} format
|
||||
entities = get_value('entities')
|
||||
entities = get_value("entities")
|
||||
if entities:
|
||||
# Validate and normalize each entity
|
||||
validated_entities = []
|
||||
@@ -703,38 +694,34 @@ Text:
|
||||
if isinstance(ent, str):
|
||||
# Normalize string to Entity object
|
||||
validated_entities.append(Entity(text=ent))
|
||||
elif isinstance(ent, dict) and 'text' in ent:
|
||||
elif isinstance(ent, dict) and "text" in ent:
|
||||
try:
|
||||
validated_entities.append(Entity.model_validate(ent))
|
||||
except Exception as e:
|
||||
logger.warning(f"Invalid entity {ent}: {e}")
|
||||
if validated_entities:
|
||||
fact_data['entities'] = validated_entities
|
||||
fact_data["entities"] = validated_entities
|
||||
|
||||
# Add causal relations if present (validate as CausalRelation objects)
|
||||
# Filter out invalid relations (missing required fields)
|
||||
causal_relations = get_value('causal_relations')
|
||||
causal_relations = get_value("causal_relations")
|
||||
if causal_relations:
|
||||
validated_relations = []
|
||||
for rel in causal_relations:
|
||||
if isinstance(rel, dict) and 'target_fact_index' in rel and 'relation_type' in rel:
|
||||
if isinstance(rel, dict) and "target_fact_index" in rel and "relation_type" in rel:
|
||||
try:
|
||||
validated_relations.append(CausalRelation.model_validate(rel))
|
||||
except Exception as e:
|
||||
logger.warning(f"Invalid causal relation {rel}: {e}")
|
||||
if validated_relations:
|
||||
fact_data['causal_relations'] = validated_relations
|
||||
fact_data["causal_relations"] = validated_relations
|
||||
|
||||
# Always set mentioned_at to the event_date (when the conversation/document occurred)
|
||||
fact_data['mentioned_at'] = event_date.isoformat()
|
||||
fact_data["mentioned_at"] = event_date.isoformat()
|
||||
|
||||
# Build Fact model instance
|
||||
try:
|
||||
fact = Fact(
|
||||
fact=combined_text,
|
||||
fact_type=fact_type,
|
||||
**fact_data
|
||||
)
|
||||
fact = Fact(fact=combined_text, fact_type=fact_type, **fact_data)
|
||||
chunk_facts.append(fact)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create Fact model for fact {i}: {e}")
|
||||
@@ -753,7 +740,9 @@ Text:
|
||||
except BadRequestError as e:
|
||||
last_error = e
|
||||
if "json_validate_failed" in str(e):
|
||||
logger.warning(f" [1.3.{chunk_index + 1}] Attempt {attempt + 1}/{max_retries} failed with JSON validation error: {e}")
|
||||
logger.warning(
|
||||
f" [1.3.{chunk_index + 1}] Attempt {attempt + 1}/{max_retries} failed with JSON validation error: {e}"
|
||||
)
|
||||
if attempt < max_retries - 1:
|
||||
logger.info(f" [1.3.{chunk_index + 1}] Retrying...")
|
||||
continue
|
||||
@@ -772,8 +761,8 @@ async def _extract_facts_with_auto_split(
|
||||
context: str,
|
||||
llm_config: LLMConfig,
|
||||
agent_name: str = None,
|
||||
extract_opinions: bool = False
|
||||
) -> List[Dict[str, str]]:
|
||||
extract_opinions: bool = False,
|
||||
) -> list[dict[str, str]]:
|
||||
"""
|
||||
Extract facts from a chunk with automatic splitting if output exceeds token limits.
|
||||
|
||||
@@ -794,6 +783,7 @@ async def _extract_facts_with_auto_split(
|
||||
List of fact dictionaries extracted from the chunk (possibly from sub-chunks)
|
||||
"""
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
try:
|
||||
@@ -806,9 +796,9 @@ async def _extract_facts_with_auto_split(
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions
|
||||
extract_opinions=extract_opinions,
|
||||
)
|
||||
except OutputTooLongError as e:
|
||||
except OutputTooLongError:
|
||||
# Output exceeded token limits - split the chunk in half and retry
|
||||
logger.warning(
|
||||
f"Output too long for chunk {chunk_index + 1}/{total_chunks} "
|
||||
@@ -824,7 +814,7 @@ async def _extract_facts_with_auto_split(
|
||||
search_start = max(0, mid_point - search_range)
|
||||
search_end = min(len(chunk), mid_point + search_range)
|
||||
|
||||
sentence_endings = ['. ', '! ', '? ', '\n\n']
|
||||
sentence_endings = [". ", "! ", "? ", "\n\n"]
|
||||
best_split = mid_point
|
||||
|
||||
for ending in sentence_endings:
|
||||
@@ -838,8 +828,7 @@ async def _extract_facts_with_auto_split(
|
||||
second_half = chunk[best_split:].strip()
|
||||
|
||||
logger.info(
|
||||
f"Split chunk {chunk_index + 1} into two sub-chunks: "
|
||||
f"{len(first_half)} chars and {len(second_half)} chars"
|
||||
f"Split chunk {chunk_index + 1} into two sub-chunks: {len(first_half)} chars and {len(second_half)} chars"
|
||||
)
|
||||
|
||||
# Process both halves recursively (in parallel)
|
||||
@@ -852,7 +841,7 @@ async def _extract_facts_with_auto_split(
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions
|
||||
extract_opinions=extract_opinions,
|
||||
),
|
||||
_extract_facts_with_auto_split(
|
||||
chunk=second_half,
|
||||
@@ -862,8 +851,8 @@ async def _extract_facts_with_auto_split(
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions
|
||||
)
|
||||
extract_opinions=extract_opinions,
|
||||
),
|
||||
]
|
||||
|
||||
sub_results = await asyncio.gather(*sub_tasks)
|
||||
@@ -873,9 +862,7 @@ async def _extract_facts_with_auto_split(
|
||||
for sub_result in sub_results:
|
||||
all_facts.extend(sub_result)
|
||||
|
||||
logger.info(
|
||||
f"Successfully extracted {len(all_facts)} facts from split chunk {chunk_index + 1}"
|
||||
)
|
||||
logger.info(f"Successfully extracted {len(all_facts)} facts from split chunk {chunk_index + 1}")
|
||||
|
||||
return all_facts
|
||||
|
||||
@@ -887,7 +874,7 @@ async def extract_facts_from_text(
|
||||
agent_name: str,
|
||||
context: str = "",
|
||||
extract_opinions: bool = False,
|
||||
) -> tuple[List[Fact], List[tuple[str, int]]]:
|
||||
) -> tuple[list[Fact], list[tuple[str, int]]]:
|
||||
"""
|
||||
Extract semantic facts from conversational or narrative text using LLM.
|
||||
|
||||
@@ -920,7 +907,7 @@ async def extract_facts_from_text(
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions
|
||||
extract_opinions=extract_opinions,
|
||||
)
|
||||
for i, chunk in enumerate(chunks)
|
||||
]
|
||||
@@ -938,8 +925,10 @@ async def extract_facts_from_text(
|
||||
# ============================================================================
|
||||
|
||||
# Import types for the orchestration layer (note: ExtractedFact here is different from the Pydantic model above)
|
||||
from .types import RetainContent, ExtractedFact as ExtractedFactType, ChunkMetadata, CausalRelation as CausalRelationType
|
||||
from typing import Tuple
|
||||
|
||||
from .types import CausalRelation as CausalRelationType
|
||||
from .types import ChunkMetadata, RetainContent
|
||||
from .types import ExtractedFact as ExtractedFactType
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -948,11 +937,8 @@ SECONDS_PER_FACT = 10
|
||||
|
||||
|
||||
async def extract_facts_from_contents(
|
||||
contents: List[RetainContent],
|
||||
llm_config,
|
||||
agent_name: str,
|
||||
extract_opinions: bool = False
|
||||
) -> Tuple[List[ExtractedFactType], List[ChunkMetadata]]:
|
||||
contents: list[RetainContent], llm_config, agent_name: str, extract_opinions: bool = False
|
||||
) -> tuple[list[ExtractedFactType], list[ChunkMetadata]]:
|
||||
"""
|
||||
Extract facts from multiple content items in parallel.
|
||||
|
||||
@@ -985,7 +971,7 @@ async def extract_facts_from_contents(
|
||||
context=item.context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions
|
||||
extract_opinions=extract_opinions,
|
||||
)
|
||||
fact_extraction_tasks.append(task)
|
||||
|
||||
@@ -993,8 +979,8 @@ async def extract_facts_from_contents(
|
||||
all_fact_results = await asyncio.gather(*fact_extraction_tasks)
|
||||
|
||||
# Step 3: Flatten and convert to typed objects
|
||||
extracted_facts: List[ExtractedFactType] = []
|
||||
chunks_metadata: List[ChunkMetadata] = []
|
||||
extracted_facts: list[ExtractedFactType] = []
|
||||
chunks_metadata: list[ChunkMetadata] = []
|
||||
|
||||
global_chunk_idx = 0
|
||||
global_fact_idx = 0
|
||||
@@ -1008,7 +994,7 @@ async def extract_facts_from_contents(
|
||||
chunk_text=chunk_text,
|
||||
fact_count=chunk_fact_count,
|
||||
content_index=content_index,
|
||||
chunk_index=global_chunk_idx
|
||||
chunk_index=global_chunk_idx,
|
||||
)
|
||||
chunks_metadata.append(chunk_metadata)
|
||||
global_chunk_idx += 1
|
||||
@@ -1029,18 +1015,21 @@ async def extract_facts_from_contents(
|
||||
fact_type=fact_from_llm.fact_type,
|
||||
entities=[e.text for e in (fact_from_llm.entities or [])],
|
||||
# occurred_start/end: from LLM only, leave None if not provided
|
||||
occurred_start=_parse_datetime(fact_from_llm.occurred_start) if fact_from_llm.occurred_start else None,
|
||||
occurred_end=_parse_datetime(fact_from_llm.occurred_end) if fact_from_llm.occurred_end else None,
|
||||
occurred_start=_parse_datetime(fact_from_llm.occurred_start)
|
||||
if fact_from_llm.occurred_start
|
||||
else None,
|
||||
occurred_end=_parse_datetime(fact_from_llm.occurred_end)
|
||||
if fact_from_llm.occurred_end
|
||||
else None,
|
||||
causal_relations=_convert_causal_relations(
|
||||
fact_from_llm.causal_relations or [],
|
||||
global_fact_idx
|
||||
fact_from_llm.causal_relations or [], global_fact_idx
|
||||
),
|
||||
content_index=content_index,
|
||||
chunk_index=chunk_global_idx,
|
||||
context=content.context,
|
||||
# mentioned_at: always the event_date (when the conversation/document occurred)
|
||||
mentioned_at=content.event_date,
|
||||
metadata=content.metadata
|
||||
metadata=content.metadata,
|
||||
)
|
||||
|
||||
extracted_facts.append(extracted_fact)
|
||||
@@ -1056,13 +1045,14 @@ async def extract_facts_from_contents(
|
||||
def _parse_datetime(date_str: str):
|
||||
"""Parse ISO datetime string."""
|
||||
from dateutil import parser as date_parser
|
||||
|
||||
try:
|
||||
return date_parser.isoparse(date_str)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> List[CausalRelationType]:
|
||||
def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> list[CausalRelationType]:
|
||||
"""
|
||||
Convert causal relations from LLM format to ExtractedFact format.
|
||||
|
||||
@@ -1073,13 +1063,13 @@ def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> List[C
|
||||
causal_relation = CausalRelationType(
|
||||
relation_type=rel.relation_type,
|
||||
target_fact_index=fact_start_idx + rel.target_fact_index,
|
||||
strength=rel.strength
|
||||
strength=rel.strength,
|
||||
)
|
||||
causal_relations.append(causal_relation)
|
||||
return causal_relations
|
||||
|
||||
|
||||
def _add_temporal_offsets(facts: List[ExtractedFactType], contents: List[RetainContent]) -> None:
|
||||
def _add_temporal_offsets(facts: list[ExtractedFactType], contents: list[RetainContent]) -> None:
|
||||
"""
|
||||
Add time offsets to preserve fact ordering within each content.
|
||||
|
||||
|
||||
@@ -3,10 +3,9 @@ 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 logging
|
||||
|
||||
from .types import ProcessedFact
|
||||
|
||||
@@ -14,11 +13,8 @@ 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.
|
||||
|
||||
@@ -62,7 +58,7 @@ async def insert_facts_batch(
|
||||
contexts.append(fact.context)
|
||||
fact_types.append(fact.fact_type)
|
||||
# confidence_score is only for opinion facts
|
||||
confidence_scores.append(1.0 if fact.fact_type == 'opinion' else None)
|
||||
confidence_scores.append(1.0 if fact.fact_type == "opinion" else None)
|
||||
access_counts.append(0) # Initial access count
|
||||
metadata_jsons.append(json.dumps(fact.metadata))
|
||||
chunk_ids.append(fact.chunk_id)
|
||||
@@ -93,10 +89,10 @@ async def insert_facts_batch(
|
||||
access_counts,
|
||||
metadata_jsons,
|
||||
chunk_ids,
|
||||
document_ids
|
||||
document_ids,
|
||||
)
|
||||
|
||||
unit_ids = [str(row['id']) for row in results]
|
||||
unit_ids = [str(row["id"]) for row in results]
|
||||
return unit_ids
|
||||
|
||||
|
||||
@@ -119,17 +115,12 @@ async def ensure_bank_exists(conn, bank_id: str) -> None:
|
||||
""",
|
||||
bank_id,
|
||||
'{"skepticism": 3, "literalism": 3, "empathy": 3}',
|
||||
""
|
||||
"",
|
||||
)
|
||||
|
||||
|
||||
async def handle_document_tracking(
|
||||
conn,
|
||||
bank_id: str,
|
||||
document_id: str,
|
||||
combined_content: str,
|
||||
is_first_batch: bool,
|
||||
retain_params: Optional[dict] = None
|
||||
conn, bank_id: str, document_id: str, combined_content: str, is_first_batch: bool, retain_params: dict | None = None
|
||||
) -> None:
|
||||
"""
|
||||
Handle document tracking in the database.
|
||||
@@ -150,10 +141,7 @@ async def handle_document_tracking(
|
||||
# Always delete old document first if it exists (cascades to units and links)
|
||||
# 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
|
||||
)
|
||||
await conn.fetchval("DELETE FROM 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(
|
||||
@@ -172,5 +160,5 @@ 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,
|
||||
)
|
||||
|
||||
@@ -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,10 +2,9 @@
|
||||
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 .types import EntityLink
|
||||
@@ -19,7 +18,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 +53,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 +100,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 +118,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 +126,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 +172,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 +200,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 +221,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 +230,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 +254,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,13 +280,14 @@ 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(
|
||||
"""
|
||||
@@ -267,25 +295,29 @@ async def extract_entities_batch_optimized(
|
||||
FROM 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 +331,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 +384,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.
|
||||
@@ -361,10 +418,13 @@ async def create_temporal_links_batch_per_fact(
|
||||
FROM 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
|
||||
@@ -383,9 +443,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,8 +471,8 @@ 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")
|
||||
|
||||
@@ -421,7 +484,7 @@ async def create_temporal_links_batch_per_fact(
|
||||
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,
|
||||
)
|
||||
_log(log_buffer, f" [7.4] Insert {len(links)} temporal links: {time_mod.time() - insert_start:.3f}s")
|
||||
|
||||
@@ -430,6 +493,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 +501,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,6 +529,7 @@ 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
|
||||
@@ -478,9 +543,12 @@ async def create_semantic_links_batch(
|
||||
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 +556,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 +606,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,9 +634,12 @@ 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()
|
||||
@@ -577,20 +649,23 @@ async def create_semantic_links_batch(
|
||||
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
|
||||
all_links,
|
||||
)
|
||||
_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 +681,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,21 +707,15 @@ 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")
|
||||
|
||||
@@ -665,8 +733,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 +762,7 @@ async def create_causal_links_batch(
|
||||
|
||||
try:
|
||||
import time as time_mod
|
||||
|
||||
create_start = time_mod.time()
|
||||
|
||||
# Build links list
|
||||
@@ -705,12 +774,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,7 +804,6 @@ 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:
|
||||
@@ -745,14 +813,16 @@ async def create_causal_links_batch(
|
||||
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 +830,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
|
||||
|
||||
@@ -3,15 +3,14 @@ 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 datetime import UTC, datetime
|
||||
|
||||
from ..search import observation_utils
|
||||
from . import embedding_utils
|
||||
from ..db_utils import acquire_with_retry
|
||||
from .types import EntityLink
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -19,12 +18,12 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
def utcnow():
|
||||
"""Get current UTC time."""
|
||||
return datetime.now(timezone.utc)
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
# Simple dataclass-like container for facts (avoid importing from memory_engine)
|
||||
class MemoryFactForObservation:
|
||||
def __init__(self, id: str, text: str, fact_type: str, context: str, occurred_start: Optional[str]):
|
||||
def __init__(self, id: str, text: str, fact_type: str, context: str, occurred_start: str | None):
|
||||
self.id = id
|
||||
self.text = text
|
||||
self.fact_type = fact_type
|
||||
@@ -33,12 +32,7 @@ class MemoryFactForObservation:
|
||||
|
||||
|
||||
async def regenerate_observations_batch(
|
||||
conn,
|
||||
embeddings_model,
|
||||
llm_config,
|
||||
bank_id: str,
|
||||
entity_links: List[EntityLink],
|
||||
log_buffer: List[str] = None
|
||||
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.
|
||||
@@ -61,7 +55,7 @@ async def regenerate_observations_batch(
|
||||
return
|
||||
|
||||
# Count mentions per entity in this batch
|
||||
entity_mention_counts: Dict[str, int] = {}
|
||||
entity_mention_counts: dict[str, int] = {}
|
||||
for link in entity_links:
|
||||
if link.entity_id:
|
||||
entity_id = str(link.entity_id)
|
||||
@@ -71,11 +65,7 @@ async def regenerate_observations_batch(
|
||||
return
|
||||
|
||||
# Sort by mention count descending and take top N
|
||||
sorted_entities = sorted(
|
||||
entity_mention_counts.items(),
|
||||
key=lambda x: x[1],
|
||||
reverse=True
|
||||
)
|
||||
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()
|
||||
@@ -89,9 +79,10 @@ async def regenerate_observations_batch(
|
||||
SELECT id, canonical_name FROM entities
|
||||
WHERE id = ANY($1) AND bank_id = $2
|
||||
""",
|
||||
entity_uuids, bank_id
|
||||
entity_uuids,
|
||||
bank_id,
|
||||
)
|
||||
entity_names = {row['id']: row['canonical_name'] for row in entity_rows}
|
||||
entity_names = {row["id"]: row["canonical_name"] for row in entity_rows}
|
||||
|
||||
# Batch query for fact counts
|
||||
fact_counts = await conn.fetch(
|
||||
@@ -102,9 +93,10 @@ async def regenerate_observations_batch(
|
||||
WHERE ue.entity_id = ANY($1) AND mu.bank_id = $2
|
||||
GROUP BY ue.entity_id
|
||||
""",
|
||||
entity_uuids, bank_id
|
||||
entity_uuids,
|
||||
bank_id,
|
||||
)
|
||||
entity_fact_counts = {row['entity_id']: row['cnt'] for row in fact_counts}
|
||||
entity_fact_counts = {row["entity_id"]: row["cnt"] for row in fact_counts}
|
||||
|
||||
# Filter entities that meet the threshold
|
||||
entities_with_names = []
|
||||
@@ -126,8 +118,7 @@ async def regenerate_observations_batch(
|
||||
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
|
||||
conn, embeddings_model, llm_config, bank_id, entity_id, entity_name
|
||||
)
|
||||
total_observations += len(obs_ids)
|
||||
except Exception as e:
|
||||
@@ -135,17 +126,14 @@ async def regenerate_observations_batch(
|
||||
|
||||
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")
|
||||
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]:
|
||||
conn, embeddings_model, llm_config, bank_id: str, entity_id: str, entity_name: str
|
||||
) -> list[str]:
|
||||
"""
|
||||
Regenerate observations for a single entity.
|
||||
|
||||
@@ -176,7 +164,8 @@ async def _regenerate_entity_observations(
|
||||
ORDER BY mu.occurred_start DESC
|
||||
LIMIT 50
|
||||
""",
|
||||
bank_id, entity_uuid
|
||||
bank_id,
|
||||
entity_uuid,
|
||||
)
|
||||
|
||||
if not rows:
|
||||
@@ -185,21 +174,19 @@ async def _regenerate_entity_observations(
|
||||
# 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
|
||||
))
|
||||
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
|
||||
)
|
||||
observations = await observation_utils.extract_observations_from_facts(llm_config, entity_name, facts)
|
||||
|
||||
if not observations:
|
||||
return []
|
||||
@@ -217,13 +204,12 @@ async def _regenerate_entity_observations(
|
||||
AND ue.entity_id = $2
|
||||
)
|
||||
""",
|
||||
bank_id, entity_uuid
|
||||
bank_id,
|
||||
entity_uuid,
|
||||
)
|
||||
|
||||
# Generate embeddings for new observations
|
||||
embeddings = await embedding_utils.generate_embeddings_batch(
|
||||
embeddings_model, observations
|
||||
)
|
||||
embeddings = await embedding_utils.generate_embeddings_batch(embeddings_model, observations)
|
||||
|
||||
# Insert new observations
|
||||
current_time = utcnow()
|
||||
@@ -247,9 +233,9 @@ async def _regenerate_entity_observations(
|
||||
current_time,
|
||||
current_time,
|
||||
current_time,
|
||||
current_time
|
||||
current_time,
|
||||
)
|
||||
obs_id = str(result['id'])
|
||||
obs_id = str(result["id"])
|
||||
created_ids.append(obs_id)
|
||||
|
||||
# Link observation to entity
|
||||
@@ -258,7 +244,8 @@ async def _regenerate_entity_observations(
|
||||
INSERT INTO unit_entities (unit_id, entity_id)
|
||||
VALUES ($1, $2)
|
||||
""",
|
||||
uuid.UUID(obs_id), entity_uuid
|
||||
uuid.UUID(obs_id),
|
||||
entity_uuid,
|
||||
)
|
||||
|
||||
return created_ids
|
||||
|
||||
@@ -3,31 +3,33 @@ 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 typing import Any
|
||||
|
||||
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 . import (
|
||||
fact_extraction,
|
||||
embedding_processing,
|
||||
deduplication,
|
||||
chunk_storage,
|
||||
fact_storage,
|
||||
deduplication,
|
||||
embedding_processing,
|
||||
entity_processing,
|
||||
fact_extraction,
|
||||
fact_storage,
|
||||
link_creation,
|
||||
observation_regeneration
|
||||
observation_regeneration,
|
||||
)
|
||||
from .types import ExtractedFact, ProcessedFact, RetainContent
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -41,12 +43,12 @@ async def retain_batch(
|
||||
format_date_fn,
|
||||
duplicate_checker_fn,
|
||||
bank_id: str,
|
||||
contents_dicts: List[Dict[str, Any]],
|
||||
document_id: Optional[str] = None,
|
||||
contents_dicts: list[dict[str, Any]],
|
||||
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,
|
||||
) -> list[list[str]]:
|
||||
"""
|
||||
Process a batch of content through the retain pipeline.
|
||||
|
||||
@@ -73,10 +75,10 @@ async def retain_batch(
|
||||
|
||||
# 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)
|
||||
@@ -89,23 +91,26 @@ async def retain_batch(
|
||||
content=item["content"],
|
||||
context=item.get("context", ""),
|
||||
event_date=item.get("event_date") or utcnow(),
|
||||
metadata=item.get("metadata", {})
|
||||
metadata=item.get("metadata", {}),
|
||||
)
|
||||
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
|
||||
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:
|
||||
total_time = time.time() - start_time
|
||||
logger.info(
|
||||
f"RETAIN_BATCH COMPLETE: 0 facts extracted from {len(contents)} contents in {total_time:.3f}s (nothing to store)"
|
||||
)
|
||||
return [[] for _ in contents]
|
||||
|
||||
# Apply fact_type_override if provided
|
||||
@@ -130,6 +135,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,7 +161,11 @@ 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"]
|
||||
|
||||
@@ -195,7 +205,11 @@ 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"]
|
||||
|
||||
@@ -205,7 +219,9 @@ async def retain_batch(
|
||||
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 +246,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,7 +283,9 @@ 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)
|
||||
@@ -293,14 +313,18 @@ 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()
|
||||
@@ -309,34 +333,22 @@ async def retain_batch(
|
||||
|
||||
# Regenerate observations INSIDE transaction for atomicity
|
||||
await observation_regeneration.regenerate_observations_batch(
|
||||
conn,
|
||||
embeddings_model,
|
||||
llm_config,
|
||||
bank_id,
|
||||
entity_links,
|
||||
log_buffer
|
||||
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
|
||||
)
|
||||
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
|
||||
)
|
||||
await _trigger_background_tasks(task_backend, bank_id, unit_ids, non_duplicate_facts)
|
||||
|
||||
# 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")
|
||||
|
||||
@@ -344,11 +356,11 @@ async def retain_batch(
|
||||
|
||||
|
||||
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.
|
||||
|
||||
@@ -376,17 +388,19 @@ def _map_results_to_contents(
|
||||
async def _trigger_background_tasks(
|
||||
task_backend,
|
||||
bank_id: str,
|
||||
unit_ids: List[str],
|
||||
facts: List[ProcessedFact],
|
||||
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
|
||||
})
|
||||
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,8 +6,7 @@ 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 uuid import UUID
|
||||
|
||||
|
||||
@@ -18,16 +17,18 @@ 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)
|
||||
event_date: datetime | None = 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)
|
||||
from datetime import datetime
|
||||
|
||||
self.event_date = datetime.now(UTC)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -37,6 +38,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 +52,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 +65,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 +78,21 @@ 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)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -97,37 +102,38 @@ 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
|
||||
|
||||
@property
|
||||
def is_duplicate(self) -> bool:
|
||||
@@ -136,10 +142,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 +155,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 +176,7 @@ class ProcessedFact:
|
||||
metadata=extracted_fact.metadata,
|
||||
entities=entities,
|
||||
causal_relations=extracted_fact.causal_relations,
|
||||
chunk_id=chunk_id
|
||||
chunk_id=chunk_id,
|
||||
)
|
||||
|
||||
|
||||
@@ -183,10 +187,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 +202,25 @@ 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
|
||||
|
||||
# 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,11 @@ 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 .types import RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -40,10 +38,10 @@ 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,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
Retrieve relevant facts via graph traversal.
|
||||
|
||||
@@ -109,10 +107,10 @@ 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,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
Retrieve facts using BFS spreading activation.
|
||||
|
||||
@@ -127,9 +125,7 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
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
|
||||
)
|
||||
return await self._retrieve_with_conn(conn, query_embedding_str, bank_id, fact_type, budget)
|
||||
|
||||
async def _retrieve_with_conn(
|
||||
self,
|
||||
@@ -138,7 +134,7 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
) -> List[RetrievalResult]:
|
||||
) -> list[RetrievalResult]:
|
||||
"""Internal implementation with connection."""
|
||||
|
||||
# Step 1: Find entry points
|
||||
@@ -155,8 +151,11 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str, bank_id, fact_type,
|
||||
self.entry_point_threshold, self.entry_point_limit
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
self.entry_point_threshold,
|
||||
self.entry_point_limit,
|
||||
)
|
||||
|
||||
if not entry_points:
|
||||
@@ -165,10 +164,7 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
# 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:
|
||||
@@ -205,7 +201,10 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
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:
|
||||
|
||||
@@ -16,13 +16,12 @@ 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 .graph_retrieval import GraphRetriever
|
||||
from .types import RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -31,9 +30,11 @@ logger = logging.getLogger(__name__)
|
||||
# Data Classes
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class EdgeTarget:
|
||||
"""A neighbor node with its edge weight."""
|
||||
|
||||
node_id: str
|
||||
weight: float
|
||||
|
||||
@@ -41,19 +42,15 @@ class EdgeTarget:
|
||||
@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)
|
||||
|
||||
def get_neighbors(self, edge_type: str, node_id: str) -> List[EdgeTarget]:
|
||||
# edge_type -> from_node_id -> list of (to_node_id, weight)
|
||||
graphs: dict[str, dict[str, list[EdgeTarget]]] = field(default_factory=dict)
|
||||
|
||||
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,45 +60,49 @@ 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]
|
||||
|
||||
|
||||
@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)
|
||||
|
||||
@@ -110,9 +111,10 @@ class SeedNode:
|
||||
# Core Algorithm
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def mpfp_traverse(
|
||||
seeds: List[SeedNode],
|
||||
pattern: List[str],
|
||||
seeds: list[SeedNode],
|
||||
pattern: list[str],
|
||||
adjacency: TypedAdjacency,
|
||||
config: MPFPConfig,
|
||||
) -> PatternResult:
|
||||
@@ -131,20 +133,18 @@ def mpfp_traverse(
|
||||
if not seeds:
|
||||
return PatternResult(pattern=pattern, scores={})
|
||||
|
||||
scores: Dict[str, float] = {}
|
||||
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
|
||||
}
|
||||
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] = {}
|
||||
next_frontier: dict[str, float] = {}
|
||||
|
||||
for node_id, mass in frontier.items():
|
||||
if mass < config.threshold:
|
||||
@@ -155,15 +155,10 @@ def mpfp_traverse(
|
||||
|
||||
# Push (1-α) to neighbors
|
||||
push_mass = (1 - config.alpha) * mass
|
||||
neighbors = adjacency.get_normalized_neighbors(
|
||||
edge_type, node_id, config.top_k_neighbors
|
||||
)
|
||||
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
|
||||
)
|
||||
next_frontier[neighbor.node_id] = next_frontier.get(neighbor.node_id, 0) + push_mass * neighbor.weight
|
||||
|
||||
frontier = next_frontier
|
||||
|
||||
@@ -176,10 +171,10 @@ def mpfp_traverse(
|
||||
|
||||
|
||||
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 +186,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,6 +208,7 @@ 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.
|
||||
@@ -237,31 +225,27 @@ async def load_typed_adjacency(pool, bank_id: str) -> TypedAdjacency:
|
||||
AND ml.weight >= 0.1
|
||||
ORDER BY ml.from_unit_id, ml.weight DESC
|
||||
""",
|
||||
bank_id
|
||||
bank_id,
|
||||
)
|
||||
|
||||
graphs: Dict[str, Dict[str, List[EdgeTarget]]] = defaultdict(
|
||||
lambda: defaultdict(list)
|
||||
)
|
||||
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']
|
||||
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)
|
||||
)
|
||||
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 []
|
||||
@@ -276,7 +260,7 @@ async def fetch_memory_units_by_ids(
|
||||
AND fact_type = $2
|
||||
""",
|
||||
node_ids,
|
||||
fact_type
|
||||
fact_type,
|
||||
)
|
||||
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
|
||||
@@ -286,6 +270,7 @@ async def fetch_memory_units_by_ids(
|
||||
# Graph Retriever Implementation
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MPFPGraphRetriever(GraphRetriever):
|
||||
"""
|
||||
Graph retrieval using Meta-Path Forward Push.
|
||||
@@ -294,7 +279,7 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
then fuses results via RRF.
|
||||
"""
|
||||
|
||||
def __init__(self, config: Optional[MPFPConfig] = None):
|
||||
def __init__(self, config: MPFPConfig | None = None):
|
||||
"""
|
||||
Initialize MPFP retriever.
|
||||
|
||||
@@ -302,7 +287,7 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
config: Algorithm configuration (uses defaults if None)
|
||||
"""
|
||||
self.config = config or MPFPConfig()
|
||||
self._adjacency_cache: Dict[str, TypedAdjacency] = {}
|
||||
self._adjacency_cache: dict[str, TypedAdjacency] = {}
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
@@ -315,10 +300,10 @@ 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,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
Retrieve facts using MPFP algorithm.
|
||||
|
||||
@@ -339,14 +324,12 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
adjacency = await load_typed_adjacency(pool, bank_id)
|
||||
|
||||
# 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:
|
||||
semantic_seed_nodes = await self._find_semantic_seeds(
|
||||
pool, query_embedding_str, bank_id, fact_type
|
||||
)
|
||||
semantic_seed_nodes = await self._find_semantic_seeds(pool, query_embedding_str, bank_id, fact_type)
|
||||
|
||||
# Run all patterns in parallel
|
||||
tasks = []
|
||||
@@ -407,9 +390,9 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
|
||||
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,7 +414,7 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
fact_type: str,
|
||||
limit: int = 20,
|
||||
threshold: float = 0.3,
|
||||
) -> List[SeedNode]:
|
||||
) -> list[SeedNode]:
|
||||
"""Fallback: find semantic seeds via embedding search."""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
rows = await conn.fetch(
|
||||
@@ -445,10 +428,11 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str, bank_id, fact_type, threshold, limit
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
threshold,
|
||||
limit,
|
||||
)
|
||||
|
||||
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]
|
||||
|
||||
@@ -6,7 +6,7 @@ about an entity, without personality influence.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import List, Dict, Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..response_models import MemoryFact
|
||||
@@ -16,18 +16,17 @@ 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"
|
||||
)
|
||||
|
||||
observations: list[Observation] = Field(default_factory=list, description="List of observations about the entity")
|
||||
|
||||
|
||||
def format_facts_for_observation_prompt(facts: List[MemoryFact]) -> str:
|
||||
def format_facts_for_observation_prompt(facts: list[MemoryFact]) -> str:
|
||||
"""Format facts as text for observation extraction prompt."""
|
||||
import json
|
||||
|
||||
@@ -35,9 +34,7 @@ def format_facts_for_observation_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:
|
||||
@@ -92,11 +89,7 @@ def get_observation_system_message() -> str:
|
||||
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]:
|
||||
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.
|
||||
|
||||
@@ -118,10 +111,10 @@ async def extract_observations_from_facts(
|
||||
result = await llm_config.call(
|
||||
messages=[
|
||||
{"role": "system", "content": get_observation_system_message()},
|
||||
{"role": "user", "content": prompt}
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
response_format=ObservationExtractionResponse,
|
||||
scope="memory_extract_observation"
|
||||
scope="memory_extract_observation",
|
||||
)
|
||||
|
||||
observations = [op.observation for op in result.observations]
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
Cross-encoder neural reranking for search results.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from .types import MergedCandidate, ScoredResult
|
||||
|
||||
|
||||
@@ -24,14 +23,11 @@ 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
|
||||
|
||||
def rerank(
|
||||
self,
|
||||
query: str,
|
||||
candidates: List[MergedCandidate]
|
||||
) -> List[ScoredResult]:
|
||||
def rerank(self, query: str, candidates: list[MergedCandidate]) -> list[ScoredResult]:
|
||||
"""
|
||||
Rerank candidates using cross-encoder scores.
|
||||
|
||||
@@ -77,6 +73,7 @@ class CrossEncoderReranker:
|
||||
# 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 +86,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)
|
||||
|
||||
|
||||
@@ -8,16 +8,17 @@ Implements:
|
||||
4. Temporal retrieval (time-aware search with spreading)
|
||||
"""
|
||||
|
||||
from typing import List, Dict, Optional
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
import asyncio
|
||||
import logging
|
||||
from ..db_utils import acquire_with_retry
|
||||
from .types import RetrievalResult
|
||||
from .graph_retrieval import GraphRetriever, BFSGraphRetriever
|
||||
from .mpfp_retrieval import MPFPGraphRetriever
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from typing import Optional
|
||||
|
||||
from ...config import get_config
|
||||
from ..db_utils import acquire_with_retry
|
||||
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
|
||||
from .mpfp_retrieval import MPFPGraphRetriever
|
||||
from .types import RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -25,16 +26,17 @@ logger = logging.getLogger(__name__)
|
||||
@dataclass
|
||||
class ParallelRetrievalResult:
|
||||
"""Result from parallel retrieval across all methods."""
|
||||
semantic: List[RetrievalResult]
|
||||
bm25: List[RetrievalResult]
|
||||
graph: List[RetrievalResult]
|
||||
temporal: Optional[List[RetrievalResult]]
|
||||
timings: Dict[str, float] = field(default_factory=dict)
|
||||
temporal_constraint: Optional[tuple] = None # (start_date, end_date)
|
||||
|
||||
semantic: list[RetrievalResult]
|
||||
bm25: list[RetrievalResult]
|
||||
graph: list[RetrievalResult]
|
||||
temporal: list[RetrievalResult] | None
|
||||
timings: dict[str, float] = field(default_factory=dict)
|
||||
temporal_constraint: tuple | None = None # (start_date, end_date)
|
||||
|
||||
|
||||
# Default graph retriever instance (can be overridden)
|
||||
_default_graph_retriever: Optional[GraphRetriever] = None
|
||||
_default_graph_retriever: GraphRetriever | None = None
|
||||
|
||||
|
||||
def get_default_graph_retriever() -> GraphRetriever:
|
||||
@@ -62,12 +64,8 @@ def set_default_graph_retriever(retriever: GraphRetriever) -> None:
|
||||
|
||||
|
||||
async def retrieve_semantic(
|
||||
conn,
|
||||
query_emb_str: str,
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
limit: int
|
||||
) -> List[RetrievalResult]:
|
||||
conn, query_emb_str: str, bank_id: str, fact_type: str, limit: int
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
Semantic retrieval via vector similarity.
|
||||
|
||||
@@ -93,18 +91,15 @@ async def retrieve_semantic(
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $4
|
||||
""",
|
||||
query_emb_str, bank_id, fact_type, limit
|
||||
query_emb_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit,
|
||||
)
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in results]
|
||||
|
||||
|
||||
async def retrieve_bm25(
|
||||
conn,
|
||||
query_text: str,
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
limit: int
|
||||
) -> List[RetrievalResult]:
|
||||
async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, limit: int) -> list[RetrievalResult]:
|
||||
"""
|
||||
BM25 keyword retrieval via full-text search.
|
||||
|
||||
@@ -122,7 +117,7 @@ async def retrieve_bm25(
|
||||
|
||||
# Sanitize query text: remove special characters that have meaning in tsquery
|
||||
# Keep only alphanumeric characters and spaces
|
||||
sanitized_text = re.sub(r'[^\w\s]', ' ', query_text.lower())
|
||||
sanitized_text = re.sub(r"[^\w\s]", " ", query_text.lower())
|
||||
|
||||
# Split and filter empty strings
|
||||
tokens = [token for token in sanitized_text.split() if token]
|
||||
@@ -146,7 +141,10 @@ async def retrieve_bm25(
|
||||
ORDER BY bm25_score DESC
|
||||
LIMIT $4
|
||||
""",
|
||||
query_tsquery, bank_id, fact_type, limit
|
||||
query_tsquery,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit,
|
||||
)
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in results]
|
||||
|
||||
@@ -159,8 +157,8 @@ async def retrieve_temporal(
|
||||
start_date: datetime,
|
||||
end_date: datetime,
|
||||
budget: int,
|
||||
semantic_threshold: float = 0.1
|
||||
) -> List[RetrievalResult]:
|
||||
semantic_threshold: float = 0.1,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
Temporal retrieval with spreading activation.
|
||||
|
||||
@@ -182,13 +180,12 @@ async def retrieve_temporal(
|
||||
Returns:
|
||||
List of RetrievalResult objects with temporal scores
|
||||
"""
|
||||
from datetime import timezone
|
||||
|
||||
# Ensure start_date and end_date are timezone-aware (UTC) to match database datetimes
|
||||
if start_date.tzinfo is None:
|
||||
start_date = start_date.replace(tzinfo=timezone.utc)
|
||||
start_date = start_date.replace(tzinfo=UTC)
|
||||
if end_date.tzinfo is None:
|
||||
end_date = end_date.replace(tzinfo=timezone.utc)
|
||||
end_date = end_date.replace(tzinfo=UTC)
|
||||
|
||||
entry_points = await conn.fetch(
|
||||
"""
|
||||
@@ -215,7 +212,12 @@ async def retrieve_temporal(
|
||||
ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, (embedding <=> $1::vector) ASC
|
||||
LIMIT 10
|
||||
""",
|
||||
query_emb_str, bank_id, fact_type, start_date, end_date, semantic_threshold
|
||||
query_emb_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
start_date,
|
||||
end_date,
|
||||
semantic_threshold,
|
||||
)
|
||||
|
||||
if not entry_points:
|
||||
@@ -258,7 +260,9 @@ async def retrieve_temporal(
|
||||
results.append(ep_result)
|
||||
|
||||
# Spread through temporal links
|
||||
queue = [(RetrievalResult.from_db_row(dict(ep)), ep["similarity"], 1.0) for ep in entry_points] # (unit, semantic_sim, temporal_score)
|
||||
queue = [
|
||||
(RetrievalResult.from_db_row(dict(ep)), ep["similarity"], 1.0) for ep in entry_points
|
||||
] # (unit, semantic_sim, temporal_score)
|
||||
budget_remaining = budget - len(entry_points)
|
||||
|
||||
while queue and budget_remaining > 0:
|
||||
@@ -283,7 +287,10 @@ async def retrieve_temporal(
|
||||
ORDER BY ml.weight DESC
|
||||
LIMIT 10
|
||||
""",
|
||||
query_emb_str, current.id, fact_type, semantic_threshold
|
||||
query_emb_str,
|
||||
current.id,
|
||||
fact_type,
|
||||
semantic_threshold,
|
||||
)
|
||||
|
||||
for n in neighbors:
|
||||
@@ -307,7 +314,9 @@ async def retrieve_temporal(
|
||||
|
||||
if neighbor_best_date:
|
||||
days_from_mid = abs((neighbor_best_date - mid_date).total_seconds() / 86400)
|
||||
neighbor_temporal_proximity = 1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
|
||||
neighbor_temporal_proximity = (
|
||||
1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
|
||||
)
|
||||
else:
|
||||
neighbor_temporal_proximity = 0.3 # Lower score if no temporal data
|
||||
|
||||
@@ -349,9 +358,9 @@ async def retrieve_parallel(
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
thinking_budget: int,
|
||||
question_date: Optional[datetime] = None,
|
||||
question_date: datetime | None = None,
|
||||
query_analyzer: Optional["QueryAnalyzer"] = None,
|
||||
graph_retriever: Optional[GraphRetriever] = None,
|
||||
graph_retriever: GraphRetriever | None = None,
|
||||
) -> ParallelRetrievalResult:
|
||||
"""
|
||||
Run 3-way or 4-way parallel retrieval (adds temporal if detected).
|
||||
@@ -372,29 +381,26 @@ async def retrieve_parallel(
|
||||
"""
|
||||
from .temporal_extraction import extract_temporal_constraint
|
||||
|
||||
temporal_constraint = extract_temporal_constraint(
|
||||
query_text, reference_date=question_date, analyzer=query_analyzer
|
||||
)
|
||||
temporal_constraint = extract_temporal_constraint(query_text, reference_date=question_date, analyzer=query_analyzer)
|
||||
|
||||
retriever = graph_retriever or get_default_graph_retriever()
|
||||
|
||||
if retriever.name == "mpfp":
|
||||
return await _retrieve_parallel_mpfp(
|
||||
pool, query_text, query_embedding_str, bank_id, fact_type,
|
||||
thinking_budget, temporal_constraint, retriever
|
||||
pool, query_text, query_embedding_str, bank_id, fact_type, thinking_budget, temporal_constraint, retriever
|
||||
)
|
||||
else:
|
||||
return await _retrieve_parallel_bfs(
|
||||
pool, query_text, query_embedding_str, bank_id, fact_type,
|
||||
thinking_budget, temporal_constraint, retriever
|
||||
pool, query_text, query_embedding_str, bank_id, fact_type, thinking_budget, temporal_constraint, retriever
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _SemanticGraphResult:
|
||||
"""Internal result from semantic→graph chain."""
|
||||
semantic: List[RetrievalResult]
|
||||
graph: List[RetrievalResult]
|
||||
|
||||
semantic: list[RetrievalResult]
|
||||
graph: list[RetrievalResult]
|
||||
semantic_time: float
|
||||
graph_time: float
|
||||
|
||||
@@ -402,7 +408,8 @@ class _SemanticGraphResult:
|
||||
@dataclass
|
||||
class _TimedResult:
|
||||
"""Internal result with timing."""
|
||||
results: List[RetrievalResult]
|
||||
|
||||
results: list[RetrievalResult]
|
||||
time: float
|
||||
|
||||
|
||||
@@ -413,7 +420,7 @@ async def _retrieve_parallel_mpfp(
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
thinking_budget: int,
|
||||
temporal_constraint: Optional[tuple],
|
||||
temporal_constraint: tuple | None,
|
||||
retriever: GraphRetriever,
|
||||
) -> ParallelRetrievalResult:
|
||||
"""
|
||||
@@ -430,9 +437,7 @@ async def _retrieve_parallel_mpfp(
|
||||
"""Chain: semantic retrieval → graph retrieval (using semantic as seeds)."""
|
||||
start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
semantic = await retrieve_semantic(
|
||||
conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget
|
||||
)
|
||||
semantic = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
|
||||
semantic_time = time.time() - start
|
||||
|
||||
# Get temporal seeds if needed (quick query, part of this chain)
|
||||
@@ -441,8 +446,7 @@ async def _retrieve_parallel_mpfp(
|
||||
tc_start, tc_end = temporal_constraint
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
temporal_seeds = await _get_temporal_entry_points(
|
||||
conn, query_embedding_str, bank_id, fact_type,
|
||||
tc_start, tc_end, limit=20
|
||||
conn, query_embedding_str, bank_id, fact_type, tc_start, tc_end, limit=20
|
||||
)
|
||||
|
||||
# Run graph with seeds
|
||||
@@ -473,8 +477,14 @@ async def _retrieve_parallel_mpfp(
|
||||
start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
results = await retrieve_temporal(
|
||||
conn, query_embedding_str, bank_id, fact_type,
|
||||
tc_start, tc_end, budget=thinking_budget, semantic_threshold=0.1
|
||||
conn,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
tc_start,
|
||||
tc_end,
|
||||
budget=thinking_budget,
|
||||
semantic_threshold=0.1,
|
||||
)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
@@ -527,14 +537,13 @@ async def _get_temporal_entry_points(
|
||||
end_date: datetime,
|
||||
limit: int = 20,
|
||||
semantic_threshold: float = 0.1,
|
||||
) -> List[RetrievalResult]:
|
||||
) -> list[RetrievalResult]:
|
||||
"""Get temporal entry points (facts in date range with semantic relevance)."""
|
||||
from datetime import timezone
|
||||
|
||||
if start_date.tzinfo is None:
|
||||
start_date = start_date.replace(tzinfo=timezone.utc)
|
||||
start_date = start_date.replace(tzinfo=UTC)
|
||||
if end_date.tzinfo is None:
|
||||
end_date = end_date.replace(tzinfo=timezone.utc)
|
||||
end_date = end_date.replace(tzinfo=UTC)
|
||||
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
@@ -557,7 +566,13 @@ async def _get_temporal_entry_points(
|
||||
(embedding <=> $1::vector) ASC
|
||||
LIMIT $7
|
||||
""",
|
||||
query_embedding_str, bank_id, fact_type, start_date, end_date, semantic_threshold, limit
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
start_date,
|
||||
end_date,
|
||||
semantic_threshold,
|
||||
limit,
|
||||
)
|
||||
|
||||
results = []
|
||||
@@ -597,7 +612,7 @@ async def _retrieve_parallel_bfs(
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
thinking_budget: int,
|
||||
temporal_constraint: Optional[tuple],
|
||||
temporal_constraint: tuple | None,
|
||||
retriever: GraphRetriever,
|
||||
) -> ParallelRetrievalResult:
|
||||
"""BFS retrieval: all methods run in parallel (original behavior)."""
|
||||
@@ -631,8 +646,14 @@ async def _retrieve_parallel_bfs(
|
||||
start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
results = await retrieve_temporal(
|
||||
conn, query_embedding_str, bank_id, fact_type,
|
||||
tc_start, tc_end, budget=thinking_budget, semantic_threshold=0.1
|
||||
conn,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
tc_start,
|
||||
tc_end,
|
||||
budget=thinking_budget,
|
||||
semantic_threshold=0.1,
|
||||
)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
|
||||
@@ -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,6 +58,7 @@ 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
|
||||
@@ -79,6 +80,7 @@ def calculate_frequency_weight(access_count: int, max_boost: float = 2.0) -> flo
|
||||
Weight between 1.0 and max_boost
|
||||
"""
|
||||
import math
|
||||
|
||||
if access_count <= 0:
|
||||
return 1.0
|
||||
|
||||
@@ -116,11 +118,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.
|
||||
|
||||
|
||||
@@ -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,35 @@
|
||||
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 datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
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"
|
||||
|
||||
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 +41,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 +49,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 +57,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 +66,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 +74,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,7 +86,7 @@ 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)
|
||||
|
||||
@@ -176,16 +168,14 @@ 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."
|
||||
|
||||
|
||||
async def extract_opinions_from_text(
|
||||
llm_config,
|
||||
text: str,
|
||||
query: str
|
||||
) -> List[Opinion]:
|
||||
async def extract_opinions_from_text(llm_config, text: str, query: str) -> list[Opinion]:
|
||||
"""
|
||||
Extract opinions with reasons and confidence from text using LLM.
|
||||
|
||||
@@ -238,11 +228,14 @@ If no genuine opinions are expressed (e.g., the response just says "I don't know
|
||||
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}
|
||||
{
|
||||
"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"
|
||||
scope="memory_extract_opinion",
|
||||
)
|
||||
|
||||
# Format opinions with confidence score and convert to first-person
|
||||
@@ -253,14 +246,18 @@ If no genuine opinions are expressed (e.g., the response just says "I don't know
|
||||
|
||||
# Replace common third-person patterns with first-person
|
||||
def singularize_verb(verb):
|
||||
if verb.endswith('es'):
|
||||
if verb.endswith("es"):
|
||||
return verb[:-1] # believes -> believe
|
||||
elif verb.endswith('s'):
|
||||
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)
|
||||
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
|
||||
@@ -268,17 +265,96 @@ If no genuine opinions are expressed (e.g., the response just says "I don't know
|
||||
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"]
|
||||
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
|
||||
))
|
||||
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 []
|
||||
|
||||
|
||||
async def reflect(
|
||||
llm_config,
|
||||
query: str,
|
||||
experience_facts: list[str] = None,
|
||||
world_facts: list[str] = None,
|
||||
opinion_facts: list[str] = None,
|
||||
name: str = "Assistant",
|
||||
disposition: DispositionTraits = None,
|
||||
background: str = "",
|
||||
context: str = None,
|
||||
) -> str:
|
||||
"""
|
||||
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 provider instance
|
||||
query: Question to answer
|
||||
experience_facts: List of experience/agent fact strings
|
||||
world_facts: List of world fact strings
|
||||
opinion_facts: List of opinion fact strings
|
||||
name: Name of the agent/persona
|
||||
disposition: Disposition traits (defaults to neutral)
|
||||
background: Background information
|
||||
context: Additional context for the prompt
|
||||
|
||||
Returns:
|
||||
Generated answer text
|
||||
"""
|
||||
# Default disposition if not provided
|
||||
if disposition is None:
|
||||
disposition = DispositionTraits(skepticism=3, literalism=3, empathy=3)
|
||||
|
||||
# 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)]
|
||||
|
||||
agent_results = to_memory_facts(experience_facts or [], "experience")
|
||||
world_results = to_memory_facts(world_facts or [], "world")
|
||||
opinion_results = to_memory_facts(opinion_facts or [], "opinion")
|
||||
|
||||
# Format facts for prompt
|
||||
agent_facts_text = format_facts_for_prompt(agent_results)
|
||||
world_facts_text = format_facts_for_prompt(world_results)
|
||||
opinion_facts_text = format_facts_for_prompt(opinion_results)
|
||||
|
||||
# Build prompt
|
||||
prompt = build_think_prompt(
|
||||
agent_facts_text=agent_facts_text,
|
||||
world_facts_text=world_facts_text,
|
||||
opinion_facts_text=opinion_facts_text,
|
||||
query=query,
|
||||
name=name,
|
||||
disposition=disposition,
|
||||
background=background,
|
||||
context=context,
|
||||
)
|
||||
|
||||
system_message = get_system_message(disposition)
|
||||
|
||||
# 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,
|
||||
)
|
||||
|
||||
return answer_text.strip()
|
||||
|
||||
@@ -4,15 +4,18 @@ 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 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")
|
||||
@@ -20,6 +23,7 @@ class QueryInfo(BaseModel):
|
||||
|
||||
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 +32,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 +48,120 @@ 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")
|
||||
event_date: datetime | None = Field(default=None, description="When the memory occurred")
|
||||
access_count: int = Field(description="Number of times accessed before this search")
|
||||
|
||||
# How this node was reached
|
||||
is_entry_point: bool = Field(description="Whether this is an entry point")
|
||||
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 +176,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 +215,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 +236,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,25 @@ 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,
|
||||
WeightComponents,
|
||||
)
|
||||
|
||||
|
||||
@@ -58,17 +59,17 @@ class SearchTracer:
|
||||
self.max_tokens = max_tokens
|
||||
|
||||
# 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] = []
|
||||
|
||||
# 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,7 +84,7 @@ 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
|
||||
|
||||
@@ -117,9 +118,9 @@ class SearchTracer:
|
||||
event_date: datetime,
|
||||
access_count: int,
|
||||
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,
|
||||
@@ -199,10 +200,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 +267,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 +287,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 +332,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 +351,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 +374,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 +398,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,7 +425,7 @@ 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,
|
||||
)
|
||||
|
||||
@@ -6,8 +6,8 @@ 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
|
||||
@@ -17,28 +17,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
|
||||
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
|
||||
access_count: int = 0
|
||||
embedding: Optional[List[float]] = None
|
||||
embedding: list[float] | None = None
|
||||
|
||||
# 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"]),
|
||||
@@ -68,13 +69,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 +91,7 @@ class ScoredResult:
|
||||
|
||||
Contains all retrieval/merge data plus reranking scores and combined score.
|
||||
"""
|
||||
|
||||
# Original merged candidate
|
||||
candidate: MergedCandidate
|
||||
|
||||
@@ -115,7 +118,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.
|
||||
|
||||
|
||||
@@ -6,10 +6,12 @@ This provides an abstraction that can be adapted to different execution models:
|
||||
- Pub/Sub architectures (future)
|
||||
- Message brokers (future)
|
||||
"""
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, Optional, Callable, Awaitable
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -29,10 +31,10 @@ class TaskBackend(ABC):
|
||||
|
||||
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.
|
||||
|
||||
@@ -49,7 +51,7 @@ class TaskBackend(ABC):
|
||||
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.
|
||||
|
||||
@@ -65,7 +67,7 @@ class TaskBackend(ABC):
|
||||
"""
|
||||
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,16 +75,17 @@ 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()
|
||||
|
||||
|
||||
@@ -94,11 +97,7 @@ class AsyncIOQueueBackend(TaskBackend):
|
||||
and a periodic consumer worker.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
batch_size: int = 100,
|
||||
batch_interval: float = 1.0
|
||||
):
|
||||
def __init__(self, batch_size: int = 100, batch_interval: float = 1.0):
|
||||
"""
|
||||
Initialize AsyncIO queue backend.
|
||||
|
||||
@@ -107,9 +106,9 @@ class AsyncIOQueueBackend(TaskBackend):
|
||||
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._queue: asyncio.Queue | None = None
|
||||
self._worker_task: asyncio.Task | None = None
|
||||
self._shutdown_event: asyncio.Event | None = None
|
||||
self._batch_size = batch_size
|
||||
self._batch_interval = batch_interval
|
||||
|
||||
@@ -124,7 +123,7 @@ class AsyncIOQueueBackend(TaskBackend):
|
||||
self._initialized = True
|
||||
logger.info("AsyncIOQueueBackend 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.
|
||||
|
||||
@@ -135,8 +134,8 @@ class AsyncIOQueueBackend(TaskBackend):
|
||||
await self.initialize()
|
||||
|
||||
await self._queue.put(task_dict)
|
||||
task_type = task_dict.get('type', 'unknown')
|
||||
task_id = task_dict.get('id')
|
||||
task_type = task_dict.get("type", "unknown")
|
||||
task_id = task_dict.get("id")
|
||||
|
||||
async def wait_for_pending_tasks(self, timeout: float = 5.0):
|
||||
"""
|
||||
@@ -200,20 +199,16 @@ class AsyncIOQueueBackend(TaskBackend):
|
||||
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
|
||||
)
|
||||
task_dict = await asyncio.wait_for(self._queue.get(), timeout=remaining_time)
|
||||
tasks.append(task_dict)
|
||||
except asyncio.TimeoutError:
|
||||
except 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
|
||||
*[self._execute_task(task_dict) for task_dict in tasks], return_exceptions=True
|
||||
)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
"""
|
||||
Utility functions for memory system.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import List, Dict, TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .llm_wrapper import LLMConfig
|
||||
@@ -12,7 +13,14 @@ if TYPE_CHECKING:
|
||||
from .retain.fact_extraction import extract_facts_from_text
|
||||
|
||||
|
||||
async def extract_facts(text: str, event_date: datetime, context: str = "", llm_config: 'LLMConfig' = None, agent_name: str = None, extract_opinions: bool = False) -> tuple[List['Fact'], List[tuple[str, int]]]:
|
||||
async def extract_facts(
|
||||
text: str,
|
||||
event_date: datetime,
|
||||
context: str = "",
|
||||
llm_config: "LLMConfig" = None,
|
||||
agent_name: str = None,
|
||||
extract_opinions: bool = False,
|
||||
) -> tuple[list["Fact"], list[tuple[str, int]]]:
|
||||
"""
|
||||
Extract semantic facts from text using LLM.
|
||||
|
||||
@@ -41,16 +49,25 @@ async def extract_facts(text: str, event_date: datetime, context: str = "", llm_
|
||||
if not text or not text.strip():
|
||||
return [], []
|
||||
|
||||
facts, chunks = await extract_facts_from_text(text, event_date, context=context, llm_config=llm_config, agent_name=agent_name, extract_opinions=extract_opinions)
|
||||
facts, chunks = await extract_facts_from_text(
|
||||
text,
|
||||
event_date,
|
||||
context=context,
|
||||
llm_config=llm_config,
|
||||
agent_name=agent_name,
|
||||
extract_opinions=extract_opinions,
|
||||
)
|
||||
|
||||
if not facts:
|
||||
logging.warning(f"LLM extracted 0 facts from text of length {len(text)}. This may indicate the text contains no meaningful information, or the LLM failed to extract facts. Full text: {text}")
|
||||
logging.warning(
|
||||
f"LLM extracted 0 facts from text of length {len(text)}. This may indicate the text contains no meaningful information, or the LLM failed to extract facts. Full text: {text}"
|
||||
)
|
||||
return [], chunks
|
||||
|
||||
return facts, chunks
|
||||
|
||||
|
||||
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.
|
||||
|
||||
@@ -100,6 +117,7 @@ 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
|
||||
@@ -121,6 +139,7 @@ def calculate_frequency_weight(access_count: int, max_boost: float = 2.0) -> flo
|
||||
Weight between 1.0 and max_boost
|
||||
"""
|
||||
import math
|
||||
|
||||
if access_count <= 0:
|
||||
return 1.0
|
||||
|
||||
@@ -158,11 +177,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.
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ Run the server with:
|
||||
|
||||
Stop with Ctrl+C.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import atexit
|
||||
@@ -13,15 +14,14 @@ import os
|
||||
import signal
|
||||
import sys
|
||||
import warnings
|
||||
from typing import Optional
|
||||
|
||||
import uvicorn
|
||||
|
||||
from . import MemoryEngine
|
||||
from .api import create_app
|
||||
from .config import get_config, HindsightConfig
|
||||
|
||||
from .banner import print_banner
|
||||
from .config import HindsightConfig, get_config
|
||||
|
||||
print()
|
||||
print_banner()
|
||||
|
||||
@@ -33,7 +33,7 @@ warnings.filterwarnings("ignore", message="websockets.server.WebSocketServerProt
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
|
||||
# Global reference for cleanup
|
||||
_memory: Optional[MemoryEngine] = None
|
||||
_memory: MemoryEngine | None = None
|
||||
|
||||
|
||||
def _cleanup():
|
||||
@@ -70,59 +70,41 @@ def main():
|
||||
|
||||
# Server options
|
||||
parser.add_argument(
|
||||
"--host", default=config.host,
|
||||
help=f"Host to bind to (default: {config.host}, env: HINDSIGHT_API_HOST)"
|
||||
"--host", default=config.host, help=f"Host to bind to (default: {config.host}, env: HINDSIGHT_API_HOST)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--port", type=int, default=config.port,
|
||||
help=f"Port to bind to (default: {config.port}, env: HINDSIGHT_API_PORT)"
|
||||
"--port",
|
||||
type=int,
|
||||
default=config.port,
|
||||
help=f"Port to bind to (default: {config.port}, env: HINDSIGHT_API_PORT)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--log-level", default=config.log_level,
|
||||
"--log-level",
|
||||
default=config.log_level,
|
||||
choices=["critical", "error", "warning", "info", "debug", "trace"],
|
||||
help=f"Log level (default: {config.log_level}, env: HINDSIGHT_API_LOG_LEVEL)"
|
||||
help=f"Log level (default: {config.log_level}, env: HINDSIGHT_API_LOG_LEVEL)",
|
||||
)
|
||||
|
||||
# Development options
|
||||
parser.add_argument(
|
||||
"--reload", action="store_true",
|
||||
help="Enable auto-reload on code changes (development only)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--workers", type=int, default=1,
|
||||
help="Number of worker processes (default: 1)"
|
||||
)
|
||||
parser.add_argument("--reload", action="store_true", help="Enable auto-reload on code changes (development only)")
|
||||
parser.add_argument("--workers", type=int, default=1, help="Number of worker processes (default: 1)")
|
||||
|
||||
# Access log options
|
||||
parser.add_argument(
|
||||
"--access-log", action="store_true",
|
||||
help="Enable access log"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no-access-log", dest="access_log", action="store_false",
|
||||
help="Disable access log (default)"
|
||||
)
|
||||
parser.add_argument("--access-log", action="store_true", help="Enable access log")
|
||||
parser.add_argument("--no-access-log", dest="access_log", action="store_false", help="Disable access log (default)")
|
||||
parser.set_defaults(access_log=False)
|
||||
|
||||
# Proxy options
|
||||
parser.add_argument(
|
||||
"--proxy-headers", action="store_true",
|
||||
help="Enable X-Forwarded-Proto, X-Forwarded-For headers"
|
||||
"--proxy-headers", action="store_true", help="Enable X-Forwarded-Proto, X-Forwarded-For headers"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--forwarded-allow-ips", default=None,
|
||||
help="Comma separated list of IPs to trust with proxy headers"
|
||||
"--forwarded-allow-ips", default=None, help="Comma separated list of IPs to trust with proxy headers"
|
||||
)
|
||||
|
||||
# SSL options
|
||||
parser.add_argument(
|
||||
"--ssl-keyfile", default=None,
|
||||
help="SSL key file"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ssl-certfile", default=None,
|
||||
help="SSL certificate file"
|
||||
)
|
||||
parser.add_argument("--ssl-keyfile", default=None, help="SSL key file")
|
||||
parser.add_argument("--ssl-certfile", default=None, help="SSL certificate file")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -145,8 +127,10 @@ def main():
|
||||
port=args.port,
|
||||
log_level=args.log_level,
|
||||
mcp_enabled=config.mcp_enabled,
|
||||
graph_retriever=config.graph_retriever,
|
||||
)
|
||||
config.configure_logging()
|
||||
config.log_config()
|
||||
|
||||
# Register cleanup handlers
|
||||
atexit.register(_cleanup)
|
||||
@@ -188,9 +172,8 @@ def main():
|
||||
if args.ssl_certfile:
|
||||
uvicorn_config["ssl_certfile"] = args.ssl_certfile
|
||||
|
||||
|
||||
|
||||
from .banner import print_startup_info
|
||||
|
||||
print_startup_info(
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
"""
|
||||
Local MCP server for use with Claude Code (stdio transport).
|
||||
|
||||
This runs a fully local Hindsight instance with embedded PostgreSQL (pg0).
|
||||
No external database or server required.
|
||||
|
||||
Run with:
|
||||
hindsight-local-mcp
|
||||
|
||||
Or with uvx:
|
||||
uvx hindsight-api@latest hindsight-local-mcp
|
||||
|
||||
Configure in Claude Code's MCP settings:
|
||||
{
|
||||
"mcpServers": {
|
||||
"hindsight": {
|
||||
"command": "uvx",
|
||||
"args": ["hindsight-api@latest", "hindsight-local-mcp"],
|
||||
"env": {
|
||||
"HINDSIGHT_API_LLM_API_KEY": "your-openai-key"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Environment variables:
|
||||
HINDSIGHT_API_LLM_API_KEY: Required. API key for LLM provider.
|
||||
HINDSIGHT_API_LLM_PROVIDER: Optional. LLM provider (default: "openai").
|
||||
HINDSIGHT_API_LLM_MODEL: Optional. LLM model (default: "gpt-4o-mini").
|
||||
HINDSIGHT_API_MCP_LOCAL_BANK_ID: Optional. Memory bank ID (default: "mcp").
|
||||
HINDSIGHT_API_LOG_LEVEL: Optional. Log level (default: "warning").
|
||||
HINDSIGHT_API_MCP_INSTRUCTIONS: Optional. Additional instructions appended to both retain and recall tools.
|
||||
|
||||
Example custom instructions (these are ADDED to the default behavior):
|
||||
To also store assistant actions:
|
||||
HINDSIGHT_API_MCP_INSTRUCTIONS="Also store every action you take, including tool calls, code written, and decisions made."
|
||||
|
||||
To also store conversation summaries:
|
||||
HINDSIGHT_API_MCP_INSTRUCTIONS="Also store summaries of important conversations and their outcomes."
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from mcp.types import Icon
|
||||
|
||||
from hindsight_api.config import (
|
||||
DEFAULT_MCP_LOCAL_BANK_ID,
|
||||
DEFAULT_MCP_RECALL_DESCRIPTION,
|
||||
DEFAULT_MCP_RETAIN_DESCRIPTION,
|
||||
ENV_MCP_INSTRUCTIONS,
|
||||
ENV_MCP_LOCAL_BANK_ID,
|
||||
)
|
||||
|
||||
# Configure logging - default to warning to avoid polluting stderr during MCP init
|
||||
# MCP clients interpret stderr output as errors, so we suppress INFO logs by default
|
||||
_log_level_str = os.environ.get("HINDSIGHT_API_LOG_LEVEL", "warning").lower()
|
||||
_log_level_map = {
|
||||
"critical": logging.CRITICAL,
|
||||
"error": logging.ERROR,
|
||||
"warning": logging.WARNING,
|
||||
"info": logging.INFO,
|
||||
"debug": logging.DEBUG,
|
||||
}
|
||||
logging.basicConfig(
|
||||
level=_log_level_map.get(_log_level_str, logging.WARNING),
|
||||
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
|
||||
stream=sys.stderr, # MCP uses stdout for protocol, logs go to stderr
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def create_local_mcp_server(bank_id: str, memory=None) -> FastMCP:
|
||||
"""
|
||||
Create a stdio MCP server with retain/recall tools.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID to use for all operations.
|
||||
memory: Optional MemoryEngine instance. If not provided, creates one with pg0.
|
||||
|
||||
Returns:
|
||||
Configured FastMCP server instance.
|
||||
"""
|
||||
# Import here to avoid slow startup if just checking --help
|
||||
from hindsight_api import MemoryEngine
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
|
||||
|
||||
# Create memory engine with pg0 embedded database if not provided
|
||||
if memory is None:
|
||||
memory = MemoryEngine(db_url="pg0://hindsight-mcp")
|
||||
|
||||
# Get custom instructions from environment variable (appended to both tools)
|
||||
extra_instructions = os.environ.get(ENV_MCP_INSTRUCTIONS, "")
|
||||
|
||||
retain_description = DEFAULT_MCP_RETAIN_DESCRIPTION
|
||||
recall_description = DEFAULT_MCP_RECALL_DESCRIPTION
|
||||
|
||||
if extra_instructions:
|
||||
retain_description = f"{DEFAULT_MCP_RETAIN_DESCRIPTION}\n\nAdditional instructions: {extra_instructions}"
|
||||
recall_description = f"{DEFAULT_MCP_RECALL_DESCRIPTION}\n\nAdditional instructions: {extra_instructions}"
|
||||
|
||||
mcp = FastMCP("hindsight")
|
||||
|
||||
@mcp.tool(description=retain_description)
|
||||
async def retain(content: str, context: str = "general") -> dict:
|
||||
"""
|
||||
Args:
|
||||
content: The fact/memory to store (be specific and include relevant details)
|
||||
context: Category for the memory (e.g., 'preferences', 'work', 'hobbies', 'family'). Default: 'general'
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
async def _retain():
|
||||
try:
|
||||
await memory.retain_batch_async(bank_id=bank_id, contents=[{"content": content, "context": context}])
|
||||
except Exception as e:
|
||||
logger.error(f"Error storing memory: {e}", exc_info=True)
|
||||
|
||||
# Fire and forget - don't block on memory storage
|
||||
asyncio.create_task(_retain())
|
||||
return {"status": "accepted", "message": "Memory storage initiated"}
|
||||
|
||||
@mcp.tool(description=recall_description)
|
||||
async def recall(query: str, max_tokens: int = 4096, budget: str = "low") -> dict:
|
||||
"""
|
||||
Args:
|
||||
query: Natural language search query (e.g., "user's food preferences", "what projects is user working on")
|
||||
max_tokens: Maximum tokens to return in results (default: 4096)
|
||||
budget: Search budget level - "low", "mid", or "high" (default: "low")
|
||||
"""
|
||||
try:
|
||||
# Map string budget to enum
|
||||
budget_map = {"low": Budget.LOW, "mid": Budget.MID, "high": Budget.HIGH}
|
||||
budget_enum = budget_map.get(budget.lower(), Budget.LOW)
|
||||
|
||||
search_result = await memory.recall_async(
|
||||
bank_id=bank_id,
|
||||
query=query,
|
||||
fact_type=list(VALID_RECALL_FACT_TYPES),
|
||||
budget=budget_enum,
|
||||
max_tokens=max_tokens,
|
||||
)
|
||||
|
||||
return search_result.model_dump()
|
||||
except Exception as e:
|
||||
logger.error(f"Error searching: {e}", exc_info=True)
|
||||
return {"error": str(e), "results": []}
|
||||
|
||||
return mcp
|
||||
|
||||
|
||||
async def _initialize_and_run(bank_id: str):
|
||||
"""Initialize memory and run the MCP server."""
|
||||
from hindsight_api import MemoryEngine
|
||||
|
||||
# Create and initialize memory engine with pg0 embedded database
|
||||
# Note: We avoid printing to stderr during init as MCP clients show it as "errors"
|
||||
memory = MemoryEngine(db_url="pg0://hindsight-mcp")
|
||||
await memory.initialize()
|
||||
|
||||
# Create and run the server
|
||||
mcp = create_local_mcp_server(bank_id, memory=memory)
|
||||
await mcp.run_stdio_async()
|
||||
|
||||
|
||||
def main():
|
||||
"""Main entry point for the stdio MCP server."""
|
||||
import asyncio
|
||||
|
||||
from hindsight_api.config import ENV_LLM_API_KEY, get_config
|
||||
|
||||
# Check for required environment variables
|
||||
config = get_config()
|
||||
if not config.llm_api_key:
|
||||
print(f"Error: {ENV_LLM_API_KEY} environment variable is required", file=sys.stderr)
|
||||
print("Set it in your MCP configuration or shell environment", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
# Get bank ID from environment, default to "mcp"
|
||||
bank_id = os.environ.get(ENV_MCP_LOCAL_BANK_ID, DEFAULT_MCP_LOCAL_BANK_ID)
|
||||
|
||||
# Note: We don't print to stderr as MCP clients display it as "error output"
|
||||
# Use HINDSIGHT_API_LOG_LEVEL=debug for verbose startup logging
|
||||
|
||||
# Run the async initialization and server
|
||||
asyncio.run(_initialize_and_run(bank_id))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -6,16 +6,15 @@ This module provides metrics for:
|
||||
- Token usage (input/output) per operation
|
||||
- Per-bank granularity via labels
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Dict, Any, Optional
|
||||
from contextlib import contextmanager
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
|
||||
from opentelemetry import metrics
|
||||
from opentelemetry.exporter.prometheus import PrometheusMetricReader
|
||||
from opentelemetry.sdk.metrics import MeterProvider
|
||||
from opentelemetry.sdk.resources import Resource
|
||||
from opentelemetry.exporter.prometheus import PrometheusMetricReader
|
||||
from prometheus_client import REGISTRY
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -39,19 +38,18 @@ def initialize_metrics(service_name: str = "hindsight-api", service_version: str
|
||||
global _meter
|
||||
|
||||
# Create resource with service information
|
||||
resource = Resource.create({
|
||||
"service.name": service_name,
|
||||
"service.version": service_version,
|
||||
})
|
||||
resource = Resource.create(
|
||||
{
|
||||
"service.name": service_name,
|
||||
"service.version": service_version,
|
||||
}
|
||||
)
|
||||
|
||||
# Create Prometheus metric reader
|
||||
prometheus_reader = PrometheusMetricReader()
|
||||
|
||||
# Create meter provider with Prometheus exporter
|
||||
provider = MeterProvider(
|
||||
resource=resource,
|
||||
metric_readers=[prometheus_reader]
|
||||
)
|
||||
provider = MeterProvider(resource=resource, metric_readers=[prometheus_reader])
|
||||
|
||||
# Set the global meter provider
|
||||
metrics.set_meter_provider(provider)
|
||||
@@ -73,11 +71,19 @@ class MetricsCollectorBase:
|
||||
"""Base class for metrics collectors."""
|
||||
|
||||
@contextmanager
|
||||
def record_operation(self, operation: str, bank_id: str, budget: Optional[str] = None, max_tokens: Optional[int] = None):
|
||||
def record_operation(self, operation: str, bank_id: str, budget: str | None = None, max_tokens: int | None = None):
|
||||
"""Context manager to record operation duration and status."""
|
||||
raise NotImplementedError
|
||||
|
||||
def record_tokens(self, operation: str, bank_id: str, input_tokens: int = 0, output_tokens: int = 0, budget: Optional[str] = None, max_tokens: Optional[int] = None):
|
||||
def record_tokens(
|
||||
self,
|
||||
operation: str,
|
||||
bank_id: str,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
budget: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
):
|
||||
"""Record token usage for an operation."""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -86,11 +92,19 @@ class NoOpMetricsCollector(MetricsCollectorBase):
|
||||
"""No-op metrics collector that does nothing. Used when metrics are disabled."""
|
||||
|
||||
@contextmanager
|
||||
def record_operation(self, operation: str, bank_id: str, budget: Optional[str] = None, max_tokens: Optional[int] = None):
|
||||
def record_operation(self, operation: str, bank_id: str, budget: str | None = None, max_tokens: int | None = None):
|
||||
"""No-op context manager."""
|
||||
yield
|
||||
|
||||
def record_tokens(self, operation: str, bank_id: str, input_tokens: int = 0, output_tokens: int = 0, budget: Optional[str] = None, max_tokens: Optional[int] = None):
|
||||
def record_tokens(
|
||||
self,
|
||||
operation: str,
|
||||
bank_id: str,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
budget: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
):
|
||||
"""No-op token recording."""
|
||||
pass
|
||||
|
||||
@@ -108,33 +122,25 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
# Operation latency histogram (in seconds)
|
||||
# Records duration of retain, recall, reflect operations
|
||||
self.operation_duration = self.meter.create_histogram(
|
||||
name="hindsight.operation.duration",
|
||||
description="Duration of Hindsight operations in seconds",
|
||||
unit="s"
|
||||
name="hindsight.operation.duration", description="Duration of Hindsight operations in seconds", unit="s"
|
||||
)
|
||||
|
||||
# Token usage counters
|
||||
self.tokens_input = self.meter.create_counter(
|
||||
name="hindsight.tokens.input",
|
||||
description="Number of input tokens consumed",
|
||||
unit="tokens"
|
||||
name="hindsight.tokens.input", description="Number of input tokens consumed", unit="tokens"
|
||||
)
|
||||
|
||||
self.tokens_output = self.meter.create_counter(
|
||||
name="hindsight.tokens.output",
|
||||
description="Number of output tokens generated",
|
||||
unit="tokens"
|
||||
name="hindsight.tokens.output", description="Number of output tokens generated", unit="tokens"
|
||||
)
|
||||
|
||||
# Operation counter (success/failure)
|
||||
self.operation_total = self.meter.create_counter(
|
||||
name="hindsight.operation.total",
|
||||
description="Total number of operations executed",
|
||||
unit="operations"
|
||||
name="hindsight.operation.total", description="Total number of operations executed", unit="operations"
|
||||
)
|
||||
|
||||
@contextmanager
|
||||
def record_operation(self, operation: str, bank_id: str, budget: Optional[str] = None, max_tokens: Optional[int] = None):
|
||||
def record_operation(self, operation: str, bank_id: str, budget: str | None = None, max_tokens: int | None = None):
|
||||
"""
|
||||
Context manager to record operation duration and status.
|
||||
|
||||
@@ -175,7 +181,15 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
# Record operation count
|
||||
self.operation_total.add(1, attributes)
|
||||
|
||||
def record_tokens(self, operation: str, bank_id: str, input_tokens: int = 0, output_tokens: int = 0, budget: Optional[str] = None, max_tokens: Optional[int] = None):
|
||||
def record_tokens(
|
||||
self,
|
||||
operation: str,
|
||||
bank_id: str,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
budget: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
):
|
||||
"""
|
||||
Record token usage for an operation.
|
||||
|
||||
|
||||
@@ -11,11 +11,10 @@ safe rolling deployments.
|
||||
|
||||
No alembic.ini required - all configuration is done programmatically.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
@@ -31,7 +30,7 @@ def _run_migrations_internal(database_url: str, script_location: str) -> None:
|
||||
"""
|
||||
Internal function to run migrations without locking.
|
||||
"""
|
||||
logger.info(f"Running database migrations to head...")
|
||||
logger.info("Running database migrations to head...")
|
||||
logger.info(f"Database URL: {database_url}")
|
||||
logger.info(f"Script location: {script_location}")
|
||||
|
||||
@@ -57,7 +56,7 @@ def _run_migrations_internal(database_url: str, script_location: str) -> None:
|
||||
logger.info("Database migrations completed successfully")
|
||||
|
||||
|
||||
def run_migrations(database_url: str, script_location: Optional[str] = None) -> None:
|
||||
def run_migrations(database_url: str, script_location: str | None = None) -> None:
|
||||
"""
|
||||
Run database migrations to the latest version using programmatic Alembic configuration.
|
||||
|
||||
@@ -97,8 +96,7 @@ def run_migrations(database_url: str, script_location: Optional[str] = None) ->
|
||||
script_path = Path(script_location)
|
||||
if not script_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Alembic script location not found at {script_location}. "
|
||||
"Database migrations cannot be run."
|
||||
f"Alembic script location not found at {script_location}. Database migrations cannot be run."
|
||||
)
|
||||
|
||||
# Use PostgreSQL advisory lock to coordinate between distributed workers
|
||||
@@ -130,7 +128,9 @@ def run_migrations(database_url: str, script_location: Optional[str] = None) ->
|
||||
raise RuntimeError("Database migration failed") from e
|
||||
|
||||
|
||||
def check_migration_status(database_url: Optional[str] = None, script_location: Optional[str] = None) -> tuple[str | None, str | None]:
|
||||
def check_migration_status(
|
||||
database_url: str | None = None, script_location: str | None = None
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""
|
||||
Check current database schema version and latest available version.
|
||||
|
||||
@@ -151,7 +151,9 @@ def check_migration_status(database_url: Optional[str] = None, script_location:
|
||||
if database_url is None:
|
||||
database_url = os.getenv("HINDSIGHT_API_DATABASE_URL")
|
||||
if not database_url:
|
||||
logger.warning("Database URL not provided and HINDSIGHT_API_DATABASE_URL not set, cannot check migration status")
|
||||
logger.warning(
|
||||
"Database URL not provided and HINDSIGHT_API_DATABASE_URL not set, cannot check migration status"
|
||||
)
|
||||
return None, None
|
||||
|
||||
# Get current revision from database
|
||||
|
||||
@@ -1,49 +1,47 @@
|
||||
"""
|
||||
SQLAlchemy models for the memory system.
|
||||
"""
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from uuid import UUID as PyUUID, uuid4
|
||||
|
||||
from datetime import datetime
|
||||
from uuid import UUID as PyUUID
|
||||
|
||||
from pgvector.sqlalchemy import Vector
|
||||
from sqlalchemy import (
|
||||
CheckConstraint,
|
||||
Column,
|
||||
Float,
|
||||
ForeignKey,
|
||||
ForeignKeyConstraint,
|
||||
Index,
|
||||
Integer,
|
||||
PrimaryKeyConstraint,
|
||||
Text,
|
||||
func,
|
||||
)
|
||||
from sqlalchemy import (
|
||||
text as sql_text,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import JSONB, TIMESTAMP, UUID
|
||||
from sqlalchemy.ext.asyncio import AsyncAttrs
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship
|
||||
from pgvector.sqlalchemy import Vector
|
||||
|
||||
|
||||
class Base(AsyncAttrs, DeclarativeBase):
|
||||
"""Base class for all models."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class Document(Base):
|
||||
"""Source documents for memory units."""
|
||||
|
||||
__tablename__ = "documents"
|
||||
|
||||
id: Mapped[str] = mapped_column(Text, primary_key=True)
|
||||
bank_id: Mapped[str] = mapped_column(Text, primary_key=True)
|
||||
original_text: Mapped[Optional[str]] = mapped_column(Text)
|
||||
content_hash: Mapped[Optional[str]] = mapped_column(Text)
|
||||
original_text: Mapped[str | None] = mapped_column(Text)
|
||||
content_hash: Mapped[str | None] = mapped_column(Text)
|
||||
doc_metadata: Mapped[dict] = mapped_column("metadata", JSONB, server_default=sql_text("'{}'::jsonb"))
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
updated_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
|
||||
# Relationships
|
||||
memory_units = relationship("MemoryUnit", back_populates="document", cascade="all, delete-orphan")
|
||||
@@ -56,45 +54,42 @@ class Document(Base):
|
||||
|
||||
class MemoryUnit(Base):
|
||||
"""Individual sentence-level memories."""
|
||||
|
||||
__tablename__ = "memory_units"
|
||||
|
||||
id: Mapped[PyUUID] = mapped_column(
|
||||
UUID(as_uuid=True), primary_key=True, server_default=sql_text("gen_random_uuid()")
|
||||
)
|
||||
bank_id: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
document_id: Mapped[Optional[str]] = mapped_column(Text)
|
||||
document_id: Mapped[str | None] = mapped_column(Text)
|
||||
text: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
embedding = mapped_column(Vector(384)) # pgvector type
|
||||
context: Mapped[Optional[str]] = mapped_column(Text)
|
||||
event_date: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), nullable=False) # Kept for backward compatibility
|
||||
occurred_start: Mapped[Optional[datetime]] = mapped_column(TIMESTAMP(timezone=True)) # When fact occurred (range start)
|
||||
occurred_end: Mapped[Optional[datetime]] = mapped_column(TIMESTAMP(timezone=True)) # When fact occurred (range end)
|
||||
mentioned_at: Mapped[Optional[datetime]] = mapped_column(TIMESTAMP(timezone=True)) # When fact was mentioned
|
||||
context: Mapped[str | None] = mapped_column(Text)
|
||||
event_date: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), nullable=False
|
||||
) # Kept for backward compatibility
|
||||
occurred_start: Mapped[datetime | None] = mapped_column(
|
||||
TIMESTAMP(timezone=True)
|
||||
) # When fact occurred (range start)
|
||||
occurred_end: Mapped[datetime | None] = mapped_column(TIMESTAMP(timezone=True)) # When fact occurred (range end)
|
||||
mentioned_at: Mapped[datetime | None] = mapped_column(TIMESTAMP(timezone=True)) # When fact was mentioned
|
||||
fact_type: Mapped[str] = mapped_column(Text, nullable=False, server_default="world")
|
||||
confidence_score: Mapped[Optional[float]] = mapped_column(Float)
|
||||
confidence_score: Mapped[float | None] = mapped_column(Float)
|
||||
access_count: Mapped[int] = mapped_column(Integer, server_default="0")
|
||||
unit_metadata: Mapped[dict] = mapped_column("metadata", JSONB, server_default=sql_text("'{}'::jsonb")) # User-defined metadata (str->str)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
unit_metadata: Mapped[dict] = mapped_column(
|
||||
"metadata", JSONB, server_default=sql_text("'{}'::jsonb")
|
||||
) # User-defined metadata (str->str)
|
||||
created_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
updated_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
|
||||
# Relationships
|
||||
document = relationship("Document", back_populates="memory_units")
|
||||
unit_entities = relationship("UnitEntity", back_populates="memory_unit", cascade="all, delete-orphan")
|
||||
outgoing_links = relationship(
|
||||
"MemoryLink",
|
||||
foreign_keys="MemoryLink.from_unit_id",
|
||||
back_populates="from_unit",
|
||||
cascade="all, delete-orphan"
|
||||
"MemoryLink", foreign_keys="MemoryLink.from_unit_id", back_populates="from_unit", cascade="all, delete-orphan"
|
||||
)
|
||||
incoming_links = relationship(
|
||||
"MemoryLink",
|
||||
foreign_keys="MemoryLink.to_unit_id",
|
||||
back_populates="to_unit",
|
||||
cascade="all, delete-orphan"
|
||||
"MemoryLink", foreign_keys="MemoryLink.to_unit_id", back_populates="to_unit", cascade="all, delete-orphan"
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
@@ -110,7 +105,7 @@ class MemoryUnit(Base):
|
||||
"(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",
|
||||
),
|
||||
Index("idx_memory_units_bank_id", "bank_id"),
|
||||
Index("idx_memory_units_document_id", "document_id"),
|
||||
@@ -119,39 +114,46 @@ class MemoryUnit(Base):
|
||||
Index("idx_memory_units_access_count", "access_count", postgresql_ops={"access_count": "DESC"}),
|
||||
Index("idx_memory_units_fact_type", "fact_type"),
|
||||
Index("idx_memory_units_bank_fact_type", "bank_id", "fact_type"),
|
||||
Index("idx_memory_units_bank_type_date", "bank_id", "fact_type", "event_date", postgresql_ops={"event_date": "DESC"}),
|
||||
Index(
|
||||
"idx_memory_units_bank_type_date",
|
||||
"bank_id",
|
||||
"fact_type",
|
||||
"event_date",
|
||||
postgresql_ops={"event_date": "DESC"},
|
||||
),
|
||||
Index(
|
||||
"idx_memory_units_opinion_confidence",
|
||||
"bank_id",
|
||||
"confidence_score",
|
||||
postgresql_where=sql_text("fact_type = 'opinion'"),
|
||||
postgresql_ops={"confidence_score": "DESC"}
|
||||
postgresql_ops={"confidence_score": "DESC"},
|
||||
),
|
||||
Index(
|
||||
"idx_memory_units_opinion_date",
|
||||
"bank_id",
|
||||
"event_date",
|
||||
postgresql_where=sql_text("fact_type = 'opinion'"),
|
||||
postgresql_ops={"event_date": "DESC"}
|
||||
postgresql_ops={"event_date": "DESC"},
|
||||
),
|
||||
Index(
|
||||
"idx_memory_units_observation_date",
|
||||
"bank_id",
|
||||
"event_date",
|
||||
postgresql_where=sql_text("fact_type = 'observation'"),
|
||||
postgresql_ops={"event_date": "DESC"}
|
||||
postgresql_ops={"event_date": "DESC"},
|
||||
),
|
||||
Index(
|
||||
"idx_memory_units_embedding",
|
||||
"embedding",
|
||||
postgresql_using="hnsw",
|
||||
postgresql_ops={"embedding": "vector_cosine_ops"}
|
||||
postgresql_ops={"embedding": "vector_cosine_ops"},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class Entity(Base):
|
||||
"""Resolved entities (people, organizations, locations, etc.)."""
|
||||
|
||||
__tablename__ = "entities"
|
||||
|
||||
id: Mapped[PyUUID] = mapped_column(
|
||||
@@ -160,12 +162,8 @@ class Entity(Base):
|
||||
canonical_name: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
bank_id: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
entity_metadata: Mapped[dict] = mapped_column("metadata", JSONB, server_default=sql_text("'{}'::jsonb"))
|
||||
first_seen: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
last_seen: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
first_seen: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
last_seen: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
mention_count: Mapped[int] = mapped_column(Integer, server_default="1")
|
||||
|
||||
# Relationships
|
||||
@@ -175,13 +173,13 @@ class Entity(Base):
|
||||
"EntityCooccurrence",
|
||||
foreign_keys="EntityCooccurrence.entity_id_1",
|
||||
back_populates="entity_1",
|
||||
cascade="all, delete-orphan"
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
cooccurrences_2 = relationship(
|
||||
"EntityCooccurrence",
|
||||
foreign_keys="EntityCooccurrence.entity_id_2",
|
||||
back_populates="entity_2",
|
||||
cascade="all, delete-orphan"
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
@@ -193,6 +191,7 @@ class Entity(Base):
|
||||
|
||||
class UnitEntity(Base):
|
||||
"""Association between memory units and entities."""
|
||||
|
||||
__tablename__ = "unit_entities"
|
||||
|
||||
unit_id: Mapped[PyUUID] = mapped_column(
|
||||
@@ -214,6 +213,7 @@ class UnitEntity(Base):
|
||||
|
||||
class EntityCooccurrence(Base):
|
||||
"""Materialized cache of entity co-occurrences."""
|
||||
|
||||
__tablename__ = "entity_cooccurrences"
|
||||
|
||||
entity_id_1: Mapped[PyUUID] = mapped_column(
|
||||
@@ -223,9 +223,7 @@ class EntityCooccurrence(Base):
|
||||
UUID(as_uuid=True), ForeignKey("entities.id", ondelete="CASCADE"), primary_key=True
|
||||
)
|
||||
cooccurrence_count: Mapped[int] = mapped_column(Integer, server_default="1")
|
||||
last_cooccurred: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
last_cooccurred: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
|
||||
# Relationships
|
||||
entity_1 = relationship("Entity", foreign_keys=[entity_id_1], back_populates="cooccurrences_1")
|
||||
@@ -241,6 +239,7 @@ class EntityCooccurrence(Base):
|
||||
|
||||
class MemoryLink(Base):
|
||||
"""Links between memory units (temporal, semantic, entity)."""
|
||||
|
||||
__tablename__ = "memory_links"
|
||||
|
||||
from_unit_id: Mapped[PyUUID] = mapped_column(
|
||||
@@ -250,13 +249,11 @@ class MemoryLink(Base):
|
||||
UUID(as_uuid=True), ForeignKey("memory_units.id", ondelete="CASCADE"), primary_key=True
|
||||
)
|
||||
link_type: Mapped[str] = mapped_column(Text, primary_key=True)
|
||||
entity_id: Mapped[Optional[PyUUID]] = mapped_column(
|
||||
entity_id: Mapped[PyUUID | None] = mapped_column(
|
||||
UUID(as_uuid=True), ForeignKey("entities.id", ondelete="CASCADE"), primary_key=True
|
||||
)
|
||||
weight: Mapped[float] = mapped_column(Float, nullable=False, server_default="1.0")
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
|
||||
# Relationships
|
||||
from_unit = relationship("MemoryUnit", foreign_keys=[from_unit_id], back_populates="outgoing_links")
|
||||
@@ -266,7 +263,7 @@ class MemoryLink(Base):
|
||||
__table_args__ = (
|
||||
CheckConstraint(
|
||||
"link_type IN ('temporal', 'semantic', 'entity', 'causes', 'caused_by', 'enables', 'prevents')",
|
||||
name="memory_links_link_type_check"
|
||||
name="memory_links_link_type_check",
|
||||
),
|
||||
CheckConstraint("weight >= 0.0 AND weight <= 1.0", name="memory_links_weight_check"),
|
||||
Index("idx_memory_links_from", "from_unit_id"),
|
||||
@@ -278,31 +275,22 @@ class MemoryLink(Base):
|
||||
"from_unit_id",
|
||||
"weight",
|
||||
postgresql_where=sql_text("weight >= 0.1"),
|
||||
postgresql_ops={"weight": "DESC"}
|
||||
postgresql_ops={"weight": "DESC"},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class Bank(Base):
|
||||
"""Memory bank profiles with disposition traits and background."""
|
||||
|
||||
__tablename__ = "banks"
|
||||
|
||||
bank_id: Mapped[str] = mapped_column(Text, primary_key=True)
|
||||
disposition: Mapped[dict] = mapped_column(
|
||||
JSONB,
|
||||
nullable=False,
|
||||
server_default=sql_text(
|
||||
'\'{"skepticism": 3, "literalism": 3, "empathy": 3}\'::jsonb'
|
||||
)
|
||||
JSONB, nullable=False, server_default=sql_text('\'{"skepticism": 3, "literalism": 3, "empathy": 3}\'::jsonb')
|
||||
)
|
||||
background: Mapped[str] = mapped_column(Text, nullable=False, server_default="")
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now()
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
updated_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
|
||||
|
||||
__table_args__ = (
|
||||
Index("idx_banks_bank_id", "bank_id"),
|
||||
)
|
||||
__table_args__ = (Index("idx_banks_bank_id", "bank_id"),)
|
||||
|
||||
@@ -1,12 +1,10 @@
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from pg0 import Pg0
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_PORT = 5555
|
||||
DEFAULT_USERNAME = "hindsight"
|
||||
DEFAULT_PASSWORD = "hindsight"
|
||||
DEFAULT_DATABASE = "hindsight"
|
||||
@@ -17,34 +15,38 @@ class EmbeddedPostgres:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
port: int = DEFAULT_PORT,
|
||||
port: int | None = None,
|
||||
username: str = DEFAULT_USERNAME,
|
||||
password: str = DEFAULT_PASSWORD,
|
||||
database: str = DEFAULT_DATABASE,
|
||||
name: str = "hindsight",
|
||||
**kwargs,
|
||||
):
|
||||
self.port = port
|
||||
self.port = port # None means pg0 will auto-assign
|
||||
self.username = username
|
||||
self.password = password
|
||||
self.database = database
|
||||
self.name = name
|
||||
self._pg0: Optional[Pg0] = None
|
||||
self._pg0: Pg0 | None = None
|
||||
|
||||
def _get_pg0(self) -> Pg0:
|
||||
if self._pg0 is None:
|
||||
self._pg0 = Pg0(
|
||||
name=self.name,
|
||||
port=self.port,
|
||||
username=self.username,
|
||||
password=self.password,
|
||||
database=self.database,
|
||||
)
|
||||
kwargs = {
|
||||
"name": self.name,
|
||||
"username": self.username,
|
||||
"password": self.password,
|
||||
"database": self.database,
|
||||
}
|
||||
# Only set port if explicitly specified
|
||||
if self.port is not None:
|
||||
kwargs["port"] = self.port
|
||||
self._pg0 = Pg0(**kwargs)
|
||||
return self._pg0
|
||||
|
||||
async def start(self, max_retries: int = 3, retry_delay: float = 2.0) -> str:
|
||||
async def start(self, max_retries: int = 5, retry_delay: float = 4.0) -> str:
|
||||
"""Start the PostgreSQL server with retry logic."""
|
||||
logger.info(f"Starting embedded PostgreSQL (name: {self.name}, port: {self.port})...")
|
||||
port_info = f"port={self.port}" if self.port else "port=auto"
|
||||
logger.info(f"Starting embedded PostgreSQL (name={self.name}, {port_info})...")
|
||||
|
||||
pg0 = self._get_pg0()
|
||||
last_error = None
|
||||
@@ -53,9 +55,9 @@ class EmbeddedPostgres:
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
info = await loop.run_in_executor(None, pg0.start)
|
||||
logger.info(f"PostgreSQL started on port {self.port}")
|
||||
# Construct URI manually since pg0-embedded may return None
|
||||
uri = info.uri if info and info.uri else f"postgresql://{self.username}:{self.password}@localhost:{self.port}/{self.database}"
|
||||
# Get URI from pg0 (includes auto-assigned port)
|
||||
uri = info.uri
|
||||
logger.info(f"PostgreSQL started: {uri}")
|
||||
return uri
|
||||
except Exception as e:
|
||||
last_error = str(e)
|
||||
@@ -68,8 +70,7 @@ class EmbeddedPostgres:
|
||||
logger.debug(f"pg0 start attempt {attempt}/{max_retries} failed: {last_error}")
|
||||
|
||||
raise RuntimeError(
|
||||
f"Failed to start embedded PostgreSQL after {max_retries} attempts. "
|
||||
f"Last error: {last_error}"
|
||||
f"Failed to start embedded PostgreSQL after {max_retries} attempts. Last error: {last_error}"
|
||||
)
|
||||
|
||||
async def stop(self) -> None:
|
||||
@@ -91,9 +92,7 @@ class EmbeddedPostgres:
|
||||
pg0 = self._get_pg0()
|
||||
loop = asyncio.get_event_loop()
|
||||
info = await loop.run_in_executor(None, pg0.info)
|
||||
# Construct URI manually since pg0-embedded may return None
|
||||
uri = info.uri if info and info.uri else f"postgresql://{self.username}:{self.password}@localhost:{self.port}/{self.database}"
|
||||
return uri
|
||||
return info.uri
|
||||
|
||||
async def is_running(self) -> bool:
|
||||
"""Check if the PostgreSQL server is currently running."""
|
||||
@@ -112,7 +111,7 @@ class EmbeddedPostgres:
|
||||
return await self.start()
|
||||
|
||||
|
||||
_default_instance: Optional[EmbeddedPostgres] = None
|
||||
_default_instance: EmbeddedPostgres | None = None
|
||||
|
||||
|
||||
def get_embedded_postgres() -> EmbeddedPostgres:
|
||||
|
||||
@@ -6,6 +6,7 @@ This module provides the ASGI app for uvicorn import string usage:
|
||||
|
||||
For CLI usage, use the hindsight-api command instead.
|
||||
"""
|
||||
|
||||
import os
|
||||
import warnings
|
||||
|
||||
@@ -29,15 +30,11 @@ config.configure_logging()
|
||||
_memory = MemoryEngine()
|
||||
|
||||
# Create unified app with both HTTP and optionally MCP
|
||||
app = create_app(
|
||||
memory=_memory,
|
||||
http_api_enabled=True,
|
||||
mcp_api_enabled=config.mcp_enabled,
|
||||
mcp_mount_path="/mcp"
|
||||
)
|
||||
app = create_app(memory=_memory, http_api_enabled=True, mcp_api_enabled=config.mcp_enabled, mcp_mount_path="/mcp")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# When run directly, delegate to the CLI
|
||||
from hindsight_api.main import main
|
||||
|
||||
main()
|
||||
|
||||
@@ -4,8 +4,8 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "hindsight-api"
|
||||
version = "0.1.5"
|
||||
description = "Temporal + Semantic + Entity Memory System for AI agents using PostgreSQL"
|
||||
version = "0.1.8"
|
||||
description = "Hindsight: Agent Memory That Works Like Human Memory"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
@@ -25,11 +25,11 @@ dependencies = [
|
||||
"greenlet>=3.2.4",
|
||||
"psycopg2-binary>=2.9.11",
|
||||
"transformers>=4.30.0,<4.46.0",
|
||||
"torch>=2.0.0,<2.6.0",
|
||||
"torch>=2.0.0",
|
||||
"tiktoken>=0.12.0",
|
||||
"httpx>=0.27.0",
|
||||
"fastmcp>=2.3.0",
|
||||
"pg0-embedded>=0.1.0",
|
||||
"pg0-embedded>=0.11.0",
|
||||
"python-dateutil>=2.8.0",
|
||||
"opentelemetry-api>=1.20.0",
|
||||
"opentelemetry-sdk>=1.20.0",
|
||||
@@ -50,6 +50,7 @@ test = [
|
||||
|
||||
[project.scripts]
|
||||
hindsight-api = "hindsight_api.main:main"
|
||||
hindsight-local-mcp = "hindsight_api.mcp_local:main"
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
packages = ["hindsight_api"]
|
||||
@@ -90,4 +91,33 @@ dev = [
|
||||
"pytest-xdist>=3.8.0",
|
||||
"python-dotenv>=1.2.1",
|
||||
"filelock>=3.0.0",
|
||||
"ruff>=0.8.0",
|
||||
]
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 120
|
||||
target-version = "py311"
|
||||
exclude = [
|
||||
"tests/",
|
||||
"**/tests/",
|
||||
]
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = [
|
||||
"E", # pycodestyle errors
|
||||
"W", # pycodestyle warnings
|
||||
"F", # Pyflakes
|
||||
"I", # isort
|
||||
]
|
||||
ignore = [
|
||||
"E501", # line too long (handled by formatter)
|
||||
"E402", # module import not at top of file
|
||||
"F401", # unused import (too noisy during development)
|
||||
"F841", # unused variable (too noisy during development)
|
||||
"F811", # redefined while unused
|
||||
"F821", # undefined name (forward references in type hints)
|
||||
]
|
||||
|
||||
[tool.ruff.format]
|
||||
quote-style = "double"
|
||||
indent-style = "space"
|
||||
|
||||
@@ -12,12 +12,12 @@ This comprehensive test suite validates that the fact extraction system:
|
||||
These are quality/accuracy tests that verify the LLM-based extraction
|
||||
produces semantically correct and complete facts.
|
||||
"""
|
||||
import pytest
|
||||
import re
|
||||
from datetime import datetime, timezone
|
||||
from hindsight_api.engine.retain.fact_extraction import extract_facts_from_text
|
||||
from hindsight_api import LLMConfig
|
||||
from datetime import UTC, datetime
|
||||
|
||||
import pytest
|
||||
|
||||
from hindsight_api import LLMConfig
|
||||
from hindsight_api.engine.retain.fact_extraction import extract_facts_from_text
|
||||
|
||||
# =============================================================================
|
||||
# DIMENSION PRESERVATION TESTS
|
||||
@@ -432,6 +432,7 @@ with a concert surrounded by music, joy and the warm summer breeze.
|
||||
assert birthday_fact is not None, "Should extract fact about birthday celebration"
|
||||
|
||||
fact_date_str = birthday_fact.occurred_start
|
||||
assert fact_date_str is not None, "occurred_start should not be None for temporal events"
|
||||
|
||||
if 'T' in fact_date_str:
|
||||
fact_date = datetime.fromisoformat(fact_date_str.replace('Z', '+00:00'))
|
||||
@@ -497,7 +498,7 @@ Yesterday I went for a morning jog for the first time in a nearby park.
|
||||
async def test_extract_facts_with_relative_dates(self):
|
||||
"""Test that relative dates are converted to absolute dates."""
|
||||
|
||||
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=timezone.utc)
|
||||
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=UTC)
|
||||
llm_config = LLMConfig.for_memory()
|
||||
|
||||
text = """
|
||||
@@ -531,7 +532,7 @@ Yesterday I went for a morning jog for the first time in a nearby park.
|
||||
async def test_extract_facts_with_no_temporal_info(self):
|
||||
"""Test that facts without temporal info are still extracted."""
|
||||
|
||||
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=timezone.utc)
|
||||
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=UTC)
|
||||
llm_config = LLMConfig.for_memory()
|
||||
|
||||
text = "Alice works at Google. She loves Python programming."
|
||||
@@ -555,7 +556,7 @@ Yesterday I went for a morning jog for the first time in a nearby park.
|
||||
async def test_extract_facts_with_absolute_dates(self):
|
||||
"""Test that absolute dates in text are preserved."""
|
||||
|
||||
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=timezone.utc)
|
||||
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=UTC)
|
||||
llm_config = LLMConfig.for_memory()
|
||||
|
||||
text = """
|
||||
@@ -1047,4 +1048,4 @@ class TestDispositionInference:
|
||||
|
||||
assert "texas" in background.lower()
|
||||
# Higher skepticism expected from "very skeptical of people"
|
||||
assert disposition["skepticism"] >= 3
|
||||
assert disposition["skepticism"] >= 3
|
||||
|
||||
@@ -426,3 +426,185 @@ async def test_document_deletion(api_client):
|
||||
f"/v1/default/banks/{test_bank_id}/documents/sales-report-q1-2024"
|
||||
)
|
||||
assert response.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_retain(api_client):
|
||||
"""Test asynchronous retain functionality.
|
||||
|
||||
When async=true is passed, the retain endpoint should:
|
||||
1. Return immediately with success and async_=true
|
||||
2. Process the content in the background
|
||||
3. Eventually store the memories
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
test_bank_id = f"async_retain_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memory with async=true
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"async": True,
|
||||
"items": [
|
||||
{
|
||||
"content": "Alice is a senior engineer at TechCorp. She has been working on the authentication system for 5 years.",
|
||||
"context": "team introduction"
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
assert result["success"] is True
|
||||
assert result["async"] is True, "Response should indicate async processing"
|
||||
assert result["items_count"] == 1
|
||||
|
||||
# Check operations endpoint to see the pending operation
|
||||
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/operations")
|
||||
assert response.status_code == 200
|
||||
ops_result = response.json()
|
||||
assert "operations" in ops_result
|
||||
|
||||
# Wait for async processing to complete (poll with timeout)
|
||||
max_wait_seconds = 30
|
||||
poll_interval = 0.5
|
||||
elapsed = 0
|
||||
memories_found = False
|
||||
|
||||
while elapsed < max_wait_seconds:
|
||||
# Check if memories are stored
|
||||
response = await api_client.get(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/list",
|
||||
params={"limit": 10}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
items = response.json()["items"]
|
||||
|
||||
if len(items) > 0:
|
||||
memories_found = True
|
||||
break
|
||||
|
||||
await asyncio.sleep(poll_interval)
|
||||
elapsed += poll_interval
|
||||
|
||||
assert memories_found, f"Async retain did not complete within {max_wait_seconds} seconds"
|
||||
|
||||
# Verify we can recall the stored memory
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/recall",
|
||||
json={
|
||||
"query": "Who works at TechCorp?",
|
||||
"thinking_budget": 30
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
search_results = response.json()
|
||||
assert len(search_results["results"]) > 0, "Should find the asynchronously stored memory"
|
||||
|
||||
# Verify Alice is mentioned
|
||||
found_alice = any("Alice" in r["text"] for r in search_results["results"])
|
||||
assert found_alice, "Should find Alice in search results"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_retain_parallel(api_client):
|
||||
"""Test multiple async retain operations running in parallel.
|
||||
|
||||
Verifies that:
|
||||
1. Multiple async operations can be submitted concurrently
|
||||
2. All operations complete successfully
|
||||
3. The exact number of documents are processed
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
test_bank_id = f"async_parallel_test_{datetime.now().timestamp()}"
|
||||
num_documents = 5
|
||||
|
||||
# Prepare multiple documents to retain
|
||||
documents = [
|
||||
{
|
||||
"content": f"Document {i}: This is test content about Person{i} who works at Company{i}.",
|
||||
"context": f"test document {i}",
|
||||
"document_id": f"doc_{i}"
|
||||
}
|
||||
for i in range(num_documents)
|
||||
]
|
||||
|
||||
# Submit all async retain operations in parallel
|
||||
async def submit_async_retain(doc):
|
||||
return await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"async": True,
|
||||
"items": [doc]
|
||||
}
|
||||
)
|
||||
|
||||
# Run all submissions concurrently
|
||||
responses = await asyncio.gather(*[submit_async_retain(doc) for doc in documents])
|
||||
|
||||
# Verify all submissions succeeded
|
||||
for i, response in enumerate(responses):
|
||||
assert response.status_code == 200, f"Document {i} submission failed"
|
||||
result = response.json()
|
||||
assert result["success"] is True
|
||||
assert result["async"] is True
|
||||
|
||||
# Check operations endpoint - should show pending operations
|
||||
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/operations")
|
||||
assert response.status_code == 200
|
||||
|
||||
# Wait for all async operations to complete (poll with timeout)
|
||||
max_wait_seconds = 60
|
||||
poll_interval = 1.0
|
||||
elapsed = 0
|
||||
all_docs_processed = False
|
||||
|
||||
while elapsed < max_wait_seconds:
|
||||
# Check document count
|
||||
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/documents")
|
||||
assert response.status_code == 200
|
||||
docs = response.json()["items"]
|
||||
|
||||
if len(docs) >= num_documents:
|
||||
all_docs_processed = True
|
||||
break
|
||||
|
||||
await asyncio.sleep(poll_interval)
|
||||
elapsed += poll_interval
|
||||
|
||||
assert all_docs_processed, f"Expected {num_documents} documents, but only {len(docs)} were processed within {max_wait_seconds} seconds"
|
||||
|
||||
# Verify exact document count
|
||||
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/documents")
|
||||
assert response.status_code == 200
|
||||
final_docs = response.json()["items"]
|
||||
assert len(final_docs) == num_documents, f"Expected exactly {num_documents} documents, got {len(final_docs)}"
|
||||
|
||||
# Verify each document exists
|
||||
doc_ids = {doc["id"] for doc in final_docs}
|
||||
for i in range(num_documents):
|
||||
assert f"doc_{i}" in doc_ids, f"Document doc_{i} not found"
|
||||
|
||||
# Verify memories were created for all documents
|
||||
response = await api_client.get(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/list",
|
||||
params={"limit": 100}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
memories = response.json()["items"]
|
||||
assert len(memories) >= num_documents, f"Expected at least {num_documents} memories, got {len(memories)}"
|
||||
|
||||
# Verify we can recall content from different documents
|
||||
for i in [0, num_documents - 1]: # Check first and last
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/recall",
|
||||
json={
|
||||
"query": f"Who works at Company{i}?",
|
||||
"thinking_budget": 30
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
results = response.json()["results"]
|
||||
assert len(results) > 0, f"Should find memories for document {i}"
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
"""
|
||||
Test LLM provider with different models and providers.
|
||||
Test LLM provider with different models using actual memory operations.
|
||||
"""
|
||||
import os
|
||||
from datetime import datetime
|
||||
import pytest
|
||||
from hindsight_api.engine.llm_wrapper import LLMProvider
|
||||
from hindsight_api.engine.utils import extract_facts
|
||||
from hindsight_api.engine.search.think_utils import reflect
|
||||
|
||||
|
||||
# Model matrix: (provider, model)
|
||||
@@ -15,13 +18,14 @@ MODEL_MATRIX = [
|
||||
("openai", "gpt-5-mini"),
|
||||
("openai", "gpt-5-nano"),
|
||||
("openai", "gpt-5"),
|
||||
("openai", "gpt-5.2"),
|
||||
# Groq models
|
||||
("groq", "llama-3.3-70b-versatile"),
|
||||
("groq", "openai/gpt-oss-120b"),
|
||||
("groq", "openai/gpt-oss-20b"),
|
||||
# Gemini models
|
||||
("gemini", "gemini-2.5-flash"),
|
||||
("gemini", "gemini-2.5-flash-lite"),
|
||||
("gemini", "gemini-3-pro-preview"),
|
||||
]
|
||||
|
||||
|
||||
@@ -38,10 +42,10 @@ def get_api_key_for_provider(provider: str) -> str | None:
|
||||
|
||||
@pytest.mark.parametrize("provider,model", MODEL_MATRIX)
|
||||
@pytest.mark.asyncio
|
||||
async def test_llm_provider_call(provider: str, model: str):
|
||||
async def test_llm_provider_memory_operations(provider: str, model: str):
|
||||
"""
|
||||
Test LLM provider can make a basic call with different models.
|
||||
Skips if the required API key is not available.
|
||||
Test LLM provider with actual memory operations: fact extraction and reflect.
|
||||
All models must pass this test.
|
||||
"""
|
||||
api_key = get_api_key_for_provider(provider)
|
||||
if not api_key:
|
||||
@@ -54,74 +58,53 @@ async def test_llm_provider_call(provider: str, model: str):
|
||||
model=model,
|
||||
)
|
||||
|
||||
# Test basic call
|
||||
response = await llm.call(
|
||||
messages=[{"role": "user", "content": "Say 'hello' and nothing else."}],
|
||||
max_completion_tokens=50,
|
||||
temperature=0.1,
|
||||
# Test 1: Fact extraction (structured output)
|
||||
test_text = """
|
||||
User: I just got back from my trip to Paris last week. The Eiffel Tower was amazing!
|
||||
Assistant: That sounds wonderful! How long were you there?
|
||||
User: About 5 days. I also visited the Louvre and saw the Mona Lisa.
|
||||
"""
|
||||
event_date = datetime(2024, 12, 10)
|
||||
|
||||
facts, chunks = await extract_facts(
|
||||
text=test_text,
|
||||
event_date=event_date,
|
||||
context="Travel conversation",
|
||||
llm_config=llm,
|
||||
)
|
||||
|
||||
print(f"\n{provider}/{model} response: {response}")
|
||||
assert response is not None, f"{provider}/{model} returned None"
|
||||
print(f"\n{provider}/{model} - Fact extraction:")
|
||||
print(f" Extracted {len(facts)} facts from {len(chunks)} chunks")
|
||||
for fact in facts:
|
||||
print(f" - {fact.fact}")
|
||||
|
||||
assert facts is not None, f"{provider}/{model} fact extraction returned None"
|
||||
assert len(facts) > 0, f"{provider}/{model} should extract at least one fact"
|
||||
|
||||
@pytest.mark.parametrize("provider,model", MODEL_MATRIX)
|
||||
@pytest.mark.asyncio
|
||||
async def test_llm_provider_verify_connection(provider: str, model: str):
|
||||
"""
|
||||
Test LLM provider verify_connection method with different models.
|
||||
Skips if the required API key is not available.
|
||||
"""
|
||||
api_key = get_api_key_for_provider(provider)
|
||||
if not api_key:
|
||||
pytest.skip(f"Skipping {provider}/{model}: no API key available")
|
||||
# Verify facts have required fields
|
||||
for fact in facts:
|
||||
assert fact.fact, f"{provider}/{model} fact missing text"
|
||||
assert fact.fact_type in ["world", "experience", "opinion"], f"{provider}/{model} invalid fact_type: {fact.fact_type}"
|
||||
|
||||
llm = LLMProvider(
|
||||
provider=provider,
|
||||
api_key=api_key,
|
||||
base_url="",
|
||||
model=model,
|
||||
# Test 2: Reflect (actual reflect function)
|
||||
response = await reflect(
|
||||
llm_config=llm,
|
||||
query="What was the highlight of my Paris trip?",
|
||||
experience_facts=[
|
||||
"I visited Paris in December 2024",
|
||||
"I saw the Eiffel Tower and it was amazing",
|
||||
"I visited the Louvre and saw the Mona Lisa",
|
||||
"The trip lasted 5 days",
|
||||
],
|
||||
world_facts=[
|
||||
"The Eiffel Tower is a famous landmark in Paris",
|
||||
"The Mona Lisa is displayed at the Louvre museum",
|
||||
],
|
||||
name="Traveler",
|
||||
)
|
||||
|
||||
# Test verify_connection
|
||||
await llm.verify_connection()
|
||||
print(f"\n{provider}/{model} connection verified")
|
||||
print(f"\n{provider}/{model} - Reflect response:")
|
||||
print(f" {response[:200]}...")
|
||||
|
||||
|
||||
# Models that support large output (65000+ tokens)
|
||||
LARGE_OUTPUT_MODELS = [
|
||||
("openai", "gpt-5-mini"),
|
||||
("openai", "gpt-5-nano"),
|
||||
("openai", "gpt-5"),
|
||||
("gemini", "gemini-2.5-flash"),
|
||||
("gemini", "gemini-2.5-flash-lite"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider,model", LARGE_OUTPUT_MODELS)
|
||||
@pytest.mark.asyncio
|
||||
async def test_llm_provider_large_output(provider: str, model: str):
|
||||
"""
|
||||
Test LLM provider with large max_completion_tokens (65000).
|
||||
Only tests models that support large outputs.
|
||||
Skips if the required API key is not available.
|
||||
"""
|
||||
api_key = get_api_key_for_provider(provider)
|
||||
if not api_key:
|
||||
pytest.skip(f"Skipping {provider}/{model}: no API key available")
|
||||
|
||||
llm = LLMProvider(
|
||||
provider=provider,
|
||||
api_key=api_key,
|
||||
base_url="",
|
||||
model=model,
|
||||
)
|
||||
|
||||
# Test call with large max_completion_tokens
|
||||
response = await llm.call(
|
||||
messages=[{"role": "user", "content": "Say 'ok'"}],
|
||||
max_completion_tokens=65000,
|
||||
)
|
||||
|
||||
print(f"\n{provider}/{model} large output response: {response}")
|
||||
assert response is not None, f"{provider}/{model} returned None"
|
||||
assert response is not None, f"{provider}/{model} reflect returned None"
|
||||
assert len(response) > 10, f"{provider}/{model} reflect response too short"
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
"""Test local MCP server."""
|
||||
|
||||
import asyncio
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_memory():
|
||||
"""Create a mock MemoryEngine."""
|
||||
memory = MagicMock()
|
||||
memory._initialized = True
|
||||
memory.retain_batch_async = AsyncMock()
|
||||
memory.recall_async = AsyncMock(return_value=MagicMock(results=[]))
|
||||
return memory
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_mcp_server_retain(mock_memory):
|
||||
"""Test that retain tool fires async and returns immediately."""
|
||||
from hindsight_api.mcp_local import create_local_mcp_server
|
||||
|
||||
bank_id = "test-bank"
|
||||
mcp_server = create_local_mcp_server(bank_id, memory=mock_memory)
|
||||
|
||||
# Get the tools
|
||||
tools = mcp_server._tool_manager._tools
|
||||
assert "retain" in tools
|
||||
|
||||
# Call retain
|
||||
retain_tool = tools["retain"]
|
||||
result = await retain_tool.fn(content="test content", context="test_context")
|
||||
|
||||
# Returns immediately with accepted status
|
||||
assert result["status"] == "accepted"
|
||||
|
||||
# Wait for background task to complete
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Verify the memory was called correctly
|
||||
mock_memory.retain_batch_async.assert_called_once()
|
||||
call_kwargs = mock_memory.retain_batch_async.call_args.kwargs
|
||||
assert call_kwargs["bank_id"] == "test-bank"
|
||||
assert call_kwargs["contents"] == [{"content": "test content", "context": "test_context"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_mcp_server_recall(mock_memory):
|
||||
"""Test that recall tool calls memory.recall_async with correct params."""
|
||||
from hindsight_api.mcp_local import create_local_mcp_server
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
|
||||
# Mock recall_async to return a proper pydantic model
|
||||
mock_result = MagicMock()
|
||||
mock_result.model_dump.return_value = {"results": []}
|
||||
mock_memory.recall_async = AsyncMock(return_value=mock_result)
|
||||
|
||||
bank_id = "test-bank"
|
||||
mcp_server = create_local_mcp_server(bank_id, memory=mock_memory)
|
||||
|
||||
# Get the tools
|
||||
tools = mcp_server._tool_manager._tools
|
||||
assert "recall" in tools
|
||||
|
||||
# Call recall with new params
|
||||
recall_tool = tools["recall"]
|
||||
result = await recall_tool.fn(query="test query", max_tokens=2048, budget="mid")
|
||||
|
||||
# Result is a dict
|
||||
assert isinstance(result, dict)
|
||||
|
||||
# Verify the memory was called correctly
|
||||
mock_memory.recall_async.assert_called_once()
|
||||
call_kwargs = mock_memory.recall_async.call_args.kwargs
|
||||
assert call_kwargs["bank_id"] == "test-bank"
|
||||
assert call_kwargs["query"] == "test query"
|
||||
assert call_kwargs["max_tokens"] == 2048
|
||||
assert call_kwargs["budget"] == Budget.MID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_mcp_server_retain_with_default_context(mock_memory):
|
||||
"""Test that retain uses default context when not provided."""
|
||||
from hindsight_api.mcp_local import create_local_mcp_server
|
||||
|
||||
bank_id = "test-bank"
|
||||
mcp_server = create_local_mcp_server(bank_id, memory=mock_memory)
|
||||
|
||||
tools = mcp_server._tool_manager._tools
|
||||
retain_tool = tools["retain"]
|
||||
|
||||
# Call retain without context
|
||||
await retain_tool.fn(content="test content")
|
||||
|
||||
# Wait for background task
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
call_kwargs = mock_memory.retain_batch_async.call_args.kwargs
|
||||
assert call_kwargs["contents"] == [{"content": "test content", "context": "general"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_mcp_server_retain_error_handling(mock_memory):
|
||||
"""Test that retain errors are logged but don't affect response."""
|
||||
from hindsight_api.mcp_local import create_local_mcp_server
|
||||
|
||||
mock_memory.retain_batch_async = AsyncMock(side_effect=Exception("Test error"))
|
||||
|
||||
mcp_server = create_local_mcp_server("test-bank", memory=mock_memory)
|
||||
|
||||
tools = mcp_server._tool_manager._tools
|
||||
retain_tool = tools["retain"]
|
||||
|
||||
# Retain returns immediately with accepted status (fire and forget)
|
||||
result = await retain_tool.fn(content="test content")
|
||||
assert result["status"] == "accepted"
|
||||
|
||||
# Wait for background task to complete (and log error)
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_mcp_server_recall_error_handling(mock_memory):
|
||||
"""Test that recall handles errors gracefully."""
|
||||
from hindsight_api.mcp_local import create_local_mcp_server
|
||||
|
||||
mock_memory.recall_async = AsyncMock(side_effect=Exception("Test error"))
|
||||
|
||||
mcp_server = create_local_mcp_server("test-bank", memory=mock_memory)
|
||||
|
||||
tools = mcp_server._tool_manager._tools
|
||||
recall_tool = tools["recall"]
|
||||
|
||||
result = await recall_tool.fn(query="test query")
|
||||
|
||||
# Result is a dict with error
|
||||
assert isinstance(result, dict)
|
||||
assert "error" in result
|
||||
assert result["results"] == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_mcp_server_recall_with_defaults(mock_memory):
|
||||
"""Test that recall uses default max_tokens and budget."""
|
||||
from hindsight_api.mcp_local import create_local_mcp_server
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.model_dump.return_value = {"results": []}
|
||||
mock_memory.recall_async = AsyncMock(return_value=mock_result)
|
||||
|
||||
mcp_server = create_local_mcp_server("test-bank", memory=mock_memory)
|
||||
|
||||
tools = mcp_server._tool_manager._tools
|
||||
recall_tool = tools["recall"]
|
||||
|
||||
# Call with defaults
|
||||
await recall_tool.fn(query="test query")
|
||||
|
||||
call_kwargs = mock_memory.recall_async.call_args.kwargs
|
||||
assert call_kwargs["max_tokens"] == 4096
|
||||
assert call_kwargs["budget"] == Budget.LOW
|
||||
@@ -8,7 +8,7 @@ from unittest.mock import AsyncMock, MagicMock
|
||||
def mock_memory():
|
||||
"""Create a mock MemoryEngine."""
|
||||
memory = MagicMock()
|
||||
memory.put_batch_async = AsyncMock()
|
||||
memory.retain_batch_async = AsyncMock()
|
||||
memory.recall_async = AsyncMock(return_value=MagicMock(results=[]))
|
||||
return memory
|
||||
|
||||
@@ -52,8 +52,8 @@ async def test_mcp_tools_use_context_bank_id(mock_memory):
|
||||
assert "successfully" in result.lower()
|
||||
|
||||
# Verify the memory was called with the context bank_id
|
||||
mock_memory.put_batch_async.assert_called_once()
|
||||
call_kwargs = mock_memory.put_batch_async.call_args.kwargs
|
||||
mock_memory.retain_batch_async.assert_called_once()
|
||||
call_kwargs = mock_memory.retain_batch_async.call_args.kwargs
|
||||
assert call_kwargs["bank_id"] == "context-bank-id"
|
||||
finally:
|
||||
_current_bank_id.reset(token)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "hindsight-cli"
|
||||
version = "0.1.5"
|
||||
version = "0.1.8"
|
||||
edition = "2021"
|
||||
authors = ["Hindsight Team"]
|
||||
description = "A beautiful CLI for Hindsight - semantic memory system"
|
||||
|
||||
@@ -91,6 +91,9 @@ enum Commands {
|
||||
#[command(alias = "tui")]
|
||||
Explore,
|
||||
|
||||
/// Launch the web-based control plane UI
|
||||
Ui,
|
||||
|
||||
/// Configure the CLI (API URL, etc.)
|
||||
#[command(after_help = "Configuration priority:\n 1. Environment variable (HINDSIGHT_API_URL) - highest priority\n 2. Config file (~/.hindsight/config)\n 3. Default (http://localhost:8888)")]
|
||||
Configure {
|
||||
@@ -373,6 +376,11 @@ fn run() -> Result<()> {
|
||||
return handle_configure(api_url, output_format);
|
||||
}
|
||||
|
||||
// Handle ui command - needs config but not API client
|
||||
if let Commands::Ui = cli.command {
|
||||
return handle_ui(output_format);
|
||||
}
|
||||
|
||||
// Load configuration
|
||||
let config = Config::from_env().unwrap_or_else(|e| {
|
||||
ui::print_error(&format!("Configuration error: {}", e));
|
||||
@@ -390,6 +398,7 @@ fn run() -> Result<()> {
|
||||
// Execute command and handle errors
|
||||
let result: Result<()> = match cli.command {
|
||||
Commands::Configure { .. } => unreachable!(), // Handled above
|
||||
Commands::Ui => unreachable!(), // Handled above
|
||||
Commands::Explore => commands::explore::run(&client),
|
||||
Commands::Bank(bank_cmd) => match bank_cmd {
|
||||
BankCommands::List => commands::bank::list(&client, verbose, output_format),
|
||||
@@ -521,3 +530,50 @@ fn handle_configure(api_url: Option<String>, output_format: OutputFormat) -> Res
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn handle_ui(output_format: OutputFormat) -> Result<()> {
|
||||
use std::process::Command;
|
||||
|
||||
// Load configuration to get the API URL
|
||||
let config = Config::load().unwrap_or_else(|e| {
|
||||
ui::print_error(&format!("Configuration error: {}", e));
|
||||
errors::print_config_help();
|
||||
std::process::exit(1);
|
||||
});
|
||||
|
||||
let api_url = config.api_url();
|
||||
|
||||
if output_format == OutputFormat::Pretty {
|
||||
ui::print_info("Launching Hindsight Control Plane UI...");
|
||||
println!();
|
||||
println!(" API URL: {}", api_url);
|
||||
println!();
|
||||
}
|
||||
|
||||
// Run npx @vectorize-io/hindsight-control-plane --api-url {api_url}
|
||||
let status = Command::new("npx")
|
||||
.arg("@vectorize-io/hindsight-control-plane")
|
||||
.arg("--api-url")
|
||||
.arg(api_url)
|
||||
.status();
|
||||
|
||||
match status {
|
||||
Ok(exit_status) => {
|
||||
if !exit_status.success() {
|
||||
if let Some(code) = exit_status.code() {
|
||||
std::process::exit(code);
|
||||
} else {
|
||||
std::process::exit(1);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
ui::print_error(&format!("Failed to launch control plane UI: {}", e));
|
||||
ui::print_info("Make sure you have Node.js and npm installed.");
|
||||
ui::print_info("You can also install the control plane globally: npm install -g @vectorize-io/hindsight-control-plane");
|
||||
std::process::exit(1);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "hindsight-client"
|
||||
version = "0.1.5"
|
||||
version = "0.1.8"
|
||||
description = "Python client for Hindsight - Semantic memory system with personality-driven thinking"
|
||||
authors = [
|
||||
{name = "Hindsight Team"}
|
||||
@@ -20,6 +20,7 @@ dependencies = [
|
||||
test = [
|
||||
"pytest>=7.0.0",
|
||||
"pytest-asyncio>=0.21.0",
|
||||
"requests>=2.28.0",
|
||||
]
|
||||
|
||||
[build-system]
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@vectorize-io/hindsight-client",
|
||||
"version": "0.1.5",
|
||||
"version": "0.1.8",
|
||||
"description": "TypeScript client for Hindsight - Semantic memory system with personality-driven thinking",
|
||||
"main": "./dist/src/index.js",
|
||||
"types": "./dist/src/index.d.ts",
|
||||
|
||||
@@ -14,6 +14,7 @@
|
||||
|
||||
# production
|
||||
/build
|
||||
/standalone
|
||||
|
||||
# misc
|
||||
.DS_Store
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
{
|
||||
"semi": true,
|
||||
"singleQuote": false,
|
||||
"tabWidth": 2,
|
||||
"trailingComma": "es5",
|
||||
"printWidth": 100
|
||||
}
|
||||
Executable
+86
@@ -0,0 +1,86 @@
|
||||
#!/usr/bin/env node
|
||||
|
||||
const { spawn } = require('child_process');
|
||||
const path = require('path');
|
||||
const fs = require('fs');
|
||||
|
||||
const args = process.argv.slice(2);
|
||||
|
||||
// Parse command line arguments
|
||||
let port = process.env.PORT || 9999;
|
||||
let hostname = process.env.HOSTNAME || '0.0.0.0';
|
||||
let apiUrl = process.env.HINDSIGHT_CP_DATAPLANE_API_URL;
|
||||
|
||||
for (let i = 0; i < args.length; i++) {
|
||||
if (args[i] === '--port' || args[i] === '-p') {
|
||||
port = args[++i];
|
||||
} else if (args[i] === '--hostname' || args[i] === '-H') {
|
||||
hostname = args[++i];
|
||||
} else if (args[i] === '--api-url' || args[i] === '-a') {
|
||||
apiUrl = args[++i];
|
||||
} else if (args[i] === '--help' || args[i] === '-h') {
|
||||
console.log(`
|
||||
Hindsight Control Plane
|
||||
|
||||
Usage: hindsight-control-plane [options]
|
||||
|
||||
Options:
|
||||
-p, --port <port> Port to listen on (default: 9999, env: PORT)
|
||||
-H, --hostname <host> Hostname to bind to (default: 0.0.0.0, env: HOSTNAME)
|
||||
-a, --api-url <url> Hindsight API URL (env: HINDSIGHT_CP_DATAPLANE_API_URL)
|
||||
-h, --help Show this help message
|
||||
|
||||
Environment Variables:
|
||||
PORT Port to listen on
|
||||
HOSTNAME Hostname to bind to
|
||||
HINDSIGHT_CP_DATAPLANE_API_URL URL of the Hindsight API server
|
||||
`);
|
||||
process.exit(0);
|
||||
}
|
||||
}
|
||||
|
||||
// Find the standalone server
|
||||
const standaloneDir = path.join(__dirname, '..', 'standalone');
|
||||
const serverPath = path.join(standaloneDir, 'server.js');
|
||||
|
||||
if (!fs.existsSync(serverPath)) {
|
||||
console.error('Error: Standalone server not found at', serverPath);
|
||||
console.error('This package may not have been built correctly.');
|
||||
process.exit(1);
|
||||
}
|
||||
|
||||
// Set up environment
|
||||
const env = {
|
||||
...process.env,
|
||||
PORT: String(port),
|
||||
HOSTNAME: hostname,
|
||||
};
|
||||
|
||||
if (apiUrl) {
|
||||
env.HINDSIGHT_CP_DATAPLANE_API_URL = apiUrl;
|
||||
}
|
||||
|
||||
console.log(`Starting Hindsight Control Plane on http://${hostname}:${port}`);
|
||||
if (apiUrl) {
|
||||
console.log(`API URL: ${apiUrl}`);
|
||||
}
|
||||
|
||||
// Run the standalone server
|
||||
const server = spawn('node', [serverPath], {
|
||||
cwd: standaloneDir,
|
||||
env,
|
||||
stdio: 'inherit',
|
||||
});
|
||||
|
||||
server.on('error', (err) => {
|
||||
console.error('Failed to start server:', err.message);
|
||||
process.exit(1);
|
||||
});
|
||||
|
||||
server.on('close', (code) => {
|
||||
process.exit(code || 0);
|
||||
});
|
||||
|
||||
// Handle signals
|
||||
process.on('SIGTERM', () => server.kill('SIGTERM'));
|
||||
process.on('SIGINT', () => server.kill('SIGINT'));
|
||||
@@ -0,0 +1,37 @@
|
||||
import js from "@eslint/js";
|
||||
import tseslint from "typescript-eslint";
|
||||
import reactPlugin from "eslint-plugin-react";
|
||||
import reactHooksPlugin from "eslint-plugin-react-hooks";
|
||||
|
||||
export default [
|
||||
js.configs.recommended,
|
||||
...tseslint.configs.recommended,
|
||||
{
|
||||
files: ["**/*.{ts,tsx}"],
|
||||
plugins: {
|
||||
react: reactPlugin,
|
||||
"react-hooks": reactHooksPlugin,
|
||||
},
|
||||
languageOptions: {
|
||||
parserOptions: {
|
||||
ecmaFeatures: {
|
||||
jsx: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
rules: {
|
||||
"@typescript-eslint/no-unused-vars": "warn",
|
||||
"@typescript-eslint/no-explicit-any": "warn",
|
||||
"react/react-in-jsx-scope": "off",
|
||||
"no-case-declarations": "off",
|
||||
},
|
||||
settings: {
|
||||
react: {
|
||||
version: "detect",
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
ignores: [".next/", "node_modules/"],
|
||||
},
|
||||
];
|
||||
@@ -1,7 +1,14 @@
|
||||
import type { NextConfig } from "next";
|
||||
import path from "path";
|
||||
|
||||
const nextConfig: NextConfig = {
|
||||
output: 'standalone',
|
||||
// Disable request logging in production
|
||||
logging: false,
|
||||
// Set the monorepo root explicitly to avoid detecting wrong lockfiles in parent directories
|
||||
turbopack: {
|
||||
root: path.resolve(__dirname, '..'),
|
||||
},
|
||||
};
|
||||
|
||||
export default nextConfig;
|
||||
|
||||
@@ -1,17 +1,26 @@
|
||||
{
|
||||
"name": "hindsight-control-plane",
|
||||
"version": "0.1.5",
|
||||
"private": true,
|
||||
"name": "@vectorize-io/hindsight-control-plane",
|
||||
"version": "0.1.8",
|
||||
"description": "Control plane for Hindsight - Semantic memory system",
|
||||
"bin": {
|
||||
"hindsight-control-plane": "./bin/cli.js"
|
||||
},
|
||||
"files": [
|
||||
"bin",
|
||||
"standalone",
|
||||
"public"
|
||||
],
|
||||
"scripts": {
|
||||
"dev": "next dev",
|
||||
"build": "next build",
|
||||
"build": "next build && npm run build:standalone",
|
||||
"build:standalone": "rm -rf standalone && STANDALONE_ROOT=$(dirname $(find .next/standalone -name 'server.js' | head -1)) && cp -r \"$STANDALONE_ROOT\" standalone && mkdir -p standalone/.next && cp -r .next/static standalone/.next/static && mkdir -p standalone/public && cp -r public/* standalone/public/ 2>/dev/null || true",
|
||||
"start": "next start",
|
||||
"lint": "next lint"
|
||||
"lint": "next lint",
|
||||
"prepublishOnly": "npm run build"
|
||||
},
|
||||
"keywords": [],
|
||||
"keywords": ["hindsight", "memory", "semantic", "ai"],
|
||||
"author": "Hindsight Team",
|
||||
"license": "ISC",
|
||||
"description": "Control plane for Hindsight - Semantic memory system",
|
||||
"dependencies": {
|
||||
"@radix-ui/react-checkbox": "^1.3.3",
|
||||
"@radix-ui/react-dialog": "^1.1.15",
|
||||
@@ -27,7 +36,6 @@
|
||||
"@types/node": "^24.10.0",
|
||||
"@types/react": "^19.2.2",
|
||||
"@types/react-dom": "^19.2.2",
|
||||
"@vectorize-io/hindsight-client": "file:../hindsight-clients/typescript",
|
||||
"autoprefixer": "^10.4.21",
|
||||
"class-variance-authority": "^0.7.1",
|
||||
"clsx": "^2.1.1",
|
||||
@@ -48,5 +56,14 @@
|
||||
"tailwindcss-animate": "^1.0.7",
|
||||
"three": "^0.182.0",
|
||||
"typescript": "^5.9.3"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@vectorize-io/hindsight-client": "file:../hindsight-clients/typescript",
|
||||
"@eslint/eslintrc": "^3.3.3",
|
||||
"@eslint/js": "^9.39.2",
|
||||
"eslint-plugin-react": "^7.37.5",
|
||||
"eslint-plugin-react-hooks": "^7.0.1",
|
||||
"prettier": "^3.7.4",
|
||||
"typescript-eslint": "^8.50.0"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,16 +1,13 @@
|
||||
import { NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET() {
|
||||
try {
|
||||
const response = await sdk.listBanks({ client: lowLevelClient });
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error fetching banks:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to fetch banks' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error fetching banks:", error);
|
||||
return NextResponse.json({ error: "Failed to fetch banks" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,10 +17,7 @@ export async function POST(request: Request) {
|
||||
const { bank_id } = body;
|
||||
|
||||
if (!bank_id) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const response = await sdk.createOrUpdateBank({
|
||||
@@ -34,10 +28,7 @@ export async function POST(request: Request) {
|
||||
|
||||
return NextResponse.json(response.data, { status: 201 });
|
||||
} catch (error) {
|
||||
console.error('Error creating bank:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to create bank' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error creating bank:", error);
|
||||
return NextResponse.json({ error: "Failed to create bank" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET(
|
||||
request: NextRequest,
|
||||
@@ -10,15 +10,12 @@ export async function GET(
|
||||
|
||||
const response = await sdk.getChunk({
|
||||
client: lowLevelClient,
|
||||
path: { chunk_id: chunkId }
|
||||
path: { chunk_id: chunkId },
|
||||
});
|
||||
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error fetching chunk:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to fetch chunk' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error fetching chunk:", error);
|
||||
return NextResponse.json({ error: "Failed to fetch chunk" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET(
|
||||
request: NextRequest,
|
||||
@@ -8,26 +8,20 @@ export async function GET(
|
||||
try {
|
||||
const { documentId } = await params;
|
||||
const searchParams = request.nextUrl.searchParams;
|
||||
const bankId = searchParams.get('bank_id');
|
||||
const bankId = searchParams.get("bank_id");
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const response = await sdk.getDocument({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: bankId, document_id: documentId }
|
||||
path: { bank_id: bankId, document_id: documentId },
|
||||
});
|
||||
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error fetching document:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to fetch document' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error fetching document:", error);
|
||||
return NextResponse.json({ error: "Failed to fetch document" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,33 +1,27 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET(request: NextRequest) {
|
||||
try {
|
||||
const searchParams = request.nextUrl.searchParams;
|
||||
const bankId = searchParams.get('bank_id');
|
||||
const bankId = searchParams.get("bank_id");
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const limit = searchParams.get('limit') ? Number(searchParams.get('limit')) : undefined;
|
||||
const offset = searchParams.get('offset') ? Number(searchParams.get('offset')) : undefined;
|
||||
const limit = searchParams.get("limit") ? Number(searchParams.get("limit")) : undefined;
|
||||
const offset = searchParams.get("offset") ? Number(searchParams.get("offset")) : undefined;
|
||||
|
||||
const response = await sdk.listDocuments({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: bankId },
|
||||
query: { limit, offset }
|
||||
query: { limit, offset },
|
||||
});
|
||||
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error fetching documents:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to fetch documents' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error fetching documents:", error);
|
||||
return NextResponse.json({ error: "Failed to fetch documents" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function POST(
|
||||
request: NextRequest,
|
||||
@@ -8,13 +8,10 @@ export async function POST(
|
||||
try {
|
||||
const { entityId } = await params;
|
||||
const searchParams = request.nextUrl.searchParams;
|
||||
const bankId = searchParams.get('bank_id');
|
||||
const bankId = searchParams.get("bank_id");
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const decodedEntityId = decodeURIComponent(entityId);
|
||||
@@ -23,15 +20,15 @@ export async function POST(
|
||||
client: lowLevelClient,
|
||||
path: {
|
||||
bank_id: bankId,
|
||||
entity_id: decodedEntityId
|
||||
}
|
||||
entity_id: decodedEntityId,
|
||||
},
|
||||
});
|
||||
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error regenerating entity observations:', error);
|
||||
console.error("Error regenerating entity observations:", error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to regenerate entity observations' },
|
||||
{ error: "Failed to regenerate entity observations" },
|
||||
{ status: 500 }
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET(
|
||||
request: NextRequest,
|
||||
@@ -8,13 +8,10 @@ export async function GET(
|
||||
try {
|
||||
const { entityId } = await params;
|
||||
const searchParams = request.nextUrl.searchParams;
|
||||
const bankId = searchParams.get('bank_id');
|
||||
const bankId = searchParams.get("bank_id");
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
// Decode URL-encoded entityId in case it contains special chars
|
||||
@@ -24,23 +21,17 @@ export async function GET(
|
||||
client: lowLevelClient,
|
||||
path: {
|
||||
bank_id: bankId,
|
||||
entity_id: decodedEntityId
|
||||
}
|
||||
entity_id: decodedEntityId,
|
||||
},
|
||||
});
|
||||
|
||||
if (response.error) {
|
||||
return NextResponse.json(
|
||||
{ error: response.error },
|
||||
{ status: 500 }
|
||||
);
|
||||
return NextResponse.json({ error: response.error }, { status: 500 });
|
||||
}
|
||||
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error getting entity:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to get entity' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error getting entity:", error);
|
||||
return NextResponse.json({ error: "Failed to get entity" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,39 +1,30 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET(request: NextRequest) {
|
||||
try {
|
||||
const searchParams = request.nextUrl.searchParams;
|
||||
const bankId = searchParams.get('bank_id');
|
||||
const bankId = searchParams.get("bank_id");
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const limit = searchParams.get('limit') ? Number(searchParams.get('limit')) : undefined;
|
||||
const limit = searchParams.get("limit") ? Number(searchParams.get("limit")) : undefined;
|
||||
|
||||
const response = await sdk.listEntities({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: bankId },
|
||||
query: { limit }
|
||||
query: { limit },
|
||||
});
|
||||
|
||||
if (response.error) {
|
||||
return NextResponse.json(
|
||||
{ error: response.error },
|
||||
{ status: 500 }
|
||||
);
|
||||
return NextResponse.json({ error: response.error }, { status: 500 });
|
||||
}
|
||||
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error listing entities:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to list entities' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error listing entities:", error);
|
||||
return NextResponse.json({ error: "Failed to list entities" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,35 +1,29 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET(request: NextRequest) {
|
||||
try {
|
||||
const searchParams = request.nextUrl.searchParams;
|
||||
const bankId = searchParams.get('bank_id') || searchParams.get('agent_id');
|
||||
const bankId = searchParams.get("bank_id") || searchParams.get("agent_id");
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
// Get optional query parameters
|
||||
const type = searchParams.get('type') || searchParams.get('fact_type') || undefined;
|
||||
const type = searchParams.get("type") || searchParams.get("fact_type") || undefined;
|
||||
|
||||
const response = await sdk.getGraph({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: bankId },
|
||||
query: {
|
||||
type: type
|
||||
}
|
||||
type: type,
|
||||
},
|
||||
});
|
||||
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error fetching graph data:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to fetch graph data' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error fetching graph data:", error);
|
||||
return NextResponse.json({ error: "Failed to fetch graph data" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
import { NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET() {
|
||||
const status: {
|
||||
status: string;
|
||||
service: string;
|
||||
dataplane?: {
|
||||
status: string;
|
||||
url: string;
|
||||
error?: string;
|
||||
};
|
||||
} = {
|
||||
status: "ok",
|
||||
service: "hindsight-control-plane",
|
||||
};
|
||||
|
||||
// Check dataplane connectivity
|
||||
const dataplaneUrl = process.env.HINDSIGHT_CP_DATAPLANE_API_URL || "http://localhost:8888";
|
||||
try {
|
||||
await sdk.listBanks({ client: lowLevelClient });
|
||||
status.dataplane = {
|
||||
status: "connected",
|
||||
url: dataplaneUrl,
|
||||
};
|
||||
} catch (error) {
|
||||
status.dataplane = {
|
||||
status: "disconnected",
|
||||
url: dataplaneUrl,
|
||||
error: error instanceof Error ? error.message : String(error),
|
||||
};
|
||||
}
|
||||
|
||||
return NextResponse.json(status, { status: 200 });
|
||||
}
|
||||
@@ -1,37 +1,31 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { hindsightClient, sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { hindsightClient, sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET(request: NextRequest) {
|
||||
try {
|
||||
const searchParams = request.nextUrl.searchParams;
|
||||
const bankId = searchParams.get('bank_id') || searchParams.get('agent_id');
|
||||
const bankId = searchParams.get("bank_id") || searchParams.get("agent_id");
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const limit = searchParams.get('limit') ? Number(searchParams.get('limit')) : undefined;
|
||||
const offset = searchParams.get('offset') ? Number(searchParams.get('offset')) : undefined;
|
||||
const type = searchParams.get('type') || searchParams.get('fact_type') || undefined;
|
||||
const q = searchParams.get('q') || undefined;
|
||||
const limit = searchParams.get("limit") ? Number(searchParams.get("limit")) : undefined;
|
||||
const offset = searchParams.get("offset") ? Number(searchParams.get("offset")) : undefined;
|
||||
const type = searchParams.get("type") || searchParams.get("fact_type") || undefined;
|
||||
const q = searchParams.get("q") || undefined;
|
||||
|
||||
const response = await hindsightClient.listMemories(bankId, {
|
||||
limit,
|
||||
offset,
|
||||
type,
|
||||
q
|
||||
q,
|
||||
});
|
||||
|
||||
return NextResponse.json(response, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error listing memory units:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to list memory units' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error listing memory units:", error);
|
||||
return NextResponse.json({ error: "Failed to list memory units" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -39,7 +33,10 @@ export async function GET(request: NextRequest) {
|
||||
// Use clearBankMemories to delete all memories for a bank instead
|
||||
export async function DELETE(request: NextRequest) {
|
||||
return NextResponse.json(
|
||||
{ error: 'Individual memory unit deletion is not yet supported. Use clear all memories instead.' },
|
||||
{
|
||||
error:
|
||||
"Individual memory unit deletion is not yet supported. Use clear all memories instead.",
|
||||
},
|
||||
{ status: 501 } // Not Implemented
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { hindsightClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { hindsightClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function POST(request: NextRequest) {
|
||||
try {
|
||||
@@ -7,10 +7,7 @@ export async function POST(request: NextRequest) {
|
||||
const bankId = body.bank_id || body.agent_id;
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const { items, document_id } = body;
|
||||
@@ -19,10 +16,7 @@ export async function POST(request: NextRequest) {
|
||||
|
||||
return NextResponse.json(response, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error batch retain:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to batch retain' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error batch retain:", error);
|
||||
return NextResponse.json({ error: "Failed to batch retain" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function POST(request: NextRequest) {
|
||||
try {
|
||||
@@ -7,10 +7,7 @@ export async function POST(request: NextRequest) {
|
||||
const bankId = body.bank_id || body.agent_id;
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'bank_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const { items } = body;
|
||||
@@ -18,15 +15,12 @@ export async function POST(request: NextRequest) {
|
||||
const response = await sdk.retainMemories({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: bankId },
|
||||
body: { items, async: true }
|
||||
body: { items, async: true },
|
||||
});
|
||||
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error batch retain async:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to batch retain async' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error batch retain async:", error);
|
||||
return NextResponse.json({ error: "Failed to batch retain async" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET(
|
||||
request: NextRequest,
|
||||
@@ -9,15 +9,12 @@ export async function GET(
|
||||
const { agentId } = await params;
|
||||
const response = await sdk.listOperations({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: agentId }
|
||||
path: { bank_id: agentId },
|
||||
});
|
||||
return NextResponse.json(response.data || {}, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error fetching operations:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to fetch operations' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error fetching operations:", error);
|
||||
return NextResponse.json({ error: "Failed to fetch operations" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -28,26 +25,20 @@ export async function DELETE(
|
||||
try {
|
||||
const { agentId } = await params;
|
||||
const searchParams = request.nextUrl.searchParams;
|
||||
const operationId = searchParams.get('operation_id');
|
||||
const operationId = searchParams.get("operation_id");
|
||||
|
||||
if (!operationId) {
|
||||
return NextResponse.json(
|
||||
{ error: 'operation_id is required' },
|
||||
{ status: 400 }
|
||||
);
|
||||
return NextResponse.json({ error: "operation_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const response = await sdk.cancelOperation({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: agentId, operation_id: operationId }
|
||||
path: { bank_id: agentId, operation_id: operationId },
|
||||
});
|
||||
|
||||
return NextResponse.json(response.data || {}, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error canceling operation:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to cancel operation' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error canceling operation:", error);
|
||||
return NextResponse.json({ error: "Failed to cancel operation" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function GET(
|
||||
request: NextRequest,
|
||||
@@ -9,15 +9,12 @@ export async function GET(
|
||||
const { bankId } = await params;
|
||||
const response = await sdk.getBankProfile({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: bankId }
|
||||
path: { bank_id: bankId },
|
||||
});
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error fetching bank profile:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to fetch bank profile' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error fetching bank profile:", error);
|
||||
return NextResponse.json({ error: "Failed to fetch bank profile" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -32,14 +29,11 @@ export async function PUT(
|
||||
const response = await sdk.createOrUpdateBank({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: bankId },
|
||||
body: body
|
||||
body: body,
|
||||
});
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error updating bank profile:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to update bank profile' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error updating bank profile:", error);
|
||||
return NextResponse.json({ error: "Failed to update bank profile" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,15 +1,12 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { lowLevelClient, sdk } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { lowLevelClient, sdk } from "@/lib/hindsight-client";
|
||||
|
||||
export async function POST(request: NextRequest) {
|
||||
try {
|
||||
const body = await request.json();
|
||||
const bankId = body.bank_id || body.agent_id || 'default';
|
||||
const bankId = body.bank_id || body.agent_id || "default";
|
||||
const { query, types, fact_type, max_tokens, trace, budget, include, query_timestamp } = body;
|
||||
|
||||
console.log('[Recall API] Request:', { bankId, query, types: types || fact_type, max_tokens, trace, budget, query_timestamp });
|
||||
console.log('[Recall API] Include options:', JSON.stringify(include, null, 2));
|
||||
|
||||
const response = await sdk.recallMemories({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: bankId },
|
||||
@@ -18,29 +15,17 @@ export async function POST(request: NextRequest) {
|
||||
types: types || fact_type,
|
||||
max_tokens,
|
||||
trace,
|
||||
budget: budget || 'mid',
|
||||
budget: budget || "mid",
|
||||
include,
|
||||
query_timestamp,
|
||||
},
|
||||
});
|
||||
|
||||
if (!response.data) {
|
||||
console.error('[Recall API] No data in response', { response, error: response.error });
|
||||
throw new Error(`API returned no data: ${JSON.stringify(response.error || 'Unknown error')}`);
|
||||
console.error("[Recall API] No data in response", { response, error: response.error });
|
||||
throw new Error(`API returned no data: ${JSON.stringify(response.error || "Unknown error")}`);
|
||||
}
|
||||
|
||||
console.log('[Recall API] Response structure:', {
|
||||
hasResults: !!response.data?.results,
|
||||
resultsCount: response.data?.results?.length,
|
||||
hasTrace: !!response.data?.trace,
|
||||
hasEntities: !!response.data?.entities,
|
||||
entitiesType: typeof response.data?.entities,
|
||||
entitiesKeys: response.data?.entities ? Object.keys(response.data.entities) : null,
|
||||
hasChunks: !!response.data?.chunks,
|
||||
chunksType: typeof response.data?.chunks,
|
||||
chunksKeys: response.data?.chunks ? Object.keys(response.data.chunks) : null,
|
||||
});
|
||||
|
||||
// Return a clean JSON object by spreading the response
|
||||
// This ensures any non-serializable properties are excluded
|
||||
const jsonResponse = {
|
||||
@@ -52,10 +37,7 @@ export async function POST(request: NextRequest) {
|
||||
|
||||
return NextResponse.json(jsonResponse, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error recalling:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to recall' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error recalling:", error);
|
||||
return NextResponse.json({ error: "Failed to recall" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,37 +1,34 @@
|
||||
import { NextRequest, NextResponse } from 'next/server';
|
||||
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
|
||||
|
||||
export async function POST(request: NextRequest) {
|
||||
try {
|
||||
const body = await request.json();
|
||||
const bankId = body.bank_id || body.agent_id || 'default';
|
||||
const bankId = body.bank_id || body.agent_id || "default";
|
||||
const { query, context, budget, thinking_budget, include_facts } = body;
|
||||
|
||||
const requestBody: any = {
|
||||
query,
|
||||
budget: budget || (thinking_budget ? 'mid' : 'low'),
|
||||
context: context || undefined
|
||||
budget: budget || (thinking_budget ? "mid" : "low"),
|
||||
context: context || undefined,
|
||||
};
|
||||
|
||||
// Add include options if specified
|
||||
if (include_facts) {
|
||||
requestBody.include = {
|
||||
facts: {}
|
||||
facts: {},
|
||||
};
|
||||
}
|
||||
|
||||
const response = await sdk.reflect({
|
||||
client: lowLevelClient,
|
||||
path: { bank_id: bankId },
|
||||
body: requestBody
|
||||
body: requestBody,
|
||||
});
|
||||
|
||||
return NextResponse.json(response.data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error('Error reflecting:', error);
|
||||
return NextResponse.json(
|
||||
{ error: 'Failed to reflect' },
|
||||
{ status: 500 }
|
||||
);
|
||||
console.error("Error reflecting:", error);
|
||||
return NextResponse.json({ error: "Failed to reflect" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user