Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
344ac8fae8 | ||
|
|
4b0c617ecf | ||
|
|
0a04770450 | ||
|
|
60574ee08f | ||
|
|
7d95a002c7 | ||
|
|
83ca669011 | ||
|
|
e798979733 | ||
|
|
43f9a8bec2 | ||
|
|
f641b30d83 | ||
|
|
90be7c6829 | ||
|
|
6eec83b20d |
@@ -127,6 +127,38 @@ API URL for control plane
|
||||
{{- printf "http://%s-api:%d" (include "hindsight.fullname" .) (.Values.api.service.port | int) }}
|
||||
{{- end }}
|
||||
|
||||
{{/*
|
||||
TEI reranker labels
|
||||
*/}}
|
||||
{{- define "hindsight.tei.reranker.labels" -}}
|
||||
{{ include "hindsight.labels" . }}
|
||||
app.kubernetes.io/component: tei-reranker
|
||||
{{- end }}
|
||||
|
||||
{{/*
|
||||
TEI reranker selector labels
|
||||
*/}}
|
||||
{{- define "hindsight.tei.reranker.selectorLabels" -}}
|
||||
{{ include "hindsight.selectorLabels" . }}
|
||||
app.kubernetes.io/component: tei-reranker
|
||||
{{- end }}
|
||||
|
||||
{{/*
|
||||
TEI embedding labels
|
||||
*/}}
|
||||
{{- define "hindsight.tei.embedding.labels" -}}
|
||||
{{ include "hindsight.labels" . }}
|
||||
app.kubernetes.io/component: tei-embedding
|
||||
{{- end }}
|
||||
|
||||
{{/*
|
||||
TEI embedding selector labels
|
||||
*/}}
|
||||
{{- define "hindsight.tei.embedding.selectorLabels" -}}
|
||||
{{ include "hindsight.selectorLabels" . }}
|
||||
app.kubernetes.io/component: tei-embedding
|
||||
{{- end }}
|
||||
|
||||
{{/*
|
||||
Get the name of the secret to use
|
||||
*/}}
|
||||
|
||||
@@ -67,6 +67,18 @@ spec:
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
{{- if .Values.tei.reranker.enabled }}
|
||||
- name: HINDSIGHT_API_RERANKER_PROVIDER
|
||||
value: "tei"
|
||||
- name: HINDSIGHT_API_RERANKER_TEI_URL
|
||||
value: "http://{{ include "hindsight.fullname" . }}-tei-reranker:{{ .Values.tei.reranker.port }}"
|
||||
{{- end }}
|
||||
{{- if .Values.tei.embedding.enabled }}
|
||||
- name: HINDSIGHT_API_EMBEDDINGS_PROVIDER
|
||||
value: "tei"
|
||||
- name: HINDSIGHT_API_EMBEDDINGS_TEI_URL
|
||||
value: "http://{{ include "hindsight.fullname" . }}-tei-embedding:{{ .Values.tei.embedding.port }}"
|
||||
{{- end }}
|
||||
{{- /* Only use api.secrets when not using existingSecret (for chart-managed secrets) */}}
|
||||
{{- if not .Values.existingSecret }}
|
||||
{{- range $key, $value := .Values.api.secrets }}
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
{{- if .Values.tei.embedding.enabled }}
|
||||
apiVersion: apps/v1
|
||||
kind: Deployment
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-tei-embedding
|
||||
labels:
|
||||
{{- include "hindsight.tei.embedding.labels" . | nindent 4 }}
|
||||
spec:
|
||||
replicas: {{ .Values.tei.embedding.replicaCount }}
|
||||
selector:
|
||||
matchLabels:
|
||||
{{- include "hindsight.tei.embedding.selectorLabels" . | nindent 6 }}
|
||||
template:
|
||||
metadata:
|
||||
{{- with .Values.podAnnotations }}
|
||||
annotations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
labels:
|
||||
{{- include "hindsight.tei.embedding.selectorLabels" . | nindent 8 }}
|
||||
spec:
|
||||
{{- if .Values.serviceAccount.create }}
|
||||
serviceAccountName: {{ include "hindsight.serviceAccountName" . }}
|
||||
{{- end }}
|
||||
securityContext:
|
||||
{{- toYaml .Values.podSecurityContext | nindent 8 }}
|
||||
containers:
|
||||
- name: tei-embedding
|
||||
securityContext:
|
||||
{{- toYaml .Values.securityContext | nindent 10 }}
|
||||
image: "{{ .Values.tei.embedding.image.repository }}:{{ .Values.tei.embedding.image.tag }}"
|
||||
imagePullPolicy: {{ .Values.tei.embedding.image.pullPolicy }}
|
||||
args:
|
||||
- "--model-id"
|
||||
- {{ .Values.tei.embedding.model | quote }}
|
||||
- "--hostname"
|
||||
- "0.0.0.0"
|
||||
{{- range .Values.tei.embedding.args }}
|
||||
- {{ . | quote }}
|
||||
{{- end }}
|
||||
ports:
|
||||
- name: http
|
||||
containerPort: {{ .Values.tei.embedding.port }}
|
||||
protocol: TCP
|
||||
env:
|
||||
- name: PORT
|
||||
value: {{ .Values.tei.embedding.port | quote }}
|
||||
{{- range $key, $value := .Values.tei.embedding.env }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
livenessProbe:
|
||||
{{- toYaml .Values.tei.embedding.livenessProbe | nindent 10 }}
|
||||
readinessProbe:
|
||||
{{- toYaml .Values.tei.embedding.readinessProbe | nindent 10 }}
|
||||
resources:
|
||||
{{- toYaml .Values.tei.embedding.resources | nindent 10 }}
|
||||
volumeMounts:
|
||||
- name: model-cache
|
||||
mountPath: /data
|
||||
volumes:
|
||||
- name: model-cache
|
||||
emptyDir: {}
|
||||
{{- with .Values.nodeSelector }}
|
||||
nodeSelector:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.affinity }}
|
||||
affinity:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.tolerations }}
|
||||
tolerations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
@@ -0,0 +1,17 @@
|
||||
{{- if .Values.tei.embedding.enabled }}
|
||||
apiVersion: v1
|
||||
kind: Service
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-tei-embedding
|
||||
labels:
|
||||
{{- include "hindsight.tei.embedding.labels" . | nindent 4 }}
|
||||
spec:
|
||||
type: ClusterIP
|
||||
ports:
|
||||
- port: {{ .Values.tei.embedding.port }}
|
||||
targetPort: http
|
||||
protocol: TCP
|
||||
name: http
|
||||
selector:
|
||||
{{- include "hindsight.tei.embedding.selectorLabels" . | nindent 4 }}
|
||||
{{- end }}
|
||||
@@ -0,0 +1,76 @@
|
||||
{{- if .Values.tei.reranker.enabled }}
|
||||
apiVersion: apps/v1
|
||||
kind: Deployment
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-tei-reranker
|
||||
labels:
|
||||
{{- include "hindsight.tei.reranker.labels" . | nindent 4 }}
|
||||
spec:
|
||||
replicas: {{ .Values.tei.reranker.replicaCount }}
|
||||
selector:
|
||||
matchLabels:
|
||||
{{- include "hindsight.tei.reranker.selectorLabels" . | nindent 6 }}
|
||||
template:
|
||||
metadata:
|
||||
{{- with .Values.podAnnotations }}
|
||||
annotations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
labels:
|
||||
{{- include "hindsight.tei.reranker.selectorLabels" . | nindent 8 }}
|
||||
spec:
|
||||
{{- if .Values.serviceAccount.create }}
|
||||
serviceAccountName: {{ include "hindsight.serviceAccountName" . }}
|
||||
{{- end }}
|
||||
securityContext:
|
||||
{{- toYaml .Values.podSecurityContext | nindent 8 }}
|
||||
containers:
|
||||
- name: tei-reranker
|
||||
securityContext:
|
||||
{{- toYaml .Values.securityContext | nindent 10 }}
|
||||
image: "{{ .Values.tei.reranker.image.repository }}:{{ .Values.tei.reranker.image.tag }}"
|
||||
imagePullPolicy: {{ .Values.tei.reranker.image.pullPolicy }}
|
||||
args:
|
||||
- "--model-id"
|
||||
- {{ .Values.tei.reranker.model | quote }}
|
||||
- "--hostname"
|
||||
- "0.0.0.0"
|
||||
{{- range .Values.tei.reranker.args }}
|
||||
- {{ . | quote }}
|
||||
{{- end }}
|
||||
ports:
|
||||
- name: http
|
||||
containerPort: {{ .Values.tei.reranker.port }}
|
||||
protocol: TCP
|
||||
env:
|
||||
- name: PORT
|
||||
value: {{ .Values.tei.reranker.port | quote }}
|
||||
{{- range $key, $value := .Values.tei.reranker.env }}
|
||||
- name: {{ $key }}
|
||||
value: {{ $value | quote }}
|
||||
{{- end }}
|
||||
livenessProbe:
|
||||
{{- toYaml .Values.tei.reranker.livenessProbe | nindent 10 }}
|
||||
readinessProbe:
|
||||
{{- toYaml .Values.tei.reranker.readinessProbe | nindent 10 }}
|
||||
resources:
|
||||
{{- toYaml .Values.tei.reranker.resources | nindent 10 }}
|
||||
volumeMounts:
|
||||
- name: model-cache
|
||||
mountPath: /data
|
||||
volumes:
|
||||
- name: model-cache
|
||||
emptyDir: {}
|
||||
{{- with .Values.nodeSelector }}
|
||||
nodeSelector:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.affinity }}
|
||||
affinity:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.tolerations }}
|
||||
tolerations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
@@ -0,0 +1,17 @@
|
||||
{{- if .Values.tei.reranker.enabled }}
|
||||
apiVersion: v1
|
||||
kind: Service
|
||||
metadata:
|
||||
name: {{ include "hindsight.fullname" . }}-tei-reranker
|
||||
labels:
|
||||
{{- include "hindsight.tei.reranker.labels" . | nindent 4 }}
|
||||
spec:
|
||||
type: ClusterIP
|
||||
ports:
|
||||
- port: {{ .Values.tei.reranker.port }}
|
||||
targetPort: http
|
||||
protocol: TCP
|
||||
name: http
|
||||
selector:
|
||||
{{- include "hindsight.tei.reranker.selectorLabels" . | nindent 4 }}
|
||||
{{- end }}
|
||||
@@ -293,6 +293,84 @@ tolerations: []
|
||||
# Affinity (applied to all components unless overridden per-component)
|
||||
affinity: {}
|
||||
|
||||
# TEI (Text Embeddings Inference) - optional standalone deployments
|
||||
# for reranking and/or embedding models
|
||||
tei:
|
||||
reranker:
|
||||
enabled: false
|
||||
replicaCount: 1
|
||||
image:
|
||||
repository: ghcr.io/huggingface/text-embeddings-inference
|
||||
tag: cpu-1.8.3
|
||||
pullPolicy: IfNotPresent
|
||||
model: "cross-encoder/ms-marco-MiniLM-L-6-v2"
|
||||
port: 8090
|
||||
args:
|
||||
- "--auto-truncate"
|
||||
env:
|
||||
PAYLOAD_LIMIT: "10000000"
|
||||
MAX_CLIENT_BATCH_SIZE: "256"
|
||||
resources:
|
||||
limits:
|
||||
cpu: 2000m
|
||||
memory: 2Gi
|
||||
requests:
|
||||
cpu: 500m
|
||||
memory: 1Gi
|
||||
livenessProbe:
|
||||
httpGet:
|
||||
path: /health
|
||||
port: 8090
|
||||
initialDelaySeconds: 30
|
||||
periodSeconds: 10
|
||||
timeoutSeconds: 5
|
||||
failureThreshold: 6
|
||||
readinessProbe:
|
||||
httpGet:
|
||||
path: /health
|
||||
port: 8090
|
||||
initialDelaySeconds: 15
|
||||
periodSeconds: 5
|
||||
timeoutSeconds: 3
|
||||
failureThreshold: 3
|
||||
|
||||
embedding:
|
||||
enabled: false
|
||||
replicaCount: 1
|
||||
image:
|
||||
repository: ghcr.io/huggingface/text-embeddings-inference
|
||||
tag: cpu-1.8.3
|
||||
pullPolicy: IfNotPresent
|
||||
model: "sentence-transformers/all-MiniLM-L6-v2"
|
||||
port: 8091
|
||||
args: []
|
||||
env:
|
||||
PAYLOAD_LIMIT: "10000000"
|
||||
MAX_CLIENT_BATCH_SIZE: "256"
|
||||
resources:
|
||||
limits:
|
||||
cpu: 2000m
|
||||
memory: 2Gi
|
||||
requests:
|
||||
cpu: 500m
|
||||
memory: 1Gi
|
||||
livenessProbe:
|
||||
httpGet:
|
||||
path: /health
|
||||
port: 8091
|
||||
initialDelaySeconds: 30
|
||||
periodSeconds: 10
|
||||
timeoutSeconds: 5
|
||||
failureThreshold: 6
|
||||
readinessProbe:
|
||||
httpGet:
|
||||
path: /health
|
||||
port: 8091
|
||||
initialDelaySeconds: 15
|
||||
periodSeconds: 5
|
||||
timeoutSeconds: 3
|
||||
failureThreshold: 3
|
||||
|
||||
# Autoscaling
|
||||
autoscaling:
|
||||
enabled: false
|
||||
|
||||
@@ -6,7 +6,6 @@ Provides both HTTP REST API and MCP (Model Context Protocol) server.
|
||||
|
||||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import FastAPI
|
||||
|
||||
@@ -46,14 +45,14 @@ def create_app(
|
||||
# Both HTTP and MCP
|
||||
app = create_app(memory, mcp_api_enabled=True)
|
||||
"""
|
||||
mcp_app = None
|
||||
mcp_servers = None
|
||||
|
||||
# Create MCP app first if enabled (we need its lifespan for chaining)
|
||||
# Create MCP servers first if enabled (we need their lifespans for chaining)
|
||||
if mcp_api_enabled:
|
||||
try:
|
||||
from .mcp import create_mcp_app
|
||||
from .mcp import MCPMiddleware, create_mcp_servers
|
||||
|
||||
mcp_app = create_mcp_app(memory=memory)
|
||||
mcp_servers = create_mcp_servers(memory=memory)
|
||||
except ImportError as e:
|
||||
logger.error(f"MCP server requested but dependencies not available: {e}")
|
||||
logger.error("Install with: pip install hindsight-api[mcp]")
|
||||
@@ -70,11 +69,9 @@ def create_app(
|
||||
app = FastAPI(title="Hindsight API", version="0.0.7")
|
||||
logger.info("HTTP REST API disabled")
|
||||
|
||||
# Mount MCP server and chain its lifespan if enabled
|
||||
if mcp_app is not None:
|
||||
# Get both MCP apps' underlying Starlette apps for lifespan access
|
||||
multi_bank_starlette_app = mcp_app.multi_bank_app
|
||||
single_bank_starlette_app = mcp_app.single_bank_app
|
||||
# Add MCP middleware and chain its lifespan if enabled
|
||||
if mcp_servers is not None:
|
||||
multi_bank_server, single_bank_server, multi_bank_starlette_app, single_bank_starlette_app = mcp_servers
|
||||
|
||||
# Store the original lifespan
|
||||
original_lifespan = app.router.lifespan_context
|
||||
@@ -94,8 +91,19 @@ def create_app(
|
||||
# Replace the app's lifespan with the chained version
|
||||
app.router.lifespan_context = chained_lifespan
|
||||
|
||||
# Mount the MCP middleware
|
||||
app.mount(mcp_mount_path, mcp_app)
|
||||
# Add MCP as a wrapping middleware — intercepts /mcp* requests directly,
|
||||
# passes everything else through to the FastAPI app. No Starlette Mount
|
||||
# means no 307 redirect for /mcp (no trailing slash).
|
||||
app.add_middleware(
|
||||
MCPMiddleware,
|
||||
memory=memory,
|
||||
prefix=mcp_mount_path,
|
||||
multi_bank_app=multi_bank_starlette_app,
|
||||
single_bank_app=single_bank_starlette_app,
|
||||
multi_bank_server=multi_bank_server,
|
||||
single_bank_server=single_bank_server,
|
||||
)
|
||||
|
||||
logger.info(f"MCP server enabled at {mcp_mount_path}/")
|
||||
|
||||
return app
|
||||
|
||||
@@ -32,9 +32,44 @@ def _parse_metadata(metadata: Any) -> dict[str, Any]:
|
||||
return {}
|
||||
|
||||
|
||||
from typing import Callable
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
|
||||
from hindsight_api import MemoryEngine
|
||||
|
||||
|
||||
def FieldWithDefault(default_factory: Callable, **kwargs) -> Any:
|
||||
"""
|
||||
Field wrapper that ensures default_factory values appear in OpenAPI schema.
|
||||
|
||||
Pydantic doesn't include default_factory in OpenAPI schemas, causing OpenAPI
|
||||
Generator to make fields Optional with default=None instead of non-optional
|
||||
with the correct default value.
|
||||
|
||||
This wrapper adds json_schema_extra to include the default in the schema.
|
||||
"""
|
||||
# Determine the default value for the schema based on the factory
|
||||
if default_factory is list:
|
||||
schema_default = []
|
||||
elif default_factory is dict:
|
||||
schema_default = {}
|
||||
else:
|
||||
# For custom factories (like IncludeOptions), use empty dict as placeholder
|
||||
schema_default = {}
|
||||
|
||||
# Add or merge json_schema_extra
|
||||
json_extra = kwargs.pop("json_schema_extra", {})
|
||||
if isinstance(json_extra, dict):
|
||||
json_extra["default"] = schema_default
|
||||
else:
|
||||
# If json_schema_extra was a function, we can't merge easily
|
||||
# Fall back to just setting default
|
||||
json_extra = {"default": schema_default}
|
||||
|
||||
return Field(default_factory=default_factory, json_schema_extra=json_extra, **kwargs)
|
||||
|
||||
|
||||
from hindsight_api.engine.db_utils import acquire_with_retry
|
||||
from hindsight_api.engine.memory_engine import Budget, _get_tiktoken_encoding, fq_table
|
||||
from hindsight_api.engine.reflect.observations import Observation
|
||||
@@ -103,8 +138,8 @@ class RecallRequest(BaseModel):
|
||||
query_timestamp: str | None = Field(
|
||||
default=None, description="ISO format date string (e.g., '2023-05-30T23:40:00')"
|
||||
)
|
||||
include: IncludeOptions = Field(
|
||||
default_factory=IncludeOptions,
|
||||
include: IncludeOptions = FieldWithDefault(
|
||||
IncludeOptions,
|
||||
description="Options for including additional data (entities are included by default)",
|
||||
)
|
||||
tags: list[str] | None = Field(
|
||||
@@ -570,18 +605,16 @@ class ReflectLLMCall(BaseModel):
|
||||
class ReflectBasedOn(BaseModel):
|
||||
"""Evidence the response is based on: memories, mental models, and directives."""
|
||||
|
||||
memories: list[ReflectFact] = Field(default_factory=list, description="Memory facts used to generate the response")
|
||||
mental_models: list[ReflectMentalModel] = Field(
|
||||
default_factory=list, description="Mental models used during reflection"
|
||||
)
|
||||
directives: list[ReflectDirective] = Field(default_factory=list, description="Directives applied during reflection")
|
||||
memories: list[ReflectFact] = FieldWithDefault(list, description="Memory facts used to generate the response")
|
||||
mental_models: list[ReflectMentalModel] = FieldWithDefault(list, description="Mental models used during reflection")
|
||||
directives: list[ReflectDirective] = FieldWithDefault(list, description="Directives applied during reflection")
|
||||
|
||||
|
||||
class ReflectTrace(BaseModel):
|
||||
"""Execution trace of LLM and tool calls during reflection."""
|
||||
|
||||
tool_calls: list[ReflectToolCall] = Field(default_factory=list, description="Tool calls made during reflection")
|
||||
llm_calls: list[ReflectLLMCall] = Field(default_factory=list, description="LLM calls made during reflection")
|
||||
tool_calls: list[ReflectToolCall] = FieldWithDefault(list, description="Tool calls made during reflection")
|
||||
llm_calls: list[ReflectLLMCall] = FieldWithDefault(list, description="LLM calls made during reflection")
|
||||
|
||||
|
||||
class ReflectResponse(BaseModel):
|
||||
@@ -942,7 +975,7 @@ class DocumentResponse(BaseModel):
|
||||
created_at: str
|
||||
updated_at: str
|
||||
memory_unit_count: int
|
||||
tags: list[str] = Field(default_factory=list, description="Tags associated with this document")
|
||||
tags: list[str] = FieldWithDefault(list, description="Tags associated with this document")
|
||||
|
||||
|
||||
class DeleteDocumentResponse(BaseModel):
|
||||
@@ -1066,7 +1099,7 @@ class DirectiveResponse(BaseModel):
|
||||
content: str
|
||||
priority: int = 0
|
||||
is_active: bool = True
|
||||
tags: list[str] = Field(default_factory=list)
|
||||
tags: list[str] = FieldWithDefault(list)
|
||||
created_at: str | None = None
|
||||
updated_at: str | None = None
|
||||
|
||||
@@ -1084,7 +1117,7 @@ class CreateDirectiveRequest(BaseModel):
|
||||
content: str = Field(description="The directive text to inject into prompts")
|
||||
priority: int = Field(default=0, description="Higher priority directives are injected first")
|
||||
is_active: bool = Field(default=True, description="Whether this directive is active")
|
||||
tags: list[str] = Field(default_factory=list, description="Tags for filtering")
|
||||
tags: list[str] = FieldWithDefault(list, description="Tags for filtering")
|
||||
|
||||
|
||||
class UpdateDirectiveRequest(BaseModel):
|
||||
@@ -1121,9 +1154,9 @@ class MentalModelResponse(BaseModel):
|
||||
content: str = Field(
|
||||
description="The mental model content as well-formatted markdown (auto-generated from reflect endpoint)"
|
||||
)
|
||||
tags: list[str] = Field(default_factory=list)
|
||||
tags: list[str] = FieldWithDefault(list)
|
||||
max_tokens: int = Field(default=2048)
|
||||
trigger: MentalModelTrigger = Field(default_factory=MentalModelTrigger)
|
||||
trigger: MentalModelTrigger = FieldWithDefault(MentalModelTrigger)
|
||||
last_refreshed_at: str | None = None
|
||||
created_at: str | None = None
|
||||
reflect_response: dict | None = Field(
|
||||
@@ -1159,9 +1192,9 @@ class CreateMentalModelRequest(BaseModel):
|
||||
)
|
||||
name: str = Field(description="Human-readable name for the mental model")
|
||||
source_query: str = Field(description="The query to run to generate content")
|
||||
tags: list[str] = Field(default_factory=list, description="Tags for scoped visibility")
|
||||
tags: list[str] = FieldWithDefault(list, description="Tags for scoped visibility")
|
||||
max_tokens: int = Field(default=2048, ge=256, le=8192, description="Maximum tokens for generated content")
|
||||
trigger: MentalModelTrigger = Field(default_factory=MentalModelTrigger, description="Trigger settings")
|
||||
trigger: MentalModelTrigger = FieldWithDefault(MentalModelTrigger, description="Trigger settings")
|
||||
|
||||
|
||||
class CreateMentalModelResponse(BaseModel):
|
||||
@@ -2354,23 +2387,6 @@ def _register_routes(app: FastAPI):
|
||||
):
|
||||
"""Get a mental model by ID."""
|
||||
try:
|
||||
# Pre-operation validation hook
|
||||
validator = app.state.memory._operation_validator
|
||||
if validator:
|
||||
from hindsight_api.extensions.operation_validator import MentalModelGetContext
|
||||
|
||||
ctx = MentalModelGetContext(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=request_context,
|
||||
)
|
||||
validation = await validator.validate_mental_model_get(ctx)
|
||||
if not validation.allowed:
|
||||
raise OperationValidationError(
|
||||
validation.reason or "Operation not allowed",
|
||||
status_code=validation.status_code,
|
||||
)
|
||||
|
||||
mental_model = await app.state.memory.get_mental_model(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
@@ -2379,25 +2395,6 @@ def _register_routes(app: FastAPI):
|
||||
if mental_model is None:
|
||||
raise HTTPException(status_code=404, detail=f"Mental model '{mental_model_id}' not found")
|
||||
|
||||
# Post-operation hook
|
||||
if validator:
|
||||
from hindsight_api.extensions.operation_validator import MentalModelGetResult
|
||||
|
||||
content = mental_model.get("content", "")
|
||||
output_tokens = len(content) // 4 if content else 0
|
||||
|
||||
result_ctx = MentalModelGetResult(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=request_context,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
)
|
||||
try:
|
||||
await validator.on_mental_model_get_complete(result_ctx)
|
||||
except Exception as hook_err:
|
||||
logger.warning(f"Post-mental-model-get hook error (non-fatal): {hook_err}")
|
||||
|
||||
return MentalModelResponse(**mental_model)
|
||||
except (AuthenticationError, HTTPException):
|
||||
raise
|
||||
@@ -2427,23 +2424,6 @@ def _register_routes(app: FastAPI):
|
||||
):
|
||||
"""Create a mental model (async - returns operation_id)."""
|
||||
try:
|
||||
# Pre-operation validation hook
|
||||
validator = app.state.memory._operation_validator
|
||||
if validator:
|
||||
from hindsight_api.extensions.operation_validator import MentalModelRefreshContext
|
||||
|
||||
ctx = MentalModelRefreshContext(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=None, # Not yet created
|
||||
request_context=request_context,
|
||||
)
|
||||
validation = await validator.validate_mental_model_refresh(ctx)
|
||||
if not validation.allowed:
|
||||
raise OperationValidationError(
|
||||
validation.reason or "Operation not allowed",
|
||||
status_code=validation.status_code,
|
||||
)
|
||||
|
||||
# 1. Create the mental model with placeholder content
|
||||
mental_model = await app.state.memory.create_mental_model(
|
||||
bank_id=bank_id,
|
||||
@@ -2491,23 +2471,6 @@ def _register_routes(app: FastAPI):
|
||||
):
|
||||
"""Refresh a mental model by re-running its source query (async)."""
|
||||
try:
|
||||
# Pre-operation validation hook
|
||||
validator = app.state.memory._operation_validator
|
||||
if validator:
|
||||
from hindsight_api.extensions.operation_validator import MentalModelRefreshContext
|
||||
|
||||
ctx = MentalModelRefreshContext(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=request_context,
|
||||
)
|
||||
validation = await validator.validate_mental_model_refresh(ctx)
|
||||
if not validation.allowed:
|
||||
raise OperationValidationError(
|
||||
validation.reason or "Operation not allowed",
|
||||
status_code=validation.status_code,
|
||||
)
|
||||
|
||||
result = await app.state.memory.submit_async_refresh_mental_model(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
|
||||
@@ -90,7 +90,19 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||
tenant_id_resolver=get_current_tenant_id, # Propagate tenant_id for usage metering
|
||||
api_key_id_resolver=get_current_api_key_id, # Propagate api_key_id for usage metering
|
||||
include_bank_id_param=multi_bank,
|
||||
tools=None if multi_bank else {"retain", "recall", "reflect"}, # Scoped tools for single-bank mode
|
||||
tools=None
|
||||
if multi_bank
|
||||
else {
|
||||
"retain",
|
||||
"recall",
|
||||
"reflect",
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"create_mental_model",
|
||||
"update_mental_model",
|
||||
"delete_mental_model",
|
||||
"refresh_mental_model",
|
||||
}, # Scoped tools for single-bank mode (excludes bank management: list_banks, create_bank)
|
||||
retain_fire_and_forget=False, # HTTP MCP supports sync/async modes
|
||||
)
|
||||
|
||||
@@ -106,7 +118,10 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||
|
||||
|
||||
class MCPMiddleware:
|
||||
"""ASGI middleware that handles authentication and routes to appropriate MCP server.
|
||||
"""ASGI middleware that intercepts MCP requests and routes to appropriate MCP server.
|
||||
|
||||
This middleware wraps the main FastAPI app and intercepts requests matching the
|
||||
configured prefix (default: /mcp). Non-MCP requests pass through to the inner app.
|
||||
|
||||
Authentication:
|
||||
1. If HINDSIGHT_API_MCP_AUTH_TOKEN is set (legacy), validates against that token
|
||||
@@ -137,27 +152,33 @@ class MCPMiddleware:
|
||||
--header "X-Bank-Id: my-bank" --header "Authorization: Bearer <token>"
|
||||
"""
|
||||
|
||||
def __init__(self, app, memory: MemoryEngine):
|
||||
def __init__(
|
||||
self,
|
||||
app,
|
||||
memory: MemoryEngine,
|
||||
prefix: str = "/mcp",
|
||||
multi_bank_app=None,
|
||||
single_bank_app=None,
|
||||
multi_bank_server=None,
|
||||
single_bank_server=None,
|
||||
):
|
||||
self.app = app
|
||||
self.prefix = prefix
|
||||
self.memory = memory
|
||||
self.tenant_extension = memory._tenant_extension
|
||||
|
||||
# Create two server instances:
|
||||
# 1. Multi-bank server (for /mcp/ root endpoint)
|
||||
self.multi_bank_server = create_mcp_server(memory, multi_bank=True)
|
||||
self.multi_bank_app = self.multi_bank_server.http_app(path="/")
|
||||
|
||||
# 2. Single-bank server (for /mcp/{bank_id}/ endpoints)
|
||||
self.single_bank_server = create_mcp_server(memory, multi_bank=False)
|
||||
self.single_bank_app = self.single_bank_server.http_app(path="/")
|
||||
|
||||
# Backward compatibility: expose multi_bank_app as mcp_app
|
||||
self.mcp_app = self.multi_bank_app
|
||||
|
||||
# Expose the lifespan for the parent app to chain (use multi-bank as default)
|
||||
self.lifespan = (
|
||||
self.multi_bank_app.lifespan_handler if hasattr(self.multi_bank_app, "lifespan_handler") else None
|
||||
)
|
||||
if multi_bank_app and single_bank_app:
|
||||
# Pre-created servers (used when called via add_middleware from create_app)
|
||||
self.multi_bank_app = multi_bank_app
|
||||
self.single_bank_app = single_bank_app
|
||||
self.multi_bank_server = multi_bank_server
|
||||
self.single_bank_server = single_bank_server
|
||||
else:
|
||||
# Create servers internally (for direct construction / tests)
|
||||
self.multi_bank_server = create_mcp_server(memory, multi_bank=True)
|
||||
self.multi_bank_app = self.multi_bank_server.http_app(path="/")
|
||||
self.single_bank_server = create_mcp_server(memory, multi_bank=False)
|
||||
self.single_bank_app = self.single_bank_server.http_app(path="/")
|
||||
|
||||
def _get_header(self, scope: dict, name: str) -> str | None:
|
||||
"""Extract a header value from ASGI scope."""
|
||||
@@ -169,9 +190,20 @@ class MCPMiddleware:
|
||||
|
||||
async def __call__(self, scope, receive, send):
|
||||
if scope["type"] != "http":
|
||||
await self.multi_bank_app(scope, receive, send)
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
path = scope.get("path", "")
|
||||
|
||||
# Check if this is an MCP request (matches prefix)
|
||||
if not (path == self.prefix or path.startswith(self.prefix + "/")):
|
||||
# Not an MCP request — pass through to the inner app
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
# Strip prefix from path
|
||||
path = path[len(self.prefix) :] or "/"
|
||||
|
||||
# Extract auth token from header (for tenant auth propagation)
|
||||
auth_header = self._get_header(scope, "Authorization")
|
||||
auth_token: str | None = None
|
||||
@@ -210,36 +242,15 @@ class MCPMiddleware:
|
||||
_current_schema.set(tenant_context.schema_name) if tenant_context and tenant_context.schema_name else None
|
||||
)
|
||||
|
||||
path = scope.get("path", "")
|
||||
|
||||
# Strip any mount prefix (e.g., /mcp) that FastAPI might not have stripped
|
||||
root_path = scope.get("root_path", "")
|
||||
if root_path and path.startswith(root_path):
|
||||
path = path[len(root_path) :] or "/"
|
||||
|
||||
# Also handle case where mount path wasn't stripped (e.g., /mcp/...)
|
||||
if path.startswith("/mcp/"):
|
||||
path = path[4:] # Remove /mcp prefix
|
||||
elif path == "/mcp":
|
||||
path = "/"
|
||||
|
||||
# Ensure path has leading slash (needed after stripping mount path)
|
||||
if path and not path.startswith("/"):
|
||||
path = "/" + path
|
||||
|
||||
# Try to get bank_id from header first (for Claude Code compatibility)
|
||||
bank_id = self._get_header(scope, "X-Bank-Id")
|
||||
bank_id_from_path = False
|
||||
|
||||
# MCP endpoint paths that should not be treated as bank_ids
|
||||
MCP_ENDPOINTS = {"sse", "messages"}
|
||||
|
||||
# If no header, try to extract from path: /{bank_id}/...
|
||||
new_path = path
|
||||
if not bank_id and path.startswith("/") and len(path) > 1:
|
||||
parts = path[1:].split("/", 1)
|
||||
# Don't treat MCP endpoints as bank_ids
|
||||
if parts[0] and parts[0] not in MCP_ENDPOINTS:
|
||||
if parts[0]:
|
||||
# First segment looks like a bank_id
|
||||
bank_id = parts[0]
|
||||
bank_id_from_path = True
|
||||
@@ -268,9 +279,19 @@ class MCPMiddleware:
|
||||
# Clear root_path since we're passing directly to the app
|
||||
new_scope["root_path"] = ""
|
||||
|
||||
# Wrap send to rewrite the SSE endpoint URL to include bank_id if using path-based routing
|
||||
# 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".
|
||||
is_sse_response = False
|
||||
|
||||
async def send_wrapper(message):
|
||||
if message["type"] == "http.response.body" and bank_id_from_path:
|
||||
nonlocal is_sse_response
|
||||
if message["type"] == "http.response.start":
|
||||
for header_name, header_value in message.get("headers", []):
|
||||
if header_name == b"content-type" and b"text/event-stream" in header_value:
|
||||
is_sse_response = True
|
||||
break
|
||||
if message["type"] == "http.response.body" and bank_id_from_path and is_sse_response:
|
||||
body = message.get("body", b"")
|
||||
if body and b"/messages" in body:
|
||||
# Rewrite /messages to /{bank_id}/messages in SSE endpoint event
|
||||
@@ -308,30 +329,19 @@ class MCPMiddleware:
|
||||
)
|
||||
|
||||
|
||||
def create_mcp_app(memory: MemoryEngine):
|
||||
"""
|
||||
Create an ASGI app that handles MCP requests with dynamic tool exposure.
|
||||
def create_mcp_servers(memory: MemoryEngine):
|
||||
"""Create multi-bank and single-bank MCP servers and their Starlette apps.
|
||||
|
||||
Authentication:
|
||||
Uses the TenantExtension from the MemoryEngine (same auth as REST API).
|
||||
|
||||
Two modes based on URL structure:
|
||||
|
||||
1. Single-bank mode (recommended for agent isolation):
|
||||
- URL: /mcp/{bank_id}/
|
||||
- Tools: retain, recall, reflect (no bank_id parameter)
|
||||
- Example: claude mcp add --transport http my-agent http://localhost:8888/mcp/my-agent-bank/
|
||||
|
||||
2. Multi-bank mode (for cross-bank operations):
|
||||
- URL: /mcp/
|
||||
- Tools: retain, recall, reflect, list_banks, create_bank (all with bank_id parameter)
|
||||
- Bank ID from: X-Bank-Id header or HINDSIGHT_MCP_BANK_ID env var (default: "default")
|
||||
- Example: claude mcp add --transport http hindsight http://localhost:8888/mcp --header "X-Bank-Id: my-bank"
|
||||
|
||||
Args:
|
||||
memory: MemoryEngine instance
|
||||
Returns the servers and apps separately so lifespans can be chained before
|
||||
the middleware wraps the main app.
|
||||
|
||||
Returns:
|
||||
ASGI application
|
||||
Tuple of (multi_bank_server, single_bank_server, multi_bank_app, single_bank_app)
|
||||
"""
|
||||
return MCPMiddleware(None, memory)
|
||||
multi_bank_server = create_mcp_server(memory, multi_bank=True)
|
||||
multi_bank_app = multi_bank_server.http_app(path="/")
|
||||
|
||||
single_bank_server = create_mcp_server(memory, multi_bank=False)
|
||||
single_bank_app = single_bank_server.http_app(path="/")
|
||||
|
||||
return multi_bank_server, single_bank_server, multi_bank_app, single_bank_app
|
||||
|
||||
@@ -66,27 +66,40 @@ ENV_CONSOLIDATION_LLM_TIMEOUT = "HINDSIGHT_API_CONSOLIDATION_LLM_TIMEOUT"
|
||||
ENV_EMBEDDINGS_PROVIDER = "HINDSIGHT_API_EMBEDDINGS_PROVIDER"
|
||||
ENV_EMBEDDINGS_LOCAL_MODEL = "HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL"
|
||||
ENV_EMBEDDINGS_LOCAL_FORCE_CPU = "HINDSIGHT_API_EMBEDDINGS_LOCAL_FORCE_CPU"
|
||||
ENV_EMBEDDINGS_LOCAL_TRUST_REMOTE_CODE = "HINDSIGHT_API_EMBEDDINGS_LOCAL_TRUST_REMOTE_CODE"
|
||||
ENV_EMBEDDINGS_TEI_URL = "HINDSIGHT_API_EMBEDDINGS_TEI_URL"
|
||||
ENV_EMBEDDINGS_OPENAI_API_KEY = "HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY"
|
||||
ENV_EMBEDDINGS_OPENAI_MODEL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL"
|
||||
ENV_EMBEDDINGS_OPENAI_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL"
|
||||
|
||||
ENV_COHERE_API_KEY = "HINDSIGHT_API_COHERE_API_KEY"
|
||||
# Cohere configuration (separate for embeddings and reranker)
|
||||
ENV_EMBEDDINGS_COHERE_API_KEY = "HINDSIGHT_API_EMBEDDINGS_COHERE_API_KEY"
|
||||
ENV_EMBEDDINGS_COHERE_MODEL = "HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL"
|
||||
ENV_EMBEDDINGS_COHERE_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL"
|
||||
ENV_RERANKER_COHERE_API_KEY = "HINDSIGHT_API_RERANKER_COHERE_API_KEY"
|
||||
ENV_RERANKER_COHERE_MODEL = "HINDSIGHT_API_RERANKER_COHERE_MODEL"
|
||||
ENV_RERANKER_COHERE_BASE_URL = "HINDSIGHT_API_RERANKER_COHERE_BASE_URL"
|
||||
|
||||
# LiteLLM gateway configuration (for embeddings and reranker via LiteLLM proxy)
|
||||
# Deprecated: Legacy shared Cohere API key (for backward compatibility)
|
||||
ENV_COHERE_API_KEY = "HINDSIGHT_API_COHERE_API_KEY"
|
||||
|
||||
# LiteLLM configuration (separate for embeddings and reranker)
|
||||
ENV_EMBEDDINGS_LITELLM_API_BASE = "HINDSIGHT_API_EMBEDDINGS_LITELLM_API_BASE"
|
||||
ENV_EMBEDDINGS_LITELLM_API_KEY = "HINDSIGHT_API_EMBEDDINGS_LITELLM_API_KEY"
|
||||
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"
|
||||
|
||||
# Deprecated: Legacy shared LiteLLM config (for backward compatibility)
|
||||
ENV_LITELLM_API_BASE = "HINDSIGHT_API_LITELLM_API_BASE"
|
||||
ENV_LITELLM_API_KEY = "HINDSIGHT_API_LITELLM_API_KEY"
|
||||
ENV_EMBEDDINGS_LITELLM_MODEL = "HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL"
|
||||
ENV_RERANKER_LITELLM_MODEL = "HINDSIGHT_API_RERANKER_LITELLM_MODEL"
|
||||
|
||||
ENV_RERANKER_PROVIDER = "HINDSIGHT_API_RERANKER_PROVIDER"
|
||||
ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
|
||||
ENV_RERANKER_LOCAL_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_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"
|
||||
@@ -190,6 +203,7 @@ DEFAULT_LLM_VERTEXAI_SERVICE_ACCOUNT_KEY = None # Optional, uses ADC if not set
|
||||
DEFAULT_EMBEDDINGS_PROVIDER = "local"
|
||||
DEFAULT_EMBEDDINGS_LOCAL_MODEL = "BAAI/bge-small-en-v1.5"
|
||||
DEFAULT_EMBEDDINGS_LOCAL_FORCE_CPU = False # Force CPU mode for local embeddings (avoids MPS/XPC issues on macOS)
|
||||
DEFAULT_EMBEDDINGS_LOCAL_TRUST_REMOTE_CODE = False # Security: disabled by default, required for some models
|
||||
DEFAULT_EMBEDDINGS_OPENAI_MODEL = "text-embedding-3-small"
|
||||
DEFAULT_EMBEDDING_DIMENSION = 384
|
||||
|
||||
@@ -197,6 +211,9 @@ DEFAULT_RERANKER_PROVIDER = "local"
|
||||
DEFAULT_RERANKER_LOCAL_MODEL = "cross-encoder/ms-marco-MiniLM-L-6-v2"
|
||||
DEFAULT_RERANKER_LOCAL_FORCE_CPU = False # Force CPU mode for local reranker (avoids MPS/XPC issues on macOS)
|
||||
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT = 4 # Limit concurrent CPU-bound reranking to prevent thrashing
|
||||
DEFAULT_RERANKER_LOCAL_TRUST_REMOTE_CODE = (
|
||||
False # Security: disabled by default, required for some models like jina-reranker-v2
|
||||
)
|
||||
DEFAULT_RERANKER_TEI_BATCH_SIZE = 128
|
||||
DEFAULT_RERANKER_TEI_MAX_CONCURRENT = 8
|
||||
DEFAULT_RERANKER_MAX_CANDIDATES = 300
|
||||
@@ -393,20 +410,32 @@ class HindsightConfig:
|
||||
embeddings_provider: str
|
||||
embeddings_local_model: str
|
||||
embeddings_local_force_cpu: bool
|
||||
embeddings_local_trust_remote_code: bool
|
||||
embeddings_tei_url: str | None
|
||||
embeddings_openai_base_url: str | None
|
||||
embeddings_cohere_api_key: str | None
|
||||
embeddings_cohere_model: str
|
||||
embeddings_cohere_base_url: str | None
|
||||
embeddings_litellm_api_base: str
|
||||
embeddings_litellm_api_key: str | None
|
||||
embeddings_litellm_model: str
|
||||
|
||||
# Reranker
|
||||
reranker_provider: str
|
||||
reranker_local_model: str
|
||||
reranker_local_force_cpu: bool
|
||||
reranker_local_max_concurrent: int
|
||||
reranker_local_trust_remote_code: bool
|
||||
reranker_tei_url: str | None
|
||||
reranker_tei_batch_size: int
|
||||
reranker_tei_max_concurrent: int
|
||||
reranker_max_candidates: int
|
||||
reranker_cohere_api_key: str | None
|
||||
reranker_cohere_model: str
|
||||
reranker_cohere_base_url: str | None
|
||||
reranker_litellm_api_base: str
|
||||
reranker_litellm_api_key: str | None
|
||||
reranker_litellm_model: str
|
||||
|
||||
# Server
|
||||
host: str
|
||||
@@ -586,9 +615,21 @@ class HindsightConfig:
|
||||
ENV_EMBEDDINGS_LOCAL_FORCE_CPU, str(DEFAULT_EMBEDDINGS_LOCAL_FORCE_CPU)
|
||||
).lower()
|
||||
in ("true", "1"),
|
||||
embeddings_local_trust_remote_code=os.getenv(
|
||||
ENV_EMBEDDINGS_LOCAL_TRUST_REMOTE_CODE, str(DEFAULT_EMBEDDINGS_LOCAL_TRUST_REMOTE_CODE)
|
||||
).lower()
|
||||
in ("true", "1"),
|
||||
embeddings_tei_url=os.getenv(ENV_EMBEDDINGS_TEI_URL),
|
||||
embeddings_openai_base_url=os.getenv(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None,
|
||||
# Cohere embeddings (with backward-compatible fallback to shared API key)
|
||||
embeddings_cohere_api_key=os.getenv(ENV_EMBEDDINGS_COHERE_API_KEY) or os.getenv(ENV_COHERE_API_KEY),
|
||||
embeddings_cohere_model=os.getenv(ENV_EMBEDDINGS_COHERE_MODEL, DEFAULT_EMBEDDINGS_COHERE_MODEL),
|
||||
embeddings_cohere_base_url=os.getenv(ENV_EMBEDDINGS_COHERE_BASE_URL) or None,
|
||||
# LiteLLM embeddings (with backward-compatible fallback to shared config)
|
||||
embeddings_litellm_api_base=os.getenv(ENV_EMBEDDINGS_LITELLM_API_BASE)
|
||||
or os.getenv(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE),
|
||||
embeddings_litellm_api_key=os.getenv(ENV_EMBEDDINGS_LITELLM_API_KEY) or os.getenv(ENV_LITELLM_API_KEY),
|
||||
embeddings_litellm_model=os.getenv(ENV_EMBEDDINGS_LITELLM_MODEL, DEFAULT_EMBEDDINGS_LITELLM_MODEL),
|
||||
# Reranker
|
||||
reranker_provider=os.getenv(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER),
|
||||
reranker_local_model=os.getenv(ENV_RERANKER_LOCAL_MODEL, DEFAULT_RERANKER_LOCAL_MODEL),
|
||||
@@ -599,13 +640,25 @@ class HindsightConfig:
|
||||
reranker_local_max_concurrent=int(
|
||||
os.getenv(ENV_RERANKER_LOCAL_MAX_CONCURRENT, str(DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT))
|
||||
),
|
||||
reranker_local_trust_remote_code=os.getenv(
|
||||
ENV_RERANKER_LOCAL_TRUST_REMOTE_CODE, str(DEFAULT_RERANKER_LOCAL_TRUST_REMOTE_CODE)
|
||||
).lower()
|
||||
in ("true", "1"),
|
||||
reranker_tei_url=os.getenv(ENV_RERANKER_TEI_URL),
|
||||
reranker_tei_batch_size=int(os.getenv(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE))),
|
||||
reranker_tei_max_concurrent=int(
|
||||
os.getenv(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT))
|
||||
),
|
||||
reranker_max_candidates=int(os.getenv(ENV_RERANKER_MAX_CANDIDATES, str(DEFAULT_RERANKER_MAX_CANDIDATES))),
|
||||
# Cohere reranker (with backward-compatible fallback to shared API key)
|
||||
reranker_cohere_api_key=os.getenv(ENV_RERANKER_COHERE_API_KEY) or os.getenv(ENV_COHERE_API_KEY),
|
||||
reranker_cohere_model=os.getenv(ENV_RERANKER_COHERE_MODEL, DEFAULT_RERANKER_COHERE_MODEL),
|
||||
reranker_cohere_base_url=os.getenv(ENV_RERANKER_COHERE_BASE_URL) or None,
|
||||
# LiteLLM reranker (with backward-compatible fallback to shared config)
|
||||
reranker_litellm_api_base=os.getenv(ENV_RERANKER_LITELLM_API_BASE)
|
||||
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),
|
||||
# Server
|
||||
host=os.getenv(ENV_HOST, DEFAULT_HOST),
|
||||
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
|
||||
|
||||
@@ -24,20 +24,18 @@ from ..config import (
|
||||
DEFAULT_RERANKER_LOCAL_FORCE_CPU,
|
||||
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||
DEFAULT_RERANKER_LOCAL_MODEL,
|
||||
DEFAULT_RERANKER_LOCAL_TRUST_REMOTE_CODE,
|
||||
DEFAULT_RERANKER_PROVIDER,
|
||||
DEFAULT_RERANKER_TEI_BATCH_SIZE,
|
||||
DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
|
||||
ENV_COHERE_API_KEY,
|
||||
ENV_LITELLM_API_BASE,
|
||||
ENV_LITELLM_API_KEY,
|
||||
ENV_RERANKER_COHERE_BASE_URL,
|
||||
ENV_RERANKER_COHERE_API_KEY,
|
||||
ENV_RERANKER_COHERE_MODEL,
|
||||
ENV_RERANKER_FLASHRANK_CACHE_DIR,
|
||||
ENV_RERANKER_FLASHRANK_MODEL,
|
||||
ENV_RERANKER_LITELLM_MODEL,
|
||||
ENV_RERANKER_LOCAL_FORCE_CPU,
|
||||
ENV_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||
ENV_RERANKER_LOCAL_MODEL,
|
||||
ENV_RERANKER_LOCAL_TRUST_REMOTE_CODE,
|
||||
ENV_RERANKER_PROVIDER,
|
||||
ENV_RERANKER_TEI_BATCH_SIZE,
|
||||
ENV_RERANKER_TEI_MAX_CONCURRENT,
|
||||
@@ -102,7 +100,13 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
_executor: ThreadPoolExecutor | None = None
|
||||
_max_concurrent: int = 4 # Limit concurrent CPU-bound reranking calls
|
||||
|
||||
def __init__(self, model_name: str | None = None, max_concurrent: int = 4, force_cpu: bool = False):
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str | None = None,
|
||||
max_concurrent: int = 4,
|
||||
force_cpu: bool = False,
|
||||
trust_remote_code: bool = False,
|
||||
):
|
||||
"""
|
||||
Initialize local SentenceTransformers cross-encoder.
|
||||
|
||||
@@ -113,9 +117,13 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
Higher values may cause CPU thrashing under load.
|
||||
force_cpu: Force CPU mode (avoids MPS/XPC issues on macOS in daemon mode).
|
||||
Default: False
|
||||
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)
|
||||
"""
|
||||
self.model_name = model_name or DEFAULT_RERANKER_LOCAL_MODEL
|
||||
self.force_cpu = force_cpu
|
||||
self.trust_remote_code = trust_remote_code
|
||||
self._model = None
|
||||
LocalSTCrossEncoder._max_concurrent = max_concurrent
|
||||
|
||||
@@ -181,6 +189,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||
self.model_name,
|
||||
device=device,
|
||||
model_kwargs={"low_cpu_mem_usage": False},
|
||||
trust_remote_code=self.trust_remote_code,
|
||||
)
|
||||
finally:
|
||||
# Restore original logging level
|
||||
@@ -847,23 +856,27 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
model_name=config.reranker_local_model,
|
||||
max_concurrent=config.reranker_local_max_concurrent,
|
||||
force_cpu=config.reranker_local_force_cpu,
|
||||
trust_remote_code=config.reranker_local_trust_remote_code,
|
||||
)
|
||||
elif provider == "cohere":
|
||||
api_key = os.environ.get(ENV_COHERE_API_KEY)
|
||||
api_key = config.reranker_cohere_api_key
|
||||
if not api_key:
|
||||
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_RERANKER_PROVIDER} is 'cohere'")
|
||||
model = os.environ.get(ENV_RERANKER_COHERE_MODEL, DEFAULT_RERANKER_COHERE_MODEL)
|
||||
base_url = os.environ.get(ENV_RERANKER_COHERE_BASE_URL) or None
|
||||
return CohereCrossEncoder(api_key=api_key, model=model, base_url=base_url)
|
||||
raise ValueError(f"{ENV_RERANKER_COHERE_API_KEY} is required when {ENV_RERANKER_PROVIDER} is 'cohere'")
|
||||
return CohereCrossEncoder(
|
||||
api_key=api_key,
|
||||
model=config.reranker_cohere_model,
|
||||
base_url=config.reranker_cohere_base_url,
|
||||
)
|
||||
elif provider == "flashrank":
|
||||
model = os.environ.get(ENV_RERANKER_FLASHRANK_MODEL, DEFAULT_RERANKER_FLASHRANK_MODEL)
|
||||
cache_dir = os.environ.get(ENV_RERANKER_FLASHRANK_CACHE_DIR, DEFAULT_RERANKER_FLASHRANK_CACHE_DIR)
|
||||
return FlashRankCrossEncoder(model_name=model, cache_dir=cache_dir)
|
||||
elif provider == "litellm":
|
||||
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
|
||||
api_key = os.environ.get(ENV_LITELLM_API_KEY)
|
||||
model = os.environ.get(ENV_RERANKER_LITELLM_MODEL, DEFAULT_RERANKER_LITELLM_MODEL)
|
||||
return LiteLLMCrossEncoder(api_base=api_base, api_key=api_key, model=model)
|
||||
return LiteLLMCrossEncoder(
|
||||
api_base=config.reranker_litellm_api_base,
|
||||
api_key=config.reranker_litellm_api_key,
|
||||
model=config.reranker_litellm_model,
|
||||
)
|
||||
elif provider == "rrf":
|
||||
return RRFPassthroughCrossEncoder()
|
||||
else:
|
||||
|
||||
@@ -21,22 +21,19 @@ from ..config import (
|
||||
DEFAULT_EMBEDDINGS_LITELLM_MODEL,
|
||||
DEFAULT_EMBEDDINGS_LOCAL_FORCE_CPU,
|
||||
DEFAULT_EMBEDDINGS_LOCAL_MODEL,
|
||||
DEFAULT_EMBEDDINGS_LOCAL_TRUST_REMOTE_CODE,
|
||||
DEFAULT_EMBEDDINGS_OPENAI_MODEL,
|
||||
DEFAULT_EMBEDDINGS_PROVIDER,
|
||||
DEFAULT_LITELLM_API_BASE,
|
||||
ENV_COHERE_API_KEY,
|
||||
ENV_EMBEDDINGS_COHERE_BASE_URL,
|
||||
ENV_EMBEDDINGS_COHERE_MODEL,
|
||||
ENV_EMBEDDINGS_LITELLM_MODEL,
|
||||
ENV_EMBEDDINGS_COHERE_API_KEY,
|
||||
ENV_EMBEDDINGS_LOCAL_FORCE_CPU,
|
||||
ENV_EMBEDDINGS_LOCAL_MODEL,
|
||||
ENV_EMBEDDINGS_LOCAL_TRUST_REMOTE_CODE,
|
||||
ENV_EMBEDDINGS_OPENAI_API_KEY,
|
||||
ENV_EMBEDDINGS_OPENAI_BASE_URL,
|
||||
ENV_EMBEDDINGS_OPENAI_MODEL,
|
||||
ENV_EMBEDDINGS_PROVIDER,
|
||||
ENV_EMBEDDINGS_TEI_URL,
|
||||
ENV_LITELLM_API_BASE,
|
||||
ENV_LITELLM_API_KEY,
|
||||
ENV_LLM_API_KEY,
|
||||
)
|
||||
|
||||
@@ -95,7 +92,7 @@ class LocalSTEmbeddings(Embeddings):
|
||||
The embedding dimension is auto-detected from the model.
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str | None = None, force_cpu: bool = False):
|
||||
def __init__(self, model_name: str | None = None, force_cpu: bool = False, trust_remote_code: bool = False):
|
||||
"""
|
||||
Initialize local SentenceTransformers embeddings.
|
||||
|
||||
@@ -104,9 +101,13 @@ class LocalSTEmbeddings(Embeddings):
|
||||
Default: BAAI/bge-small-en-v1.5
|
||||
force_cpu: Force CPU mode (avoids MPS/XPC issues on macOS in daemon mode).
|
||||
Default: False
|
||||
trust_remote_code: Allow loading models with custom code (security risk).
|
||||
Required for some models with custom architectures.
|
||||
Default: False (disabled for security)
|
||||
"""
|
||||
self.model_name = model_name or DEFAULT_EMBEDDINGS_LOCAL_MODEL
|
||||
self.force_cpu = force_cpu
|
||||
self.trust_remote_code = trust_remote_code
|
||||
self._model = None
|
||||
self._dimension: int | None = None
|
||||
|
||||
@@ -176,6 +177,7 @@ class LocalSTEmbeddings(Embeddings):
|
||||
self.model_name,
|
||||
device=device,
|
||||
model_kwargs={"low_cpu_mem_usage": False},
|
||||
trust_remote_code=self.trust_remote_code,
|
||||
)
|
||||
finally:
|
||||
# Restore original logging level
|
||||
@@ -741,6 +743,7 @@ def create_embeddings_from_env() -> Embeddings:
|
||||
return LocalSTEmbeddings(
|
||||
model_name=config.embeddings_local_model,
|
||||
force_cpu=config.embeddings_local_force_cpu,
|
||||
trust_remote_code=config.embeddings_local_trust_remote_code,
|
||||
)
|
||||
elif provider == "openai":
|
||||
# Use dedicated embeddings API key, or fall back to LLM API key
|
||||
@@ -754,17 +757,20 @@ def create_embeddings_from_env() -> Embeddings:
|
||||
base_url = os.environ.get(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None
|
||||
return OpenAIEmbeddings(api_key=api_key, model=model, base_url=base_url)
|
||||
elif provider == "cohere":
|
||||
api_key = os.environ.get(ENV_COHERE_API_KEY)
|
||||
api_key = config.embeddings_cohere_api_key
|
||||
if not api_key:
|
||||
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_EMBEDDINGS_PROVIDER} is 'cohere'")
|
||||
model = os.environ.get(ENV_EMBEDDINGS_COHERE_MODEL, DEFAULT_EMBEDDINGS_COHERE_MODEL)
|
||||
base_url = os.environ.get(ENV_EMBEDDINGS_COHERE_BASE_URL) or None
|
||||
return CohereEmbeddings(api_key=api_key, model=model, base_url=base_url)
|
||||
raise ValueError(f"{ENV_EMBEDDINGS_COHERE_API_KEY} is required when {ENV_EMBEDDINGS_PROVIDER} is 'cohere'")
|
||||
return CohereEmbeddings(
|
||||
api_key=api_key,
|
||||
model=config.embeddings_cohere_model,
|
||||
base_url=config.embeddings_cohere_base_url,
|
||||
)
|
||||
elif provider == "litellm":
|
||||
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
|
||||
api_key = os.environ.get(ENV_LITELLM_API_KEY)
|
||||
model = os.environ.get(ENV_EMBEDDINGS_LITELLM_MODEL, DEFAULT_EMBEDDINGS_LITELLM_MODEL)
|
||||
return LiteLLMEmbeddings(api_base=api_base, api_key=api_key, model=model)
|
||||
return LiteLLMEmbeddings(
|
||||
api_base=config.embeddings_litellm_api_base,
|
||||
api_key=config.embeddings_litellm_api_key,
|
||||
model=config.embeddings_litellm_model,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei', 'openai', 'cohere', 'litellm'"
|
||||
|
||||
@@ -545,16 +545,19 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
f"[BATCH_RETAIN_TASK] Starting background batch retain for bank_id={bank_id}, {len(contents)} items"
|
||||
)
|
||||
|
||||
# Restore tenant_id/api_key_id from task payload so downstream operations
|
||||
# (e.g., consolidation and mental model refreshes) can attribute usage.
|
||||
# Restore tenant_id/api_key_id from task payload so extensions
|
||||
# (e.g., operation validators) can attribute the operation correctly.
|
||||
# internal=True to skip extension auth (worker has no API key),
|
||||
# user_initiated=True so extensions know this originated from a user request.
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
internal_context = RequestContext(
|
||||
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.retain_batch_async(bank_id=bank_id, contents=contents, request_context=internal_context)
|
||||
await self.retain_batch_async(bank_id=bank_id, contents=contents, request_context=context)
|
||||
|
||||
logger.info(f"[BATCH_RETAIN_TASK] Completed background batch retain for bank_id={bank_id}")
|
||||
|
||||
@@ -1484,6 +1487,9 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
unit_ids=result,
|
||||
success=True,
|
||||
error=None,
|
||||
llm_input_tokens=total_usage.input_tokens,
|
||||
llm_output_tokens=total_usage.output_tokens,
|
||||
llm_total_tokens=total_usage.total_tokens,
|
||||
)
|
||||
try:
|
||||
await self._operation_validator.on_retain_complete(result_ctx)
|
||||
@@ -4690,6 +4696,18 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
Pinned mental model dict or None if not found
|
||||
"""
|
||||
await self._authenticate_tenant(request_context)
|
||||
|
||||
# Pre-operation validation (credit check / usage metering)
|
||||
if self._operation_validator:
|
||||
from hindsight_api.extensions.operation_validator import MentalModelGetContext
|
||||
|
||||
ctx = MentalModelGetContext(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=request_context,
|
||||
)
|
||||
await self._validate_operation(self._operation_validator.validate_mental_model_get(ctx))
|
||||
|
||||
pool = await self._get_pool()
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
@@ -4705,7 +4723,28 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
mental_model_id,
|
||||
)
|
||||
|
||||
return self._row_to_mental_model(row) if row else None
|
||||
result = self._row_to_mental_model(row) if row else None
|
||||
|
||||
# Post-operation hook (usage recording)
|
||||
if result and self._operation_validator:
|
||||
from hindsight_api.extensions.operation_validator import MentalModelGetResult
|
||||
|
||||
content = result.get("content", "")
|
||||
output_tokens = len(content) // 4 if content else 0
|
||||
|
||||
result_ctx = MentalModelGetResult(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=request_context,
|
||||
output_tokens=output_tokens,
|
||||
success=True,
|
||||
)
|
||||
try:
|
||||
await self._operation_validator.on_mental_model_get_complete(result_ctx)
|
||||
except Exception as hook_err:
|
||||
logger.warning(f"Post-mental-model-get hook error (non-fatal): {hook_err}")
|
||||
|
||||
return result
|
||||
|
||||
async def create_mental_model(
|
||||
self,
|
||||
@@ -5696,6 +5735,17 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
"""
|
||||
await self._authenticate_tenant(request_context)
|
||||
|
||||
# Pre-operation validation (credit check)
|
||||
if self._operation_validator:
|
||||
from hindsight_api.extensions.operation_validator import MentalModelRefreshContext
|
||||
|
||||
ctx = MentalModelRefreshContext(
|
||||
bank_id=bank_id,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=request_context,
|
||||
)
|
||||
await self._validate_operation(self._operation_validator.validate_mental_model_refresh(ctx))
|
||||
|
||||
# Verify mental model exists
|
||||
mental_model = await self.get_mental_model(bank_id, mental_model_id, request_context=request_context)
|
||||
if not mental_model:
|
||||
|
||||
@@ -132,6 +132,10 @@ class RetainResult:
|
||||
unit_ids: list[list[str]] # List of unit IDs per content item
|
||||
success: bool = True
|
||||
error: str | None = None
|
||||
# Actual LLM token usage (populated by engine when available)
|
||||
llm_input_tokens: int | None = None
|
||||
llm_output_tokens: int | None = None
|
||||
llm_total_tokens: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -197,18 +197,30 @@ def main():
|
||||
embeddings_provider=config.embeddings_provider,
|
||||
embeddings_local_model=config.embeddings_local_model,
|
||||
embeddings_local_force_cpu=config.embeddings_local_force_cpu,
|
||||
embeddings_local_trust_remote_code=config.embeddings_local_trust_remote_code,
|
||||
embeddings_tei_url=config.embeddings_tei_url,
|
||||
embeddings_openai_base_url=config.embeddings_openai_base_url,
|
||||
embeddings_cohere_api_key=config.embeddings_cohere_api_key,
|
||||
embeddings_cohere_model=config.embeddings_cohere_model,
|
||||
embeddings_cohere_base_url=config.embeddings_cohere_base_url,
|
||||
embeddings_litellm_api_base=config.embeddings_litellm_api_base,
|
||||
embeddings_litellm_api_key=config.embeddings_litellm_api_key,
|
||||
embeddings_litellm_model=config.embeddings_litellm_model,
|
||||
reranker_provider=config.reranker_provider,
|
||||
reranker_local_model=config.reranker_local_model,
|
||||
reranker_local_force_cpu=config.reranker_local_force_cpu,
|
||||
reranker_local_max_concurrent=config.reranker_local_max_concurrent,
|
||||
reranker_local_trust_remote_code=config.reranker_local_trust_remote_code,
|
||||
reranker_tei_url=config.reranker_tei_url,
|
||||
reranker_tei_batch_size=config.reranker_tei_batch_size,
|
||||
reranker_tei_max_concurrent=config.reranker_tei_max_concurrent,
|
||||
reranker_max_candidates=config.reranker_max_candidates,
|
||||
reranker_cohere_api_key=config.reranker_cohere_api_key,
|
||||
reranker_cohere_model=config.reranker_cohere_model,
|
||||
reranker_cohere_base_url=config.reranker_cohere_base_url,
|
||||
reranker_litellm_api_base=config.reranker_litellm_api_base,
|
||||
reranker_litellm_api_key=config.reranker_litellm_api_key,
|
||||
reranker_litellm_model=config.reranker_litellm_model,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
log_level=args.log_level,
|
||||
|
||||
@@ -127,7 +127,19 @@ def register_mcp_tools(
|
||||
memory: MemoryEngine instance
|
||||
config: Tool configuration
|
||||
"""
|
||||
tools_to_register = config.tools or {"retain", "recall", "reflect", "list_banks", "create_bank"}
|
||||
tools_to_register = config.tools or {
|
||||
"retain",
|
||||
"recall",
|
||||
"reflect",
|
||||
"list_banks",
|
||||
"create_bank",
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"create_mental_model",
|
||||
"update_mental_model",
|
||||
"delete_mental_model",
|
||||
"refresh_mental_model",
|
||||
}
|
||||
|
||||
if "retain" in tools_to_register:
|
||||
_register_retain(mcp, memory, config)
|
||||
@@ -144,6 +156,25 @@ def register_mcp_tools(
|
||||
if "create_bank" in tools_to_register:
|
||||
_register_create_bank(mcp, memory, config)
|
||||
|
||||
# Mental model tools
|
||||
if "list_mental_models" in tools_to_register:
|
||||
_register_list_mental_models(mcp, memory, config)
|
||||
|
||||
if "get_mental_model" in tools_to_register:
|
||||
_register_get_mental_model(mcp, memory, config)
|
||||
|
||||
if "create_mental_model" in tools_to_register:
|
||||
_register_create_mental_model(mcp, memory, config)
|
||||
|
||||
if "update_mental_model" in tools_to_register:
|
||||
_register_update_mental_model(mcp, memory, config)
|
||||
|
||||
if "delete_mental_model" in tools_to_register:
|
||||
_register_delete_mental_model(mcp, memory, config)
|
||||
|
||||
if "refresh_mental_model" in tools_to_register:
|
||||
_register_refresh_mental_model(mcp, memory, config)
|
||||
|
||||
|
||||
def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the retain tool."""
|
||||
@@ -519,3 +550,567 @@ def _register_create_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating bank: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}"}}'
|
||||
|
||||
|
||||
def _validate_mental_model_inputs(
|
||||
name: str | None = None, source_query: str | None = None, max_tokens: int | None = None
|
||||
) -> str | None:
|
||||
"""Validate mental model inputs, returning an error message or None if valid."""
|
||||
if name is not None and not name.strip():
|
||||
return "name cannot be empty"
|
||||
if source_query is not None and not source_query.strip():
|
||||
return "source_query cannot be empty"
|
||||
if max_tokens is not None and (max_tokens < 256 or max_tokens > 8192):
|
||||
return f"max_tokens must be between 256 and 8192, got {max_tokens}"
|
||||
return None
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# MENTAL MODEL TOOLS
|
||||
# =========================================================================
|
||||
|
||||
|
||||
def _register_list_mental_models(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the list_mental_models tool."""
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool()
|
||||
async def list_mental_models(
|
||||
tags: list[str] | None = None,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
List mental models (pinned reflections) for a memory bank.
|
||||
|
||||
Mental models are living documents that stay current by periodically re-running
|
||||
a source query through reflect. Use them to maintain up-to-date summaries,
|
||||
preferences, or synthesized knowledge.
|
||||
|
||||
Args:
|
||||
tags: Optional tags to filter by (returns models matching any tag)
|
||||
bank_id: Optional bank to list from (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return '{"error": "No bank_id configured", "items": []}'
|
||||
|
||||
models = await memory.list_mental_models(
|
||||
bank_id=target_bank,
|
||||
tags=tags,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
return json.dumps({"items": models}, indent=2, default=str)
|
||||
except Exception as e:
|
||||
logger.error(f"Error listing mental models: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}", "items": []}}'
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool()
|
||||
async def list_mental_models(
|
||||
tags: list[str] | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
List mental models (pinned reflections) for this memory bank.
|
||||
|
||||
Mental models are living documents that stay current by periodically re-running
|
||||
a source query through reflect. Use them to maintain up-to-date summaries,
|
||||
preferences, or synthesized knowledge.
|
||||
|
||||
Args:
|
||||
tags: Optional tags to filter by (returns models matching any tag)
|
||||
"""
|
||||
try:
|
||||
target_bank = config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return {"error": "No bank_id configured", "items": []}
|
||||
|
||||
models = await memory.list_mental_models(
|
||||
bank_id=target_bank,
|
||||
tags=tags,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
return {"items": models}
|
||||
except Exception as e:
|
||||
logger.error(f"Error listing mental models: {e}", exc_info=True)
|
||||
return {"error": str(e), "items": []}
|
||||
|
||||
|
||||
def _register_get_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the get_mental_model tool."""
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool()
|
||||
async def get_mental_model(
|
||||
mental_model_id: str,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get a specific mental model by ID.
|
||||
|
||||
Returns the full mental model including its generated content, source query,
|
||||
and metadata. Use list_mental_models first to discover available model IDs.
|
||||
|
||||
Args:
|
||||
mental_model_id: The ID of the mental model to retrieve
|
||||
bank_id: Optional bank (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return '{"error": "No bank_id configured"}'
|
||||
|
||||
model = await memory.get_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
if model is None:
|
||||
return json.dumps({"error": f"Mental model '{mental_model_id}' not found in bank '{target_bank}'"})
|
||||
return json.dumps(model, indent=2, default=str)
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting mental model: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}"}}'
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool()
|
||||
async def get_mental_model(
|
||||
mental_model_id: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Get a specific mental model by ID.
|
||||
|
||||
Returns the full mental model including its generated content, source query,
|
||||
and metadata. Use list_mental_models first to discover available model IDs.
|
||||
|
||||
Args:
|
||||
mental_model_id: The ID of the mental model to retrieve
|
||||
"""
|
||||
try:
|
||||
target_bank = config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return {"error": "No bank_id configured"}
|
||||
|
||||
model = await memory.get_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
if model is None:
|
||||
return {"error": f"Mental model '{mental_model_id}' not found in bank '{target_bank}'"}
|
||||
return model
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting mental model: {e}", exc_info=True)
|
||||
return {"error": str(e)}
|
||||
|
||||
|
||||
def _register_create_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the create_mental_model tool."""
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool()
|
||||
async def create_mental_model(
|
||||
name: str,
|
||||
source_query: str,
|
||||
mental_model_id: str | None = None,
|
||||
tags: list[str] | None = None,
|
||||
max_tokens: int = 2048,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Create a new mental model (pinned reflection).
|
||||
|
||||
A mental model is a living document generated by running the source_query through
|
||||
reflect. The content is auto-generated asynchronously - use the returned operation_id
|
||||
to track progress.
|
||||
|
||||
EXAMPLES:
|
||||
- name="Coding Preferences", source_query="What coding patterns and tools does the user prefer?"
|
||||
- name="Project Goals", source_query="What are the user's current project goals and priorities?"
|
||||
- name="Communication Style", source_query="How does the user prefer to communicate?"
|
||||
|
||||
Args:
|
||||
name: Human-readable name for the mental model
|
||||
source_query: The query to run through reflect to generate content
|
||||
mental_model_id: Optional custom ID (alphanumeric lowercase with hyphens). Auto-generated if not provided.
|
||||
tags: Optional tags for scoped visibility filtering
|
||||
max_tokens: Maximum tokens for generated content (256-8192, default: 2048)
|
||||
bank_id: Optional bank (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return '{"error": "No bank_id configured"}'
|
||||
|
||||
validation_error = _validate_mental_model_inputs(
|
||||
name=name, source_query=source_query, max_tokens=max_tokens
|
||||
)
|
||||
if validation_error:
|
||||
return json.dumps({"error": validation_error})
|
||||
|
||||
request_context = _get_request_context(config)
|
||||
|
||||
# Create with placeholder content
|
||||
model = await memory.create_mental_model(
|
||||
bank_id=target_bank,
|
||||
name=name,
|
||||
source_query=source_query,
|
||||
content="Generating content...",
|
||||
mental_model_id=mental_model_id,
|
||||
tags=tags,
|
||||
max_tokens=max_tokens,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
# Schedule async refresh to generate actual content
|
||||
result = await memory.submit_async_refresh_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=model["id"],
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
return json.dumps(
|
||||
{
|
||||
"mental_model_id": model["id"],
|
||||
"operation_id": result["operation_id"],
|
||||
"status": "created",
|
||||
"message": f"Mental model '{name}' created. Content is being generated asynchronously.",
|
||||
}
|
||||
)
|
||||
except ValueError as e:
|
||||
return json.dumps({"error": str(e)})
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating mental model: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}"}}'
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool()
|
||||
async def create_mental_model(
|
||||
name: str,
|
||||
source_query: str,
|
||||
mental_model_id: str | None = None,
|
||||
tags: list[str] | None = None,
|
||||
max_tokens: int = 2048,
|
||||
) -> dict:
|
||||
"""
|
||||
Create a new mental model (pinned reflection).
|
||||
|
||||
A mental model is a living document generated by running the source_query through
|
||||
reflect. The content is auto-generated asynchronously - use the returned operation_id
|
||||
to track progress.
|
||||
|
||||
EXAMPLES:
|
||||
- name="Coding Preferences", source_query="What coding patterns and tools does the user prefer?"
|
||||
- name="Project Goals", source_query="What are the user's current project goals and priorities?"
|
||||
- name="Communication Style", source_query="How does the user prefer to communicate?"
|
||||
|
||||
Args:
|
||||
name: Human-readable name for the mental model
|
||||
source_query: The query to run through reflect to generate content
|
||||
mental_model_id: Optional custom ID (alphanumeric lowercase with hyphens). Auto-generated if not provided.
|
||||
tags: Optional tags for scoped visibility filtering
|
||||
max_tokens: Maximum tokens for generated content (256-8192, default: 2048)
|
||||
"""
|
||||
try:
|
||||
target_bank = config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return {"error": "No bank_id configured"}
|
||||
|
||||
validation_error = _validate_mental_model_inputs(
|
||||
name=name, source_query=source_query, max_tokens=max_tokens
|
||||
)
|
||||
if validation_error:
|
||||
return {"error": validation_error}
|
||||
|
||||
request_context = _get_request_context(config)
|
||||
|
||||
model = await memory.create_mental_model(
|
||||
bank_id=target_bank,
|
||||
name=name,
|
||||
source_query=source_query,
|
||||
content="Generating content...",
|
||||
mental_model_id=mental_model_id,
|
||||
tags=tags,
|
||||
max_tokens=max_tokens,
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
result = await memory.submit_async_refresh_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=model["id"],
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
return {
|
||||
"mental_model_id": model["id"],
|
||||
"operation_id": result["operation_id"],
|
||||
"status": "created",
|
||||
"message": f"Mental model '{name}' created. Content is being generated asynchronously.",
|
||||
}
|
||||
except ValueError as e:
|
||||
return {"error": str(e)}
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating mental model: {e}", exc_info=True)
|
||||
return {"error": str(e)}
|
||||
|
||||
|
||||
def _register_update_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the update_mental_model tool."""
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool()
|
||||
async def update_mental_model(
|
||||
mental_model_id: str,
|
||||
name: str | None = None,
|
||||
source_query: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
tags: list[str] | None = None,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Update a mental model's metadata.
|
||||
|
||||
Changes the name, source query, or tags of an existing mental model.
|
||||
To regenerate the content, use refresh_mental_model after updating the source query.
|
||||
|
||||
Args:
|
||||
mental_model_id: The ID of the mental model to update
|
||||
name: New name (leave None to keep current)
|
||||
source_query: New source query (leave None to keep current)
|
||||
max_tokens: New max tokens for content generation (256-8192, leave None to keep current)
|
||||
tags: New tags (leave None to keep current)
|
||||
bank_id: Optional bank (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return '{"error": "No bank_id configured"}'
|
||||
|
||||
validation_error = _validate_mental_model_inputs(
|
||||
name=name, source_query=source_query, max_tokens=max_tokens
|
||||
)
|
||||
if validation_error:
|
||||
return json.dumps({"error": validation_error})
|
||||
|
||||
model = await memory.update_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=mental_model_id,
|
||||
name=name,
|
||||
source_query=source_query,
|
||||
max_tokens=max_tokens,
|
||||
tags=tags,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
if model is None:
|
||||
return json.dumps({"error": f"Mental model '{mental_model_id}' not found in bank '{target_bank}'"})
|
||||
return json.dumps(model, indent=2, default=str)
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating mental model: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}"}}'
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool()
|
||||
async def update_mental_model(
|
||||
mental_model_id: str,
|
||||
name: str | None = None,
|
||||
source_query: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
tags: list[str] | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Update a mental model's metadata.
|
||||
|
||||
Changes the name, source query, or tags of an existing mental model.
|
||||
To regenerate the content, use refresh_mental_model after updating the source query.
|
||||
|
||||
Args:
|
||||
mental_model_id: The ID of the mental model to update
|
||||
name: New name (leave None to keep current)
|
||||
source_query: New source query (leave None to keep current)
|
||||
max_tokens: New max tokens for content generation (256-8192, leave None to keep current)
|
||||
tags: New tags (leave None to keep current)
|
||||
"""
|
||||
try:
|
||||
target_bank = config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return {"error": "No bank_id configured"}
|
||||
|
||||
validation_error = _validate_mental_model_inputs(
|
||||
name=name, source_query=source_query, max_tokens=max_tokens
|
||||
)
|
||||
if validation_error:
|
||||
return {"error": validation_error}
|
||||
|
||||
model = await memory.update_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=mental_model_id,
|
||||
name=name,
|
||||
source_query=source_query,
|
||||
max_tokens=max_tokens,
|
||||
tags=tags,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
if model is None:
|
||||
return {"error": f"Mental model '{mental_model_id}' not found in bank '{target_bank}'"}
|
||||
return model
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating mental model: {e}", exc_info=True)
|
||||
return {"error": str(e)}
|
||||
|
||||
|
||||
def _register_delete_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the delete_mental_model tool."""
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool()
|
||||
async def delete_mental_model(
|
||||
mental_model_id: str,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Delete a mental model.
|
||||
|
||||
Permanently removes a mental model and its generated content.
|
||||
|
||||
Args:
|
||||
mental_model_id: The ID of the mental model to delete
|
||||
bank_id: Optional bank (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return '{"error": "No bank_id configured"}'
|
||||
|
||||
deleted = await memory.delete_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
if not deleted:
|
||||
return json.dumps({"error": f"Mental model '{mental_model_id}' not found in bank '{target_bank}'"})
|
||||
return json.dumps({"status": "deleted", "mental_model_id": mental_model_id})
|
||||
except Exception as e:
|
||||
logger.error(f"Error deleting mental model: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}"}}'
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool()
|
||||
async def delete_mental_model(
|
||||
mental_model_id: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Delete a mental model.
|
||||
|
||||
Permanently removes a mental model and its generated content.
|
||||
|
||||
Args:
|
||||
mental_model_id: The ID of the mental model to delete
|
||||
"""
|
||||
try:
|
||||
target_bank = config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return {"error": "No bank_id configured"}
|
||||
|
||||
deleted = await memory.delete_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
if not deleted:
|
||||
return {"error": f"Mental model '{mental_model_id}' not found in bank '{target_bank}'"}
|
||||
return {"status": "deleted", "mental_model_id": mental_model_id}
|
||||
except Exception as e:
|
||||
logger.error(f"Error deleting mental model: {e}", exc_info=True)
|
||||
return {"error": str(e)}
|
||||
|
||||
|
||||
def _register_refresh_mental_model(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
|
||||
"""Register the refresh_mental_model tool."""
|
||||
|
||||
if config.include_bank_id_param:
|
||||
|
||||
@mcp.tool()
|
||||
async def refresh_mental_model(
|
||||
mental_model_id: str,
|
||||
bank_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Refresh a mental model by re-running its source query.
|
||||
|
||||
Schedules an async task to re-run the source query through reflect and update the
|
||||
mental model's content with fresh results. Use this after adding new memories or
|
||||
when the mental model's content may be stale.
|
||||
|
||||
Args:
|
||||
mental_model_id: The ID of the mental model to refresh
|
||||
bank_id: Optional bank (defaults to session bank). Use for cross-bank operations.
|
||||
"""
|
||||
try:
|
||||
target_bank = bank_id or config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return '{"error": "No bank_id configured"}'
|
||||
|
||||
result = await memory.submit_async_refresh_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
return json.dumps(
|
||||
{
|
||||
"operation_id": result["operation_id"],
|
||||
"status": "queued",
|
||||
"message": f"Refresh queued for mental model '{mental_model_id}'.",
|
||||
}
|
||||
)
|
||||
except ValueError as e:
|
||||
return json.dumps({"error": str(e)})
|
||||
except Exception as e:
|
||||
logger.error(f"Error refreshing mental model: {e}", exc_info=True)
|
||||
return f'{{"error": "{e}"}}'
|
||||
|
||||
else:
|
||||
|
||||
@mcp.tool()
|
||||
async def refresh_mental_model(
|
||||
mental_model_id: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Refresh a mental model by re-running its source query.
|
||||
|
||||
Schedules an async task to re-run the source query through reflect and update the
|
||||
mental model's content with fresh results. Use this after adding new memories or
|
||||
when the mental model's content may be stale.
|
||||
|
||||
Args:
|
||||
mental_model_id: The ID of the mental model to refresh
|
||||
"""
|
||||
try:
|
||||
target_bank = config.bank_id_resolver()
|
||||
if target_bank is None:
|
||||
return {"error": "No bank_id configured"}
|
||||
|
||||
result = await memory.submit_async_refresh_mental_model(
|
||||
bank_id=target_bank,
|
||||
mental_model_id=mental_model_id,
|
||||
request_context=_get_request_context(config),
|
||||
)
|
||||
return {
|
||||
"operation_id": result["operation_id"],
|
||||
"status": "queued",
|
||||
"message": f"Refresh queued for mental model '{mental_model_id}'.",
|
||||
}
|
||||
except ValueError as e:
|
||||
return {"error": str(e)}
|
||||
except Exception as e:
|
||||
logger.error(f"Error refreshing mental model: {e}", exc_info=True)
|
||||
return {"error": str(e)}
|
||||
|
||||
@@ -20,7 +20,8 @@ class RequestContext:
|
||||
api_key: str | None = None
|
||||
api_key_id: str | None = None # UUID of the API key used for authentication
|
||||
tenant_id: str | None = None # Tenant identifier (set by extension after auth)
|
||||
internal: bool = False # True for background/internal operations (not user-visible)
|
||||
internal: bool = False # True for background/internal operations (skips extension auth)
|
||||
user_initiated: bool = False # True for async operations that originated from a user request
|
||||
|
||||
|
||||
from pgvector.sqlalchemy import Vector
|
||||
|
||||
@@ -353,6 +353,14 @@ class TestOperationHooksParameters:
|
||||
assert post_result.error is None
|
||||
assert post_result.unit_ids == result # Should match the return value
|
||||
|
||||
# Verify actual LLM token usage is populated
|
||||
assert post_result.llm_input_tokens is not None
|
||||
assert post_result.llm_input_tokens > 0
|
||||
assert post_result.llm_output_tokens is not None
|
||||
assert post_result.llm_output_tokens > 0
|
||||
assert post_result.llm_total_tokens is not None
|
||||
assert post_result.llm_total_tokens == post_result.llm_input_tokens + post_result.llm_output_tokens
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_pre_hook_receives_all_parameters(self, memory_with_tracking_validator):
|
||||
"""Pre-recall hook receives all user-provided parameters."""
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""Integration test for MCP endpoint routing.
|
||||
|
||||
This test verifies that /mcp/ and /mcp/{bank_id}/ expose different tool sets.
|
||||
This test verifies that /mcp/ and /mcp/{bank_id}/ expose different tool sets,
|
||||
and that URLs with or without trailing slashes both work (no 307 redirect).
|
||||
"""
|
||||
|
||||
import httpx
|
||||
@@ -39,12 +40,18 @@ async def test_mcp_endpoint_routing_integration(memory):
|
||||
|
||||
multi_tools = {t.name for t in multi_result.tools}
|
||||
|
||||
# Multi-bank should have all tools including bank management
|
||||
# Multi-bank should have all tools including bank management and mental models
|
||||
assert "retain" in multi_tools
|
||||
assert "recall" in multi_tools
|
||||
assert "reflect" in multi_tools
|
||||
assert "list_banks" in multi_tools, "Multi-bank should expose list_banks"
|
||||
assert "create_bank" in multi_tools, "Multi-bank should expose create_bank"
|
||||
assert "list_mental_models" in multi_tools, "Multi-bank should expose list_mental_models"
|
||||
assert "create_mental_model" in multi_tools, "Multi-bank should expose create_mental_model"
|
||||
assert "get_mental_model" in multi_tools, "Multi-bank should expose get_mental_model"
|
||||
assert "update_mental_model" in multi_tools, "Multi-bank should expose update_mental_model"
|
||||
assert "delete_mental_model" in multi_tools, "Multi-bank should expose delete_mental_model"
|
||||
assert "refresh_mental_model" in multi_tools, "Multi-bank should expose refresh_mental_model"
|
||||
|
||||
# Multi-bank retain should have bank_id parameter
|
||||
retain_tool = next((t for t in multi_result.tools if t.name == "retain"), None)
|
||||
@@ -64,10 +71,12 @@ async def test_mcp_endpoint_routing_integration(memory):
|
||||
|
||||
single_tools = {t.name for t in single_result.tools}
|
||||
|
||||
# Single-bank should only have scoped tools (no bank management)
|
||||
# Single-bank should have scoped tools including mental models (no bank management)
|
||||
assert "retain" in single_tools
|
||||
assert "recall" in single_tools
|
||||
assert "reflect" in single_tools
|
||||
assert "list_mental_models" in single_tools, "Single-bank should expose list_mental_models"
|
||||
assert "create_mental_model" in single_tools, "Single-bank should expose create_mental_model"
|
||||
assert "list_banks" not in single_tools, "Single-bank should NOT expose list_banks"
|
||||
assert "create_bank" not in single_tools, "Single-bank should NOT expose create_bank"
|
||||
|
||||
@@ -76,3 +85,196 @@ async def test_mcp_endpoint_routing_integration(memory):
|
||||
assert retain_tool is not None
|
||||
single_params = set(retain_tool.inputSchema.get("properties", {}).keys())
|
||||
assert "bank_id" not in single_params, "Single-bank retain should NOT have bank_id parameter"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_no_trailing_slash_works(memory):
|
||||
"""Test that /mcp (no trailing slash) discovers tools without 307 redirect.
|
||||
|
||||
Starlette's Mount class redirects /mcp to /mcp/ with a 307 Temporary Redirect.
|
||||
Many MCP clients don't follow POST redirects, causing 0 tools to be discovered.
|
||||
MCPMiddleware wraps the app directly (no Mount), so the redirect never happens.
|
||||
"""
|
||||
from hindsight_api.api import create_app
|
||||
|
||||
app = create_app(memory, mcp_api_enabled=True, initialize_memory=False)
|
||||
|
||||
async with app.router.lifespan_context(app):
|
||||
from httpx import ASGITransport
|
||||
|
||||
async with httpx.AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as http_client:
|
||||
# /mcp (no slash) should work the same as /mcp/
|
||||
async with streamable_http_client("http://test/mcp", http_client=http_client) as (
|
||||
read_stream,
|
||||
write_stream,
|
||||
_,
|
||||
):
|
||||
async with ClientSession(read_stream, write_stream) as session:
|
||||
await session.initialize()
|
||||
result = await session.list_tools()
|
||||
|
||||
tools = {t.name for t in result.tools}
|
||||
assert len(tools) >= 11, f"Expected at least 11 tools from /mcp, got {len(tools)}: {tools}"
|
||||
assert "retain" in tools
|
||||
assert "recall" in tools
|
||||
assert "list_banks" in tools
|
||||
|
||||
# /mcp/my-bank (single-bank, no slash) should also work
|
||||
async with streamable_http_client("http://test/mcp/my-bank", http_client=http_client) as (
|
||||
read_stream,
|
||||
write_stream,
|
||||
_,
|
||||
):
|
||||
async with ClientSession(read_stream, write_stream) as session:
|
||||
await session.initialize()
|
||||
result = await session.list_tools()
|
||||
|
||||
tools = {t.name for t in result.tools}
|
||||
assert "retain" in tools
|
||||
assert "list_banks" not in tools, "Single-bank /mcp/my-bank should NOT expose list_banks"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_tool_execution_through_client(memory):
|
||||
"""Test that tools can be called (not just discovered) through the MCP client.
|
||||
|
||||
This verifies the full pipeline: HTTP → middleware → FastMCP → tool → engine → response.
|
||||
Previous tests only checked tool discovery (list_tools), not actual execution.
|
||||
"""
|
||||
from httpx import ASGITransport
|
||||
|
||||
from hindsight_api.api import create_app
|
||||
|
||||
app = create_app(memory, mcp_api_enabled=True, initialize_memory=False)
|
||||
|
||||
async with app.router.lifespan_context(app):
|
||||
async with httpx.AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as http_client:
|
||||
async with streamable_http_client("http://test/mcp/", http_client=http_client) as (
|
||||
read_stream,
|
||||
write_stream,
|
||||
_,
|
||||
):
|
||||
async with ClientSession(read_stream, write_stream) as session:
|
||||
await session.initialize()
|
||||
|
||||
# Execute list_banks tool
|
||||
result = await session.call_tool("list_banks", arguments={})
|
||||
assert result is not None
|
||||
assert len(result.content) > 0
|
||||
# The result text should be valid JSON with a "banks" key
|
||||
import json
|
||||
|
||||
response_text = result.content[0].text
|
||||
parsed = json.loads(response_text)
|
||||
assert "banks" in parsed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_mental_model_validation_through_client(memory):
|
||||
"""Test that input validation works through the real MCP transport.
|
||||
|
||||
Verifies that invalid inputs return error messages without crashing,
|
||||
and that the engine is never called with invalid data.
|
||||
"""
|
||||
from httpx import ASGITransport
|
||||
|
||||
from hindsight_api.api import create_app
|
||||
|
||||
app = create_app(memory, mcp_api_enabled=True, initialize_memory=False)
|
||||
|
||||
async with app.router.lifespan_context(app):
|
||||
async with httpx.AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as http_client:
|
||||
async with streamable_http_client("http://test/mcp/", http_client=http_client) as (
|
||||
read_stream,
|
||||
write_stream,
|
||||
_,
|
||||
):
|
||||
async with ClientSession(read_stream, write_stream) as session:
|
||||
await session.initialize()
|
||||
|
||||
# Test: empty name should return validation error
|
||||
import json
|
||||
|
||||
result = await session.call_tool(
|
||||
"create_mental_model",
|
||||
arguments={"name": "", "source_query": "test query"},
|
||||
)
|
||||
assert result is not None
|
||||
parsed = json.loads(result.content[0].text)
|
||||
assert "error" in parsed
|
||||
assert "name cannot be empty" in parsed["error"]
|
||||
|
||||
# Test: max_tokens out of range should return validation error
|
||||
result = await session.call_tool(
|
||||
"create_mental_model",
|
||||
arguments={"name": "Test", "source_query": "test query", "max_tokens": 0},
|
||||
)
|
||||
parsed = json.loads(result.content[0].text)
|
||||
assert "error" in parsed
|
||||
assert "max_tokens must be between 256 and 8192" in parsed["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_bank_named_sse_routes_to_single_bank(memory):
|
||||
"""Test that a bank named 'sse' routes to single-bank mode.
|
||||
|
||||
Regression test: the old MCP_ENDPOINTS blocklist prevented banks named 'sse'
|
||||
or 'messages' from being accessed via path routing. They fell through to
|
||||
multi-bank mode instead.
|
||||
"""
|
||||
from httpx import ASGITransport
|
||||
|
||||
from hindsight_api.api import create_app
|
||||
|
||||
app = create_app(memory, mcp_api_enabled=True, initialize_memory=False)
|
||||
|
||||
async with app.router.lifespan_context(app):
|
||||
async with httpx.AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as http_client:
|
||||
async with streamable_http_client("http://test/mcp/sse/", http_client=http_client) as (
|
||||
read_stream,
|
||||
write_stream,
|
||||
_,
|
||||
):
|
||||
async with ClientSession(read_stream, write_stream) as session:
|
||||
await session.initialize()
|
||||
result = await session.list_tools()
|
||||
tools = {t.name for t in result.tools}
|
||||
|
||||
# Should be single-bank mode (no bank management tools)
|
||||
assert "retain" in tools
|
||||
assert "recall" in tools
|
||||
assert "list_banks" not in tools, "Bank 'sse' should route to single-bank mode"
|
||||
assert "create_bank" not in tools
|
||||
|
||||
# retain should NOT have bank_id parameter (single-bank mode)
|
||||
retain_tool = next(t for t in result.tools if t.name == "retain")
|
||||
params = set(retain_tool.inputSchema.get("properties", {}).keys())
|
||||
assert "bank_id" not in params
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_bank_named_messages_routes_to_single_bank(memory):
|
||||
"""Test that a bank named 'messages' routes to single-bank mode.
|
||||
|
||||
Same regression test as test_mcp_bank_named_sse_routes_to_single_bank but for 'messages'.
|
||||
"""
|
||||
from httpx import ASGITransport
|
||||
|
||||
from hindsight_api.api import create_app
|
||||
|
||||
app = create_app(memory, mcp_api_enabled=True, initialize_memory=False)
|
||||
|
||||
async with app.router.lifespan_context(app):
|
||||
async with httpx.AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as http_client:
|
||||
async with streamable_http_client("http://test/mcp/messages/", http_client=http_client) as (
|
||||
read_stream,
|
||||
write_stream,
|
||||
_,
|
||||
):
|
||||
async with ClientSession(read_stream, write_stream) as session:
|
||||
await session.initialize()
|
||||
result = await session.list_tools()
|
||||
tools = {t.name for t in result.tools}
|
||||
|
||||
assert "retain" in tools
|
||||
assert "list_banks" not in tools, "Bank 'messages' should route to single-bank mode"
|
||||
|
||||
@@ -165,5 +165,5 @@ class TestMCPExtensionIntegration:
|
||||
assert "create_bank" in tools
|
||||
# Extension tool also present
|
||||
assert "test_extension_tool" in tools
|
||||
# Total: 5 core + 1 extension = 6 tools
|
||||
assert len(tools) == 6
|
||||
# At least 11 core + 1 extension = 12 tools (may grow as new tools are added)
|
||||
assert len(tools) >= 12
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
"""Test MCP server routing with dynamic bank_id."""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_memory():
|
||||
@@ -17,7 +18,7 @@ def mock_memory():
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_context_variable():
|
||||
"""Test that context variable works correctly."""
|
||||
from hindsight_api.api.mcp import get_current_bank_id, _current_bank_id
|
||||
from hindsight_api.api.mcp import _current_bank_id, get_current_bank_id
|
||||
|
||||
# Initially None
|
||||
assert get_current_bank_id() is None
|
||||
@@ -36,7 +37,7 @@ async def test_mcp_context_variable():
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_tools_use_context_bank_id(mock_memory):
|
||||
"""Test that MCP tools use bank_id from context."""
|
||||
from hindsight_api.api.mcp import create_mcp_server, _current_bank_id
|
||||
from hindsight_api.api.mcp import _current_bank_id, create_mcp_server
|
||||
|
||||
mcp_server = create_mcp_server(mock_memory)
|
||||
|
||||
@@ -62,6 +63,7 @@ async def test_mcp_tools_use_context_bank_id(mock_memory):
|
||||
|
||||
def test_path_parsing_logic():
|
||||
"""Test the path parsing logic for bank_id extraction."""
|
||||
|
||||
def parse_path(path):
|
||||
"""Simulate the path parsing logic from MCPMiddleware."""
|
||||
if not path.startswith("/") or len(path) <= 1:
|
||||
@@ -102,7 +104,7 @@ def test_path_parsing_logic():
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_key_context_variable():
|
||||
"""Test that API key context variable works correctly."""
|
||||
from hindsight_api.api.mcp import get_current_api_key, _current_api_key
|
||||
from hindsight_api.api.mcp import _current_api_key, get_current_api_key
|
||||
|
||||
# Initially None
|
||||
assert get_current_api_key() is None
|
||||
@@ -121,7 +123,7 @@ async def test_api_key_context_variable():
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_tools_propagate_api_key(mock_memory):
|
||||
"""Test that MCP tools propagate API key to RequestContext."""
|
||||
from hindsight_api.api.mcp import create_mcp_server, _current_bank_id, _current_api_key
|
||||
from hindsight_api.api.mcp import _current_api_key, _current_bank_id, create_mcp_server
|
||||
|
||||
mcp_server = create_mcp_server(mock_memory)
|
||||
tools = mcp_server._tool_manager._tools
|
||||
@@ -147,8 +149,10 @@ async def test_mcp_tools_propagate_api_key(mock_memory):
|
||||
async def test_tenant_id_context_variable():
|
||||
"""Test that tenant_id and api_key_id context variables work correctly."""
|
||||
from hindsight_api.api.mcp import (
|
||||
get_current_tenant_id, _current_tenant_id,
|
||||
get_current_api_key_id, _current_api_key_id,
|
||||
_current_api_key_id,
|
||||
_current_tenant_id,
|
||||
get_current_api_key_id,
|
||||
get_current_tenant_id,
|
||||
)
|
||||
|
||||
# Initially None
|
||||
@@ -179,9 +183,11 @@ async def test_mcp_tools_propagate_tenant_id_and_api_key_id(mock_memory):
|
||||
MCP operations get tenant_id="unknown" and billing is skipped entirely.
|
||||
"""
|
||||
from hindsight_api.api.mcp import (
|
||||
_current_api_key,
|
||||
_current_api_key_id,
|
||||
_current_bank_id,
|
||||
_current_tenant_id,
|
||||
create_mcp_server,
|
||||
_current_bank_id, _current_api_key,
|
||||
_current_tenant_id, _current_api_key_id,
|
||||
)
|
||||
|
||||
mcp_server = create_mcp_server(mock_memory)
|
||||
@@ -210,20 +216,28 @@ async def test_mcp_tools_propagate_tenant_id_and_api_key_id(mock_memory):
|
||||
|
||||
|
||||
def test_multi_bank_mode_exposes_all_tools(mock_memory):
|
||||
"""Test that multi-bank mode exposes all tools including bank management."""
|
||||
"""Test that multi-bank mode exposes all tools including bank management and mental models."""
|
||||
from hindsight_api.api.mcp import create_mcp_server
|
||||
|
||||
# Create server in multi-bank mode (default)
|
||||
mcp_server = create_mcp_server(mock_memory, multi_bank=True)
|
||||
tools = mcp_server._tool_manager._tools
|
||||
|
||||
# Should have all tools
|
||||
# Core tools
|
||||
assert "retain" in tools
|
||||
assert "recall" in tools
|
||||
assert "reflect" in tools
|
||||
assert "list_banks" in tools
|
||||
assert "create_bank" in tools
|
||||
|
||||
# Mental model tools
|
||||
assert "list_mental_models" in tools
|
||||
assert "get_mental_model" in tools
|
||||
assert "create_mental_model" in tools
|
||||
assert "update_mental_model" in tools
|
||||
assert "delete_mental_model" in tools
|
||||
assert "refresh_mental_model" in tools
|
||||
|
||||
|
||||
def test_single_bank_mode_excludes_bank_management_tools(mock_memory):
|
||||
"""Test that single-bank mode only exposes bank-scoped tools."""
|
||||
@@ -233,11 +247,19 @@ def test_single_bank_mode_excludes_bank_management_tools(mock_memory):
|
||||
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
|
||||
tools = mcp_server._tool_manager._tools
|
||||
|
||||
# Should only have bank-scoped tools
|
||||
# Should have bank-scoped tools
|
||||
assert "retain" in tools
|
||||
assert "recall" in tools
|
||||
assert "reflect" in tools
|
||||
|
||||
# Mental model tools should also be present (they're bank-scoped)
|
||||
assert "list_mental_models" in tools
|
||||
assert "get_mental_model" in tools
|
||||
assert "create_mental_model" in tools
|
||||
assert "update_mental_model" in tools
|
||||
assert "delete_mental_model" in tools
|
||||
assert "refresh_mental_model" in tools
|
||||
|
||||
# Should NOT have bank management tools
|
||||
assert "list_banks" not in tools
|
||||
assert "create_bank" not in tools
|
||||
@@ -245,46 +267,56 @@ def test_single_bank_mode_excludes_bank_management_tools(mock_memory):
|
||||
|
||||
def test_multi_bank_mode_tools_have_bank_id_param(mock_memory):
|
||||
"""Test that multi-bank mode tools include bank_id parameter."""
|
||||
from hindsight_api.api.mcp import create_mcp_server
|
||||
import inspect
|
||||
|
||||
from hindsight_api.api.mcp import create_mcp_server
|
||||
|
||||
mcp_server = create_mcp_server(mock_memory, multi_bank=True)
|
||||
tools = mcp_server._tool_manager._tools
|
||||
|
||||
# Check that tools have bank_id parameter
|
||||
retain_tool = tools["retain"]
|
||||
retain_sig = inspect.signature(retain_tool.fn)
|
||||
assert "bank_id" in retain_sig.parameters
|
||||
|
||||
recall_tool = tools["recall"]
|
||||
recall_sig = inspect.signature(recall_tool.fn)
|
||||
assert "bank_id" in recall_sig.parameters
|
||||
|
||||
reflect_tool = tools["reflect"]
|
||||
reflect_sig = inspect.signature(reflect_tool.fn)
|
||||
assert "bank_id" in reflect_sig.parameters
|
||||
# All bank-scoped tools should have bank_id parameter in multi-bank mode
|
||||
bank_scoped_tools = [
|
||||
"retain",
|
||||
"recall",
|
||||
"reflect",
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"create_mental_model",
|
||||
"update_mental_model",
|
||||
"delete_mental_model",
|
||||
"refresh_mental_model",
|
||||
]
|
||||
for tool_name in bank_scoped_tools:
|
||||
tool = tools[tool_name]
|
||||
sig = inspect.signature(tool.fn)
|
||||
assert "bank_id" in sig.parameters, f"{tool_name} should have bank_id param in multi-bank mode"
|
||||
|
||||
|
||||
def test_single_bank_mode_tools_no_bank_id_param(mock_memory):
|
||||
"""Test that single-bank mode tools do NOT include bank_id parameter."""
|
||||
from hindsight_api.api.mcp import create_mcp_server
|
||||
import inspect
|
||||
|
||||
from hindsight_api.api.mcp import create_mcp_server
|
||||
|
||||
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
|
||||
tools = mcp_server._tool_manager._tools
|
||||
|
||||
# Check that tools do NOT have bank_id parameter
|
||||
retain_tool = tools["retain"]
|
||||
retain_sig = inspect.signature(retain_tool.fn)
|
||||
assert "bank_id" not in retain_sig.parameters
|
||||
|
||||
recall_tool = tools["recall"]
|
||||
recall_sig = inspect.signature(recall_tool.fn)
|
||||
assert "bank_id" not in recall_sig.parameters
|
||||
|
||||
reflect_tool = tools["reflect"]
|
||||
reflect_sig = inspect.signature(reflect_tool.fn)
|
||||
assert "bank_id" not in reflect_sig.parameters
|
||||
# No bank-scoped tool should have bank_id parameter in single-bank mode
|
||||
bank_scoped_tools = [
|
||||
"retain",
|
||||
"recall",
|
||||
"reflect",
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"create_mental_model",
|
||||
"update_mental_model",
|
||||
"delete_mental_model",
|
||||
"refresh_mental_model",
|
||||
]
|
||||
for tool_name in bank_scoped_tools:
|
||||
tool = tools[tool_name]
|
||||
sig = inspect.signature(tool.fn)
|
||||
assert "bank_id" not in sig.parameters, f"{tool_name} should NOT have bank_id param in single-bank mode"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -308,19 +340,26 @@ async def test_middleware_handles_both_endpoints(mock_memory):
|
||||
assert "recall" in multi_bank_tools
|
||||
assert "list_banks" in multi_bank_tools
|
||||
assert "create_bank" in multi_bank_tools
|
||||
assert "list_mental_models" in multi_bank_tools
|
||||
assert "create_mental_model" in multi_bank_tools
|
||||
|
||||
# Single-bank should only have scoped tools
|
||||
assert "retain" in single_bank_tools
|
||||
assert "recall" in single_bank_tools
|
||||
assert "list_mental_models" in single_bank_tools
|
||||
assert "create_mental_model" in single_bank_tools
|
||||
assert "list_banks" not in single_bank_tools
|
||||
assert "create_bank" not in single_bank_tools
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_logic_from_url_path():
|
||||
"""Test that routing correctly selects server based on URL structure."""
|
||||
"""Test that routing correctly selects server based on URL structure.
|
||||
|
||||
Simulates the path parsing logic from MCPMiddleware.__call__ after the
|
||||
prefix has been stripped. Any first path segment is treated as a bank_id.
|
||||
"""
|
||||
from hindsight_api.api.mcp import MCPMiddleware
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
# Mock memory
|
||||
mock_memory = MagicMock()
|
||||
@@ -329,28 +368,23 @@ async def test_routing_logic_from_url_path():
|
||||
middleware = MCPMiddleware(None, mock_memory)
|
||||
|
||||
# Simulate different URL patterns and verify routing
|
||||
# Path is what remains after stripping the /mcp prefix
|
||||
test_cases = [
|
||||
# (path_after_stripping_mcp, expected_bank_id_from_path, expected_bank_id, description)
|
||||
# (path_after_prefix_strip, expected_bank_id_from_path, expected_bank_id, description)
|
||||
("/alice/messages", True, "alice", "Bank ID in path with endpoint"),
|
||||
("/my-agent-123/", True, "my-agent-123", "Bank ID in path with trailing slash"),
|
||||
("ciccio/messages", True, "ciccio", "Bank ID without leading slash (after mount strip)"),
|
||||
("bob", True, "bob", "Bank ID only, no leading slash"),
|
||||
("/messages", False, None, "MCP endpoint, no bank ID"),
|
||||
("/sse/", True, "sse", "Bank named 'sse' routes to single-bank"),
|
||||
("/messages/", True, "messages", "Bank named 'messages' routes to single-bank"),
|
||||
("/", False, None, "Root path, no bank ID"),
|
||||
]
|
||||
|
||||
for path, expected_bank_from_path, expected_bank_id, description in test_cases:
|
||||
# Simulate the path parsing logic with leading slash normalization
|
||||
if path and not path.startswith("/"):
|
||||
path = "/" + path
|
||||
|
||||
bank_id = None
|
||||
bank_id_from_path = False
|
||||
MCP_ENDPOINTS = {"sse", "messages"}
|
||||
|
||||
if path.startswith("/") and len(path) > 1:
|
||||
parts = path[1:].split("/", 1)
|
||||
if parts[0] and parts[0] not in MCP_ENDPOINTS:
|
||||
if parts[0]:
|
||||
bank_id = parts[0]
|
||||
bank_id_from_path = True
|
||||
|
||||
|
||||
@@ -1,10 +1,17 @@
|
||||
"""Tests for the shared MCP tools module."""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from hindsight_api.mcp_tools import build_content_dict, parse_timestamp
|
||||
from hindsight_api.mcp_tools import (
|
||||
MCPToolsConfig,
|
||||
_validate_mental_model_inputs,
|
||||
build_content_dict,
|
||||
parse_timestamp,
|
||||
register_mcp_tools,
|
||||
)
|
||||
|
||||
|
||||
class TestParseTimestamp:
|
||||
@@ -61,3 +68,579 @@ class TestBuildContentDict:
|
||||
result, error = build_content_dict("test content", "test_context", None)
|
||||
assert error is None
|
||||
assert "event_date" not in result
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Mental Model MCP Tool Tests
|
||||
# =========================================================================
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_memory():
|
||||
"""Create a mock MemoryEngine with mental model methods."""
|
||||
memory = MagicMock()
|
||||
memory.list_mental_models = AsyncMock(
|
||||
return_value=[
|
||||
{"id": "mm-1", "name": "Coding Prefs", "source_query": "coding preferences?", "content": "Prefers Python"},
|
||||
{"id": "mm-2", "name": "Goals", "source_query": "current goals?", "content": "Ship v2"},
|
||||
]
|
||||
)
|
||||
memory.get_mental_model = AsyncMock(
|
||||
return_value={
|
||||
"id": "mm-1",
|
||||
"name": "Coding Prefs",
|
||||
"source_query": "coding preferences?",
|
||||
"content": "Prefers Python",
|
||||
}
|
||||
)
|
||||
memory.create_mental_model = AsyncMock(return_value={"id": "mm-new"})
|
||||
memory.submit_async_refresh_mental_model = AsyncMock(return_value={"operation_id": "op-123"})
|
||||
memory.update_mental_model = AsyncMock(
|
||||
return_value={
|
||||
"id": "mm-1",
|
||||
"name": "Updated Name",
|
||||
"source_query": "new query?",
|
||||
"content": "Updated",
|
||||
}
|
||||
)
|
||||
memory.delete_mental_model = AsyncMock(return_value=True)
|
||||
return memory
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mcp_server_with_mental_models(mock_memory):
|
||||
"""Create a FastMCP server with mental model tools registered (multi-bank mode)."""
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP("test", stateless_http=True)
|
||||
config = MCPToolsConfig(
|
||||
bank_id_resolver=lambda: "test-bank",
|
||||
include_bank_id_param=True,
|
||||
tools={
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"create_mental_model",
|
||||
"update_mental_model",
|
||||
"delete_mental_model",
|
||||
"refresh_mental_model",
|
||||
},
|
||||
)
|
||||
register_mcp_tools(mcp, mock_memory, config)
|
||||
return mcp
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mcp_server_single_bank(mock_memory):
|
||||
"""Create a FastMCP server with mental model tools registered (single-bank mode)."""
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP("test")
|
||||
config = MCPToolsConfig(
|
||||
bank_id_resolver=lambda: "fixed-bank",
|
||||
include_bank_id_param=False,
|
||||
tools={
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"create_mental_model",
|
||||
"update_mental_model",
|
||||
"delete_mental_model",
|
||||
"refresh_mental_model",
|
||||
},
|
||||
)
|
||||
register_mcp_tools(mcp, mock_memory, config)
|
||||
return mcp
|
||||
|
||||
|
||||
class TestMentalModelToolRegistration:
|
||||
"""Test that mental model tools are registered correctly."""
|
||||
|
||||
def test_tools_registered_multi_bank(self, mcp_server_with_mental_models):
|
||||
tools = mcp_server_with_mental_models._tool_manager._tools
|
||||
expected = {
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"create_mental_model",
|
||||
"update_mental_model",
|
||||
"delete_mental_model",
|
||||
"refresh_mental_model",
|
||||
}
|
||||
assert expected == set(tools.keys())
|
||||
|
||||
def test_tools_registered_single_bank(self, mcp_server_single_bank):
|
||||
tools = mcp_server_single_bank._tool_manager._tools
|
||||
expected = {
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"create_mental_model",
|
||||
"update_mental_model",
|
||||
"delete_mental_model",
|
||||
"refresh_mental_model",
|
||||
}
|
||||
assert expected == set(tools.keys())
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_mental_models_propagates_request_context(self, mock_memory):
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP("test", stateless_http=True)
|
||||
config = MCPToolsConfig(
|
||||
bank_id_resolver=lambda: "test-bank",
|
||||
api_key_resolver=lambda: "test-api-key",
|
||||
include_bank_id_param=True,
|
||||
tools={"list_mental_models"},
|
||||
)
|
||||
register_mcp_tools(mcp, mock_memory, config)
|
||||
await _tools(mcp)["list_mental_models"].fn()
|
||||
request_context = mock_memory.list_mental_models.call_args.kwargs["request_context"]
|
||||
assert request_context.api_key == "test-api-key"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_mental_model_propagates_request_context(self, mock_memory):
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP("test", stateless_http=True)
|
||||
config = MCPToolsConfig(
|
||||
bank_id_resolver=lambda: "test-bank",
|
||||
api_key_resolver=lambda: "test-api-key",
|
||||
include_bank_id_param=True,
|
||||
tools={"create_mental_model"},
|
||||
)
|
||||
register_mcp_tools(mcp, mock_memory, config)
|
||||
await _tools(mcp)["create_mental_model"].fn(name="Test", source_query="query")
|
||||
request_context = mock_memory.create_mental_model.call_args.kwargs["request_context"]
|
||||
assert request_context.api_key == "test-api-key"
|
||||
|
||||
def test_mental_model_tools_in_default_set(self):
|
||||
"""Mental model tools should be in the default tools set when config.tools is None."""
|
||||
from fastmcp import FastMCP
|
||||
|
||||
memory = MagicMock()
|
||||
# Mock all engine methods that tools reference
|
||||
memory.retain_batch_async = AsyncMock()
|
||||
memory.submit_async_retain = AsyncMock(return_value={"operation_id": "op"})
|
||||
memory.recall_async = AsyncMock(return_value=MagicMock(results=[]))
|
||||
memory.reflect_async = AsyncMock()
|
||||
memory.list_banks = AsyncMock(return_value=[])
|
||||
memory.get_bank_profile = AsyncMock(return_value={})
|
||||
memory.update_bank = AsyncMock()
|
||||
memory.list_mental_models = AsyncMock(return_value=[])
|
||||
memory.get_mental_model = AsyncMock()
|
||||
memory.create_mental_model = AsyncMock()
|
||||
memory.submit_async_refresh_mental_model = AsyncMock()
|
||||
memory.update_mental_model = AsyncMock()
|
||||
memory.delete_mental_model = AsyncMock()
|
||||
|
||||
mcp = FastMCP("test", stateless_http=True)
|
||||
config = MCPToolsConfig(
|
||||
bank_id_resolver=lambda: "bank",
|
||||
include_bank_id_param=True,
|
||||
tools=None, # Default - all tools
|
||||
)
|
||||
register_mcp_tools(mcp, memory, config)
|
||||
tools = mcp._tool_manager._tools
|
||||
assert "list_mental_models" in tools
|
||||
assert "create_mental_model" in tools
|
||||
assert "refresh_mental_model" in tools
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def no_bank_mcp_server(mock_memory):
|
||||
"""Create a multi-bank MCP server where bank_id_resolver returns None."""
|
||||
from fastmcp import FastMCP
|
||||
|
||||
mcp = FastMCP("test", stateless_http=True)
|
||||
config = MCPToolsConfig(
|
||||
bank_id_resolver=lambda: None,
|
||||
include_bank_id_param=True,
|
||||
tools={
|
||||
"list_mental_models",
|
||||
"get_mental_model",
|
||||
"create_mental_model",
|
||||
"update_mental_model",
|
||||
"delete_mental_model",
|
||||
"refresh_mental_model",
|
||||
},
|
||||
)
|
||||
register_mcp_tools(mcp, mock_memory, config)
|
||||
return mcp
|
||||
|
||||
|
||||
def _tools(mcp_server):
|
||||
"""Helper to get tools dict from MCP server."""
|
||||
return mcp_server._tool_manager._tools
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestListMentalModels:
|
||||
async def test_list_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["list_mental_models"].fn()
|
||||
assert '"mm-1"' in result
|
||||
assert '"mm-2"' in result
|
||||
mock_memory.list_mental_models.assert_called_once()
|
||||
assert mock_memory.list_mental_models.call_args.kwargs["bank_id"] == "test-bank"
|
||||
|
||||
async def test_list_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory):
|
||||
"""Explicit bank_id should override the resolver."""
|
||||
await _tools(mcp_server_with_mental_models)["list_mental_models"].fn(bank_id="other-bank")
|
||||
assert mock_memory.list_mental_models.call_args.kwargs["bank_id"] == "other-bank"
|
||||
|
||||
async def test_list_with_tags(self, mcp_server_with_mental_models, mock_memory):
|
||||
await _tools(mcp_server_with_mental_models)["list_mental_models"].fn(tags=["work"])
|
||||
assert mock_memory.list_mental_models.call_args.kwargs["tags"] == ["work"]
|
||||
|
||||
async def test_list_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
result = await _tools(mcp_server_single_bank)["list_mental_models"].fn()
|
||||
assert isinstance(result, dict)
|
||||
assert len(result["items"]) == 2
|
||||
assert mock_memory.list_mental_models.call_args.kwargs["bank_id"] == "fixed-bank"
|
||||
|
||||
async def test_list_no_bank_returns_error(self, no_bank_mcp_server):
|
||||
result = await _tools(no_bank_mcp_server)["list_mental_models"].fn()
|
||||
assert "error" in result
|
||||
|
||||
async def test_list_engine_error_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.list_mental_models.side_effect = RuntimeError("DB connection lost")
|
||||
result = await _tools(mcp_server_with_mental_models)["list_mental_models"].fn()
|
||||
assert "error" in result
|
||||
assert "DB connection lost" in result
|
||||
|
||||
async def test_list_engine_error_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
mock_memory.list_mental_models.side_effect = RuntimeError("DB connection lost")
|
||||
result = await _tools(mcp_server_single_bank)["list_mental_models"].fn()
|
||||
assert isinstance(result, dict)
|
||||
assert "error" in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestGetMentalModel:
|
||||
async def test_get_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["get_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert '"mm-1"' in result
|
||||
assert mock_memory.get_mental_model.call_args.kwargs["mental_model_id"] == "mm-1"
|
||||
|
||||
async def test_get_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory):
|
||||
await _tools(mcp_server_with_mental_models)["get_mental_model"].fn(mental_model_id="mm-1", bank_id="other-bank")
|
||||
assert mock_memory.get_mental_model.call_args.kwargs["bank_id"] == "other-bank"
|
||||
|
||||
async def test_get_not_found_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.get_mental_model.return_value = None
|
||||
result = await _tools(mcp_server_with_mental_models)["get_mental_model"].fn(mental_model_id="missing")
|
||||
assert "not found" in result
|
||||
|
||||
async def test_get_not_found_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
mock_memory.get_mental_model.return_value = None
|
||||
result = await _tools(mcp_server_single_bank)["get_mental_model"].fn(mental_model_id="missing")
|
||||
assert isinstance(result, dict)
|
||||
assert "not found" in result["error"]
|
||||
|
||||
async def test_get_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
result = await _tools(mcp_server_single_bank)["get_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert isinstance(result, dict)
|
||||
assert result["id"] == "mm-1"
|
||||
|
||||
async def test_get_no_bank_returns_error(self, no_bank_mcp_server):
|
||||
result = await _tools(no_bank_mcp_server)["get_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert "error" in result
|
||||
|
||||
async def test_get_engine_error(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.get_mental_model.side_effect = RuntimeError("DB error")
|
||||
result = await _tools(mcp_server_with_mental_models)["get_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert "error" in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestCreateMentalModel:
|
||||
async def test_create_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["create_mental_model"].fn(
|
||||
name="Test Model",
|
||||
source_query="What are the user's preferences?",
|
||||
)
|
||||
assert '"mm-new"' in result
|
||||
assert '"op-123"' in result
|
||||
mock_memory.create_mental_model.assert_called_once()
|
||||
call_kwargs = mock_memory.create_mental_model.call_args.kwargs
|
||||
assert call_kwargs["name"] == "Test Model"
|
||||
assert call_kwargs["source_query"] == "What are the user's preferences?"
|
||||
assert call_kwargs["content"] == "Generating content..."
|
||||
# Verify async refresh was scheduled
|
||||
mock_memory.submit_async_refresh_mental_model.assert_called_once()
|
||||
assert mock_memory.submit_async_refresh_mental_model.call_args.kwargs["mental_model_id"] == "mm-new"
|
||||
|
||||
async def test_create_with_custom_id(self, mcp_server_with_mental_models, mock_memory):
|
||||
await _tools(mcp_server_with_mental_models)["create_mental_model"].fn(
|
||||
name="Test", source_query="query", mental_model_id="custom-id"
|
||||
)
|
||||
assert mock_memory.create_mental_model.call_args.kwargs["mental_model_id"] == "custom-id"
|
||||
|
||||
async def test_create_with_tags_and_max_tokens(self, mcp_server_with_mental_models, mock_memory):
|
||||
await _tools(mcp_server_with_mental_models)["create_mental_model"].fn(
|
||||
name="Test", source_query="query", tags=["work", "coding"], max_tokens=4096
|
||||
)
|
||||
call_kwargs = mock_memory.create_mental_model.call_args.kwargs
|
||||
assert call_kwargs["tags"] == ["work", "coding"]
|
||||
assert call_kwargs["max_tokens"] == 4096
|
||||
|
||||
async def test_create_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory):
|
||||
await _tools(mcp_server_with_mental_models)["create_mental_model"].fn(
|
||||
name="Test", source_query="query", bank_id="other-bank"
|
||||
)
|
||||
assert mock_memory.create_mental_model.call_args.kwargs["bank_id"] == "other-bank"
|
||||
assert mock_memory.submit_async_refresh_mental_model.call_args.kwargs["bank_id"] == "other-bank"
|
||||
|
||||
async def test_create_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
result = await _tools(mcp_server_single_bank)["create_mental_model"].fn(name="Test", source_query="query")
|
||||
assert isinstance(result, dict)
|
||||
assert result["mental_model_id"] == "mm-new"
|
||||
assert result["operation_id"] == "op-123"
|
||||
|
||||
async def test_create_no_bank_returns_error(self, no_bank_mcp_server):
|
||||
result = await _tools(no_bank_mcp_server)["create_mental_model"].fn(name="Test", source_query="query")
|
||||
assert "error" in result
|
||||
|
||||
async def test_create_value_error_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
"""ValueError from engine (e.g. invalid ID format) should return error, not crash."""
|
||||
mock_memory.create_mental_model.side_effect = ValueError("ID must be alphanumeric lowercase")
|
||||
result = await _tools(mcp_server_with_mental_models)["create_mental_model"].fn(
|
||||
name="Test", source_query="query", mental_model_id="INVALID!!"
|
||||
)
|
||||
assert "alphanumeric" in result
|
||||
|
||||
async def test_create_value_error_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
mock_memory.create_mental_model.side_effect = ValueError("ID must be alphanumeric lowercase")
|
||||
result = await _tools(mcp_server_single_bank)["create_mental_model"].fn(
|
||||
name="Test", source_query="query", mental_model_id="INVALID!!"
|
||||
)
|
||||
assert isinstance(result, dict)
|
||||
assert "alphanumeric" in result["error"]
|
||||
|
||||
async def test_create_engine_error(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.create_mental_model.side_effect = RuntimeError("DB error")
|
||||
result = await _tools(mcp_server_with_mental_models)["create_mental_model"].fn(
|
||||
name="Test", source_query="query"
|
||||
)
|
||||
assert "error" in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestUpdateMentalModel:
|
||||
async def test_update_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["update_mental_model"].fn(
|
||||
mental_model_id="mm-1", name="Updated Name"
|
||||
)
|
||||
assert '"Updated Name"' in result
|
||||
call_kwargs = mock_memory.update_mental_model.call_args.kwargs
|
||||
assert call_kwargs["name"] == "Updated Name"
|
||||
assert call_kwargs["source_query"] is None # Not updated
|
||||
|
||||
async def test_update_multiple_fields(self, mcp_server_with_mental_models, mock_memory):
|
||||
await _tools(mcp_server_with_mental_models)["update_mental_model"].fn(
|
||||
mental_model_id="mm-1", name="New Name", source_query="new query?", tags=["updated"], max_tokens=4096
|
||||
)
|
||||
call_kwargs = mock_memory.update_mental_model.call_args.kwargs
|
||||
assert call_kwargs["name"] == "New Name"
|
||||
assert call_kwargs["source_query"] == "new query?"
|
||||
assert call_kwargs["tags"] == ["updated"]
|
||||
assert call_kwargs["max_tokens"] == 4096
|
||||
|
||||
async def test_update_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory):
|
||||
await _tools(mcp_server_with_mental_models)["update_mental_model"].fn(
|
||||
mental_model_id="mm-1", name="X", bank_id="other-bank"
|
||||
)
|
||||
assert mock_memory.update_mental_model.call_args.kwargs["bank_id"] == "other-bank"
|
||||
|
||||
async def test_update_not_found_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.update_mental_model.return_value = None
|
||||
result = await _tools(mcp_server_with_mental_models)["update_mental_model"].fn(
|
||||
mental_model_id="missing", name="X"
|
||||
)
|
||||
assert "not found" in result
|
||||
|
||||
async def test_update_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
result = await _tools(mcp_server_single_bank)["update_mental_model"].fn(mental_model_id="mm-1", name="Updated")
|
||||
assert isinstance(result, dict)
|
||||
assert mock_memory.update_mental_model.call_args.kwargs["bank_id"] == "fixed-bank"
|
||||
|
||||
async def test_update_not_found_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
mock_memory.update_mental_model.return_value = None
|
||||
result = await _tools(mcp_server_single_bank)["update_mental_model"].fn(mental_model_id="missing", name="X")
|
||||
assert isinstance(result, dict)
|
||||
assert "not found" in result["error"]
|
||||
|
||||
async def test_update_no_bank_returns_error(self, no_bank_mcp_server):
|
||||
result = await _tools(no_bank_mcp_server)["update_mental_model"].fn(mental_model_id="mm-1", name="X")
|
||||
assert "error" in result
|
||||
|
||||
async def test_update_engine_error(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.update_mental_model.side_effect = RuntimeError("DB error")
|
||||
result = await _tools(mcp_server_with_mental_models)["update_mental_model"].fn(mental_model_id="mm-1", name="X")
|
||||
assert "error" in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestDeleteMentalModel:
|
||||
async def test_delete_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["delete_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert '"deleted"' in result
|
||||
assert mock_memory.delete_mental_model.call_args.kwargs["mental_model_id"] == "mm-1"
|
||||
|
||||
async def test_delete_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory):
|
||||
await _tools(mcp_server_with_mental_models)["delete_mental_model"].fn(
|
||||
mental_model_id="mm-1", bank_id="other-bank"
|
||||
)
|
||||
assert mock_memory.delete_mental_model.call_args.kwargs["bank_id"] == "other-bank"
|
||||
|
||||
async def test_delete_not_found_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.delete_mental_model.return_value = False
|
||||
result = await _tools(mcp_server_with_mental_models)["delete_mental_model"].fn(mental_model_id="missing")
|
||||
assert "not found" in result
|
||||
|
||||
async def test_delete_not_found_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
mock_memory.delete_mental_model.return_value = False
|
||||
result = await _tools(mcp_server_single_bank)["delete_mental_model"].fn(mental_model_id="missing")
|
||||
assert isinstance(result, dict)
|
||||
assert "not found" in result["error"]
|
||||
|
||||
async def test_delete_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
result = await _tools(mcp_server_single_bank)["delete_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert isinstance(result, dict)
|
||||
assert result["status"] == "deleted"
|
||||
|
||||
async def test_delete_no_bank_returns_error(self, no_bank_mcp_server):
|
||||
result = await _tools(no_bank_mcp_server)["delete_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert "error" in result
|
||||
|
||||
async def test_delete_engine_error(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.delete_mental_model.side_effect = RuntimeError("DB error")
|
||||
result = await _tools(mcp_server_with_mental_models)["delete_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert "error" in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestRefreshMentalModel:
|
||||
async def test_refresh_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["refresh_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert '"op-123"' in result
|
||||
assert '"queued"' in result
|
||||
|
||||
async def test_refresh_with_bank_id_override(self, mcp_server_with_mental_models, mock_memory):
|
||||
await _tools(mcp_server_with_mental_models)["refresh_mental_model"].fn(
|
||||
mental_model_id="mm-1", bank_id="other-bank"
|
||||
)
|
||||
assert mock_memory.submit_async_refresh_mental_model.call_args.kwargs["bank_id"] == "other-bank"
|
||||
|
||||
async def test_refresh_not_found_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.submit_async_refresh_mental_model.side_effect = ValueError("Mental model 'missing' not found")
|
||||
result = await _tools(mcp_server_with_mental_models)["refresh_mental_model"].fn(mental_model_id="missing")
|
||||
assert "not found" in result
|
||||
|
||||
async def test_refresh_not_found_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
mock_memory.submit_async_refresh_mental_model.side_effect = ValueError("not found")
|
||||
result = await _tools(mcp_server_single_bank)["refresh_mental_model"].fn(mental_model_id="missing")
|
||||
assert isinstance(result, dict)
|
||||
assert "not found" in result["error"]
|
||||
|
||||
async def test_refresh_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
result = await _tools(mcp_server_single_bank)["refresh_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert isinstance(result, dict)
|
||||
assert result["operation_id"] == "op-123"
|
||||
|
||||
async def test_refresh_no_bank_returns_error(self, no_bank_mcp_server):
|
||||
result = await _tools(no_bank_mcp_server)["refresh_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert "error" in result
|
||||
|
||||
async def test_refresh_engine_error(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.submit_async_refresh_mental_model.side_effect = RuntimeError("DB error")
|
||||
result = await _tools(mcp_server_with_mental_models)["refresh_mental_model"].fn(mental_model_id="mm-1")
|
||||
assert "error" in result
|
||||
|
||||
|
||||
class TestValidateMentalModelInputs:
|
||||
"""Tests for the _validate_mental_model_inputs helper."""
|
||||
|
||||
def test_valid_inputs(self):
|
||||
assert _validate_mental_model_inputs(name="Test", source_query="query", max_tokens=2048) is None
|
||||
|
||||
def test_none_inputs(self):
|
||||
assert _validate_mental_model_inputs() is None
|
||||
|
||||
def test_empty_name(self):
|
||||
result = _validate_mental_model_inputs(name="")
|
||||
assert result == "name cannot be empty"
|
||||
|
||||
def test_whitespace_name(self):
|
||||
result = _validate_mental_model_inputs(name=" ")
|
||||
assert result == "name cannot be empty"
|
||||
|
||||
def test_empty_source_query(self):
|
||||
result = _validate_mental_model_inputs(source_query="")
|
||||
assert result == "source_query cannot be empty"
|
||||
|
||||
def test_whitespace_source_query(self):
|
||||
result = _validate_mental_model_inputs(source_query=" \t ")
|
||||
assert result == "source_query cannot be empty"
|
||||
|
||||
def test_max_tokens_too_low(self):
|
||||
result = _validate_mental_model_inputs(max_tokens=0)
|
||||
assert "max_tokens must be between 256 and 8192" in result
|
||||
|
||||
def test_max_tokens_too_high(self):
|
||||
result = _validate_mental_model_inputs(max_tokens=10000)
|
||||
assert "max_tokens must be between 256 and 8192" in result
|
||||
|
||||
def test_max_tokens_at_lower_bound(self):
|
||||
assert _validate_mental_model_inputs(max_tokens=256) is None
|
||||
|
||||
def test_max_tokens_at_upper_bound(self):
|
||||
assert _validate_mental_model_inputs(max_tokens=8192) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestMentalModelInputValidation:
|
||||
"""Tests that validation is applied in create/update tools before engine calls."""
|
||||
|
||||
async def test_create_empty_name_returns_error_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["create_mental_model"].fn(name="", source_query="query")
|
||||
assert "name cannot be empty" in result
|
||||
mock_memory.create_mental_model.assert_not_called()
|
||||
|
||||
async def test_create_empty_source_query_returns_error_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["create_mental_model"].fn(name="Test", source_query="")
|
||||
assert "source_query cannot be empty" in result
|
||||
mock_memory.create_mental_model.assert_not_called()
|
||||
|
||||
async def test_create_max_tokens_too_low_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["create_mental_model"].fn(
|
||||
name="Test", source_query="query", max_tokens=0
|
||||
)
|
||||
assert "max_tokens must be between 256 and 8192" in result
|
||||
mock_memory.create_mental_model.assert_not_called()
|
||||
|
||||
async def test_create_max_tokens_too_high_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
result = await _tools(mcp_server_single_bank)["create_mental_model"].fn(
|
||||
name="Test", source_query="query", max_tokens=10000
|
||||
)
|
||||
assert isinstance(result, dict)
|
||||
assert "max_tokens must be between 256 and 8192" in result["error"]
|
||||
mock_memory.create_mental_model.assert_not_called()
|
||||
|
||||
async def test_update_empty_name_returns_error_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
result = await _tools(mcp_server_with_mental_models)["update_mental_model"].fn(mental_model_id="mm-1", name="")
|
||||
assert "name cannot be empty" in result
|
||||
mock_memory.update_mental_model.assert_not_called()
|
||||
|
||||
async def test_update_empty_name_returns_error_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
result = await _tools(mcp_server_single_bank)["update_mental_model"].fn(mental_model_id="mm-1", name=" ")
|
||||
assert isinstance(result, dict)
|
||||
assert "name cannot be empty" in result["error"]
|
||||
mock_memory.update_mental_model.assert_not_called()
|
||||
|
||||
async def test_not_found_error_includes_bank_id_multi_bank(self, mcp_server_with_mental_models, mock_memory):
|
||||
mock_memory.get_mental_model.return_value = None
|
||||
result = await _tools(mcp_server_with_mental_models)["get_mental_model"].fn(mental_model_id="missing")
|
||||
assert "test-bank" in result
|
||||
|
||||
async def test_not_found_error_includes_bank_id_single_bank(self, mcp_server_single_bank, mock_memory):
|
||||
mock_memory.get_mental_model.return_value = None
|
||||
result = await _tools(mcp_server_single_bank)["get_mental_model"].fn(mental_model_id="missing")
|
||||
assert isinstance(result, dict)
|
||||
assert "fixed-bank" in result["error"]
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
"""
|
||||
Test reflect endpoint with empty based_on (no memories scenario).
|
||||
|
||||
This test verifies that the API returns the correct based_on format:
|
||||
- v0.3.0 (old): returned based_on as list []
|
||||
- v0.4.0+ (current): returns based_on as object {"memories": [], "mental_models": [], "directives": []}
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
import httpx
|
||||
from hindsight_api.api import create_app
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def api_client(memory):
|
||||
"""Create an async test client for the FastAPI app."""
|
||||
app = create_app(memory, initialize_memory=False)
|
||||
transport = httpx.ASGITransport(app=app)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
yield client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reflect_with_no_memories_empty_bank(api_client):
|
||||
"""Test reflect on an empty bank (no memories) with include.facts enabled."""
|
||||
bank_id = "test_empty_bank"
|
||||
|
||||
# Reflect on empty bank with facts requested
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/reflect",
|
||||
json={
|
||||
"query": "What do you know about machine learning?",
|
||||
"budget": "low",
|
||||
"include": {
|
||||
"facts": {} # Request facts but bank is empty
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
|
||||
# DEBUG: Print what the API actually returned
|
||||
import json
|
||||
print("\n" + "="*80)
|
||||
print("API Response:")
|
||||
print(json.dumps(data, indent=2))
|
||||
print("="*80 + "\n")
|
||||
|
||||
# Verify response structure
|
||||
assert "text" in data
|
||||
assert "based_on" in data
|
||||
|
||||
# The API should return based_on as either:
|
||||
# 1. null/None (if include.facts not set)
|
||||
# 2. {"memories": [], "mental_models": [], "directives": []} (if include.facts set but empty)
|
||||
# It should NEVER return based_on: []
|
||||
|
||||
based_on = data.get("based_on")
|
||||
if based_on is not None:
|
||||
assert isinstance(based_on, dict), f"based_on should be dict or null, got {type(based_on)}: {based_on}"
|
||||
assert not isinstance(based_on, list), f"based_on should NEVER be a list! Got: {based_on}"
|
||||
assert "memories" in based_on
|
||||
assert "mental_models" in based_on
|
||||
assert "directives" in based_on
|
||||
# All should be empty lists
|
||||
assert based_on["memories"] == []
|
||||
assert based_on["mental_models"] == []
|
||||
assert based_on["directives"] == []
|
||||
|
||||
# Verify the structure is parseable as proper types
|
||||
assert isinstance(data["text"], str)
|
||||
if based_on is not None:
|
||||
# Verify it's the v0.4.0+ format (object with arrays)
|
||||
assert isinstance(based_on["memories"], list)
|
||||
assert isinstance(based_on["mental_models"], list)
|
||||
assert isinstance(based_on["directives"], list)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reflect_without_include_facts(api_client):
|
||||
"""Test reflect without requesting facts (based_on should be None)."""
|
||||
bank_id = "test_no_facts"
|
||||
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/reflect",
|
||||
json={
|
||||
"query": "Hello world",
|
||||
"budget": "low"
|
||||
# No include.facts
|
||||
}
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
|
||||
# When include.facts is not set, based_on should not be in response (or be null)
|
||||
based_on = data.get("based_on")
|
||||
assert based_on is None, f"based_on should be None when not requested, got {type(based_on)}: {based_on}"
|
||||
|
||||
# Verify structure
|
||||
assert isinstance(data["text"], str)
|
||||
@@ -0,0 +1,125 @@
|
||||
"""
|
||||
Test ReflectResponse parsing for different API versions.
|
||||
|
||||
This tests the client's ability to parse reflect responses from:
|
||||
- v0.3.0 API (based_on as list)
|
||||
- v0.4.0+ API (based_on as object)
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from hindsight_client_api.models.reflect_response import ReflectResponse
|
||||
from hindsight_client_api.models.reflect_based_on import ReflectBasedOn
|
||||
|
||||
|
||||
def test_parse_v4_format_with_empty_based_on():
|
||||
"""Test parsing v0.4.0+ format with empty based_on object."""
|
||||
response_data = {
|
||||
"text": "I don't have any information about that.",
|
||||
"based_on": {
|
||||
"memories": [],
|
||||
"mental_models": [],
|
||||
"directives": []
|
||||
}
|
||||
}
|
||||
|
||||
response = ReflectResponse.from_dict(response_data)
|
||||
assert response is not None
|
||||
assert response.text == "I don't have any information about that."
|
||||
assert response.based_on is not None
|
||||
assert isinstance(response.based_on, ReflectBasedOn)
|
||||
assert response.based_on.memories == []
|
||||
assert response.based_on.mental_models == []
|
||||
assert response.based_on.directives == []
|
||||
|
||||
|
||||
def test_parse_v4_format_with_null_based_on():
|
||||
"""Test parsing v0.4.0+ format with null based_on (include.facts not set)."""
|
||||
response_data = {
|
||||
"text": "Hello!",
|
||||
"based_on": None
|
||||
}
|
||||
|
||||
response = ReflectResponse.from_dict(response_data)
|
||||
assert response is not None
|
||||
assert response.text == "Hello!"
|
||||
assert response.based_on is None
|
||||
|
||||
|
||||
def test_parse_v4_format_with_populated_based_on():
|
||||
"""Test parsing v0.4.0+ format with actual facts."""
|
||||
response_data = {
|
||||
"text": "Based on my knowledge, AI is transformative.",
|
||||
"based_on": {
|
||||
"memories": [
|
||||
{
|
||||
"id": "mem-123",
|
||||
"text": "AI is used in healthcare",
|
||||
"type": "world",
|
||||
"context": None,
|
||||
"occurred_start": None,
|
||||
"occurred_end": None
|
||||
}
|
||||
],
|
||||
"mental_models": [
|
||||
{
|
||||
"id": "mm-456",
|
||||
"text": "AI transforms industries",
|
||||
"context": "technology trends"
|
||||
}
|
||||
],
|
||||
"directives": [
|
||||
{
|
||||
"id": "dir-789",
|
||||
"name": "Be concise",
|
||||
"content": "Keep responses brief"
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
response = ReflectResponse.from_dict(response_data)
|
||||
assert response is not None
|
||||
assert response.text == "Based on my knowledge, AI is transformative."
|
||||
assert response.based_on is not None
|
||||
assert len(response.based_on.memories) == 1
|
||||
assert response.based_on.memories[0].id == "mem-123"
|
||||
assert len(response.based_on.mental_models) == 1
|
||||
assert response.based_on.mental_models[0].id == "mm-456"
|
||||
assert len(response.based_on.directives) == 1
|
||||
assert response.based_on.directives[0].id == "dir-789"
|
||||
|
||||
|
||||
def test_parse_v3_format_with_empty_list_fails():
|
||||
"""
|
||||
Test that v0.3.0 format (based_on as list) fails validation.
|
||||
|
||||
This is a BREAKING CHANGE from v0.3.0 to v0.4.0.
|
||||
Clients using v0.4.x SDK cannot parse v0.3.0 API responses.
|
||||
|
||||
Users must either:
|
||||
- Upgrade API to v0.4.0+
|
||||
- Use v0.3.0 client with v0.3.0 API
|
||||
"""
|
||||
response_data = {
|
||||
"text": "No information available.",
|
||||
"based_on": [] # v0.3.0 format - incompatible with v0.4.0+ client
|
||||
}
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
ReflectResponse.from_dict(response_data)
|
||||
|
||||
# Should fail with validation error
|
||||
assert "ValidationError" in str(type(exc_info.value).__name__) or "validation" in str(exc_info.value).lower()
|
||||
|
||||
|
||||
def test_parse_missing_based_on_field():
|
||||
"""Test parsing response when based_on field is omitted entirely."""
|
||||
response_data = {
|
||||
"text": "Hello!"
|
||||
# based_on field not present
|
||||
}
|
||||
|
||||
response = ReflectResponse.from_dict(response_data)
|
||||
assert response is not None
|
||||
assert response.text == "Hello!"
|
||||
assert response.based_on is None
|
||||
@@ -269,15 +269,16 @@ export HINDSIGHT_API_RETAIN_LLM_MAX_BACKOFF=120.0 # Cap at 2min instead of 1m
|
||||
|----------|-------------|---------|
|
||||
| `HINDSIGHT_API_EMBEDDINGS_PROVIDER` | Provider: `local`, `tei`, `openai`, `cohere`, or `litellm` | `local` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL` | Model for local provider | `BAAI/bge-small-en-v1.5` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_LOCAL_TRUST_REMOTE_CODE` | Allow loading models with custom code (security risk, disabled by default) | `false` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_TEI_URL` | TEI server URL | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY` | OpenAI API key (falls back to `HINDSIGHT_API_LLM_API_KEY`) | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL` | OpenAI embedding model | `text-embedding-3-small` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL` | Custom base URL for OpenAI-compatible API (e.g., Azure OpenAI) | - |
|
||||
| `HINDSIGHT_API_COHERE_API_KEY` | Cohere API key (shared for embeddings and reranker) | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_COHERE_API_KEY` | Cohere API key for embeddings | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL` | Cohere embedding model | `embed-english-v3.0` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL` | Custom base URL for Cohere-compatible API (e.g., Azure-hosted) | - |
|
||||
| `HINDSIGHT_API_LITELLM_API_BASE` | LiteLLM proxy base URL (shared for embeddings and reranker) | `http://localhost:4000` |
|
||||
| `HINDSIGHT_API_LITELLM_API_KEY` | LiteLLM proxy API key (optional, depends on proxy config) | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_LITELLM_API_BASE` | LiteLLM proxy base URL for embeddings | `http://localhost:4000` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_LITELLM_API_KEY` | LiteLLM proxy API key for embeddings (optional, depends on proxy config) | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL` | LiteLLM embedding model (use provider prefix, e.g., `cohere/embed-english-v3.0`) | `text-embedding-3-small` |
|
||||
|
||||
```bash
|
||||
@@ -285,6 +286,11 @@ export HINDSIGHT_API_RETAIN_LLM_MAX_BACKOFF=120.0 # Cap at 2min instead of 1m
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=local
|
||||
export HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL=BAAI/bge-small-en-v1.5
|
||||
|
||||
# Local with custom model requiring trust_remote_code
|
||||
# WARNING: Only enable trust_remote_code for models you trust (security risk)
|
||||
# export HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL=your-custom-model
|
||||
# export HINDSIGHT_API_EMBEDDINGS_LOCAL_TRUST_REMOTE_CODE=true
|
||||
|
||||
# OpenAI - cloud-based embeddings
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=openai
|
||||
export HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=sk-xxxxxxxxxxxx # or reuses HINDSIGHT_API_LLM_API_KEY
|
||||
@@ -302,19 +308,19 @@ export HINDSIGHT_API_EMBEDDINGS_TEI_URL=http://localhost:8080
|
||||
|
||||
# Cohere - cloud-based embeddings
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-api-key
|
||||
export HINDSIGHT_API_EMBEDDINGS_COHERE_API_KEY=your-api-key
|
||||
export HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL=embed-english-v3.0 # 1024 dimensions
|
||||
|
||||
# Azure-hosted Cohere - embeddings via custom endpoint
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-azure-api-key
|
||||
export HINDSIGHT_API_EMBEDDINGS_COHERE_API_KEY=your-azure-api-key
|
||||
export HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL=embed-english-v3.0
|
||||
export HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL=https://your-azure-cohere-endpoint.com
|
||||
|
||||
# LiteLLM proxy - unified gateway for multiple providers
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=litellm
|
||||
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
|
||||
export HINDSIGHT_API_LITELLM_API_KEY=your-litellm-key # optional
|
||||
export HINDSIGHT_API_EMBEDDINGS_LITELLM_API_BASE=http://localhost:4000
|
||||
export HINDSIGHT_API_EMBEDDINGS_LITELLM_API_KEY=your-litellm-key # optional
|
||||
export HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL=text-embedding-3-small # or cohere/embed-english-v3.0
|
||||
```
|
||||
|
||||
@@ -341,11 +347,15 @@ Supported OpenAI embedding dimensions:
|
||||
| `HINDSIGHT_API_RERANKER_PROVIDER` | Provider: `local`, `tei`, `cohere`, `flashrank`, `litellm`, or `rrf` | `local` |
|
||||
| `HINDSIGHT_API_RERANKER_LOCAL_MODEL` | Model for local provider | `cross-encoder/ms-marco-MiniLM-L-6-v2` |
|
||||
| `HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT` | Max concurrent local reranking (prevents CPU thrashing under load) | `4` |
|
||||
| `HINDSIGHT_API_RERANKER_LOCAL_TRUST_REMOTE_CODE` | Allow loading models with custom code (security risk, disabled by default) | `false` |
|
||||
| `HINDSIGHT_API_RERANKER_TEI_URL` | TEI server URL | - |
|
||||
| `HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE` | Batch size for TEI reranking | `128` |
|
||||
| `HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT` | Max concurrent TEI reranking requests | `8` |
|
||||
| `HINDSIGHT_API_RERANKER_COHERE_API_KEY` | Cohere API key for reranking | - |
|
||||
| `HINDSIGHT_API_RERANKER_COHERE_MODEL` | Cohere rerank model | `rerank-english-v3.0` |
|
||||
| `HINDSIGHT_API_RERANKER_COHERE_BASE_URL` | Custom base URL for Cohere-compatible API (e.g., Azure-hosted) | - |
|
||||
| `HINDSIGHT_API_RERANKER_LITELLM_API_BASE` | LiteLLM proxy base URL for reranking | `http://localhost:4000` |
|
||||
| `HINDSIGHT_API_RERANKER_LITELLM_API_KEY` | LiteLLM proxy API key for reranking (optional, depends on proxy config) | - |
|
||||
| `HINDSIGHT_API_RERANKER_LITELLM_MODEL` | LiteLLM rerank model (use provider prefix, e.g., `cohere/rerank-english-v3.0`) | `cohere/rerank-english-v3.0` |
|
||||
| `HINDSIGHT_API_RERANKER_FLASHRANK_MODEL` | FlashRank model for fast CPU-based reranking | `ms-marco-MiniLM-L-12-v2` |
|
||||
| `HINDSIGHT_API_RERANKER_FLASHRANK_CACHE_DIR` | Cache directory for FlashRank models | System default |
|
||||
@@ -355,25 +365,31 @@ Supported OpenAI embedding dimensions:
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=local
|
||||
export HINDSIGHT_API_RERANKER_LOCAL_MODEL=cross-encoder/ms-marco-MiniLM-L-6-v2
|
||||
|
||||
# Local with custom model requiring trust_remote_code (e.g., jina-reranker-v2)
|
||||
# WARNING: Only enable trust_remote_code for models you trust (security risk)
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=local
|
||||
export HINDSIGHT_API_RERANKER_LOCAL_MODEL=jinaai/jina-reranker-v2-base-multilingual
|
||||
export HINDSIGHT_API_RERANKER_LOCAL_TRUST_REMOTE_CODE=true
|
||||
|
||||
# TEI - for high-performance inference
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=tei
|
||||
export HINDSIGHT_API_RERANKER_TEI_URL=http://localhost:8081
|
||||
|
||||
# Cohere - cloud-based reranking
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-api-key # shared with embeddings
|
||||
export HINDSIGHT_API_RERANKER_COHERE_API_KEY=your-api-key
|
||||
export HINDSIGHT_API_RERANKER_COHERE_MODEL=rerank-english-v3.0
|
||||
|
||||
# Azure-hosted Cohere - reranking via custom endpoint
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-azure-api-key
|
||||
export HINDSIGHT_API_RERANKER_COHERE_API_KEY=your-azure-api-key
|
||||
export HINDSIGHT_API_RERANKER_COHERE_MODEL=rerank-english-v3.0
|
||||
export HINDSIGHT_API_RERANKER_COHERE_BASE_URL=https://your-azure-cohere-endpoint.com
|
||||
|
||||
# LiteLLM proxy - unified gateway for multiple reranking providers
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=litellm
|
||||
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
|
||||
export HINDSIGHT_API_LITELLM_API_KEY=your-litellm-key # optional
|
||||
export HINDSIGHT_API_RERANKER_LITELLM_API_BASE=http://localhost:4000
|
||||
export HINDSIGHT_API_RERANKER_LITELLM_API_KEY=your-litellm-key # optional
|
||||
export HINDSIGHT_API_RERANKER_LITELLM_MODEL=cohere/rerank-english-v3.0 # or voyage/rerank-2, together_ai/...
|
||||
```
|
||||
|
||||
|
||||
@@ -3560,7 +3560,8 @@
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Tags",
|
||||
"description": "Tags for filtering"
|
||||
"description": "Tags for filtering",
|
||||
"default": []
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -3601,7 +3602,8 @@
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Tags",
|
||||
"description": "Tags for scoped visibility"
|
||||
"description": "Tags for scoped visibility",
|
||||
"default": []
|
||||
},
|
||||
"max_tokens": {
|
||||
"type": "integer",
|
||||
@@ -3613,7 +3615,8 @@
|
||||
},
|
||||
"trigger": {
|
||||
"$ref": "#/components/schemas/MentalModelTrigger",
|
||||
"description": "Trigger settings"
|
||||
"description": "Trigger settings",
|
||||
"default": {}
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -3789,7 +3792,8 @@
|
||||
"type": "string"
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Tags"
|
||||
"title": "Tags",
|
||||
"default": []
|
||||
},
|
||||
"created_at": {
|
||||
"anyOf": [
|
||||
@@ -3905,7 +3909,8 @@
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Tags",
|
||||
"description": "Tags associated with this document"
|
||||
"description": "Tags associated with this document",
|
||||
"default": []
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -4686,7 +4691,8 @@
|
||||
"type": "string"
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Tags"
|
||||
"title": "Tags",
|
||||
"default": []
|
||||
},
|
||||
"max_tokens": {
|
||||
"type": "integer",
|
||||
@@ -4694,7 +4700,8 @@
|
||||
"default": 2048
|
||||
},
|
||||
"trigger": {
|
||||
"$ref": "#/components/schemas/MentalModelTrigger"
|
||||
"$ref": "#/components/schemas/MentalModelTrigger",
|
||||
"default": {}
|
||||
},
|
||||
"last_refreshed_at": {
|
||||
"anyOf": [
|
||||
@@ -5008,7 +5015,8 @@
|
||||
},
|
||||
"include": {
|
||||
"$ref": "#/components/schemas/IncludeOptions",
|
||||
"description": "Options for including additional data (entities are included by default)"
|
||||
"description": "Options for including additional data (entities are included by default)",
|
||||
"default": {}
|
||||
},
|
||||
"tags": {
|
||||
"anyOf": [
|
||||
@@ -5333,7 +5341,8 @@
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Memories",
|
||||
"description": "Memory facts used to generate the response"
|
||||
"description": "Memory facts used to generate the response",
|
||||
"default": []
|
||||
},
|
||||
"mental_models": {
|
||||
"items": {
|
||||
@@ -5341,7 +5350,8 @@
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Mental Models",
|
||||
"description": "Mental models used during reflection"
|
||||
"description": "Mental models used during reflection",
|
||||
"default": []
|
||||
},
|
||||
"directives": {
|
||||
"items": {
|
||||
@@ -5349,7 +5359,8 @@
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Directives",
|
||||
"description": "Directives applied during reflection"
|
||||
"description": "Directives applied during reflection",
|
||||
"default": []
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -5825,7 +5836,8 @@
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Tool Calls",
|
||||
"description": "Tool calls made during reflection"
|
||||
"description": "Tool calls made during reflection",
|
||||
"default": []
|
||||
},
|
||||
"llm_calls": {
|
||||
"items": {
|
||||
@@ -5833,7 +5845,8 @@
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Llm Calls",
|
||||
"description": "LLM calls made during reflection"
|
||||
"description": "LLM calls made during reflection",
|
||||
"default": []
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
|
||||
Reference in New Issue
Block a user