Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f515b8207b | ||
|
|
0e3ccd038b |
+1
-6
@@ -2,7 +2,7 @@
|
||||
# Copy this file to .env and fill in your values
|
||||
|
||||
# LLM Configuration (Required)
|
||||
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio, vertexai, minimax
|
||||
# Supported providers: openai, groq, ollama, gemini, anthropic, lmstudio, vertexai
|
||||
HINDSIGHT_API_LLM_PROVIDER=openai
|
||||
HINDSIGHT_API_LLM_API_KEY=your-api-key-here
|
||||
HINDSIGHT_API_LLM_MODEL=gpt-4o-mini
|
||||
@@ -20,11 +20,6 @@ HINDSIGHT_API_LLM_BASE_URL=https://api.openai.com/v1
|
||||
# HINDSIGHT_API_LLM_VERTEXAI_REGION=us-central1
|
||||
# HINDSIGHT_API_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY=/path/to/service-account-key.json # Optional, uses ADC if not set
|
||||
|
||||
# Example: MiniMax configuration (204K context window)
|
||||
# HINDSIGHT_API_LLM_PROVIDER=minimax
|
||||
# HINDSIGHT_API_LLM_API_KEY=your-minimax-api-key
|
||||
# HINDSIGHT_API_LLM_MODEL=MiniMax-M2.5
|
||||
|
||||
# Example: LM Studio local configuration (Qwen 2.5 32B recommended)
|
||||
# HINDSIGHT_API_LLM_PROVIDER=lmstudio
|
||||
# HINDSIGHT_API_LLM_API_KEY=lmstudio
|
||||
|
||||
@@ -1,6 +0,0 @@
|
||||
version: 2
|
||||
updates:
|
||||
- package-ecosystem: "github-actions"
|
||||
directory: "/"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
@@ -21,20 +21,20 @@ jobs:
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/setup-node@v6
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: 20
|
||||
cache: npm
|
||||
cache-dependency-path: package-lock.json
|
||||
- uses: astral-sh/setup-uv@v7
|
||||
- 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
|
||||
env:
|
||||
UMAMI_URL: https://analytics.hindsight.vectorize.io
|
||||
UMAMI_WEBSITE_ID: ${{ secrets.UMAMI_WEBSITE_ID }}
|
||||
- uses: actions/upload-pages-artifact@v4
|
||||
- uses: actions/upload-pages-artifact@v3
|
||||
with:
|
||||
path: hindsight-docs/build
|
||||
deploy:
|
||||
|
||||
@@ -13,15 +13,15 @@ jobs:
|
||||
id-token: write
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
@@ -30,20 +30,12 @@ jobs:
|
||||
working-directory: ./hindsight-clients/python
|
||||
run: uv build --out-dir dist
|
||||
|
||||
- name: Build hindsight-api-slim
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv build --out-dir dist
|
||||
|
||||
- name: Build hindsight-api
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build --out-dir dist
|
||||
|
||||
- name: Build hindsight-all
|
||||
working-directory: ./hindsight-all
|
||||
run: uv build --out-dir dist
|
||||
|
||||
- name: Build hindsight-all-slim
|
||||
working-directory: ./hindsight-all-slim
|
||||
working-directory: ./hindsight
|
||||
run: uv build --out-dir dist
|
||||
|
||||
- name: Build hindsight-litellm
|
||||
@@ -62,19 +54,13 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/pydantic-ai
|
||||
run: uv build --out-dir dist
|
||||
|
||||
# Publish in order (client and api-slim first, then api/all wrappers which depend on them)
|
||||
# Publish in order (client and api first, then hindsight-all which depends on them)
|
||||
- name: Publish hindsight-client to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: ./hindsight-clients/python/dist
|
||||
skip-existing: true
|
||||
|
||||
- name: Publish hindsight-api-slim to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: ./hindsight-api-slim/dist
|
||||
skip-existing: true
|
||||
|
||||
- name: Publish hindsight-api to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
@@ -84,13 +70,7 @@ jobs:
|
||||
- name: Publish hindsight-all to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: ./hindsight-all/dist
|
||||
skip-existing: true
|
||||
|
||||
- name: Publish hindsight-all-slim to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: ./hindsight-all-slim/dist
|
||||
packages-dir: ./hindsight/dist
|
||||
skip-existing: true
|
||||
|
||||
- name: Publish hindsight-litellm to PyPI
|
||||
@@ -119,15 +99,13 @@ jobs:
|
||||
|
||||
# Upload artifacts for GitHub release
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v7
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: python-packages
|
||||
path: |
|
||||
hindsight-clients/python/dist/*
|
||||
hindsight-api-slim/dist/*
|
||||
hindsight-api/dist/*
|
||||
hindsight-all/dist/*
|
||||
hindsight-all-slim/dist/*
|
||||
hindsight/dist/*
|
||||
hindsight-integrations/litellm/dist/*
|
||||
hindsight-embed/dist/*
|
||||
hindsight-integrations/crewai/dist/*
|
||||
@@ -139,10 +117,10 @@ jobs:
|
||||
environment: npm
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
registry-url: 'https://registry.npmjs.org'
|
||||
@@ -177,7 +155,7 @@ jobs:
|
||||
run: npm pack
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v7
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: typescript-client
|
||||
path: hindsight-clients/typescript/*.tgz
|
||||
@@ -188,10 +166,10 @@ jobs:
|
||||
environment: npm
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
registry-url: 'https://registry.npmjs.org'
|
||||
@@ -226,7 +204,7 @@ jobs:
|
||||
run: npm pack
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v7
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: openclaw-integration
|
||||
path: hindsight-integrations/openclaw/*.tgz
|
||||
@@ -237,10 +215,10 @@ jobs:
|
||||
environment: npm
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
registry-url: 'https://registry.npmjs.org'
|
||||
@@ -275,7 +253,7 @@ jobs:
|
||||
run: npm pack
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v7
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ai-sdk-integration
|
||||
path: hindsight-integrations/ai-sdk/*.tgz
|
||||
@@ -286,10 +264,10 @@ jobs:
|
||||
environment: npm
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
registry-url: 'https://registry.npmjs.org'
|
||||
@@ -324,7 +302,7 @@ jobs:
|
||||
run: npm pack
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v7
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: chat-integration
|
||||
path: hindsight-integrations/chat/*.tgz
|
||||
@@ -335,10 +313,10 @@ jobs:
|
||||
environment: npm
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
registry-url: 'https://registry.npmjs.org'
|
||||
@@ -386,7 +364,7 @@ jobs:
|
||||
run: npm pack
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v7
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: control-plane
|
||||
path: hindsight-control-plane/*.tgz
|
||||
@@ -415,7 +393,7 @@ jobs:
|
||||
asset_name: hindsight-linux-arm64
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
@@ -433,7 +411,7 @@ jobs:
|
||||
chmod +x artifacts/${{ matrix.asset_name }}
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v7
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: rust-cli-${{ matrix.asset_name }}
|
||||
path: artifacts/${{ matrix.asset_name }}
|
||||
@@ -474,7 +452,7 @@ jobs:
|
||||
PRELOAD_ML_MODELS=false
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Free Disk Space
|
||||
uses: jlumbroso/free-disk-space@main
|
||||
@@ -488,13 +466,13 @@ jobs:
|
||||
swap-storage: true
|
||||
|
||||
- name: Set up QEMU
|
||||
uses: docker/setup-qemu-action@v4
|
||||
uses: docker/setup-qemu-action@v3
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v4
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Log in to GitHub Container Registry
|
||||
uses: docker/login-action@v4
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.actor }}
|
||||
@@ -506,7 +484,7 @@ jobs:
|
||||
|
||||
- name: Extract metadata for release tags
|
||||
id: meta
|
||||
uses: docker/metadata-action@v6
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ghcr.io/${{ github.repository_owner }}/${{ matrix.image_name }}
|
||||
flavor: |
|
||||
@@ -522,7 +500,7 @@ jobs:
|
||||
# # 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@v7
|
||||
# uses: docker/build-push-action@v6
|
||||
# with:
|
||||
# context: .
|
||||
# file: docker/standalone/Dockerfile
|
||||
@@ -541,7 +519,7 @@ jobs:
|
||||
|
||||
# Build multi-platform and push to release tags
|
||||
- name: Build and push release images
|
||||
uses: docker/build-push-action@v7
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: docker/standalone/Dockerfile
|
||||
@@ -559,7 +537,7 @@ jobs:
|
||||
packages: write
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install Helm
|
||||
uses: azure/setup-helm@v4
|
||||
@@ -579,7 +557,7 @@ jobs:
|
||||
run: helm push helm-packages/*.tgz oci://ghcr.io/${{ github.repository_owner }}/charts
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v7
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: helm-chart
|
||||
path: helm-packages/*.tgz
|
||||
@@ -592,68 +570,68 @@ jobs:
|
||||
contents: write
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Extract version from tag
|
||||
id: get_version
|
||||
run: echo "VERSION=${GITHUB_REF#refs/tags/v}" >> $GITHUB_OUTPUT
|
||||
|
||||
- name: Download Python packages
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: python-packages
|
||||
path: ./artifacts/python-packages
|
||||
|
||||
- name: Download TypeScript client
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: typescript-client
|
||||
path: ./artifacts/typescript-client
|
||||
|
||||
- name: Download OpenClaw Integration
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: openclaw-integration
|
||||
path: ./artifacts/openclaw-integration
|
||||
|
||||
- name: Download AI SDK Integration
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: ai-sdk-integration
|
||||
path: ./artifacts/ai-sdk-integration
|
||||
|
||||
- name: Download Chat Integration
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: chat-integration
|
||||
path: ./artifacts/chat-integration
|
||||
|
||||
- name: Download Control Plane
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: control-plane
|
||||
path: ./artifacts/control-plane
|
||||
|
||||
- name: Download Rust CLI (Linux)
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: rust-cli-hindsight-linux-amd64
|
||||
path: ./artifacts/rust-cli-linux
|
||||
|
||||
- name: Download Rust CLI (macOS Intel)
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: rust-cli-hindsight-darwin-amd64
|
||||
path: ./artifacts/rust-cli-darwin-amd64
|
||||
|
||||
- name: Download Rust CLI (macOS ARM)
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: rust-cli-hindsight-darwin-arm64
|
||||
path: ./artifacts/rust-cli-darwin-arm64
|
||||
|
||||
- name: Download Helm chart
|
||||
uses: actions/download-artifact@v8
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: helm-chart
|
||||
path: ./artifacts/helm-chart
|
||||
@@ -663,10 +641,8 @@ jobs:
|
||||
mkdir -p release-assets
|
||||
# Python packages
|
||||
cp artifacts/python-packages/hindsight-clients/python/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight-api-slim/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight-api/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight-all/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight-all-slim/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight-integrations/litellm/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight-integrations/pydantic-ai/dist/* release-assets/ || true
|
||||
cp artifacts/python-packages/hindsight-embed/dist/* release-assets/ || true
|
||||
|
||||
+156
-220
@@ -16,30 +16,30 @@ jobs:
|
||||
python-version: ['3.11', '3.12', '3.13', '3.14']
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Set up Python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Build hindsight-api
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
build-typescript-client:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
cache: 'npm'
|
||||
@@ -55,10 +55,10 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
|
||||
@@ -78,10 +78,10 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
|
||||
@@ -101,10 +101,10 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
|
||||
@@ -124,10 +124,10 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
cache: 'npm'
|
||||
@@ -176,10 +176,10 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
cache: 'npm'
|
||||
@@ -202,7 +202,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -214,7 +214,7 @@ jobs:
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
|
||||
- name: Cache cargo
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
@@ -231,41 +231,41 @@ jobs:
|
||||
run: cargo build --release
|
||||
|
||||
- name: Upload CLI artifact
|
||||
uses: actions/upload-artifact@v7
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: hindsight-cli
|
||||
path: hindsight-cli/target/release/hindsight
|
||||
retention-days: 1
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -316,7 +316,7 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install Helm
|
||||
uses: azure/setup-helm@v4
|
||||
@@ -358,7 +358,7 @@ jobs:
|
||||
PRELOAD_ML_MODELS=false
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Free Disk Space
|
||||
uses: jlumbroso/free-disk-space@main
|
||||
@@ -372,10 +372,10 @@ jobs:
|
||||
swap-storage: true
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v4
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Build ${{ matrix.name }} image (${{ matrix.variant }})
|
||||
uses: docker/build-push-action@v7
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: docker/standalone/Dockerfile
|
||||
@@ -424,7 +424,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -433,34 +433,34 @@ jobs:
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --extra test --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -472,7 +472,7 @@ jobs:
|
||||
"
|
||||
|
||||
- name: Run tests
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv run pytest tests -v
|
||||
|
||||
test-python-client:
|
||||
@@ -487,7 +487,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -496,18 +496,18 @@ jobs:
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Build Python client
|
||||
@@ -519,19 +519,19 @@ jobs:
|
||||
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -590,7 +590,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -599,28 +599,28 @@ jobs:
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install TypeScript client dependencies
|
||||
working-directory: ./hindsight-clients/typescript
|
||||
@@ -631,15 +631,15 @@ jobs:
|
||||
run: npm run build
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -690,7 +690,7 @@ jobs:
|
||||
runs-on: ubuntu-24.04-arm
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
@@ -698,7 +698,7 @@ jobs:
|
||||
targets: aarch64-unknown-linux-gnu
|
||||
|
||||
- name: Cache cargo
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
@@ -722,7 +722,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -731,13 +731,13 @@ jobs:
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
@@ -745,7 +745,7 @@ jobs:
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
|
||||
- name: Cache cargo
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
@@ -754,23 +754,23 @@ jobs:
|
||||
key: ${{ runner.os }}-cargo-client-${{ hashFiles('hindsight-clients/rust/Cargo.lock') }}
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -829,7 +829,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -838,40 +838,40 @@ jobs:
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v6
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version: '1.23'
|
||||
cache-dependency-path: hindsight-clients/go/go.sum
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -934,7 +934,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -943,43 +943,43 @@ jobs:
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '22'
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install embed dependencies
|
||||
working-directory: ./hindsight-embed
|
||||
run: uv sync --frozen --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -1041,7 +1041,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -1050,38 +1050,38 @@ jobs:
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build API
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: uv build
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install integration test dependencies
|
||||
working-directory: ./hindsight-integration-tests
|
||||
run: uv sync --frozen
|
||||
|
||||
- name: Install API dependencies
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -1132,16 +1132,16 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
@@ -1161,16 +1161,16 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
@@ -1190,16 +1190,16 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
@@ -1215,66 +1215,6 @@ jobs:
|
||||
working-directory: ./hindsight-integrations/pydantic-ai
|
||||
run: uv run pytest tests -v
|
||||
|
||||
test-pip-slim:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
HINDSIGHT_API_LLM_PROVIDER: vertexai
|
||||
HINDSIGHT_API_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY: /tmp/gcp-credentials.json
|
||||
HINDSIGHT_API_LLM_MODEL: google/gemini-2.5-flash-lite
|
||||
HINDSIGHT_API_EMBEDDINGS_PROVIDER: cohere
|
||||
HINDSIGHT_API_RERANKER_PROVIDER: cohere
|
||||
HINDSIGHT_API_COHERE_API_KEY: ${{ secrets.COHERE_API_KEY }}
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
printf '%s' '${{ secrets.GCP_VERTEXAI_CREDENTIALS }}' > /tmp/gcp-credentials.json
|
||||
PROJECT_ID=$(jq -r '.project_id' /tmp/gcp-credentials.json)
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Install hindsight-api-slim (embedded-db only, no local ML)
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --extra embedded-db --index-strategy unsafe-best-match
|
||||
|
||||
- name: Start API server
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: |
|
||||
uv run hindsight-api --port 8888 > /tmp/slim-api-server.log 2>&1 &
|
||||
for i in $(seq 1 60); do
|
||||
if curl -s http://localhost:8888/health | grep -q "healthy"; 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/slim-api-server.log
|
||||
exit 1
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
- name: Smoke test - retain and recall
|
||||
run: ./scripts/smoke-test-slim.sh http://localhost:8888
|
||||
|
||||
- name: Show API server logs
|
||||
if: always()
|
||||
run: |
|
||||
echo "=== API Server Logs ==="
|
||||
cat /tmp/slim-api-server.log 2>/dev/null || true
|
||||
|
||||
test-embed:
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
@@ -1285,7 +1225,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -1294,13 +1234,13 @@ jobs:
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
@@ -1308,12 +1248,8 @@ jobs:
|
||||
working-directory: ./hindsight-embed
|
||||
run: uv sync --frozen --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install API dependencies (with local-ml and embedded-db for smoke test)
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-embed-${{ hashFiles('hindsight-embed/pyproject.toml') }}
|
||||
@@ -1343,7 +1279,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -1352,35 +1288,35 @@ jobs:
|
||||
echo "HINDSIGHT_API_LLM_VERTEXAI_PROJECT_ID=$PROJECT_ID" >> $GITHUB_ENV
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Build hindsight-all
|
||||
working-directory: ./hindsight-all
|
||||
working-directory: ./hindsight
|
||||
run: uv build
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: ./hindsight-all
|
||||
working-directory: ./hindsight
|
||||
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-all-${{ hashFiles('hindsight-all/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-all-${{ hashFiles('hindsight/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-all-
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Run unit tests
|
||||
working-directory: ./hindsight-all
|
||||
working-directory: ./hindsight
|
||||
run: uv run pytest tests/ -v
|
||||
|
||||
test-doc-examples:
|
||||
@@ -1399,7 +1335,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Setup GCP credentials
|
||||
run: |
|
||||
@@ -1413,7 +1349,7 @@ jobs:
|
||||
|
||||
- name: Cache cargo
|
||||
if: matrix.language == 'cli'
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
@@ -1429,35 +1365,35 @@ jobs:
|
||||
cp target/release/hindsight /usr/local/bin/hindsight
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Set up Node.js
|
||||
if: matrix.language == 'node'
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: package-lock.json
|
||||
|
||||
- name: Build and install API
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv build
|
||||
uv sync --frozen --no-install-project --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install Python client dependencies
|
||||
if: matrix.language == 'python'
|
||||
working-directory: ./hindsight-clients/python
|
||||
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
|
||||
|
||||
- name: Build and install API
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: |
|
||||
uv build
|
||||
uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install TypeScript client
|
||||
if: matrix.language == 'node'
|
||||
run: |
|
||||
@@ -1465,15 +1401,15 @@ jobs:
|
||||
npm run build --workspace=hindsight-clients/typescript
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -1533,7 +1469,7 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0 # Full history needed for git clone of tags
|
||||
|
||||
@@ -1547,21 +1483,21 @@ jobs:
|
||||
run: git fetch --tags
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
prune-cache: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Cache HuggingFace models
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/huggingface
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api-slim/pyproject.toml') }}
|
||||
key: ${{ runner.os }}-huggingface-${{ hashFiles('hindsight-api/pyproject.toml') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-huggingface-
|
||||
|
||||
@@ -1570,11 +1506,11 @@ jobs:
|
||||
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
|
||||
|
||||
- name: Install current hindsight-api
|
||||
working-directory: ./hindsight-api-slim
|
||||
run: uv sync --frozen --all-extras --index-strategy unsafe-best-match
|
||||
working-directory: ./hindsight-api
|
||||
run: uv sync --frozen --index-strategy unsafe-best-match
|
||||
|
||||
- name: Pre-download models
|
||||
working-directory: ./hindsight-api-slim
|
||||
working-directory: ./hindsight-api
|
||||
run: |
|
||||
uv run python -c "
|
||||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||||
@@ -1607,20 +1543,20 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@v6
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
cache: 'npm'
|
||||
@@ -1630,7 +1566,7 @@ jobs:
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
|
||||
- name: Cache cargo
|
||||
uses: actions/cache@v5
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
@@ -1679,17 +1615,17 @@ jobs:
|
||||
UV_INDEX: pytorch=https://download.pytorch.org/whl/cpu
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0 # Fetch full git history to access base branch
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v7
|
||||
uses: astral-sh/setup-uv@v5
|
||||
with:
|
||||
enable-cache: true
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version-file: ".python-version"
|
||||
|
||||
|
||||
@@ -17,20 +17,20 @@ Hindsight is an agent memory system that provides long-term memory for AI agents
|
||||
./scripts/dev/start-api.sh
|
||||
|
||||
# Run all tests (parallelized with pytest-xdist)
|
||||
cd hindsight-api-slim && uv run pytest tests/
|
||||
cd hindsight-api && uv run pytest tests/
|
||||
|
||||
# Run specific test file
|
||||
cd hindsight-api-slim && uv run pytest tests/test_http_api_integration.py -v
|
||||
cd hindsight-api && uv run pytest tests/test_http_api_integration.py -v
|
||||
|
||||
# Run single test function
|
||||
cd hindsight-api-slim && uv run pytest tests/test_retain.py::test_retain_simple -v
|
||||
cd hindsight-api && uv run pytest tests/test_retain.py::test_retain_simple -v
|
||||
|
||||
# Lint and format
|
||||
cd hindsight-api-slim && uv run ruff check .
|
||||
cd hindsight-api-slim && uv run ruff format .
|
||||
cd hindsight-api && uv run ruff check .
|
||||
cd hindsight-api && uv run ruff format .
|
||||
|
||||
# Type checking (uses ty - extremely fast type checker from Astral)
|
||||
cd hindsight-api-slim && uv run ty check hindsight_api/
|
||||
cd hindsight-api && uv run ty check hindsight_api/
|
||||
```
|
||||
|
||||
### Control Plane (Next.js)
|
||||
@@ -72,7 +72,7 @@ cd hindsight-control-plane && npm run dev
|
||||
## Architecture
|
||||
|
||||
### Monorepo Structure
|
||||
- **hindsight-api-slim/**: Core FastAPI server with memory engine (Python, uv)
|
||||
- **hindsight-api/**: Core FastAPI server with memory engine (Python, uv)
|
||||
- **hindsight/**: Embedded Python bundle (hindsight-all package)
|
||||
- **hindsight-control-plane/**: Admin UI (Next.js, npm)
|
||||
- **hindsight-cli/**: CLI tool (Rust, cargo, uses progenitor for API client)
|
||||
@@ -81,9 +81,9 @@ cd hindsight-control-plane && npm run dev
|
||||
- **hindsight-integrations/**: Framework integrations (LiteLLM, OpenAI)
|
||||
- **hindsight-dev/**: Development tools and benchmarks
|
||||
|
||||
### Core Engine (hindsight-api-slim/hindsight_api/engine/)
|
||||
### Core Engine (hindsight-api/hindsight_api/engine/)
|
||||
- `memory_engine.py`: Main orchestrator (~170KB) for retain/recall/reflect operations
|
||||
- `llm_wrapper.py`: LLM abstraction supporting OpenAI, Anthropic, Gemini, Groq, MiniMax, Ollama, LM Studio
|
||||
- `llm_wrapper.py`: LLM abstraction supporting OpenAI, Anthropic, Gemini, Groq, Ollama, LM Studio
|
||||
- `embeddings.py`: Embedding generation (local sentence-transformers or TEI)
|
||||
- `cross_encoder.py`: Reranking (local or TEI)
|
||||
- `entity_resolver.py`: Entity extraction and normalization
|
||||
@@ -101,7 +101,7 @@ cd hindsight-control-plane && npm run dev
|
||||
- `fusion.py`: Reciprocal rank fusion for combining results
|
||||
- `reranking.py`: Cross-encoder reranking
|
||||
|
||||
### API Layer (hindsight-api-slim/hindsight_api/api/)
|
||||
### API Layer (hindsight-api/hindsight_api/api/)
|
||||
- `http.py`: FastAPI HTTP routers (~80KB) for all REST endpoints
|
||||
- `mcp.py`: Model Context Protocol server implementation
|
||||
|
||||
@@ -111,13 +111,13 @@ Main operations:
|
||||
- **Reflect**: Disposition-aware reasoning using memories and mental models.
|
||||
|
||||
### Database
|
||||
PostgreSQL with pgvector. Schema managed via Alembic migrations in `hindsight-api-slim/hindsight_api/alembic/`. Migrations run automatically on API startup.
|
||||
PostgreSQL with pgvector. Schema managed via Alembic migrations in `hindsight-api/hindsight_api/alembic/`. Migrations run automatically on API startup.
|
||||
|
||||
Key tables: `banks`, `memory_units`, `documents`, `entities`, `entity_links`
|
||||
|
||||
### Adding Database Migrations
|
||||
|
||||
1. **Create a new migration file** in `hindsight-api-slim/hindsight_api/alembic/versions/`:
|
||||
1. **Create a new migration file** in `hindsight-api/hindsight_api/alembic/versions/`:
|
||||
- File name format: `<revision_id>_<description>.py` (e.g., `f1a2b3c4d5e6_add_new_index.py`)
|
||||
- Use a unique hex revision ID (12 chars)
|
||||
- Set `down_revision` to the previous migration's revision ID
|
||||
@@ -154,7 +154,7 @@ Key tables: `banks`, `memory_units`, `documents`, `entities`, `entity_links`
|
||||
|
||||
3. **Run migrations locally**:
|
||||
```bash
|
||||
# Set database URL and run migrations for the base schema plus all tenants
|
||||
# Set database URL and run migrations
|
||||
uv run hindsight-admin run-db-migration
|
||||
|
||||
# Run on a specific tenant schema
|
||||
@@ -251,7 +251,7 @@ Fields must be categorized as either **hierarchical** (can be overridden per-ten
|
||||
|
||||
#### Adding a New Configuration Field
|
||||
|
||||
1. **config.py** (`hindsight-api-slim/hindsight_api/config.py`):
|
||||
1. **config.py** (`hindsight-api/hindsight_api/config.py`):
|
||||
- Add `ENV_*` constant for the environment variable name (e.g., `ENV_MY_SETTING = "HINDSIGHT_API_MY_SETTING"`)
|
||||
- Add `DEFAULT_*` constant for the default value
|
||||
- Add field to `HindsightConfig` dataclass with type annotation
|
||||
@@ -268,7 +268,7 @@ Fields must be categorized as either **hierarchical** (can be overridden per-ten
|
||||
# Static field - just don't add to _HIERARCHICAL_FIELDS
|
||||
```
|
||||
|
||||
2. **main.py** (`hindsight-api-slim/hindsight_api/main.py`):
|
||||
2. **main.py** (`hindsight-api/hindsight_api/main.py`):
|
||||
- Add field to the manual `HindsightConfig()` constructor call (search for "CLI override")
|
||||
|
||||
3. **Use hierarchical config in MemoryEngine**:
|
||||
@@ -308,14 +308,14 @@ cp .env.example .env
|
||||
# Edit .env with LLM API key
|
||||
|
||||
# Python deps
|
||||
uv sync --directory hindsight-api-slim/
|
||||
uv sync --directory hindsight-api/
|
||||
|
||||
# Node deps (uses npm workspaces)
|
||||
npm install
|
||||
```
|
||||
|
||||
Required env vars:
|
||||
- `HINDSIGHT_API_LLM_PROVIDER`: openai, anthropic, gemini, groq, minimax, ollama, lmstudio
|
||||
- `HINDSIGHT_API_LLM_PROVIDER`: openai, anthropic, gemini, groq, ollama, lmstudio
|
||||
- `HINDSIGHT_API_LLM_API_KEY`: Your API key
|
||||
- `HINDSIGHT_API_LLM_MODEL`: Model name (e.g., gpt-4o-mini, claude-sonnet-4-20250514)
|
||||
|
||||
|
||||
@@ -9,9 +9,8 @@
|
||||
[](https://opensource.org/licenses/MIT)
|
||||

|
||||

|
||||
<br/>
|
||||
|
||||
<a href="https://trendshift.io/repositories/15603" target="_blank"><img src="https://trendshift.io/api/badge/repositories/15603" alt="vectorize-io%2Fhindsight | Trendshift" style="width: 250px; height: 55px;" width="250" height="55"/></a>
|
||||
|
||||
</div>
|
||||
|
||||
---
|
||||
@@ -70,7 +69,7 @@ docker run --rm -it --pull always -p 8888:8888 -p 9999:9999 \
|
||||
>API: http://localhost:8888
|
||||
>UI: http://localhost:9999
|
||||
|
||||
You can modify the LLM provider by setting `HINDSIGHT_API_LLM_PROVIDER`. Valid options are `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio`, and `minimax`. The documentation provides more details on [supported models](https://hindsight.vectorize.io/developer/models).
|
||||
You can modify the LLM provider by setting `HINDSIGHT_API_LLM_PROVIDER`. Valid options are `openai`, `anthropic`, `gemini`, `groq`, `ollama`, and `lmstudio`. The documentation provides more details on [supported models](https://hindsight.vectorize.io/developer/models).
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -42,22 +42,25 @@ RUN apt-get update && apt-get install -y \
|
||||
&& pip install --no-cache-dir uv
|
||||
|
||||
# Copy dependency files and README (required by pyproject.toml)
|
||||
COPY hindsight-api-slim/pyproject.toml ./api/
|
||||
COPY hindsight-api-slim/README.md ./api/
|
||||
COPY hindsight-api/pyproject.toml ./api/
|
||||
COPY hindsight-api/README.md ./api/
|
||||
|
||||
WORKDIR /app/api
|
||||
|
||||
# Sync dependencies using appropriate extras based on INCLUDE_LOCAL_MODELS
|
||||
# local-ml: torch, sentence-transformers, transformers, einops, flashrank, mlx (optional)
|
||||
# embedded-db: pg0-embedded (always included for embedded PostgreSQL support)
|
||||
RUN if [ "$INCLUDE_LOCAL_MODELS" = "true" ]; then \
|
||||
uv sync --extra local-ml --extra embedded-db; \
|
||||
else \
|
||||
uv sync --extra embedded-db; \
|
||||
# Remove local ML model dependencies if INCLUDE_LOCAL_MODELS=false
|
||||
# This creates a smaller image when using external providers (TEI, OpenAI, Cohere)
|
||||
RUN if [ "$INCLUDE_LOCAL_MODELS" != "true" ]; then \
|
||||
echo "Removing local-models dependencies (sentence-transformers, torch, transformers)..." && \
|
||||
sed -i '/"sentence-transformers/d' pyproject.toml && \
|
||||
sed -i '/"transformers/d' pyproject.toml && \
|
||||
sed -i '/"torch/d' pyproject.toml; \
|
||||
fi
|
||||
|
||||
# Sync dependencies (will create lock file if needed)
|
||||
RUN uv sync
|
||||
|
||||
# Copy source code (alembic migrations are inside hindsight_api/)
|
||||
COPY hindsight-api-slim/hindsight_api ./hindsight_api
|
||||
COPY hindsight-api/hindsight_api ./hindsight_api
|
||||
|
||||
# Install the local package (uv sync only installed dependencies, not the package itself)
|
||||
RUN uv pip install -e .
|
||||
|
||||
@@ -77,32 +77,18 @@ PIDS=()
|
||||
# Start API if enabled
|
||||
if [ "$ENABLE_API" = "true" ]; then
|
||||
cd /app/api
|
||||
API_HEALTH_URL="${HINDSIGHT_API_HEALTH_URL:-http://localhost:8888/health}"
|
||||
API_STARTUP_WAIT_SECONDS="${HINDSIGHT_API_STARTUP_WAIT_SECONDS:-300}"
|
||||
|
||||
# Run API directly - Python's PYTHONUNBUFFERED=1 handles output buffering
|
||||
hindsight-api &
|
||||
API_PID=$!
|
||||
PIDS+=($API_PID)
|
||||
|
||||
# Wait for API to be ready
|
||||
api_ready=false
|
||||
for ((i=1; i<=API_STARTUP_WAIT_SECONDS; i++)); do
|
||||
if ! kill -0 "$API_PID" 2>/dev/null; then
|
||||
wait "$API_PID"
|
||||
exit $?
|
||||
fi
|
||||
if curl -sf "$API_HEALTH_URL" &>/dev/null; then
|
||||
api_ready=true
|
||||
for i in {1..60}; do
|
||||
if curl -sf http://localhost:8888/health &>/dev/null; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
if [ "$api_ready" != "true" ]; then
|
||||
echo "❌ API did not become healthy within ${API_STARTUP_WAIT_SECONDS}s"
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
echo "API disabled (HINDSIGHT_ENABLE_API=false)"
|
||||
fi
|
||||
@@ -111,7 +97,6 @@ fi
|
||||
if [ "$ENABLE_CP" = "true" ]; then
|
||||
echo "🎛️ Starting Control Plane..."
|
||||
cd /app/control-plane
|
||||
export HOSTNAME="${HINDSIGHT_CP_HOSTNAME:-0.0.0.0}"
|
||||
PORT="${HINDSIGHT_CP_PORT:-9999}" node server.js &
|
||||
CP_PID=$!
|
||||
PIDS+=($CP_PID)
|
||||
|
||||
@@ -49,9 +49,6 @@
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
REPO_ROOT="$(dirname "$SCRIPT_DIR")"
|
||||
|
||||
# Colors for output
|
||||
RED='\033[0;31m'
|
||||
GREEN='\033[0;32m'
|
||||
@@ -181,21 +178,6 @@ for i in $(seq 1 "$TIMEOUT"); do
|
||||
echo "=== Health Response ==="
|
||||
curl -s "http://localhost:${HEALTH_PORT}${HEALTH_PATH}" | python3 -m json.tool 2>/dev/null || curl -s "http://localhost:${HEALTH_PORT}${HEALTH_PATH}"
|
||||
echo ""
|
||||
|
||||
# Run retain/recall smoke test for API targets
|
||||
if [ "$TARGET" != "cp-only" ]; then
|
||||
echo ""
|
||||
echo "=== Retain/Recall Smoke Test ==="
|
||||
if ! "$REPO_ROOT/scripts/smoke-test-slim.sh" "http://localhost:${HEALTH_PORT}"; then
|
||||
echo ""
|
||||
echo "=== Container Logs (last 50 lines) ==="
|
||||
docker logs "$CONTAINER_NAME" 2>&1 | tail -50
|
||||
echo ""
|
||||
echo -e "${RED}Smoke test FAILED${NC}"
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=== Container Logs (last 50 lines) ==="
|
||||
docker logs "$CONTAINER_NAME" 2>&1 | tail -50
|
||||
|
||||
@@ -2,8 +2,8 @@ apiVersion: v2
|
||||
name: hindsight
|
||||
description: Hindsight helm chart
|
||||
type: application
|
||||
version: 0.4.18
|
||||
appVersion: "0.4.18"
|
||||
version: 0.4.16
|
||||
appVersion: "0.4.16"
|
||||
keywords:
|
||||
- ai
|
||||
- memory
|
||||
|
||||
@@ -1,33 +0,0 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=61"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "hindsight-all-slim"
|
||||
version = "0.4.18"
|
||||
description = "Hindsight: Agent Memory That Works Like Human Memory - Slim All-in-One Bundle"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
"hindsight-api-slim>=0.4.17",
|
||||
"hindsight-client>=0.0.7",
|
||||
"hindsight-embed>=0.1.0",
|
||||
]
|
||||
|
||||
[tool.uv.sources]
|
||||
hindsight-api-slim = { workspace = true }
|
||||
hindsight-client = { workspace = true }
|
||||
hindsight-embed = { workspace = true }
|
||||
|
||||
[project.optional-dependencies]
|
||||
test = [
|
||||
"pytest>=7.0.0",
|
||||
"pytest-asyncio>=0.21.0",
|
||||
]
|
||||
|
||||
[tool.setuptools]
|
||||
packages = []
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
asyncio_mode = "auto"
|
||||
asyncio_default_fixture_loop_scope = "function"
|
||||
@@ -1,48 +0,0 @@
|
||||
# hindsight-all
|
||||
|
||||
All-in-one package for Hindsight - Agent Memory That Works Like Human Memory
|
||||
|
||||
## Quick Start
|
||||
|
||||
```python
|
||||
from hindsight import start_server, HindsightClient
|
||||
|
||||
# Start server with embedded PostgreSQL
|
||||
server = start_server(
|
||||
llm_provider="groq",
|
||||
llm_api_key="your-api-key",
|
||||
llm_model="openai/gpt-oss-120b"
|
||||
)
|
||||
|
||||
# Create client
|
||||
client = HindsightClient(base_url=server.url)
|
||||
|
||||
# Store memories
|
||||
client.put(agent_id="assistant", content="User prefers Python for data analysis")
|
||||
|
||||
# Search memories
|
||||
results = client.search(agent_id="assistant", query="programming preferences")
|
||||
|
||||
# Generate contextual response
|
||||
response = client.think(agent_id="assistant", query="What languages should I recommend?")
|
||||
|
||||
# Stop server when done
|
||||
server.stop()
|
||||
```
|
||||
|
||||
## Using Context Manager
|
||||
|
||||
```python
|
||||
from hindsight import HindsightServer, HindsightClient
|
||||
|
||||
with HindsightServer(llm_provider="groq", llm_api_key="...") as server:
|
||||
client = HindsightClient(base_url=server.url)
|
||||
# ... use client ...
|
||||
# Server automatically stops
|
||||
```
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
pip install hindsight-all
|
||||
```
|
||||
@@ -1,423 +0,0 @@
|
||||
"""
|
||||
Wrapper for Hindsight client that adds API namespaces.
|
||||
|
||||
Provides organized access to different parts of the Hindsight API through
|
||||
namespaces like .banks, .mental_models, etc.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from hindsight_client import Hindsight
|
||||
|
||||
|
||||
class BanksAPI:
|
||||
"""Namespace for bank-related operations.
|
||||
|
||||
Provides methods to create, delete, and manage memory banks.
|
||||
"""
|
||||
|
||||
def __init__(self, client: Hindsight):
|
||||
self._client = client
|
||||
|
||||
def create(
|
||||
self,
|
||||
bank_id: str,
|
||||
name: str | None = None,
|
||||
mission: str | None = None,
|
||||
disposition: dict[str, Any] | None = None,
|
||||
) -> Any:
|
||||
"""Create a new bank.
|
||||
|
||||
Args:
|
||||
bank_id: Unique identifier for the bank.
|
||||
name: Optional display name for the bank.
|
||||
mission: Optional mission statement for the bank.
|
||||
disposition: Optional disposition configuration dict.
|
||||
|
||||
Returns:
|
||||
Bank creation response from the API.
|
||||
"""
|
||||
return self._client.create_bank(
|
||||
bank_id=bank_id,
|
||||
name=name,
|
||||
mission=mission,
|
||||
disposition=disposition,
|
||||
)
|
||||
|
||||
def delete(self, bank_id: str) -> Any:
|
||||
"""Delete a bank.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank to delete.
|
||||
|
||||
Returns:
|
||||
Deletion response from the API.
|
||||
"""
|
||||
return self._client.delete_bank(bank_id=bank_id)
|
||||
|
||||
def set_mission(self, bank_id: str, mission: str) -> Any:
|
||||
"""Set or update the mission for a bank.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
mission: The mission statement to set.
|
||||
|
||||
Returns:
|
||||
API response confirming the update.
|
||||
"""
|
||||
return self._client.set_mission(bank_id=bank_id, mission=mission)
|
||||
|
||||
def set_disposition(self, bank_id: str, disposition: dict[str, Any]) -> Any:
|
||||
"""Set or update the disposition for a bank.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
disposition: The disposition configuration dict.
|
||||
|
||||
Returns:
|
||||
API response confirming the update.
|
||||
"""
|
||||
return self._client.set_disposition(bank_id=bank_id, disposition=disposition)
|
||||
|
||||
def list(self) -> Any:
|
||||
"""List all banks.
|
||||
|
||||
Returns:
|
||||
List of banks from the API.
|
||||
"""
|
||||
from hindsight_client.hindsight_client import _run_async
|
||||
|
||||
return _run_async(self._client._banks_api.list_banks())
|
||||
|
||||
|
||||
class MentalModelsAPI:
|
||||
"""Namespace for mental model operations.
|
||||
|
||||
Mental models are reusable knowledge structures that guide agent behavior.
|
||||
"""
|
||||
|
||||
def __init__(self, client: Hindsight):
|
||||
self._client = client
|
||||
|
||||
def create(
|
||||
self,
|
||||
bank_id: str,
|
||||
name: str,
|
||||
content: str,
|
||||
tags: list[str] | None = None,
|
||||
) -> Any:
|
||||
"""Create a new mental model.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank to add the model to.
|
||||
name: Name for the mental model.
|
||||
content: The content/instructions for the mental model.
|
||||
tags: Optional list of tags for categorization.
|
||||
|
||||
Returns:
|
||||
Creation response from the API.
|
||||
"""
|
||||
return self._client.create_mental_model(
|
||||
bank_id=bank_id,
|
||||
name=name,
|
||||
content=content,
|
||||
tags=tags,
|
||||
)
|
||||
|
||||
def list(self, bank_id: str, tags: list[str] | None = None) -> Any:
|
||||
"""List all mental models for a bank.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
tags: Optional filter by tags.
|
||||
|
||||
Returns:
|
||||
List of mental models.
|
||||
"""
|
||||
return self._client.list_mental_models(bank_id=bank_id, tags=tags)
|
||||
|
||||
def get(self, bank_id: str, mental_model_id: str) -> Any:
|
||||
"""Get a specific mental model.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
mental_model_id: The ID of the mental model.
|
||||
|
||||
Returns:
|
||||
The mental model details.
|
||||
"""
|
||||
return self._client.get_mental_model(bank_id=bank_id, mental_model_id=mental_model_id)
|
||||
|
||||
def refresh(self, bank_id: str, mental_model_id: str) -> Any:
|
||||
"""Refresh a mental model.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
mental_model_id: The ID of the mental model to refresh.
|
||||
|
||||
Returns:
|
||||
Refresh response from the API.
|
||||
"""
|
||||
return self._client.refresh_mental_model(bank_id=bank_id, mental_model_id=mental_model_id)
|
||||
|
||||
def update(
|
||||
self,
|
||||
bank_id: str,
|
||||
mental_model_id: str,
|
||||
name: str | None = None,
|
||||
content: str | None = None,
|
||||
tags: list[str] | None = None,
|
||||
) -> Any:
|
||||
"""Update a mental model.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
mental_model_id: The ID of the mental model to update.
|
||||
name: Optional new name.
|
||||
content: Optional new content.
|
||||
tags: Optional new tags list.
|
||||
|
||||
Returns:
|
||||
Update response from the API.
|
||||
"""
|
||||
return self._client.update_mental_model(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
name=name,
|
||||
content=content,
|
||||
tags=tags,
|
||||
)
|
||||
|
||||
def delete(self, bank_id: str, mental_model_id: str) -> Any:
|
||||
"""Delete a mental model.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
mental_model_id: The ID of the mental model to delete.
|
||||
|
||||
Returns:
|
||||
Deletion response from the API.
|
||||
"""
|
||||
return self._client.delete_mental_model(bank_id=bank_id, mental_model_id=mental_model_id)
|
||||
|
||||
|
||||
class DirectivesAPI:
|
||||
"""Namespace for directive operations.
|
||||
|
||||
Directives are explicit instructions that guide agent behavior.
|
||||
"""
|
||||
|
||||
def __init__(self, client: Hindsight):
|
||||
self._client = client
|
||||
|
||||
def create(
|
||||
self,
|
||||
bank_id: str,
|
||||
name: str,
|
||||
content: str,
|
||||
tags: list[str] | None = None,
|
||||
) -> Any:
|
||||
"""Create a new directive.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank to add the directive to.
|
||||
name: Name for the directive.
|
||||
content: The directive content/instructions.
|
||||
tags: Optional list of tags for categorization.
|
||||
|
||||
Returns:
|
||||
Creation response from the API.
|
||||
"""
|
||||
return self._client.create_directive(
|
||||
bank_id=bank_id,
|
||||
name=name,
|
||||
content=content,
|
||||
tags=tags,
|
||||
)
|
||||
|
||||
def list(self, bank_id: str, tags: list[str] | None = None) -> Any:
|
||||
"""List all directives for a bank.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
tags: Optional filter by tags.
|
||||
|
||||
Returns:
|
||||
List of directives.
|
||||
"""
|
||||
return self._client.list_directives(bank_id=bank_id, tags=tags)
|
||||
|
||||
def get(self, bank_id: str, directive_id: str) -> Any:
|
||||
"""Get a specific directive.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
directive_id: The ID of the directive.
|
||||
|
||||
Returns:
|
||||
The directive details.
|
||||
"""
|
||||
return self._client.get_directive(bank_id=bank_id, directive_id=directive_id)
|
||||
|
||||
def update(
|
||||
self,
|
||||
bank_id: str,
|
||||
directive_id: str,
|
||||
name: str | None = None,
|
||||
content: str | None = None,
|
||||
tags: list[str] | None = None,
|
||||
) -> Any:
|
||||
"""Update a directive.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
directive_id: The ID of the directive to update.
|
||||
name: Optional new name.
|
||||
content: Optional new content.
|
||||
tags: Optional new tags list.
|
||||
|
||||
Returns:
|
||||
Update response from the API.
|
||||
"""
|
||||
return self._client.update_directive(
|
||||
bank_id=bank_id,
|
||||
directive_id=directive_id,
|
||||
name=name,
|
||||
content=content,
|
||||
tags=tags,
|
||||
)
|
||||
|
||||
def delete(self, bank_id: str, directive_id: str) -> Any:
|
||||
"""Delete a directive.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank.
|
||||
directive_id: The ID of the directive to delete.
|
||||
|
||||
Returns:
|
||||
Deletion response from the API.
|
||||
"""
|
||||
return self._client.delete_directive(bank_id=bank_id, directive_id=directive_id)
|
||||
|
||||
|
||||
class MemoriesAPI:
|
||||
"""Namespace for memory operations.
|
||||
|
||||
Provides methods to query and retrieve stored memories.
|
||||
"""
|
||||
|
||||
def __init__(self, client: Hindsight):
|
||||
self._client = client
|
||||
|
||||
def list(
|
||||
self,
|
||||
bank_id: str,
|
||||
type: str | None = None,
|
||||
search_query: str | None = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
) -> Any:
|
||||
"""List memories in a bank.
|
||||
|
||||
Args:
|
||||
bank_id: The ID of the bank to query.
|
||||
type: Optional filter by memory type.
|
||||
search_query: Optional search query for filtering.
|
||||
limit: Maximum number of results to return (default: 100).
|
||||
offset: Number of results to skip for pagination (default: 0).
|
||||
|
||||
Returns:
|
||||
List of memories matching the criteria.
|
||||
"""
|
||||
return self._client.list_memories(
|
||||
bank_id=bank_id,
|
||||
type=type,
|
||||
search_query=search_query,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
|
||||
class HindsightClient(Hindsight):
|
||||
"""
|
||||
Enhanced Hindsight client with organized API namespaces.
|
||||
|
||||
This wrapper extends the auto-generated Hindsight client with organized
|
||||
access to different parts of the API through namespaces.
|
||||
|
||||
Example:
|
||||
```python
|
||||
from hindsight import HindsightClient
|
||||
|
||||
client = HindsightClient(base_url="http://localhost:8888")
|
||||
|
||||
# Core operations (inherited from Hindsight)
|
||||
client.retain(bank_id="test", content="Hello")
|
||||
results = client.recall(bank_id="test", query="Hello")
|
||||
|
||||
# Organized API access through namespaces
|
||||
client.banks.create(bank_id="test", name="Test Bank")
|
||||
models = client.mental_models.list(bank_id="test")
|
||||
directives = client.directives.list(bank_id="test")
|
||||
memories = client.memories.list(bank_id="test")
|
||||
```
|
||||
|
||||
Attributes:
|
||||
banks: Namespace for bank management operations.
|
||||
mental_models: Namespace for mental model operations.
|
||||
directives: Namespace for directive operations.
|
||||
memories: Namespace for memory listing operations.
|
||||
"""
|
||||
|
||||
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
self._banks_namespace: BanksAPI | None = None
|
||||
self._mental_models_namespace: MentalModelsAPI | None = None
|
||||
self._directives_namespace: DirectivesAPI | None = None
|
||||
self._memories_namespace: MemoriesAPI | None = None
|
||||
|
||||
@property
|
||||
def banks(self) -> BanksAPI:
|
||||
"""Access bank management operations.
|
||||
|
||||
Returns:
|
||||
BanksAPI instance for bank operations.
|
||||
"""
|
||||
if self._banks_namespace is None:
|
||||
self._banks_namespace = BanksAPI(self)
|
||||
return self._banks_namespace
|
||||
|
||||
@property
|
||||
def mental_models(self) -> MentalModelsAPI:
|
||||
"""Access mental model operations.
|
||||
|
||||
Returns:
|
||||
MentalModelsAPI instance for mental model operations.
|
||||
"""
|
||||
if self._mental_models_namespace is None:
|
||||
self._mental_models_namespace = MentalModelsAPI(self)
|
||||
return self._mental_models_namespace
|
||||
|
||||
@property
|
||||
def directives(self) -> DirectivesAPI:
|
||||
"""Access directive operations.
|
||||
|
||||
Returns:
|
||||
DirectivesAPI instance for directive operations.
|
||||
"""
|
||||
if self._directives_namespace is None:
|
||||
self._directives_namespace = DirectivesAPI(self)
|
||||
return self._directives_namespace
|
||||
|
||||
@property
|
||||
def memories(self) -> MemoriesAPI:
|
||||
"""Access memory listing operations.
|
||||
|
||||
Returns:
|
||||
MemoriesAPI instance for memory operations.
|
||||
"""
|
||||
if self._memories_namespace is None:
|
||||
self._memories_namespace = MemoriesAPI(self)
|
||||
return self._memories_namespace
|
||||
@@ -1,137 +0,0 @@
|
||||
# Hindsight API
|
||||
|
||||
**Memory System for AI Agents** — Temporal + Semantic + Entity Memory Architecture using PostgreSQL with pgvector.
|
||||
|
||||
Hindsight gives AI agents persistent memory that works like human memory: it stores facts, tracks entities and relationships, handles temporal reasoning ("what happened last spring?"), and forms opinions based on configurable disposition traits.
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
pip install hindsight-api
|
||||
```
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Run the Server
|
||||
|
||||
```bash
|
||||
# Set your LLM provider
|
||||
export HINDSIGHT_API_LLM_PROVIDER=openai
|
||||
export HINDSIGHT_API_LLM_API_KEY=sk-xxxxxxxxxxxx
|
||||
|
||||
# Start the server (uses embedded PostgreSQL by default)
|
||||
hindsight-api
|
||||
```
|
||||
|
||||
The server starts at http://localhost:8888 with:
|
||||
- REST API for memory operations
|
||||
- MCP server at `/mcp` for tool-use integration
|
||||
|
||||
### Use the Python API
|
||||
|
||||
```python
|
||||
from hindsight_api import MemoryEngine
|
||||
|
||||
# Create and initialize the memory engine
|
||||
memory = MemoryEngine()
|
||||
await memory.initialize()
|
||||
|
||||
# Create a memory bank for your agent
|
||||
bank = await memory.create_memory_bank(
|
||||
name="my-assistant",
|
||||
background="A helpful coding assistant"
|
||||
)
|
||||
|
||||
# Store a memory
|
||||
await memory.retain(
|
||||
memory_bank_id=bank.id,
|
||||
content="The user prefers Python for data science projects"
|
||||
)
|
||||
|
||||
# Recall memories
|
||||
results = await memory.recall(
|
||||
memory_bank_id=bank.id,
|
||||
query="What programming language does the user prefer?"
|
||||
)
|
||||
|
||||
# Reflect with reasoning
|
||||
response = await memory.reflect(
|
||||
memory_bank_id=bank.id,
|
||||
query="Should I recommend Python or R for this ML project?"
|
||||
)
|
||||
```
|
||||
|
||||
## CLI Options
|
||||
|
||||
```bash
|
||||
hindsight-api --help
|
||||
|
||||
# Common options
|
||||
hindsight-api --port 9000 # Custom port (default: 8888)
|
||||
hindsight-api --host 127.0.0.1 # Bind to localhost only
|
||||
hindsight-api --workers 4 # Multiple worker processes
|
||||
hindsight-api --log-level debug # Verbose logging
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
Configure via environment variables:
|
||||
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `HINDSIGHT_API_DATABASE_URL` | PostgreSQL connection string | `pg0` (embedded) |
|
||||
| `HINDSIGHT_API_LLM_PROVIDER` | `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio` | `openai` |
|
||||
| `HINDSIGHT_API_LLM_API_KEY` | API key for LLM provider | - |
|
||||
| `HINDSIGHT_API_LLM_MODEL` | Model name | `gpt-4o-mini` |
|
||||
| `HINDSIGHT_API_HOST` | Server bind address | `0.0.0.0` |
|
||||
| `HINDSIGHT_API_PORT` | Server port | `8888` |
|
||||
|
||||
### Example with External PostgreSQL
|
||||
|
||||
```bash
|
||||
export HINDSIGHT_API_DATABASE_URL=postgresql://user:pass@localhost:5432/hindsight
|
||||
export HINDSIGHT_API_LLM_PROVIDER=groq
|
||||
export HINDSIGHT_API_LLM_API_KEY=gsk_xxxxxxxxxxxx
|
||||
|
||||
hindsight-api
|
||||
```
|
||||
|
||||
## Docker
|
||||
|
||||
```bash
|
||||
docker run --rm -it -p 8888:8888 \
|
||||
-e HINDSIGHT_API_LLM_API_KEY=$OPENAI_API_KEY \
|
||||
-v $HOME/.hindsight-docker:/home/hindsight/.pg0 \
|
||||
ghcr.io/vectorize-io/hindsight:latest
|
||||
```
|
||||
|
||||
## MCP Server
|
||||
|
||||
For local MCP integration without running the full API server:
|
||||
|
||||
```bash
|
||||
hindsight-local-mcp
|
||||
```
|
||||
|
||||
This runs a stdio-based MCP server that can be used directly with MCP-compatible clients.
|
||||
|
||||
## Key Features
|
||||
|
||||
- **Multi-Strategy Retrieval (TEMPR)** — Semantic, keyword, graph, and temporal search combined with RRF fusion
|
||||
- **Entity Graph** — Automatic entity extraction and relationship tracking
|
||||
- **Temporal Reasoning** — Native support for time-based queries
|
||||
- **Disposition Traits** — Configurable skepticism, literalism, and empathy influence opinion formation
|
||||
- **Three Memory Types** — World facts, bank actions, and formed opinions with confidence scores
|
||||
|
||||
## Documentation
|
||||
|
||||
Full documentation: [https://hindsight.vectorize.io](https://hindsight.vectorize.io)
|
||||
|
||||
- [Installation Guide](https://hindsight.vectorize.io/developer/installation)
|
||||
- [Configuration Reference](https://hindsight.vectorize.io/developer/configuration)
|
||||
- [API Reference](https://hindsight.vectorize.io/api-reference)
|
||||
- [Python SDK](https://hindsight.vectorize.io/sdks/python)
|
||||
|
||||
## License
|
||||
|
||||
Apache 2.0
|
||||
-30
@@ -1,30 +0,0 @@
|
||||
"""Add history column to mental_models
|
||||
|
||||
Revision ID: c3d4e5f6g7h8
|
||||
Revises: a2b3c4d5e6f7, a2b3c4d5e6f8
|
||||
Create Date: 2026-03-06
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "c3d4e5f6g7h8"
|
||||
down_revision: str | Sequence[str] | None = ("a2b3c4d5e6f7", "a2b3c4d5e6f8")
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
op.execute(f"ALTER TABLE {schema}mental_models ADD COLUMN IF NOT EXISTS history JSONB DEFAULT '[]'::jsonb")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
op.execute(f"ALTER TABLE {schema}mental_models DROP COLUMN IF EXISTS history")
|
||||
-53
@@ -1,53 +0,0 @@
|
||||
"""Recreate idx_memory_units_source_memory_ids GIN index with fastupdate=off
|
||||
|
||||
GIN indexes use a "fastupdate" pending list by default: small writes are
|
||||
buffered there and flushed to the main GIN tree in bulk. Flushing requires
|
||||
AccessExclusiveLock on the index. Under high insert concurrency (e.g. 8
|
||||
parallel pytest-xdist workers all calling retain_async) two transactions can
|
||||
each trigger a flush simultaneously and deadlock.
|
||||
|
||||
Disabling fastupdate makes every insert write directly to the GIN tree
|
||||
(slightly slower per insert, but no pending-list lock cycles).
|
||||
|
||||
Revision ID: d4e5f6g7h8i9
|
||||
Revises: d5e6f7a8b9c0
|
||||
Create Date: 2026-03-11
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "d4e5f6g7h8i9"
|
||||
down_revision: str | Sequence[str] | None = "d5e6f7a8b9c0"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
# DROP + CREATE CONCURRENTLY must run outside a transaction block.
|
||||
op.execute("COMMIT")
|
||||
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}idx_memory_units_source_memory_ids")
|
||||
op.execute(
|
||||
f"CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_memory_units_source_memory_ids "
|
||||
f"ON {schema}memory_units USING GIN (source_memory_ids) "
|
||||
f"WITH (fastupdate=off) "
|
||||
f"WHERE source_memory_ids IS NOT NULL"
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
op.execute("COMMIT")
|
||||
op.execute(f"DROP INDEX CONCURRENTLY IF EXISTS {schema}idx_memory_units_source_memory_ids")
|
||||
op.execute(
|
||||
f"CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_memory_units_source_memory_ids "
|
||||
f"ON {schema}memory_units USING GIN (source_memory_ids) "
|
||||
f"WHERE source_memory_ids IS NOT NULL"
|
||||
)
|
||||
-131
@@ -1,131 +0,0 @@
|
||||
"""Add internal_id to banks and per-(bank, fact_type) partial HNSW indexes
|
||||
|
||||
Revision ID: d5e6f7a8b9c0
|
||||
Revises: a3b4c5d6e7f8
|
||||
Create Date: 2026-03-11
|
||||
|
||||
This migration:
|
||||
1. Adds internal_id UUID column to banks (stable identifier for index naming)
|
||||
2. Drops the global HNSW index (competes with per-bank partial indexes)
|
||||
3. Creates per-(bank_id, fact_type) partial HNSW indexes for all existing banks
|
||||
(new banks get indexes created at bank-creation time via bank_utils.create_bank_hnsw_indexes)
|
||||
|
||||
Why per-(bank, fact_type) indexes:
|
||||
- fact_type-only partial indexes are never chosen by the planner when bank_id is in the WHERE
|
||||
clause, because the idx_memory_units_bank_id B-tree index always wins at planning time.
|
||||
- Per-(bank, fact_type) partial indexes have both predicates matching → planner selects them.
|
||||
- The global HNSW index competes for larger partitions (world, observation) and must be dropped.
|
||||
|
||||
For large deployments, create indexes CONCURRENTLY before running this migration:
|
||||
SELECT internal_id, bank_id FROM banks;
|
||||
-- for each bank and each fact_type in (world, experience, observation):
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_mu_emb_{ft}_{uid16}
|
||||
ON memory_units USING hnsw (embedding vector_cosine_ops)
|
||||
WHERE fact_type = '{ft}' AND bank_id = '{bank_id}';
|
||||
DROP INDEX CONCURRENTLY IF EXISTS idx_memory_units_embedding;
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
from sqlalchemy import text
|
||||
|
||||
revision: str = "d5e6f7a8b9c0"
|
||||
down_revision: str | Sequence[str] | None = "c3d4e5f6g7h8"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
_HNSW_FACT_TYPES: dict[str, str] = {
|
||||
"world": "worl",
|
||||
"experience": "expr",
|
||||
"observation": "obsv",
|
||||
}
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# 1. Add internal_id column to banks
|
||||
op.execute(
|
||||
f"ALTER TABLE {schema}banks ADD COLUMN IF NOT EXISTS internal_id UUID DEFAULT gen_random_uuid() NOT NULL"
|
||||
)
|
||||
op.execute(f"ALTER TABLE {schema}banks ADD CONSTRAINT banks_internal_id_unique UNIQUE (internal_id)")
|
||||
|
||||
# 2. Drop any fact_type-only partial HNSW indexes that may exist from prior migrations
|
||||
# (bank_id B-tree always wins over them when bank_id is in the WHERE clause)
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_mu_emb_world")
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_mu_emb_observation")
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_mu_emb_experience")
|
||||
|
||||
# 4. Drop global HNSW index (competes with per-bank partial indexes)
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_embedding")
|
||||
|
||||
# 5. Create per-(bank, fact_type) partial HNSW indexes for all existing banks
|
||||
bind = op.get_bind()
|
||||
schema_name = context.config.get_main_option("target_schema")
|
||||
table_ref = f'"{schema_name}".memory_units' if schema_name else "memory_units"
|
||||
banks_ref = f'"{schema_name}".banks' if schema_name else "banks"
|
||||
|
||||
rows = bind.execute(text(f"SELECT bank_id, internal_id FROM {banks_ref}")).fetchall() # noqa: S608
|
||||
for row in rows:
|
||||
bank_id = row[0]
|
||||
internal_id = str(row[1]).replace("-", "")[:16]
|
||||
escaped_bank_id = bank_id.replace("'", "''")
|
||||
for ft, ft_short in _HNSW_FACT_TYPES.items():
|
||||
idx_name = f"idx_mu_emb_{ft_short}_{internal_id}"
|
||||
# Index name is schema-unqualified (indexes live in the schema of their table)
|
||||
bind.execute(
|
||||
text(
|
||||
f"CREATE INDEX IF NOT EXISTS {idx_name} "
|
||||
f"ON {table_ref} USING hnsw (embedding vector_cosine_ops) "
|
||||
f"WHERE fact_type = '{ft}' AND bank_id = '{escaped_bank_id}'"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Drop per-bank HNSW indexes (iterate existing banks)
|
||||
bind = op.get_bind()
|
||||
schema_name = context.config.get_main_option("target_schema")
|
||||
banks_ref = f'"{schema_name}".banks' if schema_name else "banks"
|
||||
|
||||
rows = bind.execute(text(f"SELECT internal_id FROM {banks_ref}")).fetchall() # noqa: S608
|
||||
for row in rows:
|
||||
internal_id = str(row[0]).replace("-", "")[:16]
|
||||
for ft_short in _HNSW_FACT_TYPES.values():
|
||||
idx_name = f"idx_mu_emb_{ft_short}_{internal_id}"
|
||||
bind.execute(text(f"DROP INDEX IF EXISTS {schema}{idx_name}"))
|
||||
|
||||
# Restore the global HNSW index
|
||||
table_ref = f'"{schema_name}".memory_units' if schema_name else "memory_units"
|
||||
op.execute(
|
||||
f"CREATE INDEX IF NOT EXISTS idx_memory_units_embedding ON {table_ref} USING hnsw (embedding vector_cosine_ops)"
|
||||
)
|
||||
|
||||
# Restore old fact_type-only partial indexes
|
||||
op.execute(
|
||||
f"CREATE INDEX IF NOT EXISTS idx_mu_emb_world "
|
||||
f"ON {table_ref} USING hnsw (embedding vector_cosine_ops) "
|
||||
f"WHERE fact_type = 'world'"
|
||||
)
|
||||
op.execute(
|
||||
f"CREATE INDEX IF NOT EXISTS idx_mu_emb_observation "
|
||||
f"ON {table_ref} USING hnsw (embedding vector_cosine_ops) "
|
||||
f"WHERE fact_type = 'observation'"
|
||||
)
|
||||
op.execute(
|
||||
f"CREATE INDEX IF NOT EXISTS idx_mu_emb_experience "
|
||||
f"ON {table_ref} USING hnsw (embedding vector_cosine_ops) "
|
||||
f"WHERE fact_type = 'experience'"
|
||||
)
|
||||
|
||||
# Drop internal_id column
|
||||
op.execute(f"ALTER TABLE {schema}banks DROP CONSTRAINT IF EXISTS banks_internal_id_unique")
|
||||
op.execute(f"ALTER TABLE {schema}banks DROP COLUMN IF EXISTS internal_id")
|
||||
-73
@@ -1,73 +0,0 @@
|
||||
"""Add CASCADE DELETE FK from async_operations and webhooks to banks.
|
||||
|
||||
When a bank is deleted, all its async_operations and webhooks rows are
|
||||
automatically deleted by the database. This ensures that any in-flight
|
||||
worker tasks detect the deletion via _check_op_alive() and abort early.
|
||||
|
||||
Revision ID: e5f6g7h8i9j0
|
||||
Revises: d4e5f6g7h8i9
|
||||
Create Date: 2026-03-11
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "e5f6g7h8i9j0"
|
||||
down_revision: str | Sequence[str] | None = "d4e5f6g7h8i9"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Remove orphaned async_operations rows whose bank no longer exists
|
||||
# (can happen because there was no FK before this migration).
|
||||
op.execute(
|
||||
f"""
|
||||
DELETE FROM {schema}async_operations
|
||||
WHERE bank_id IS NOT NULL
|
||||
AND bank_id NOT IN (SELECT bank_id FROM {schema}banks)
|
||||
"""
|
||||
)
|
||||
|
||||
# Remove orphaned webhooks rows whose bank no longer exists.
|
||||
op.execute(
|
||||
f"""
|
||||
DELETE FROM {schema}webhooks
|
||||
WHERE bank_id IS NOT NULL
|
||||
AND bank_id NOT IN (SELECT bank_id FROM {schema}banks)
|
||||
"""
|
||||
)
|
||||
|
||||
# Add FK with ON DELETE CASCADE so that deleting a bank automatically
|
||||
# cleans up all its pending/processing operations and webhook configs.
|
||||
op.execute(
|
||||
f"""
|
||||
ALTER TABLE {schema}async_operations
|
||||
ADD CONSTRAINT fk_async_operations_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id)
|
||||
ON DELETE CASCADE
|
||||
"""
|
||||
)
|
||||
|
||||
op.execute(
|
||||
f"""
|
||||
ALTER TABLE {schema}webhooks
|
||||
ADD CONSTRAINT fk_webhooks_bank_id
|
||||
FOREIGN KEY (bank_id) REFERENCES {schema}banks(bank_id)
|
||||
ON DELETE CASCADE
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
op.execute(f"ALTER TABLE {schema}async_operations DROP CONSTRAINT IF EXISTS fk_async_operations_bank_id")
|
||||
op.execute(f"ALTER TABLE {schema}webhooks DROP CONSTRAINT IF EXISTS fk_webhooks_bank_id")
|
||||
-38
@@ -1,38 +0,0 @@
|
||||
"""chunk_fk_cascade_delete
|
||||
|
||||
Revision ID: f6g7h8i9j0k1
|
||||
Revises: e5f6g7h8i9j0
|
||||
Create Date: 2026-03-16 00:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "f6g7h8i9j0k1"
|
||||
down_revision: str | Sequence[str] | None = "e5f6g7h8i9j0"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Change memory_units.chunk_id FK from SET NULL to CASCADE.
|
||||
|
||||
When a document is deleted the CASCADE reaches chunks first; with SET NULL
|
||||
the memory_units rows survived with chunk_id = NULL, leaving ghost records.
|
||||
Switching to CASCADE ensures they are removed together with their chunk.
|
||||
"""
|
||||
op.drop_constraint("memory_units_chunk_fkey", "memory_units", type_="foreignkey")
|
||||
op.create_foreign_key(
|
||||
"memory_units_chunk_fkey", "memory_units", "chunks", ["chunk_id"], ["chunk_id"], ondelete="CASCADE"
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Revert to SET NULL behaviour."""
|
||||
op.drop_constraint("memory_units_chunk_fkey", "memory_units", type_="foreignkey")
|
||||
op.create_foreign_key(
|
||||
"memory_units_chunk_fkey", "memory_units", "chunks", ["chunk_id"], ["chunk_id"], ondelete="SET NULL"
|
||||
)
|
||||
-71
@@ -1,71 +0,0 @@
|
||||
"""backsweep_orphan_memory_units
|
||||
|
||||
Two-pass cleanup of memory_units rows that were never removed by earlier bugs:
|
||||
|
||||
Pass 1 — any fact_type, bank gone:
|
||||
memory_units whose bank_id no longer exists in banks. These accumulate when
|
||||
a bank is deleted without a proper cascade (no FK from memory_units to banks
|
||||
exists in the schema).
|
||||
|
||||
Pass 2 — observations only, all sources gone:
|
||||
observation rows whose bank still exists but every source_memory_id points
|
||||
to a deleted memory unit. These were left behind before PR #580 fixed the
|
||||
chunk FK cascade and before delete_document() called
|
||||
_delete_stale_observations_for_memories.
|
||||
|
||||
Revision ID: g7h8i9j0k1l2
|
||||
Revises: f6g7h8i9j0k1
|
||||
Create Date: 2026-03-16
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision: str = "g7h8i9j0k1l2"
|
||||
down_revision: str | Sequence[str] | None = "f6g7h8i9j0k1"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
schema = _get_schema_prefix()
|
||||
mu = f"{schema}memory_units"
|
||||
banks = f"{schema}banks"
|
||||
|
||||
# Pass 1: delete all memory_units (any fact_type) whose bank no longer exists.
|
||||
# There is no FK from memory_units to banks, so these never cascade away.
|
||||
op.execute(
|
||||
f"""
|
||||
DELETE FROM {mu}
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM {banks} b WHERE b.bank_id = {mu}.bank_id
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
# Pass 2: delete orphaned observations whose bank still exists but every
|
||||
# source_memory_id refers to a now-deleted memory unit (or the array is
|
||||
# empty). Observations with at least one surviving source are left alone.
|
||||
op.execute(
|
||||
f"""
|
||||
DELETE FROM {mu} orphan
|
||||
WHERE orphan.fact_type = 'observation'
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM {mu} src
|
||||
WHERE src.id = ANY(orphan.source_memory_ids)
|
||||
AND src.bank_id = orphan.bank_id
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Deleted rows cannot be restored.
|
||||
pass
|
||||
@@ -1,144 +0,0 @@
|
||||
"""
|
||||
MLX implementation of jina-reranker-v3 for Apple Silicon.
|
||||
|
||||
This file is adapted from the official model repository:
|
||||
https://huggingface.co/jinaai/jina-reranker-v3-mlx/blob/main/rerank.py
|
||||
|
||||
License: CC BY-NC 4.0 (contact Jina AI for commercial usage)
|
||||
|
||||
Changes from upstream:
|
||||
- Removed the __main__ example block
|
||||
- Type annotations added to public methods
|
||||
- top_n parameter added to rerank() (upstream only exposed it implicitly)
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
class _MLPProjector:
|
||||
def __init__(self):
|
||||
import mlx.nn as nn
|
||||
|
||||
self.linear1 = nn.Linear(1024, 512, bias=False)
|
||||
self.linear2 = nn.Linear(512, 512, bias=False)
|
||||
|
||||
def __call__(self, x):
|
||||
import mlx.nn as nn
|
||||
|
||||
x = self.linear1(x)
|
||||
x = nn.relu(x)
|
||||
x = self.linear2(x)
|
||||
return x
|
||||
|
||||
|
||||
def _load_projector(projector_path: str) -> _MLPProjector:
|
||||
import mlx.core as mx
|
||||
from safetensors import safe_open
|
||||
|
||||
projector = _MLPProjector()
|
||||
with safe_open(projector_path, framework="numpy") as f:
|
||||
projector.linear1.weight = mx.array(f.get_tensor("linear1.weight"))
|
||||
projector.linear2.weight = mx.array(f.get_tensor("linear2.weight"))
|
||||
return projector
|
||||
|
||||
|
||||
def _sanitize(text: str, special_tokens: dict[str, str]) -> str:
|
||||
for token in special_tokens.values():
|
||||
text = text.replace(token, "")
|
||||
return text
|
||||
|
||||
|
||||
def _format_prompt(query: str, docs: list[str], special_tokens: dict[str, str]) -> str:
|
||||
query = _sanitize(query, special_tokens)
|
||||
docs = [_sanitize(d, special_tokens) for d in docs]
|
||||
|
||||
doc_token = special_tokens["doc_embed_token"]
|
||||
query_token = special_tokens["query_embed_token"]
|
||||
|
||||
prefix = (
|
||||
"<|im_start|>system\n"
|
||||
"You are a search relevance expert who can determine a ranking of the passages based on how relevant they are to the query. "
|
||||
"If the query is a question, how relevant a passage is depends on how well it answers the question. "
|
||||
"If not, try to analyze the intent of the query and assess how well each passage satisfies the intent. "
|
||||
"If an instruction is provided, you should follow the instruction when determining the ranking."
|
||||
"<|im_end|>\n<|im_start|>user\n"
|
||||
)
|
||||
suffix = "<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n"
|
||||
|
||||
body = (
|
||||
f"I will provide you with {len(docs)} passages, each indicated by a numerical identifier. "
|
||||
f"Rank the passages based on their relevance to query: {query}\n"
|
||||
)
|
||||
body += "\n".join(f'<passage id="{i}">\n{doc}{doc_token}\n</passage>' for i, doc in enumerate(docs))
|
||||
body += f"\n<query>\n{query}{query_token}\n</query>"
|
||||
return prefix + body + suffix
|
||||
|
||||
|
||||
class MLXReranker:
|
||||
"""
|
||||
MLX-accelerated jina-reranker-v3 for Apple Silicon.
|
||||
|
||||
Loads the model from a local directory (use huggingface_hub.snapshot_download
|
||||
to fetch jinaai/jina-reranker-v3-mlx if you don't have it already).
|
||||
"""
|
||||
|
||||
_SPECIAL_TOKENS = {
|
||||
"query_embed_token": "<|rerank_token|>",
|
||||
"doc_embed_token": "<|embed_token|>",
|
||||
}
|
||||
_DOC_TOKEN_ID = 151670
|
||||
_QUERY_TOKEN_ID = 151671
|
||||
|
||||
def __init__(self, model_path: str, projector_path: str):
|
||||
from mlx_lm import load
|
||||
|
||||
self.model, self.tokenizer = load(model_path)
|
||||
self.model.eval()
|
||||
self.projector = _load_projector(projector_path)
|
||||
|
||||
def rerank(self, query: str, documents: list[str], top_n: int | None = None) -> list[dict]:
|
||||
"""
|
||||
Rank documents by relevance to a query.
|
||||
|
||||
Returns a list of dicts with keys: document, relevance_score, index.
|
||||
Sorted by descending relevance_score.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
prompt = _format_prompt(query, documents, self._SPECIAL_TOKENS)
|
||||
input_ids = self.tokenizer.encode(prompt)
|
||||
hidden_states = self.model.model([input_ids])[0] # [seq_len, hidden_size]
|
||||
|
||||
input_ids_np = np.array(input_ids)
|
||||
query_positions = np.where(input_ids_np == self._QUERY_TOKEN_ID)[0]
|
||||
doc_positions = np.where(input_ids_np == self._DOC_TOKEN_ID)[0]
|
||||
|
||||
if len(query_positions) == 0:
|
||||
raise ValueError("Query embed token not found in prompt")
|
||||
if len(doc_positions) == 0:
|
||||
raise ValueError("Document embed tokens not found in prompt")
|
||||
|
||||
query_hidden = mx.expand_dims(hidden_states[int(query_positions[0])], axis=0)
|
||||
doc_hidden = mx.stack([hidden_states[int(p)] for p in doc_positions])
|
||||
|
||||
query_emb = self.projector(query_hidden) # [1, 512]
|
||||
doc_emb = self.projector(doc_hidden) # [num_docs, 512]
|
||||
|
||||
query_exp = mx.broadcast_to(mx.expand_dims(query_emb, 0), (1, len(documents), 512))
|
||||
doc_exp = mx.expand_dims(doc_emb, 0)
|
||||
|
||||
scores = mx.sum(doc_exp * query_exp, axis=-1) / (
|
||||
mx.sqrt(mx.sum(doc_exp * doc_exp, axis=-1)) * mx.sqrt(mx.sum(query_exp * query_exp, axis=-1))
|
||||
) # [1, num_docs]
|
||||
scores_np = np.array(scores[0])
|
||||
|
||||
order = np.argsort(scores_np)[::-1]
|
||||
n = min(top_n, len(documents)) if top_n is not None else len(documents)
|
||||
return [
|
||||
{
|
||||
"document": documents[order[i]],
|
||||
"relevance_score": float(scores_np[order[i]]),
|
||||
"index": int(order[i]),
|
||||
}
|
||||
for i in range(n)
|
||||
]
|
||||
@@ -1,128 +0,0 @@
|
||||
"""File parser implementations."""
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .base import FileParser, UnsupportedFileTypeError
|
||||
from .iris import IrisParser
|
||||
from .markitdown import MarkitdownParser
|
||||
|
||||
__all__ = [
|
||||
"FileParser",
|
||||
"UnsupportedFileTypeError",
|
||||
"IrisParser",
|
||||
"MarkitdownParser",
|
||||
"FileParserRegistry",
|
||||
"ConvertResult",
|
||||
]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConvertResult:
|
||||
"""Result of a successful file conversion."""
|
||||
|
||||
content: str
|
||||
parser_name: str
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class FileParserRegistry:
|
||||
"""Registry for file parsers with auto-detection."""
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize empty parser registry."""
|
||||
self._parsers: dict[str, FileParser] = {}
|
||||
|
||||
def register(self, parser: FileParser):
|
||||
"""
|
||||
Register a parser.
|
||||
|
||||
Args:
|
||||
parser: FileParser instance
|
||||
"""
|
||||
self._parsers[parser.name()] = parser
|
||||
|
||||
def get_parser(
|
||||
self,
|
||||
name: str | None,
|
||||
filename: str,
|
||||
content_type: str | None = None,
|
||||
) -> FileParser:
|
||||
"""
|
||||
Get parser by name or auto-detect.
|
||||
|
||||
Args:
|
||||
name: Parser name (e.g., "markitdown") or None for auto-detect
|
||||
filename: File name for auto-detection
|
||||
content_type: MIME type (optional)
|
||||
|
||||
Returns:
|
||||
FileParser instance
|
||||
|
||||
Raises:
|
||||
ValueError: If no suitable parser found
|
||||
"""
|
||||
if name:
|
||||
# Explicit parser requested — return it directly, let the parser
|
||||
# raise UnsupportedFileTypeError from convert() if needed
|
||||
if name not in self._parsers:
|
||||
raise ValueError(f"Parser '{name}' not found. Available: {list(self._parsers.keys())}")
|
||||
return self._parsers[name]
|
||||
|
||||
# Auto-detect parser
|
||||
for parser in self._parsers.values():
|
||||
if parser.supports(filename, content_type):
|
||||
return parser
|
||||
|
||||
raise ValueError(f"No parser found for {filename}. Available parsers: {list(self._parsers.keys())}")
|
||||
|
||||
async def convert_with_fallback(
|
||||
self,
|
||||
parsers: list[str],
|
||||
file_data: bytes,
|
||||
filename: str,
|
||||
content_type: str | None = None,
|
||||
) -> ConvertResult:
|
||||
"""
|
||||
Try each parser in order, falling back on failure or empty content.
|
||||
|
||||
Moves to the next parser if the current one raises UnsupportedFileTypeError
|
||||
or returns empty content. Any other exception (RuntimeError, network error,
|
||||
etc.) also triggers a fallback so the chain is exhausted before failing.
|
||||
|
||||
Args:
|
||||
parsers: Ordered list of parser names to try
|
||||
file_data: Raw file bytes
|
||||
filename: Original filename
|
||||
content_type: MIME type (optional)
|
||||
|
||||
Returns:
|
||||
ConvertResult with the parsed content and the name of the parser that succeeded
|
||||
|
||||
Raises:
|
||||
ValueError: If a parser name is not registered
|
||||
RuntimeError: If all parsers fail or return empty content
|
||||
"""
|
||||
last_error: Exception | None = None
|
||||
for name in parsers:
|
||||
parser = self.get_parser(name, filename, content_type)
|
||||
try:
|
||||
content = await parser.convert(file_data, filename)
|
||||
if content and content.strip():
|
||||
return ConvertResult(content=content, parser_name=name)
|
||||
logger.warning(f"Parser '{name}' returned empty content for '{filename}', trying next")
|
||||
last_error = RuntimeError(f"Parser '{name}' returned no content for '{filename}'")
|
||||
except UnsupportedFileTypeError as e:
|
||||
logger.warning(f"Parser '{name}' does not support '{filename}', trying next: {e}")
|
||||
last_error = e
|
||||
except Exception as e:
|
||||
logger.warning(f"Parser '{name}' failed for '{filename}', trying next: {e}")
|
||||
last_error = e
|
||||
|
||||
raise last_error or RuntimeError(f"No parsers available for '{filename}'")
|
||||
|
||||
def list_parsers(self) -> list[str]:
|
||||
"""Get list of registered parser names."""
|
||||
return list(self._parsers.keys())
|
||||
@@ -1,544 +0,0 @@
|
||||
"""
|
||||
Main orchestrator for the retain pipeline.
|
||||
|
||||
Coordinates all retain pipeline modules to store memories efficiently.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from ..db_utils import acquire_with_retry, retry_with_backoff
|
||||
from . import bank_utils
|
||||
|
||||
|
||||
def utcnow():
|
||||
"""Get current UTC time."""
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
def parse_datetime_flexible(value: Any) -> datetime:
|
||||
"""
|
||||
Parse a datetime value that could be either a datetime object or an ISO string.
|
||||
|
||||
This handles datetime values from both direct Python calls and deserialized JSON
|
||||
(where datetime objects are serialized as ISO strings).
|
||||
|
||||
Args:
|
||||
value: Either a datetime object or an ISO format string
|
||||
|
||||
Returns:
|
||||
datetime object (timezone-aware)
|
||||
|
||||
Raises:
|
||||
TypeError: If value is neither datetime nor string
|
||||
ValueError: If string is not a valid ISO datetime
|
||||
"""
|
||||
if isinstance(value, datetime):
|
||||
# Ensure timezone-aware
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=UTC)
|
||||
return value
|
||||
elif isinstance(value, str):
|
||||
# Parse ISO format string (handles both 'Z' and '+00:00' timezone formats)
|
||||
dt = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
# Ensure timezone-aware
|
||||
if dt.tzinfo is None:
|
||||
return dt.replace(tzinfo=UTC)
|
||||
return dt
|
||||
else:
|
||||
raise TypeError(f"Expected datetime or string, got {type(value).__name__}")
|
||||
|
||||
|
||||
import asyncpg
|
||||
|
||||
from ..response_models import TokenUsage
|
||||
from . import (
|
||||
chunk_storage,
|
||||
embedding_processing,
|
||||
entity_processing,
|
||||
fact_extraction,
|
||||
fact_storage,
|
||||
link_creation,
|
||||
)
|
||||
from .types import EntityLink, ExtractedFact, ProcessedFact, RetainContent, RetainContentDict
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def retain_batch(
|
||||
pool,
|
||||
embeddings_model,
|
||||
llm_config,
|
||||
entity_resolver,
|
||||
format_date_fn,
|
||||
bank_id: str,
|
||||
contents_dicts: list[RetainContentDict],
|
||||
config,
|
||||
document_id: str | None = None,
|
||||
is_first_batch: bool = True,
|
||||
fact_type_override: str | None = None,
|
||||
confidence_score: float | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
operation_id: str | None = None,
|
||||
schema: str | None = None,
|
||||
outbox_callback: Callable[["asyncpg.Connection"], Awaitable[None]] | None = None,
|
||||
) -> tuple[list[list[str]], TokenUsage]:
|
||||
"""
|
||||
Process a batch of content through the retain pipeline.
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
embeddings_model: Embeddings model for generating embeddings
|
||||
llm_config: LLM configuration for fact extraction
|
||||
entity_resolver: Entity resolver for entity processing
|
||||
format_date_fn: Function to format datetime to readable string
|
||||
bank_id: Bank identifier
|
||||
contents_dicts: List of content dictionaries
|
||||
config: Resolved HindsightConfig for this bank
|
||||
document_id: Optional document ID
|
||||
is_first_batch: Whether this is the first batch
|
||||
fact_type_override: Override fact type for all facts
|
||||
confidence_score: Confidence score for opinions
|
||||
document_tags: Tags applied to all items in this batch
|
||||
|
||||
Returns:
|
||||
Tuple of (unit ID lists, token usage for fact extraction)
|
||||
"""
|
||||
start_time = time.time()
|
||||
total_chars = sum(len(item.get("content", "")) for item in contents_dicts)
|
||||
|
||||
# Buffer all logs
|
||||
log_buffer = []
|
||||
log_buffer.append(f"{'=' * 60}")
|
||||
log_buffer.append(f"RETAIN_BATCH START: {bank_id}")
|
||||
log_buffer.append(f"Batch size: {len(contents_dicts)} content items, {total_chars:,} chars")
|
||||
log_buffer.append(f"{'=' * 60}")
|
||||
|
||||
# Get bank profile
|
||||
profile = await bank_utils.get_bank_profile(pool, bank_id)
|
||||
agent_name = profile["name"]
|
||||
|
||||
# Convert dicts to RetainContent objects
|
||||
contents = []
|
||||
for item in contents_dicts:
|
||||
# Merge item-level tags with document-level tags
|
||||
item_tags = item.get("tags", []) or []
|
||||
merged_tags = list(set(item_tags + (document_tags or [])))
|
||||
|
||||
# Handle event_date: distinguish "not provided" (default to now) from
|
||||
# "explicitly None" (caller opted into no timestamp).
|
||||
if "event_date" in item and item["event_date"] is None:
|
||||
event_date_value = None # Caller explicitly signalled "unknown date"
|
||||
elif item.get("event_date"):
|
||||
event_date_value = parse_datetime_flexible(item["event_date"])
|
||||
else:
|
||||
event_date_value = utcnow() # Backward-compatible default
|
||||
|
||||
content = RetainContent(
|
||||
content=item["content"],
|
||||
context=item.get("context", ""),
|
||||
event_date=event_date_value,
|
||||
metadata=item.get("metadata", {}),
|
||||
entities=item.get("entities", []),
|
||||
tags=merged_tags,
|
||||
observation_scopes=item.get("observation_scopes"),
|
||||
)
|
||||
contents.append(content)
|
||||
|
||||
# Step 1: Extract facts from all contents
|
||||
step_start = time.time()
|
||||
|
||||
extracted_facts, chunks, usage = await fact_extraction.extract_facts_from_contents(
|
||||
contents, llm_config, agent_name, config, pool, operation_id, schema
|
||||
)
|
||||
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:
|
||||
# Still need to create document if document_id was provided or chunks exist
|
||||
from collections import defaultdict
|
||||
|
||||
docs_tracked = 0
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
# Group contents by document_id (consistent with normal path)
|
||||
contents_by_doc_early = defaultdict(list)
|
||||
for idx, content_dict in enumerate(contents_dicts):
|
||||
doc_id = content_dict.get("document_id")
|
||||
contents_by_doc_early[doc_id].append((idx, content_dict))
|
||||
|
||||
if document_id:
|
||||
# Legacy: single document_id parameter
|
||||
combined_content = "\n".join([c.get("content", "") for c in contents_dicts])
|
||||
# Collect tags from all content items and merge with document_tags
|
||||
all_tags = set(document_tags or [])
|
||||
for item in contents_dicts:
|
||||
item_tags = item.get("tags", []) or []
|
||||
all_tags.update(item_tags)
|
||||
merged_tags = list(all_tags)
|
||||
|
||||
retain_params = {}
|
||||
if contents_dicts:
|
||||
first_item = contents_dicts[0]
|
||||
if first_item.get("context"):
|
||||
retain_params["context"] = first_item["context"]
|
||||
if first_item.get("event_date"):
|
||||
retain_params["event_date"] = (
|
||||
first_item["event_date"].isoformat()
|
||||
if hasattr(first_item["event_date"], "isoformat")
|
||||
else str(first_item["event_date"])
|
||||
)
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, merged_tags
|
||||
)
|
||||
docs_tracked += 1
|
||||
else:
|
||||
# Handle per-item document_ids and/or chunks (mirrors normal path logic)
|
||||
has_any_doc_ids = any(item.get("document_id") for item in contents_dicts)
|
||||
|
||||
if has_any_doc_ids or chunks:
|
||||
for original_doc_id, doc_contents in contents_by_doc_early.items():
|
||||
should_create_doc = (original_doc_id is not None) or chunks
|
||||
if not should_create_doc:
|
||||
continue
|
||||
|
||||
actual_doc_id = original_doc_id
|
||||
if actual_doc_id is None:
|
||||
# No document_id but have chunks - generate one
|
||||
actual_doc_id = str(uuid.uuid4())
|
||||
|
||||
combined_content = "\n".join([c.get("content", "") for _, c in doc_contents])
|
||||
all_tags = set(document_tags or [])
|
||||
for _, item in doc_contents:
|
||||
item_tags = item.get("tags", []) or []
|
||||
all_tags.update(item_tags)
|
||||
merged_tags = list(all_tags)
|
||||
|
||||
retain_params = {}
|
||||
if doc_contents:
|
||||
first_item = doc_contents[0][1]
|
||||
if first_item.get("context"):
|
||||
retain_params["context"] = first_item["context"]
|
||||
if first_item.get("event_date"):
|
||||
retain_params["event_date"] = (
|
||||
first_item["event_date"].isoformat()
|
||||
if hasattr(first_item["event_date"], "isoformat")
|
||||
else str(first_item["event_date"])
|
||||
)
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn,
|
||||
bank_id,
|
||||
actual_doc_id,
|
||||
combined_content,
|
||||
is_first_batch,
|
||||
retain_params,
|
||||
merged_tags,
|
||||
)
|
||||
docs_tracked += 1
|
||||
|
||||
total_time = time.time() - start_time
|
||||
doc_status = f"{docs_tracked} document(s) tracked" if docs_tracked > 0 else "no document tracked"
|
||||
logger.info(
|
||||
f"RETAIN_BATCH COMPLETE: 0 facts extracted from {len(contents)} contents in {total_time:.3f}s ({doc_status}, no facts)"
|
||||
)
|
||||
return [[] for _ in contents], usage
|
||||
|
||||
# Apply fact_type_override if provided
|
||||
if fact_type_override:
|
||||
for fact in extracted_facts:
|
||||
fact.fact_type = fact_type_override
|
||||
|
||||
# Step 2: Augment texts and generate embeddings
|
||||
step_start = time.time()
|
||||
augmented_texts = embedding_processing.augment_texts_with_dates(extracted_facts, format_date_fn)
|
||||
embeddings = await embedding_processing.generate_embeddings_batch(embeddings_model, augmented_texts)
|
||||
log_buffer.append(f"[2] Generate embeddings: {len(embeddings)} embeddings in {time.time() - step_start:.3f}s")
|
||||
|
||||
# Step 3: Convert to ProcessedFact objects (without chunk_ids yet)
|
||||
processed_facts = [
|
||||
ProcessedFact.from_extracted_fact(extracted_fact, embedding)
|
||||
for extracted_fact, embedding in zip(extracted_facts, embeddings)
|
||||
]
|
||||
|
||||
# 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")
|
||||
contents_by_doc[doc_id].append((idx, content_dict))
|
||||
|
||||
# Step 4: Database transaction (retried on deadlock)
|
||||
result_unit_ids: list[list[str]] = []
|
||||
|
||||
log_buffer_pre_db = len(log_buffer)
|
||||
|
||||
async def _run_db_work() -> None:
|
||||
nonlocal result_unit_ids
|
||||
|
||||
# Reset per-fact mutations and log buffer so each retry attempt starts clean
|
||||
del log_buffer[log_buffer_pre_db:]
|
||||
document_ids_added: list[str] = []
|
||||
for pf in processed_facts:
|
||||
pf.document_id = None
|
||||
pf.chunk_id = None
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
# Handle document tracking for all documents
|
||||
step_start = time.time()
|
||||
# Map None document_id to generated UUIDs
|
||||
doc_id_mapping = {} # Maps original doc_id (including None) to actual doc_id used
|
||||
|
||||
if document_id:
|
||||
# Legacy: single document_id parameter
|
||||
combined_content = "\n".join([c.get("content", "") for c in contents_dicts])
|
||||
retain_params = {}
|
||||
# Collect tags from all content items and merge with document_tags
|
||||
all_tags = set(document_tags or [])
|
||||
for item in contents_dicts:
|
||||
item_tags = item.get("tags", []) or []
|
||||
all_tags.update(item_tags)
|
||||
merged_tags = list(all_tags)
|
||||
|
||||
if contents_dicts:
|
||||
first_item = contents_dicts[0]
|
||||
if first_item.get("context"):
|
||||
retain_params["context"] = first_item["context"]
|
||||
if first_item.get("event_date"):
|
||||
retain_params["event_date"] = (
|
||||
first_item["event_date"].isoformat()
|
||||
if hasattr(first_item["event_date"], "isoformat")
|
||||
else str(first_item["event_date"])
|
||||
)
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, merged_tags
|
||||
)
|
||||
document_ids_added.append(document_id)
|
||||
doc_id_mapping[None] = document_id # For backwards compatibility
|
||||
else:
|
||||
# Handle per-item document_ids (create documents if any item has document_id or if chunks exist)
|
||||
has_any_doc_ids = any(item.get("document_id") for item in contents_dicts)
|
||||
|
||||
if has_any_doc_ids or chunks:
|
||||
for original_doc_id, doc_contents in contents_by_doc.items():
|
||||
actual_doc_id = original_doc_id
|
||||
|
||||
# Only create document record if:
|
||||
# 1. Item has explicit document_id, OR
|
||||
# 2. There are chunks (need document for chunk storage)
|
||||
should_create_doc = (original_doc_id is not None) or chunks
|
||||
|
||||
if should_create_doc:
|
||||
if actual_doc_id is None:
|
||||
# No document_id but have chunks - generate one
|
||||
actual_doc_id = str(uuid.uuid4())
|
||||
|
||||
# Store mapping for later use
|
||||
doc_id_mapping[original_doc_id] = actual_doc_id
|
||||
|
||||
# Combine content for this document
|
||||
combined_content = "\n".join([c.get("content", "") for _, c in doc_contents])
|
||||
|
||||
# Collect tags from all content items for this document and merge with document_tags
|
||||
all_tags = set(document_tags or [])
|
||||
for _, item in doc_contents:
|
||||
item_tags = item.get("tags", []) or []
|
||||
all_tags.update(item_tags)
|
||||
merged_tags = list(all_tags)
|
||||
|
||||
# Extract retain params from first content item
|
||||
retain_params = {}
|
||||
if doc_contents:
|
||||
first_item = doc_contents[0][1]
|
||||
if first_item.get("context"):
|
||||
retain_params["context"] = first_item["context"]
|
||||
if first_item.get("event_date"):
|
||||
retain_params["event_date"] = (
|
||||
first_item["event_date"].isoformat()
|
||||
if hasattr(first_item["event_date"], "isoformat")
|
||||
else str(first_item["event_date"])
|
||||
)
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn,
|
||||
bank_id,
|
||||
actual_doc_id,
|
||||
combined_content,
|
||||
is_first_batch,
|
||||
retain_params,
|
||||
merged_tags,
|
||||
)
|
||||
document_ids_added.append(actual_doc_id)
|
||||
|
||||
if document_ids_added:
|
||||
log_buffer.append(
|
||||
f"[2.5] Document tracking: {len(document_ids_added)} documents in {time.time() - step_start:.3f}s"
|
||||
)
|
||||
|
||||
# Store chunks and map to facts for all documents
|
||||
step_start = time.time()
|
||||
chunk_id_map_by_doc = {} # Maps (doc_id, chunk_index) -> chunk_id
|
||||
|
||||
if chunks:
|
||||
# Group chunks by their source document
|
||||
chunks_by_doc = defaultdict(list)
|
||||
for chunk in chunks:
|
||||
# chunk.content_index tells us which content this chunk came from
|
||||
original_doc_id = contents_dicts[chunk.content_index].get("document_id")
|
||||
# Map to actual document_id (handles None -> generated UUID mapping)
|
||||
actual_doc_id = doc_id_mapping.get(original_doc_id, original_doc_id)
|
||||
if actual_doc_id is None and document_id:
|
||||
actual_doc_id = document_id
|
||||
chunks_by_doc[actual_doc_id].append(chunk)
|
||||
|
||||
# Store chunks for each document
|
||||
for doc_id, doc_chunks in chunks_by_doc.items():
|
||||
chunk_id_map = await chunk_storage.store_chunks_batch(conn, bank_id, doc_id, doc_chunks)
|
||||
# Store mapping with document context
|
||||
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"
|
||||
)
|
||||
|
||||
# Map chunk_ids and document_ids to facts
|
||||
for fact, processed_fact in zip(extracted_facts, processed_facts):
|
||||
# Get the original document_id for this fact's source content
|
||||
original_doc_id = contents_dicts[fact.content_index].get("document_id")
|
||||
# Map to actual document_id (handles None -> generated UUID mapping)
|
||||
actual_doc_id = doc_id_mapping.get(original_doc_id, original_doc_id)
|
||||
if actual_doc_id is None and document_id:
|
||||
actual_doc_id = document_id
|
||||
|
||||
# Set document_id on the fact
|
||||
processed_fact.document_id = actual_doc_id
|
||||
|
||||
# Map chunk_id if this fact came from a chunk
|
||||
if fact.chunk_index is not None:
|
||||
# Look up chunk_id using (doc_id, chunk_index)
|
||||
chunk_id = chunk_id_map_by_doc.get((actual_doc_id, fact.chunk_index))
|
||||
if chunk_id:
|
||||
processed_fact.chunk_id = chunk_id
|
||||
else:
|
||||
# No chunks - still need to set document_id on facts
|
||||
for fact, processed_fact in zip(extracted_facts, processed_facts):
|
||||
original_doc_id = contents_dicts[fact.content_index].get("document_id")
|
||||
# Map to actual document_id (handles None -> generated UUID mapping)
|
||||
actual_doc_id = doc_id_mapping.get(original_doc_id, original_doc_id)
|
||||
if actual_doc_id is None and document_id:
|
||||
actual_doc_id = document_id
|
||||
processed_fact.document_id = actual_doc_id
|
||||
|
||||
non_duplicate_facts = processed_facts
|
||||
|
||||
# Insert facts (document_id is now stored per-fact)
|
||||
step_start = time.time()
|
||||
unit_ids = await fact_storage.insert_facts_batch(conn, bank_id, non_duplicate_facts)
|
||||
log_buffer.append(f"[5] Insert facts: {len(unit_ids)} units in {time.time() - step_start:.3f}s")
|
||||
|
||||
# Process entities
|
||||
step_start = time.time()
|
||||
# Build map of content_index -> user entities for merging
|
||||
user_entities_per_content = {
|
||||
idx: content.entities for idx, content in enumerate(contents) if content.entities
|
||||
}
|
||||
entity_links = await entity_processing.process_entities_batch(
|
||||
entity_resolver,
|
||||
conn,
|
||||
bank_id,
|
||||
unit_ids,
|
||||
non_duplicate_facts,
|
||||
log_buffer,
|
||||
user_entities_per_content=user_entities_per_content,
|
||||
entity_labels=getattr(config, "entity_labels", None),
|
||||
)
|
||||
log_buffer.append(f"[6] Process entities: {len(entity_links)} links in {time.time() - step_start:.3f}s")
|
||||
|
||||
# Create temporal links
|
||||
step_start = time.time()
|
||||
temporal_link_count = await link_creation.create_temporal_links_batch(conn, bank_id, unit_ids)
|
||||
log_buffer.append(f"[7] Temporal links: {temporal_link_count} links in {time.time() - step_start:.3f}s")
|
||||
|
||||
# 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
|
||||
)
|
||||
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"
|
||||
)
|
||||
|
||||
# Create causal links
|
||||
step_start = time.time()
|
||||
causal_link_count = await link_creation.create_causal_links_batch(conn, unit_ids, non_duplicate_facts)
|
||||
log_buffer.append(f"[10] Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s")
|
||||
|
||||
# Map results back to original content items
|
||||
result_unit_ids = _map_results_to_contents(contents, extracted_facts, unit_ids)
|
||||
|
||||
# Transactional outbox: queue any side-effect tasks (e.g. webhook deliveries)
|
||||
# inside the same transaction so they are atomically committed with the retain data.
|
||||
if outbox_callback:
|
||||
await outbox_callback(conn)
|
||||
|
||||
# Flush entity stats (mention_count / last_seen) now that the transaction
|
||||
# has committed. Uses a fresh pool connection — no locks held.
|
||||
await entity_resolver.flush_pending_stats()
|
||||
|
||||
# Log final summary
|
||||
total_time = time.time() - start_time
|
||||
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}")
|
||||
|
||||
logger.info("\n" + "\n".join(log_buffer) + "\n")
|
||||
|
||||
await retry_with_backoff(_run_db_work)
|
||||
return result_unit_ids, usage
|
||||
|
||||
|
||||
def _map_results_to_contents(
|
||||
contents: list[RetainContent],
|
||||
extracted_facts: list[ExtractedFact],
|
||||
unit_ids: list[str],
|
||||
) -> list[list[str]]:
|
||||
"""Map created unit IDs back to original content items."""
|
||||
facts_by_content: dict[int, list[int]] = {i: [] for i in range(len(contents))}
|
||||
for i, fact in enumerate(extracted_facts):
|
||||
facts_by_content[fact.content_index].append(i)
|
||||
|
||||
result_unit_ids = []
|
||||
unit_idx = 0
|
||||
for content_index in range(len(contents)):
|
||||
content_unit_ids = []
|
||||
for _ in facts_by_content[content_index]:
|
||||
content_unit_ids.append(unit_ids[unit_idx])
|
||||
unit_idx += 1
|
||||
result_unit_ids.append(content_unit_ids)
|
||||
|
||||
return result_unit_ids
|
||||
@@ -1,390 +0,0 @@
|
||||
"""
|
||||
Tags filtering utilities for retrieval.
|
||||
|
||||
Provides SQL building functions for filtering memories by tags.
|
||||
Supports four matching modes via TagsMatch enum:
|
||||
- "any": OR matching, includes untagged memories (default, backward compatible)
|
||||
- "all": AND matching, includes untagged memories
|
||||
- "any_strict": OR matching, excludes untagged memories
|
||||
- "all_strict": AND matching, excludes untagged memories
|
||||
|
||||
OR matching (any/any_strict): Memory matches if ANY of its tags overlap with request tags
|
||||
AND matching (all/all_strict): Memory matches if ALL request tags are present in its tags
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
TagsMatch = Literal["any", "all", "any_strict", "all_strict"]
|
||||
|
||||
|
||||
def _parse_tags_match(match: TagsMatch) -> tuple[str, bool]:
|
||||
"""
|
||||
Parse TagsMatch into operator and include_untagged flag.
|
||||
|
||||
Returns:
|
||||
Tuple of (operator, include_untagged)
|
||||
- operator: "&&" for any/any_strict, "@>" for all/all_strict
|
||||
- include_untagged: True for any/all, False for any_strict/all_strict
|
||||
"""
|
||||
if match == "any":
|
||||
return "&&", True
|
||||
elif match == "all":
|
||||
return "@>", True
|
||||
elif match == "any_strict":
|
||||
return "&&", False
|
||||
elif match == "all_strict":
|
||||
return "@>", False
|
||||
else:
|
||||
# Default to "any" behavior
|
||||
return "&&", True
|
||||
|
||||
|
||||
def build_tags_where_clause(
|
||||
tags: list[str] | None,
|
||||
param_offset: int = 1,
|
||||
table_alias: str = "",
|
||||
match: TagsMatch = "any",
|
||||
) -> tuple[str, list, int]:
|
||||
"""
|
||||
Build a SQL WHERE clause for filtering by tags.
|
||||
|
||||
Supports four matching modes:
|
||||
- "any" (default): OR matching, includes untagged memories
|
||||
- "all": AND matching, includes untagged memories
|
||||
- "any_strict": OR matching, excludes untagged memories
|
||||
- "all_strict": AND matching, excludes untagged memories
|
||||
|
||||
Args:
|
||||
tags: List of tags to filter by. If None or empty, returns empty clause (no filtering).
|
||||
param_offset: Starting parameter number for SQL placeholders (default 1).
|
||||
table_alias: Optional table alias prefix (e.g., "mu." for "memory_units mu").
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
Tuple of (sql_clause, params, next_param_offset):
|
||||
- sql_clause: SQL WHERE clause string
|
||||
- params: List of parameter values to bind
|
||||
- next_param_offset: Next available parameter number
|
||||
|
||||
Example:
|
||||
>>> clause, params, next_offset = build_tags_where_clause(['user_a'], 3, 'mu.', 'any_strict')
|
||||
>>> print(clause) # "AND mu.tags IS NOT NULL AND mu.tags != '{}' AND mu.tags && $3"
|
||||
"""
|
||||
if not tags:
|
||||
return "", [], param_offset
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
operator, include_untagged = _parse_tags_match(match)
|
||||
|
||||
if include_untagged:
|
||||
# Include untagged memories (NULL or empty array) OR matching tags
|
||||
clause = f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_offset})"
|
||||
else:
|
||||
# Strict: only memories with matching tags (exclude NULL and empty)
|
||||
clause = f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_offset}"
|
||||
|
||||
return clause, [tags], param_offset + 1
|
||||
|
||||
|
||||
def build_tags_where_clause_simple(
|
||||
tags: list[str] | None,
|
||||
param_num: int,
|
||||
table_alias: str = "",
|
||||
match: TagsMatch = "any",
|
||||
) -> str:
|
||||
"""
|
||||
Build a simple SQL WHERE clause for tags filtering.
|
||||
|
||||
This is a convenience version that returns just the clause string,
|
||||
assuming the caller will add the tags array to their params list.
|
||||
|
||||
Args:
|
||||
tags: List of tags to filter by. If None or empty, returns empty string.
|
||||
param_num: Parameter number to use in the clause.
|
||||
table_alias: Optional table alias prefix.
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
SQL clause string or empty string.
|
||||
"""
|
||||
if not tags:
|
||||
return ""
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
operator, include_untagged = _parse_tags_match(match)
|
||||
|
||||
if include_untagged:
|
||||
# Include untagged memories (NULL or empty array) OR matching tags
|
||||
return f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_num})"
|
||||
else:
|
||||
# Strict: only memories with matching tags (exclude NULL and empty)
|
||||
return f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_num}"
|
||||
|
||||
|
||||
def filter_results_by_tags(
|
||||
results: list,
|
||||
tags: list[str] | None,
|
||||
match: TagsMatch = "any",
|
||||
) -> list:
|
||||
"""
|
||||
Filter retrieval results by tags in Python (for post-processing).
|
||||
|
||||
Used when SQL filtering isn't possible (e.g., graph traversal results).
|
||||
|
||||
Args:
|
||||
results: List of RetrievalResult objects with a 'tags' attribute.
|
||||
tags: List of tags to filter by. If None or empty, returns all results.
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
Filtered list of results.
|
||||
"""
|
||||
if not tags:
|
||||
return results
|
||||
|
||||
_, include_untagged = _parse_tags_match(match)
|
||||
is_any_match = match in ("any", "any_strict")
|
||||
|
||||
tags_set = set(tags)
|
||||
filtered = []
|
||||
|
||||
for result in results:
|
||||
result_tags = getattr(result, "tags", None)
|
||||
|
||||
# Check if untagged
|
||||
is_untagged = result_tags is None or len(result_tags) == 0
|
||||
|
||||
if is_untagged:
|
||||
if include_untagged:
|
||||
filtered.append(result)
|
||||
# else: skip untagged
|
||||
else:
|
||||
result_tags_set = set(result_tags)
|
||||
if is_any_match:
|
||||
# Any overlap
|
||||
if result_tags_set & tags_set:
|
||||
filtered.append(result)
|
||||
else:
|
||||
# All tags must be present
|
||||
if tags_set <= result_tags_set:
|
||||
filtered.append(result)
|
||||
|
||||
return filtered
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Compound tag group models (recursive boolean expressions)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TagGroupLeaf(BaseModel):
|
||||
"""A leaf tag filter: matches memories by tag list and match mode."""
|
||||
|
||||
tags: list[str]
|
||||
match: TagsMatch = "any_strict"
|
||||
|
||||
|
||||
class TagGroupAnd(BaseModel):
|
||||
"""Compound AND group: all child filters must match."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
filters: list[TagGroup] = Field(alias="and")
|
||||
|
||||
|
||||
class TagGroupOr(BaseModel):
|
||||
"""Compound OR group: at least one child filter must match."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
filters: list[TagGroup] = Field(alias="or")
|
||||
|
||||
|
||||
class TagGroupNot(BaseModel):
|
||||
"""Compound NOT group: child filter must NOT match."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
filter: TagGroup = Field(alias="not")
|
||||
|
||||
|
||||
# TagGroup is a discriminated union; Pydantic will try left-to-right.
|
||||
# TagGroupLeaf is identified by the presence of 'tags'.
|
||||
# TagGroupAnd / TagGroupOr / TagGroupNot are compound (no 'tags' key).
|
||||
TagGroup = Annotated[
|
||||
TagGroupLeaf | TagGroupAnd | TagGroupOr | TagGroupNot,
|
||||
Field(union_mode="left_to_right"),
|
||||
]
|
||||
|
||||
# Rebuild forward-reference models so recursive TagGroup is resolved.
|
||||
TagGroupAnd.model_rebuild()
|
||||
TagGroupOr.model_rebuild()
|
||||
TagGroupNot.model_rebuild()
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# SQL builder for compound tag groups
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def _build_group_clause(
|
||||
group: TagGroup,
|
||||
param_offset: int,
|
||||
table_alias: str,
|
||||
) -> tuple[str, list, int]:
|
||||
"""
|
||||
Recursively build an inner SQL clause (no leading AND/OR) for a single TagGroup.
|
||||
|
||||
Returns:
|
||||
(inner_clause, params, next_param_offset)
|
||||
"""
|
||||
if isinstance(group, TagGroupLeaf):
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
operator, include_untagged = _parse_tags_match(group.match)
|
||||
if include_untagged:
|
||||
clause = f"({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_offset})"
|
||||
else:
|
||||
clause = f"({column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_offset})"
|
||||
return clause, [group.tags], param_offset + 1
|
||||
|
||||
elif isinstance(group, TagGroupAnd):
|
||||
parts = []
|
||||
params: list = []
|
||||
offset = param_offset
|
||||
for child in group.filters:
|
||||
child_clause, child_params, offset = _build_group_clause(child, offset, table_alias)
|
||||
parts.append(child_clause)
|
||||
params.extend(child_params)
|
||||
inner = " AND ".join(parts)
|
||||
return f"({inner})", params, offset
|
||||
|
||||
elif isinstance(group, TagGroupOr):
|
||||
parts = []
|
||||
params = []
|
||||
offset = param_offset
|
||||
for child in group.filters:
|
||||
child_clause, child_params, offset = _build_group_clause(child, offset, table_alias)
|
||||
parts.append(child_clause)
|
||||
params.extend(child_params)
|
||||
inner = " OR ".join(parts)
|
||||
return f"({inner})", params, offset
|
||||
|
||||
elif isinstance(group, TagGroupNot):
|
||||
child_clause, child_params, next_offset = _build_group_clause(group.filter, param_offset, table_alias)
|
||||
return f"NOT {child_clause}", child_params, next_offset
|
||||
|
||||
else:
|
||||
# Should never happen with proper Pydantic validation
|
||||
return "", [], param_offset
|
||||
|
||||
|
||||
def build_tag_groups_where_clause(
|
||||
tag_groups: list[TagGroup] | None,
|
||||
param_offset: int,
|
||||
table_alias: str = "",
|
||||
) -> tuple[str, list, int]:
|
||||
"""
|
||||
Build a SQL WHERE clause for compound tag group filtering.
|
||||
|
||||
Top-level groups are AND-ed together. Each group is a recursive boolean
|
||||
expression (leaf, and, or, not).
|
||||
|
||||
Args:
|
||||
tag_groups: List of TagGroup objects. If None or empty, returns empty clause.
|
||||
param_offset: Starting parameter number for SQL placeholders.
|
||||
table_alias: Optional table alias prefix (e.g., "mu." for "memory_units mu").
|
||||
|
||||
Returns:
|
||||
Tuple of (sql_clause, params, next_param_offset):
|
||||
- sql_clause: SQL WHERE clause string starting with "AND" (or empty string)
|
||||
- params: List of parameter values to bind (one per leaf node)
|
||||
- next_param_offset: Next available parameter number
|
||||
|
||||
Example:
|
||||
>>> groups = [TagGroupLeaf(tags=["user:alice"], match="all_strict")]
|
||||
>>> clause, params, next_offset = build_tag_groups_where_clause(groups, 3)
|
||||
>>> print(clause) # "AND (tags IS NOT NULL AND tags != '{}' AND tags @> $3)"
|
||||
"""
|
||||
if not tag_groups:
|
||||
return "", [], param_offset
|
||||
|
||||
all_params: list = []
|
||||
all_clauses: list[str] = []
|
||||
offset = param_offset
|
||||
|
||||
for group in tag_groups:
|
||||
inner_clause, group_params, offset = _build_group_clause(group, offset, table_alias)
|
||||
all_clauses.append(inner_clause)
|
||||
all_params.extend(group_params)
|
||||
|
||||
combined = " AND ".join(all_clauses)
|
||||
return f"AND {combined}", all_params, offset
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Python-side filter for compound tag groups (post-retrieval filtering)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def _match_group(result: object, group: TagGroup) -> bool:
|
||||
"""
|
||||
Recursively evaluate a TagGroup against a retrieval result.
|
||||
|
||||
Args:
|
||||
result: Any object with a 'tags' attribute (list[str] or None).
|
||||
group: The TagGroup to evaluate.
|
||||
|
||||
Returns:
|
||||
True if the result matches the group, False otherwise.
|
||||
"""
|
||||
if isinstance(group, TagGroupLeaf):
|
||||
result_tags = getattr(result, "tags", None)
|
||||
is_untagged = result_tags is None or len(result_tags) == 0
|
||||
_, include_untagged = _parse_tags_match(group.match)
|
||||
is_any_match = group.match in ("any", "any_strict")
|
||||
tags_set = set(group.tags)
|
||||
|
||||
if is_untagged:
|
||||
return include_untagged
|
||||
else:
|
||||
result_tags_set = set(result_tags)
|
||||
if is_any_match:
|
||||
return bool(result_tags_set & tags_set)
|
||||
else:
|
||||
return tags_set <= result_tags_set
|
||||
|
||||
elif isinstance(group, TagGroupAnd):
|
||||
return all(_match_group(result, child) for child in group.filters)
|
||||
|
||||
elif isinstance(group, TagGroupOr):
|
||||
return any(_match_group(result, child) for child in group.filters)
|
||||
|
||||
elif isinstance(group, TagGroupNot):
|
||||
return not _match_group(result, group.filter)
|
||||
|
||||
else:
|
||||
return True
|
||||
|
||||
|
||||
def filter_results_by_tag_groups(
|
||||
results: list,
|
||||
tag_groups: list[TagGroup] | None,
|
||||
) -> list:
|
||||
"""
|
||||
Filter retrieval results by compound tag groups in Python (for post-processing).
|
||||
|
||||
Used when SQL filtering isn't possible (e.g., graph traversal results).
|
||||
Top-level groups are AND-ed together.
|
||||
|
||||
Args:
|
||||
results: List of RetrievalResult objects with a 'tags' attribute.
|
||||
tag_groups: List of TagGroup objects. If None or empty, returns all results.
|
||||
|
||||
Returns:
|
||||
Filtered list of results where ALL top-level groups match.
|
||||
"""
|
||||
if not tag_groups:
|
||||
return results
|
||||
|
||||
return [r for r in results if all(_match_group(r, group) for group in tag_groups)]
|
||||
@@ -1,105 +0,0 @@
|
||||
"""Google Cloud Storage backend using obstore."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import obstore as obs
|
||||
from obstore.store import GCSStore
|
||||
|
||||
from .base import FileStorage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _make_google_auth_credential_provider():
|
||||
"""Create a credential provider using google.auth (supports all credential types).
|
||||
|
||||
obstore's built-in credential parsing only supports service_account and
|
||||
authorized_user JSON types. This provider uses the google-auth library
|
||||
which additionally handles external_account (Workload Identity Federation),
|
||||
impersonated credentials, and metadata-server credentials.
|
||||
"""
|
||||
import google.auth
|
||||
import google.auth.transport.requests
|
||||
|
||||
credentials, _ = google.auth.default(scopes=["https://www.googleapis.com/auth/cloud-platform"])
|
||||
request = google.auth.transport.requests.Request()
|
||||
|
||||
def _provide():
|
||||
credentials.refresh(request)
|
||||
expiry = credentials.expiry
|
||||
if expiry and expiry.tzinfo is None:
|
||||
expiry = expiry.replace(tzinfo=timezone.utc)
|
||||
return {"token": credentials.token, "expires_at": expiry}
|
||||
|
||||
return _provide
|
||||
|
||||
|
||||
class GCSFileStorage(FileStorage):
|
||||
"""
|
||||
Google Cloud Storage backend.
|
||||
|
||||
Uses obstore (Rust-backed) for high-throughput async access to GCS.
|
||||
Supports Application Default Credentials, service account keys, and explicit credentials.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
bucket: str,
|
||||
service_account_key: str | None = None,
|
||||
):
|
||||
kwargs: dict = {}
|
||||
if service_account_key:
|
||||
kwargs["service_account_key"] = service_account_key
|
||||
else:
|
||||
# Use google.auth credential provider for broad credential type support
|
||||
# (service_account, authorized_user, external_account, metadata server, etc.)
|
||||
try:
|
||||
kwargs["credential_provider"] = _make_google_auth_credential_provider()
|
||||
logger.info("Using google.auth credential provider for GCS")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to create google.auth credential provider, falling back to obstore defaults: {e}"
|
||||
)
|
||||
|
||||
# Workaround for https://github.com/developmentseed/obstore/issues/605
|
||||
# obstore's Rust layer doesn't support external_account credentials (Workload
|
||||
# Identity Federation) and eagerly parses GOOGLE_APPLICATION_CREDENTIALS even
|
||||
# when credential_provider is given. Per the obstore maintainer's guidance,
|
||||
# remove env vars so the Rust code doesn't try to authenticate itself.
|
||||
# google.auth (used by credential_provider above) has already loaded credentials.
|
||||
gac = os.environ.pop("GOOGLE_APPLICATION_CREDENTIALS", None)
|
||||
try:
|
||||
self._store = GCSStore(bucket, **kwargs)
|
||||
finally:
|
||||
if gac is not None:
|
||||
os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = gac
|
||||
logger.info(f"Initialized GCS file storage: bucket={bucket}")
|
||||
|
||||
async def store(self, file_data: bytes, key: str, metadata: dict[str, str] | None = None) -> str:
|
||||
await obs.put_async(self._store, key, file_data)
|
||||
logger.debug(f"Stored file {key} ({len(file_data)} bytes) in GCS")
|
||||
return key
|
||||
|
||||
async def retrieve(self, key: str) -> bytes:
|
||||
try:
|
||||
response = await obs.get_async(self._store, key)
|
||||
return await response.bytes_async()
|
||||
except Exception as e:
|
||||
if "not found" in str(e).lower():
|
||||
raise FileNotFoundError(f"File not found: {key}") from e
|
||||
raise
|
||||
|
||||
async def delete(self, key: str) -> None:
|
||||
await obs.delete_async(self._store, key)
|
||||
|
||||
async def exists(self, key: str) -> bool:
|
||||
try:
|
||||
await obs.head_async(self._store, key)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
async def get_download_url(self, key: str, expires_in: int = 3600) -> str:
|
||||
return await obs.sign_async(self._store, "GET", key, timedelta(seconds=expires_in))
|
||||
@@ -1,198 +0,0 @@
|
||||
[build-system]
|
||||
requires = ["hatchling"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "hindsight-api-slim"
|
||||
version = "0.4.18"
|
||||
description = "Hindsight: Agent Memory That Works Like Human Memory"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
"asyncpg>=0.29.0",
|
||||
"python-dotenv>=1.0.0",
|
||||
"openai>=1.0.0",
|
||||
"pydantic>=2.0.0",
|
||||
"rich>=13.0.0",
|
||||
"langchain-text-splitters>=0.3.0",
|
||||
"fastapi[standard]>=0.120.3",
|
||||
"uvicorn>=0.38.0",
|
||||
"wsproto>=1.0.0",
|
||||
"sqlalchemy>=2.0.44",
|
||||
"alembic>=1.17.1",
|
||||
"pgvector>=0.4.1",
|
||||
"greenlet>=3.2.4",
|
||||
"psycopg2-binary>=2.9.11",
|
||||
"tiktoken>=0.12.0",
|
||||
"httpx>=0.27.0",
|
||||
"PyJWT[crypto]>=2.8.0",
|
||||
"fastmcp>=2.14.0", # CVE-2025-66416
|
||||
"python-dateutil>=2.8.0",
|
||||
"opentelemetry-api>=1.20.0",
|
||||
"opentelemetry-sdk>=1.20.0",
|
||||
"opentelemetry-instrumentation-fastapi>=0.41b0",
|
||||
"opentelemetry-exporter-prometheus>=0.41b0",
|
||||
"opentelemetry-exporter-otlp-proto-http>=1.20.0",
|
||||
"opentelemetry-semantic-conventions>=0.41b0",
|
||||
"dateparser>=1.2.2",
|
||||
"google-genai>=1.0.0",
|
||||
"google-auth>=2.0.0",
|
||||
"anthropic>=0.40.0",
|
||||
"typer>=0.9.0",
|
||||
"cohere>=5.0.0",
|
||||
"litellm>=1.0.0",
|
||||
"markitdown[pdf,docx,pptx,xlsx,xls]>=0.1.4", # File to markdown conversion
|
||||
"obstore>=0.4.0", # S3/GCS/Azure object storage client (Rust-backed)
|
||||
"uvloop>=0.22.1",
|
||||
# Transitive dependency security fixes
|
||||
"pyasn1>=0.6.2", # DoS vulnerability fix
|
||||
"urllib3>=2.6.3", # Decompression-bomb safeguards bypass fix
|
||||
"langchain-core>=1.2.11", # Serialization injection + SSRF vulnerability fix
|
||||
"langsmith>=0.6.3", # SSRF via tracing header injection fix
|
||||
"protobuf>=6.33.5", # JSON recursion depth bypass fix
|
||||
"pillow>=12.1.1", # Out-of-bounds write in PSD image loading fix
|
||||
"cryptography>=46.0.5", # Subgroup attack vulnerability fix
|
||||
"filelock>=3.20.1", # TOCTOU race condition fix
|
||||
"authlib>=1.6.6", # Account takeover vulnerability fix
|
||||
"aiohttp>=3.13.3", # Multiple DoS vulnerabilities
|
||||
"claude-agent-sdk>=0.1.27",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
local-ml = [
|
||||
# Local ML models for embeddings/reranking
|
||||
"sentence-transformers>=3.3.0",
|
||||
"transformers>=4.53.0", # Security fixes for ReDoS vulnerabilities
|
||||
"torch>=2.6.0", # CVE fix for remote code execution
|
||||
"einops>=0.8.2",
|
||||
"flashrank>=0.2.0",
|
||||
# Apple Silicon local inference
|
||||
"mlx>=0.31.0",
|
||||
"mlx-lm>=0.31.1",
|
||||
"safetensors>=0.6.2",
|
||||
]
|
||||
embedded-db = [
|
||||
"pg0-embedded>=0.11.0",
|
||||
]
|
||||
all = [
|
||||
"hindsight-api-slim[local-ml,embedded-db]",
|
||||
]
|
||||
test = [
|
||||
"pytest>=7.0.0",
|
||||
"pytest-asyncio>=0.21.0",
|
||||
"pytest-timeout>=2.4.0",
|
||||
"pytest-xdist>=3.0.0",
|
||||
"filelock>=3.20.1", # TOCTOU race condition fix
|
||||
"testcontainers>=4.0.0",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
hindsight-api = "hindsight_api.main:main"
|
||||
hindsight-worker = "hindsight_api.worker.main:main"
|
||||
hindsight-local-mcp = "hindsight_api.mcp_local:main"
|
||||
hindsight-admin = "hindsight_api.admin.cli:main"
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
packages = ["hindsight_api"]
|
||||
|
||||
[tool.hatch.build.targets.wheel.sources]
|
||||
"hindsight_api" = "hindsight_api"
|
||||
|
||||
[tool.hatch.build.targets.sdist]
|
||||
include = [
|
||||
"hindsight_api/**/*",
|
||||
]
|
||||
|
||||
[tool.hatch.build]
|
||||
include = [
|
||||
"hindsight_api/**/*.py",
|
||||
"hindsight_api/alembic/**/*",
|
||||
]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
log_cli = true
|
||||
log_cli_level = "INFO"
|
||||
log_cli_format = "%(asctime)s - %(levelname)s - %(name)s - %(message)s"
|
||||
log_cli_date_format = "%Y-%m-%d %H:%M:%S"
|
||||
addopts = "--timeout 300 -n 8 --dist loadgroup --durations=10 -v"
|
||||
asyncio_mode = "auto"
|
||||
asyncio_default_fixture_loop_scope = "function"
|
||||
log_auto_indent = true
|
||||
filterwarnings = [
|
||||
"ignore:The @wait_container_is_ready decorator is deprecated:DeprecationWarning",
|
||||
"ignore::RuntimeWarning:asyncio",
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
dev = [
|
||||
"pytest>=9.0.0",
|
||||
"pytest-asyncio>=1.3.0",
|
||||
"pytest-timeout>=2.4.0",
|
||||
"pytest-xdist>=3.8.0",
|
||||
"python-dotenv>=1.2.1",
|
||||
"filelock>=3.20.1", # TOCTOU race condition fix
|
||||
"ruff>=0.8.0",
|
||||
"ty>=0.0.1",
|
||||
"testcontainers>=4.0.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.lint.isort]
|
||||
known-third-party = ["alembic"]
|
||||
|
||||
[tool.ruff.format]
|
||||
quote-style = "double"
|
||||
indent-style = "space"
|
||||
|
||||
[tool.uv]
|
||||
# Allow uv to search all configured indexes for packages, not just the first one
|
||||
# This prevents dependency resolution failures when using pytorch index + PyPI
|
||||
index-strategy = "unsafe-best-match"
|
||||
|
||||
[tool.ty]
|
||||
# Type checking configuration
|
||||
# ty is an extremely fast Python type checker from Astral (same team as ruff/uv)
|
||||
|
||||
[tool.ty.environment]
|
||||
python-version = "3.11"
|
||||
|
||||
[tool.ty.src]
|
||||
exclude = [
|
||||
"tests/",
|
||||
"hindsight_api/alembic/",
|
||||
]
|
||||
|
||||
[tool.ty.rules]
|
||||
# Disable noisy rules while keeping important ones
|
||||
invalid-argument-type = "ignore" # False positives with **kwargs patterns
|
||||
invalid-return-type = "ignore" # Often intentional in async code
|
||||
invalid-parameter-default = "ignore" # Optional params with None default
|
||||
possibly-missing-attribute = "ignore" # Common with Optional types
|
||||
invalid-raise = "ignore" # False positives with exception tracking
|
||||
call-non-callable = "ignore" # False positives with Optional types
|
||||
invalid-key = "ignore" # Pydantic ConfigDict not understood
|
||||
invalid-method-override = "ignore" # Intentional signature differences
|
||||
unresolved-reference = "ignore" # Forward references not always resolved
|
||||
@@ -1,193 +0,0 @@
|
||||
"""
|
||||
Tests for per-bank HNSW index lifecycle and UNION ALL retrieval.
|
||||
|
||||
Covers:
|
||||
- _hnsw_index_name deterministic naming
|
||||
- Per-bank HNSW indexes created on bank creation (retain_async / ensure_bank_exists)
|
||||
- Per-bank HNSW indexes dropped on bank deletion
|
||||
- retrieve_semantic_bm25_combined groups results correctly by fact_type and source
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from hindsight_api.engine.retain.bank_utils import _HNSW_FACT_TYPES, _hnsw_index_name
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unit tests — no DB required
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestHnswIndexName:
|
||||
def test_deterministic(self):
|
||||
uid = "550e8400-e29b-41d4-a716-446655440000"
|
||||
assert _hnsw_index_name("world", uid) == _hnsw_index_name("world", uid)
|
||||
|
||||
def test_strips_dashes(self):
|
||||
uid = "550e8400-e29b-41d4-a716-446655440000"
|
||||
name = _hnsw_index_name("world", uid)
|
||||
# uid16 should be hex chars only
|
||||
assert "-" not in name
|
||||
|
||||
def test_uses_first_16_hex_chars(self):
|
||||
uid = "550e8400-e29b-41d4-a716-446655440000"
|
||||
uid16 = uid.replace("-", "")[:16] # "550e8400e29b41d4"
|
||||
assert name_ends_with(name=_hnsw_index_name("world", uid), suffix=uid16)
|
||||
|
||||
def test_suffix_per_fact_type(self):
|
||||
uid = "550e8400-e29b-41d4-a716-446655440000"
|
||||
names = {ft: _hnsw_index_name(ft, uid) for ft in _HNSW_FACT_TYPES}
|
||||
# All three names must be distinct
|
||||
assert len(set(names.values())) == 3
|
||||
|
||||
def test_all_fact_types_covered(self):
|
||||
assert set(_HNSW_FACT_TYPES) == {"world", "experience", "observation"}
|
||||
|
||||
def test_fits_pg_identifier_limit(self):
|
||||
# PostgreSQL max identifier length is 63 chars
|
||||
uid = "f" * 32 # simulated UUID without dashes
|
||||
for ft in _HNSW_FACT_TYPES:
|
||||
assert len(_hnsw_index_name(ft, uid)) <= 63
|
||||
|
||||
|
||||
def name_ends_with(name: str, suffix: str) -> bool:
|
||||
return name.endswith(suffix)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Integration tests — require DB (memory fixture)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _get_bank_hnsw_indexes(pool, bank_id: str) -> list[str]:
|
||||
"""Return index names for memory_units that match the per-bank pattern."""
|
||||
async with pool.acquire() as conn:
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT indexname
|
||||
FROM pg_indexes
|
||||
WHERE tablename = 'memory_units'
|
||||
AND indexname LIKE 'idx_mu_emb_%'
|
||||
AND indexdef LIKE $1
|
||||
ORDER BY indexname
|
||||
""",
|
||||
f"%bank_id = '{bank_id}'%",
|
||||
)
|
||||
return [row["indexname"] for row in rows]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_creates_per_bank_hnsw_indexes(memory, request_context):
|
||||
"""retain_async on a new bank must create 3 per-(bank, fact_type) HNSW indexes."""
|
||||
bank_id = f"test_hnsw_create_{uuid.uuid4().hex[:8]}"
|
||||
try:
|
||||
await memory.retain_async(
|
||||
bank_id=bank_id,
|
||||
content="Alice is a software engineer.",
|
||||
request_context=request_context,
|
||||
)
|
||||
indexes = await _get_bank_hnsw_indexes(memory._pool, bank_id)
|
||||
assert len(indexes) == 3, f"Expected 3 per-bank HNSW indexes, got: {indexes}"
|
||||
for ft_short in _HNSW_FACT_TYPES.values():
|
||||
assert any(ft_short in idx for idx in indexes), (
|
||||
f"Missing index for fact_type short '{ft_short}' in {indexes}"
|
||||
)
|
||||
finally:
|
||||
await memory.delete_bank(bank_id, request_context=request_context)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_bank_drops_hnsw_indexes(memory, request_context):
|
||||
"""delete_bank must drop all per-bank HNSW indexes."""
|
||||
bank_id = f"test_hnsw_drop_{uuid.uuid4().hex[:8]}"
|
||||
|
||||
await memory.retain_async(
|
||||
bank_id=bank_id,
|
||||
content="Bob is a data scientist.",
|
||||
request_context=request_context,
|
||||
)
|
||||
# Verify indexes exist before deletion
|
||||
indexes_before = await _get_bank_hnsw_indexes(memory._pool, bank_id)
|
||||
assert len(indexes_before) == 3
|
||||
|
||||
await memory.delete_bank(bank_id, request_context=request_context)
|
||||
|
||||
indexes_after = await _get_bank_hnsw_indexes(memory._pool, bank_id)
|
||||
assert indexes_after == [], f"Indexes should be dropped after bank deletion, got: {indexes_after}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_idempotent_bank_creation(memory, request_context):
|
||||
"""Retaining into the same bank twice must not error and still have exactly 3 indexes."""
|
||||
bank_id = f"test_hnsw_idem_{uuid.uuid4().hex[:8]}"
|
||||
try:
|
||||
await memory.retain_async(
|
||||
bank_id=bank_id,
|
||||
content="Carol is a product manager.",
|
||||
request_context=request_context,
|
||||
)
|
||||
await memory.retain_async(
|
||||
bank_id=bank_id,
|
||||
content="Carol joined the company in 2022.",
|
||||
request_context=request_context,
|
||||
)
|
||||
indexes = await _get_bank_hnsw_indexes(memory._pool, bank_id)
|
||||
assert len(indexes) == 3
|
||||
finally:
|
||||
await memory.delete_bank(bank_id, request_context=request_context)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve_semantic_bm25_grouped_by_fact_type(memory, request_context):
|
||||
"""
|
||||
retrieve_semantic_bm25_combined must return a dict keyed by fact_type with
|
||||
(semantic_list, bm25_list) tuples. All returned facts must belong to their
|
||||
declared fact_type.
|
||||
"""
|
||||
from hindsight_api.engine.search.retrieval import retrieve_semantic_bm25_combined
|
||||
|
||||
bank_id = f"test_retrieval_{uuid.uuid4().hex[:8]}"
|
||||
try:
|
||||
await memory.retain_async(
|
||||
bank_id=bank_id,
|
||||
content=(
|
||||
"Alice is a software engineer at TechCorp. "
|
||||
"She visited Paris in 2023 for a conference."
|
||||
),
|
||||
context="background",
|
||||
event_date=datetime(2023, 6, 1, tzinfo=timezone.utc),
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
query_emb = memory.embeddings.encode(["software engineer Alice"])
|
||||
query_emb_str = str(query_emb[0])
|
||||
|
||||
fact_types = ["world", "experience"]
|
||||
async with memory._pool.acquire() as conn:
|
||||
results = await retrieve_semantic_bm25_combined(
|
||||
conn=conn,
|
||||
query_emb_str=query_emb_str,
|
||||
query_text="software engineer Alice",
|
||||
bank_id=bank_id,
|
||||
fact_types=fact_types,
|
||||
limit=5,
|
||||
)
|
||||
|
||||
# Must return an entry for every requested fact_type
|
||||
assert set(results.keys()) == set(fact_types)
|
||||
|
||||
for ft, (sem, bm25) in results.items():
|
||||
# Semantic and BM25 lists must be lists
|
||||
assert isinstance(sem, list)
|
||||
assert isinstance(bm25, list)
|
||||
# All semantic results must declare the correct fact_type
|
||||
for r in sem:
|
||||
assert r.fact_type == ft, f"Semantic result has wrong fact_type: {r.fact_type}"
|
||||
# All BM25 results must declare the correct fact_type
|
||||
for r in bm25:
|
||||
assert r.fact_type == ft, f"BM25 result has wrong fact_type: {r.fact_type}"
|
||||
|
||||
finally:
|
||||
await memory.delete_bank(bank_id, request_context=request_context)
|
||||
@@ -1,37 +0,0 @@
|
||||
import pytest
|
||||
|
||||
from hindsight_api.engine.llm_wrapper import sanitize_llm_output
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"input_text, expected",
|
||||
[
|
||||
# Null bytes stripped
|
||||
("hello\x00world", "helloworld"),
|
||||
("FIRST\u0000PAGE", "FIRSTPAGE"),
|
||||
# Multiple null bytes
|
||||
("\x00\x00text\x00", "text"),
|
||||
# Other control characters stripped (non-whitespace)
|
||||
("text\x01\x02\x03end", "textend"),
|
||||
("text\x08end", "textend"), # backspace
|
||||
("text\x0cend", "textend"), # form feed
|
||||
("text\x0bend", "textend"), # vertical tab
|
||||
("text\x1fend", "textend"), # unit separator
|
||||
("text\x7fend", "textend"), # DEL
|
||||
# Whitespace preserved
|
||||
("hello\tworld", "hello\tworld"),
|
||||
("hello\nworld", "hello\nworld"),
|
||||
("hello\r\nworld", "hello\r\nworld"),
|
||||
# Unicode surrogates stripped
|
||||
("text\ud800end", "textend"),
|
||||
("text\udfffend", "textend"),
|
||||
# Clean text unchanged
|
||||
("normal text", "normal text"),
|
||||
("unicode: café naïve", "unicode: café naïve"),
|
||||
# Edge cases
|
||||
("", ""),
|
||||
(None, None),
|
||||
],
|
||||
)
|
||||
def test_sanitize_llm_output(input_text, expected):
|
||||
assert sanitize_llm_output(input_text) == expected
|
||||
@@ -1,321 +0,0 @@
|
||||
"""
|
||||
Reproduce issue #520: Reflect fails with LM Studio due to unsupported tool_choice format.
|
||||
|
||||
The reflect agent forces tool selection via named tool_choice dicts on the first few iterations:
|
||||
{"type": "function", "function": {"name": "search_mental_models"}}
|
||||
|
||||
LM Studio (and Ollama) reject this format with HTTP 400:
|
||||
"Tool choice of type 'function' is not supported. Use 'auto', 'none', or 'required'."
|
||||
|
||||
The fix should convert named tool_choice to "required" and filter the tools list
|
||||
to only the requested tool for providers that don't support named tool_choice.
|
||||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from openai import APIStatusError
|
||||
|
||||
from hindsight_api.engine.providers.openai_compatible_llm import OpenAICompatibleLLM
|
||||
|
||||
# Reflect agent tools (subset matching what agent.py uses)
|
||||
REFLECT_TOOLS = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "search_mental_models",
|
||||
"description": "Search consolidated mental models",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"query": {"type": "string"}},
|
||||
"required": ["query"],
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "search_observations",
|
||||
"description": "Search raw observations",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"query": {"type": "string"}},
|
||||
"required": ["query"],
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "recall",
|
||||
"description": "Recall semantic memories",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"query": {"type": "string"}},
|
||||
"required": ["query"],
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "done",
|
||||
"description": "Finish and return the answer",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"answer": {"type": "string"}},
|
||||
"required": ["answer"],
|
||||
},
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _make_lmstudio_llm() -> OpenAICompatibleLLM:
|
||||
return OpenAICompatibleLLM(
|
||||
provider="lmstudio",
|
||||
api_key="local",
|
||||
base_url="http://localhost:1234/v1",
|
||||
model="openai/gpt-oss-20b",
|
||||
)
|
||||
|
||||
|
||||
def _lmstudio_400_error(msg: str = "Tool choice of type 'function' is not supported. Use 'auto', 'none', or 'required'.") -> APIStatusError:
|
||||
"""Simulate the HTTP 400 LM Studio returns for unsupported tool_choice format."""
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 400
|
||||
mock_response.headers = {}
|
||||
return APIStatusError(
|
||||
message=msg,
|
||||
response=mock_response,
|
||||
body={"error": {"message": msg, "type": "invalid_request_error"}},
|
||||
)
|
||||
|
||||
|
||||
def _make_tool_call_response(tool_name: str, arguments: dict) -> MagicMock:
|
||||
"""Build a mock successful tool call response from the LLM API."""
|
||||
mock_tc = MagicMock()
|
||||
mock_tc.id = "call_abc123"
|
||||
mock_tc.function.name = tool_name
|
||||
mock_tc.function.arguments = json.dumps(arguments)
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.usage.prompt_tokens = 120
|
||||
mock_response.usage.completion_tokens = 40
|
||||
mock_response.usage.total_tokens = 160
|
||||
mock_response.choices[0].finish_reason = "tool_calls"
|
||||
mock_response.choices[0].message.content = None
|
||||
mock_response.choices[0].message.tool_calls = [mock_tc]
|
||||
return mock_response
|
||||
|
||||
|
||||
class TestLMStudioNamedToolChoiceBug:
|
||||
"""
|
||||
Reproduces issue #520.
|
||||
|
||||
The reflect agent (agent.py lines 546-555) sets tool_choice to a named dict
|
||||
on the first iterations to force sequential retrieval:
|
||||
|
||||
iteration=0, has_mental_models=True → {"type": "function", "function": {"name": "search_mental_models"}}
|
||||
iteration=0, has_mental_models=False → {"type": "function", "function": {"name": "search_observations"}}
|
||||
iteration=1, has_mental_models=True → {"type": "function", "function": {"name": "search_observations"}}
|
||||
iteration=1 or (2 with models) → {"type": "function", "function": {"name": "recall"}}
|
||||
|
||||
LM Studio rejects these dict formats with HTTP 400.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lmstudio_named_tool_choice_no_longer_causes_400(self):
|
||||
"""
|
||||
Regression test for issue #520: named tool_choice dict is converted to
|
||||
"required" + filtered tools before the API call, so LM Studio never
|
||||
sees the unsupported format and the 400 error no longer occurs.
|
||||
"""
|
||||
llm = _make_lmstudio_llm()
|
||||
named_tool_choice = {"type": "function", "function": {"name": "search_mental_models"}}
|
||||
success_response = _make_tool_call_response("search_mental_models", {"query": "user name"})
|
||||
|
||||
with patch.object(llm._client.chat.completions, "create", new_callable=AsyncMock) as mock_create:
|
||||
mock_create.return_value = success_response
|
||||
|
||||
# Should succeed — no 400 because the dict is converted before sending
|
||||
result = await llm.call_with_tools(
|
||||
messages=[{"role": "user", "content": "What is the user's name?"}],
|
||||
tools=REFLECT_TOOLS,
|
||||
tool_choice=named_tool_choice,
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
assert len(result.tool_calls) == 1
|
||||
assert result.tool_calls[0].name == "search_mental_models"
|
||||
|
||||
sent_kwargs = mock_create.call_args.kwargs
|
||||
assert sent_kwargs["tool_choice"] == "required"
|
||||
assert len(sent_kwargs["tools"]) == 1
|
||||
assert sent_kwargs["tools"][0]["function"]["name"] == "search_mental_models"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"forced_tool_name",
|
||||
["search_mental_models", "search_observations", "recall"],
|
||||
)
|
||||
async def test_all_reflect_forced_tools_fail_on_lmstudio(self, forced_tool_name: str):
|
||||
"""
|
||||
Each named tool_choice the reflect agent uses on iterations 0-2 triggers
|
||||
the same 400 error on LM Studio.
|
||||
"""
|
||||
llm = _make_lmstudio_llm()
|
||||
named_tool_choice = {"type": "function", "function": {"name": forced_tool_name}}
|
||||
|
||||
with patch.object(llm._client.chat.completions, "create", new_callable=AsyncMock) as mock_create:
|
||||
mock_create.side_effect = _lmstudio_400_error()
|
||||
|
||||
with pytest.raises(APIStatusError) as exc_info:
|
||||
await llm.call_with_tools(
|
||||
messages=[{"role": "user", "content": "Test query"}],
|
||||
tools=REFLECT_TOOLS,
|
||||
tool_choice=named_tool_choice,
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lmstudio_string_tool_choice_works_fine(self):
|
||||
"""
|
||||
String tool_choice values ("auto", "none", "required") ARE supported by LM Studio.
|
||||
Only the dict format {"type": "function", "function": {"name": "..."}} fails.
|
||||
This test confirms the control case works.
|
||||
"""
|
||||
llm = _make_lmstudio_llm()
|
||||
success_response = _make_tool_call_response("search_mental_models", {"query": "user name"})
|
||||
|
||||
with patch.object(llm._client.chat.completions, "create", new_callable=AsyncMock) as mock_create:
|
||||
mock_create.return_value = success_response
|
||||
|
||||
result = await llm.call_with_tools(
|
||||
messages=[{"role": "user", "content": "What is the user's name?"}],
|
||||
tools=REFLECT_TOOLS,
|
||||
tool_choice="required", # string form — LM Studio accepts this
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
assert len(result.tool_calls) == 1
|
||||
assert result.tool_calls[0].name == "search_mental_models"
|
||||
|
||||
# Confirm "required" was sent, not a dict
|
||||
sent_kwargs = mock_create.call_args.kwargs
|
||||
assert sent_kwargs["tool_choice"] == "required"
|
||||
|
||||
|
||||
class TestExpectedFixBehavior:
|
||||
"""
|
||||
Tests that document the EXPECTED behavior after the fix is applied.
|
||||
|
||||
For lmstudio (and ollama) providers, when tool_choice is a named dict:
|
||||
{"type": "function", "function": {"name": "search_mental_models"}}
|
||||
|
||||
The fix should:
|
||||
1. Convert tool_choice to "required"
|
||||
2. Filter tools to only the requested tool
|
||||
|
||||
These tests currently FAIL (because the fix is not yet implemented).
|
||||
After the fix is applied, they should PASS.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fix_converts_named_tool_choice_to_required(self):
|
||||
"""
|
||||
After fix: named tool_choice dict is converted to "required" for lmstudio.
|
||||
The API receives tool_choice="required" instead of the unsupported dict.
|
||||
"""
|
||||
llm = _make_lmstudio_llm()
|
||||
named_tool_choice = {"type": "function", "function": {"name": "search_mental_models"}}
|
||||
success_response = _make_tool_call_response("search_mental_models", {"query": "user name"})
|
||||
|
||||
with patch.object(llm._client.chat.completions, "create", new_callable=AsyncMock) as mock_create:
|
||||
mock_create.return_value = success_response
|
||||
|
||||
result = await llm.call_with_tools(
|
||||
messages=[{"role": "user", "content": "What is the user's name?"}],
|
||||
tools=REFLECT_TOOLS,
|
||||
tool_choice=named_tool_choice,
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
assert len(result.tool_calls) == 1
|
||||
assert result.tool_calls[0].name == "search_mental_models"
|
||||
|
||||
sent_kwargs = mock_create.call_args.kwargs
|
||||
# Fix: dict was converted to "required"
|
||||
assert sent_kwargs["tool_choice"] == "required", (
|
||||
f"Expected tool_choice='required', got {sent_kwargs['tool_choice']!r}"
|
||||
)
|
||||
# Fix: tools filtered to just the requested one
|
||||
assert len(sent_kwargs["tools"]) == 1
|
||||
assert sent_kwargs["tools"][0]["function"]["name"] == "search_mental_models"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"forced_tool_name",
|
||||
["search_mental_models", "search_observations", "recall"],
|
||||
)
|
||||
async def test_fix_filters_tools_to_requested_tool(self, forced_tool_name: str):
|
||||
"""
|
||||
After fix: tools list is filtered to only the forced tool so the model
|
||||
can only call that one tool (equivalent to the named tool_choice behavior).
|
||||
"""
|
||||
llm = _make_lmstudio_llm()
|
||||
named_tool_choice = {"type": "function", "function": {"name": forced_tool_name}}
|
||||
success_response = _make_tool_call_response(forced_tool_name, {"query": "test"})
|
||||
|
||||
with patch.object(llm._client.chat.completions, "create", new_callable=AsyncMock) as mock_create:
|
||||
mock_create.return_value = success_response
|
||||
|
||||
await llm.call_with_tools(
|
||||
messages=[{"role": "user", "content": "Test query"}],
|
||||
tools=REFLECT_TOOLS,
|
||||
tool_choice=named_tool_choice,
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
sent_kwargs = mock_create.call_args.kwargs
|
||||
assert sent_kwargs["tool_choice"] == "required"
|
||||
assert len(sent_kwargs["tools"]) == 1
|
||||
assert sent_kwargs["tools"][0]["function"]["name"] == forced_tool_name
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fix_also_applies_to_openai_provider(self):
|
||||
"""
|
||||
The fix is generalized: all providers convert named tool_choice to
|
||||
"required" + filtered tools. OpenAI natively supports the dict format
|
||||
too, so the behaviour is semantically identical either way.
|
||||
"""
|
||||
from hindsight_api.engine.providers.openai_compatible_llm import OpenAICompatibleLLM
|
||||
|
||||
openai_llm = OpenAICompatibleLLM(
|
||||
provider="openai",
|
||||
api_key="sk-test",
|
||||
base_url="",
|
||||
model="gpt-4o-mini",
|
||||
)
|
||||
|
||||
named_tool_choice = {"type": "function", "function": {"name": "search_mental_models"}}
|
||||
success_response = _make_tool_call_response("search_mental_models", {"query": "test"})
|
||||
|
||||
with patch.object(openai_llm._client.chat.completions, "create", new_callable=AsyncMock) as mock_create:
|
||||
mock_create.return_value = success_response
|
||||
|
||||
await openai_llm.call_with_tools(
|
||||
messages=[{"role": "user", "content": "Test"}],
|
||||
tools=REFLECT_TOOLS,
|
||||
tool_choice=named_tool_choice,
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
sent_kwargs = mock_create.call_args.kwargs
|
||||
# Generalized fix applies to OpenAI too
|
||||
assert sent_kwargs["tool_choice"] == "required"
|
||||
assert len(sent_kwargs["tools"]) == 1
|
||||
assert sent_kwargs["tools"][0]["function"]["name"] == "search_mental_models"
|
||||
@@ -1,154 +0,0 @@
|
||||
"""Tests for migration g7h8i9j0k1l2 (backsweep orphaned memory_units).
|
||||
|
||||
Uses a dedicated pg0 instance (port 5562) so the test can control exactly
|
||||
which migrations have run before inserting the orphan seed data.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
from sqlalchemy import create_engine, text
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_SCRIPT_LOCATION = str(Path(__file__).parent.parent / "hindsight_api" / "alembic")
|
||||
|
||||
|
||||
def _alembic_cfg(db_url: str) -> Config:
|
||||
cfg = Config()
|
||||
cfg.set_main_option("script_location", _SCRIPT_LOCATION)
|
||||
cfg.set_main_option("sqlalchemy.url", db_url)
|
||||
cfg.set_main_option("prepend_sys_path", ".")
|
||||
cfg.set_main_option("path_separator", "os")
|
||||
return cfg
|
||||
|
||||
|
||||
def _upgrade(db_url: str, revision: str) -> None:
|
||||
command.upgrade(_alembic_cfg(db_url), revision)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixture: fresh database at the revision just before the backsweep
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def pre_backsweep_db_url():
|
||||
"""
|
||||
Spin up a dedicated pg0 instance and run all migrations up to (but not
|
||||
including) the backsweep revision so each test can seed orphan data and
|
||||
then apply the backsweep itself.
|
||||
"""
|
||||
from hindsight_api.pg0 import EmbeddedPostgres
|
||||
|
||||
pg0 = EmbeddedPostgres(name="hindsight-backsweep-test", port=5562)
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
url = loop.run_until_complete(pg0.ensure_running())
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
# Migrate up to the revision just before the backsweep.
|
||||
_upgrade(url, "f6g7h8i9j0k1")
|
||||
return url
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# The test
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_backsweep_removes_orphans_and_preserves_legit_rows(pre_backsweep_db_url):
|
||||
"""
|
||||
Seed four kinds of rows then apply the backsweep migration and verify:
|
||||
|
||||
Rows that MUST be deleted
|
||||
─────────────────────────
|
||||
A. Any fact_type, bank_id missing from banks
|
||||
→ Pass 1 deletes these regardless of fact_type or source links.
|
||||
|
||||
B. observation, bank exists, but ALL source_memory_ids are gone
|
||||
→ Pass 2 deletes these.
|
||||
|
||||
Rows that MUST survive
|
||||
──────────────────────
|
||||
C. observation, bank exists, at least ONE source_memory_id still live
|
||||
→ Pass 2 must not touch these.
|
||||
|
||||
D. Non-observation (world), bank exists, no sources (not relevant)
|
||||
→ Pass 1 must not touch these (bank exists).
|
||||
"""
|
||||
db_url = pre_backsweep_db_url
|
||||
engine = create_engine(db_url)
|
||||
|
||||
alive_bank = f"bank_{uuid.uuid4().hex[:8]}"
|
||||
ghost_bank = f"bank_{uuid.uuid4().hex[:8]}" # never inserted into banks
|
||||
|
||||
# UUIDs for memory units
|
||||
id_pass1_world = uuid.uuid4() # A: world unit, ghost bank
|
||||
id_pass1_obs = uuid.uuid4() # A: observation, ghost bank
|
||||
id_pass2_obs = uuid.uuid4() # B: observation, all sources gone
|
||||
id_keep_obs = uuid.uuid4() # C: observation with one live source
|
||||
id_keep_world = uuid.uuid4() # D: world unit, alive bank
|
||||
id_live_source = uuid.uuid4() # live source for C
|
||||
|
||||
with engine.connect() as conn:
|
||||
# --- banks ---
|
||||
conn.execute(text("INSERT INTO banks (bank_id) VALUES (:b)"), {"b": alive_bank})
|
||||
|
||||
# --- seed memory_units ---
|
||||
def insert_mu(uid, bank, fact_type, sources=None):
|
||||
src_arr = "{" + ",".join(str(s) for s in (sources or [])) + "}"
|
||||
conn.execute(
|
||||
text(
|
||||
"""
|
||||
INSERT INTO memory_units
|
||||
(id, bank_id, text, fact_type, source_memory_ids)
|
||||
VALUES
|
||||
(:id, :bank, :text, :ft, CAST(:src AS uuid[]))
|
||||
"""
|
||||
),
|
||||
{"id": uid, "bank": bank, "text": "test", "ft": fact_type, "src": src_arr},
|
||||
)
|
||||
|
||||
# A: ghost-bank rows (Pass 1 targets)
|
||||
insert_mu(id_pass1_world, ghost_bank, "world")
|
||||
insert_mu(id_pass1_obs, ghost_bank, "observation", sources=[uuid.uuid4()])
|
||||
|
||||
# B: observation with all-dead sources (Pass 2 target)
|
||||
insert_mu(id_pass2_obs, alive_bank, "observation", sources=[uuid.uuid4(), uuid.uuid4()])
|
||||
|
||||
# C: observation with one live source (must survive)
|
||||
insert_mu(id_live_source, alive_bank, "world")
|
||||
insert_mu(id_keep_obs, alive_bank, "observation", sources=[id_live_source, uuid.uuid4()])
|
||||
|
||||
# D: world unit in alive bank (must survive)
|
||||
insert_mu(id_keep_world, alive_bank, "world")
|
||||
|
||||
conn.commit()
|
||||
|
||||
# --- apply the backsweep ---
|
||||
_upgrade(db_url, "g7h8i9j0k1l2")
|
||||
|
||||
# --- verify ---
|
||||
with engine.connect() as conn:
|
||||
def exists(uid):
|
||||
return conn.execute(
|
||||
text("SELECT 1 FROM memory_units WHERE id = :id"), {"id": uid}
|
||||
).fetchone() is not None
|
||||
|
||||
# Must be gone
|
||||
assert not exists(id_pass1_world), "Pass 1: world unit with ghost bank should be deleted"
|
||||
assert not exists(id_pass1_obs), "Pass 1: observation with ghost bank should be deleted"
|
||||
assert not exists(id_pass2_obs), "Pass 2: observation with all-dead sources should be deleted"
|
||||
|
||||
# Must survive
|
||||
assert exists(id_keep_obs), "observation with a live source must not be deleted"
|
||||
assert exists(id_keep_world), "world unit in alive bank must not be deleted"
|
||||
assert exists(id_live_source), "live source memory unit must not be deleted"
|
||||
|
||||
engine.dispose()
|
||||
@@ -1,48 +0,0 @@
|
||||
import threading
|
||||
import time
|
||||
|
||||
from hindsight_api import migrations
|
||||
|
||||
|
||||
def test_run_migrations_internal_serializes_alembic_upgrade(monkeypatch):
|
||||
max_concurrent_upgrades = 0
|
||||
active_upgrades = 0
|
||||
active_lock = threading.Lock()
|
||||
start_barrier = threading.Barrier(2)
|
||||
|
||||
def fake_upgrade(_cfg, _revision):
|
||||
nonlocal max_concurrent_upgrades, active_upgrades
|
||||
with active_lock:
|
||||
active_upgrades += 1
|
||||
max_concurrent_upgrades = max(max_concurrent_upgrades, active_upgrades)
|
||||
time.sleep(0.05)
|
||||
with active_lock:
|
||||
active_upgrades -= 1
|
||||
|
||||
monkeypatch.setattr(migrations.command, "upgrade", fake_upgrade)
|
||||
|
||||
errors = []
|
||||
|
||||
def run_in_thread(schema):
|
||||
try:
|
||||
start_barrier.wait()
|
||||
migrations._run_migrations_internal(
|
||||
"postgresql://user:pass@localhost/db",
|
||||
"/tmp/alembic",
|
||||
schema=schema,
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - diagnostic path
|
||||
errors.append(exc)
|
||||
|
||||
threads = [
|
||||
threading.Thread(target=run_in_thread, args=("tenant_alpha",)),
|
||||
threading.Thread(target=run_in_thread, args=("tenant_beta",)),
|
||||
]
|
||||
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
for thread in threads:
|
||||
thread.join()
|
||||
|
||||
assert not errors
|
||||
assert max_concurrent_upgrades == 1
|
||||
@@ -1,311 +0,0 @@
|
||||
"""Tests for operation cancellation when a bank is deleted.
|
||||
|
||||
Covers:
|
||||
- CASCADE DELETE: deleting a bank removes async_operations and webhooks rows
|
||||
- _check_op_alive: returns True when op exists, False when deleted
|
||||
- _mark_operation_completed / _mark_operation_failed: graceful no-op when row is gone
|
||||
- Consolidation checkpoint: stops early after a batch commit if op was deleted
|
||||
- Retain checkpoint: stops between sub-batches if op was deleted
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from hindsight_api.engine.memory_engine import MemoryEngine
|
||||
|
||||
|
||||
pytestmark = pytest.mark.xdist_group("op_cancellation_tests")
|
||||
|
||||
_BANK_PREFIX = "test-op-cancel"
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def pool(pg0_db_url):
|
||||
import asyncpg
|
||||
from hindsight_api.pg0 import resolve_database_url
|
||||
|
||||
resolved_url = await resolve_database_url(pg0_db_url)
|
||||
p = await asyncpg.create_pool(resolved_url, min_size=1, max_size=5, command_timeout=30)
|
||||
yield p
|
||||
await p.close()
|
||||
|
||||
|
||||
@pytest_asyncio.fixture(autouse=True)
|
||||
async def cleanup(pool):
|
||||
"""Remove test rows before and after each test."""
|
||||
await pool.execute(f"DELETE FROM banks WHERE bank_id LIKE '{_BANK_PREFIX}%'")
|
||||
yield
|
||||
await pool.execute(f"DELETE FROM banks WHERE bank_id LIKE '{_BANK_PREFIX}%'")
|
||||
|
||||
|
||||
async def _insert_bank(pool, bank_id: str):
|
||||
await pool.execute(
|
||||
"INSERT INTO banks (bank_id, name) VALUES ($1, $2) ON CONFLICT DO NOTHING",
|
||||
bank_id,
|
||||
bank_id,
|
||||
)
|
||||
|
||||
|
||||
async def _insert_op(pool, bank_id: str, op_id: uuid.UUID | None = None) -> uuid.UUID:
|
||||
op_id = op_id or uuid.uuid4()
|
||||
await pool.execute(
|
||||
"""
|
||||
INSERT INTO async_operations (operation_id, bank_id, operation_type, status)
|
||||
VALUES ($1, $2, 'consolidation', 'processing')
|
||||
""",
|
||||
op_id,
|
||||
bank_id,
|
||||
)
|
||||
return op_id
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# CASCADE DELETE tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCascadeDeleteOnBankDeletion:
|
||||
@pytest.mark.asyncio
|
||||
async def test_bank_deletion_cascades_to_async_operations(self, pool):
|
||||
bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}"
|
||||
await _insert_bank(pool, bank_id)
|
||||
op_id = await _insert_op(pool, bank_id)
|
||||
|
||||
# Verify op exists
|
||||
row = await pool.fetchrow("SELECT operation_id FROM async_operations WHERE operation_id = $1", op_id)
|
||||
assert row is not None
|
||||
|
||||
# Delete the bank — should cascade to async_operations
|
||||
await pool.execute("DELETE FROM banks WHERE bank_id = $1", bank_id)
|
||||
|
||||
row = await pool.fetchrow("SELECT operation_id FROM async_operations WHERE operation_id = $1", op_id)
|
||||
assert row is None, "async_operations row should be deleted by CASCADE"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bank_deletion_cascades_to_webhooks(self, pool):
|
||||
bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}"
|
||||
await _insert_bank(pool, bank_id)
|
||||
webhook_id = uuid.uuid4()
|
||||
await pool.execute(
|
||||
"""
|
||||
INSERT INTO webhooks (id, bank_id, url, event_types)
|
||||
VALUES ($1, $2, 'https://example.com/hook', '{}')
|
||||
""",
|
||||
webhook_id,
|
||||
bank_id,
|
||||
)
|
||||
|
||||
row = await pool.fetchrow("SELECT id FROM webhooks WHERE id = $1", webhook_id)
|
||||
assert row is not None
|
||||
|
||||
await pool.execute("DELETE FROM banks WHERE bank_id = $1", bank_id)
|
||||
|
||||
row = await pool.fetchrow("SELECT id FROM webhooks WHERE id = $1", webhook_id)
|
||||
assert row is None, "webhooks row should be deleted by CASCADE"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _check_op_alive tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCheckOpAlive:
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_true_when_op_exists(self, memory: MemoryEngine, request_context):
|
||||
bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}"
|
||||
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
|
||||
|
||||
op_id = uuid.uuid4()
|
||||
async with memory._pool.acquire() as conn:
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO async_operations (operation_id, bank_id, operation_type, status)
|
||||
VALUES ($1, $2, 'consolidation', 'processing')
|
||||
""",
|
||||
op_id,
|
||||
bank_id,
|
||||
)
|
||||
|
||||
assert await memory._check_op_alive(str(op_id)) is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_false_when_op_deleted(self, memory: MemoryEngine, request_context):
|
||||
bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}"
|
||||
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
|
||||
|
||||
op_id = uuid.uuid4()
|
||||
async with memory._pool.acquire() as conn:
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO async_operations (operation_id, bank_id, operation_type, status)
|
||||
VALUES ($1, $2, 'consolidation', 'processing')
|
||||
""",
|
||||
op_id,
|
||||
bank_id,
|
||||
)
|
||||
await conn.execute("DELETE FROM async_operations WHERE operation_id = $1", op_id)
|
||||
|
||||
assert await memory._check_op_alive(str(op_id)) is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_false_after_bank_cascade_delete(self, memory: MemoryEngine, request_context):
|
||||
bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}"
|
||||
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
|
||||
|
||||
op_id = uuid.uuid4()
|
||||
async with memory._pool.acquire() as conn:
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO async_operations (operation_id, bank_id, operation_type, status)
|
||||
VALUES ($1, $2, 'consolidation', 'processing')
|
||||
""",
|
||||
op_id,
|
||||
bank_id,
|
||||
)
|
||||
|
||||
# Delete the bank — cascades to the op row
|
||||
await memory.delete_bank(bank_id=bank_id, request_context=request_context)
|
||||
|
||||
assert await memory._check_op_alive(str(op_id)) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _mark_operation_completed / _mark_operation_failed graceful no-op
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMarkOperationGracefulOnMissingRow:
|
||||
@pytest.mark.asyncio
|
||||
async def test_mark_completed_does_not_raise_when_row_missing(self, memory: MemoryEngine):
|
||||
# Row never existed — should log and return cleanly
|
||||
missing_id = str(uuid.uuid4())
|
||||
await memory._mark_operation_completed(missing_id) # no exception
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mark_failed_does_not_raise_when_row_missing(self, memory: MemoryEngine):
|
||||
missing_id = str(uuid.uuid4())
|
||||
await memory._mark_operation_failed(missing_id, "some error", "traceback here") # no exception
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mark_completed_and_fire_webhook_does_not_raise_when_row_missing(
|
||||
self, memory: MemoryEngine
|
||||
):
|
||||
missing_id = str(uuid.uuid4())
|
||||
await memory._mark_operation_completed_and_fire_webhook(
|
||||
operation_id=missing_id,
|
||||
bank_id="nonexistent-bank",
|
||||
status="completed",
|
||||
result=None,
|
||||
) # no exception
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Consolidation checkpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConsolidationCheckpoint:
|
||||
@pytest.mark.asyncio
|
||||
async def test_consolidation_stops_early_when_op_cancelled(self, memory: MemoryEngine, request_context):
|
||||
"""Consolidation returns 'cancelled' status after the first batch if _check_op_alive is False."""
|
||||
from hindsight_api.config import _get_raw_config
|
||||
from hindsight_api.engine.consolidation.consolidator import run_consolidation_job
|
||||
|
||||
config = _get_raw_config()
|
||||
original = config.enable_observations
|
||||
config.enable_observations = True
|
||||
|
||||
try:
|
||||
bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}"
|
||||
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
|
||||
|
||||
# Insert a few unconsolidated memories directly so we control the batch without LLM
|
||||
async with memory._pool.acquire() as conn:
|
||||
for i in range(3):
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO memory_units
|
||||
(id, bank_id, text, fact_type, created_at, updated_at)
|
||||
VALUES (gen_random_uuid(), $1, $2, 'experience', NOW(), NOW())
|
||||
""",
|
||||
bank_id,
|
||||
f"Test memory {i} for cancellation test",
|
||||
)
|
||||
|
||||
op_id = str(uuid.uuid4())
|
||||
call_count = 0
|
||||
|
||||
async def _fake_check(operation_id: str) -> bool:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
# Return False on the very first checkpoint call
|
||||
return False
|
||||
|
||||
with patch.object(memory, "_check_op_alive", side_effect=_fake_check):
|
||||
result = await run_consolidation_job(
|
||||
memory_engine=memory,
|
||||
bank_id=bank_id,
|
||||
request_context=request_context,
|
||||
operation_id=op_id,
|
||||
)
|
||||
|
||||
assert result["status"] == "cancelled"
|
||||
assert call_count >= 1
|
||||
finally:
|
||||
config.enable_observations = original
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Retain checkpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRetainCheckpoint:
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_stops_between_sub_batches_when_cancelled(
|
||||
self, memory: MemoryEngine, request_context
|
||||
):
|
||||
"""retain_batch_async returns partial results if _check_op_alive is False between sub-batches."""
|
||||
from hindsight_api.config import _get_raw_config
|
||||
|
||||
bank_id = f"{_BANK_PREFIX}-{uuid.uuid4().hex[:8]}"
|
||||
await memory.get_bank_profile(bank_id=bank_id, request_context=request_context)
|
||||
|
||||
# Force sub-batch splitting by temporarily lowering the token threshold
|
||||
config = _get_raw_config()
|
||||
original_tokens = config.retain_batch_tokens
|
||||
# Set threshold very low so each item becomes its own sub-batch
|
||||
config.retain_batch_tokens = 1
|
||||
|
||||
try:
|
||||
op_id = str(uuid.uuid4())
|
||||
check_calls = 0
|
||||
|
||||
async def _fake_check(operation_id: str) -> bool:
|
||||
nonlocal check_calls
|
||||
check_calls += 1
|
||||
# Cancel after the first sub-batch completes
|
||||
return check_calls <= 1
|
||||
|
||||
contents = [
|
||||
{"content": f"Memory item {i} about something interesting."} for i in range(4)
|
||||
]
|
||||
|
||||
with patch.object(memory, "_check_op_alive", side_effect=_fake_check):
|
||||
result = await memory.retain_batch_async(
|
||||
bank_id=bank_id,
|
||||
contents=contents,
|
||||
request_context=request_context,
|
||||
operation_id=op_id,
|
||||
)
|
||||
|
||||
# Should have stopped early: fewer results than total items
|
||||
assert len(result) < len(contents), (
|
||||
f"Expected early stop but got {len(result)}/{len(contents)} results"
|
||||
)
|
||||
assert check_calls >= 1
|
||||
finally:
|
||||
config.retain_batch_tokens = original_tokens
|
||||
@@ -1,170 +0,0 @@
|
||||
"""Tests for source_facts token limiting in recall.
|
||||
|
||||
Covers:
|
||||
- max_source_facts_tokens: total token budget across all source facts
|
||||
- max_source_facts_tokens_per_observation: per-observation cap
|
||||
|
||||
Both parameters are tested at the recall_async level and verified to produce
|
||||
fewer source facts when the budget is tight vs. unlimited.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from hindsight_api.config import _get_raw_config
|
||||
from hindsight_api.engine.memory_engine import Budget
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def enable_observations():
|
||||
config = _get_raw_config()
|
||||
original = config.enable_observations
|
||||
config.enable_observations = True
|
||||
yield
|
||||
config.enable_observations = original
|
||||
|
||||
|
||||
async def _setup_bank_with_observations(memory, bank_id, request_context):
|
||||
"""Retain several memories and trigger consolidation to produce observations with source facts."""
|
||||
contents = [
|
||||
"Alice is a software engineer who loves Python programming.",
|
||||
"Alice has been working at TechCorp for 5 years.",
|
||||
"Alice recently completed a machine learning certification course.",
|
||||
"Alice mentors junior developers on the team.",
|
||||
"Alice prefers functional programming patterns in her code.",
|
||||
]
|
||||
for content in contents:
|
||||
await memory.retain_async(
|
||||
bank_id=bank_id,
|
||||
content=content,
|
||||
request_context=request_context,
|
||||
)
|
||||
await memory.run_consolidation(bank_id=bank_id, request_context=request_context)
|
||||
|
||||
|
||||
class TestRecallSourceFactsPerObservationCap:
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_observation_cap_reduces_source_facts(self, memory, request_context):
|
||||
"""A tight per-observation token cap should return fewer source facts than unlimited."""
|
||||
bank_id = "test-sf-per-obs-cap"
|
||||
try:
|
||||
await _setup_bank_with_observations(memory, bank_id, request_context)
|
||||
|
||||
result_limited = await memory.recall_async(
|
||||
bank_id=bank_id,
|
||||
query="Alice engineer",
|
||||
fact_type=["observation"],
|
||||
max_tokens=4096,
|
||||
include_source_facts=True,
|
||||
max_source_facts_tokens_per_observation=1, # Effectively cuts all source facts
|
||||
budget=Budget.MID,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
result_unlimited = await memory.recall_async(
|
||||
bank_id=bank_id,
|
||||
query="Alice engineer",
|
||||
fact_type=["observation"],
|
||||
max_tokens=4096,
|
||||
include_source_facts=True,
|
||||
max_source_facts_tokens_per_observation=-1,
|
||||
budget=Budget.MID,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
unlimited_count = len(result_unlimited.source_facts) if result_unlimited.source_facts else 0
|
||||
limited_count = len(result_limited.source_facts) if result_limited.source_facts else 0
|
||||
|
||||
if unlimited_count > 0:
|
||||
assert limited_count <= unlimited_count, (
|
||||
f"Per-observation cap should yield fewer source facts ({limited_count} <= {unlimited_count})"
|
||||
)
|
||||
finally:
|
||||
await memory.delete_bank(bank_id, request_context=request_context)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_per_observation_cap_does_not_mix_between_observations(self, memory, request_context):
|
||||
"""Each observation's source facts are capped independently — not as a shared pool."""
|
||||
bank_id = "test-sf-per-obs-independent"
|
||||
try:
|
||||
await _setup_bank_with_observations(memory, bank_id, request_context)
|
||||
|
||||
# With a generous per-observation limit each observation can have facts;
|
||||
# with a global limit of 1 token the first observation would consume the whole budget.
|
||||
result_per_obs = await memory.recall_async(
|
||||
bank_id=bank_id,
|
||||
query="Alice engineer",
|
||||
fact_type=["observation"],
|
||||
max_tokens=4096,
|
||||
include_source_facts=True,
|
||||
max_source_facts_tokens=4096, # large global budget
|
||||
max_source_facts_tokens_per_observation=512, # reasonable per-obs limit
|
||||
budget=Budget.MID,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
# Should not raise; source_facts may be populated for multiple observations
|
||||
assert result_per_obs.source_facts is not None or len(result_per_obs.results) == 0
|
||||
finally:
|
||||
await memory.delete_bank(bank_id, request_context=request_context)
|
||||
|
||||
|
||||
class TestRecallSourceFactsTotalBudget:
|
||||
@pytest.mark.asyncio
|
||||
async def test_total_budget_limits_source_facts(self, memory, request_context):
|
||||
"""A tight total token budget should return fewer source facts than unlimited."""
|
||||
bank_id = "test-sf-total-budget"
|
||||
try:
|
||||
await _setup_bank_with_observations(memory, bank_id, request_context)
|
||||
|
||||
result_tight = await memory.recall_async(
|
||||
bank_id=bank_id,
|
||||
query="Alice engineer",
|
||||
fact_type=["observation"],
|
||||
max_tokens=4096,
|
||||
include_source_facts=True,
|
||||
max_source_facts_tokens=1, # Effectively cuts all source facts
|
||||
budget=Budget.MID,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
result_unlimited = await memory.recall_async(
|
||||
bank_id=bank_id,
|
||||
query="Alice engineer",
|
||||
fact_type=["observation"],
|
||||
max_tokens=4096,
|
||||
include_source_facts=True,
|
||||
max_source_facts_tokens=-1,
|
||||
budget=Budget.MID,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
unlimited_count = len(result_unlimited.source_facts) if result_unlimited.source_facts else 0
|
||||
tight_count = len(result_tight.source_facts) if result_tight.source_facts else 0
|
||||
|
||||
if unlimited_count > 0:
|
||||
assert tight_count <= unlimited_count, (
|
||||
f"Total budget should yield fewer source facts ({tight_count} <= {unlimited_count})"
|
||||
)
|
||||
finally:
|
||||
await memory.delete_bank(bank_id, request_context=request_context)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_source_facts_without_flag(self, memory, request_context):
|
||||
"""source_facts should be None when include_source_facts is not set."""
|
||||
bank_id = "test-sf-no-flag"
|
||||
try:
|
||||
await _setup_bank_with_observations(memory, bank_id, request_context)
|
||||
|
||||
result = await memory.recall_async(
|
||||
bank_id=bank_id,
|
||||
query="Alice engineer",
|
||||
fact_type=["observation"],
|
||||
max_tokens=4096,
|
||||
include_source_facts=False, # default
|
||||
budget=Budget.MID,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
assert result.source_facts is None or len(result.source_facts) == 0
|
||||
finally:
|
||||
await memory.delete_bank(bank_id, request_context=request_context)
|
||||
+1
-1
@@ -46,4 +46,4 @@ __all__ = [
|
||||
"RemoteTEICrossEncoder",
|
||||
"LLMConfig",
|
||||
]
|
||||
__version__ = "0.4.18"
|
||||
__version__ = "0.4.16"
|
||||
+9
-81
@@ -14,8 +14,7 @@ from typing import Any
|
||||
import asyncpg
|
||||
import typer
|
||||
|
||||
from ..config import DEFAULT_DATABASE_SCHEMA, HindsightConfig
|
||||
from ..extensions import TenantExtension, load_extension
|
||||
from ..config import HindsightConfig
|
||||
from ..pg0 import parse_pg0_url, resolve_database_url
|
||||
|
||||
|
||||
@@ -215,81 +214,20 @@ def restore(
|
||||
typer.echo("Restore complete")
|
||||
|
||||
|
||||
async def _run_migration(
|
||||
db_url: str,
|
||||
schema: str | None = None,
|
||||
base_schema: str = DEFAULT_DATABASE_SCHEMA,
|
||||
embedding_dimension: int | None = None,
|
||||
) -> list[str]:
|
||||
"""Resolve database URL and run migrations for one schema or all discovered schemas."""
|
||||
from ..migrations import (
|
||||
ensure_embedding_dimension,
|
||||
ensure_text_search_extension,
|
||||
ensure_vector_extension,
|
||||
run_migrations,
|
||||
)
|
||||
async def _run_migration(db_url: str, schema: str = "public") -> None:
|
||||
"""Resolve database URL and run migrations."""
|
||||
from ..migrations import run_migrations
|
||||
|
||||
is_pg0, instance_name, _ = parse_pg0_url(db_url)
|
||||
if is_pg0:
|
||||
typer.echo(f"Starting embedded PostgreSQL (instance: {instance_name})...")
|
||||
resolved_url = await resolve_database_url(db_url)
|
||||
|
||||
config = HindsightConfig.from_env()
|
||||
if schema:
|
||||
schemas = [schema]
|
||||
else:
|
||||
tenant_extension = load_extension("TENANT", TenantExtension)
|
||||
|
||||
schemas = [base_schema or DEFAULT_DATABASE_SCHEMA]
|
||||
if tenant_extension:
|
||||
tenants = await tenant_extension.list_tenants()
|
||||
schemas.extend(tenant.schema for tenant in tenants if tenant.schema)
|
||||
|
||||
# Preserve order while removing duplicates.
|
||||
schemas = list(dict.fromkeys(schemas))
|
||||
|
||||
for schema in schemas:
|
||||
run_migrations(resolved_url, schema=schema)
|
||||
|
||||
if embedding_dimension is not None:
|
||||
for schema in schemas:
|
||||
ensure_embedding_dimension(
|
||||
resolved_url,
|
||||
embedding_dimension,
|
||||
schema=schema,
|
||||
vector_extension=config.vector_extension,
|
||||
)
|
||||
|
||||
for schema in schemas:
|
||||
ensure_vector_extension(
|
||||
resolved_url,
|
||||
vector_extension=config.vector_extension,
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
for schema in schemas:
|
||||
ensure_text_search_extension(
|
||||
resolved_url,
|
||||
text_search_extension=config.text_search_extension,
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
return schemas
|
||||
run_migrations(resolved_url, schema=schema)
|
||||
|
||||
|
||||
@app.command(name="run-db-migration")
|
||||
def run_db_migration(
|
||||
schema: str | None = typer.Option(
|
||||
None,
|
||||
"--schema",
|
||||
"-s",
|
||||
help="Database schema to run migrations on. If omitted, migrate the base schema and all discovered tenant schemas.",
|
||||
),
|
||||
embedding_dimension: int | None = typer.Option(
|
||||
None,
|
||||
"--embedding-dimension",
|
||||
help="Expected embedding dimension to enforce after migrations. Omit to skip dimension sync.",
|
||||
),
|
||||
schema: str = typer.Option("public", "--schema", "-s", help="Database schema to run migrations on"),
|
||||
):
|
||||
"""Run database migrations to the latest version."""
|
||||
config = HindsightConfig.from_env()
|
||||
@@ -299,21 +237,11 @@ def run_db_migration(
|
||||
typer.echo("Set HINDSIGHT_API_DATABASE_URL environment variable.", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
if schema:
|
||||
typer.echo(f"Running database migrations for schema: {schema}...")
|
||||
else:
|
||||
typer.echo("Running database migrations for base schema and all discovered tenant schemas...")
|
||||
typer.echo(f"Running database migrations (schema: {schema})...")
|
||||
|
||||
schemas = asyncio.run(
|
||||
_run_migration(
|
||||
config.database_url,
|
||||
schema=schema,
|
||||
base_schema=config.database_schema,
|
||||
embedding_dimension=embedding_dimension,
|
||||
)
|
||||
)
|
||||
asyncio.run(_run_migration(config.database_url, schema))
|
||||
|
||||
typer.echo(f"Database migrations completed successfully for {len(schemas)} schema(s)")
|
||||
typer.echo("Database migrations completed successfully")
|
||||
|
||||
|
||||
async def _decommission_worker(db_url: str, worker_id: str, schema: str = "public") -> int:
|
||||
+1
-1
@@ -34,7 +34,7 @@ def upgrade() -> None:
|
||||
# Create file_storage table (minimal: just key + data)
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {schema}file_storage (
|
||||
CREATE TABLE {schema}file_storage (
|
||||
storage_key TEXT PRIMARY KEY,
|
||||
data BYTEA NOT NULL
|
||||
)
|
||||
+1
-1
@@ -35,7 +35,7 @@ def upgrade() -> None:
|
||||
|
||||
# Add GIN index for JSONB containment queries (@> operator)
|
||||
op.execute(f"""
|
||||
CREATE INDEX IF NOT EXISTS idx_async_operations_result_metadata
|
||||
CREATE INDEX idx_async_operations_result_metadata
|
||||
ON {schema}async_operations
|
||||
USING gin(result_metadata)
|
||||
""")
|
||||
+22
-296
@@ -34,7 +34,7 @@ def _parse_metadata(metadata: Any) -> dict[str, Any]:
|
||||
|
||||
from typing import Callable
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
|
||||
@@ -73,13 +73,15 @@ def FieldWithDefault(default_factory: Callable, **kwargs) -> Any:
|
||||
from hindsight_api.config import get_config
|
||||
from hindsight_api.engine.memory_engine import Budget, _current_schema, _get_tiktoken_encoding, fq_table
|
||||
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES, MemoryFact, TokenUsage
|
||||
from hindsight_api.engine.search.tags import TagGroup, TagsMatch
|
||||
from hindsight_api.engine.search.tags import TagsMatch
|
||||
from hindsight_api.extensions import HttpExtension, OperationValidationError, load_extension
|
||||
from hindsight_api.metrics import create_metrics_collector, get_metrics_collector, initialize_metrics
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MAX_QUERY_TOKENS = 500 # Maximum tokens allowed in recall query
|
||||
|
||||
|
||||
class EntityIncludeOptions(BaseModel):
|
||||
"""Options for including entity observations in recall results."""
|
||||
@@ -96,12 +98,7 @@ class ChunkIncludeOptions(BaseModel):
|
||||
class SourceFactsIncludeOptions(BaseModel):
|
||||
"""Options for including source facts for observation-type results."""
|
||||
|
||||
max_tokens: int = Field(
|
||||
default=4096, description="Maximum total tokens for source facts across all observations (-1 = unlimited)"
|
||||
)
|
||||
max_tokens_per_observation: int = Field(
|
||||
default=-1, description="Maximum tokens of source facts per observation (-1 = unlimited)"
|
||||
)
|
||||
max_tokens: int = Field(default=4096, description="Maximum tokens for source facts")
|
||||
|
||||
|
||||
class IncludeOptions(BaseModel):
|
||||
@@ -163,17 +160,6 @@ class RecallRequest(BaseModel):
|
||||
description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), "
|
||||
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
|
||||
)
|
||||
tag_groups: list[TagGroup] | None = Field(
|
||||
default=None,
|
||||
description="Compound tag filter using boolean groups. Groups in the list are AND-ed. "
|
||||
"Each group is a leaf {tags, match} or compound {and: [...]}, {or: [...]}, {not: ...}.",
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_tags_exclusive(self) -> "RecallRequest":
|
||||
if self.tags is not None and self.tag_groups is not None:
|
||||
raise ValueError("'tags' and 'tag_groups' are mutually exclusive. Use 'tag_groups' for compound filtering.")
|
||||
return self
|
||||
|
||||
|
||||
class RecallResult(BaseModel):
|
||||
@@ -486,11 +472,6 @@ class FileRetainMetadata(BaseModel):
|
||||
metadata: dict[str, Any] | None = Field(default=None, description="Additional metadata")
|
||||
tags: list[str] | None = Field(default=None, description="Tags for this file")
|
||||
timestamp: str | None = Field(default=None, description="ISO timestamp")
|
||||
parser: str | list[str] | None = Field(
|
||||
default=None,
|
||||
description="Parser or ordered fallback chain for this file (overrides request-level parser). "
|
||||
"E.g. 'iris' or ['iris', 'markitdown'].",
|
||||
)
|
||||
|
||||
|
||||
class FileRetainRequest(BaseModel):
|
||||
@@ -499,21 +480,14 @@ class FileRetainRequest(BaseModel):
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"parser": "iris",
|
||||
"files_metadata": [
|
||||
{"document_id": "report_2024", "tags": ["quarterly"]},
|
||||
{"context": "meeting notes", "parser": ["iris", "markitdown"]},
|
||||
{"context": "meeting notes"},
|
||||
],
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
parser: str | list[str] | None = Field(
|
||||
default=None,
|
||||
description="Default parser or ordered fallback chain for all files in this request. "
|
||||
"E.g. 'markitdown' or ['iris', 'markitdown']. Falls back to server default if not set. "
|
||||
"Per-file 'parser' in files_metadata takes precedence over this value.",
|
||||
)
|
||||
files_metadata: list[FileRetainMetadata] | None = Field(
|
||||
default=None,
|
||||
description="Metadata for each file (optional, must match number of files if provided)",
|
||||
@@ -650,17 +624,6 @@ class ReflectRequest(BaseModel):
|
||||
description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), "
|
||||
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
|
||||
)
|
||||
tag_groups: list[TagGroup] | None = Field(
|
||||
default=None,
|
||||
description="Compound tag filter using boolean groups. Groups in the list are AND-ed. "
|
||||
"Each group is a leaf {tags, match} or compound {and: [...]}, {or: [...]}, {not: ...}.",
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_tags_exclusive(self) -> "ReflectRequest":
|
||||
if self.tags is not None and self.tag_groups is not None:
|
||||
raise ValueError("'tags' and 'tag_groups' are mutually exclusive. Use 'tag_groups' for compound filtering.")
|
||||
return self
|
||||
|
||||
|
||||
class ReflectFact(BaseModel):
|
||||
@@ -1226,30 +1189,6 @@ class DocumentResponse(BaseModel):
|
||||
tags: list[str] = FieldWithDefault(list, description="Tags associated with this document")
|
||||
|
||||
|
||||
class UpdateDocumentRequest(BaseModel):
|
||||
"""Request model for updating a document's mutable fields."""
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"tags": ["team-a", "team-b"],
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
tags: list[str] | None = Field(
|
||||
default=None,
|
||||
description="New tags for the document and its memory units. "
|
||||
"Triggers observation invalidation and re-consolidation.",
|
||||
)
|
||||
|
||||
|
||||
class UpdateDocumentResponse(BaseModel):
|
||||
"""Response model for update document endpoint."""
|
||||
|
||||
success: bool = True
|
||||
|
||||
|
||||
class DeleteDocumentResponse(BaseModel):
|
||||
"""Response model for delete document endpoint."""
|
||||
|
||||
@@ -1578,24 +1517,6 @@ class CancelOperationResponse(BaseModel):
|
||||
operation_id: str
|
||||
|
||||
|
||||
class RetryOperationResponse(BaseModel):
|
||||
"""Response model for retry operation endpoint."""
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"success": True,
|
||||
"message": "Operation 550e8400-e29b-41d4-a716-446655440000 queued for retry",
|
||||
"operation_id": "550e8400-e29b-41d4-a716-446655440000",
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
success: bool
|
||||
message: str
|
||||
operation_id: str
|
||||
|
||||
|
||||
class ChildOperationStatus(BaseModel):
|
||||
"""Status of a child operation (for batch operations)."""
|
||||
|
||||
@@ -2199,7 +2120,7 @@ def _register_routes(app: FastAPI):
|
||||
@app.get(
|
||||
"/v1/default/banks/{bank_id}/memories/{memory_id}",
|
||||
summary="Get memory unit",
|
||||
description="Get a single memory unit by ID with all its metadata including entities and tags. Note: the 'history' field is deprecated and always returns an empty list - use GET /memories/{memory_id}/history instead.",
|
||||
description="Get a single memory unit by ID with all its metadata including entities and tags.",
|
||||
operation_id="get_memory",
|
||||
tags=["Memory"],
|
||||
)
|
||||
@@ -2229,39 +2150,6 @@ def _register_routes(app: FastAPI):
|
||||
logger.error(f"Error in /v1/default/banks/{bank_id}/memories/{memory_id}: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.get(
|
||||
"/v1/default/banks/{bank_id}/memories/{memory_id}/history",
|
||||
summary="Get observation history",
|
||||
description="Get the full history of an observation, with each change's source facts resolved to their text.",
|
||||
operation_id="get_observation_history",
|
||||
tags=["Memory"],
|
||||
)
|
||||
async def api_get_observation_history(
|
||||
bank_id: str,
|
||||
memory_id: str,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""Get the history of a single observation by ID."""
|
||||
try:
|
||||
data = await app.state.memory.get_observation_history(
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
request_context=request_context,
|
||||
)
|
||||
if data is None:
|
||||
raise HTTPException(status_code=404, detail=f"Memory unit '{memory_id}' not found")
|
||||
return data
|
||||
except OperationValidationError as e:
|
||||
raise HTTPException(status_code=e.status_code, detail=e.reason)
|
||||
except (AuthenticationError, HTTPException):
|
||||
raise
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
|
||||
logger.error(f"Error in /v1/default/banks/{bank_id}/memories/{memory_id}/history: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.post(
|
||||
"/v1/default/banks/{bank_id}/memories/recall",
|
||||
response_model=RecallResponse,
|
||||
@@ -2283,13 +2171,12 @@ def _register_routes(app: FastAPI):
|
||||
metrics = get_metrics_collector()
|
||||
|
||||
# Validate query length to prevent expensive operations on oversized queries
|
||||
max_query_tokens = get_config().recall_max_query_tokens
|
||||
encoding = _get_tiktoken_encoding()
|
||||
query_tokens = len(encoding.encode(request.query))
|
||||
if query_tokens > max_query_tokens:
|
||||
if query_tokens > MAX_QUERY_TOKENS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Query too long: {query_tokens} tokens exceeds maximum of {max_query_tokens}. Please shorten your query.",
|
||||
detail=f"Query too long: {query_tokens} tokens exceeds maximum of {MAX_QUERY_TOKENS}. Please shorten your query.",
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -2318,9 +2205,6 @@ def _register_routes(app: FastAPI):
|
||||
# Determine source facts inclusion settings
|
||||
include_source_facts = request.include.source_facts is not None
|
||||
max_source_facts_tokens = request.include.source_facts.max_tokens if include_source_facts else 4096
|
||||
max_source_facts_tokens_per_observation = (
|
||||
request.include.source_facts.max_tokens_per_observation if include_source_facts else -1
|
||||
)
|
||||
|
||||
pre_recall = time.time() - handler_start
|
||||
# Run recall with tracing (record metrics)
|
||||
@@ -2342,11 +2226,9 @@ def _register_routes(app: FastAPI):
|
||||
max_chunk_tokens=max_chunk_tokens,
|
||||
include_source_facts=include_source_facts,
|
||||
max_source_facts_tokens=max_source_facts_tokens,
|
||||
max_source_facts_tokens_per_observation=max_source_facts_tokens_per_observation,
|
||||
request_context=request_context,
|
||||
tags=request.tags,
|
||||
tags_match=request.tags_match,
|
||||
tag_groups=request.tag_groups,
|
||||
)
|
||||
|
||||
# Convert core MemoryFact objects to API RecallResult objects (excluding internal metrics)
|
||||
@@ -2482,7 +2364,6 @@ def _register_routes(app: FastAPI):
|
||||
request_context=request_context,
|
||||
tags=request.tags,
|
||||
tags_match=request.tags_match,
|
||||
tag_groups=request.tag_groups,
|
||||
)
|
||||
|
||||
# Build based_on (memories + mental_models + directives) if facts are requested
|
||||
@@ -2813,41 +2694,6 @@ def _register_routes(app: FastAPI):
|
||||
logger.error(f"Error in GET /v1/default/banks/{bank_id}/mental-models/{mental_model_id}: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.get(
|
||||
"/v1/default/banks/{bank_id}/mental-models/{mental_model_id}/history",
|
||||
summary="Get mental model history",
|
||||
description="Get the refresh history of a mental model, showing content changes over time.",
|
||||
operation_id="get_mental_model_history",
|
||||
tags=["Mental Models"],
|
||||
)
|
||||
async def api_get_mental_model_history(
|
||||
bank_id: str,
|
||||
mental_model_id: str,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""Get the refresh history of a mental model."""
|
||||
try:
|
||||
data = await app.state.memory.get_mental_model_history(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=request_context,
|
||||
)
|
||||
if data is None:
|
||||
raise HTTPException(status_code=404, detail=f"Mental model '{mental_model_id}' not found")
|
||||
return data
|
||||
except (AuthenticationError, HTTPException):
|
||||
raise
|
||||
except OperationValidationError as e:
|
||||
raise HTTPException(status_code=e.status_code, detail=e.reason)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
|
||||
logger.error(
|
||||
f"Error in GET /v1/default/banks/{bank_id}/mental-models/{mental_model_id}/history: {error_detail}"
|
||||
)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.post(
|
||||
"/v1/default/banks/{bank_id}/mental-models",
|
||||
response_model=CreateMentalModelResponse,
|
||||
@@ -3369,55 +3215,6 @@ def _register_routes(app: FastAPI):
|
||||
logger.error(f"Error in /v1/default/chunks/{chunk_id}: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.patch(
|
||||
"/v1/default/banks/{bank_id}/documents/{document_id:path}",
|
||||
response_model=UpdateDocumentResponse,
|
||||
summary="Update document",
|
||||
description="Update mutable fields on a document without re-processing its content.\n\n"
|
||||
"**Tags** (`tags`): Propagated to all associated memory units. Observations derived from "
|
||||
"those units are invalidated and queued for re-consolidation under the new tags. "
|
||||
"Co-source memories from other documents that shared those observations are also reset.\n\n"
|
||||
"At least one field must be provided.",
|
||||
operation_id="update_document",
|
||||
tags=["Documents"],
|
||||
)
|
||||
async def api_update_document(
|
||||
bank_id: str,
|
||||
document_id: str,
|
||||
body: UpdateDocumentRequest,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""
|
||||
Update mutable fields on a document without re-processing its content.
|
||||
|
||||
Args:
|
||||
bank_id: Memory Bank ID (from path)
|
||||
document_id: Document ID (from path)
|
||||
body: Fields to update (tags, metadata, context)
|
||||
"""
|
||||
if body.tags is None:
|
||||
raise HTTPException(status_code=422, detail="At least one field (tags) must be provided")
|
||||
try:
|
||||
result = await app.state.memory.update_document(
|
||||
document_id,
|
||||
bank_id,
|
||||
tags=body.tags,
|
||||
request_context=request_context,
|
||||
)
|
||||
if not result:
|
||||
raise HTTPException(status_code=404, detail="Document not found")
|
||||
return UpdateDocumentResponse(success=True)
|
||||
except OperationValidationError as e:
|
||||
raise HTTPException(status_code=e.status_code, detail=e.reason)
|
||||
except (AuthenticationError, HTTPException):
|
||||
raise
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
|
||||
logger.error(f"Error in PATCH /v1/default/banks/{bank_id}/documents/{document_id}: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.delete(
|
||||
"/v1/default/banks/{bank_id}/documents/{document_id:path}",
|
||||
response_model=DeleteDocumentResponse,
|
||||
@@ -3468,17 +3265,13 @@ def _register_routes(app: FastAPI):
|
||||
"/v1/default/banks/{bank_id}/operations",
|
||||
response_model=OperationsListResponse,
|
||||
summary="List async operations",
|
||||
description="Get a list of async operations for a specific agent, with optional filtering by status and operation type. Results are sorted by most recent first.",
|
||||
description="Get a list of async operations for a specific agent, with optional filtering by status. Results are sorted by most recent first.",
|
||||
operation_id="list_operations",
|
||||
tags=["Operations"],
|
||||
)
|
||||
async def api_list_operations(
|
||||
bank_id: str,
|
||||
status: str | None = Query(default=None, description="Filter by status: pending, completed, or failed"),
|
||||
type: str | None = Query(
|
||||
default=None,
|
||||
description="Filter by operation type: retain, consolidation, refresh_mental_model, file_convert_retain, webhook_delivery",
|
||||
),
|
||||
limit: int = Query(default=20, ge=1, le=100, description="Maximum number of operations to return"),
|
||||
offset: int = Query(default=0, ge=0, description="Number of operations to skip"),
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
@@ -3486,7 +3279,7 @@ def _register_routes(app: FastAPI):
|
||||
"""List async operations for a memory bank with optional filtering and pagination."""
|
||||
try:
|
||||
result = await app.state.memory.list_operations(
|
||||
bank_id, status=status, task_type=type, limit=limit, offset=offset, request_context=request_context
|
||||
bank_id, status=status, limit=limit, offset=offset, request_context=request_context
|
||||
)
|
||||
return OperationsListResponse(
|
||||
bank_id=bank_id,
|
||||
@@ -3573,39 +3366,6 @@ def _register_routes(app: FastAPI):
|
||||
logger.error(f"Error in /v1/default/banks/{bank_id}/operations/{operation_id}: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.post(
|
||||
"/v1/default/banks/{bank_id}/operations/{operation_id}/retry",
|
||||
response_model=RetryOperationResponse,
|
||||
summary="Retry a failed async operation",
|
||||
description="Re-queue a failed async operation so the worker picks it up again",
|
||||
operation_id="retry_operation",
|
||||
tags=["Operations"],
|
||||
)
|
||||
async def api_retry_operation(
|
||||
bank_id: str, operation_id: str, request_context: RequestContext = Depends(get_request_context)
|
||||
):
|
||||
"""Retry a failed async operation."""
|
||||
try:
|
||||
try:
|
||||
uuid.UUID(operation_id)
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=400, detail=f"Invalid operation_id format: {operation_id}")
|
||||
|
||||
result = await app.state.memory.retry_operation(bank_id, operation_id, request_context=request_context)
|
||||
return RetryOperationResponse(**result)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
except OperationValidationError as e:
|
||||
raise HTTPException(status_code=e.status_code, detail=e.reason)
|
||||
except (AuthenticationError, HTTPException):
|
||||
raise
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
|
||||
logger.error(f"Error in POST /v1/default/banks/{bank_id}/operations/{operation_id}/retry: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.get(
|
||||
"/v1/default/banks/{bank_id}/profile",
|
||||
response_model=BankProfileResponse,
|
||||
@@ -4115,10 +3875,6 @@ def _register_routes(app: FastAPI):
|
||||
try:
|
||||
pool = await app.state.memory._get_pool()
|
||||
from hindsight_api.engine.memory_engine import fq_table
|
||||
from hindsight_api.engine.retain import bank_utils
|
||||
|
||||
# Ensure the bank row exists before inserting into webhooks (FK constraint).
|
||||
await bank_utils.get_bank_profile(pool, bank_id)
|
||||
|
||||
webhook_id = uuid.uuid4()
|
||||
now = datetime.utcnow().isoformat() + "Z"
|
||||
@@ -4556,14 +4312,8 @@ def _register_routes(app: FastAPI):
|
||||
"Use the operations endpoint to monitor progress.\n\n"
|
||||
"**Request format:** multipart/form-data with:\n"
|
||||
"- `files`: One or more files to upload\n"
|
||||
"- `request`: JSON string with FileRetainRequest model\n\n"
|
||||
"**Parser selection:**\n"
|
||||
"- Set `parser` in the request body to override the server default for all files.\n"
|
||||
"- Set `parser` inside a `files_metadata` entry for per-file control.\n"
|
||||
"- Pass a list (e.g. `['iris', 'markitdown']`) to define an ordered fallback chain — "
|
||||
"each parser is tried in sequence until one succeeds.\n"
|
||||
"- Falls back to the server default (`HINDSIGHT_API_FILE_PARSER`) if not specified.\n"
|
||||
"- Only parsers enabled on the server may be requested; others return HTTP 400.",
|
||||
"- `request`: JSON string with FileRetainRequest model (files_metadata)\n\n"
|
||||
"**Note:** File parser is configured server-side via `HINDSIGHT_API_FILE_PARSER` (default: markitdown).",
|
||||
operation_id="file_retain",
|
||||
tags=["Files"],
|
||||
)
|
||||
@@ -4609,39 +4359,20 @@ def _register_routes(app: FastAPI):
|
||||
detail=f"files_metadata count ({len(request_data.files_metadata)}) must match files count ({len(files)})",
|
||||
)
|
||||
|
||||
# Resolve the registered parser names for allowlist validation
|
||||
registered_parsers = app.state.memory._parser_registry.list_parsers()
|
||||
allowlist = config.file_parser_allowlist if config.file_parser_allowlist is not None else registered_parsers
|
||||
|
||||
def _resolve_parser(raw: str | list[str] | None) -> list[str]:
|
||||
"""Normalize parser value to a non-empty list of names."""
|
||||
if raw is None:
|
||||
return config.file_parser
|
||||
return [raw] if isinstance(raw, str) else list(raw)
|
||||
|
||||
def _validate_parsers(parsers: list[str], context: str) -> None:
|
||||
"""Raise HTTP 400 if any parser name is not in the allowlist."""
|
||||
disallowed = [p for p in parsers if p not in allowlist]
|
||||
if disallowed:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Parser(s) not available ({context}): {disallowed}. Available: {allowlist}",
|
||||
)
|
||||
|
||||
# Validate request-level parser early (before reading files)
|
||||
if request_data.parser is not None:
|
||||
_validate_parsers(_resolve_parser(request_data.parser), "request-level parser")
|
||||
|
||||
# Prepare file items and calculate total batch size
|
||||
import io
|
||||
|
||||
file_items = []
|
||||
total_batch_size = 0
|
||||
|
||||
for i, file in enumerate(files):
|
||||
# Read file content to check size
|
||||
file_content = await file.read()
|
||||
total_batch_size += len(file_content)
|
||||
size = len(file_content)
|
||||
total_batch_size += size
|
||||
|
||||
# Create a temporary file-like object from the bytes
|
||||
import io
|
||||
|
||||
file_obj = io.BytesIO(file_content)
|
||||
|
||||
# Create a mock UploadFile with the necessary attributes
|
||||
class FileWrapper:
|
||||
@@ -4649,6 +4380,7 @@ def _register_routes(app: FastAPI):
|
||||
self._content = content
|
||||
self.filename = filename
|
||||
self.content_type = content_type
|
||||
self._buffer = io.BytesIO(content)
|
||||
|
||||
async def read(self):
|
||||
return self._content
|
||||
@@ -4659,12 +4391,6 @@ def _register_routes(app: FastAPI):
|
||||
file_meta = request_data.files_metadata[i] if request_data.files_metadata else FileRetainMetadata()
|
||||
doc_id = file_meta.document_id or f"file_{uuid.uuid4()}"
|
||||
|
||||
# Resolve and validate per-file parser chain
|
||||
# Priority: per-file > request-level > server default
|
||||
raw_parser = file_meta.parser if file_meta.parser is not None else request_data.parser
|
||||
parser_chain = _resolve_parser(raw_parser)
|
||||
_validate_parsers(parser_chain, f"file '{file.filename}'")
|
||||
|
||||
item = {
|
||||
"file": wrapped_file,
|
||||
"document_id": doc_id,
|
||||
@@ -4672,7 +4398,6 @@ def _register_routes(app: FastAPI):
|
||||
"metadata": file_meta.metadata or {},
|
||||
"tags": file_meta.tags or [],
|
||||
"timestamp": file_meta.timestamp,
|
||||
"parser": parser_chain,
|
||||
}
|
||||
file_items.append(item)
|
||||
|
||||
@@ -4687,6 +4412,7 @@ def _register_routes(app: FastAPI):
|
||||
result = await app.state.memory.submit_async_file_retain(
|
||||
bank_id=bank_id,
|
||||
file_items=file_items,
|
||||
parser=config.file_parser,
|
||||
document_tags=None,
|
||||
request_context=request_context,
|
||||
)
|
||||
@@ -381,15 +381,6 @@ class MCPMiddleware:
|
||||
# Clear root_path since we're passing directly to the app
|
||||
new_scope["root_path"] = ""
|
||||
|
||||
# Ensure Accept header includes required MIME types for MCP SDK.
|
||||
# Some clients (e.g., Claude Code) don't send Accept, causing
|
||||
# the SDK to reject with 406 Not Acceptable.
|
||||
accept_header = self._get_header(new_scope, "accept")
|
||||
if not accept_header or "text/event-stream" not in accept_header:
|
||||
headers = [(k, v) for k, v in new_scope.get("headers", []) if k.lower() != b"accept"]
|
||||
headers.append((b"accept", b"application/json, text/event-stream"))
|
||||
new_scope["headers"] = headers
|
||||
|
||||
# Wrap send to rewrite the SSE endpoint URL to include bank_id if using path-based routing.
|
||||
# Only rewrite SSE (text/event-stream) responses to avoid corrupting tool results
|
||||
# that might contain the literal string "data: /messages".
|
||||
@@ -193,7 +193,6 @@ ENV_EMBEDDINGS_LITELLM_MODEL = "HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL"
|
||||
ENV_RERANKER_LITELLM_API_BASE = "HINDSIGHT_API_RERANKER_LITELLM_API_BASE"
|
||||
ENV_RERANKER_LITELLM_API_KEY = "HINDSIGHT_API_RERANKER_LITELLM_API_KEY"
|
||||
ENV_RERANKER_LITELLM_MODEL = "HINDSIGHT_API_RERANKER_LITELLM_MODEL"
|
||||
ENV_RERANKER_LITELLM_MAX_TOKENS_PER_DOC = "HINDSIGHT_API_RERANKER_LITELLM_MAX_TOKENS_PER_DOC"
|
||||
|
||||
# LiteLLM SDK configuration (direct API access, no proxy needed)
|
||||
ENV_EMBEDDINGS_LITELLM_SDK_API_KEY = "HINDSIGHT_API_EMBEDDINGS_LITELLM_SDK_API_KEY"
|
||||
@@ -212,9 +211,6 @@ ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
|
||||
ENV_RERANKER_LOCAL_FORCE_CPU = "HINDSIGHT_API_RERANKER_LOCAL_FORCE_CPU"
|
||||
ENV_RERANKER_LOCAL_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT"
|
||||
ENV_RERANKER_LOCAL_TRUST_REMOTE_CODE = "HINDSIGHT_API_RERANKER_LOCAL_TRUST_REMOTE_CODE"
|
||||
ENV_RERANKER_LOCAL_FP16 = "HINDSIGHT_API_RERANKER_LOCAL_FP16"
|
||||
ENV_RERANKER_LOCAL_BUCKET_BATCHING = "HINDSIGHT_API_RERANKER_LOCAL_BUCKET_BATCHING"
|
||||
ENV_RERANKER_LOCAL_BATCH_SIZE = "HINDSIGHT_API_RERANKER_LOCAL_BATCH_SIZE"
|
||||
ENV_RERANKER_TEI_URL = "HINDSIGHT_API_RERANKER_TEI_URL"
|
||||
ENV_RERANKER_TEI_BATCH_SIZE = "HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE"
|
||||
ENV_RERANKER_TEI_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT"
|
||||
@@ -242,7 +238,6 @@ ENV_GRAPH_RETRIEVER = "HINDSIGHT_API_GRAPH_RETRIEVER"
|
||||
ENV_MPFP_TOP_K_NEIGHBORS = "HINDSIGHT_API_MPFP_TOP_K_NEIGHBORS"
|
||||
ENV_RECALL_MAX_CONCURRENT = "HINDSIGHT_API_RECALL_MAX_CONCURRENT"
|
||||
ENV_RECALL_CONNECTION_BUDGET = "HINDSIGHT_API_RECALL_CONNECTION_BUDGET"
|
||||
ENV_RECALL_MAX_QUERY_TOKENS = "HINDSIGHT_API_RECALL_MAX_QUERY_TOKENS"
|
||||
ENV_MENTAL_MODEL_REFRESH_CONCURRENCY = "HINDSIGHT_API_MENTAL_MODEL_REFRESH_CONCURRENCY"
|
||||
|
||||
# OpenTelemetry tracing configuration
|
||||
@@ -285,7 +280,6 @@ ENV_FILE_STORAGE_AZURE_CONTAINER = "HINDSIGHT_API_FILE_STORAGE_AZURE_CONTAINER"
|
||||
ENV_FILE_STORAGE_AZURE_ACCOUNT_NAME = "HINDSIGHT_API_FILE_STORAGE_AZURE_ACCOUNT_NAME"
|
||||
ENV_FILE_STORAGE_AZURE_ACCOUNT_KEY = "HINDSIGHT_API_FILE_STORAGE_AZURE_ACCOUNT_KEY"
|
||||
ENV_FILE_PARSER = "HINDSIGHT_API_FILE_PARSER"
|
||||
ENV_FILE_PARSER_ALLOWLIST = "HINDSIGHT_API_FILE_PARSER_ALLOWLIST"
|
||||
ENV_FILE_PARSER_IRIS_TOKEN = "HINDSIGHT_API_FILE_PARSER_IRIS_TOKEN"
|
||||
ENV_FILE_PARSER_IRIS_ORG_ID = "HINDSIGHT_API_FILE_PARSER_IRIS_ORG_ID"
|
||||
ENV_FILE_CONVERSION_MAX_BATCH_SIZE_MB = "HINDSIGHT_API_FILE_CONVERSION_MAX_BATCH_SIZE_MB"
|
||||
@@ -298,13 +292,7 @@ ENV_ENABLE_OBSERVATIONS = "HINDSIGHT_API_ENABLE_OBSERVATIONS"
|
||||
ENV_CONSOLIDATION_BATCH_SIZE = "HINDSIGHT_API_CONSOLIDATION_BATCH_SIZE"
|
||||
ENV_CONSOLIDATION_LLM_BATCH_SIZE = "HINDSIGHT_API_CONSOLIDATION_LLM_BATCH_SIZE"
|
||||
ENV_CONSOLIDATION_MAX_TOKENS = "HINDSIGHT_API_CONSOLIDATION_MAX_TOKENS"
|
||||
ENV_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS = "HINDSIGHT_API_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS"
|
||||
ENV_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS_PER_OBSERVATION = (
|
||||
"HINDSIGHT_API_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS_PER_OBSERVATION"
|
||||
)
|
||||
ENV_OBSERVATIONS_MISSION = "HINDSIGHT_API_OBSERVATIONS_MISSION"
|
||||
ENV_ENABLE_OBSERVATION_HISTORY = "HINDSIGHT_API_ENABLE_OBSERVATION_HISTORY"
|
||||
ENV_ENABLE_MENTAL_MODEL_HISTORY = "HINDSIGHT_API_ENABLE_MENTAL_MODEL_HISTORY"
|
||||
|
||||
# Webhook configuration (global, static - server-level only)
|
||||
ENV_WEBHOOK_URL = "HINDSIGHT_API_WEBHOOK_URL"
|
||||
@@ -355,7 +343,6 @@ PROVIDER_DEFAULT_MODELS = {
|
||||
"anthropic": "claude-haiku-4-5-20251001",
|
||||
"gemini": "gemini-2.5-flash",
|
||||
"groq": "openai/gpt-oss-120b",
|
||||
"minimax": "MiniMax-M2.5",
|
||||
"ollama": "gemma3:12b",
|
||||
"lmstudio": "local-model",
|
||||
"vertexai": "google/gemini-2.5-flash-lite",
|
||||
@@ -392,9 +379,6 @@ DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT = 4 # Limit concurrent CPU-bound rerankin
|
||||
DEFAULT_RERANKER_LOCAL_TRUST_REMOTE_CODE = (
|
||||
False # Security: disabled by default, required for some models like jina-reranker-v2
|
||||
)
|
||||
DEFAULT_RERANKER_LOCAL_FP16 = False # FP16 inference: opt-in, faster on MPS/CUDA (not CPU)
|
||||
DEFAULT_RERANKER_LOCAL_BUCKET_BATCHING = False # Length-sorted bucket batching: opt-in, 36-54% speedup
|
||||
DEFAULT_RERANKER_LOCAL_BATCH_SIZE = 32 # Batch size for local reranker predict() calls
|
||||
DEFAULT_RERANKER_TEI_BATCH_SIZE = 128
|
||||
DEFAULT_RERANKER_TEI_MAX_CONCURRENT = 8
|
||||
DEFAULT_RERANKER_MAX_CANDIDATES = 300
|
||||
@@ -416,7 +400,6 @@ DEFAULT_TEXT_SEARCH_EXTENSION = "native" # Options: "native", "vchord", "pg_tex
|
||||
DEFAULT_LITELLM_API_BASE = "http://localhost:4000"
|
||||
DEFAULT_EMBEDDINGS_LITELLM_MODEL = "text-embedding-3-small"
|
||||
DEFAULT_RERANKER_LITELLM_MODEL = "cohere/rerank-english-v3.0"
|
||||
DEFAULT_RERANKER_LITELLM_MAX_TOKENS_PER_DOC: int | None = None
|
||||
|
||||
# LiteLLM SDK defaults
|
||||
DEFAULT_EMBEDDINGS_LITELLM_SDK_MODEL = "cohere/embed-english-v3.0"
|
||||
@@ -435,7 +418,6 @@ DEFAULT_GRAPH_RETRIEVER = "link_expansion" # Options: "link_expansion", "mpfp",
|
||||
DEFAULT_MPFP_TOP_K_NEIGHBORS = 20 # Fan-out limit per node in MPFP graph traversal
|
||||
DEFAULT_RECALL_MAX_CONCURRENT = 32 # Max concurrent recall operations per worker
|
||||
DEFAULT_RECALL_CONNECTION_BUDGET = 4 # Max concurrent DB connections per recall operation
|
||||
DEFAULT_RECALL_MAX_QUERY_TOKENS = 500 # Maximum tokens allowed in recall query
|
||||
DEFAULT_MENTAL_MODEL_REFRESH_CONCURRENCY = 8 # Max concurrent mental model refreshes
|
||||
|
||||
# Retain settings
|
||||
@@ -453,8 +435,7 @@ DEFAULT_RETAIN_BATCH_POLL_INTERVAL_SECONDS = 60 # Batch API polling interval in
|
||||
|
||||
# File storage defaults
|
||||
DEFAULT_FILE_STORAGE_TYPE = "native" # PostgreSQL BYTEA storage
|
||||
DEFAULT_FILE_PARSER = "markitdown" # Default parser fallback chain (comma-separated, e.g. "iris,markitdown")
|
||||
DEFAULT_FILE_PARSER_ALLOWLIST = None # Allowlist of parsers clients may request (None = all registered parsers)
|
||||
DEFAULT_FILE_PARSER = "markitdown" # File parser to use (markitdown is the only supported parser)
|
||||
DEFAULT_FILE_CONVERSION_MAX_BATCH_SIZE_MB = 100 # Max total batch size in MB (all files combined)
|
||||
DEFAULT_FILE_CONVERSION_MAX_BATCH_SIZE = 10 # Max files per batch upload
|
||||
DEFAULT_ENABLE_FILE_UPLOAD_API = True # Enable file upload endpoint
|
||||
@@ -462,17 +443,9 @@ DEFAULT_FILE_DELETE_AFTER_RETAIN = True # Delete file bytes after retain (saves
|
||||
|
||||
# Observations defaults (consolidated knowledge from facts)
|
||||
DEFAULT_ENABLE_OBSERVATIONS = True # Observations enabled by default
|
||||
DEFAULT_ENABLE_OBSERVATION_HISTORY = True # Observation history tracking enabled by default
|
||||
DEFAULT_ENABLE_MENTAL_MODEL_HISTORY = True # Mental model history tracking enabled by default
|
||||
DEFAULT_CONSOLIDATION_BATCH_SIZE = 50 # Memories to load per batch (internal memory optimization)
|
||||
DEFAULT_CONSOLIDATION_LLM_BATCH_SIZE = 8 # Facts per LLM call (1 = no batching; >1 = batch mode)
|
||||
DEFAULT_CONSOLIDATION_MAX_TOKENS = 512 # Max tokens for recall when finding related observations
|
||||
DEFAULT_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS = (
|
||||
-1
|
||||
) # Total token budget for source facts in consolidation recall (-1 = unlimited)
|
||||
DEFAULT_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS_PER_OBSERVATION = (
|
||||
256 # Max tokens of source facts per observation in consolidation prompt (-1 = unlimited)
|
||||
)
|
||||
DEFAULT_OBSERVATIONS_MISSION = None # Declarative spec of what observations are for this bank
|
||||
|
||||
# Database migrations
|
||||
@@ -567,11 +540,6 @@ class JsonFormatter(logging.Formatter):
|
||||
return json.dumps(log_entry)
|
||||
|
||||
|
||||
def _parse_str_list(value: str) -> list[str]:
|
||||
"""Parse a comma-separated string into a non-empty list of stripped tokens."""
|
||||
return [v.strip() for v in value.split(",") if v.strip()]
|
||||
|
||||
|
||||
def _validate_extraction_mode(mode: str) -> str:
|
||||
"""Validate and normalize extraction mode."""
|
||||
mode_lower = mode.lower()
|
||||
@@ -674,9 +642,6 @@ class HindsightConfig:
|
||||
reranker_local_force_cpu: bool
|
||||
reranker_local_max_concurrent: int
|
||||
reranker_local_trust_remote_code: bool
|
||||
reranker_local_fp16: bool
|
||||
reranker_local_bucket_batching: bool
|
||||
reranker_local_batch_size: int
|
||||
reranker_tei_url: str | None
|
||||
reranker_tei_batch_size: int
|
||||
reranker_tei_max_concurrent: int
|
||||
@@ -687,7 +652,6 @@ class HindsightConfig:
|
||||
reranker_litellm_api_base: str
|
||||
reranker_litellm_api_key: str | None
|
||||
reranker_litellm_model: str
|
||||
reranker_litellm_max_tokens_per_doc: int | None
|
||||
reranker_litellm_sdk_api_key: str | None
|
||||
reranker_litellm_sdk_model: str
|
||||
reranker_litellm_sdk_api_base: str | None
|
||||
@@ -709,7 +673,6 @@ class HindsightConfig:
|
||||
mpfp_top_k_neighbors: int
|
||||
recall_max_concurrent: int
|
||||
recall_connection_budget: int
|
||||
recall_max_query_tokens: int
|
||||
mental_model_refresh_concurrency: int
|
||||
|
||||
# Retain settings
|
||||
@@ -736,8 +699,7 @@ class HindsightConfig:
|
||||
file_storage_azure_container: str | None # Azure container name (required for azure storage)
|
||||
file_storage_azure_account_name: str | None # Azure storage account name
|
||||
file_storage_azure_account_key: str | None # Azure storage account key
|
||||
file_parser: list[str] # Ordered fallback chain of parsers (e.g. ["iris", "markitdown"])
|
||||
file_parser_allowlist: list[str] | None # Parsers clients may request (None = all registered)
|
||||
file_parser: str # File parser to use (e.g., "markitdown", "iris")
|
||||
file_parser_iris_token: str | None # Vectorize API token for iris parser (VECTORIZE_TOKEN)
|
||||
file_parser_iris_org_id: str | None # Vectorize org ID for iris parser (VECTORIZE_ORG_ID)
|
||||
file_conversion_max_batch_size_mb: int # Max total batch size in MB (all files combined)
|
||||
@@ -747,13 +709,9 @@ class HindsightConfig:
|
||||
|
||||
# Observations settings (consolidated knowledge from facts)
|
||||
enable_observations: bool
|
||||
enable_observation_history: bool
|
||||
enable_mental_model_history: bool
|
||||
consolidation_batch_size: int
|
||||
consolidation_llm_batch_size: int
|
||||
consolidation_max_tokens: int
|
||||
consolidation_source_facts_max_tokens: int
|
||||
consolidation_source_facts_max_tokens_per_observation: int
|
||||
observations_mission: str | None
|
||||
|
||||
# Entity labels (controlled vocabulary of key:value classification labels extracted at retain time)
|
||||
@@ -854,9 +812,6 @@ class HindsightConfig:
|
||||
"entities_allow_free_form",
|
||||
# Consolidation settings
|
||||
"enable_observations",
|
||||
"consolidation_llm_batch_size",
|
||||
"consolidation_source_facts_max_tokens",
|
||||
"consolidation_source_facts_max_tokens_per_observation",
|
||||
"observations_mission",
|
||||
# Reflect settings
|
||||
"reflect_mission",
|
||||
@@ -1101,17 +1056,6 @@ class HindsightConfig:
|
||||
ENV_RERANKER_LOCAL_TRUST_REMOTE_CODE, str(DEFAULT_RERANKER_LOCAL_TRUST_REMOTE_CODE)
|
||||
).lower()
|
||||
in ("true", "1"),
|
||||
reranker_local_fp16=os.getenv(
|
||||
ENV_RERANKER_LOCAL_FP16, str(DEFAULT_RERANKER_LOCAL_FP16)
|
||||
).lower()
|
||||
in ("true", "1"),
|
||||
reranker_local_bucket_batching=os.getenv(
|
||||
ENV_RERANKER_LOCAL_BUCKET_BATCHING, str(DEFAULT_RERANKER_LOCAL_BUCKET_BATCHING)
|
||||
).lower()
|
||||
in ("true", "1"),
|
||||
reranker_local_batch_size=int(
|
||||
os.getenv(ENV_RERANKER_LOCAL_BATCH_SIZE, str(DEFAULT_RERANKER_LOCAL_BATCH_SIZE))
|
||||
),
|
||||
reranker_tei_url=os.getenv(ENV_RERANKER_TEI_URL),
|
||||
reranker_tei_batch_size=int(os.getenv(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE))),
|
||||
reranker_tei_max_concurrent=int(
|
||||
@@ -1127,9 +1071,6 @@ class HindsightConfig:
|
||||
or os.getenv(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE),
|
||||
reranker_litellm_api_key=os.getenv(ENV_RERANKER_LITELLM_API_KEY) or os.getenv(ENV_LITELLM_API_KEY),
|
||||
reranker_litellm_model=os.getenv(ENV_RERANKER_LITELLM_MODEL, DEFAULT_RERANKER_LITELLM_MODEL),
|
||||
reranker_litellm_max_tokens_per_doc=int(v)
|
||||
if (v := os.getenv(ENV_RERANKER_LITELLM_MAX_TOKENS_PER_DOC))
|
||||
else DEFAULT_RERANKER_LITELLM_MAX_TOKENS_PER_DOC,
|
||||
# LiteLLM SDK reranker (direct API access)
|
||||
reranker_litellm_sdk_api_key=os.getenv(ENV_RERANKER_LITELLM_SDK_API_KEY),
|
||||
reranker_litellm_sdk_model=os.getenv(ENV_RERANKER_LITELLM_SDK_MODEL, DEFAULT_RERANKER_LITELLM_SDK_MODEL),
|
||||
@@ -1156,7 +1097,6 @@ class HindsightConfig:
|
||||
recall_connection_budget=int(
|
||||
os.getenv(ENV_RECALL_CONNECTION_BUDGET, str(DEFAULT_RECALL_CONNECTION_BUDGET))
|
||||
),
|
||||
recall_max_query_tokens=int(os.getenv(ENV_RECALL_MAX_QUERY_TOKENS, str(DEFAULT_RECALL_MAX_QUERY_TOKENS))),
|
||||
mental_model_refresh_concurrency=int(
|
||||
os.getenv(ENV_MENTAL_MODEL_REFRESH_CONCURRENCY, str(DEFAULT_MENTAL_MODEL_REFRESH_CONCURRENCY))
|
||||
),
|
||||
@@ -1196,10 +1136,7 @@ class HindsightConfig:
|
||||
file_storage_azure_container=os.getenv(ENV_FILE_STORAGE_AZURE_CONTAINER) or None,
|
||||
file_storage_azure_account_name=os.getenv(ENV_FILE_STORAGE_AZURE_ACCOUNT_NAME) or None,
|
||||
file_storage_azure_account_key=os.getenv(ENV_FILE_STORAGE_AZURE_ACCOUNT_KEY) or None,
|
||||
file_parser=_parse_str_list(os.getenv(ENV_FILE_PARSER, DEFAULT_FILE_PARSER)),
|
||||
file_parser_allowlist=_parse_str_list(os.getenv(ENV_FILE_PARSER_ALLOWLIST))
|
||||
if os.getenv(ENV_FILE_PARSER_ALLOWLIST)
|
||||
else None,
|
||||
file_parser=os.getenv(ENV_FILE_PARSER, DEFAULT_FILE_PARSER),
|
||||
file_parser_iris_token=os.getenv(ENV_FILE_PARSER_IRIS_TOKEN) or None,
|
||||
file_parser_iris_org_id=os.getenv(ENV_FILE_PARSER_IRIS_ORG_ID) or None,
|
||||
file_conversion_max_batch_size_mb=int(
|
||||
@@ -1216,14 +1153,6 @@ class HindsightConfig:
|
||||
== "true",
|
||||
# Observations settings (consolidated knowledge from facts)
|
||||
enable_observations=os.getenv(ENV_ENABLE_OBSERVATIONS, str(DEFAULT_ENABLE_OBSERVATIONS)).lower() == "true",
|
||||
enable_observation_history=os.getenv(
|
||||
ENV_ENABLE_OBSERVATION_HISTORY, str(DEFAULT_ENABLE_OBSERVATION_HISTORY)
|
||||
).lower()
|
||||
== "true",
|
||||
enable_mental_model_history=os.getenv(
|
||||
ENV_ENABLE_MENTAL_MODEL_HISTORY, str(DEFAULT_ENABLE_MENTAL_MODEL_HISTORY)
|
||||
).lower()
|
||||
== "true",
|
||||
consolidation_batch_size=int(
|
||||
os.getenv(ENV_CONSOLIDATION_BATCH_SIZE, str(DEFAULT_CONSOLIDATION_BATCH_SIZE))
|
||||
),
|
||||
@@ -1233,15 +1162,6 @@ class HindsightConfig:
|
||||
consolidation_max_tokens=int(
|
||||
os.getenv(ENV_CONSOLIDATION_MAX_TOKENS, str(DEFAULT_CONSOLIDATION_MAX_TOKENS))
|
||||
),
|
||||
consolidation_source_facts_max_tokens=int(
|
||||
os.getenv(ENV_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS, str(DEFAULT_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS))
|
||||
),
|
||||
consolidation_source_facts_max_tokens_per_observation=int(
|
||||
os.getenv(
|
||||
ENV_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS_PER_OBSERVATION,
|
||||
str(DEFAULT_CONSOLIDATION_SOURCE_FACTS_MAX_TOKENS_PER_OBSERVATION),
|
||||
)
|
||||
),
|
||||
observations_mission=os.getenv(ENV_OBSERVATIONS_MISSION) or DEFAULT_OBSERVATIONS_MISSION,
|
||||
entity_labels=None,
|
||||
entities_allow_free_form=True,
|
||||
+18
-49
@@ -24,10 +24,9 @@ from datetime import datetime, timezone
|
||||
from itertools import combinations
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from pydantic import BaseModel, field_validator
|
||||
from pydantic import BaseModel
|
||||
|
||||
from ...config import get_config
|
||||
from ..llm_wrapper import sanitize_llm_output
|
||||
from ..memory_engine import fq_table
|
||||
from ..retain import embedding_utils
|
||||
from .prompts import build_batch_consolidation_prompt
|
||||
@@ -46,22 +45,12 @@ class _CreateAction(BaseModel):
|
||||
text: str
|
||||
source_fact_ids: list[str] # memory UUIDs from the NEW FACTS list
|
||||
|
||||
@field_validator("text", mode="before")
|
||||
@classmethod
|
||||
def sanitize_text(cls, v: str) -> str:
|
||||
return sanitize_llm_output(v) or ""
|
||||
|
||||
|
||||
class _UpdateAction(BaseModel):
|
||||
text: str
|
||||
observation_id: str # UUID of the existing observation to update
|
||||
source_fact_ids: list[str] # memory UUIDs from the NEW FACTS list
|
||||
|
||||
@field_validator("text", mode="before")
|
||||
@classmethod
|
||||
def sanitize_text(cls, v: str) -> str:
|
||||
return sanitize_llm_output(v) or ""
|
||||
|
||||
|
||||
class _DeleteAction(BaseModel):
|
||||
observation_id: str # UUID of the observation to remove
|
||||
@@ -161,7 +150,6 @@ async def run_consolidation_job(
|
||||
memory_engine: "MemoryEngine",
|
||||
bank_id: str,
|
||||
request_context: "RequestContext",
|
||||
operation_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Run consolidation job for a bank.
|
||||
@@ -387,13 +375,6 @@ async def run_consolidation_job(
|
||||
[(m["id"],) for m in llm_batch],
|
||||
)
|
||||
|
||||
# Checkpoint: abort if the operation (and thus the bank) was deleted mid-run.
|
||||
if operation_id and not await memory_engine._check_op_alive(operation_id):
|
||||
logger.info(
|
||||
f"[CONSOLIDATION] bank={bank_id} operation {operation_id} cancelled (bank deleted), stopping early"
|
||||
)
|
||||
return {"status": "cancelled", "bank_id": bank_id, **stats}
|
||||
|
||||
for result in results:
|
||||
stats["memories_processed"] += 1
|
||||
action = result.get("action")
|
||||
@@ -785,17 +766,13 @@ async def _execute_update_action(
|
||||
logger.debug(f"Update skipped: observation {observation_id} not found in recall results")
|
||||
return
|
||||
|
||||
from ...config import get_config
|
||||
|
||||
history_entry = {
|
||||
"previous_text": model.text,
|
||||
"previous_tags": list(model.tags or []),
|
||||
"previous_occurred_start": model.occurred_start,
|
||||
"previous_occurred_end": model.occurred_end,
|
||||
"previous_mentioned_at": model.mentioned_at,
|
||||
"changed_at": datetime.now(timezone.utc).isoformat(),
|
||||
"new_source_memory_ids": [str(mid) for mid in source_memory_ids],
|
||||
}
|
||||
history = [
|
||||
{
|
||||
"previous_text": model.text,
|
||||
"changed_at": datetime.now(timezone.utc).isoformat(),
|
||||
"source_memory_ids": [str(mid) for mid in source_memory_ids],
|
||||
}
|
||||
]
|
||||
|
||||
source_ids = list(model.source_fact_ids or []) + source_memory_ids
|
||||
|
||||
@@ -810,18 +787,13 @@ async def _execute_update_action(
|
||||
if perf:
|
||||
perf.record_timing("embedding", time.time() - t0)
|
||||
|
||||
config = get_config()
|
||||
history_clause = (
|
||||
"history = COALESCE(history, '[]'::jsonb) || $3::jsonb," if config.enable_observation_history else ""
|
||||
)
|
||||
|
||||
t0 = time.time()
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("memory_units")}
|
||||
SET text = $1,
|
||||
embedding = $2::vector,
|
||||
{history_clause}
|
||||
history = $3,
|
||||
source_memory_ids = $4,
|
||||
proof_count = $5,
|
||||
tags = $10,
|
||||
@@ -833,7 +805,7 @@ async def _execute_update_action(
|
||||
""",
|
||||
new_text,
|
||||
embedding_str,
|
||||
json.dumps([history_entry]),
|
||||
json.dumps(history),
|
||||
source_ids,
|
||||
len(source_ids),
|
||||
uuid.UUID(observation_id),
|
||||
@@ -945,9 +917,10 @@ async def _find_related_observations(
|
||||
"""
|
||||
# Use recall to find related observations with token budget
|
||||
# max_tokens naturally limits how many observations are returned
|
||||
from ...config import get_config
|
||||
from ...tracing import get_tracer, is_tracing_enabled
|
||||
|
||||
config = await memory_engine._config_resolver.resolve_full_config(bank_id, request_context)
|
||||
config = get_config()
|
||||
|
||||
# SECURITY: Use all_strict matching if tags provided to prevent cross-scope consolidation
|
||||
tags_match = "all_strict" if tags else "any"
|
||||
@@ -972,8 +945,7 @@ async def _find_related_observations(
|
||||
tags=tags, # Filter by source memory's tags
|
||||
tags_match=tags_match, # Use strict matching for security
|
||||
include_source_facts=True, # Embed source facts so we avoid a separate DB fetch
|
||||
max_source_facts_tokens=config.consolidation_source_facts_max_tokens,
|
||||
max_source_facts_tokens_per_observation=config.consolidation_source_facts_max_tokens_per_observation,
|
||||
max_source_facts_tokens=-1, # No token limit — we need all source facts for consolidation
|
||||
_quiet=True, # Suppress logging
|
||||
)
|
||||
finally:
|
||||
@@ -1037,17 +1009,14 @@ async def _consolidate_batch_with_llm(
|
||||
observations_text = "[]"
|
||||
|
||||
def _fact_line(m: dict[str, Any]) -> str:
|
||||
text = f"[{m['id']}] {m['text']}"
|
||||
temporal_parts = []
|
||||
parts = [f"[{m['id']}] {m['text']}"]
|
||||
if m.get("occurred_start"):
|
||||
temporal_parts.append(f"occurred_start={m['occurred_start']}")
|
||||
parts.append(f"occurred_start={m['occurred_start']}")
|
||||
if m.get("occurred_end"):
|
||||
temporal_parts.append(f"occurred_end={m['occurred_end']}")
|
||||
parts.append(f"occurred_end={m['occurred_end']}")
|
||||
if m.get("mentioned_at"):
|
||||
temporal_parts.append(f"mentioned_at={m['mentioned_at']}")
|
||||
if temporal_parts:
|
||||
text += f" ({', '.join(temporal_parts)})"
|
||||
return text
|
||||
parts.append(f"mentioned_at={m['mentioned_at']}")
|
||||
return " | ".join(parts)
|
||||
|
||||
facts_lines = "\n".join(_fact_line(m) for m in memories)
|
||||
|
||||
+3
-20
@@ -29,31 +29,14 @@ Compare the facts against existing observations:
|
||||
- Same topic as an existing observation → UPDATE it (observation_id + source_fact_ids)
|
||||
- New topic with durable knowledge → CREATE a new observation (source_fact_ids)
|
||||
- Cross-reference facts within the batch: a later fact may resolve a vague reference in an earlier one
|
||||
- Purely ephemeral facts → omit them unless the MISSION above explicitly targets such data (e.g. timestamped events, session state, screen content)"""
|
||||
- Purely ephemeral facts → omit them (no create/update needed)"""
|
||||
|
||||
# Output format — JSON braces escaped as {{ }} so .format() leaves them literal
|
||||
_BATCH_OUTPUT_FORMAT = """
|
||||
Output a JSON object with three arrays.
|
||||
|
||||
## EXAMPLE
|
||||
|
||||
Input facts:
|
||||
[a1b2c3d4-e5f6-7890-abcd-ef1234567890] Alice mentioned she works long hours, often past midnight | Involving: Alice (occurred_start=2024-01-15, mentioned_at=2024-01-15)
|
||||
[b2c3d4e5-f6a7-8901-bcde-f12345678901] Alice said she's exhausted from the project deadlines | Involving: Alice (occurred_start=2024-01-20, mentioned_at=2024-01-20)
|
||||
|
||||
Good observation text — clean prose, no metadata, each fact tracked distinctly:
|
||||
"Alice works long hours, often past midnight."
|
||||
"Alice feels exhausted from project deadlines."
|
||||
|
||||
Bad observation text — NEVER do this (verbatim copy of fact text with metadata):
|
||||
"Alice mentioned she works long hours, often past midnight | Involving: Alice (occurred_start=2024-01-15, mentioned_at=2024-01-15)"
|
||||
|
||||
Observation text rules:
|
||||
- Write clean prose — NEVER copy raw fact lines or their metadata (temporal fields, "Involving:", "When:" labels, UUIDs).
|
||||
- Parenthesized metadata like (occurred_start=...) and pipe-separated labels like "| Involving: ..." are fact formatting — strip them entirely from observation text.
|
||||
- How many observations to create and how much to aggregate is driven by the MISSION above.
|
||||
|
||||
{{"creates": [{{"text": "Alice works long hours, often past midnight.", "source_fact_ids": ["a1b2c3d4-e5f6-7890-abcd-ef1234567890"]}}, {{"text": "Alice feels exhausted from project deadlines.", "source_fact_ids": ["b2c3d4e5-f6a7-8901-bcde-f12345678901"]}}],
|
||||
Example (showing the required UUID format for all IDs):
|
||||
{{"creates": [{{"text": "Alice lives in Berlin", "source_fact_ids": ["a1b2c3d4-e5f6-7890-abcd-ef1234567890", "b2c3d4e5-f6a7-8901-bcde-f12345678901"]}}],
|
||||
"updates": [{{"text": "Alice works at Acme Corp as a senior engineer", "observation_id": "c3d4e5f6-a7b8-9012-cdef-123456789012", "source_fact_ids": ["d4e5f6a7-b8c9-0123-defa-234567890123"]}}],
|
||||
"deletes": [{{"observation_id": "e5f6a7b8-c9d0-1234-efab-345678901234"}}]}}
|
||||
|
||||
+3
-189
@@ -20,10 +20,8 @@ from ..config import (
|
||||
DEFAULT_RERANKER_COHERE_MODEL,
|
||||
DEFAULT_RERANKER_FLASHRANK_CACHE_DIR,
|
||||
DEFAULT_RERANKER_FLASHRANK_MODEL,
|
||||
DEFAULT_RERANKER_LITELLM_MAX_TOKENS_PER_DOC,
|
||||
DEFAULT_RERANKER_LITELLM_MODEL,
|
||||
DEFAULT_RERANKER_LITELLM_SDK_MODEL,
|
||||
DEFAULT_RERANKER_LOCAL_BATCH_SIZE,
|
||||
DEFAULT_RERANKER_LOCAL_FORCE_CPU,
|
||||
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||
DEFAULT_RERANKER_LOCAL_MODEL,
|
||||
@@ -112,9 +110,6 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
max_concurrent: int = 4,
|
||||
force_cpu: bool = False,
|
||||
trust_remote_code: bool = False,
|
||||
fp16: bool = False,
|
||||
bucket_batching: bool = False,
|
||||
batch_size: int = DEFAULT_RERANKER_LOCAL_BATCH_SIZE,
|
||||
):
|
||||
"""
|
||||
Initialize local SentenceTransformers cross-encoder.
|
||||
@@ -129,20 +124,10 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
trust_remote_code: Allow loading models with custom code (security risk).
|
||||
Required for some models like jina-reranker-v2-base-multilingual.
|
||||
Default: False (disabled for security)
|
||||
fp16: Use FP16 (half precision) inference. Faster on MPS and CUDA,
|
||||
may be slower on CPU. Default: False (opt-in via env var).
|
||||
bucket_batching: Sort pairs by token length before batching to reduce
|
||||
padding waste. 36-54% speedup, quality-identical.
|
||||
Default: False (opt-in via env var).
|
||||
batch_size: Batch size for predict() calls. Optimal values vary by
|
||||
hardware and model (MPS: 32, CUDA: 128+). Default: 32.
|
||||
"""
|
||||
self.model_name = model_name or DEFAULT_RERANKER_LOCAL_MODEL
|
||||
self.force_cpu = force_cpu
|
||||
self.trust_remote_code = trust_remote_code
|
||||
self.fp16 = fp16
|
||||
self.bucket_batching = bucket_batching
|
||||
self.batch_size = batch_size
|
||||
self._model = None
|
||||
LocalSTCrossEncoder._max_concurrent = max_concurrent
|
||||
|
||||
@@ -190,24 +175,6 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to detect GPU/MPS, falling back to CPU: {e}")
|
||||
|
||||
# Patch transformers 5.x compatibility for models using XLM-RoBERTa
|
||||
# (e.g., jina-reranker-v2-base-multilingual). transformers 5.x removed
|
||||
# create_position_ids_from_input_ids as a module-level function; the custom
|
||||
# code in these models still references it. This monkey-patch restores it.
|
||||
try:
|
||||
import transformers.models.xlm_roberta.modeling_xlm_roberta as xlm_module
|
||||
from transformers.models.xlm_roberta.modeling_xlm_roberta import XLMRobertaEmbeddings
|
||||
|
||||
if not hasattr(xlm_module, "create_position_ids_from_input_ids"):
|
||||
setattr(
|
||||
xlm_module,
|
||||
"create_position_ids_from_input_ids",
|
||||
XLMRobertaEmbeddings.create_position_ids_from_input_ids,
|
||||
)
|
||||
logger.info("Reranker: applied transformers 5.x compatibility patch for XLM-RoBERTa")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Suppress verbose transformers warnings during model loading
|
||||
# This suppresses the "UNEXPECTED" warnings from CrossEncoder which are harmless
|
||||
# but look alarming to users (e.g., "embeddings.position_ids | UNEXPECTED")
|
||||
@@ -232,12 +199,6 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
# Restore original logging level
|
||||
transformers_logger.setLevel(original_level)
|
||||
|
||||
# FP16 inference: convert model weights to half precision.
|
||||
# Empirically validated: 27-36% faster on MPS, quality-identical (20/20 overlap).
|
||||
if self.fp16 and device != "cpu":
|
||||
self._model.model.half()
|
||||
logger.info("Reranker: FP16 inference enabled")
|
||||
|
||||
# Initialize shared executor (limited workers naturally limits concurrency)
|
||||
if LocalSTCrossEncoder._executor is None:
|
||||
LocalSTCrossEncoder._executor = ThreadPoolExecutor(
|
||||
@@ -249,32 +210,8 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
logger.info("Reranker: local provider initialized (using existing executor)")
|
||||
|
||||
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""Synchronous prediction wrapper for thread pool execution.
|
||||
|
||||
Supports two optimizations (controlled via .env):
|
||||
- bucket_batching: sort pairs by token length to reduce padding waste (36-54% speedup)
|
||||
- batch_size: explicit batch size for predict() calls (MPS optimal: 32)
|
||||
"""
|
||||
import numpy as np
|
||||
|
||||
if self.bucket_batching and len(pairs) > 1:
|
||||
# Sort pairs by approximate token length to create homogeneous batches.
|
||||
# This eliminates padding waste — short pairs aren't padded to the length
|
||||
# of the longest pair in the batch. Quality-identical by construction.
|
||||
lengths = [len(pairs[i][0]) + len(pairs[i][1]) for i in range(len(pairs))]
|
||||
sorted_indices = sorted(range(len(pairs)), key=lambda i: lengths[i])
|
||||
sorted_pairs = [pairs[i] for i in sorted_indices]
|
||||
|
||||
sorted_scores = self._model.predict(sorted_pairs, batch_size=self.batch_size, show_progress_bar=False)
|
||||
sorted_scores = sorted_scores.tolist() if hasattr(sorted_scores, "tolist") else list(sorted_scores)
|
||||
|
||||
# Restore original order
|
||||
scores = [0.0] * len(pairs)
|
||||
for new_pos, orig_idx in enumerate(sorted_indices):
|
||||
scores[orig_idx] = sorted_scores[new_pos]
|
||||
return scores
|
||||
|
||||
scores = self._model.predict(pairs, batch_size=self.batch_size, show_progress_bar=False)
|
||||
"""Synchronous prediction wrapper for thread pool execution."""
|
||||
scores = self._model.predict(pairs, show_progress_bar=False)
|
||||
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
@@ -883,17 +820,6 @@ class FlashRankCrossEncoder(CrossEncoderModel):
|
||||
return await loop.run_in_executor(FlashRankCrossEncoder._executor, self._predict_sync, pairs)
|
||||
|
||||
|
||||
def _truncate_to_tokens(text: str, max_tokens: int) -> str:
|
||||
"""Truncate text to at most max_tokens using the shared tiktoken encoder."""
|
||||
from .memory_engine import _get_tiktoken_encoding
|
||||
|
||||
enc = _get_tiktoken_encoding()
|
||||
tokens = enc.encode(text)
|
||||
if len(tokens) <= max_tokens:
|
||||
return text
|
||||
return enc.decode(tokens[:max_tokens])
|
||||
|
||||
|
||||
class LiteLLMCrossEncoder(CrossEncoderModel):
|
||||
"""
|
||||
LiteLLM cross-encoder implementation using LiteLLM proxy's /rerank endpoint.
|
||||
@@ -917,7 +843,6 @@ class LiteLLMCrossEncoder(CrossEncoderModel):
|
||||
api_key: str | None = None,
|
||||
model: str = DEFAULT_RERANKER_LITELLM_MODEL,
|
||||
timeout: float = 60.0,
|
||||
max_tokens_per_doc: int | None = DEFAULT_RERANKER_LITELLM_MAX_TOKENS_PER_DOC,
|
||||
):
|
||||
"""
|
||||
Initialize LiteLLM cross-encoder client.
|
||||
@@ -928,15 +853,11 @@ class LiteLLMCrossEncoder(CrossEncoderModel):
|
||||
model: Reranking model name (default: cohere/rerank-english-v3.0)
|
||||
Use provider prefix (e.g., cohere/, together_ai/, voyage/)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
max_tokens_per_doc: If set, truncate each document to this many tokens before
|
||||
sending to the reranker (uses tiktoken cl100k_base encoding).
|
||||
Useful for models with small context windows (e.g. 1024 tokens).
|
||||
"""
|
||||
self.api_base = api_base.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self.max_tokens_per_doc = max_tokens_per_doc
|
||||
self._async_client: httpx.AsyncClient | None = None
|
||||
|
||||
@property
|
||||
@@ -984,8 +905,6 @@ class LiteLLMCrossEncoder(CrossEncoderModel):
|
||||
|
||||
for query, indexed_texts in query_groups.items():
|
||||
texts = [text for _, text in indexed_texts]
|
||||
if self.max_tokens_per_doc is not None:
|
||||
texts = [_truncate_to_tokens(t, self.max_tokens_per_doc) for t in texts]
|
||||
indices = [idx for idx, _ in indexed_texts]
|
||||
|
||||
# LiteLLM /rerank follows Cohere API format
|
||||
@@ -1031,7 +950,6 @@ class LiteLLMSDKCrossEncoder(CrossEncoderModel):
|
||||
model: str = DEFAULT_RERANKER_LITELLM_SDK_MODEL,
|
||||
api_base: str | None = None,
|
||||
timeout: float = 60.0,
|
||||
max_tokens_per_doc: int | None = DEFAULT_RERANKER_LITELLM_MAX_TOKENS_PER_DOC,
|
||||
):
|
||||
"""
|
||||
Initialize LiteLLM SDK cross-encoder client.
|
||||
@@ -1041,15 +959,11 @@ class LiteLLMSDKCrossEncoder(CrossEncoderModel):
|
||||
model: Model name with provider prefix (e.g., "deepinfra/Qwen3-reranker-8B")
|
||||
api_base: Custom base URL for API (optional)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
max_tokens_per_doc: If set, truncate each document to this many tokens before
|
||||
sending to the reranker (uses tiktoken cl100k_base encoding).
|
||||
Useful for models with small context windows (e.g. 1024 tokens).
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.api_base = api_base
|
||||
self.timeout = timeout
|
||||
self.max_tokens_per_doc = max_tokens_per_doc
|
||||
self._initialized = False
|
||||
self._litellm = None # Will be set during initialization
|
||||
|
||||
@@ -1103,8 +1017,6 @@ class LiteLLMSDKCrossEncoder(CrossEncoderModel):
|
||||
|
||||
for query, indexed_texts in query_groups.items():
|
||||
texts = [text for _, text in indexed_texts]
|
||||
if self.max_tokens_per_doc is not None:
|
||||
texts = [_truncate_to_tokens(t, self.max_tokens_per_doc) for t in texts]
|
||||
indices = [idx for idx, _ in indexed_texts]
|
||||
|
||||
# Build kwargs for rerank call
|
||||
@@ -1138,97 +1050,6 @@ class LiteLLMSDKCrossEncoder(CrossEncoderModel):
|
||||
return all_scores
|
||||
|
||||
|
||||
class JinaMLXCrossEncoder(CrossEncoderModel):
|
||||
"""
|
||||
Jina Reranker v3 MLX implementation for Apple Silicon.
|
||||
|
||||
Uses jinaai/jina-reranker-v3-mlx — a 0.6B parameter multilingual listwise reranker
|
||||
optimized for Apple Silicon via the MLX framework. No transformers/PyTorch dependency.
|
||||
|
||||
The model is downloaded automatically from HuggingFace Hub on first use.
|
||||
Requires: mlx>=0.31.0, mlx-lm>=0.31.1, safetensors>=0.6.2
|
||||
"""
|
||||
|
||||
HF_REPO_ID = "jinaai/jina-reranker-v3-mlx"
|
||||
|
||||
def __init__(self, model_path: str | None = None):
|
||||
"""
|
||||
Args:
|
||||
model_path: Local path to the downloaded model directory.
|
||||
If None, the model is downloaded from HuggingFace Hub.
|
||||
"""
|
||||
self.model_path = model_path
|
||||
self._reranker = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "jina-mlx"
|
||||
|
||||
async def initialize(self) -> None:
|
||||
if self._reranker is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
import mlx.core # noqa: F401
|
||||
import mlx_lm # noqa: F401
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"mlx and mlx-lm are required for JinaMLXCrossEncoder. "
|
||||
"Install with: pip install mlx>=0.31.0 mlx-lm>=0.31.1 safetensors>=0.6.2"
|
||||
)
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
await loop.run_in_executor(None, self._load_model)
|
||||
|
||||
def _load_model(self) -> None:
|
||||
"""Download (if needed) and load the MLX reranker. Runs in a thread."""
|
||||
import os
|
||||
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
from .jina_mlx_reranker import MLXReranker
|
||||
|
||||
model_path = self.model_path
|
||||
if model_path is None:
|
||||
logger.info(f"Reranker: downloading {self.HF_REPO_ID} from HuggingFace Hub...")
|
||||
model_path = snapshot_download(repo_id=self.HF_REPO_ID)
|
||||
|
||||
logger.info(f"Reranker: loading jina-reranker-v3-mlx from {model_path}")
|
||||
self._reranker = MLXReranker(
|
||||
model_path=model_path,
|
||||
projector_path=os.path.join(model_path, "projector.safetensors"),
|
||||
)
|
||||
logger.info("Reranker: jina-mlx provider initialized")
|
||||
|
||||
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""Score pairs grouped by query. Runs in a thread."""
|
||||
if not pairs:
|
||||
return []
|
||||
|
||||
query_groups: dict[str, list[tuple[int, str]]] = {}
|
||||
for idx, (query, doc) in enumerate(pairs):
|
||||
query_groups.setdefault(query, []).append((idx, doc))
|
||||
|
||||
all_scores = [0.0] * len(pairs)
|
||||
|
||||
for query, indexed_docs in query_groups.items():
|
||||
docs = [doc for _, doc in indexed_docs]
|
||||
indices = [idx for idx, _ in indexed_docs]
|
||||
results = self._reranker.rerank(query, docs)
|
||||
for result in results:
|
||||
original_idx = result["index"]
|
||||
all_scores[indices[original_idx]] = result["relevance_score"]
|
||||
|
||||
return all_scores
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
if self._reranker is None:
|
||||
raise RuntimeError("Reranker not initialized. Call initialize() first.")
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
return await loop.run_in_executor(None, self._predict_sync, pairs)
|
||||
|
||||
|
||||
def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
"""
|
||||
Create a CrossEncoderModel instance based on configuration.
|
||||
@@ -1258,9 +1079,6 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
max_concurrent=config.reranker_local_max_concurrent,
|
||||
force_cpu=config.reranker_local_force_cpu,
|
||||
trust_remote_code=config.reranker_local_trust_remote_code,
|
||||
fp16=config.reranker_local_fp16,
|
||||
bucket_batching=config.reranker_local_bucket_batching,
|
||||
batch_size=config.reranker_local_batch_size,
|
||||
)
|
||||
elif provider == "cohere":
|
||||
api_key = config.reranker_cohere_api_key
|
||||
@@ -1280,7 +1098,6 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
api_base=config.reranker_litellm_api_base,
|
||||
api_key=config.reranker_litellm_api_key,
|
||||
model=config.reranker_litellm_model,
|
||||
max_tokens_per_doc=config.reranker_litellm_max_tokens_per_doc,
|
||||
)
|
||||
elif provider == "litellm-sdk":
|
||||
api_key = config.reranker_litellm_sdk_api_key
|
||||
@@ -1292,7 +1109,6 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
api_key=api_key,
|
||||
model=config.reranker_litellm_sdk_model,
|
||||
api_base=config.reranker_litellm_sdk_api_base,
|
||||
max_tokens_per_doc=config.reranker_litellm_max_tokens_per_doc,
|
||||
)
|
||||
elif provider == "zeroentropy":
|
||||
api_key = config.reranker_zeroentropy_api_key
|
||||
@@ -1306,9 +1122,7 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
)
|
||||
elif provider == "rrf":
|
||||
return RRFPassthroughCrossEncoder()
|
||||
elif provider == "jina-mlx":
|
||||
return JinaMLXCrossEncoder()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'zeroentropy', 'flashrank', 'litellm', 'litellm-sdk', 'rrf', 'jina-mlx'"
|
||||
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'zeroentropy', 'flashrank', 'litellm', 'litellm-sdk', 'rrf'"
|
||||
)
|
||||
+4
-10
@@ -58,16 +58,10 @@ async def retry_with_backoff(
|
||||
last_exception = e
|
||||
if attempt < max_retries:
|
||||
delay = min(base_delay * (2**attempt), max_delay)
|
||||
if isinstance(e, asyncpg.exceptions.DeadlockDetectedError):
|
||||
logger.warning(
|
||||
f"Deadlock detected during parallel document processing — this is expected and will resolve automatically "
|
||||
f"(attempt {attempt + 1}/{max_retries + 1}, retrying in {delay:.1f}s)"
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"Database operation failed (attempt {attempt + 1}/{max_retries + 1}): {e}. "
|
||||
f"Retrying in {delay:.1f}s..."
|
||||
)
|
||||
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}")
|
||||
+3
-7
@@ -459,12 +459,10 @@ class EntityResolver:
|
||||
entity_dates = [g.event_date for _, g in sorted_groups]
|
||||
|
||||
# INSERT ... ON CONFLICT DO NOTHING — no row lock on already-existing entities.
|
||||
# mention_count starts at 0 here; flush_pending_stats() is the sole source of
|
||||
# truth for mention counting (one stat per original mention in the batch).
|
||||
inserted_rows = await conn.fetch(
|
||||
f"""
|
||||
INSERT INTO {fq_table("entities")} (bank_id, canonical_name, first_seen, last_seen, mention_count)
|
||||
SELECT $1, name, COALESCE(event_date, now()), COALESCE(event_date, now()), 0
|
||||
SELECT $1, name, COALESCE(event_date, now()), COALESCE(event_date, now()), 1
|
||||
FROM unnest($2::text[], $3::timestamptz[]) AS t(name, event_date)
|
||||
ON CONFLICT (bank_id, LOWER(canonical_name))
|
||||
DO NOTHING
|
||||
@@ -491,15 +489,13 @@ class EntityResolver:
|
||||
for row in existing_rows:
|
||||
id_by_name[row["name_lower"]] = row["id"]
|
||||
|
||||
# Assign entity IDs back and queue one stat per original mention so that
|
||||
# flush_pending_stats() increments mention_count by the true mention count,
|
||||
# not just 1 per unique name.
|
||||
# Assign entity IDs back and queue for post-txn stats flush.
|
||||
for name_lower, g in sorted_groups:
|
||||
entity_id = id_by_name.get(name_lower)
|
||||
if entity_id:
|
||||
for original_idx in g.indices:
|
||||
entity_ids[original_idx] = entity_id
|
||||
pending.append(_EntityStat(entity_id=entity_id, event_date=g.event_date))
|
||||
pending.append(_EntityStat(entity_id=entity_id, event_date=g.event_date))
|
||||
|
||||
# Accumulate into the resolver's pending list; the orchestrator flushes
|
||||
# these with await entity_resolver.flush_pending_stats() after the txn.
|
||||
+1
-26
@@ -48,28 +48,6 @@ _llm_max_concurrent = int(os.getenv(ENV_LLM_MAX_CONCURRENT, str(DEFAULT_LLM_MAX_
|
||||
_global_llm_semaphore = asyncio.Semaphore(_llm_max_concurrent)
|
||||
|
||||
|
||||
def sanitize_llm_output(text: str | None) -> str | None:
|
||||
"""
|
||||
Sanitize text by removing characters that break downstream systems.
|
||||
|
||||
Removes:
|
||||
- ASCII control characters (0x00-0x08, 0x0B-0x0C, 0x0E-0x1F, 0x7F): break
|
||||
json.loads and PostgreSQL UTF-8 encoding; tab (0x09), newline (0x0A), and
|
||||
carriage return (0x0D) are preserved as they are valid in text and JSON.
|
||||
- Unicode surrogates (U+D800-U+DFFF): Invalid in UTF-8, break LLM APIs
|
||||
|
||||
Surrogate characters are used in UTF-16 encoding but cannot be encoded
|
||||
in UTF-8. They can appear in Python strings from improperly decoded data
|
||||
(e.g., from JavaScript or broken files). Control characters commonly appear
|
||||
in LLM output embedded inside JSON string values.
|
||||
"""
|
||||
if text is None:
|
||||
return None
|
||||
if not text:
|
||||
return text
|
||||
return re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f\x7f\ud800-\udfff]", "", text)
|
||||
|
||||
|
||||
class OutputTooLongError(Exception):
|
||||
"""
|
||||
Bridge exception raised when LLM output exceeds token limits.
|
||||
@@ -227,7 +205,7 @@ def create_llm_provider(
|
||||
reasoning_effort=reasoning_effort,
|
||||
)
|
||||
|
||||
elif provider_lower in ("openai", "groq", "ollama", "lmstudio", "minimax"):
|
||||
elif provider_lower in ("openai", "groq", "ollama", "lmstudio"):
|
||||
return OpenAICompatibleLLM(
|
||||
provider=provider,
|
||||
api_key=api_key,
|
||||
@@ -296,7 +274,6 @@ class LLMProvider:
|
||||
"openai-codex",
|
||||
"claude-code",
|
||||
"mock",
|
||||
"minimax",
|
||||
]
|
||||
if self.provider not in valid_providers:
|
||||
raise ValueError(f"Invalid LLM provider: {self.provider}. Must be one of: {', '.join(valid_providers)}")
|
||||
@@ -309,8 +286,6 @@ class LLMProvider:
|
||||
self.base_url = "http://localhost:11434/v1"
|
||||
elif self.provider == "lmstudio":
|
||||
self.base_url = "http://localhost:1234/v1"
|
||||
elif self.provider == "minimax":
|
||||
self.base_url = "https://api.minimax.io/v1"
|
||||
|
||||
# Prepare Vertex AI config (if applicable)
|
||||
vertexai_project_id = None
|
||||
+38
-511
@@ -16,7 +16,7 @@ import logging
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from datetime import UTC, datetime, timedelta, timezone
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import asyncpg
|
||||
@@ -51,9 +51,13 @@ def get_current_schema() -> str:
|
||||
return schema
|
||||
|
||||
|
||||
# Initialize tiktoken encoder once at module level for efficiency
|
||||
_tiktoken_encoder = tiktoken.get_encoding("cl100k_base") # GPT-4/GPT-3.5-turbo encoding
|
||||
|
||||
|
||||
def count_tokens(text: str) -> int:
|
||||
"""Count tokens in text using tiktoken (cl100k_base encoding for GPT-4/3.5)."""
|
||||
return len(_get_tiktoken_encoding().encode(text))
|
||||
return len(_tiktoken_encoder.encode(text))
|
||||
|
||||
|
||||
def fq_table(table_name: str) -> str:
|
||||
@@ -164,7 +168,7 @@ from enum import Enum
|
||||
from ..metrics import get_metrics_collector
|
||||
from ..pg0 import EmbeddedPostgres, parse_pg0_url
|
||||
from .entity_resolver import EntityResolver
|
||||
from .llm_wrapper import LLMConfig, requires_api_key, sanitize_llm_output
|
||||
from .llm_wrapper import LLMConfig, requires_api_key
|
||||
from .query_analyzer import QueryAnalyzer
|
||||
from .reflect import run_reflect_agent
|
||||
from .reflect.tools import tool_expand, tool_recall, tool_search_mental_models, tool_search_observations
|
||||
@@ -184,7 +188,7 @@ from .retain import bank_utils, embedding_utils
|
||||
from .retain.types import RetainContentDict
|
||||
from .search import think_utils
|
||||
from .search.reranking import CrossEncoderReranker, apply_combined_scoring
|
||||
from .search.tags import TagGroup, TagsMatch, build_tags_where_clause
|
||||
from .search.tags import TagsMatch, build_tags_where_clause
|
||||
from .task_backend import BrokerTaskBackend, SyncTaskBackend, TaskBackend
|
||||
|
||||
|
||||
@@ -204,6 +208,8 @@ def utcnow():
|
||||
# Logger for memory system
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
import tiktoken
|
||||
|
||||
from .db_utils import acquire_with_retry
|
||||
|
||||
# Cache tiktoken encoding for token budget filtering (module-level singleton)
|
||||
@@ -647,19 +653,13 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
# Retrieve file from storage
|
||||
file_data = await self._file_storage.retrieve(storage_key)
|
||||
|
||||
# Convert to markdown using the ordered fallback chain stored in the task payload.
|
||||
# task_dict["parser"] is always a list[str] set at submission time.
|
||||
parser_chain: list[str] = task_dict.get("parser") or []
|
||||
if not parser_chain:
|
||||
raise ValueError("No parser chain defined for file_convert_retain task")
|
||||
convert_result = await self._parser_registry.convert_with_fallback(
|
||||
parsers=parser_chain,
|
||||
file_data=file_data,
|
||||
# Convert to markdown
|
||||
parser = self._parser_registry.get_parser(
|
||||
name=task_dict.get("parser"),
|
||||
filename=filename,
|
||||
content_type=task_dict.get("content_type"),
|
||||
)
|
||||
markdown_content = sanitize_llm_output(convert_result.content) or ""
|
||||
winning_parser = convert_result.parser_name
|
||||
markdown_content = await parser.convert(file_data, filename)
|
||||
except Exception as e:
|
||||
# Re-raise with filename context for better error reporting
|
||||
error_msg = f"Failed to parse file '{filename}': {str(e)}"
|
||||
@@ -671,31 +671,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
f"document_id={document_id}, {len(markdown_content)} chars. Submitting retain task."
|
||||
)
|
||||
|
||||
# Fire file conversion hook (e.g., for Iris billing)
|
||||
if self._operation_validator:
|
||||
try:
|
||||
from hindsight_api.extensions.operation_validator import FileConvertResult
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
convert_context = RequestContext(
|
||||
internal=True,
|
||||
user_initiated=True,
|
||||
tenant_id=task_dict.get("_tenant_id"),
|
||||
api_key_id=task_dict.get("_api_key_id"),
|
||||
)
|
||||
await self._operation_validator.on_file_convert_complete(
|
||||
FileConvertResult(
|
||||
bank_id=bank_id,
|
||||
parser_name=winning_parser,
|
||||
filename=filename,
|
||||
output_chars=len(markdown_content),
|
||||
output_text=markdown_content,
|
||||
request_context=convert_context,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"[FILE_CONVERT_RETAIN] on_file_convert_complete hook failed: {e}")
|
||||
|
||||
# Build retain task payload
|
||||
retain_contents = [
|
||||
{
|
||||
@@ -814,7 +789,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
memory_engine=self,
|
||||
bank_id=bank_id,
|
||||
request_context=internal_context,
|
||||
operation_id=task_dict.get("operation_id"),
|
||||
)
|
||||
|
||||
logger.info(f"[CONSOLIDATION] bank={bank_id} completed: {result.get('memories_processed', 0)} processed")
|
||||
@@ -1134,13 +1108,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
)
|
||||
|
||||
async def _callback(conn: asyncpg.Connection) -> None:
|
||||
# Resolve schema at call time (not at callback creation time) because
|
||||
# _current_schema contextvar may not yet be set when the callback is built
|
||||
# from the HTTP path (http.py calls _build_retain_outbox_callback before
|
||||
# retain_batch_async which is where _authenticate_tenant sets the schema).
|
||||
resolved_schema = schema or _current_schema.get()
|
||||
for event in events:
|
||||
await webhook_manager.fire_event_with_conn(event, conn, schema=resolved_schema)
|
||||
await webhook_manager.fire_event_with_conn(event, conn, schema=schema)
|
||||
|
||||
return _callback
|
||||
|
||||
@@ -1244,24 +1213,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to delete async operation record {operation_id}: {e}")
|
||||
|
||||
async def _check_op_alive(self, operation_id: str) -> bool:
|
||||
"""Return False if the operation row no longer exists (e.g. bank was deleted via CASCADE).
|
||||
|
||||
Long-running operations should call this at natural checkpoints (e.g. after each
|
||||
committed batch) to detect bank deletion early and abort cleanly.
|
||||
"""
|
||||
try:
|
||||
pool = await self._get_pool()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
row = await conn.fetchrow(
|
||||
f"SELECT operation_id FROM {fq_table('async_operations')} WHERE operation_id = $1",
|
||||
uuid.UUID(operation_id),
|
||||
)
|
||||
return row is not None
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to check operation liveness {operation_id}: {e}")
|
||||
return True # Assume alive on DB error to avoid false-positive aborts
|
||||
|
||||
async def _mark_operation_failed(self, operation_id: str, error_message: str, error_traceback: str):
|
||||
"""Helper to mark an operation as failed in the database.
|
||||
|
||||
@@ -1277,19 +1228,15 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
# Mark this operation as failed
|
||||
row = await conn.fetchrow(
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("async_operations")}
|
||||
SET status = 'failed', error_message = $2, updated_at = NOW()
|
||||
WHERE operation_id = $1
|
||||
RETURNING operation_id
|
||||
""",
|
||||
uuid.UUID(operation_id),
|
||||
truncated_error,
|
||||
)
|
||||
if row is None:
|
||||
logger.info(f"Operation {operation_id} no longer exists (bank deleted), skipping mark-failed")
|
||||
return
|
||||
logger.info(f"Marked async operation as failed: {operation_id}")
|
||||
|
||||
# Check if this is a child operation and update parent if all siblings are done
|
||||
@@ -1309,20 +1256,14 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
# Mark this operation as completed
|
||||
row = await conn.fetchrow(
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("async_operations")}
|
||||
SET status = 'completed', updated_at = NOW(), completed_at = NOW()
|
||||
WHERE operation_id = $1
|
||||
RETURNING operation_id
|
||||
""",
|
||||
uuid.UUID(operation_id),
|
||||
)
|
||||
if row is None:
|
||||
logger.info(
|
||||
f"Operation {operation_id} no longer exists (bank deleted), skipping mark-completed"
|
||||
)
|
||||
return
|
||||
logger.info(f"Marked async operation as completed: {operation_id}")
|
||||
|
||||
# Check if this is a child operation and update parent if all siblings are done
|
||||
@@ -1352,20 +1293,14 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
pool = await self._get_pool()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
row = await conn.fetchrow(
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("async_operations")}
|
||||
SET status = 'completed', updated_at = NOW(), completed_at = NOW()
|
||||
WHERE operation_id = $1
|
||||
RETURNING operation_id
|
||||
""",
|
||||
uuid.UUID(operation_id),
|
||||
)
|
||||
if row is None:
|
||||
logger.info(
|
||||
f"Operation {operation_id} no longer exists (bank deleted), skipping mark-completed"
|
||||
)
|
||||
return
|
||||
logger.info(f"Marked async operation as completed: {operation_id}")
|
||||
await self._maybe_update_parent_operation(operation_id, conn)
|
||||
|
||||
@@ -1663,15 +1598,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
# Create connection pool
|
||||
# For read-heavy workloads with many parallel think/search operations,
|
||||
# we need a larger pool. Read operations don't need strong isolation.
|
||||
async def _init_connection(conn: asyncpg.Connection) -> None:
|
||||
# SET (not SET LOCAL) so it persists for the connection lifetime.
|
||||
# ef_search=200 improves HNSW recall quality for the per-fact_type
|
||||
# semantic queries in retrieve_semantic_bm25_combined().
|
||||
try:
|
||||
await conn.execute("SET hnsw.ef_search = 200")
|
||||
except Exception:
|
||||
logger.debug("Could not set hnsw.ef_search — extension may not support it")
|
||||
|
||||
self._pool = await asyncpg.create_pool(
|
||||
self.db_url,
|
||||
min_size=self._pool_min_size,
|
||||
@@ -1679,7 +1605,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
command_timeout=self._db_command_timeout,
|
||||
statement_cache_size=0, # Disable prepared statement cache
|
||||
timeout=self._db_acquire_timeout, # Connection acquisition timeout (seconds)
|
||||
init=_init_connection,
|
||||
)
|
||||
|
||||
# Initialize entity resolver with pool and configured lookup strategy
|
||||
@@ -2086,15 +2011,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
# Process each sub-batch
|
||||
all_results = []
|
||||
for i, sub_batch in enumerate(sub_batches, 1):
|
||||
# Checkpoint: abort if the operation was deleted (bank was deleted) between sub-batches.
|
||||
if operation_id and not await self._check_op_alive(operation_id):
|
||||
logger.info(
|
||||
f"[BATCH_RETAIN] bank={bank_id} operation {operation_id} cancelled (bank deleted), stopping after {i - 1}/{len(sub_batches)} sub-batches"
|
||||
)
|
||||
if return_usage:
|
||||
return all_results, total_usage
|
||||
return all_results
|
||||
|
||||
sub_batch_tokens = sum(count_tokens(item.get("content", "")) for item in sub_batch)
|
||||
logger.info(
|
||||
f"Processing sub-batch {i}/{len(sub_batches)}: {len(sub_batch)} items, {sub_batch_tokens:,} tokens"
|
||||
@@ -2235,7 +2151,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
document_tags=document_tags,
|
||||
config=resolved_config,
|
||||
operation_id=operation_id,
|
||||
schema=_current_schema.get(),
|
||||
schema=request_context.tenant_id if request_context else None,
|
||||
outbox_callback=outbox_callback,
|
||||
)
|
||||
|
||||
@@ -2296,11 +2212,9 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_chunk_tokens: int = 8192,
|
||||
include_source_facts: bool = False,
|
||||
max_source_facts_tokens: int = 4096,
|
||||
max_source_facts_tokens_per_observation: int = -1,
|
||||
request_context: "RequestContext",
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
tag_groups: list[TagGroup] | None = None,
|
||||
_connection_budget: int | None = None,
|
||||
_quiet: bool = False,
|
||||
) -> RecallResultModel:
|
||||
@@ -2435,12 +2349,10 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
semaphore_wait=semaphore_wait,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
tag_groups=tag_groups,
|
||||
connection_budget=_connection_budget,
|
||||
quiet=_quiet,
|
||||
include_source_facts=include_source_facts,
|
||||
max_source_facts_tokens=max_source_facts_tokens,
|
||||
max_source_facts_tokens_per_observation=max_source_facts_tokens_per_observation,
|
||||
)
|
||||
break # Success - exit retry loop
|
||||
except Exception as e:
|
||||
@@ -2563,12 +2475,10 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
semaphore_wait: float = 0.0,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
tag_groups: list[TagGroup] | None = None,
|
||||
connection_budget: int | None = None,
|
||||
quiet: bool = False,
|
||||
include_source_facts: bool = False,
|
||||
max_source_facts_tokens: int = 4096,
|
||||
max_source_facts_tokens_per_observation: int = -1,
|
||||
) -> RecallResultModel:
|
||||
"""
|
||||
Search implementation with modular retrieval and reranking.
|
||||
@@ -2683,7 +2593,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
self.query_analyzer,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
tag_groups=tag_groups,
|
||||
)
|
||||
parallel_duration = time.time() - parallel_start
|
||||
finally:
|
||||
@@ -3191,9 +3100,18 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
|
||||
encoding = _get_tiktoken_encoding()
|
||||
source_facts_dict = {}
|
||||
|
||||
def _make_source_fact(sid: str, r: Any) -> MemoryFact:
|
||||
return MemoryFact(
|
||||
total_source_tokens = 0
|
||||
for sid in source_ids_ordered:
|
||||
if sid not in source_row_by_id:
|
||||
continue
|
||||
r = source_row_by_id[sid]
|
||||
fact_tokens = len(encoding.encode(r["text"]))
|
||||
if (
|
||||
max_source_facts_tokens >= 0
|
||||
and total_source_tokens + fact_tokens > max_source_facts_tokens
|
||||
):
|
||||
break
|
||||
source_facts_dict[sid] = MemoryFact(
|
||||
id=sid,
|
||||
text=r["text"],
|
||||
fact_type=r["fact_type"],
|
||||
@@ -3205,37 +3123,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
chunk_id=str(r["chunk_id"]) if r["chunk_id"] else None,
|
||||
tags=r["tags"] or None,
|
||||
)
|
||||
|
||||
if max_source_facts_tokens_per_observation >= 0:
|
||||
# Per-observation capping: each observation independently selects
|
||||
# source facts up to its token budget.
|
||||
for obs_id, sids in source_fact_ids_by_obs.items():
|
||||
obs_tokens = 0
|
||||
for sid in sids:
|
||||
if sid not in source_row_by_id:
|
||||
continue
|
||||
r = source_row_by_id[sid]
|
||||
fact_tokens = len(encoding.encode(r["text"]))
|
||||
if obs_tokens + fact_tokens > max_source_facts_tokens_per_observation:
|
||||
break
|
||||
obs_tokens += fact_tokens
|
||||
if sid not in source_facts_dict:
|
||||
source_facts_dict[sid] = _make_source_fact(sid, r)
|
||||
else:
|
||||
# Global budget: fill in order of first appearance until exhausted.
|
||||
total_source_tokens = 0
|
||||
for sid in source_ids_ordered:
|
||||
if sid not in source_row_by_id:
|
||||
continue
|
||||
r = source_row_by_id[sid]
|
||||
fact_tokens = len(encoding.encode(r["text"]))
|
||||
if (
|
||||
max_source_facts_tokens >= 0
|
||||
and total_source_tokens + fact_tokens > max_source_facts_tokens
|
||||
):
|
||||
break
|
||||
source_facts_dict[sid] = _make_source_fact(sid, r)
|
||||
total_source_tokens += fact_tokens
|
||||
total_source_tokens += fact_tokens
|
||||
|
||||
# Get entities for each fact if include_entities is requested
|
||||
fact_entity_map = {} # unit_id -> list of (entity_id, entity_name)
|
||||
@@ -3497,140 +3385,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
|
||||
return result
|
||||
|
||||
async def update_document(
|
||||
self,
|
||||
document_id: str,
|
||||
bank_id: str,
|
||||
*,
|
||||
tags: list[str] | None = None,
|
||||
request_context: "RequestContext",
|
||||
) -> bool:
|
||||
"""
|
||||
Update mutable fields on a document without re-processing its content.
|
||||
|
||||
Tag changes propagate to all associated memory units and trigger observation
|
||||
invalidation + re-consolidation (same semantics as delete_document):
|
||||
- Observations referencing the document's memory units are deleted.
|
||||
- The document's own units and any co-source memories from other documents
|
||||
have consolidated_at reset so they are re-consolidated under the new tags.
|
||||
|
||||
Args:
|
||||
document_id: Document ID to update
|
||||
bank_id: Bank ID that owns the document
|
||||
tags: New tags to apply to the document and all its memory units (optional)
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
True if the document was found and updated, False if not found
|
||||
"""
|
||||
await self._authenticate_tenant(request_context)
|
||||
if self._operation_validator:
|
||||
from hindsight_api.extensions import BankWriteContext
|
||||
|
||||
ctx = BankWriteContext(bank_id=bank_id, operation="update_document", request_context=request_context)
|
||||
await self._validate_operation(self._operation_validator.validate_bank_write(ctx))
|
||||
pool = await self._get_pool()
|
||||
invalidated_obs = 0
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
async with conn.transaction():
|
||||
set_parts: list[str] = ["updated_at = now()"]
|
||||
params: list[Any] = []
|
||||
p = 1
|
||||
|
||||
if tags is not None:
|
||||
set_parts.append(f"tags = ${p}")
|
||||
params.append(tags)
|
||||
p += 1
|
||||
|
||||
params.extend([document_id, bank_id])
|
||||
doc_id_found = await conn.fetchval(
|
||||
f"""
|
||||
UPDATE {fq_table("documents")}
|
||||
SET {", ".join(set_parts)}
|
||||
WHERE id = ${p} AND bank_id = ${p + 1}
|
||||
RETURNING id
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
if not doc_id_found:
|
||||
return False
|
||||
|
||||
if tags is not None:
|
||||
unit_rows = await conn.fetch(
|
||||
f"SELECT id FROM {fq_table('memory_units')} WHERE document_id = $1 AND fact_type IN ('experience', 'world')",
|
||||
document_id,
|
||||
)
|
||||
unit_ids = [str(row["id"]) for row in unit_rows]
|
||||
|
||||
await conn.execute(
|
||||
f"UPDATE {fq_table('memory_units')} SET tags = $1 WHERE document_id = $2",
|
||||
tags,
|
||||
document_id,
|
||||
)
|
||||
|
||||
if unit_ids:
|
||||
import uuid as uuid_module
|
||||
|
||||
unit_uuids = [uuid_module.UUID(uid) for uid in unit_ids]
|
||||
unit_uuid_set = {str(u) for u in unit_uuids}
|
||||
affected_obs = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, source_memory_ids FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $1
|
||||
AND fact_type = 'observation'
|
||||
AND source_memory_ids && $2::uuid[]
|
||||
""",
|
||||
bank_id,
|
||||
unit_uuids,
|
||||
)
|
||||
if affected_obs:
|
||||
obs_ids = [obs["id"] for obs in affected_obs]
|
||||
|
||||
seen: set[str] = set()
|
||||
other_source_uuids: list[uuid_module.UUID] = []
|
||||
for obs in affected_obs:
|
||||
for src_id in obs["source_memory_ids"] or []:
|
||||
src_str = str(src_id)
|
||||
if src_str not in unit_uuid_set and src_str not in seen:
|
||||
other_source_uuids.append(src_id)
|
||||
seen.add(src_str)
|
||||
|
||||
await conn.execute(
|
||||
f"DELETE FROM {fq_table('memory_units')} WHERE id = ANY($1::uuid[])",
|
||||
obs_ids,
|
||||
)
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("memory_units")}
|
||||
SET consolidated_at = NULL
|
||||
WHERE id = ANY($1::uuid[])
|
||||
AND fact_type IN ('experience', 'world')
|
||||
""",
|
||||
unit_uuids,
|
||||
)
|
||||
if other_source_uuids:
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("memory_units")}
|
||||
SET consolidated_at = NULL
|
||||
WHERE id = ANY($1::uuid[])
|
||||
AND fact_type IN ('experience', 'world')
|
||||
""",
|
||||
other_source_uuids,
|
||||
)
|
||||
invalidated_obs = len(obs_ids)
|
||||
logger.info(
|
||||
f"[OBSERVATIONS] Deleted {invalidated_obs} observations, reset "
|
||||
f"{len(unit_ids)} document source memories and "
|
||||
f"{len(other_source_uuids)} co-source memories for re-consolidation "
|
||||
f"after document update on '{document_id}' in bank {bank_id}"
|
||||
)
|
||||
|
||||
if invalidated_obs > 0:
|
||||
await self.submit_async_consolidation(bank_id=bank_id, request_context=request_context)
|
||||
|
||||
return True
|
||||
|
||||
async def delete_memory_unit(
|
||||
self,
|
||||
unit_id: str,
|
||||
@@ -3728,7 +3482,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
pool = await self._get_pool()
|
||||
invalidated_obs = 0
|
||||
result: dict[str, int] = {}
|
||||
bank_internal_id: str | None = None
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Ensure connection is not in read-only mode (can happen with connection poolers)
|
||||
await conn.execute("SET SESSION CHARACTERISTICS AS TRANSACTION READ WRITE")
|
||||
@@ -3784,12 +3537,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
# Delete entities (cascades to unit_entities, entity_cooccurrences, memory_links with entity_id)
|
||||
await conn.execute(f"DELETE FROM {fq_table('entities')} WHERE bank_id = $1", bank_id)
|
||||
|
||||
# Delete the bank profile and retrieve internal_id for HNSW index cleanup
|
||||
internal_id = await conn.fetchval(
|
||||
f"DELETE FROM {fq_table('banks')} WHERE bank_id = $1 RETURNING internal_id", bank_id
|
||||
)
|
||||
if internal_id:
|
||||
bank_internal_id = str(internal_id)
|
||||
# Delete the bank profile itself
|
||||
await conn.execute(f"DELETE FROM {fq_table('banks')} WHERE bank_id = $1", bank_id)
|
||||
|
||||
result = {
|
||||
"memory_units_deleted": units_count,
|
||||
@@ -3801,12 +3550,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to delete agent data: {str(e)}")
|
||||
|
||||
# Drop per-bank HNSW indexes AFTER the transaction commits to avoid
|
||||
# AccessExclusiveLock deadlocks with concurrent bank deletions.
|
||||
# (DROP INDEX on memory_units conflicts with RowExclusiveLock from DELETE inside tx)
|
||||
if bank_internal_id:
|
||||
await bank_utils.drop_bank_hnsw_indexes(conn, bank_internal_id)
|
||||
|
||||
if invalidated_obs > 0:
|
||||
await self.submit_async_consolidation(bank_id=bank_id, request_context=request_context)
|
||||
|
||||
@@ -4547,11 +4290,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
"observation_scopes": row["observation_scopes"] if row["observation_scopes"] else None,
|
||||
}
|
||||
|
||||
# For observations, include source_memory_ids
|
||||
# history is deprecated here - use GET /memories/{id}/history instead
|
||||
if row["fact_type"] == "observation":
|
||||
result["history"] = []
|
||||
|
||||
# For observations, include source_memory_ids and fetch source_memories
|
||||
if row["fact_type"] == "observation" and row["source_memory_ids"]:
|
||||
source_ids = row["source_memory_ids"]
|
||||
result["source_memory_ids"] = [str(sid) for sid in source_ids]
|
||||
@@ -4580,95 +4319,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
|
||||
return result
|
||||
|
||||
async def get_observation_history(
|
||||
self,
|
||||
bank_id: str,
|
||||
memory_id: str,
|
||||
request_context: "RequestContext",
|
||||
) -> list[dict] | None:
|
||||
"""
|
||||
Get the history of an observation, with source facts resolved to their text.
|
||||
|
||||
Returns None if the memory is not found or is not an observation.
|
||||
Returns a list of history entries (most recent first), each with source_facts resolved.
|
||||
"""
|
||||
await self._authenticate_tenant(request_context)
|
||||
if self._operation_validator:
|
||||
from hindsight_api.extensions import BankReadContext
|
||||
|
||||
ctx = BankReadContext(bank_id=bank_id, operation="get_observation_history", request_context=request_context)
|
||||
await self._validate_operation(self._operation_validator.validate_bank_read(ctx))
|
||||
pool = await self._get_pool()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
row = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT fact_type, history, source_memory_ids
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = $1 AND bank_id = $2
|
||||
""",
|
||||
uuid.UUID(memory_id),
|
||||
bank_id,
|
||||
)
|
||||
if not row:
|
||||
return None
|
||||
if row["fact_type"] != "observation":
|
||||
return []
|
||||
|
||||
raw_history = row["history"]
|
||||
if isinstance(raw_history, str):
|
||||
raw_history = json.loads(raw_history)
|
||||
if not raw_history:
|
||||
return []
|
||||
|
||||
# Collect all source memory IDs (current full set + all historical new ones)
|
||||
current_source_ids: list[str] = [str(sid) for sid in (row["source_memory_ids"] or [])]
|
||||
all_source_ids: set[uuid.UUID] = set(uuid.UUID(sid) for sid in current_source_ids)
|
||||
for entry in raw_history:
|
||||
for sid in entry.get("new_source_memory_ids", []):
|
||||
try:
|
||||
all_source_ids.add(uuid.UUID(sid))
|
||||
except (ValueError, AttributeError):
|
||||
pass
|
||||
|
||||
# Resolve all source memories in one query
|
||||
source_map: dict[str, dict] = {}
|
||||
if all_source_ids:
|
||||
source_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, fact_type, context
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = ANY($1::uuid[])
|
||||
""",
|
||||
list(all_source_ids),
|
||||
)
|
||||
for r in source_rows:
|
||||
source_map[str(r["id"])] = {
|
||||
"id": str(r["id"]),
|
||||
"text": r["text"],
|
||||
"type": r["fact_type"],
|
||||
"context": r["context"] or None,
|
||||
}
|
||||
|
||||
# Reconstruct cumulative source IDs per change by working backwards from current state.
|
||||
# Source IDs are only ever accumulated (never removed), so:
|
||||
# after_change_N = before_change_N + new_source_memory_ids_N
|
||||
cumulative_ids: list[str] = list(current_source_ids)
|
||||
enriched: list[dict] = []
|
||||
for entry in reversed(raw_history):
|
||||
new_ids_in_entry: set[str] = set(entry.get("new_source_memory_ids", []))
|
||||
source_facts = []
|
||||
for sid in cumulative_ids:
|
||||
fact = source_map.get(sid, {"id": sid, "text": None, "type": None, "context": None})
|
||||
source_facts.append({**fact, "is_new": sid in new_ids_in_entry})
|
||||
enriched_entry = dict(entry)
|
||||
enriched_entry["source_facts"] = source_facts
|
||||
enriched.append(enriched_entry)
|
||||
# Step back: remove the new IDs added by this change to get the prior state
|
||||
cumulative_ids = [sid for sid in cumulative_ids if sid not in new_ids_in_entry]
|
||||
|
||||
enriched.reverse()
|
||||
return enriched
|
||||
|
||||
async def list_documents(
|
||||
self,
|
||||
bank_id: str,
|
||||
@@ -5044,7 +4694,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
request_context: "RequestContext",
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
tag_groups: list[TagGroup] | None = None,
|
||||
exclude_mental_model_ids: list[str] | None = None,
|
||||
_skip_span: bool = False,
|
||||
) -> ReflectResult:
|
||||
@@ -5147,7 +4796,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_results=max_results,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
tag_groups=tag_groups,
|
||||
exclude_ids=exclude_mental_model_ids,
|
||||
pending_consolidation=pending_consolidation,
|
||||
)
|
||||
@@ -5161,7 +4809,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_tokens=max_tokens,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
tag_groups=tag_groups,
|
||||
last_consolidated_at=last_consolidated_at,
|
||||
pending_consolidation=pending_consolidation,
|
||||
)
|
||||
@@ -5175,7 +4822,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_tokens=max_tokens,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
tag_groups=tag_groups,
|
||||
max_chunk_tokens=max_chunk_tokens,
|
||||
)
|
||||
|
||||
@@ -6205,39 +5851,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
|
||||
return result
|
||||
|
||||
async def get_mental_model_history(
|
||||
self,
|
||||
bank_id: str,
|
||||
mental_model_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> list[dict] | None:
|
||||
"""Get the refresh history of a mental model.
|
||||
|
||||
Returns None if the mental model is not found.
|
||||
Returns a list of history entries (most recent first), each with previous_content and changed_at.
|
||||
"""
|
||||
await self._authenticate_tenant(request_context)
|
||||
pool = await self._get_pool()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
row = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT history
|
||||
FROM {fq_table("mental_models")}
|
||||
WHERE bank_id = $1 AND id = $2
|
||||
""",
|
||||
bank_id,
|
||||
mental_model_id,
|
||||
)
|
||||
if row is None:
|
||||
return None
|
||||
raw_history = row["history"]
|
||||
if isinstance(raw_history, str):
|
||||
raw_history = json.loads(raw_history)
|
||||
if not raw_history:
|
||||
return []
|
||||
return list(reversed(raw_history))
|
||||
|
||||
async def create_mental_model(
|
||||
self,
|
||||
bank_id: str,
|
||||
@@ -6458,17 +6071,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
pool = await self._get_pool()
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# If content is changing, fetch current content first to record history
|
||||
previous_content: str | None = None
|
||||
if content is not None:
|
||||
current_row = await conn.fetchrow(
|
||||
f"SELECT content FROM {fq_table('mental_models')} WHERE bank_id = $1 AND id = $2",
|
||||
bank_id,
|
||||
mental_model_id,
|
||||
)
|
||||
if current_row:
|
||||
previous_content = current_row["content"]
|
||||
|
||||
# Build dynamic update
|
||||
updates = []
|
||||
params: list[Any] = [bank_id, mental_model_id]
|
||||
@@ -6484,14 +6086,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
params.append(content)
|
||||
param_idx += 1
|
||||
updates.append("last_refreshed_at = NOW()")
|
||||
# Record history entry with the previous content
|
||||
if get_config().enable_mental_model_history:
|
||||
history_entry = json.dumps(
|
||||
[{"previous_content": previous_content, "changed_at": datetime.now(timezone.utc).isoformat()}]
|
||||
)
|
||||
updates.append(f"history = COALESCE(history, '[]'::jsonb) || ${param_idx}::jsonb")
|
||||
params.append(history_entry)
|
||||
param_idx += 1
|
||||
# Also update embedding (convert to string for asyncpg vector type)
|
||||
embedding_text = f"{name or ''} {content}"
|
||||
embedding = await embedding_utils.generate_embeddings_batch(self.embeddings, [embedding_text])
|
||||
@@ -6912,7 +6506,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
bank_id: str,
|
||||
*,
|
||||
status: str | None = None,
|
||||
task_type: str | None = None,
|
||||
limit: int = 20,
|
||||
offset: int = 0,
|
||||
request_context: "RequestContext",
|
||||
@@ -6922,7 +6515,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
Args:
|
||||
bank_id: Bank identifier
|
||||
status: Optional status filter (pending, completed, failed)
|
||||
task_type: Optional operation type filter (retain, consolidation, etc.)
|
||||
limit: Maximum number of operations to return (default 20)
|
||||
offset: Number of operations to skip (default 0)
|
||||
request_context: Request context for authentication
|
||||
@@ -6951,10 +6543,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
where_conditions.append(f"status = ${len(params) + 1}")
|
||||
params.append(status)
|
||||
|
||||
if task_type:
|
||||
where_conditions.append(f"operation_type = ${len(params) + 1}")
|
||||
params.append(task_type)
|
||||
|
||||
where_clause = " AND ".join(where_conditions)
|
||||
|
||||
# Get total count (with filter)
|
||||
@@ -7185,64 +6773,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
"bank_id": bank_id,
|
||||
}
|
||||
|
||||
async def retry_operation(
|
||||
self,
|
||||
bank_id: str,
|
||||
operation_id: str,
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""Re-queue a failed async operation."""
|
||||
await self._authenticate_tenant(request_context)
|
||||
from hindsight_api.extensions import OperationValidationError
|
||||
|
||||
if self._operation_validator:
|
||||
from hindsight_api.extensions import BankWriteContext
|
||||
|
||||
ctx = BankWriteContext(bank_id=bank_id, operation="retry_operation", request_context=request_context)
|
||||
await self._validate_operation(self._operation_validator.validate_bank_write(ctx))
|
||||
pool = await self._get_pool()
|
||||
|
||||
op_uuid = uuid.UUID(operation_id)
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
row = await conn.fetchrow(
|
||||
f"SELECT bank_id, status FROM {fq_table('async_operations')} WHERE operation_id = $1 AND bank_id = $2",
|
||||
op_uuid,
|
||||
bank_id,
|
||||
)
|
||||
|
||||
if not row:
|
||||
raise ValueError(f"Operation {operation_id} not found for bank {bank_id}")
|
||||
|
||||
if row["status"] != "failed":
|
||||
raise OperationValidationError(
|
||||
f"Operation {operation_id} cannot be retried: status is '{row['status']}', expected 'failed'",
|
||||
409,
|
||||
)
|
||||
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("async_operations")}
|
||||
SET status = 'pending',
|
||||
error_message = NULL,
|
||||
completed_at = NULL,
|
||||
next_retry_at = NULL,
|
||||
worker_id = NULL,
|
||||
claimed_at = NULL,
|
||||
retry_count = 0,
|
||||
updated_at = NOW()
|
||||
WHERE operation_id = $1
|
||||
""",
|
||||
op_uuid,
|
||||
)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"Operation {operation_id} queued for retry",
|
||||
"operation_id": operation_id,
|
||||
}
|
||||
|
||||
async def update_bank(
|
||||
self,
|
||||
bank_id: str,
|
||||
@@ -7448,10 +6978,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
parent_operation_id = uuid.uuid4()
|
||||
pool = await self._get_pool()
|
||||
|
||||
# Ensure the bank row exists before inserting async_operations (which now has a FK).
|
||||
# Banks are created lazily on first retain, but the FK requires the row to exist first.
|
||||
await bank_utils.get_bank_profile(pool, bank_id)
|
||||
|
||||
# Create typed metadata for parent operation
|
||||
parent_metadata = BatchRetainParentMetadata(
|
||||
items_count=len(contents),
|
||||
@@ -7518,6 +7044,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
self,
|
||||
bank_id: str,
|
||||
file_items: list[dict[str, Any]],
|
||||
parser: str,
|
||||
document_tags: list[str] | None,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
@@ -7536,7 +7063,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
- metadata: Optional metadata dict
|
||||
- tags: Optional tags list
|
||||
- timestamp: Optional timestamp
|
||||
- parser: Ordered list of parser names to try (fallback chain)
|
||||
parser: Parser name (e.g., "markitdown")
|
||||
document_tags: Tags applied to all documents
|
||||
request_context: Request context for authentication
|
||||
|
||||
@@ -7592,7 +7119,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
"storage_key": storage_key,
|
||||
"original_filename": file.filename,
|
||||
"content_type": file.content_type or "application/octet-stream",
|
||||
"parser": item["parser"],
|
||||
"parser": parser,
|
||||
"context": item.get("context"),
|
||||
"metadata": item.get("metadata", {}),
|
||||
"tags": item.get("tags", []),
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user