Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 51 additions & 2 deletions mesa_api/v4_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
import hashlib
import logging
from datetime import datetime
from typing import Callable
from typing import Any, Callable, Literal

from fastapi import APIRouter, Depends, Header, HTTPException, Request
from pydantic import BaseModel, ConfigDict, Field, model_validator
Expand Down Expand Up @@ -248,6 +248,10 @@ class V4SearchRequest(BaseModel):
valid_at: datetime | None = None
valid_from: datetime | None = None
valid_to: datetime | None = None
graph_mode: Literal["enabled", "disabled"] = Field(
default="enabled",
description="Enable or disable only the graph retrieval lane for matched ablation.",
)

@model_validator(mode="after")
def validate_temporal_range(self) -> "V4SearchRequest":
Expand All @@ -256,6 +260,45 @@ def validate_temporal_range(self) -> "V4SearchRequest":
return self


class V4ScopeAudit(BaseModel):
"""Versioned proof that the canonical candidate pool was scoped pre-rank."""

model_config = ConfigDict(frozen=True)

contract_version: str
enforcement_stage: Literal["pre_rank"]
requested_scope: dict[str, Any]
requested_scope_identity: str
query_identity: str
evaluated_candidate_count: int = Field(ge=0)
excluded_candidate_count: int = Field(ge=0)
eligible_candidate_count: int = Field(ge=0)
exclusion_audit_hash: str


class V4GraphAblation(BaseModel):
"""Versioned identity for a supported graph ON/OFF matched execution."""

model_config = ConfigDict(frozen=True)

contract_version: str
mode: Literal["enabled", "disabled"]
pair_identity: str
query_identity: str
retrieval_config_identity: str
scope_identity: str


class V4SearchResponse(BaseModel):
model_config = ConfigDict(frozen=True)

session_id: str
dataset_ids: list[str]
results: list[dict[str, Any]]
scope_audit: V4ScopeAudit | None = None
graph_ablation: V4GraphAblation | None = None


def _active_principal(request: Request):
principal = getattr(request.state, "principal", None)
if principal is None or getattr(principal, "status", None) != "active":
Expand Down Expand Up @@ -1032,7 +1075,7 @@ async def insert_memory(
response["duplicate"] = True
return response

@router.post("/memory/search")
@router.post("/memory/search", response_model=V4SearchResponse)
async def search_memory(
request: Request,
payload: V4SearchRequest,
Expand All @@ -1048,6 +1091,8 @@ async def search_memory(
status_code=403, detail="Dataset is outside session scope"
)
try:
certification_metadata: dict[str, Any] = {}
principal = _active_principal(request)
results = await dao.search_v4_memory(
tenant_id=str(session["tenant_id"]),
agent_id=str(session["agent_id"]),
Expand All @@ -1060,6 +1105,9 @@ async def search_memory(
payload.valid_from.isoformat() if payload.valid_from else None
),
valid_to=payload.valid_to.isoformat() if payload.valid_to else None,
graph_enabled=payload.graph_mode == "enabled",
request_principal_id=str(principal.principal_id),
certification_metadata=certification_metadata,
)
except EmbeddingMigrationRequiredError:
raise HTTPException(
Expand All @@ -1086,6 +1134,7 @@ async def search_memory(
"session_id": payload.session_id,
"dataset_ids": datasets,
"results": results,
**certification_metadata,
}

@router.get("/mutations/{mutation_id}", response_model=V4MutationStatusResponse)
Expand Down
6 changes: 5 additions & 1 deletion mesa_client/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
import asyncio
import logging
import time
from typing import Any, Callable, Optional, TypeVar
from typing import Any, Callable, Literal, Optional, TypeVar
from urllib.parse import quote

import httpx
Expand Down Expand Up @@ -619,6 +619,7 @@ def search( # type: ignore[override]
valid_at: str | None = None,
valid_from: str | None = None,
valid_to: str | None = None,
graph_mode: Literal["enabled", "disabled"] = "enabled",
) -> dict[str, Any]:
return self._request(
"POST",
Expand All @@ -632,6 +633,7 @@ def search( # type: ignore[override]
"valid_at": valid_at,
"valid_from": valid_from,
"valid_to": valid_to,
**({"graph_mode": graph_mode} if graph_mode != "enabled" else {}),
},
)

Expand Down Expand Up @@ -943,6 +945,7 @@ async def search( # type: ignore[override]
valid_at: str | None = None,
valid_from: str | None = None,
valid_to: str | None = None,
graph_mode: Literal["enabled", "disabled"] = "enabled",
) -> dict[str, Any]:
return await self._request(
"POST",
Expand All @@ -956,6 +959,7 @@ async def search( # type: ignore[override]
"valid_at": valid_at,
"valid_from": valid_from,
"valid_to": valid_to,
**({"graph_mode": graph_mode} if graph_mode != "enabled" else {}),
},
)

Expand Down
Loading
Loading