from datetime import UTC, datetime import httpx from ai_orchestrator.application.services.execution import ( DefaultFinalizer, FinalizerNodeRunner, GraphExecutionEngine, NoOpReviewer, PlannerNodeRunner, RunnerRegistry, ) from ai_orchestrator.application.services.orchestrator import CreateTaskRequest, OrchestratorService from ai_orchestrator.application.services.router import ConfigurableModelRouter, ProviderResponse from ai_orchestrator.application.services.workers import ( CapabilityAwareWorkerGateway, InMemoryWorkerConnectionManager, RegisterWorkerRequest, WorkerCommandResult, WorkerService, ) from ai_orchestrator.config import ( ExecutionConfig, McpServerConfig, ModelsConfig, ModelSlotConfig, PolicyConfig, ProjectConfig, ProjectMetadata, ) from ai_orchestrator.infrastructure.mcp_client import McpHttpClient, McpToolGateway from ai_orchestrator.infrastructure.model_router import StaticMockModelProvider from ai_orchestrator.infrastructure.policy import StaticProjectPolicyEvaluator from ai_orchestrator.infrastructure.storage.memory import ( InMemoryConfirmationRepository, InMemoryEventStore, InMemoryGraphRepository, InMemoryInvocationStore, InMemoryTaskRepository, InMemoryWorkerRepository, ) def test_task_can_plan_execute_and_finalize_over_two_scheduler_passes() -> None: task_repository = InMemoryTaskRepository() graph_repository = InMemoryGraphRepository() confirmation_repository = InMemoryConfirmationRepository() event_store = InMemoryEventStore() orchestrator = OrchestratorService( task_repository=task_repository, graph_repository=graph_repository, confirmation_repository=confirmation_repository, event_store=event_store, policy_evaluator=StaticProjectPolicyEvaluator( projects={ "default": ProjectConfig( project=ProjectMetadata(id="default", name="Default"), models=ModelsConfig( weak=ModelSlotConfig(provider="local", model="weak"), strong=ModelSlotConfig(provider="local", model="strong"), vision=ModelSlotConfig(provider="disabled"), embedding=ModelSlotConfig(provider="local", model="embed"), ), execution=ExecutionConfig(), policy=PolicyConfig(), ) } ), ) registry = RunnerRegistry() registry.register(PlannerNodeRunner()) registry.register(FinalizerNodeRunner(finalizer=DefaultFinalizer())) engine = GraphExecutionEngine( event_store=event_store, runner_registry=registry, reviewer=NoOpReviewer(), ) task = orchestrator.create_task( CreateTaskRequest(project_id="default", goal="Scenario", inputs={}) ) graph = orchestrator.plan_task(task.task_id) engine.execute_ready_nodes(task=task, graph=graph) engine.execute_ready_nodes(task=task, graph=graph) assert graph.nodes[-1].output_data["status"] == "completed" def test_worker_gateway_result_becomes_artifact_ready_invocation() -> None: worker_repository = InMemoryWorkerRepository() event_store = InMemoryEventStore() invocation_store = InMemoryInvocationStore() worker_service = WorkerService(worker_repository=worker_repository, event_store=event_store) gateway = CapabilityAwareWorkerGateway( worker_repository=worker_repository, event_store=event_store, invocation_store=invocation_store, connection_manager=InMemoryWorkerConnectionManager(), ) response = worker_service.register( RegisterWorkerRequest( worker_id="worker_home_pc", name="Home PC", machine="DESKTOP-1", os="windows", version="0.1.0", capabilities=["file.read"], ) ) queued = gateway.dispatch( project_id="default", worker_session_id=response.worker.session_id, command_name="file.read", args={"path": "D:/artifact.txt"}, task_context={"task_id": "task_1"}, ) polled = gateway.poll_commands(response.worker.session_id) gateway.complete_command( response.worker.session_id, WorkerCommandResult( command_id=polled[0].command_id, task_id="task_1", tool="file.read", status="success", started_at=polled[0].policy_context.get("started_at") or datetime.now(UTC), finished_at=datetime.now(UTC), duration_ms=100, result={"content": "ok"}, artifacts=[{"type": "file", "path": "D:/artifact.txt"}], ), ) assert queued.status == "queued" assert ( invocation_store.list_tool_invocations("task_1")[-1].response["artifacts"][0]["path"] == "D:/artifact.txt" ) def test_model_and_mcp_paths_can_run_in_same_runtime_context() -> None: project = ProjectConfig( project=ProjectMetadata(id="default", name="Default"), models=ModelsConfig( weak=ModelSlotConfig(provider="local", model="weak"), strong=ModelSlotConfig(provider="external", model="strong"), vision=ModelSlotConfig(provider="disabled"), embedding=ModelSlotConfig(provider="local", model="embed"), ), mcp_servers={"one_c": McpServerConfig(base_url="http://mcp.test")}, execution=ExecutionConfig(allow_external_models=True), policy=PolicyConfig(), ) invocation_store = InMemoryInvocationStore() router = ConfigurableModelRouter( project_configs={"default": project}, providers={ "local": StaticMockModelProvider( provider_name="local", responses=[ ProviderResponse( status="success", message={"content": ""}, tool_calls=[], usage={}, ), ProviderResponse( status="success", message={"content": ""}, tool_calls=[], usage={}, ), ], ), "external": StaticMockModelProvider( provider_name="external", responses=[ ProviderResponse( status="success", message={"content": "strong"}, tool_calls=[], usage={"input_tokens": 1}, ) ], ), }, invocation_store=invocation_store, ) def handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, json={ "jsonrpc": "2.0", "id": "1", "result": { "status": "success", "content": {"rows": [{"id": 1}]}, "artifacts": [], "logs": [], "error": None, }, }, ) gateway = McpToolGateway( project_configs={"default": project}, client=McpHttpClient(transport=httpx.MockTransport(handler)), invocation_store=invocation_store, ) model_result = router.run( project_id="default", slot="weak", messages=[{"role": "user", "content": "Hi"}], task_context={"task_id": "task_1"}, ) tool_result = gateway.call( project_id="default", server_id="one_c", tool_name="one_c.run_sql", args={"query": "select 1"}, task_context={"task_id": "task_1"}, ) assert model_result.message["content"] == "strong" assert tool_result.content["rows"][0]["id"] == 1