Compare commits

..
Author SHA1 Message Date
Nicolò Boschi f515b8207b chore: run generate scripts after dead code removal 2026-03-06 14:25:16 +01:00
Nicolò Boschi 0e3ccd038b refactor: remove dead code and clarify observations vs mental models
- Delete engine/mental_models/ module (stale Pydantic models with wrong
  schema, describing an old design where mental models were directives;
  had no importers outside itself)
- Remove unused imports in api/http.py (acquire_with_retry, Observation)
- Remove unused Pydantic models in api/http.py (BanksResponse,
  ObservationEvidenceResponse)
- Add clarifying NOTE to consolidation/consolidator.py distinguishing
  observations (auto-generated bottom-up) from mental models (user-defined
  pinned reflections refreshed via reflect)
2026-03-06 14:00:26 +01:00
652 changed files with 3488 additions and 21039 deletions
+1 -6
View File
@@ -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
-6
View File
@@ -1,6 +0,0 @@
version: 2
updates:
- package-ecosystem: "github-actions"
directory: "/"
schedule:
interval: "weekly"
+4 -4
View File
@@ -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:
+46 -70
View File
@@ -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
View File
@@ -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 -17
View File
@@ -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)
+2 -3
View File
@@ -9,9 +9,8 @@
[![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
![PyPI - Downloads](https://img.shields.io/pypi/dm/hindsight-api?label=PyPI)
![NPM Downloads](https://img.shields.io/npm/dm/%40vectorize-io%2Fhindsight-client?logoColor=orange&label=NPM&color=blue&link=https%3A%2F%2Fwww.npmjs.com%2Fpackage%2F%40vectorize-io%2Fhindsight-client)
<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).
+13 -10
View File
@@ -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 .
+2 -17
View File
@@ -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)
-18
View File
@@ -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 -2
View File
@@ -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
-33
View File
@@ -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"
-48
View File
@@ -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
```
-423
View File
@@ -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
-137
View File
@@ -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
@@ -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")
@@ -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"
)
@@ -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")
@@ -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")
@@ -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"
)
@@ -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))
-198
View File
@@ -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)
@@ -46,4 +46,4 @@ __all__ = [
"RemoteTEICrossEncoder",
"LLMConfig",
]
__version__ = "0.4.18"
__version__ = "0.4.16"
@@ -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:
@@ -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
)
@@ -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)
""")
@@ -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,
@@ -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)
@@ -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"}}]}}
@@ -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'"
)
@@ -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}")
@@ -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.
@@ -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
@@ -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