From e66b62226822e818297396ba2c0e5df6b9346c94 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Yasin=20B=C3=BCy=C3=BCktepe?= Date: Sun, 27 Sep 2026 14:07:02 +0300 Subject: [PATCH 1/3] feat(v4): expose pre-rank scope audit contract --- mesa_api/v4_router.py | 34 +++- mesa_storage/dao.py | 165 ++++++++++++---- mesa_storage/retrieval_scope.py | 14 ++ tests/test_v4_api_contract.py | 2 + tests/test_v4_certification_contracts.py | 240 +++++++++++++++++++++++ 5 files changed, 418 insertions(+), 37 deletions(-) create mode 100644 tests/test_v4_certification_contracts.py diff --git a/mesa_api/v4_router.py b/mesa_api/v4_router.py index a4f2357..a423fe2 100644 --- a/mesa_api/v4_router.py +++ b/mesa_api/v4_router.py @@ -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 @@ -256,6 +256,31 @@ 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 V4SearchResponse(BaseModel): + model_config = ConfigDict(frozen=True) + + session_id: str + dataset_ids: list[str] + results: list[dict[str, Any]] + scope_audit: V4ScopeAudit | None = None + + def _active_principal(request: Request): principal = getattr(request.state, "principal", None) if principal is None or getattr(principal, "status", None) != "active": @@ -1032,7 +1057,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, @@ -1048,6 +1073,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"]), @@ -1060,6 +1087,8 @@ 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, + request_principal_id=str(principal.principal_id), + certification_metadata=certification_metadata, ) except EmbeddingMigrationRequiredError: raise HTTPException( @@ -1086,6 +1115,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) diff --git a/mesa_storage/dao.py b/mesa_storage/dao.py index 77875bd..cd41b7d 100644 --- a/mesa_storage/dao.py +++ b/mesa_storage/dao.py @@ -81,8 +81,10 @@ from mesa_storage.retrieval_scope import ( V4_RRF_DEFAULT_K, V4_RRF_LANE_WEIGHTS, + V4_SCOPE_AUDIT_CONTRACT_VERSION, build_v4_lexical_query, rrf_fuse_lanes, + stable_contract_hash, ) from mesa_storage.sqlite_engine import AsyncEngine from mesa_storage.vector_engine import SemanticRuntimeDisabledError, VectorEngine @@ -457,10 +459,9 @@ def classify_graph_object( cleaned_tail, flags=re.IGNORECASE, ) - has_sentence_period = ( - bool(re.search(r"\.\s+[a-zA-ZçğıöşüÇĞİÖŞÜ]", tail_no_abbr)) - or (tail_no_abbr.strip().endswith(".") and len(words) >= 3) - ) + has_sentence_period = bool( + re.search(r"\.\s+[a-zA-ZçğıöşüÇĞİÖŞÜ]", tail_no_abbr) + ) or (tail_no_abbr.strip().endswith(".") and len(words) >= 3) if ( len(cleaned_tail) > 60 or len(words) > 6 @@ -5517,6 +5518,8 @@ async def search_v4_memory( valid_to: str | None = None, rrf_k: int | None = None, rrf_weights: Mapping[str, float] | None = None, + request_principal_id: str | None = None, + certification_metadata: dict[str, Any] | None = None, ) -> list[dict[str, Any]]: """Dataset-filter every retrieval lane, then combine ranks with RRF.""" _assert_valid_agent_id(agent_id) @@ -5616,6 +5619,92 @@ async def search_v4_memory( ) as cursor: in_scope_assertion_rows = await cursor.fetchall() + audit_rows: list[aiosqlite.Row] = [] + if certification_metadata is not None: + audit_query = ( + "SELECT a.assertion_id, a.tenant_id, a.dataset_id, " + "COALESCE(c.external_id, '') AS external_dataset_id, " + "m.agent_id, a.jurisdiction, a.status, a.valid_from, a.valid_to " + "FROM v4_assertions a " + "JOIN memory_mutations m ON m.mutation_id = a.mutation_id " + "LEFT JOIN v4_catalog_identities c ON c.tenant_id = a.tenant_id " + "AND c.kind = 'dataset' AND c.physical_id = a.dataset_id " + f"WHERE a.tenant_id = ? AND {assertion_status_sql} " + "AND m.state = 'COMMITTED' " + + ( + "AND " + " AND ".join(provenance_filters) + if provenance_filters + else "" + ) + + " ORDER BY a.assertion_id, a.tenant_id, m.agent_id" + ) + async with db.execute( + audit_query, (tenant_id, *provenance_params) + ) as cursor: + audit_rows = await cursor.fetchall() + + if certification_metadata is not None: + requested_scope = { + "tenant_id": tenant_id, + "dataset_ids": sorted(set(dataset_ids)), + "agent_id": agent_id, + **( + {"principal_id": request_principal_id} + if request_principal_id is not None + else {} + ), + "jurisdiction": jurisdiction, + "valid_at": valid_at, + "valid_from": valid_from, + "valid_to": valid_to, + } + requested_scope_identity = stable_contract_hash(requested_scope) + query_identity = stable_contract_hash({"query": query}) + audit_decisions = [] + eligible_count = 0 + for row in audit_rows: + eligible = ( + str(row["tenant_id"]) == tenant_id + and str(row["dataset_id"]) in datasets + and str(row["agent_id"]) == agent_id + ) + eligible_count += int(eligible) + audit_decisions.append( + { + "assertion_id": str(row["assertion_id"]), + "tenant_id": str(row["tenant_id"]), + "dataset_id": str(row["external_dataset_id"]), + "agent_id": str(row["agent_id"]), + "jurisdiction": str(row["jurisdiction"] or ""), + "status": str(row["status"]), + "valid_from": str(row["valid_from"] or ""), + "valid_to": str(row["valid_to"] or ""), + "eligible": eligible, + } + ) + certification_metadata.update( + { + "scope_audit": { + "contract_version": V4_SCOPE_AUDIT_CONTRACT_VERSION, + "enforcement_stage": "pre_rank", + "requested_scope": requested_scope, + "requested_scope_identity": requested_scope_identity, + "query_identity": query_identity, + "evaluated_candidate_count": len(audit_decisions), + "excluded_candidate_count": len(audit_decisions) + - eligible_count, + "eligible_candidate_count": eligible_count, + "exclusion_audit_hash": stable_contract_hash( + { + "query_identity": query_identity, + "requested_scope_identity": requested_scope_identity, + "decisions": audit_decisions, + } + ), + }, + } + ) + raw_allowed_entity_ids = { str(row[1]) for row in artifact_rows if row[0] == "ENTITY" } @@ -5678,7 +5767,9 @@ async def search_v4_memory( if node_id in allowed_vector_ids and node_id not in seen_vec: ranked_assertion_ids.append(node_id) seen_vec.add(node_id) - vector_raw_distances[node_id] = float(v_row.get("_distance", 0.0)) + vector_raw_distances[node_id] = float( + v_row.get("_distance", 0.0) + ) if ranked_assertion_ids: vector_placeholders = ",".join("?" for _ in ranked_assertion_ids) async with self._sql.connection() as db: @@ -6557,37 +6648,41 @@ async def search_v4_memory( cand_prov = cand["materialized_provenance"] if not cand_prov: continue - results.append( - { - "entity": entities.get(eid, {}), - "candidate_id": cand["candidate_id"], - "evidence_id": cand["assertion_id"], - "assertion_id": cand["assertion_id"], - "source_chunk_id": cand["source_chunk_id"], - "document_id": cand["document_id"], - "evidence_span": cand["evidence_span"], - "raw_score": ( - cand["raw_scores"].get("vector") - if "vector" in cand["raw_scores"] - else cand["raw_scores"].get("assertion") - ), - "rrf_score": cand["rrf_score"], - "legal_factor": cand["legal_factor"], - "final_score": cand["rrf_score"] * cand["legal_factor"], - "provenance": cand_prov, - "matched_assertions": [ - p - for p in cand_prov - if p["assertion_id"] == cand["assertion_id"] - ], - "supporting_assertions": [ - p - for p in cand_prov - if p["assertion_id"] != cand["assertion_id"] - ], - "retrieval_provenance": retrieval_provenance, + result = { + "entity": entities.get(eid, {}), + "candidate_id": cand["candidate_id"], + "evidence_id": cand["assertion_id"], + "assertion_id": cand["assertion_id"], + "source_chunk_id": cand["source_chunk_id"], + "document_id": cand["document_id"], + "evidence_span": cand["evidence_span"], + "raw_score": ( + cand["raw_scores"].get("vector") + if "vector" in cand["raw_scores"] + else cand["raw_scores"].get("assertion") + ), + "rrf_score": cand["rrf_score"], + "legal_factor": cand["legal_factor"], + "final_score": cand["rrf_score"] * cand["legal_factor"], + "provenance": cand_prov, + "matched_assertions": [ + p for p in cand_prov if p["assertion_id"] == cand["assertion_id"] + ], + "supporting_assertions": [ + p for p in cand_prov if p["assertion_id"] != cand["assertion_id"] + ], + "retrieval_provenance": retrieval_provenance, + } + if certification_metadata is not None: + first_provenance = cand_prov[0] + result["scope_identity"] = { + "tenant_id": tenant_id, + "dataset_id": str(first_provenance.get("dataset_id") or ""), + "agent_id": agent_id, + "jurisdiction": str(first_provenance.get("jurisdiction") or ""), + "status": str(first_provenance.get("status") or ""), } - ) + results.append(result) return results async def count_active_memories( diff --git a/mesa_storage/retrieval_scope.py b/mesa_storage/retrieval_scope.py index 643d9d7..4148743 100644 --- a/mesa_storage/retrieval_scope.py +++ b/mesa_storage/retrieval_scope.py @@ -1,10 +1,13 @@ """Dataset ownership filtering shared by live and rebuild retrieval paths.""" +import hashlib +import json from collections.abc import Iterable, Mapping, Sequence from typing import Any V4_RRF_DEFAULT_K = 60 V4_RRF_LANE_ORDER = ("vector", "bm25", "assertion", "graph") +V4_SCOPE_AUDIT_CONTRACT_VERSION = "mesa.scope-audit.v1" V4_RRF_LANE_WEIGHTS: dict[str, float] = { "vector": 1.0, "bm25": 1.0, @@ -13,6 +16,17 @@ } +def stable_contract_hash(value: Any) -> str: + """Hash a JSON contract value using deterministic, type-preserving encoding.""" + encoded = json.dumps( + value, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ).encode("utf-8") + return f"sha256:{hashlib.sha256(encoded).hexdigest()}" + + def compute_rrf_lane_score( rank: int, *, diff --git a/tests/test_v4_api_contract.py b/tests/test_v4_api_contract.py index 58c2054..50f4b09 100644 --- a/tests/test_v4_api_contract.py +++ b/tests/test_v4_api_contract.py @@ -1058,6 +1058,8 @@ async def test_v4_catalog_search_mutation_and_session_lifecycle_contracts( valid_at=None, valid_from=None, valid_to=None, + request_principal_id="principal-a", + certification_metadata={}, ) status = await client.get("/v4/mutations/mutation-a") diff --git a/tests/test_v4_certification_contracts.py b/tests/test_v4_certification_contracts.py new file mode 100644 index 0000000..a837de9 --- /dev/null +++ b/tests/test_v4_certification_contracts.py @@ -0,0 +1,240 @@ +"""Public contracts used by MESA E2E scope and graph certification.""" + +import asyncio +import inspect +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest +from fastapi import Depends, FastAPI, Request +from test_v4_retrieval_hardening_phase1 import _create_committed_mutation + +from mesa_api.v4_router import create_v4_router +from mesa_storage.dao import MemoryDAO +from mesa_storage.schemas import initialize_schema +from mesa_storage.sqlite_engine import AsyncEngine + + +def test_scope_audit_is_emitted_before_canonical_fusion(): + source = inspect.getsource(MemoryDAO.search_v4_memory) + assert source.index("certification_metadata.update") < source.index( + "rrf_fuse_lanes(" + ) + + +async def _environment(tmp_path): + sql = AsyncEngine(str(tmp_path / "certification.sqlite")) + await sql.initialize() + await initialize_schema(sql) + vector = SimpleNamespace( + compute_embedding=AsyncMock(return_value=[1.0, 0.0]), + compute_query_embedding=AsyncMock(return_value=[1.0, 0.0]), + upsert=AsyncMock(), + search=AsyncMock(return_value=[]), + ) + graph = SimpleNamespace( + insert_node=AsyncMock(), + insert_assertion=AsyncMock(), + link_assertions=AsyncMock(), + search_v4_graph=AsyncMock(return_value=[]), + is_operational=True, + ) + return sql, vector, graph, MemoryDAO(sql, vector, graph) + + +async def _add(dao, number, *, tenant, agent, dataset, subject="Shared policy"): + return await asyncio.wait_for( + _create_committed_mutation( + dao, + raw_log_id=number, + tenant_id=tenant, + agent_id=agent, + dataset_id=dataset, + chunk_id=f"chunk-{number}", + document_id=f"document-{number}", + content=f"{subject} applies", + subject=subject, + predicate="applies", + object_value=f"Target {number}", + evidence_span=f"{subject} applies", + ), + timeout=10, + ) + + +@pytest.mark.asyncio +async def test_scope_audit_is_pre_rank_deterministic_and_scope_sensitive(tmp_path): + sql, _vector, _graph, dao = await _environment(tmp_path) + try: + allowed = await _add( + dao, 1, tenant="tenant-a", agent="agent-a", dataset="dataset-a" + ) + baseline = await dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-a", + dataset_ids=["dataset-a"], + query="Shared policy", + ) + await _add(dao, 2, tenant="tenant-a", agent="agent-b", dataset="dataset-a") + before_cross_tenant = {} + await dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-a", + dataset_ids=["dataset-a"], + query="Shared policy", + certification_metadata=before_cross_tenant, + ) + await _add(dao, 3, tenant="tenant-b", agent="agent-b", dataset="dataset-b") + + first_metadata = {} + first = await asyncio.wait_for( + dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-a", + dataset_ids=["dataset-a"], + query="Shared policy", + certification_metadata=first_metadata, + ), + timeout=10, + ) + second_metadata = {} + second = await asyncio.wait_for( + dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-a", + dataset_ids=["dataset-a"], + query="Shared policy", + certification_metadata=second_metadata, + ), + timeout=10, + ) + + assert [row["assertion_id"] for row in first] == [allowed["assertion_id"]] + assert [ + (row["assertion_id"], row["rrf_score"], row["final_score"]) + for row in first + ] == [ + (row["assertion_id"], row["rrf_score"], row["final_score"]) + for row in baseline + ] + assert first == second + assert first_metadata == second_metadata + assert first_metadata == before_cross_tenant + audit = first_metadata["scope_audit"] + assert audit["contract_version"] == "mesa.scope-audit.v1" + assert audit["enforcement_stage"] == "pre_rank" + assert audit["evaluated_candidate_count"] == 2 + assert audit["eligible_candidate_count"] == 1 + assert audit["excluded_candidate_count"] == 1 + assert first[0]["scope_identity"] == { + "tenant_id": "tenant-a", + "dataset_id": "dataset-a", + "agent_id": "agent-a", + "jurisdiction": "", + "status": "ACTIVE", + } + serialized = repr({"results": first, **first_metadata}) + assert "dataset_physical_id" not in serialized + assert "registry_id" not in serialized + assert "mutation_id" not in first[0]["scope_identity"] + + unaudited = await dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-a", + dataset_ids=["dataset-a"], + query="Shared policy", + ) + assert [ + (row["assertion_id"], row["rrf_score"], row["final_score"]) for row in first + ] == [ + (row["assertion_id"], row["rrf_score"], row["final_score"]) + for row in unaudited + ] + + changed_metadata = {} + await asyncio.wait_for( + dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-b", + dataset_ids=["dataset-a"], + query="Shared policy", + certification_metadata=changed_metadata, + ), + timeout=10, + ) + assert ( + changed_metadata["scope_audit"]["requested_scope_identity"] + != audit["requested_scope_identity"] + ) + assert ( + changed_metadata["scope_audit"]["exclusion_audit_hash"] + != audit["exclusion_audit_hash"] + ) + finally: + await sql.close() + + +@pytest.mark.asyncio +async def test_http_contract_serializes_scope_certification(): + dao = MagicMock() + dao.rebuild_admission.is_pending = AsyncMock(return_value=False) + dao.get_v4_session = AsyncMock( + return_value={ + "tenant_id": "tenant-a", + "workspace_id": "workspace-a", + "dataset_ids": ["dataset-a"], + "agent_id": "agent-a", + "session_id": "session-a", + "status": "ACTIVE", + } + ) + + async def search_v4_memory(**kwargs): + metadata = kwargs["certification_metadata"] + metadata.update( + { + "scope_audit": { + "contract_version": "mesa.scope-audit.v1", + "enforcement_stage": "pre_rank", + "requested_scope": {"tenant_id": "tenant-a"}, + "requested_scope_identity": "sha256:scope", + "query_identity": "sha256:query", + "evaluated_candidate_count": 1, + "excluded_candidate_count": 0, + "eligible_candidate_count": 1, + "exclusion_audit_hash": "sha256:audit", + }, + } + ) + return [{"retrieval_provenance": {}}] + + dao.search_v4_memory = AsyncMock(side_effect=search_v4_memory) + access = MagicMock() + access.check_principal_session_access = AsyncMock(return_value=True) + access.check_access = AsyncMock(return_value=True) + access.check_scope_role = AsyncMock(return_value=True) + access.check_dataset_permission = AsyncMock(return_value=True) + + async def principal(request: Request): + request.state.principal = SimpleNamespace( + principal_id="principal-a", status="active" + ) + + app = FastAPI(dependencies=[Depends(principal)]) + app.include_router( + create_v4_router(get_dao=lambda: dao, get_access_control=lambda: access) + ) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), base_url="http://test" + ) as client: + response = await client.post( + "/v4/memory/search", + json={"session_id": "session-a", "query": "q"}, + ) + + assert response.status_code == 200 + body = response.json() + assert body["scope_audit"]["contract_version"] == "mesa.scope-audit.v1" + kwargs = dao.search_v4_memory.await_args.kwargs + assert kwargs["request_principal_id"] == "principal-a" From 5b552b5eb31514a2f78441ceeba648b09b551134 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Yasin=20B=C3=BCy=C3=BCktepe?= Date: Sun, 27 Sep 2026 14:09:51 +0300 Subject: [PATCH 2/3] feat(v4): add supported graph ablation contract --- mesa_api/v4_router.py | 19 ++++ mesa_client/client.py | 6 +- mesa_storage/dao.py | 34 +++++- mesa_storage/retrieval_scope.py | 19 ++++ tests/test_v4_api_contract.py | 1 + tests/test_v4_certification_contracts.py | 138 ++++++++++++++++++++++- tests/test_v4_sdk_contract.py | 31 +++++ 7 files changed, 241 insertions(+), 7 deletions(-) diff --git a/mesa_api/v4_router.py b/mesa_api/v4_router.py index a423fe2..ebf8982 100644 --- a/mesa_api/v4_router.py +++ b/mesa_api/v4_router.py @@ -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": @@ -272,6 +276,19 @@ class V4ScopeAudit(BaseModel): 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) @@ -279,6 +296,7 @@ class V4SearchResponse(BaseModel): dataset_ids: list[str] results: list[dict[str, Any]] scope_audit: V4ScopeAudit | None = None + graph_ablation: V4GraphAblation | None = None def _active_principal(request: Request): @@ -1087,6 +1105,7 @@ 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, ) diff --git a/mesa_client/client.py b/mesa_client/client.py index 2ef45a0..b87292c 100644 --- a/mesa_client/client.py +++ b/mesa_client/client.py @@ -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 @@ -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", @@ -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 {}), }, ) @@ -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", @@ -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 {}), }, ) diff --git a/mesa_storage/dao.py b/mesa_storage/dao.py index cd41b7d..ef13abf 100644 --- a/mesa_storage/dao.py +++ b/mesa_storage/dao.py @@ -79,12 +79,14 @@ build_v4_assertion_vector_payload, ) from mesa_storage.retrieval_scope import ( + V4_GRAPH_ABLATION_CONTRACT_VERSION, V4_RRF_DEFAULT_K, V4_RRF_LANE_WEIGHTS, V4_SCOPE_AUDIT_CONTRACT_VERSION, build_v4_lexical_query, rrf_fuse_lanes, stable_contract_hash, + stable_graph_path_id, ) from mesa_storage.sqlite_engine import AsyncEngine from mesa_storage.vector_engine import SemanticRuntimeDisabledError, VectorEngine @@ -5518,6 +5520,7 @@ async def search_v4_memory( valid_to: str | None = None, rrf_k: int | None = None, rrf_weights: Mapping[str, float] | None = None, + graph_enabled: bool = True, request_principal_id: str | None = None, certification_metadata: dict[str, Any] | None = None, ) -> list[dict[str, Any]]: @@ -5660,6 +5663,13 @@ async def search_v4_memory( } requested_scope_identity = stable_contract_hash(requested_scope) query_identity = stable_contract_hash({"query": query}) + retrieval_config_identity = stable_contract_hash( + { + "limit": limit, + "rrf_k": effective_k, + "rrf_weights": dict(effective_weights), + } + ) audit_decisions = [] eligible_count = 0 for row in audit_rows: @@ -5702,6 +5712,20 @@ async def search_v4_memory( } ), }, + "graph_ablation": { + "contract_version": V4_GRAPH_ABLATION_CONTRACT_VERSION, + "mode": "enabled" if graph_enabled else "disabled", + "pair_identity": stable_contract_hash( + { + "query_identity": query_identity, + "retrieval_config_identity": retrieval_config_identity, + "scope_identity": requested_scope_identity, + } + ), + "query_identity": query_identity, + "retrieval_config_identity": retrieval_config_identity, + "scope_identity": requested_scope_identity, + }, } ) @@ -6179,7 +6203,8 @@ async def search_v4_memory( allowed_graph_assertion_ids = in_scope_assertion_ids graph_provider = self._graph if ( - graph_provider is not None + graph_enabled + and graph_provider is not None and self.graph_implementation_available and graph_seed_ids and allowed_graph_assertion_ids @@ -6260,6 +6285,12 @@ async def search_v4_memory( "seed_id": path_entities[0], "score": float(path.get("score") or 0.0), } + normalized_path["graph_path_id"] = stable_graph_path_id( + assertion_ids=path_aids, + entity_ids=path_entities, + edge_directions=directions, + predicates=predicates, + ) path_key = (tuple(path_entities), tuple(path_aids)) if path_key not in seen_graph_paths: seen_graph_paths.add(path_key) @@ -6291,6 +6322,7 @@ async def search_v4_memory( "graph_target_entity_ids": targets, "graph_path_assertion_ids": list(best["assertion_ids"]), "graph_path_entity_ids": list(best["entity_ids"]), + "graph_path_id": best["graph_path_id"], "graph_edge_directions": list(best["edge_directions"]), "graph_predicates": list(best["predicates"]), "graph_direction": ( diff --git a/mesa_storage/retrieval_scope.py b/mesa_storage/retrieval_scope.py index 4148743..eb72dfc 100644 --- a/mesa_storage/retrieval_scope.py +++ b/mesa_storage/retrieval_scope.py @@ -8,6 +8,7 @@ V4_RRF_DEFAULT_K = 60 V4_RRF_LANE_ORDER = ("vector", "bm25", "assertion", "graph") V4_SCOPE_AUDIT_CONTRACT_VERSION = "mesa.scope-audit.v1" +V4_GRAPH_ABLATION_CONTRACT_VERSION = "mesa.graph-ablation.v1" V4_RRF_LANE_WEIGHTS: dict[str, float] = { "vector": 1.0, "bm25": 1.0, @@ -27,6 +28,24 @@ def stable_contract_hash(value: Any) -> str: return f"sha256:{hashlib.sha256(encoded).hexdigest()}" +def stable_graph_path_id( + *, + assertion_ids: Sequence[str], + entity_ids: Sequence[str], + edge_directions: Sequence[str], + predicates: Sequence[str], +) -> str: + """Return the stable semantic identity of one ordered graph path.""" + return stable_contract_hash( + { + "assertion_ids": list(assertion_ids), + "entity_ids": list(entity_ids), + "edge_directions": list(edge_directions), + "predicates": list(predicates), + } + ) + + def compute_rrf_lane_score( rank: int, *, diff --git a/tests/test_v4_api_contract.py b/tests/test_v4_api_contract.py index 50f4b09..d1947d6 100644 --- a/tests/test_v4_api_contract.py +++ b/tests/test_v4_api_contract.py @@ -1058,6 +1058,7 @@ async def test_v4_catalog_search_mutation_and_session_lifecycle_contracts( valid_at=None, valid_from=None, valid_to=None, + graph_enabled=True, request_principal_id="principal-a", certification_metadata={}, ) diff --git a/tests/test_v4_certification_contracts.py b/tests/test_v4_certification_contracts.py index a837de9..d60c9c2 100644 --- a/tests/test_v4_certification_contracts.py +++ b/tests/test_v4_certification_contracts.py @@ -12,6 +12,7 @@ from mesa_api.v4_router import create_v4_router from mesa_storage.dao import MemoryDAO +from mesa_storage.retrieval_scope import stable_graph_path_id from mesa_storage.schemas import initialize_schema from mesa_storage.sqlite_engine import AsyncEngine @@ -112,8 +113,7 @@ async def test_scope_audit_is_pre_rank_deterministic_and_scope_sensitive(tmp_pat assert [row["assertion_id"] for row in first] == [allowed["assertion_id"]] assert [ - (row["assertion_id"], row["rrf_score"], row["final_score"]) - for row in first + (row["assertion_id"], row["rrf_score"], row["final_score"]) for row in first ] == [ (row["assertion_id"], row["rrf_score"], row["final_score"]) for row in baseline @@ -176,7 +176,116 @@ async def test_scope_audit_is_pre_rank_deterministic_and_scope_sensitive(tmp_pat @pytest.mark.asyncio -async def test_http_contract_serializes_scope_certification(): +async def test_graph_ablation_uses_same_fusion_path_and_preserves_pair_identity( + tmp_path, monkeypatch +): + sql, _vector, graph, dao = await _environment(tmp_path) + try: + fact = await _add( + dao, 1, tenant="tenant-a", agent="agent-a", dataset="dataset-a" + ) + graph.search_v4_graph.return_value = [ + { + "entity_id": fact["object_entity_id"], + "path_assertion_ids": [fact["assertion_id"]], + "path_entity_ids": [fact["subject_id"], fact["object_entity_id"]], + "score": 1.0, + } + ] + + import mesa_storage.dao as dao_module + + original_fuse = dao_module.rrf_fuse_lanes + fused_lanes = [] + + def recording_fuse(lanes, **kwargs): + fused_lanes.append({name: list(values) for name, values in lanes.items()}) + return original_fuse(lanes, **kwargs) + + monkeypatch.setattr(dao_module, "rrf_fuse_lanes", recording_fuse) + enabled_metadata = {} + enabled = await dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-a", + dataset_ids=["dataset-a"], + query="Shared policy", + graph_enabled=True, + certification_metadata=enabled_metadata, + ) + disabled_metadata = {} + disabled = await dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-a", + dataset_ids=["dataset-a"], + query="Shared policy", + graph_enabled=False, + certification_metadata=disabled_metadata, + ) + + assert graph.search_v4_graph.await_count == 1 + assert len(fused_lanes) == 2 + assert fused_lanes[1]["graph"] == [] + assert { + key: fused_lanes[0][key] for key in ("vector", "bm25", "assertion") + } == {key: fused_lanes[1][key] for key in ("vector", "bm25", "assertion")} + assert [row["assertion_id"] for row in enabled] == [ + row["assertion_id"] for row in disabled + ] + enabled_contract = enabled_metadata["graph_ablation"] + disabled_contract = disabled_metadata["graph_ablation"] + assert enabled_contract["mode"] == "enabled" + assert disabled_contract["mode"] == "disabled" + assert enabled_contract["pair_identity"] == disabled_contract["pair_identity"] + changed_metadata = {} + await dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-a", + dataset_ids=["dataset-a"], + query="Different query", + graph_enabled=False, + certification_metadata=changed_metadata, + ) + assert ( + changed_metadata["graph_ablation"]["pair_identity"] + != disabled_contract["pair_identity"] + ) + changed_config_metadata = {} + await dao.search_v4_memory( + tenant_id="tenant-a", + agent_id="agent-a", + dataset_ids=["dataset-a"], + query="Shared policy", + limit=2, + graph_enabled=False, + certification_metadata=changed_config_metadata, + ) + assert ( + changed_config_metadata["graph_ablation"]["pair_identity"] + != disabled_contract["pair_identity"] + ) + path = enabled[0]["retrieval_provenance"]["graph_paths"][0] + assert ( + path["graph_path_id"] == enabled[0]["retrieval_provenance"]["graph_path_id"] + ) + finally: + await sql.close() + + +def test_graph_path_identity_is_stable_and_direction_sensitive(): + path = { + "assertion_ids": ["assertion-a"], + "entity_ids": ["entity-a", "entity-b"], + "edge_directions": ["forward"], + "predicates": ["applies"], + } + assert stable_graph_path_id(**path) == stable_graph_path_id(**path) + assert stable_graph_path_id(**path) != stable_graph_path_id( + **{**path, "edge_directions": ["reverse"]} + ) + + +@pytest.mark.asyncio +async def test_http_contract_validates_graph_mode_and_serializes_certification(): dao = MagicMock() dao.rebuild_admission.is_pending = AsyncMock(return_value=False) dao.get_v4_session = AsyncMock( @@ -205,9 +314,17 @@ async def search_v4_memory(**kwargs): "eligible_candidate_count": 1, "exclusion_audit_hash": "sha256:audit", }, + "graph_ablation": { + "contract_version": "mesa.graph-ablation.v1", + "mode": "disabled", + "pair_identity": "sha256:pair", + "query_identity": "sha256:query", + "retrieval_config_identity": "sha256:config", + "scope_identity": "sha256:scope", + }, } ) - return [{"retrieval_provenance": {}}] + return [{"retrieval_provenance": {"graph_path_id": "sha256:path"}}] dao.search_v4_memory = AsyncMock(side_effect=search_v4_memory) access = MagicMock() @@ -228,13 +345,24 @@ async def principal(request: Request): async with httpx.AsyncClient( transport=httpx.ASGITransport(app=app), base_url="http://test" ) as client: + invalid = await client.post( + "/v4/memory/search", + json={"session_id": "session-a", "query": "q", "graph_mode": "maybe"}, + ) response = await client.post( "/v4/memory/search", - json={"session_id": "session-a", "query": "q"}, + json={ + "session_id": "session-a", + "query": "q", + "graph_mode": "disabled", + }, ) + assert invalid.status_code == 422 assert response.status_code == 200 body = response.json() assert body["scope_audit"]["contract_version"] == "mesa.scope-audit.v1" + assert body["graph_ablation"]["mode"] == "disabled" kwargs = dao.search_v4_memory.await_args.kwargs + assert kwargs["graph_enabled"] is False assert kwargs["request_principal_id"] == "principal-a" diff --git a/tests/test_v4_sdk_contract.py b/tests/test_v4_sdk_contract.py index 29958e9..99ba69c 100644 --- a/tests/test_v4_sdk_contract.py +++ b/tests/test_v4_sdk_contract.py @@ -11,6 +11,37 @@ ) +def test_sync_v4_search_exposes_graph_ablation_without_changing_default_wire( + monkeypatch, +) -> None: + request = MagicMock(return_value={"results": []}) + monkeypatch.setattr(MesaV4Client, "_request", request) + client = MesaV4Client(base_url="http://mesa.invalid", api_key="test-key") + try: + client.search(session_id="session-a", query="q") + client.search(session_id="session-a", query="q", graph_mode="disabled") + finally: + client.close() + + default_payload = request.call_args_list[0].kwargs["json"] + disabled_payload = request.call_args_list[1].kwargs["json"] + assert "graph_mode" not in default_payload + assert disabled_payload["graph_mode"] == "disabled" + + +@pytest.mark.asyncio +async def test_async_v4_search_exposes_graph_ablation(monkeypatch) -> None: + request = AsyncMock(return_value={"results": []}) + monkeypatch.setattr(AsyncMesaV4Client, "_request", request) + client = AsyncMesaV4Client(base_url="http://mesa.invalid", api_key="test-key") + try: + await client.search(session_id="session-a", query="q", graph_mode="disabled") + finally: + await client.aclose() + + assert request.await_args.kwargs["json"]["graph_mode"] == "disabled" + + def test_sync_v4_capability_uses_versioned_contract(monkeypatch) -> None: request = MagicMock(return_value={"api_version": "v4"}) monkeypatch.setattr(MesaV4Client, "_request", request) From 2ef0a577df6f32b737b865ef1ac0928d2e6ce17f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Yasin=20B=C3=BCy=C3=BCktepe?= Date: Sun, 27 Sep 2026 14:44:01 +0300 Subject: [PATCH 3/3] fix(types): materialize scope audit rows --- mesa_storage/dao.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mesa_storage/dao.py b/mesa_storage/dao.py index ef13abf..cbcc8cd 100644 --- a/mesa_storage/dao.py +++ b/mesa_storage/dao.py @@ -5644,7 +5644,7 @@ async def search_v4_memory( async with db.execute( audit_query, (tenant_id, *provenance_params) ) as cursor: - audit_rows = await cursor.fetchall() + audit_rows = list(await cursor.fetchall()) if certification_metadata is not None: requested_scope = {