kontextual-engine/src/kontextual_engine/services/retrieval_service.py

1993 lines
73 KiB
Python

"""Governed retrieval contracts over the asset registry."""
from __future__ import annotations
from dataclasses import dataclass, field
from time import perf_counter
from typing import Any
from kontextual_engine.core import (
AuditEvent,
AuditOutcome,
AssetRepresentation,
ContextEntity,
ContextEntityType,
CoreRelationship,
KnowledgeAsset,
LifecycleState,
MetadataRecord,
OperationContext,
PolicyDecision,
RelationshipTargetKind,
RepresentationKind,
RetrievalFeedbackLabel,
RetrievalFeedbackRecord,
Sensitivity,
)
from kontextual_engine.errors import Diagnostic
from kontextual_engine.ports import AllowAllPolicyGateway, AssetRegistryRepository, PolicyGateway
SUPPORTED_SORT_KEYS = {
"asset_id",
"asset_type",
"created_at",
"lifecycle",
"title",
"updated_at",
}
SUPPORTED_CONTEXT_ENTITY_SORT_KEYS = {
"entity_id",
"entity_type",
"external_ref",
"name",
}
SUPPORTED_RELATIONSHIP_SORT_KEYS = {
"created_at",
"predicate",
"relationship_id",
"source_id",
"target_id",
"target_kind",
}
@dataclass(frozen=True)
class AssetQueryRequest:
text: str | None = None
asset_type: str | None = None
lifecycle: LifecycleState | str | None = None
sensitivity: Sensitivity | str | None = None
owner: str | None = None
topic: str | None = None
tags: tuple[str, ...] = ()
collection: str | None = None
review_state: str | None = None
metadata_filters: dict[str, Any] = field(default_factory=dict)
confirmed_metadata_only: bool = False
source_system: str | None = None
source_path: str | None = None
context_entity_id: str | None = None
context_entity_type: ContextEntityType | str | None = None
context_entity_name: str | None = None
context_entity_external_ref: str | None = None
workflow_run_id: str | None = None
related_asset_id: str | None = None
relationship_predicate: str | None = None
relationship_direction: str = "both"
include_relationships: bool = False
include_snippets: bool = False
max_snippets: int = 3
snippet_radius: int = 80
created_after: str | None = None
created_before: str | None = None
updated_after: str | None = None
updated_before: str | None = None
representation_kind: RepresentationKind | str | None = None
sort_by: str = "title"
sort_order: str = "asc"
limit: int = 50
offset: int = 0
def to_dict(self) -> dict[str, Any]:
return {
"text": self.text,
"asset_type": self.asset_type,
"lifecycle": self.lifecycle.value if isinstance(self.lifecycle, LifecycleState) else self.lifecycle,
"sensitivity": self.sensitivity.value if isinstance(self.sensitivity, Sensitivity) else self.sensitivity,
"owner": self.owner,
"topic": self.topic,
"tags": list(self.tags),
"collection": self.collection,
"review_state": self.review_state,
"metadata_filters": dict(self.metadata_filters),
"confirmed_metadata_only": self.confirmed_metadata_only,
"source_system": self.source_system,
"source_path": self.source_path,
"context_entity_id": self.context_entity_id,
"context_entity_type": self.context_entity_type.value
if isinstance(self.context_entity_type, ContextEntityType)
else self.context_entity_type,
"context_entity_name": self.context_entity_name,
"context_entity_external_ref": self.context_entity_external_ref,
"workflow_run_id": self.workflow_run_id,
"related_asset_id": self.related_asset_id,
"relationship_predicate": self.relationship_predicate,
"relationship_direction": self.relationship_direction,
"include_relationships": self.include_relationships,
"include_snippets": self.include_snippets,
"max_snippets": self.max_snippets,
"snippet_radius": self.snippet_radius,
"created_after": self.created_after,
"created_before": self.created_before,
"updated_after": self.updated_after,
"updated_before": self.updated_before,
"representation_kind": self.representation_kind.value
if isinstance(self.representation_kind, RepresentationKind)
else self.representation_kind,
"sort_by": self.sort_by,
"sort_order": self.sort_order,
"limit": self.limit,
"offset": self.offset,
}
@dataclass(frozen=True)
class AssetQueryItem:
asset: KnowledgeAsset
representations: tuple[AssetRepresentation, ...] = ()
metadata_records: tuple[MetadataRecord, ...] = ()
relationships: tuple[CoreRelationship, ...] = ()
context_entities: tuple[ContextEntity, ...] = ()
snippets: tuple["RetrievalSnippet", ...] = ()
diagnostics: tuple[Diagnostic, ...] = ()
relevance: dict[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
object.__setattr__(self, "representations", tuple(self.representations))
object.__setattr__(self, "metadata_records", tuple(self.metadata_records))
object.__setattr__(self, "relationships", tuple(self.relationships))
object.__setattr__(self, "context_entities", tuple(self.context_entities))
object.__setattr__(self, "snippets", tuple(self.snippets))
object.__setattr__(self, "diagnostics", tuple(self.diagnostics))
def to_dict(self) -> dict[str, Any]:
return {
"asset_id": self.asset.id,
"title": self.asset.title,
"lifecycle": self.asset.lifecycle.value,
"classification": self.asset.classification.to_dict(),
"current_version_id": self.asset.current_version_id,
"aliases": list(self.asset.aliases),
"asset_metadata": dict(self.asset.metadata),
"source_refs": [source_ref.to_dict() for source_ref in self.asset.source_refs],
"representations": [representation.to_dict() for representation in self.representations],
"metadata_records": [record.to_dict() for record in self.metadata_records],
"relationships": [relationship.to_dict() for relationship in self.relationships],
"context_entities": [entity.to_dict() for entity in self.context_entities],
"snippets": [snippet.to_dict() for snippet in self.snippets],
"relevance": dict(self.relevance),
"diagnostics": [diagnostic.to_dict() for diagnostic in self.diagnostics],
}
@dataclass(frozen=True)
class AssetQueryResult:
request: AssetQueryRequest
correlation_id: str
total: int
items: tuple[AssetQueryItem, ...] = ()
diagnostics: tuple[Diagnostic, ...] = ()
metadata: dict[str, Any] = field(default_factory=dict)
success: bool = True
def __post_init__(self) -> None:
object.__setattr__(self, "items", tuple(self.items))
object.__setattr__(self, "diagnostics", tuple(self.diagnostics))
@property
def result_count(self) -> int:
return len(self.items)
@property
def next_offset(self) -> int | None:
next_offset = self.request.offset + self.result_count
return next_offset if next_offset < self.total else None
def to_dict(self) -> dict[str, Any]:
return {
"query": self.request.to_dict(),
"correlation_id": self.correlation_id,
"success": self.success,
"total": self.total,
"result_count": self.result_count,
"limit": self.request.limit,
"offset": self.request.offset,
"next_offset": self.next_offset,
"sort": {
"by": self.request.sort_by,
"order": self.request.sort_order,
},
"metadata": dict(self.metadata),
"results": [item.to_dict() for item in self.items],
"diagnostics": [diagnostic.to_dict() for diagnostic in self.diagnostics],
}
@dataclass(frozen=True)
class LexicalIndexRefreshResult:
indexed_assets: int
indexed_representations: int
def to_dict(self) -> dict[str, int]:
return {
"indexed_assets": self.indexed_assets,
"indexed_representations": self.indexed_representations,
}
@dataclass(frozen=True)
class RetrievalSnippet:
asset_id: str
representation_id: str
text: str
start_offset: int
end_offset: int
match_text: str
media_type: str
source_ref_id: str | None = None
storage_ref: str | None = None
provenance: dict[str, Any] = field(default_factory=dict)
def to_dict(self) -> dict[str, Any]:
return {
"asset_id": self.asset_id,
"representation_id": self.representation_id,
"source_ref_id": self.source_ref_id,
"storage_ref": self.storage_ref,
"media_type": self.media_type,
"text": self.text,
"start_offset": self.start_offset,
"end_offset": self.end_offset,
"match_text": self.match_text,
"provenance": dict(self.provenance),
}
@dataclass(frozen=True)
class ContextEntityQueryRequest:
entity_id: str | None = None
entity_type: ContextEntityType | str | None = None
name: str | None = None
external_ref: str | None = None
metadata_filters: dict[str, Any] = field(default_factory=dict)
sort_by: str = "name"
sort_order: str = "asc"
limit: int = 50
offset: int = 0
def to_dict(self) -> dict[str, Any]:
return {
"entity_id": self.entity_id,
"entity_type": self.entity_type.value if isinstance(self.entity_type, ContextEntityType) else self.entity_type,
"name": self.name,
"external_ref": self.external_ref,
"metadata_filters": dict(self.metadata_filters),
"sort_by": self.sort_by,
"sort_order": self.sort_order,
"limit": self.limit,
"offset": self.offset,
}
@dataclass(frozen=True)
class ContextEntityQueryItem:
entity: ContextEntity
asset_ids: tuple[str, ...] = ()
relationship_count: int = 0
def __post_init__(self) -> None:
object.__setattr__(self, "asset_ids", tuple(self.asset_ids))
def to_dict(self) -> dict[str, Any]:
return {
**self.entity.to_dict(),
"asset_ids": list(self.asset_ids),
"relationship_count": self.relationship_count,
}
@dataclass(frozen=True)
class ContextEntityQueryResult:
request: ContextEntityQueryRequest
correlation_id: str
total: int
items: tuple[ContextEntityQueryItem, ...] = ()
diagnostics: tuple[Diagnostic, ...] = ()
success: bool = True
def __post_init__(self) -> None:
object.__setattr__(self, "items", tuple(self.items))
object.__setattr__(self, "diagnostics", tuple(self.diagnostics))
@property
def result_count(self) -> int:
return len(self.items)
@property
def next_offset(self) -> int | None:
next_offset = self.request.offset + self.result_count
return next_offset if next_offset < self.total else None
def to_dict(self) -> dict[str, Any]:
return {
"query": self.request.to_dict(),
"correlation_id": self.correlation_id,
"success": self.success,
"total": self.total,
"result_count": self.result_count,
"limit": self.request.limit,
"offset": self.request.offset,
"next_offset": self.next_offset,
"sort": {
"by": self.request.sort_by,
"order": self.request.sort_order,
},
"results": [item.to_dict() for item in self.items],
"diagnostics": [diagnostic.to_dict() for diagnostic in self.diagnostics],
}
@dataclass(frozen=True)
class RelationshipQueryRequest:
source_id: str | None = None
target_id: str | None = None
asset_id: str | None = None
context_entity_id: str | None = None
context_entity_type: ContextEntityType | str | None = None
context_entity_name: str | None = None
context_entity_external_ref: str | None = None
workflow_run_id: str | None = None
target_kind: RelationshipTargetKind | str | None = None
predicate: str | None = None
direction: str = "both"
sort_by: str = "source_id"
sort_order: str = "asc"
limit: int = 50
offset: int = 0
def to_dict(self) -> dict[str, Any]:
return {
"source_id": self.source_id,
"target_id": self.target_id,
"asset_id": self.asset_id,
"context_entity_id": self.context_entity_id,
"context_entity_type": self.context_entity_type.value
if isinstance(self.context_entity_type, ContextEntityType)
else self.context_entity_type,
"context_entity_name": self.context_entity_name,
"context_entity_external_ref": self.context_entity_external_ref,
"workflow_run_id": self.workflow_run_id,
"target_kind": self.target_kind.value if isinstance(self.target_kind, RelationshipTargetKind) else self.target_kind,
"predicate": self.predicate,
"direction": self.direction,
"sort_by": self.sort_by,
"sort_order": self.sort_order,
"limit": self.limit,
"offset": self.offset,
}
@dataclass(frozen=True)
class RelationshipQueryItem:
relationship: CoreRelationship
source_asset: KnowledgeAsset | None = None
target_asset: KnowledgeAsset | None = None
target_entity: ContextEntity | None = None
def to_dict(self) -> dict[str, Any]:
payload = self.relationship.to_dict()
if self.source_asset is not None:
payload["source_asset"] = self.source_asset.to_dict()
if self.target_asset is not None:
payload["target_asset"] = self.target_asset.to_dict()
if self.target_entity is not None:
payload["target_entity"] = self.target_entity.to_dict()
return payload
@dataclass(frozen=True)
class RelationshipQueryResult:
request: RelationshipQueryRequest
correlation_id: str
total: int
items: tuple[RelationshipQueryItem, ...] = ()
diagnostics: tuple[Diagnostic, ...] = ()
success: bool = True
def __post_init__(self) -> None:
object.__setattr__(self, "items", tuple(self.items))
object.__setattr__(self, "diagnostics", tuple(self.diagnostics))
@property
def result_count(self) -> int:
return len(self.items)
@property
def next_offset(self) -> int | None:
next_offset = self.request.offset + self.result_count
return next_offset if next_offset < self.total else None
def to_dict(self) -> dict[str, Any]:
return {
"query": self.request.to_dict(),
"correlation_id": self.correlation_id,
"success": self.success,
"total": self.total,
"result_count": self.result_count,
"limit": self.request.limit,
"offset": self.request.offset,
"next_offset": self.next_offset,
"sort": {
"by": self.request.sort_by,
"order": self.request.sort_order,
},
"results": [item.to_dict() for item in self.items],
"diagnostics": [diagnostic.to_dict() for diagnostic in self.diagnostics],
}
@dataclass(frozen=True)
class RetrievalFeedbackRequest:
label: RetrievalFeedbackLabel | str
query: dict[str, Any]
result_ref: dict[str, Any] = field(default_factory=dict)
notes: str | None = None
metadata: dict[str, Any] = field(default_factory=dict)
def to_dict(self) -> dict[str, Any]:
return {
"label": self.label.value if isinstance(self.label, RetrievalFeedbackLabel) else self.label,
"query": dict(self.query),
"result_ref": dict(self.result_ref),
"notes": self.notes,
"metadata": dict(self.metadata),
}
@dataclass(frozen=True)
class RetrievalFeedbackResult:
record: RetrievalFeedbackRecord | None
correlation_id: str
diagnostics: tuple[Diagnostic, ...] = ()
success: bool = True
def __post_init__(self) -> None:
object.__setattr__(self, "diagnostics", tuple(self.diagnostics))
def to_dict(self) -> dict[str, Any]:
return {
"success": self.success,
"correlation_id": self.correlation_id,
"record": self.record.to_dict() if self.record else None,
"diagnostics": [diagnostic.to_dict() for diagnostic in self.diagnostics],
}
@dataclass(frozen=True)
class RetrievalQualityMetrics:
query_count: int
zero_result_count: int
zero_result_rate: float
feedback_count: int
useful_count: int
unsafe_count: int
low_confidence_count: int
precision_at_k: float | None
citation_precision: float | None
permission_filter_observation_count: int
average_permission_filter_duration_ms: float | None
def to_dict(self) -> dict[str, Any]:
return {
"query_count": self.query_count,
"zero_result_count": self.zero_result_count,
"zero_result_rate": self.zero_result_rate,
"feedback_count": self.feedback_count,
"useful_count": self.useful_count,
"unsafe_count": self.unsafe_count,
"low_confidence_count": self.low_confidence_count,
"precision_at_k": self.precision_at_k,
"citation_precision": self.citation_precision,
"permission_filter_observation_count": self.permission_filter_observation_count,
"average_permission_filter_duration_ms": self.average_permission_filter_duration_ms,
}
@dataclass(frozen=True)
class _LexicalDocument:
asset_id: str
representation_id: str
text: str
media_type: str
source_ref_id: str | None = None
storage_ref: str | None = None
provenance: dict[str, Any] = field(default_factory=dict)
class AssetRetrievalService:
def __init__(
self,
repository: AssetRegistryRepository,
*,
policy_gateway: PolicyGateway | None = None,
) -> None:
self.repository = repository
self.policy_gateway = policy_gateway or AllowAllPolicyGateway()
self._lexical_index: tuple[_LexicalDocument, ...] = ()
self._last_refresh = LexicalIndexRefreshResult(indexed_assets=0, indexed_representations=0)
def refresh_index(self) -> LexicalIndexRefreshResult:
documents: list[_LexicalDocument] = []
for asset in self.repository.list_assets():
for representation in self.repository.list_representations(
asset_id=asset.id,
kind=RepresentationKind.NORMALIZED,
):
text = representation.metadata.get("search_text")
if not isinstance(text, str) or not text:
continue
documents.append(
_LexicalDocument(
asset_id=asset.id,
representation_id=representation.representation_id,
text=text,
media_type=representation.media_type,
source_ref_id=representation.source_ref_id,
storage_ref=representation.storage_ref,
provenance=_snippet_provenance(representation),
)
)
self._lexical_index = tuple(documents)
self._last_refresh = LexicalIndexRefreshResult(
indexed_assets=len({document.asset_id for document in documents}),
indexed_representations=len(documents),
)
return self._last_refresh
def query_assets(
self,
request: AssetQueryRequest,
context: OperationContext,
) -> AssetQueryResult:
diagnostics, normalized = _validate_request(request)
if diagnostics:
return AssetQueryResult(
request=request,
correlation_id=context.correlation_id,
total=0,
diagnostics=tuple(diagnostics),
success=False,
)
scope_decision = self._authorize_for_retrieval(
context,
"retrieval.assets.query",
"retrieval:assets",
resource_metadata={"query": normalized.to_dict()},
)
if not scope_decision.allowed:
self._audit_retrieval(
"retrieval.assets.query",
"retrieval:assets",
AuditOutcome.DENIED,
context,
scope_decision,
details={"query": normalized.to_dict()},
)
return AssetQueryResult(
request=normalized,
correlation_id=context.correlation_id,
total=0,
diagnostics=(
_permission_denied_diagnostic(scope_decision),
),
success=False,
)
assets = self.repository.list_assets(
lifecycle=normalized.lifecycle,
asset_type=normalized.asset_type,
sensitivity=normalized.sensitivity,
owner=normalized.owner,
topic=normalized.topic,
review_state=normalized.review_state,
metadata_filters=normalized.metadata_filters or None,
confirmed_metadata_only=normalized.confirmed_metadata_only,
)
assets = [
asset
for asset in assets
if _source_matches(
asset,
source_system=normalized.source_system,
source_path=normalized.source_path,
)
]
assets = [
asset
for asset in assets
if _collection_matches(
asset,
self.repository.list_metadata_records(asset.id),
normalized.collection,
)
and _tags_match(
asset,
self.repository.list_metadata_records(asset.id),
normalized.tags,
)
and _timestamp_matches(asset, normalized)
]
if normalized.representation_kind is not None:
assets = [
asset
for asset in assets
if self.repository.list_representations(
asset_id=asset.id,
kind=normalized.representation_kind,
)
]
relationship_context_by_asset: dict[str, tuple[CoreRelationship, ...]] = {}
context_entities_by_asset: dict[str, tuple[ContextEntity, ...]] = {}
if _asset_query_has_graph_filter(normalized):
graph_matches: list[KnowledgeAsset] = []
for asset in assets:
relationships, entities = self._relationship_context_for_asset(asset.id, normalized, context)
if relationships:
graph_matches.append(asset)
relationship_context_by_asset[asset.id] = relationships
context_entities_by_asset[asset.id] = entities
assets = graph_matches
relevance_by_asset: dict[str, dict[str, Any]] = {}
snippets_by_asset: dict[str, tuple[RetrievalSnippet, ...]] = {}
if normalized.text:
relevance_by_asset, snippets_by_asset = self._lexical_matches(
normalized.text,
include_snippets=normalized.include_snippets,
max_snippets=normalized.max_snippets,
snippet_radius=normalized.snippet_radius,
)
assets = [asset for asset in assets if asset.id in relevance_by_asset]
permission_filter_started = perf_counter()
asset_count_before_policy = len(assets)
assets = [asset for asset in assets if self._asset_allowed(asset, context)]
permission_filtered_count = asset_count_before_policy - len(assets)
permission_filter_duration_ms = _elapsed_ms(permission_filter_started)
ordered = _sort_assets(assets, normalized.sort_by, normalized.sort_order)
total = len(ordered)
page = ordered[normalized.offset : normalized.offset + normalized.limit]
items = tuple(
self._item_for_asset(
asset,
normalized.representation_kind,
relevance=relevance_by_asset.get(asset.id, {}),
include_relationships=normalized.include_relationships,
relationships=relationship_context_by_asset.get(asset.id),
context_entities=context_entities_by_asset.get(asset.id),
snippets=snippets_by_asset.get(asset.id, ()),
context=context,
)
for asset in page
)
result = AssetQueryResult(
request=normalized,
correlation_id=context.correlation_id,
total=total,
items=items,
metadata={
"zero_result": total == 0,
"graph_filter": _asset_query_has_graph_filter(normalized),
"policy_enforced": True,
"lexical_index": self._last_refresh.to_dict(),
},
)
self._audit_retrieval(
"retrieval.assets.query",
"retrieval:assets",
AuditOutcome.PARTIAL if permission_filtered_count else AuditOutcome.SUCCESS,
context,
scope_decision,
details={
"query": normalized.to_dict(),
"result_count": result.result_count,
"total": result.total,
"permission_filtered_count": permission_filtered_count,
"permission_filter_duration_ms": permission_filter_duration_ms,
},
)
return result
def query_context_entities(
self,
request: ContextEntityQueryRequest,
context: OperationContext,
) -> ContextEntityQueryResult:
diagnostics, normalized = _validate_context_entity_request(request)
if diagnostics:
return ContextEntityQueryResult(
request=request,
correlation_id=context.correlation_id,
total=0,
diagnostics=tuple(diagnostics),
success=False,
)
scope_decision = self._authorize_for_retrieval(
context,
"retrieval.context_entities.query",
"retrieval:context_entities",
resource_metadata={"query": normalized.to_dict()},
)
if not scope_decision.allowed:
self._audit_retrieval(
"retrieval.context_entities.query",
"retrieval:context_entities",
AuditOutcome.DENIED,
context,
scope_decision,
details={"query": normalized.to_dict()},
)
return ContextEntityQueryResult(
request=normalized,
correlation_id=context.correlation_id,
total=0,
diagnostics=(
_permission_denied_diagnostic(scope_decision),
),
success=False,
)
entities = [
entity
for entity in self.repository.list_context_entities()
if _context_entity_matches(
entity,
entity_id=normalized.entity_id,
entity_type=normalized.entity_type,
name=normalized.name,
external_ref=normalized.external_ref,
metadata_filters=normalized.metadata_filters,
)
]
permission_filter_started = perf_counter()
entity_count_before_policy = len(entities)
entities = [entity for entity in entities if self._context_entity_allowed(entity, context)]
permission_filtered_count = entity_count_before_policy - len(entities)
permission_filter_duration_ms = _elapsed_ms(permission_filter_started)
ordered = _sort_context_entities(entities, normalized.sort_by, normalized.sort_order)
total = len(ordered)
page = ordered[normalized.offset : normalized.offset + normalized.limit]
items = tuple(self._context_entity_item(entity, context) for entity in page)
result = ContextEntityQueryResult(
request=normalized,
correlation_id=context.correlation_id,
total=total,
items=items,
)
self._audit_retrieval(
"retrieval.context_entities.query",
"retrieval:context_entities",
AuditOutcome.PARTIAL if permission_filtered_count else AuditOutcome.SUCCESS,
context,
scope_decision,
details={
"query": normalized.to_dict(),
"result_count": result.result_count,
"total": result.total,
"permission_filtered_count": permission_filtered_count,
"permission_filter_duration_ms": permission_filter_duration_ms,
},
)
return result
def query_relationships(
self,
request: RelationshipQueryRequest,
context: OperationContext,
) -> RelationshipQueryResult:
diagnostics, normalized = _validate_relationship_request(request)
if diagnostics:
return RelationshipQueryResult(
request=request,
correlation_id=context.correlation_id,
total=0,
diagnostics=tuple(diagnostics),
success=False,
)
scope_decision = self._authorize_for_retrieval(
context,
"retrieval.relationships.query",
"retrieval:relationships",
resource_metadata={"query": normalized.to_dict()},
)
if not scope_decision.allowed:
self._audit_retrieval(
"retrieval.relationships.query",
"retrieval:relationships",
AuditOutcome.DENIED,
context,
scope_decision,
details={"query": normalized.to_dict()},
)
return RelationshipQueryResult(
request=normalized,
correlation_id=context.correlation_id,
total=0,
diagnostics=(
_permission_denied_diagnostic(scope_decision),
),
success=False,
)
relationships = self._relationships_for_request(normalized)
permission_filter_started = perf_counter()
relationship_count_before_policy = len(relationships)
relationships = [
relationship for relationship in relationships if self._relationship_allowed(relationship, context)
]
permission_filtered_count = relationship_count_before_policy - len(relationships)
permission_filter_duration_ms = _elapsed_ms(permission_filter_started)
ordered = _sort_relationships(relationships, normalized.sort_by, normalized.sort_order)
total = len(ordered)
page = ordered[normalized.offset : normalized.offset + normalized.limit]
items = tuple(self._relationship_item(relationship) for relationship in page)
result = RelationshipQueryResult(
request=normalized,
correlation_id=context.correlation_id,
total=total,
items=items,
)
self._audit_retrieval(
"retrieval.relationships.query",
"retrieval:relationships",
AuditOutcome.PARTIAL if permission_filtered_count else AuditOutcome.SUCCESS,
context,
scope_decision,
details={
"query": normalized.to_dict(),
"result_count": result.result_count,
"total": result.total,
"permission_filtered_count": permission_filtered_count,
"permission_filter_duration_ms": permission_filter_duration_ms,
},
)
return result
def record_feedback(
self,
request: RetrievalFeedbackRequest,
context: OperationContext,
) -> RetrievalFeedbackResult:
diagnostics: list[Diagnostic] = []
label = _parse_feedback_label(request.label, diagnostics)
if diagnostics:
return RetrievalFeedbackResult(
record=None,
correlation_id=context.correlation_id,
diagnostics=tuple(diagnostics),
success=False,
)
decision = self._authorize_for_retrieval(
context,
"retrieval.feedback.record",
"retrieval:feedback",
resource_metadata=request.to_dict(),
)
if not decision.allowed:
self._audit_retrieval(
"retrieval.feedback.record",
"retrieval:feedback",
AuditOutcome.DENIED,
context,
decision,
details=request.to_dict(),
)
return RetrievalFeedbackResult(
record=None,
correlation_id=context.correlation_id,
diagnostics=(
_permission_denied_diagnostic(decision),
),
success=False,
)
record = RetrievalFeedbackRecord(
label=label or RetrievalFeedbackLabel.LOW_CONFIDENCE,
query=dict(request.query),
result_ref=dict(request.result_ref),
actor_id=context.actor.id,
correlation_id=context.correlation_id,
notes=request.notes,
metadata=dict(request.metadata),
)
saved = self.repository.save_retrieval_feedback(record)
self._audit_retrieval(
"retrieval.feedback.record",
f"retrieval_feedback:{saved.feedback_id}",
AuditOutcome.SUCCESS,
context,
decision,
details={
"feedback_id": saved.feedback_id,
"label": saved.label.value,
"result_ref": dict(saved.result_ref),
},
)
return RetrievalFeedbackResult(record=saved, correlation_id=context.correlation_id)
def list_feedback(
self,
*,
correlation_id: str | None = None,
label: RetrievalFeedbackLabel | str | None = None,
) -> tuple[RetrievalFeedbackRecord, ...]:
label_value = label.value if isinstance(label, RetrievalFeedbackLabel) else label
return tuple(self.repository.list_retrieval_feedback(correlation_id=correlation_id, label=label_value))
def quality_metrics(
self,
*,
query_results: tuple[AssetQueryResult, ...] = (),
precision_at_k: int = 5,
) -> RetrievalQualityMetrics:
feedback = self.repository.list_retrieval_feedback()
query_count = len(query_results)
zero_result_count = sum(1 for result in query_results if result.total == 0)
ranked_feedback = [
record
for record in feedback
if _feedback_rank(record) is not None and (_feedback_rank(record) or 0) <= precision_at_k
]
useful_ranked = [record for record in ranked_feedback if record.label == RetrievalFeedbackLabel.USEFUL]
citation_feedback = [record for record in feedback if _feedback_has_citation_ref(record)]
useful_citations = [record for record in citation_feedback if record.label == RetrievalFeedbackLabel.USEFUL]
permission_latencies = [
float(event.details["permission_filter_duration_ms"])
for event in self.repository.list_audit_events()
if event.operation.startswith("retrieval.")
and "permission_filter_duration_ms" in event.details
]
return RetrievalQualityMetrics(
query_count=query_count,
zero_result_count=zero_result_count,
zero_result_rate=zero_result_count / query_count if query_count else 0.0,
feedback_count=len(feedback),
useful_count=sum(1 for record in feedback if record.label == RetrievalFeedbackLabel.USEFUL),
unsafe_count=sum(1 for record in feedback if record.label == RetrievalFeedbackLabel.UNSAFE),
low_confidence_count=sum(1 for record in feedback if record.label == RetrievalFeedbackLabel.LOW_CONFIDENCE),
precision_at_k=len(useful_ranked) / len(ranked_feedback) if ranked_feedback else None,
citation_precision=len(useful_citations) / len(citation_feedback) if citation_feedback else None,
permission_filter_observation_count=len(permission_latencies),
average_permission_filter_duration_ms=sum(permission_latencies) / len(permission_latencies)
if permission_latencies
else None,
)
def _item_for_asset(
self,
asset: KnowledgeAsset,
representation_kind: RepresentationKind | None,
*,
relevance: dict[str, Any] | None = None,
include_relationships: bool = False,
relationships: tuple[CoreRelationship, ...] | None = None,
context_entities: tuple[ContextEntity, ...] | None = None,
snippets: tuple[RetrievalSnippet, ...] = (),
context: OperationContext | None = None,
) -> AssetQueryItem:
if relationships is None and include_relationships:
relationships = self._relationships_for_asset(asset.id, "both")
if context is not None:
relationships = tuple(
relationship for relationship in relationships if self._relationship_allowed(relationship, context)
)
if context_entities is None and relationships:
context_entities = self._context_entities_for_relationships(relationships)
return AssetQueryItem(
asset=asset,
representations=tuple(
self.repository.list_representations(
asset_id=asset.id,
kind=representation_kind,
)
),
metadata_records=tuple(self.repository.list_metadata_records(asset.id)),
relationships=tuple(relationships or ()),
context_entities=tuple(context_entities or ()),
snippets=tuple(snippets),
relevance=dict(relevance or {}),
)
def _lexical_matches(
self,
text: str,
*,
include_snippets: bool,
max_snippets: int,
snippet_radius: int,
) -> tuple[dict[str, dict[str, Any]], dict[str, tuple[RetrievalSnippet, ...]]]:
if not self._lexical_index:
self.refresh_index()
needle = text.casefold()
matches: dict[str, dict[str, Any]] = {}
snippets: dict[str, list[RetrievalSnippet]] = {}
for document in self._lexical_index:
haystack = document.text.casefold()
count = haystack.count(needle)
if count <= 0:
continue
current = matches.setdefault(
document.asset_id,
{
"strategy": "lexical_substring",
"query": text,
"match_count": 0,
"representation_ids": [],
},
)
current["match_count"] += count
current["representation_ids"].append(document.representation_id)
if include_snippets:
asset_snippets = snippets.setdefault(document.asset_id, [])
remaining = max_snippets - len(asset_snippets)
if remaining > 0:
asset_snippets.extend(
_snippets_for_document(
document,
text,
max_snippets=remaining,
snippet_radius=snippet_radius,
)
)
current["snippet_count"] = len(asset_snippets)
return matches, {asset_id: tuple(items) for asset_id, items in snippets.items()}
def _context_entity_item(self, entity: ContextEntity, context: OperationContext) -> ContextEntityQueryItem:
relationships = [
relationship
for relationship in self.repository.list_relationships(target_id=entity.entity_id)
if relationship.target_kind == RelationshipTargetKind.CONTEXT_ENTITY
and self._relationship_allowed(relationship, context)
]
return ContextEntityQueryItem(
entity=entity,
asset_ids=tuple(sorted({relationship.source_id for relationship in relationships})),
relationship_count=len(relationships),
)
def _relationship_item(self, relationship: CoreRelationship) -> RelationshipQueryItem:
source_asset = self.repository.get_asset(relationship.source_id)
target_asset = None
target_entity = None
if relationship.target_kind == RelationshipTargetKind.ASSET:
target_asset = self.repository.get_asset(relationship.target_id)
else:
target_entity = self.repository.get_context_entity(relationship.target_id)
return RelationshipQueryItem(
relationship=relationship,
source_asset=source_asset,
target_asset=target_asset,
target_entity=target_entity,
)
def _relationship_context_for_asset(
self,
asset_id: str,
request: AssetQueryRequest,
context: OperationContext,
) -> tuple[tuple[CoreRelationship, ...], tuple[ContextEntity, ...]]:
relationships = self._relationships_for_asset(asset_id, request.relationship_direction)
entity_ids = self._context_entity_ids_for_asset_query(request)
filtered: list[CoreRelationship] = []
for relationship in relationships:
if request.relationship_predicate is not None and relationship.predicate != request.relationship_predicate:
continue
if request.related_asset_id is not None and not _relationship_connects_asset(
relationship,
asset_id=asset_id,
related_asset_id=request.related_asset_id,
):
continue
if entity_ids is not None and not (
relationship.target_kind == RelationshipTargetKind.CONTEXT_ENTITY
and relationship.target_id in entity_ids
):
continue
if not self._relationship_allowed(relationship, context):
continue
filtered.append(relationship)
ordered = tuple(_sort_relationships(_unique_relationships(filtered), "source_id", "asc"))
return ordered, self._context_entities_for_relationships(ordered)
def _relationships_for_request(self, request: RelationshipQueryRequest) -> list[CoreRelationship]:
relationships = self._relationships_for_relationship_query_base(request)
entity_ids = self._context_entity_ids_for_relationship_query(request)
filtered: list[CoreRelationship] = []
for relationship in relationships:
if request.target_kind is not None and relationship.target_kind != request.target_kind:
continue
if request.predicate is not None and relationship.predicate != request.predicate:
continue
if entity_ids is not None and not (
relationship.target_kind == RelationshipTargetKind.CONTEXT_ENTITY
and relationship.target_id in entity_ids
):
continue
filtered.append(relationship)
return _unique_relationships(filtered)
def _relationships_for_relationship_query_base(
self,
request: RelationshipQueryRequest,
) -> list[CoreRelationship]:
if request.asset_id is not None:
return list(self._relationships_for_asset(request.asset_id, request.direction))
if request.source_id is not None and request.target_id is not None:
return [
relationship
for relationship in self.repository.list_relationships(source_id=request.source_id)
if relationship.target_id == request.target_id
]
if request.source_id is not None:
return list(self.repository.list_relationships(source_id=request.source_id))
if request.target_id is not None:
return list(self.repository.list_relationships(target_id=request.target_id))
return _unique_relationships(
[
relationship
for asset in self.repository.list_assets()
for relationship in self.repository.list_relationships(source_id=asset.id)
]
)
def _relationships_for_asset(self, asset_id: str, direction: str) -> tuple[CoreRelationship, ...]:
relationships: list[CoreRelationship] = []
if direction in {"outbound", "both"}:
relationships.extend(self.repository.list_relationships(source_id=asset_id))
if direction in {"inbound", "both"}:
relationships.extend(
relationship
for relationship in self.repository.list_relationships(target_id=asset_id)
if relationship.target_kind == RelationshipTargetKind.ASSET
)
return tuple(_unique_relationships(relationships))
def _context_entity_ids_for_asset_query(self, request: AssetQueryRequest) -> set[str] | None:
if not _asset_query_has_context_entity_filter(request):
return None
return self._resolve_context_entity_ids(
entity_id=request.context_entity_id,
entity_type=request.context_entity_type,
name=request.context_entity_name,
external_ref=request.context_entity_external_ref,
workflow_run_id=request.workflow_run_id,
)
def _context_entity_ids_for_relationship_query(self, request: RelationshipQueryRequest) -> set[str] | None:
if not _relationship_query_has_context_entity_filter(request):
return None
return self._resolve_context_entity_ids(
entity_id=request.context_entity_id,
entity_type=request.context_entity_type,
name=request.context_entity_name,
external_ref=request.context_entity_external_ref,
workflow_run_id=request.workflow_run_id,
)
def _resolve_context_entity_ids(
self,
*,
entity_id: str | None,
entity_type: ContextEntityType | None,
name: str | None,
external_ref: str | None,
workflow_run_id: str | None,
) -> set[str]:
entities = self.repository.list_context_entities()
if workflow_run_id is not None:
entity_type = ContextEntityType.WORKFLOW_RUN
return {
entity.entity_id
for entity in entities
if _context_entity_matches(
entity,
entity_id=entity_id,
entity_type=entity_type,
name=name,
external_ref=external_ref,
metadata_filters={},
)
and _workflow_run_matches(entity, workflow_run_id)
}
def _context_entities_for_relationships(
self,
relationships: tuple[CoreRelationship, ...],
) -> tuple[ContextEntity, ...]:
entities: list[ContextEntity] = []
for relationship in relationships:
if relationship.target_kind != RelationshipTargetKind.CONTEXT_ENTITY:
continue
entities.append(self.repository.get_context_entity(relationship.target_id))
return tuple(_unique_context_entities(entities))
def _asset_allowed(self, asset: KnowledgeAsset, context: OperationContext) -> bool:
decision = self._authorize_for_retrieval(
context,
"asset.retrieve",
f"asset:{asset.id}",
resource_metadata={
"asset_id": asset.id,
"asset_type": asset.classification.asset_type,
"lifecycle": asset.lifecycle.value,
"sensitivity": asset.classification.sensitivity.value,
"owner": asset.classification.owner,
},
)
return decision.allowed
def _context_entity_allowed(self, entity: ContextEntity, context: OperationContext) -> bool:
decision = self._authorize_for_retrieval(
context,
"context_entity.retrieve",
f"context_entity:{entity.entity_id}",
resource_metadata={
"entity_id": entity.entity_id,
"entity_type": entity.entity_type.value,
"external_ref": entity.external_ref,
},
)
return decision.allowed
def _relationship_allowed(self, relationship: CoreRelationship, context: OperationContext) -> bool:
decision = self._authorize_for_retrieval(
context,
"asset.relationship.retrieve",
f"relationship:{relationship.relationship_id}",
resource_metadata={
"relationship_id": relationship.relationship_id,
"source_id": relationship.source_id,
"target_id": relationship.target_id,
"target_kind": relationship.target_kind.value,
"predicate": relationship.predicate,
},
)
if not decision.allowed:
return False
source_asset = self.repository.get_asset(relationship.source_id)
if not self._asset_allowed(source_asset, context):
return False
if relationship.target_kind == RelationshipTargetKind.ASSET:
target_asset = self.repository.get_asset(relationship.target_id)
return self._asset_allowed(target_asset, context)
target_entity = self.repository.get_context_entity(relationship.target_id)
return self._context_entity_allowed(target_entity, context)
def _authorize_for_retrieval(
self,
context: OperationContext,
action: str,
resource: str,
*,
resource_metadata: dict[str, Any] | None = None,
) -> PolicyDecision:
self.repository.save_actor(context.actor)
try:
return self.policy_gateway.authorize(
context,
action,
resource,
resource_metadata=resource_metadata,
)
except Exception as exc:
return PolicyDecision.fail_closed(
context.actor.id,
action,
resource,
reason=str(exc) or "Retrieval policy gateway failed",
context={
"gateway_error": type(exc).__name__,
"resource_metadata": resource_metadata or {},
},
)
def _audit_retrieval(
self,
operation: str,
target: str,
outcome: AuditOutcome,
context: OperationContext,
policy_decision: PolicyDecision,
*,
details: dict[str, Any] | None = None,
) -> AuditEvent:
event = AuditEvent.from_context(
operation,
target,
outcome,
context,
policy_decision=policy_decision,
details=details,
)
return self.repository.save_audit_event(event)
def _validate_request(request: AssetQueryRequest) -> tuple[list[Diagnostic], AssetQueryRequest]:
diagnostics: list[Diagnostic] = []
lifecycle = _parse_lifecycle(request.lifecycle, diagnostics)
sensitivity = _parse_sensitivity(request.sensitivity, diagnostics)
representation_kind = _parse_representation_kind(request.representation_kind, diagnostics)
context_entity_type = _parse_context_entity_type(request.context_entity_type, diagnostics)
relationship_direction = _parse_relationship_direction(request.relationship_direction, diagnostics)
if request.limit < 1 or request.limit > 500:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.limit_invalid",
message="Query limit must be between 1 and 500",
details={"limit": request.limit},
)
)
if request.offset < 0:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.offset_invalid",
message="Query offset must be zero or greater",
details={"offset": request.offset},
)
)
if request.max_snippets < 0 or request.max_snippets > 20:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.max_snippets_invalid",
message="Max snippets must be between 0 and 20",
details={"max_snippets": request.max_snippets},
)
)
if request.snippet_radius < 0 or request.snippet_radius > 1000:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.snippet_radius_invalid",
message="Snippet radius must be between 0 and 1000",
details={"snippet_radius": request.snippet_radius},
)
)
if request.sort_by not in SUPPORTED_SORT_KEYS:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.sort_invalid",
message="Query sort key is not supported",
details={"sort_by": request.sort_by, "supported": sorted(SUPPORTED_SORT_KEYS)},
)
)
if request.sort_order not in {"asc", "desc"}:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.sort_order_invalid",
message="Query sort order must be asc or desc",
details={"sort_order": request.sort_order},
)
)
return diagnostics, AssetQueryRequest(
text=request.text,
asset_type=request.asset_type,
lifecycle=lifecycle,
sensitivity=sensitivity,
owner=request.owner,
topic=request.topic,
tags=tuple(request.tags),
collection=request.collection,
review_state=request.review_state,
metadata_filters=dict(request.metadata_filters),
confirmed_metadata_only=request.confirmed_metadata_only,
source_system=request.source_system,
source_path=request.source_path,
context_entity_id=request.context_entity_id,
context_entity_type=context_entity_type,
context_entity_name=request.context_entity_name,
context_entity_external_ref=request.context_entity_external_ref,
workflow_run_id=request.workflow_run_id,
related_asset_id=request.related_asset_id,
relationship_predicate=request.relationship_predicate,
relationship_direction=relationship_direction,
include_relationships=request.include_relationships,
include_snippets=request.include_snippets,
max_snippets=request.max_snippets,
snippet_radius=request.snippet_radius,
created_after=request.created_after,
created_before=request.created_before,
updated_after=request.updated_after,
updated_before=request.updated_before,
representation_kind=representation_kind,
sort_by=request.sort_by,
sort_order=request.sort_order,
limit=request.limit,
offset=request.offset,
)
def _permission_denied_diagnostic(decision: PolicyDecision) -> Diagnostic:
return Diagnostic(
severity="error",
code="retrieval.permission_denied",
message="Retrieval query denied by policy",
details={"policy_decision": decision.to_dict()},
)
def _parse_feedback_label(
value: RetrievalFeedbackLabel | str,
diagnostics: list[Diagnostic],
) -> RetrievalFeedbackLabel | None:
if isinstance(value, RetrievalFeedbackLabel):
return value
try:
return RetrievalFeedbackLabel(value)
except ValueError:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.feedback_label_invalid",
message="Retrieval feedback label is not supported",
details={"label": value, "supported": [item.value for item in RetrievalFeedbackLabel]},
)
)
return None
def _feedback_rank(record: RetrievalFeedbackRecord) -> int | None:
rank = record.result_ref.get("rank")
if isinstance(rank, int):
return rank
if isinstance(rank, str) and rank.isdigit():
return int(rank)
return None
def _feedback_has_citation_ref(record: RetrievalFeedbackRecord) -> bool:
return any(
record.result_ref.get(key)
for key in (
"snippet_id",
"representation_id",
"source_ref_id",
)
) or bool(record.metadata.get("citation"))
def _elapsed_ms(started_at: float) -> float:
return round((perf_counter() - started_at) * 1000, 3)
def _parse_lifecycle(value: LifecycleState | str | None, diagnostics: list[Diagnostic]) -> LifecycleState | None:
if value is None or isinstance(value, LifecycleState):
return value
try:
return LifecycleState(value)
except ValueError:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.lifecycle_invalid",
message="Lifecycle filter is not supported",
details={"lifecycle": value, "supported": [item.value for item in LifecycleState]},
)
)
return None
def _parse_sensitivity(value: Sensitivity | str | None, diagnostics: list[Diagnostic]) -> Sensitivity | None:
if value is None or isinstance(value, Sensitivity):
return value
try:
return Sensitivity(value)
except ValueError:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.sensitivity_invalid",
message="Sensitivity filter is not supported",
details={"sensitivity": value, "supported": [item.value for item in Sensitivity]},
)
)
return None
def _parse_representation_kind(
value: RepresentationKind | str | None,
diagnostics: list[Diagnostic],
) -> RepresentationKind | None:
if value is None or isinstance(value, RepresentationKind):
return value
try:
return RepresentationKind(value)
except ValueError:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.representation_kind_invalid",
message="Representation kind filter is not supported",
details={"representation_kind": value, "supported": [item.value for item in RepresentationKind]},
)
)
return None
def _parse_context_entity_type(
value: ContextEntityType | str | None,
diagnostics: list[Diagnostic],
) -> ContextEntityType | None:
if value is None or isinstance(value, ContextEntityType):
return value
try:
return ContextEntityType(value)
except ValueError:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.context_entity_type_invalid",
message="Context entity type filter is not supported",
details={"context_entity_type": value, "supported": [item.value for item in ContextEntityType]},
)
)
return None
def _parse_relationship_target_kind(
value: RelationshipTargetKind | str | None,
diagnostics: list[Diagnostic],
) -> RelationshipTargetKind | None:
if value is None or isinstance(value, RelationshipTargetKind):
return value
try:
return RelationshipTargetKind(value)
except ValueError:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.relationship_target_kind_invalid",
message="Relationship target kind is not supported",
details={"target_kind": value, "supported": [item.value for item in RelationshipTargetKind]},
)
)
return None
def _parse_relationship_direction(value: str, diagnostics: list[Diagnostic]) -> str:
if value in {"outbound", "inbound", "both"}:
return value
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.relationship_direction_invalid",
message="Relationship direction must be outbound, inbound, or both",
details={"direction": value},
)
)
return "both"
def _validate_context_entity_request(
request: ContextEntityQueryRequest,
) -> tuple[list[Diagnostic], ContextEntityQueryRequest]:
diagnostics: list[Diagnostic] = []
entity_type = _parse_context_entity_type(request.entity_type, diagnostics)
_validate_window_and_sort(
diagnostics,
limit=request.limit,
offset=request.offset,
sort_by=request.sort_by,
sort_order=request.sort_order,
supported_sort_keys=SUPPORTED_CONTEXT_ENTITY_SORT_KEYS,
)
return diagnostics, ContextEntityQueryRequest(
entity_id=request.entity_id,
entity_type=entity_type,
name=request.name,
external_ref=request.external_ref,
metadata_filters=dict(request.metadata_filters),
sort_by=request.sort_by,
sort_order=request.sort_order,
limit=request.limit,
offset=request.offset,
)
def _validate_relationship_request(
request: RelationshipQueryRequest,
) -> tuple[list[Diagnostic], RelationshipQueryRequest]:
diagnostics: list[Diagnostic] = []
context_entity_type = _parse_context_entity_type(request.context_entity_type, diagnostics)
target_kind = _parse_relationship_target_kind(request.target_kind, diagnostics)
direction = _parse_relationship_direction(request.direction, diagnostics)
_validate_window_and_sort(
diagnostics,
limit=request.limit,
offset=request.offset,
sort_by=request.sort_by,
sort_order=request.sort_order,
supported_sort_keys=SUPPORTED_RELATIONSHIP_SORT_KEYS,
)
return diagnostics, RelationshipQueryRequest(
source_id=request.source_id,
target_id=request.target_id,
asset_id=request.asset_id,
context_entity_id=request.context_entity_id,
context_entity_type=context_entity_type,
context_entity_name=request.context_entity_name,
context_entity_external_ref=request.context_entity_external_ref,
workflow_run_id=request.workflow_run_id,
target_kind=target_kind,
predicate=request.predicate,
direction=direction,
sort_by=request.sort_by,
sort_order=request.sort_order,
limit=request.limit,
offset=request.offset,
)
def _validate_window_and_sort(
diagnostics: list[Diagnostic],
*,
limit: int,
offset: int,
sort_by: str,
sort_order: str,
supported_sort_keys: set[str],
) -> None:
if limit < 1 or limit > 500:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.limit_invalid",
message="Query limit must be between 1 and 500",
details={"limit": limit},
)
)
if offset < 0:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.offset_invalid",
message="Query offset must be zero or greater",
details={"offset": offset},
)
)
if sort_by not in supported_sort_keys:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.sort_invalid",
message="Query sort key is not supported",
details={"sort_by": sort_by, "supported": sorted(supported_sort_keys)},
)
)
if sort_order not in {"asc", "desc"}:
diagnostics.append(
Diagnostic(
severity="error",
code="retrieval.sort_order_invalid",
message="Query sort order must be asc or desc",
details={"sort_order": sort_order},
)
)
def _source_matches(
asset: KnowledgeAsset,
*,
source_system: str | None,
source_path: str | None,
) -> bool:
if source_system is None and source_path is None:
return True
for source_ref in asset.source_refs:
if source_system is not None and source_ref.source_system != source_system:
continue
if source_path is not None and source_ref.path != source_path:
continue
return True
return False
def _collection_matches(
asset: KnowledgeAsset,
metadata_records: list[MetadataRecord],
collection: str | None,
) -> bool:
if collection is None:
return True
if asset.metadata.get("collection") == collection:
return True
if asset.classification.metadata.get("collection") == collection:
return True
return any(record.key == "collection" and record.value == collection for record in metadata_records)
def _tags_match(
asset: KnowledgeAsset,
metadata_records: list[MetadataRecord],
tags: tuple[str, ...],
) -> bool:
if not tags:
return True
values = set(asset.classification.topics)
for record in metadata_records:
if record.key not in {"tag", "tags"}:
continue
if isinstance(record.value, list):
values.update(str(item) for item in record.value)
elif record.value is not None:
values.add(str(record.value))
return all(tag in values for tag in tags)
def _timestamp_matches(asset: KnowledgeAsset, request: AssetQueryRequest) -> bool:
if request.created_after and asset.created_at < request.created_after:
return False
if request.created_before and asset.created_at > request.created_before:
return False
if request.updated_after and asset.updated_at < request.updated_after:
return False
if request.updated_before and asset.updated_at > request.updated_before:
return False
return True
def _asset_query_has_graph_filter(request: AssetQueryRequest) -> bool:
return (
_asset_query_has_context_entity_filter(request)
or request.related_asset_id is not None
or request.relationship_predicate is not None
)
def _asset_query_has_context_entity_filter(request: AssetQueryRequest) -> bool:
return any(
value is not None
for value in (
request.context_entity_id,
request.context_entity_type,
request.context_entity_name,
request.context_entity_external_ref,
request.workflow_run_id,
)
)
def _relationship_query_has_context_entity_filter(request: RelationshipQueryRequest) -> bool:
return any(
value is not None
for value in (
request.context_entity_id,
request.context_entity_type,
request.context_entity_name,
request.context_entity_external_ref,
request.workflow_run_id,
)
)
def _context_entity_matches(
entity: ContextEntity,
*,
entity_id: str | None,
entity_type: ContextEntityType | None,
name: str | None,
external_ref: str | None,
metadata_filters: dict[str, Any],
) -> bool:
if entity_id is not None and entity.entity_id != entity_id:
return False
if entity_type is not None and entity.entity_type != entity_type:
return False
if name is not None and entity.name != name:
return False
if external_ref is not None and entity.external_ref != external_ref:
return False
for key, expected in metadata_filters.items():
if not _metadata_value_matches(entity.metadata.get(key), expected):
return False
return True
def _workflow_run_matches(entity: ContextEntity, workflow_run_id: str | None) -> bool:
if workflow_run_id is None:
return True
if entity.entity_id == workflow_run_id or entity.external_ref == workflow_run_id:
return True
return entity.metadata.get("workflow_run_id") == workflow_run_id
def _snippet_provenance(representation: AssetRepresentation) -> dict[str, Any]:
provenance: dict[str, Any] = {}
for key in (
"adapter_provenance",
"context_span",
"extractor",
"markitect_selector",
"snapshot",
"source_span",
):
value = representation.metadata.get(key)
if value:
provenance[key] = value
return provenance
def _snippets_for_document(
document: _LexicalDocument,
query_text: str,
*,
max_snippets: int,
snippet_radius: int,
) -> list[RetrievalSnippet]:
if max_snippets <= 0:
return []
needle = query_text.casefold()
haystack = document.text.casefold()
snippets: list[RetrievalSnippet] = []
cursor = 0
while len(snippets) < max_snippets:
match_start = haystack.find(needle, cursor)
if match_start < 0:
break
match_end = match_start + len(query_text)
snippet_start = max(0, match_start - snippet_radius)
snippet_end = min(len(document.text), match_end + snippet_radius)
snippet_text = document.text[snippet_start:snippet_end].strip()
if snippet_start > 0:
snippet_text = "..." + snippet_text
if snippet_end < len(document.text):
snippet_text = snippet_text + "..."
snippets.append(
RetrievalSnippet(
asset_id=document.asset_id,
representation_id=document.representation_id,
source_ref_id=document.source_ref_id,
storage_ref=document.storage_ref,
media_type=document.media_type,
text=snippet_text,
start_offset=snippet_start,
end_offset=snippet_end,
match_text=document.text[match_start:match_end],
provenance=dict(document.provenance),
)
)
cursor = match_end
return snippets
def _relationship_connects_asset(
relationship: CoreRelationship,
*,
asset_id: str,
related_asset_id: str,
) -> bool:
if relationship.target_kind != RelationshipTargetKind.ASSET:
return False
return (
relationship.source_id == asset_id
and relationship.target_id == related_asset_id
or relationship.source_id == related_asset_id
and relationship.target_id == asset_id
)
def _metadata_value_matches(value: Any, expected: Any) -> bool:
if isinstance(value, list):
return expected in value
return value == expected
def _unique_relationships(relationships: list[CoreRelationship]) -> list[CoreRelationship]:
seen: set[str] = set()
unique: list[CoreRelationship] = []
for relationship in relationships:
if relationship.relationship_id in seen:
continue
seen.add(relationship.relationship_id)
unique.append(relationship)
return unique
def _unique_context_entities(entities: list[ContextEntity]) -> list[ContextEntity]:
seen: set[str] = set()
unique: list[ContextEntity] = []
for entity in entities:
if entity.entity_id in seen:
continue
seen.add(entity.entity_id)
unique.append(entity)
return sorted(unique, key=lambda item: (item.entity_type.value, item.name, item.entity_id))
def _sort_assets(
assets: list[KnowledgeAsset],
sort_by: str,
sort_order: str,
) -> list[KnowledgeAsset]:
reverse = sort_order == "desc"
def sort_key(asset: KnowledgeAsset) -> tuple[Any, str]:
if sort_by == "asset_id":
primary = asset.id
elif sort_by == "asset_type":
primary = asset.classification.asset_type
elif sort_by == "created_at":
primary = asset.created_at
elif sort_by == "lifecycle":
primary = asset.lifecycle.value
elif sort_by == "updated_at":
primary = asset.updated_at
else:
primary = asset.title.casefold()
return primary, asset.id
return sorted(assets, key=sort_key, reverse=reverse)
def _sort_context_entities(
entities: list[ContextEntity],
sort_by: str,
sort_order: str,
) -> list[ContextEntity]:
reverse = sort_order == "desc"
def sort_key(entity: ContextEntity) -> tuple[Any, str]:
if sort_by == "entity_id":
primary = entity.entity_id
elif sort_by == "entity_type":
primary = entity.entity_type.value
elif sort_by == "external_ref":
primary = entity.external_ref or ""
else:
primary = entity.name.casefold()
return primary, entity.entity_id
return sorted(entities, key=sort_key, reverse=reverse)
def _sort_relationships(
relationships: list[CoreRelationship],
sort_by: str,
sort_order: str,
) -> list[CoreRelationship]:
reverse = sort_order == "desc"
def sort_key(relationship: CoreRelationship) -> tuple[Any, str]:
if sort_by == "created_at":
primary = relationship.created_at
elif sort_by == "predicate":
primary = relationship.predicate
elif sort_by == "relationship_id":
primary = relationship.relationship_id
elif sort_by == "target_id":
primary = relationship.target_id
elif sort_by == "target_kind":
primary = relationship.target_kind.value
else:
primary = relationship.source_id
return primary, relationship.relationship_id
return sorted(relationships, key=sort_key, reverse=reverse)