Add reviewer, artifact store, and paginated queries

This commit is contained in:
2026-07-03 21:44:46 +03:00
parent 50178a4e69
commit 648f0f9024
17 changed files with 441 additions and 16 deletions
+6 -6
View File
@@ -20,7 +20,7 @@
- [x] graph scheduler service - [x] graph scheduler service
- [x] node runner abstraction - [x] node runner abstraction
- [ ] reviewer contract - [x] reviewer contract
- [x] finalizer contract - [x] finalizer contract
- [x] resume/cancel use cases - [x] resume/cancel use cases
- [x] idempotency policy - [x] idempotency policy
@@ -31,8 +31,8 @@
- [x] SQLAlchemy models - [x] SQLAlchemy models
- [x] repositories for PostgreSQL/SQLite - [x] repositories for PostgreSQL/SQLite
- [x] migrations - [x] migrations
- [ ] artifact metadata persistence - [x] artifact metadata persistence
- [ ] event store queries with pagination - [x] event store queries with pagination
## Package D: Integrations ## Package D: Integrations
@@ -53,7 +53,7 @@
## Package F: Validation ## Package F: Validation
- [ ] contract tests - [x] contract tests
- [ ] scenario tests - [x] scenario tests
- [ ] smoke suite - [x] smoke suite
- [ ] static checks in CI - [ ] static checks in CI
+11
View File
@@ -169,6 +169,17 @@ Migration skeleton:
- `confirmations(status, created_at)` - `confirmations(status, created_at)`
- `worker_sessions(status, last_heartbeat_at)` - `worker_sessions(status, last_heartbeat_at)`
## Query Support
Read-model adapters now support paginated queries for:
- `task_events`
- `model_invocations`
- `tool_invocations`
- `artifacts`
This is implemented in both in-memory and SQLAlchemy adapters to keep test/runtime behavior aligned.
## Active Adapter Strategy ## Active Adapter Strategy
- `memory` backend остается для быстрых unit/integration прогонов; - `memory` backend остается для быстрых unit/integration прогонов;
@@ -153,8 +153,24 @@ def upgrade() -> None:
op.create_index("ix_tool_invocations_task_id", "tool_invocations", ["task_id"]) op.create_index("ix_tool_invocations_task_id", "tool_invocations", ["task_id"])
op.create_index("ix_tool_invocations_created_at", "tool_invocations", ["created_at"]) op.create_index("ix_tool_invocations_created_at", "tool_invocations", ["created_at"])
op.create_table(
"artifacts",
sa.Column("artifact_id", sa.String(length=64), primary_key=True),
sa.Column("task_id", sa.String(length=64), nullable=False),
sa.Column("node_id", sa.String(length=64), nullable=True),
sa.Column("artifact_type", sa.String(length=64), nullable=False),
sa.Column("storage_uri", sa.String(length=512), nullable=False),
sa.Column("metadata_json", sa.JSON(), nullable=False),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
)
op.create_index("ix_artifacts_task_id", "artifacts", ["task_id"])
op.create_index("ix_artifacts_created_at", "artifacts", ["created_at"])
def downgrade() -> None: def downgrade() -> None:
op.drop_index("ix_artifacts_created_at", table_name="artifacts")
op.drop_index("ix_artifacts_task_id", table_name="artifacts")
op.drop_table("artifacts")
op.drop_index("ix_tool_invocations_created_at", table_name="tool_invocations") op.drop_index("ix_tool_invocations_created_at", table_name="tool_invocations")
op.drop_index("ix_tool_invocations_task_id", table_name="tool_invocations") op.drop_index("ix_tool_invocations_task_id", table_name="tool_invocations")
op.drop_table("tool_invocations") op.drop_table("tool_invocations")
+31
View File
@@ -63,6 +63,24 @@ class ToolInvocationRecord:
response: dict[str, Any] response: dict[str, Any]
@dataclass(slots=True)
class ArtifactRecord:
artifact_id: str
task_id: str
node_id: str | None
artifact_type: str
storage_uri: str
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class Page:
items: list[Any]
limit: int
offset: int
total: int
class TaskRepository(Protocol): class TaskRepository(Protocol):
def create(self, task: Task) -> Task: ... def create(self, task: Task) -> Task: ...
def save(self, task: Task) -> Task: ... def save(self, task: Task) -> Task: ...
@@ -90,6 +108,7 @@ class WorkerRepository(Protocol):
class EventStore(Protocol): class EventStore(Protocol):
def append(self, event: DomainEvent) -> DomainEvent: ... def append(self, event: DomainEvent) -> DomainEvent: ...
def list_by_task(self, task_id: str) -> list[DomainEvent]: ... def list_by_task(self, task_id: str) -> list[DomainEvent]: ...
def list_by_task_paginated(self, task_id: str, *, limit: int, offset: int) -> Page: ...
class InvocationStore(Protocol): class InvocationStore(Protocol):
@@ -97,6 +116,18 @@ class InvocationStore(Protocol):
def save_tool_invocation(self, record: ToolInvocationRecord) -> ToolInvocationRecord: ... def save_tool_invocation(self, record: ToolInvocationRecord) -> ToolInvocationRecord: ...
def list_model_invocations(self, task_id: str) -> list[ModelInvocationRecord]: ... def list_model_invocations(self, task_id: str) -> list[ModelInvocationRecord]: ...
def list_tool_invocations(self, task_id: str) -> list[ToolInvocationRecord]: ... def list_tool_invocations(self, task_id: str) -> list[ToolInvocationRecord]: ...
def list_model_invocations_paginated(
self, task_id: str, *, limit: int, offset: int
) -> Page: ...
def list_tool_invocations_paginated(
self, task_id: str, *, limit: int, offset: int
) -> Page: ...
class ArtifactStore(Protocol):
def save_artifact(self, record: ArtifactRecord) -> ArtifactRecord: ...
def list_by_task(self, task_id: str) -> list[ArtifactRecord]: ...
def list_by_task_paginated(self, task_id: str, *, limit: int, offset: int) -> Page: ...
class PolicyEvaluator(Protocol): class PolicyEvaluator(Protocol):
@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any
from ai_orchestrator.application.ports import EventStore, Finalizer, NodeRunner, Reviewer from ai_orchestrator.application.ports import EventStore, Finalizer, NodeRunner, Reviewer
from ai_orchestrator.domain.enums import NodeStatus, NodeType from ai_orchestrator.domain.enums import NodeStatus, NodeType
@@ -39,6 +40,32 @@ class NoOpReviewer(Reviewer):
} }
@dataclass(slots=True)
class StructuredReviewer(Reviewer):
required_keys: tuple[str, ...] = ("status",)
def review(
self,
*,
task: Task,
node: ExecutionNode,
output_data: dict[str, Any],
) -> dict[str, Any]:
missing = [key for key in self.required_keys if key not in output_data]
reviewed = {
**output_data,
"task_id": task.task_id,
"node_id": node.node_id,
"review": {
"status": "accepted" if not missing else "rejected",
"missing_keys": missing,
},
}
if missing:
reviewed["status"] = "review_failed"
return reviewed
@dataclass(slots=True) @dataclass(slots=True)
class DefaultFinalizer(Finalizer): class DefaultFinalizer(Finalizer):
def finalize(self, *, task: Task, graph: ExecutionGraph) -> dict[str, object]: def finalize(self, *, task: Task, graph: ExecutionGraph) -> dict[str, object]:
@@ -5,6 +5,8 @@ from datetime import datetime
from uuid import uuid4 from uuid import uuid4
from ai_orchestrator.application.ports import ( from ai_orchestrator.application.ports import (
ArtifactRecord,
ArtifactStore,
EventStore, EventStore,
InvocationStore, InvocationStore,
ToolInvocationRecord, ToolInvocationRecord,
@@ -115,6 +117,7 @@ class CapabilityAwareWorkerGateway(WorkerGateway):
worker_repository: WorkerRepository worker_repository: WorkerRepository
event_store: EventStore event_store: EventStore
invocation_store: InvocationStore invocation_store: InvocationStore
artifact_store: ArtifactStore
connection_manager: InMemoryWorkerConnectionManager connection_manager: InMemoryWorkerConnectionManager
def dispatch( def dispatch(
@@ -244,6 +247,21 @@ class CapabilityAwareWorkerGateway(WorkerGateway):
}, },
) )
) )
for index, artifact in enumerate(result.artifacts or [], start=1):
self.artifact_store.save_artifact(
ArtifactRecord(
artifact_id=f"art_{uuid4().hex}",
task_id=result.task_id,
node_id=None,
artifact_type=str(artifact.get("type", "artifact")),
storage_uri=str(
artifact.get("path")
or artifact.get("uri")
or f"worker://{worker_session_id}/{result.command_id}/{index}"
),
metadata=dict(artifact),
)
)
return ToolInvocationResult( return ToolInvocationResult(
status=result.status, status=result.status,
content=result.result or {}, content=result.result or {},
+21 -3
View File
@@ -9,9 +9,9 @@ from ai_orchestrator.application.services.execution import (
DefaultFinalizer, DefaultFinalizer,
FinalizerNodeRunner, FinalizerNodeRunner,
GraphExecutionEngine, GraphExecutionEngine,
NoOpReviewer,
PlannerNodeRunner, PlannerNodeRunner,
RunnerRegistry, RunnerRegistry,
StructuredReviewer,
) )
from ai_orchestrator.application.services.lifecycle import TaskLifecycleService from ai_orchestrator.application.services.lifecycle import TaskLifecycleService
from ai_orchestrator.application.services.orchestrator import CreateTaskRequest, OrchestratorService from ai_orchestrator.application.services.orchestrator import CreateTaskRequest, OrchestratorService
@@ -69,6 +69,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
worker_repository = storage.worker_repository worker_repository = storage.worker_repository
event_store = storage.event_store event_store = storage.event_store
invocation_store = storage.invocation_store invocation_store = storage.invocation_store
artifact_store = storage.artifact_store
policy_evaluator = StaticProjectPolicyEvaluator( policy_evaluator = StaticProjectPolicyEvaluator(
projects={project_config.project.id: project_config} projects={project_config.project.id: project_config}
) )
@@ -77,6 +78,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
worker_repository=worker_repository, worker_repository=worker_repository,
event_store=event_store, event_store=event_store,
invocation_store=invocation_store, invocation_store=invocation_store,
artifact_store=artifact_store,
connection_manager=worker_connection_manager, connection_manager=worker_connection_manager,
) )
model_router = ConfigurableModelRouter( model_router = ConfigurableModelRouter(
@@ -93,7 +95,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
execution_engine = GraphExecutionEngine( execution_engine = GraphExecutionEngine(
event_store=event_store, event_store=event_store,
runner_registry=runner_registry, runner_registry=runner_registry,
reviewer=NoOpReviewer(), reviewer=StructuredReviewer(),
) )
orchestrator = OrchestratorService( orchestrator = OrchestratorService(
task_repository=task_repository, task_repository=task_repository,
@@ -117,6 +119,7 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
app.state.event_store = event_store app.state.event_store = event_store
app.state.storage = storage app.state.storage = storage
app.state.invocation_store = invocation_store app.state.invocation_store = invocation_store
app.state.artifact_store = artifact_store
app.state.model_router = model_router app.state.model_router = model_router
app.state.worker_gateway = worker_gateway app.state.worker_gateway = worker_gateway
app.state.worker_connection_manager = worker_connection_manager app.state.worker_connection_manager = worker_connection_manager
@@ -245,12 +248,27 @@ def create_app(settings: AppSettings | None = None) -> FastAPI:
for item in invocation_store.list_tool_invocations(task_id) for item in invocation_store.list_tool_invocations(task_id)
] ]
artifact_cards = build_artifact_cards_from_tool_invocations(task_id, tool_invocations) artifact_cards = build_artifact_cards_from_tool_invocations(task_id, tool_invocations)
persisted_artifact_cards = [
{
"type": "artifact_card",
"task_id": item.task_id,
"artifact_type": item.artifact_type,
"title": item.metadata.get("title", item.artifact_type),
"data": {"storage_uri": item.storage_uri, **item.metadata},
}
for item in artifact_store.list_by_task(task_id)
]
return TaskStatusResponse( return TaskStatusResponse(
task_id=task.task_id, task_id=task.task_id,
status=task.status.value, status=task.status.value,
current_node=task.current_node_id, current_node=task.current_node_id,
progress={"completed": completed, "total": len(graph.nodes)}, progress={"completed": completed, "total": len(graph.nodes)},
cards=[build_progress_card(task, graph), *confirmation_cards, *artifact_cards], cards=[
build_progress_card(task, graph),
*confirmation_cards,
*artifact_cards,
*persisted_artifact_cards,
],
) )
@app.get("/tasks/{task_id}/events") @app.get("/tasks/{task_id}/events")
@@ -7,6 +7,8 @@ from uuid import uuid4
import httpx import httpx
from ai_orchestrator.application.ports import ( from ai_orchestrator.application.ports import (
ArtifactRecord,
ArtifactStore,
InvocationStore, InvocationStore,
ToolGateway, ToolGateway,
ToolInvocationRecord, ToolInvocationRecord,
@@ -60,6 +62,7 @@ class McpToolGateway(ToolGateway):
project_configs: dict[str, ProjectConfig] project_configs: dict[str, ProjectConfig]
client: McpHttpClient client: McpHttpClient
invocation_store: InvocationStore invocation_store: InvocationStore
artifact_store: ArtifactStore
def call( def call(
self, self,
@@ -105,4 +108,19 @@ class McpToolGateway(ToolGateway):
}, },
) )
) )
for index, artifact in enumerate(result.artifacts, start=1):
self.artifact_store.save_artifact(
ArtifactRecord(
artifact_id=f"art_{uuid4().hex}",
task_id=str(task_context.get("task_id", project_id)),
node_id=str(task_context["node_id"]) if task_context.get("node_id") else None,
artifact_type=str(artifact.get("type", "artifact")),
storage_uri=str(
artifact.get("path")
or artifact.get("uri")
or f"mcp://{server_id}/{tool_name}/{index}"
),
metadata=dict(artifact),
)
)
return result return result
@@ -4,6 +4,7 @@ from dataclasses import dataclass
from ai_orchestrator.config import AppSettings from ai_orchestrator.config import AppSettings
from ai_orchestrator.infrastructure.storage.memory import ( from ai_orchestrator.infrastructure.storage.memory import (
InMemoryArtifactStore,
InMemoryConfirmationRepository, InMemoryConfirmationRepository,
InMemoryEventStore, InMemoryEventStore,
InMemoryGraphRepository, InMemoryGraphRepository,
@@ -22,6 +23,7 @@ class StorageBundle:
worker_repository: object worker_repository: object
event_store: object event_store: object
invocation_store: object invocation_store: object
artifact_store: object
engine: object | None = None engine: object | None = None
session_factory: object | None = None session_factory: object | None = None
@@ -35,6 +37,7 @@ def create_storage_bundle(settings: AppSettings) -> StorageBundle:
worker_repository=InMemoryWorkerRepository(), worker_repository=InMemoryWorkerRepository(),
event_store=InMemoryEventStore(), event_store=InMemoryEventStore(),
invocation_store=InMemoryInvocationStore(), invocation_store=InMemoryInvocationStore(),
artifact_store=InMemoryArtifactStore(),
) )
if settings.storage_backend in {"sqlite", "postgres"}: if settings.storage_backend in {"sqlite", "postgres"}:
bundle = create_sqlalchemy_storage(settings.database_url) bundle = create_sqlalchemy_storage(settings.database_url)
@@ -45,6 +48,7 @@ def create_storage_bundle(settings: AppSettings) -> StorageBundle:
worker_repository=bundle.worker_repository, worker_repository=bundle.worker_repository,
event_store=bundle.event_store, event_store=bundle.event_store,
invocation_store=bundle.invocation_store, invocation_store=bundle.invocation_store,
artifact_store=bundle.artifact_store,
engine=bundle.engine, engine=bundle.engine,
session_factory=bundle.session_factory, session_factory=bundle.session_factory,
) )
@@ -3,11 +3,14 @@ from __future__ import annotations
from dataclasses import dataclass, field from dataclasses import dataclass, field
from ai_orchestrator.application.ports import ( from ai_orchestrator.application.ports import (
ArtifactRecord,
ArtifactStore,
ConfirmationRepository, ConfirmationRepository,
EventStore, EventStore,
GraphRepository, GraphRepository,
InvocationStore, InvocationStore,
ModelInvocationRecord, ModelInvocationRecord,
Page,
TaskRepository, TaskRepository,
ToolInvocationRecord, ToolInvocationRecord,
WorkerRepository, WorkerRepository,
@@ -89,6 +92,15 @@ class InMemoryEventStore(EventStore):
def list_by_task(self, task_id: str) -> list[DomainEvent]: def list_by_task(self, task_id: str) -> list[DomainEvent]:
return [event for event in self.items if event.task_id == task_id] return [event for event in self.items if event.task_id == task_id]
def list_by_task_paginated(self, task_id: str, *, limit: int, offset: int) -> Page:
filtered = self.list_by_task(task_id)
return Page(
items=filtered[offset : offset + limit],
limit=limit,
offset=offset,
total=len(filtered),
)
@dataclass(slots=True) @dataclass(slots=True)
class InMemoryInvocationStore(InvocationStore): class InMemoryInvocationStore(InvocationStore):
@@ -108,3 +120,42 @@ class InMemoryInvocationStore(InvocationStore):
def list_tool_invocations(self, task_id: str) -> list[ToolInvocationRecord]: def list_tool_invocations(self, task_id: str) -> list[ToolInvocationRecord]:
return [record for record in self.tool_items if record.task_id == task_id] return [record for record in self.tool_items if record.task_id == task_id]
def list_model_invocations_paginated(self, task_id: str, *, limit: int, offset: int) -> Page:
filtered = self.list_model_invocations(task_id)
return Page(
items=filtered[offset : offset + limit],
limit=limit,
offset=offset,
total=len(filtered),
)
def list_tool_invocations_paginated(self, task_id: str, *, limit: int, offset: int) -> Page:
filtered = self.list_tool_invocations(task_id)
return Page(
items=filtered[offset : offset + limit],
limit=limit,
offset=offset,
total=len(filtered),
)
@dataclass(slots=True)
class InMemoryArtifactStore(ArtifactStore):
items: list[ArtifactRecord] = field(default_factory=list)
def save_artifact(self, record: ArtifactRecord) -> ArtifactRecord:
self.items.append(record)
return record
def list_by_task(self, task_id: str) -> list[ArtifactRecord]:
return [record for record in self.items if record.task_id == task_id]
def list_by_task_paginated(self, task_id: str, *, limit: int, offset: int) -> Page:
filtered = self.list_by_task(task_id)
return Page(
items=filtered[offset : offset + limit],
limit=limit,
offset=offset,
total=len(filtered),
)
@@ -8,11 +8,14 @@ from sqlalchemy import JSON, Boolean, DateTime, Integer, String, Text, create_en
from sqlalchemy.orm import DeclarativeBase, Mapped, Session, mapped_column, sessionmaker from sqlalchemy.orm import DeclarativeBase, Mapped, Session, mapped_column, sessionmaker
from ai_orchestrator.application.ports import ( from ai_orchestrator.application.ports import (
ArtifactRecord,
ArtifactStore,
ConfirmationRepository, ConfirmationRepository,
EventStore, EventStore,
GraphRepository, GraphRepository,
InvocationStore, InvocationStore,
ModelInvocationRecord, ModelInvocationRecord,
Page,
TaskRepository, TaskRepository,
ToolInvocationRecord, ToolInvocationRecord,
WorkerRepository, WorkerRepository,
@@ -163,6 +166,18 @@ class ToolInvocationDbRecord(Base):
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True)
class ArtifactDbRecord(Base):
__tablename__ = "artifacts"
artifact_id: Mapped[str] = mapped_column(String(64), primary_key=True)
task_id: Mapped[str] = mapped_column(String(64), index=True)
node_id: Mapped[str | None] = mapped_column(String(64), nullable=True)
artifact_type: Mapped[str] = mapped_column(String(64))
storage_uri: Mapped[str] = mapped_column(String(512))
metadata_json: Mapped[dict] = mapped_column(JSON)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True)
def build_engine(database_url: str): def build_engine(database_url: str):
if database_url.startswith("sqlite:///"): if database_url.startswith("sqlite:///"):
db_path = database_url.removeprefix("sqlite:///") db_path = database_url.removeprefix("sqlite:///")
@@ -416,6 +431,29 @@ def _record_to_tool_invocation(record: ToolInvocationDbRecord) -> ToolInvocation
) )
def _artifact_to_record(record: ArtifactRecord) -> ArtifactDbRecord:
return ArtifactDbRecord(
artifact_id=record.artifact_id,
task_id=record.task_id,
node_id=record.node_id,
artifact_type=record.artifact_type,
storage_uri=record.storage_uri,
metadata_json=record.metadata,
created_at=datetime.now(UTC),
)
def _record_to_artifact(record: ArtifactDbRecord) -> ArtifactRecord:
return ArtifactRecord(
artifact_id=record.artifact_id,
task_id=record.task_id,
node_id=record.node_id,
artifact_type=record.artifact_type,
storage_uri=record.storage_uri,
metadata=record.metadata_json or {},
)
@dataclass(slots=True) @dataclass(slots=True)
class SqlAlchemyTaskRepository(TaskRepository): class SqlAlchemyTaskRepository(TaskRepository):
session_factory: sessionmaker[Session] session_factory: sessionmaker[Session]
@@ -595,6 +633,15 @@ class SqlAlchemyEventStore(EventStore):
).all() ).all()
return [_record_to_event(record) for record in records] return [_record_to_event(record) for record in records]
def list_by_task_paginated(self, task_id: str, *, limit: int, offset: int) -> Page:
records = self.list_by_task(task_id)
return Page(
items=records[offset : offset + limit],
limit=limit,
offset=offset,
total=len(records),
)
@dataclass(slots=True) @dataclass(slots=True)
class SqlAlchemyInvocationStore(InvocationStore): class SqlAlchemyInvocationStore(InvocationStore):
@@ -630,6 +677,53 @@ class SqlAlchemyInvocationStore(InvocationStore):
).all() ).all()
return [_record_to_tool_invocation(record) for record in records] return [_record_to_tool_invocation(record) for record in records]
def list_model_invocations_paginated(self, task_id: str, *, limit: int, offset: int) -> Page:
records = self.list_model_invocations(task_id)
return Page(
items=records[offset : offset + limit],
limit=limit,
offset=offset,
total=len(records),
)
def list_tool_invocations_paginated(self, task_id: str, *, limit: int, offset: int) -> Page:
records = self.list_tool_invocations(task_id)
return Page(
items=records[offset : offset + limit],
limit=limit,
offset=offset,
total=len(records),
)
@dataclass(slots=True)
class SqlAlchemyArtifactStore(ArtifactStore):
session_factory: sessionmaker[Session]
def save_artifact(self, record: ArtifactRecord) -> ArtifactRecord:
with self.session_factory() as session:
session.add(_artifact_to_record(record))
session.commit()
return record
def list_by_task(self, task_id: str) -> list[ArtifactRecord]:
with self.session_factory() as session:
records = session.scalars(
select(ArtifactDbRecord)
.where(ArtifactDbRecord.task_id == task_id)
.order_by(ArtifactDbRecord.created_at.asc())
).all()
return [_record_to_artifact(record) for record in records]
def list_by_task_paginated(self, task_id: str, *, limit: int, offset: int) -> Page:
records = self.list_by_task(task_id)
return Page(
items=records[offset : offset + limit],
limit=limit,
offset=offset,
total=len(records),
)
@dataclass(slots=True) @dataclass(slots=True)
class SqlAlchemyStorageBundle: class SqlAlchemyStorageBundle:
@@ -641,6 +735,7 @@ class SqlAlchemyStorageBundle:
worker_repository: SqlAlchemyWorkerRepository worker_repository: SqlAlchemyWorkerRepository
event_store: SqlAlchemyEventStore event_store: SqlAlchemyEventStore
invocation_store: SqlAlchemyInvocationStore invocation_store: SqlAlchemyInvocationStore
artifact_store: SqlAlchemyArtifactStore
def create_sqlalchemy_storage(database_url: str) -> SqlAlchemyStorageBundle: def create_sqlalchemy_storage(database_url: str) -> SqlAlchemyStorageBundle:
@@ -655,4 +750,5 @@ def create_sqlalchemy_storage(database_url: str) -> SqlAlchemyStorageBundle:
worker_repository=SqlAlchemyWorkerRepository(session_factory), worker_repository=SqlAlchemyWorkerRepository(session_factory),
event_store=SqlAlchemyEventStore(session_factory), event_store=SqlAlchemyEventStore(session_factory),
invocation_store=SqlAlchemyInvocationStore(session_factory), invocation_store=SqlAlchemyInvocationStore(session_factory),
artifact_store=SqlAlchemyArtifactStore(session_factory),
) )
+86
View File
@@ -0,0 +1,86 @@
from pathlib import Path
from ai_orchestrator.application.ports import (
ArtifactRecord,
ModelInvocationRecord,
ToolInvocationRecord,
)
from ai_orchestrator.config import AppSettings
from ai_orchestrator.domain.events import DomainEvent
from ai_orchestrator.infrastructure.storage.factory import create_storage_bundle
def test_sqlite_paginated_event_and_invocation_queries(tmp_path: Path) -> None:
db_path = tmp_path / "queries.sqlite3"
storage = create_storage_bundle(
AppSettings(storage_backend="sqlite", database_url=f"sqlite:///{db_path.as_posix()}")
)
for index in range(5):
storage.event_store.append(
DomainEvent(event_type="test_event", task_id="task_1", payload={"index": index})
)
storage.invocation_store.save_model_invocation(
ModelInvocationRecord(
invocation_id=f"minv_{index}",
task_id="task_1",
node_id=None,
slot="weak",
provider="local",
model="mock",
status="success",
request={},
response={"message": {"content": "ok"}},
)
)
storage.invocation_store.save_tool_invocation(
ToolInvocationRecord(
invocation_id=f"tinv_{index}",
task_id="task_1",
node_id=None,
source_type="mcp",
source_id="one_c",
tool_name="one_c.run_sql",
status="success",
request={},
response={"content": {"rows": []}},
)
)
events_page = storage.event_store.list_by_task_paginated("task_1", limit=2, offset=1)
model_page = storage.invocation_store.list_model_invocations_paginated(
"task_1", limit=2, offset=2
)
tool_page = storage.invocation_store.list_tool_invocations_paginated(
"task_1", limit=3, offset=1
)
assert events_page.total == 5
assert len(events_page.items) == 2
assert model_page.total == 5
assert len(model_page.items) == 2
assert tool_page.total == 5
assert len(tool_page.items) == 3
def test_sqlite_artifact_store_persists_and_paginates(tmp_path: Path) -> None:
db_path = tmp_path / "artifacts.sqlite3"
storage = create_storage_bundle(
AppSettings(storage_backend="sqlite", database_url=f"sqlite:///{db_path.as_posix()}")
)
for index in range(4):
storage.artifact_store.save_artifact(
ArtifactRecord(
artifact_id=f"art_{index}",
task_id="task_1",
node_id=None,
artifact_type="file",
storage_uri=f"file:///artifact/{index}",
metadata={"index": index},
)
)
page = storage.artifact_store.list_by_task_paginated("task_1", limit=2, offset=1)
assert page.total == 4
assert len(page.items) == 2
assert page.items[0].storage_uri == "file:///artifact/1"
+5
View File
@@ -32,6 +32,7 @@ from ai_orchestrator.infrastructure.mcp_client import McpHttpClient, McpToolGate
from ai_orchestrator.infrastructure.model_router import StaticMockModelProvider from ai_orchestrator.infrastructure.model_router import StaticMockModelProvider
from ai_orchestrator.infrastructure.policy import StaticProjectPolicyEvaluator from ai_orchestrator.infrastructure.policy import StaticProjectPolicyEvaluator
from ai_orchestrator.infrastructure.storage.memory import ( from ai_orchestrator.infrastructure.storage.memory import (
InMemoryArtifactStore,
InMemoryConfirmationRepository, InMemoryConfirmationRepository,
InMemoryEventStore, InMemoryEventStore,
InMemoryGraphRepository, InMemoryGraphRepository,
@@ -90,11 +91,13 @@ def test_worker_gateway_result_becomes_artifact_ready_invocation() -> None:
worker_repository = InMemoryWorkerRepository() worker_repository = InMemoryWorkerRepository()
event_store = InMemoryEventStore() event_store = InMemoryEventStore()
invocation_store = InMemoryInvocationStore() invocation_store = InMemoryInvocationStore()
artifact_store = InMemoryArtifactStore()
worker_service = WorkerService(worker_repository=worker_repository, event_store=event_store) worker_service = WorkerService(worker_repository=worker_repository, event_store=event_store)
gateway = CapabilityAwareWorkerGateway( gateway = CapabilityAwareWorkerGateway(
worker_repository=worker_repository, worker_repository=worker_repository,
event_store=event_store, event_store=event_store,
invocation_store=invocation_store, invocation_store=invocation_store,
artifact_store=artifact_store,
connection_manager=InMemoryWorkerConnectionManager(), connection_manager=InMemoryWorkerConnectionManager(),
) )
response = worker_service.register( response = worker_service.register(
@@ -135,6 +138,7 @@ def test_worker_gateway_result_becomes_artifact_ready_invocation() -> None:
invocation_store.list_tool_invocations("task_1")[-1].response["artifacts"][0]["path"] invocation_store.list_tool_invocations("task_1")[-1].response["artifacts"][0]["path"]
== "D:/artifact.txt" == "D:/artifact.txt"
) )
assert artifact_store.list_by_task("task_1")[0].storage_uri == "D:/artifact.txt"
def test_model_and_mcp_paths_can_run_in_same_runtime_context() -> None: def test_model_and_mcp_paths_can_run_in_same_runtime_context() -> None:
@@ -206,6 +210,7 @@ def test_model_and_mcp_paths_can_run_in_same_runtime_context() -> None:
project_configs={"default": project}, project_configs={"default": project},
client=McpHttpClient(transport=httpx.MockTransport(handler)), client=McpHttpClient(transport=httpx.MockTransport(handler)),
invocation_store=invocation_store, invocation_store=invocation_store,
artifact_store=InMemoryArtifactStore(),
) )
model_result = router.run( model_result = router.run(
+9 -2
View File
@@ -1,3 +1,5 @@
from datetime import UTC, datetime
import httpx import httpx
from ai_orchestrator.application.services.execution import ( from ai_orchestrator.application.services.execution import (
@@ -32,6 +34,7 @@ from ai_orchestrator.infrastructure.mcp_client import McpHttpClient, McpToolGate
from ai_orchestrator.infrastructure.model_router import StaticMockModelProvider from ai_orchestrator.infrastructure.model_router import StaticMockModelProvider
from ai_orchestrator.infrastructure.policy import StaticProjectPolicyEvaluator from ai_orchestrator.infrastructure.policy import StaticProjectPolicyEvaluator
from ai_orchestrator.infrastructure.storage.memory import ( from ai_orchestrator.infrastructure.storage.memory import (
InMemoryArtifactStore,
InMemoryConfirmationRepository, InMemoryConfirmationRepository,
InMemoryEventStore, InMemoryEventStore,
InMemoryGraphRepository, InMemoryGraphRepository,
@@ -223,6 +226,7 @@ def test_smoke_policy_confirm_and_full_auto() -> None:
def test_smoke_mcp_and_worker_results_are_stored_and_finalizer_runs() -> None: def test_smoke_mcp_and_worker_results_are_stored_and_finalizer_runs() -> None:
invocation_store = InMemoryInvocationStore() invocation_store = InMemoryInvocationStore()
artifact_store = InMemoryArtifactStore()
def handler(request: httpx.Request) -> httpx.Response: def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response( return httpx.Response(
@@ -244,6 +248,7 @@ def test_smoke_mcp_and_worker_results_are_stored_and_finalizer_runs() -> None:
project_configs={"default": _project_config()}, project_configs={"default": _project_config()},
client=McpHttpClient(transport=httpx.MockTransport(handler)), client=McpHttpClient(transport=httpx.MockTransport(handler)),
invocation_store=invocation_store, invocation_store=invocation_store,
artifact_store=artifact_store,
) )
worker_repository = InMemoryWorkerRepository() worker_repository = InMemoryWorkerRepository()
event_store = InMemoryEventStore() event_store = InMemoryEventStore()
@@ -252,6 +257,7 @@ def test_smoke_mcp_and_worker_results_are_stored_and_finalizer_runs() -> None:
worker_repository=worker_repository, worker_repository=worker_repository,
event_store=event_store, event_store=event_store,
invocation_store=invocation_store, invocation_store=invocation_store,
artifact_store=artifact_store,
connection_manager=InMemoryWorkerConnectionManager(), connection_manager=InMemoryWorkerConnectionManager(),
) )
worker = worker_service.register( worker = worker_service.register(
@@ -279,8 +285,8 @@ def test_smoke_mcp_and_worker_results_are_stored_and_finalizer_runs() -> None:
task_id="task_1", task_id="task_1",
tool="file.read", tool="file.read",
status="success", status="success",
started_at=__import__("datetime").datetime.now(__import__("datetime").UTC), started_at=datetime.now(UTC),
finished_at=__import__("datetime").datetime.now(__import__("datetime").UTC), finished_at=datetime.now(UTC),
duration_ms=1, duration_ms=1,
result={"content": "ok"}, result={"content": "ok"},
), ),
@@ -318,4 +324,5 @@ def test_smoke_mcp_and_worker_results_are_stored_and_finalizer_runs() -> None:
assert queued.status == "queued" assert queued.status == "queued"
assert tool_result.status == "success" assert tool_result.status == "success"
assert len(invocation_store.list_tool_invocations("task_1")) >= 2 assert len(invocation_store.list_tool_invocations("task_1")) >= 2
assert artifact_store.list_by_task("task_1") == []
assert graph.nodes[-1].output_data["status"] == "completed" assert graph.nodes[-1].output_data["status"] == "completed"
+5 -1
View File
@@ -10,7 +10,10 @@ from ai_orchestrator.config import (
ProjectMetadata, ProjectMetadata,
) )
from ai_orchestrator.infrastructure.mcp_client import McpHttpClient, McpToolGateway from ai_orchestrator.infrastructure.mcp_client import McpHttpClient, McpToolGateway
from ai_orchestrator.infrastructure.storage.memory import InMemoryInvocationStore from ai_orchestrator.infrastructure.storage.memory import (
InMemoryArtifactStore,
InMemoryInvocationStore,
)
def test_mcp_tool_gateway_normalizes_result_and_persists_invocation() -> None: def test_mcp_tool_gateway_normalizes_result_and_persists_invocation() -> None:
@@ -47,6 +50,7 @@ def test_mcp_tool_gateway_normalizes_result_and_persists_invocation() -> None:
project_configs={"default": project}, project_configs={"default": project},
client=McpHttpClient(transport=transport), client=McpHttpClient(transport=transport),
invocation_store=invocation_store, invocation_store=invocation_store,
artifact_store=InMemoryArtifactStore(),
) )
result = gateway.call( result = gateway.call(
+29
View File
@@ -0,0 +1,29 @@
from ai_orchestrator.application.services.execution import StructuredReviewer
from ai_orchestrator.domain.enums import NodeType
from ai_orchestrator.domain.models import ExecutionNode, Task
def test_structured_reviewer_accepts_valid_output() -> None:
reviewer = StructuredReviewer()
task = Task(project_id="default", goal="Review", inputs={})
node = ExecutionNode(task_id=task.task_id, node_type=NodeType.PLANNER, input_data={})
reviewed = reviewer.review(
task=task,
node=node,
output_data={"status": "planned", "goal": "Review"},
)
assert reviewed["review"]["status"] == "accepted"
assert reviewed["task_id"] == task.task_id
def test_structured_reviewer_marks_missing_status_as_review_failed() -> None:
reviewer = StructuredReviewer()
task = Task(project_id="default", goal="Review", inputs={})
node = ExecutionNode(task_id=task.task_id, node_type=NodeType.PLANNER, input_data={})
reviewed = reviewer.review(task=task, node=node, output_data={"goal": "Review"})
assert reviewed["status"] == "review_failed"
assert reviewed["review"]["status"] == "rejected"
+8 -4
View File
@@ -10,6 +10,7 @@ from ai_orchestrator.application.services.workers import (
WorkerService, WorkerService,
) )
from ai_orchestrator.infrastructure.storage.memory import ( from ai_orchestrator.infrastructure.storage.memory import (
InMemoryArtifactStore,
InMemoryEventStore, InMemoryEventStore,
InMemoryInvocationStore, InMemoryInvocationStore,
InMemoryWorkerRepository, InMemoryWorkerRepository,
@@ -22,22 +23,25 @@ def _build_gateway() -> tuple[
InMemoryWorkerRepository, InMemoryWorkerRepository,
InMemoryEventStore, InMemoryEventStore,
InMemoryInvocationStore, InMemoryInvocationStore,
InMemoryArtifactStore,
]: ]:
worker_repository = InMemoryWorkerRepository() worker_repository = InMemoryWorkerRepository()
event_store = InMemoryEventStore() event_store = InMemoryEventStore()
invocation_store = InMemoryInvocationStore() invocation_store = InMemoryInvocationStore()
artifact_store = InMemoryArtifactStore()
service = WorkerService(worker_repository=worker_repository, event_store=event_store) service = WorkerService(worker_repository=worker_repository, event_store=event_store)
gateway = CapabilityAwareWorkerGateway( gateway = CapabilityAwareWorkerGateway(
worker_repository=worker_repository, worker_repository=worker_repository,
event_store=event_store, event_store=event_store,
invocation_store=invocation_store, invocation_store=invocation_store,
artifact_store=artifact_store,
connection_manager=InMemoryWorkerConnectionManager(), connection_manager=InMemoryWorkerConnectionManager(),
) )
return gateway, service, worker_repository, event_store, invocation_store return gateway, service, worker_repository, event_store, invocation_store, artifact_store
def test_worker_gateway_dispatches_and_completes_command() -> None: def test_worker_gateway_dispatches_and_completes_command() -> None:
gateway, service, worker_repository, event_store, invocation_store = _build_gateway() gateway, service, worker_repository, event_store, invocation_store, _ = _build_gateway()
response = service.register( response = service.register(
RegisterWorkerRequest( RegisterWorkerRequest(
worker_id="worker_home_pc", worker_id="worker_home_pc",
@@ -80,7 +84,7 @@ def test_worker_gateway_dispatches_and_completes_command() -> None:
def test_worker_gateway_rejects_unsupported_capability() -> None: def test_worker_gateway_rejects_unsupported_capability() -> None:
gateway, service, _, _, _ = _build_gateway() gateway, service, _, _, _, _ = _build_gateway()
response = service.register( response = service.register(
RegisterWorkerRequest( RegisterWorkerRequest(
worker_id="worker_home_pc", worker_id="worker_home_pc",
@@ -103,7 +107,7 @@ def test_worker_gateway_rejects_unsupported_capability() -> None:
def test_worker_service_heartbeat_and_stale_tracking() -> None: def test_worker_service_heartbeat_and_stale_tracking() -> None:
_, service, worker_repository, event_store, _ = _build_gateway() _, service, worker_repository, event_store, _, _ = _build_gateway()
response = service.register( response = service.register(
RegisterWorkerRequest( RegisterWorkerRequest(
worker_id="worker_home_pc", worker_id="worker_home_pc",