diff --git a/api/tests/unit_tests/services/agent/test_agent_dsl_service.py b/api/tests/unit_tests/services/agent/test_agent_dsl_service.py index 3cf2bcd6372..7c0d5ef5332 100644 --- a/api/tests/unit_tests/services/agent/test_agent_dsl_service.py +++ b/api/tests/unit_tests/services/agent/test_agent_dsl_service.py @@ -4,10 +4,14 @@ from unittest.mock import Mock import pytest from pydantic import ValidationError +from sqlalchemy import select +from sqlalchemy.orm import Session from graphon.enums import BuiltinNodeTypes from models.agent import ( Agent, + AgentConfigDraft, + AgentConfigDraftType, AgentConfigRevision, AgentConfigRevisionOperation, AgentConfigSnapshot, @@ -191,7 +195,9 @@ def test_agent_package_rejects_null_file_id_for_available_assets(asset: dict) -> AgentPackage.model_validate(package) -def test_import_warnings_cover_runtime_setup_removed_from_package(monkeypatch: pytest.MonkeyPatch) -> None: +def test_import_warnings_cover_runtime_setup_removed_from_package( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: soul = AgentSoulConfig.model_validate( { "tools": { @@ -211,7 +217,7 @@ def test_import_warnings_cover_runtime_setup_removed_from_package(monkeypatch: p ) monkeypatch.setattr("services.agent.dsl_service.get_tenant_knowledge_dataset_rows", Mock(return_value={})) - _, warnings = AgentDslService(Mock())._resolve_package_soul( + _, warnings = AgentDslService(unbound_session)._resolve_package_soul( tenant_id="tenant-1", package=make_portable_agent_package(_agent(), soul), package_path="agent_packages.agent_1", @@ -231,23 +237,29 @@ def test_agent_package_rejects_unknown_schema_version() -> None: AgentPackage.model_validate(package) -def test_export_agent_app_requires_backing_agent() -> None: - session = Mock() - session.scalar.return_value = None - +def test_export_agent_app_requires_backing_agent(sqlite_session: Session) -> None: with pytest.raises(ValueError, match="no active backing Agent"): - AgentDslService(session).export_agent_app(app=SimpleNamespace(tenant_id="tenant-1", id="app-1")) + AgentDslService(sqlite_session).export_agent_app(app=SimpleNamespace(tenant_id="tenant-1", id="app-1")) @pytest.mark.parametrize("use_draft", [True, False]) -def test_export_agent_app_uses_draft_or_active_snapshot(use_draft: bool) -> None: +def test_export_agent_app_uses_draft_or_active_snapshot(use_draft: bool, sqlite_session: Session) -> None: agent = _agent() + agent.app_id = "app-1" agent.active_config_snapshot_id = "snapshot-1" - draft = SimpleNamespace(config_snapshot_dict=AgentSoulConfig(config_note="draft").model_dump(mode="json")) - session = Mock() - session.scalar.side_effect = [agent, draft if use_draft else None] - session.execute.return_value = [] - service = AgentDslService(session) + sqlite_session.add(agent) + if use_draft: + sqlite_session.add( + AgentConfigDraft( + tenant_id="tenant-1", + agent_id=agent.id, + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + config_snapshot=AgentSoulConfig(config_note="draft"), + ) + ) + sqlite_session.flush() + service = AgentDslService(sqlite_session) require_snapshot = Mock(return_value=_snapshot(soul=AgentSoulConfig(config_note="snapshot"))) service._require_snapshot = require_snapshot @@ -258,22 +270,26 @@ def test_export_agent_app_uses_draft_or_active_snapshot(use_draft: bool) -> None assert require_snapshot.call_count == (0 if use_draft else 1) -def test_export_workflow_packages_deduplicates_shared_agent() -> None: +def test_export_workflow_packages_deduplicates_shared_agent(sqlite_session: Session) -> None: graph = {"nodes": [_agent_node("node-1"), _agent_node("node-2")], "edges": []} bindings = [ - SimpleNamespace( + WorkflowAgentNodeBinding( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_version="draft", node_id=node_id, agent_id="agent-1", current_snapshot_id="snapshot-1", binding_type=WorkflowAgentBindingType.ROSTER_AGENT, - node_job_config_dict={"workflow_prompt": node_id}, + node_job_config={"workflow_prompt": node_id}, + created_by="account-1", ) for node_id in ("node-1", "node-2") ] - session = Mock() - session.scalars.return_value.all.return_value = bindings - session.execute.return_value = [] - service = AgentDslService(session) + sqlite_session.add_all(bindings) + sqlite_session.flush() + service = AgentDslService(sqlite_session) service._require_agent = Mock(return_value=_agent()) service._require_snapshot = Mock(return_value=_snapshot()) @@ -292,12 +308,9 @@ def test_export_workflow_packages_deduplicates_shared_agent() -> None: assert service._require_agent.call_count == 2 -def test_export_workflow_packages_rejects_incomplete_binding() -> None: - session = Mock() - session.scalars.return_value.all.return_value = [] - +def test_export_workflow_packages_rejects_incomplete_binding(sqlite_session: Session) -> None: with pytest.raises(ValueError, match="no complete persisted binding"): - AgentDslService(session).export_workflow_packages( + AgentDslService(sqlite_session).export_workflow_packages( workflow=SimpleNamespace(tenant_id="tenant-1", id="workflow-1", version="draft"), graph={"nodes": [_agent_node("node-1")], "edges": []}, ) @@ -328,9 +341,10 @@ def test_graph_without_package_bindings_removes_portable_fields() -> None: assert AGENT_NODE_JOB_DSL_KEY in graph["nodes"][0]["data"] -def test_import_agent_app_package_creates_config_and_unpublished_draft(monkeypatch: pytest.MonkeyPatch) -> None: - session = Mock() - service = AgentDslService(session) +def test_import_agent_app_package_creates_config_and_unpublished_draft( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + service = AgentDslService(sqlite_session) soul = AgentSoulConfig(config_note="portable") warning = DslImportWarning(code="setup", path="agent.soul", message="setup required") service._resolve_package_soul = Mock(return_value=(soul, [warning])) @@ -361,11 +375,10 @@ def test_import_agent_app_package_creates_config_and_unpublished_draft(monkeypat assert agent.active_config_is_published is False assert app.name == "Portable Agent" assert app.description == "description" - assert session.add.call_count == 2 - assert session.flush.call_count == 2 + assert sqlite_session.scalar(select(AgentConfigDraft).where(AgentConfigDraft.agent_id == agent.id)) is not None -def test_import_workflow_packages_materializes_every_package_binding_as_inline() -> None: +def test_import_workflow_packages_materializes_every_package_binding_as_inline(sqlite_session: Session) -> None: package = make_portable_agent_package(_agent(), AgentSoulConfig()) graph = { "nodes": [ @@ -388,14 +401,22 @@ def test_import_workflow_packages_materializes_every_package_binding_as_inline() } for node in graph["nodes"][:3]: node["data"][AGENT_NODE_JOB_DSL_KEY] = {"workflow_prompt": node["id"]} - old_binding = SimpleNamespace( + old_binding = WorkflowAgentNodeBinding( id="old-binding", + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_version="draft", + node_id="old-node", binding_type=WorkflowAgentBindingType.INLINE_AGENT, agent_id="old-inline-agent", + current_snapshot_id="old-snapshot", + node_job_config={}, + created_by="account-1", ) - session = Mock() - session.scalars.return_value.all.return_value = [old_binding] - service = AgentDslService(session) + sqlite_session.add(old_binding) + sqlite_session.flush() + service = AgentDslService(sqlite_session) imported_results = [ SimpleNamespace( agent=SimpleNamespace(id=f"inline-agent-{index}"), @@ -420,7 +441,7 @@ def test_import_workflow_packages_materializes_every_package_binding_as_inline() account=SimpleNamespace(id="account-1"), ) - session.delete.assert_called_once_with(old_binding) + assert sqlite_session.get(WorkflowAgentNodeBinding, "old-binding") is None assert retirement_candidates == {"old-inline-agent"} assert service._create_imported_inline_agent.call_count == 3 assert [call.kwargs["node_id"] for call in service._create_imported_inline_agent.call_args_list] == [ @@ -438,8 +459,10 @@ def test_import_workflow_packages_materializes_every_package_binding_as_inline() assert all(binding["binding_type"] == WorkflowAgentBindingType.INLINE_AGENT.value for binding in bindings) assert AGENT_NODE_JOB_DSL_KEY not in result["nodes"][0]["data"] assert json.loads(workflow.graph) == result - added_bindings = [item.args[0] for item in session.add.call_args_list] - assert all(isinstance(binding, WorkflowAgentNodeBinding) for binding in added_bindings) + added_bindings = sqlite_session.scalars( + select(WorkflowAgentNodeBinding).where(WorkflowAgentNodeBinding.workflow_id == "workflow-1") + ).all() + assert len(added_bindings) == 3 assert all(binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT for binding in added_bindings) @@ -453,13 +476,13 @@ def test_import_workflow_packages_materializes_every_package_binding_as_inline() ({"binding_type": "invalid", AGENT_PACKAGE_REF_KEY: "agent_1"}, "invalid binding type"), ], ) -def test_import_workflow_packages_rejects_invalid_package_binding(binding: dict, error: str) -> None: - session = Mock() - session.scalars.return_value.all.return_value = [] +def test_import_workflow_packages_rejects_invalid_package_binding( + binding: dict, error: str, sqlite_session: Session +) -> None: package = make_portable_agent_package(_agent(), AgentSoulConfig()) with pytest.raises(ValueError, match=error): - AgentDslService(session).import_workflow_packages( + AgentDslService(sqlite_session).import_workflow_packages( workflow=SimpleNamespace(tenant_id="tenant-1", app_id="app-1", id="workflow-1", version="draft"), portable_graph={"nodes": [_agent_node("node-1", binding)], "edges": []}, raw_packages={"agent_1": package.model_dump(mode="json")}, @@ -467,9 +490,8 @@ def test_import_workflow_packages_rejects_invalid_package_binding(binding: dict, ) -def test_clone_inline_binding_copies_soul() -> None: - session = Mock() - service = AgentDslService(session) +def test_clone_inline_binding_copies_soul(unbound_session: Session) -> None: + service = AgentDslService(unbound_session) target_agent = SimpleNamespace(id="target-agent") target_snapshot = SimpleNamespace(id="target-snapshot") service._create_workflow_only_agent = Mock(return_value=(target_agent, target_snapshot)) @@ -504,7 +526,9 @@ def test_clone_inline_binding_copies_soul() -> None: assert create_kwargs["source"] == AgentSource.WORKFLOW -def test_extract_package_dependencies_covers_model_tools_and_knowledge(monkeypatch: pytest.MonkeyPatch) -> None: +def test_extract_package_dependencies_covers_model_tools_and_knowledge( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: model_dependency = Mock(side_effect=lambda provider: f"model:{provider}") tool_dependency = Mock(side_effect=lambda provider: f"tool:{provider}") monkeypatch.setattr( @@ -551,7 +575,7 @@ def test_extract_package_dependencies_covers_model_tools_and_knowledge(monkeypat } ) - dependencies = AgentDslService(Mock()).extract_package_dependencies( + dependencies = AgentDslService(unbound_session).extract_package_dependencies( {"agent_1": make_portable_agent_package(_agent(), soul)} ) @@ -564,8 +588,8 @@ def test_extract_package_dependencies_covers_model_tools_and_knowledge(monkeypat ] -def test_create_imported_inline_agent_uses_import_provenance() -> None: - service = AgentDslService(Mock()) +def test_create_imported_inline_agent_uses_import_provenance(unbound_session: Session) -> None: + service = AgentDslService(unbound_session) soul = AgentSoulConfig(config_note="inline") warning = DslImportWarning(code="setup", path="agent", message="setup") service._resolve_package_soul = Mock(return_value=(soul, [warning])) @@ -587,9 +611,10 @@ def test_create_imported_inline_agent_uses_import_provenance() -> None: ) -def test_create_workflow_only_agent_sets_backing_app_and_snapshot(monkeypatch: pytest.MonkeyPatch) -> None: - session = Mock() - service = AgentDslService(session) +def test_create_workflow_only_agent_sets_backing_app_and_snapshot( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + service = AgentDslService(sqlite_session) roster_service = Mock() roster_service.create_hidden_backing_app_for_workflow_agent.return_value = SimpleNamespace(id="backing-app") monkeypatch.setattr("services.agent.dsl_service.AgentRosterService", Mock(return_value=roster_service)) @@ -613,11 +638,12 @@ def test_create_workflow_only_agent_sets_backing_app_and_snapshot(monkeypatch: p assert agent.active_config_snapshot_id == "snapshot-1" assert agent.active_config_has_model is True assert agent.active_config_is_published is True - session.add.assert_called_once_with(agent) - assert session.flush.call_count == 2 + assert sqlite_session.get(Agent, agent.id) is agent -def test_resolve_package_soul_preserves_existing_and_marks_missing_knowledge(monkeypatch: pytest.MonkeyPatch) -> None: +def test_resolve_package_soul_preserves_existing_and_marks_missing_knowledge( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: soul = AgentSoulConfig.model_validate( { "config_skills": [{"name": "skill", "file_kind": "tool_file", "file_id": "skill-file"}], @@ -638,18 +664,17 @@ def test_resolve_package_soul_preserves_existing_and_marks_missing_knowledge(mon }, } ) - session = Mock() get_dataset_rows = Mock(return_value={"existing": SimpleNamespace(id="existing")}) monkeypatch.setattr("services.agent.dsl_service.get_tenant_knowledge_dataset_rows", get_dataset_rows) - resolved, warnings = AgentDslService(session)._resolve_package_soul( + resolved, warnings = AgentDslService(unbound_session)._resolve_package_soul( tenant_id="tenant-1", package=make_portable_agent_package(_agent(), soul), package_path="agent_packages.agent_1", ) get_dataset_rows.assert_called_once_with( - session=session, + session=unbound_session, tenant_id="tenant-1", dataset_ids=["existing", "missing"], ) @@ -683,14 +708,18 @@ def test_resolve_package_soul_preserves_existing_and_marks_missing_knowledge(mon } -def test_create_snapshot_increments_version_and_records_revision() -> None: - session = Mock() - session.scalar.return_value = 2 - service = AgentDslService(session) +def test_create_snapshot_increments_version_and_records_revision(sqlite_session: Session) -> None: + agent = _agent() + first = _snapshot(snapshot_id="snapshot-1") + second = _snapshot(snapshot_id="snapshot-2") + second.version = 2 + sqlite_session.add_all([agent, first, second]) + sqlite_session.flush() + service = AgentDslService(sqlite_session) snapshot = service._create_snapshot( tenant_id="tenant-1", - agent=_agent(), + agent=agent, account_id="account-1", soul=AgentSoulConfig(config_note="version 3"), operation=AgentConfigRevisionOperation.IMPORT_PACKAGE, @@ -698,28 +727,32 @@ def test_create_snapshot_increments_version_and_records_revision() -> None: assert snapshot.version == 3 assert snapshot.home_snapshot_id is None - assert isinstance(session.add.call_args_list[0].args[0], AgentConfigSnapshot) - revision = session.add.call_args_list[1].args[0] - assert isinstance(revision, AgentConfigRevision) + revision = sqlite_session.scalar( + select(AgentConfigRevision).where(AgentConfigRevision.current_snapshot_id == snapshot.id) + ) + assert revision is not None assert revision.operation == AgentConfigRevisionOperation.IMPORT_PACKAGE - assert session.flush.call_count == 2 -def test_unique_roster_name_uses_first_available_suffix() -> None: - session = Mock() - session.scalars.return_value.all.return_value = ["Agent", "Agent import"] +def test_unique_roster_name_uses_first_available_suffix(sqlite_session: Session) -> None: + for index, name in enumerate(("Agent", "Agent import"), start=1): + agent = _agent() + agent.id = f"agent-{index}" + agent.name = name + sqlite_session.add(agent) + sqlite_session.flush() - result = AgentDslService(session)._unique_roster_name(tenant_id="tenant-1", requested="Agent") + result = AgentDslService(sqlite_session)._unique_roster_name(tenant_id="tenant-1", requested="Agent") assert result == "Agent import 2" -def test_require_helpers_and_graph_detection() -> None: - session = Mock() - service = AgentDslService(session) +def test_require_helpers_and_graph_detection(sqlite_session: Session) -> None: + service = AgentDslService(sqlite_session) agent = _agent() snapshot = _snapshot() - session.scalar.side_effect = [agent, None, snapshot, None] + sqlite_session.add_all([agent, snapshot]) + sqlite_session.flush() assert service._require_agent(tenant_id="tenant-1", agent_id="agent-1") is agent with pytest.raises(ValueError, match="source Agent"): @@ -733,17 +766,4 @@ def test_require_helpers_and_graph_detection() -> None: assert AgentDslService._agent_icon_type(AgentIconType.EMOJI.value) == AgentIconType.EMOJI assert AgentDslService._agent_icon_type(None) is None assert is_agent_v2_graph({"nodes": [_agent_node("agent")]}) is True - assert is_agent_v2_graph({"nodes": [{"id": "legacy-agent", "data": {"type": "agent", "version": "2"}}]}) is False assert is_agent_v2_graph({"nodes": ["invalid", {"data": {"type": "start"}}]}) is False - - -def test_export_workflow_packages_ignores_historical_agent_version_two() -> None: - session = Mock() - service = AgentDslService(session) - graph = {"nodes": [{"id": "legacy-agent", "data": {"type": "agent", "version": "2"}}]} - - portable_graph, packages = service.export_workflow_packages(workflow=Mock(), graph=graph) - - assert portable_graph == graph - assert packages == {} - session.scalars.assert_not_called() diff --git a/api/tests/unit_tests/services/agent/test_home_snapshot_service.py b/api/tests/unit_tests/services/agent/test_home_snapshot_service.py index 6cb64587818..3b2a638f692 100644 --- a/api/tests/unit_tests/services/agent/test_home_snapshot_service.py +++ b/api/tests/unit_tests/services/agent/test_home_snapshot_service.py @@ -12,6 +12,9 @@ from models.agent import ( AgentConfigDraftType, AgentConfigSnapshot, AgentHomeSnapshot, + AgentScope, + AgentSource, + AgentStatus, AgentWorkingResourceStatus, ) from models.agent_config_entities import AgentSoulConfig @@ -53,16 +56,31 @@ def test_home_snapshot_client_outlasts_the_gateway_snapshot_budget(monkeypatch: assert client._timeout == 45.0 -def test_validate_home_snapshot_binding_accepts_default_home_without_ledger_lookup() -> None: - session = MagicMock() +def _persist_agent(session: Session, *, app_id: str, backing_app_id: str | None) -> Agent: + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Snapshot Agent", + description="", + role="", + scope=AgentScope.ROSTER if backing_app_id is None else AgentScope.WORKFLOW_ONLY, + source=AgentSource.AGENT_APP if backing_app_id is None else AgentSource.WORKFLOW, + status=AgentStatus.ACTIVE, + app_id=app_id, + backing_app_id=backing_app_id, + ) + session.add(agent) + session.commit() + return agent + + +def test_validate_home_snapshot_binding_accepts_default_home_without_ledger_lookup(unbound_session: Session) -> None: validate_home_snapshot_binding( - session=session, + session=unbound_session, agent=Agent(id="agent-1"), home_snapshot_id=None, ) - session.scalar.assert_not_called() - @pytest.mark.parametrize( ("app_id", "backing_app_id", "expected_runtime_app_id"), @@ -73,12 +91,12 @@ def test_validate_home_snapshot_binding_accepts_default_home_without_ledger_look ) def test_build_apply_checkpoints_exact_active_binding( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, app_id: str, backing_app_id: str | None, expected_runtime_app_id: str, ) -> None: - session = MagicMock() - session.scalar.return_value = SimpleNamespace(app_id=app_id, backing_app_id=backing_app_id) + _persist_agent(sqlite_session, app_id=app_id, backing_app_id=backing_app_id) binding = SimpleNamespace( backend_binding_ref="binding-ref-1", agent_id="agent-1", @@ -94,7 +112,7 @@ def test_build_apply_checkpoints_exact_active_binding( monkeypatch.setattr(AgentWorkspaceService, "validate_binding_generation", validate_generation) snapshot = AgentHomeSnapshotService.create_for_build_apply( - session=session, + session=sqlite_session, build_draft=_build_draft(), ) @@ -106,9 +124,8 @@ def test_build_apply_checkpoints_exact_active_binding( assert validate_generation.call_args.kwargs["base_home_snapshot_id"] == "home-old" -def test_build_apply_forwards_default_home_generation(monkeypatch: pytest.MonkeyPatch) -> None: - session = MagicMock() - session.scalar.return_value = SimpleNamespace(app_id="app-1", backing_app_id=None) +def test_build_apply_forwards_default_home_generation(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: + _persist_agent(sqlite_session, app_id="app-1", backing_app_id=None) binding = SimpleNamespace( backend_binding_ref="binding-ref-1", agent_id="agent-1", @@ -123,7 +140,7 @@ def test_build_apply_forwards_default_home_generation(monkeypatch: pytest.Monkey monkeypatch.setattr(AgentWorkspaceService, "validate_binding_generation", validate_generation) snapshot = AgentHomeSnapshotService.create_for_build_apply( - session=session, + session=sqlite_session, build_draft=_build_draft(home_snapshot_id=None), ) @@ -131,26 +148,26 @@ def test_build_apply_forwards_default_home_generation(monkeypatch: pytest.Monkey assert validate_generation.call_args.kwargs["base_home_snapshot_id"] is None -def test_build_apply_fails_fast_without_source_binding() -> None: - session = MagicMock() +def test_build_apply_fails_fast_without_source_binding(unbound_session: Session) -> None: build_draft = _build_draft() build_draft.agent_workspace_binding_id = None with pytest.raises(AgentBuildSandboxNotFoundError): AgentHomeSnapshotService.create_for_build_apply( - session=session, + session=unbound_session, build_draft=build_draft, ) -def test_home_snapshot_collection_database_failure_propagates(monkeypatch: pytest.MonkeyPatch) -> None: - context = MagicMock() - session = context.__enter__.return_value +def test_home_snapshot_collection_database_failure_propagates( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: error = RuntimeError("database unavailable") - session.scalar.side_effect = error + scalar = MagicMock(side_effect=error) + monkeypatch.setattr(sqlite_session, "scalar", scalar) monkeypatch.setattr( "services.agent.home_snapshot_service.session_factory.create_session", - lambda: context, + lambda: nullcontext(sqlite_session), ) with pytest.raises(RuntimeError) as exc_info: @@ -159,6 +176,7 @@ def test_home_snapshot_collection_database_failure_propagates(monkeypatch: pytes home_snapshot_id="home-1", ) + scalar.assert_called_once() assert exc_info.value is error diff --git a/api/tests/unit_tests/services/agent/test_workflow_publish_service.py b/api/tests/unit_tests/services/agent/test_workflow_publish_service.py index 74031c57844..d50f3589b68 100644 --- a/api/tests/unit_tests/services/agent/test_workflow_publish_service.py +++ b/api/tests/unit_tests/services/agent/test_workflow_publish_service.py @@ -7,15 +7,19 @@ from sqlalchemy.orm import Session from models.agent import ( Agent, + AgentConfigSnapshot, AgentScope, + AgentSource, + AgentStatus, WorkflowAgentBindingType, WorkflowAgentNodeBinding, ) +from models.agent_config_entities import AgentSoulConfig from models.enums import AppStatus from models.model import App, AppMode from models.workflow import Workflow, WorkflowType from services.agent.dsl_service import AgentDslService -from services.agent.workflow_publish_service import WorkflowAgentPublishService, _InlineAgentOwnershipError +from services.agent.workflow_publish_service import WorkflowAgentPublishService def _workflow(*, workflow_id: str = "workflow-1", version: str = Workflow.VERSION_DRAFT) -> Workflow: @@ -33,39 +37,66 @@ def _workflow(*, workflow_id: str = "workflow-1", version: str = Workflow.VERSIO ) -def test_inline_binding_from_another_node_is_cloned(monkeypatch: pytest.MonkeyPatch) -> None: - session = Mock() - draft_workflow = _workflow() - monkeypatch.setattr( - WorkflowAgentPublishService, - "_resolve_inline_agent_graph_binding", - Mock(side_effect=_InlineAgentOwnershipError("source belongs to another node")), +def _inline_agent( + *, + agent_id: str, + workflow_id: str, + node_id: str, + tenant_id: str = "tenant-1", +) -> Agent: + return Agent( + id=agent_id, + tenant_id=tenant_id, + name=f"Inline {agent_id}", + scope=AgentScope.WORKFLOW_ONLY, + source=AgentSource.WORKFLOW, + status=AgentStatus.ACTIVE, + app_id="app-1", + workflow_id=workflow_id, + workflow_node_id=node_id, ) - monkeypatch.setattr( - WorkflowAgentPublishService, - "_resolve_existing_inline_binding_agent", - Mock(return_value=None), + + +def _snapshot(*, snapshot_id: str, agent_id: str, version: int = 1) -> AgentConfigSnapshot: + return AgentConfigSnapshot( + id=snapshot_id, + tenant_id="tenant-1", + agent_id=agent_id, + version=version, + config_snapshot=AgentSoulConfig(), ) - clone = Mock(return_value=(SimpleNamespace(id="target-agent"), "target-snapshot")) - monkeypatch.setattr(WorkflowAgentPublishService, "_clone_inline_graph_binding_for_node", clone) + + +def test_inline_binding_from_another_node_is_cloned(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: + source_agent = _inline_agent(agent_id="source-agent", workflow_id="workflow-1", node_id="source-node") + source_snapshot = _snapshot(snapshot_id="source-snapshot", agent_id=source_agent.id) + sqlite_session.add_all([source_agent, source_snapshot]) + sqlite_session.commit() + target_agent = _inline_agent(agent_id="target-agent", workflow_id="workflow-1", node_id="pasted-node") + target_snapshot = _snapshot(snapshot_id="target-snapshot", agent_id=target_agent.id) + clone = Mock(return_value=(target_agent, target_snapshot)) + monkeypatch.setattr(AgentDslService, "clone_inline_binding_for_node", clone) WorkflowAgentPublishService._sync_agent_binding_for_node( - session=session, - draft_workflow=draft_workflow, + session=sqlite_session, + draft_workflow=_workflow(), node_id="pasted-node", node_data={"agent_task": "Summarize the input"}, node_binding={ "binding_type": WorkflowAgentBindingType.INLINE_AGENT.value, - "agent_id": "source-agent", - "current_snapshot_id": "source-snapshot", + "agent_id": source_agent.id, + "current_snapshot_id": source_snapshot.id, }, existing_binding=None, account_id="account-1", ) + sqlite_session.flush() clone.assert_called_once() - binding = session.add.call_args.args[0] - assert isinstance(binding, WorkflowAgentNodeBinding) + binding = sqlite_session.scalar( + select(WorkflowAgentNodeBinding).where(WorkflowAgentNodeBinding.node_id == "pasted-node") + ) + assert binding is not None assert binding.agent_id == "target-agent" assert binding.current_snapshot_id == "target-snapshot" assert binding.node_job_config.workflow_prompt == "Summarize the input" @@ -103,8 +134,9 @@ def test_draft_sync_resolves_roster_agents() -> None: assert {call.args[0].agent_id for call in session.add.call_args_list} == {"agent-a", "agent-b"} -def test_restore_replaces_bindings_and_returns_only_replaced_inline_agent() -> None: +def test_restore_replaces_bindings_and_returns_only_replaced_inline_agent(sqlite_session: Session) -> None: existing_inline = WorkflowAgentNodeBinding( + id="existing-inline", tenant_id="tenant-1", app_id="app-1", workflow_id="draft-workflow", @@ -117,6 +149,7 @@ def test_restore_replaces_bindings_and_returns_only_replaced_inline_agent() -> N created_by="account-1", ) existing_roster = WorkflowAgentNodeBinding( + id="existing-roster", tenant_id="tenant-1", app_id="app-1", workflow_id="draft-workflow", @@ -129,6 +162,7 @@ def test_restore_replaces_bindings_and_returns_only_replaced_inline_agent() -> N created_by="account-1", ) source = WorkflowAgentNodeBinding( + id="source-roster", tenant_id="tenant-1", app_id="app-1", workflow_id="published-workflow", @@ -140,30 +174,34 @@ def test_restore_replaces_bindings_and_returns_only_replaced_inline_agent() -> N node_job_config={"workflow_prompt": "Use the roster agent"}, created_by="account-1", ) - session = Mock() - session.scalars.side_effect = [ - SimpleNamespace(all=lambda: [existing_inline, existing_roster]), - SimpleNamespace(all=lambda: [source]), - ] - session.scalar.return_value = SimpleNamespace( + roster_agent = Agent( id="roster-agent", + tenant_id="tenant-1", + name="Roster Agent", scope=AgentScope.ROSTER, + source=AgentSource.ROSTER, + status=AgentStatus.ACTIVE, + app_id="roster-app", active_config_snapshot_id="published-snapshot", ) + sqlite_session.add_all([existing_inline, existing_roster, source, roster_agent]) + sqlite_session.commit() retirement_candidates = WorkflowAgentPublishService.restore_agent_node_bindings_to_draft( - session=session, + session=sqlite_session, source_workflow=_workflow(workflow_id="published-workflow", version="2026-07-13 00:00:00"), draft_workflow=_workflow(workflow_id="draft-workflow"), account_id="account-2", ) - assert {item.args[0].agent_id for item in session.delete.call_args_list} == { - "old-inline-agent", - "old-roster-agent", - } - restored = session.add.call_args.args[0] - assert isinstance(restored, WorkflowAgentNodeBinding) - assert restored.workflow_id == "draft-workflow" + assert sqlite_session.get(WorkflowAgentNodeBinding, existing_inline.id) is None + assert sqlite_session.get(WorkflowAgentNodeBinding, existing_roster.id) is None + restored = sqlite_session.scalar( + select(WorkflowAgentNodeBinding).where( + WorkflowAgentNodeBinding.workflow_id == "draft-workflow", + WorkflowAgentNodeBinding.node_id == "agent-node", + ) + ) + assert restored is not None assert restored.workflow_version == Workflow.VERSION_DRAFT assert restored.agent_id == "roster-agent" assert restored.current_snapshot_id == "published-snapshot" @@ -284,6 +322,7 @@ def test_publish_binding_copy_keeps_previous_published_owner( draft_workflow=draft_workflow, published_workflow=published_workflow, ) + sqlite_session.flush() assert result is True assert sqlite_session.get(WorkflowAgentNodeBinding, previous_inline_binding.id) is previous_inline_binding @@ -299,55 +338,50 @@ def test_publish_binding_copy_keeps_previous_published_owner( assert copied.current_snapshot_id == "draft-inline-snapshot" -def test_inline_binding_reuses_existing_node_owned_agent(monkeypatch: pytest.MonkeyPatch) -> None: - session = Mock() - draft_workflow = _workflow() +def test_inline_binding_reuses_existing_node_owned_agent(sqlite_session: Session) -> None: + existing_agent = _inline_agent(agent_id="existing-agent", workflow_id="workflow-1", node_id="pasted-node") + existing_snapshot = _snapshot(snapshot_id="existing-snapshot", agent_id=existing_agent.id) existing_binding = WorkflowAgentNodeBinding( + id="existing-binding", tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", workflow_version=Workflow.VERSION_DRAFT, node_id="pasted-node", binding_type=WorkflowAgentBindingType.INLINE_AGENT, - agent_id="existing-agent", - current_snapshot_id="existing-snapshot", + agent_id=existing_agent.id, + current_snapshot_id=existing_snapshot.id, node_job_config={}, created_by="account-1", ) - existing_agent = SimpleNamespace(id="existing-agent") - monkeypatch.setattr( - WorkflowAgentPublishService, - "_resolve_inline_agent_graph_binding", - Mock(side_effect=_InlineAgentOwnershipError("source belongs to another node")), - ) - monkeypatch.setattr( - WorkflowAgentPublishService, - "_resolve_existing_inline_binding_agent", - Mock(return_value=existing_agent), - ) - clone = Mock() - monkeypatch.setattr(WorkflowAgentPublishService, "_clone_inline_graph_binding_for_node", clone) + sqlite_session.add_all([existing_agent, existing_snapshot, existing_binding]) + sqlite_session.commit() WorkflowAgentPublishService._sync_agent_binding_for_node( - session=session, - draft_workflow=draft_workflow, + session=sqlite_session, + draft_workflow=_workflow(), node_id="pasted-node", node_data={"agent_task": "Summarize"}, node_binding={ "binding_type": WorkflowAgentBindingType.INLINE_AGENT.value, - "agent_id": "source-agent", - "current_snapshot_id": "source-snapshot", + "agent_id": "unavailable-source-agent", + "current_snapshot_id": "unavailable-source-snapshot", }, existing_binding=existing_binding, account_id="account-1", ) + sqlite_session.flush() - assert existing_binding.agent_id == "existing-agent" - assert existing_binding.current_snapshot_id == "existing-snapshot" - clone.assert_not_called() + stored = sqlite_session.get(WorkflowAgentNodeBinding, existing_binding.id) + assert stored is not None + assert stored.agent_id == "existing-agent" + assert stored.current_snapshot_id == "existing-snapshot" + assert stored.node_job_config.workflow_prompt == "Summarize" -def test_resolve_existing_inline_binding_agent_returns_valid_agent_or_none(monkeypatch: pytest.MonkeyPatch) -> None: +def test_resolve_existing_inline_binding_agent_returns_valid_agent_or_none( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: binding = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", @@ -360,13 +394,13 @@ def test_resolve_existing_inline_binding_agent_returns_valid_agent_or_none(monke node_job_config={}, created_by="account-1", ) - resolved = SimpleNamespace(id="agent-1") + resolved = _inline_agent(agent_id="agent-1", workflow_id="workflow-1", node_id="node-1") resolver = Mock(return_value=resolved) monkeypatch.setattr(WorkflowAgentPublishService, "_resolve_inline_agent_graph_binding", resolver) assert ( WorkflowAgentPublishService._resolve_existing_inline_binding_agent( - session=Mock(), + session=unbound_session, draft_workflow=_workflow(), node_id="node-1", existing_binding=binding, @@ -377,7 +411,7 @@ def test_resolve_existing_inline_binding_agent_returns_valid_agent_or_none(monke resolver.side_effect = ValueError("stale") assert ( WorkflowAgentPublishService._resolve_existing_inline_binding_agent( - session=Mock(), + session=unbound_session, draft_workflow=_workflow(), node_id="node-1", existing_binding=binding, @@ -386,30 +420,42 @@ def test_resolve_existing_inline_binding_agent_returns_valid_agent_or_none(monke ) -def test_resolve_roster_binding_rejects_unpublished_agent() -> None: - session = Mock() - session.scalar.return_value = None +def test_resolve_roster_binding_rejects_unpublished_agent(sqlite_session: Session) -> None: + sqlite_session.add( + Agent( + id="decoy-agent", + tenant_id="tenant-1", + name="Decoy", + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + app_id="decoy-app", + ) + ) + sqlite_session.commit() with pytest.raises(ValueError, match="unavailable or unpublished roster agent"): WorkflowAgentPublishService._resolve_roster_agent_graph_binding( - session=session, + session=sqlite_session, draft_workflow=_workflow(), node_id="agent-node", agent_id="agent-1", ) -def test_clone_inline_graph_binding_for_node_clones_source(monkeypatch: pytest.MonkeyPatch) -> None: - session = Mock() - source_agent = SimpleNamespace(id="source-agent") - source_snapshot = SimpleNamespace(id="source-snapshot") - session.scalar.side_effect = [source_agent, source_snapshot] - target_agent = SimpleNamespace(id="target-agent") - target_snapshot = SimpleNamespace(id="target-snapshot") +def test_clone_inline_graph_binding_for_node_clones_source( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + source_agent = _inline_agent(agent_id="source-agent", workflow_id="source-workflow", node_id="source-node") + source_snapshot = _snapshot(snapshot_id="source-snapshot", agent_id=source_agent.id) + sqlite_session.add_all([source_agent, source_snapshot]) + sqlite_session.commit() + target_agent = _inline_agent(agent_id="target-agent", workflow_id="workflow-1", node_id="target-node") + target_snapshot = _snapshot(snapshot_id="target-snapshot", agent_id=target_agent.id) clone = Mock(return_value=(target_agent, target_snapshot)) monkeypatch.setattr(AgentDslService, "clone_inline_binding_for_node", clone) result = WorkflowAgentPublishService._clone_inline_graph_binding_for_node( - session=session, + session=sqlite_session, draft_workflow=_workflow(), node_id="target-node", source_agent_id="source-agent", @@ -427,14 +473,17 @@ def test_clone_inline_graph_binding_for_node_clones_source(monkeypatch: pytest.M ) -@pytest.mark.parametrize("scalar_results", [[None], [SimpleNamespace(id="source-agent"), None]]) -def test_clone_inline_graph_binding_for_node_rejects_missing_source(scalar_results: list[object | None]) -> None: - session = Mock() - session.scalar.side_effect = scalar_results +@pytest.mark.parametrize("persist_source_agent", [False, True]) +def test_clone_inline_graph_binding_for_node_rejects_missing_source( + sqlite_session: Session, persist_source_agent: bool +) -> None: + if persist_source_agent: + sqlite_session.add(_inline_agent(agent_id="source-agent", workflow_id="source-workflow", node_id="source-node")) + sqlite_session.commit() with pytest.raises(ValueError, match="unavailable inline agent|missing inline agent config snapshot"): WorkflowAgentPublishService._clone_inline_graph_binding_for_node( - session=session, + session=sqlite_session, draft_workflow=_workflow(), node_id="target-node", source_agent_id="source-agent", @@ -443,37 +492,45 @@ def test_clone_inline_graph_binding_for_node_rejects_missing_source(scalar_resul ) -def test_restore_clones_inline_binding_owned_by_published_workflow(monkeypatch: pytest.MonkeyPatch) -> None: +def test_restore_clones_inline_binding_owned_by_published_workflow( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + source_agent = _inline_agent(agent_id="published-agent", workflow_id="published-workflow", node_id="agent-node") + source_snapshot = _snapshot(snapshot_id="published-snapshot", agent_id=source_agent.id) source = WorkflowAgentNodeBinding( + id="published-binding", tenant_id="tenant-1", app_id="app-1", workflow_id="published-workflow", workflow_version="published", node_id="agent-node", binding_type=WorkflowAgentBindingType.INLINE_AGENT, - agent_id="published-agent", - current_snapshot_id="published-snapshot", + agent_id=source_agent.id, + current_snapshot_id=source_snapshot.id, node_job_config={"workflow_prompt": "work"}, created_by="account-1", ) - session = Mock() - session.scalars.side_effect = [SimpleNamespace(all=lambda: []), SimpleNamespace(all=lambda: [source])] - monkeypatch.setattr( - WorkflowAgentPublishService, - "_resolve_inline_agent_graph_binding", - Mock(side_effect=ValueError("owned by published workflow")), - ) - clone = Mock(return_value=(SimpleNamespace(id="draft-agent"), "draft-snapshot")) - monkeypatch.setattr(WorkflowAgentPublishService, "_clone_inline_graph_binding_for_node", clone) + sqlite_session.add_all([source_agent, source_snapshot, source]) + sqlite_session.commit() + target_agent = _inline_agent(agent_id="draft-agent", workflow_id="draft-workflow", node_id="agent-node") + target_snapshot = _snapshot(snapshot_id="draft-snapshot", agent_id=target_agent.id) + clone = Mock(return_value=(target_agent, target_snapshot)) + monkeypatch.setattr(AgentDslService, "clone_inline_binding_for_node", clone) WorkflowAgentPublishService.restore_agent_node_bindings_to_draft( - session=session, + session=sqlite_session, source_workflow=_workflow(workflow_id="published-workflow", version="published"), draft_workflow=_workflow(workflow_id="draft-workflow"), account_id="account-2", ) clone.assert_called_once() - restored = session.add.call_args.args[0] + restored = sqlite_session.scalar( + select(WorkflowAgentNodeBinding).where( + WorkflowAgentNodeBinding.workflow_id == "draft-workflow", + WorkflowAgentNodeBinding.workflow_version == Workflow.VERSION_DRAFT, + ) + ) + assert restored is not None assert restored.agent_id == "draft-agent" assert restored.current_snapshot_id == "draft-snapshot" diff --git a/api/tests/unit_tests/services/agent/test_workspace_service.py b/api/tests/unit_tests/services/agent/test_workspace_service.py index 733c09896de..2a239c6e726 100644 --- a/api/tests/unit_tests/services/agent/test_workspace_service.py +++ b/api/tests/unit_tests/services/agent/test_workspace_service.py @@ -105,11 +105,6 @@ def test_workspace_client_honors_the_configured_snapshot_timeout(monkeypatch: py assert client._timeout == 123.5 -@pytest.mark.parametrize( - "sqlite_session", - [(AgentHomeSnapshot, AgentWorkspace, AgentWorkspaceBinding)], - indirect=True, -) def test_create_binding_success_persists_new_workspace_and_binding( monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ) -> None: @@ -148,11 +143,6 @@ def test_create_binding_success_persists_new_workspace_and_binding( assert request.home_snapshot_ref == "home-ref" -@pytest.mark.parametrize( - "sqlite_session", - [(AgentHomeSnapshot, AgentWorkspace, AgentWorkspaceBinding)], - indirect=True, -) def test_create_binding_without_home_snapshot_uses_backend_default( monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ) -> None: @@ -176,11 +166,6 @@ def test_create_binding_without_home_snapshot_uses_backend_default( assert request.home_snapshot_ref is None -@pytest.mark.parametrize( - "sqlite_session", - [(AgentHomeSnapshot, AgentWorkspace, AgentWorkspaceBinding)], - indirect=True, -) def test_create_binding_rejects_missing_explicit_home_snapshot_before_backend_call( monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ) -> None: @@ -200,11 +185,6 @@ def test_create_binding_rejects_missing_explicit_home_snapshot_before_backend_ca client.create_execution_binding_sync.assert_not_called() -@pytest.mark.parametrize( - "sqlite_session", - [(AgentHomeSnapshot, AgentWorkspace, AgentWorkspaceBinding)], - indirect=True, -) def test_create_second_binding_reuses_existing_workspace( monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ) -> None: @@ -243,7 +223,6 @@ def test_create_second_binding_reuses_existing_workspace( assert request.workspace_id == workspace.id -@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) def test_get_active_binding_resolves_exact_participant(sqlite_session: Session) -> None: conversation_workspace = _workspace(workspace_id="workspace-conversation") build_workspace = _workspace( @@ -274,7 +253,6 @@ def test_get_active_binding_resolves_exact_participant(sqlite_session: Session) assert resolved.id == conversation_binding.id -@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) def test_get_active_binding_rejects_wrong_owner(sqlite_session: Session) -> None: build_workspace = _workspace( workspace_id="workspace-build", @@ -299,7 +277,6 @@ def test_get_active_binding_rejects_wrong_owner(sqlite_session: Session) -> None assert resolved is None -@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) def test_retire_non_final_binding_keeps_workspace_active(sqlite_session: Session) -> None: binding = _binding() other_binding = _binding(binding_id="binding-2", agent_id="agent-2") @@ -320,7 +297,6 @@ def test_retire_non_final_binding_keeps_workspace_active(sqlite_session: Session assert other_binding.status is AgentWorkingResourceStatus.ACTIVE -@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) def test_retire_final_binding_retires_workspace(sqlite_session: Session) -> None: binding = _binding() workspace = _workspace() @@ -335,15 +311,14 @@ def test_retire_final_binding_retires_workspace(sqlite_session: Session) -> None assert workspace.retired_at == binding.retired_at -def test_retire_workspace_retires_all_active_bindings() -> None: +def test_retire_workspace_retires_all_active_bindings(sqlite_session: Session) -> None: workspace = _workspace() bindings = [_binding(), _binding(binding_id="binding-2", agent_id="agent-2")] - session = MagicMock() - session.scalar.return_value = workspace - session.scalars.return_value.all.return_value = bindings + sqlite_session.add_all([workspace, *bindings]) + sqlite_session.flush() retired_id = AgentWorkspaceService.retire_workspace( - session=session, + session=sqlite_session, tenant_id="tenant-1", workspace_id=workspace.id, ) @@ -354,7 +329,6 @@ def test_retire_workspace_retires_all_active_bindings() -> None: assert all(binding.retired_at == workspace.retired_at for binding in bindings) -@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) def test_retire_all_for_app_retires_only_active_workspaces_for_that_app(sqlite_session: Session) -> None: active = _workspace(workspace_id="workspace-active", owner_id="conversation-active") already_retired = _workspace( @@ -390,7 +364,6 @@ def test_retire_all_for_app_retires_only_active_workspaces_for_that_app(sqlite_s assert other_binding.status is AgentWorkingResourceStatus.ACTIVE -@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) def test_collect_binding_without_retired_workspace_destroys_binding_only( monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ) -> None: @@ -414,7 +387,6 @@ def test_collect_binding_without_retired_workspace_destroys_binding_only( assert sqlite_session.get(AgentWorkspace, workspace.id) is not None -@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) def test_collect_workspace_destroys_workspace_then_remaining_bindings( monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ) -> None: diff --git a/api/tests/unit_tests/services/test_app_generate_service.py b/api/tests/unit_tests/services/test_app_generate_service.py index b0bf1a2fd4e..613c6c2000e 100644 --- a/api/tests/unit_tests/services/test_app_generate_service.py +++ b/api/tests/unit_tests/services/test_app_generate_service.py @@ -21,6 +21,7 @@ from unittest.mock import MagicMock import pytest from pytest_mock import MockerFixture +from sqlalchemy.orm import Session import services.app_generate_service as ags_module from core.app.entities.app_invoke_entities import InvokeFrom @@ -79,6 +80,12 @@ def _make_user() -> MagicMock: return user +class _RealSessionTest: + @pytest.fixture(autouse=True) + def _bind_unbound_session(self, unbound_session: Session) -> None: + self.session = unbound_session + + def _make_workflow(*, workflow_id: str = "workflow-id", created_by: str = "owner-id") -> MagicMock: workflow = MagicMock() workflow.id = workflow_id @@ -251,7 +258,7 @@ class TestGetMaxActiveRequests: # --------------------------------------------------------------------------- # generate – every AppMode branch # --------------------------------------------------------------------------- -class TestGenerate: +class TestGenerate(_RealSessionTest): """Tests for AppGenerateService.generate covering each mode.""" @pytest.fixture(autouse=True) @@ -280,7 +287,7 @@ class TestGenerate: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=self.session, ) assert result == {"result": "ok"} gen_spy.assert_called_once() @@ -301,7 +308,7 @@ class TestGenerate: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=self.session, ) assert result == {"result": "agent"} gen_spy.assert_called_once() @@ -317,7 +324,7 @@ class TestGenerate: side_effect=lambda x: x, ) app = _make_app(AppMode.CHAT, is_agent=True) - session = MagicMock() + session = self.session result = AppGenerateService.generate( app_model=app, user=_make_user(), @@ -340,7 +347,7 @@ class TestGenerate: "services.app_generate_service.AgentAppGenerator.convert_to_event_stream", side_effect=lambda x: x, ) - session = MagicMock() + session = self.session result = AppGenerateService.generate( app_model=_make_app(AppMode.AGENT), @@ -371,7 +378,7 @@ class TestGenerate: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=self.session, ) assert result == {"result": "chat"} gen_spy.assert_called_once() @@ -391,7 +398,7 @@ class TestGenerate: side_effect=lambda x: x, ) - session = MagicMock() + session = self.session result = AppGenerateService.generate( app_model=_make_app(AppMode.ADVANCED_CHAT), user=_make_user(), @@ -430,7 +437,7 @@ class TestGenerate: args={"workflow_id": None, "query": "hi", "inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=True, - session=MagicMock(), + session=self.session, ) # In streaming mode it should go through retrieve_events, not generate gen_instance.retrieve_events.assert_called_once() @@ -453,7 +460,7 @@ class TestGenerate: side_effect=lambda x: x, ) - session = MagicMock() + session = self.session result = AppGenerateService.generate( app_model=_make_app(AppMode.WORKFLOW), user=_make_user(), @@ -492,7 +499,7 @@ class TestGenerate: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=True, - session=MagicMock(), + session=self.session, ) retrieve_spy.assert_called_once() # Dispatch is gated on subscribe; simulate the SSE layer entering the @@ -511,14 +518,14 @@ class TestGenerate: args={}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=self.session, ) # --------------------------------------------------------------------------- # generate – billing / quota # --------------------------------------------------------------------------- -class TestGenerateBilling: +class TestGenerateBilling(_RealSessionTest): @pytest.fixture(autouse=True) def _common(self, mocker: MockerFixture): mocker.patch("services.app_generate_service.RateLimit", _DummyRateLimit) @@ -549,7 +556,7 @@ class TestGenerateBilling: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=self.session, ) reserve_mock.assert_called_once_with(QuotaType.WORKFLOW, "tenant-id") quota_charge.commit.assert_called_once() @@ -573,7 +580,7 @@ class TestGenerateBilling: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=self.session, ) def test_exception_refunds_quota_and_exits_rate_limit( @@ -601,7 +608,7 @@ class TestGenerateBilling: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=self.session, ) quota_charge.refund.assert_called_once() @@ -633,7 +640,7 @@ class TestGenerateBilling: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=self.session, ) # exit is called in finally block for non-streaming assert exit_calls == ["dummy-request-id"] @@ -664,7 +671,7 @@ class TestGenerateBilling: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=self.session, ) quota_charge.refund.assert_called_once() @@ -698,7 +705,7 @@ class TestGenerateBilling: args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=True, - session=MagicMock(), + session=self.session, ) quota_charge.refund.assert_called_once() @@ -708,14 +715,16 @@ class TestGenerateBilling: # --------------------------------------------------------------------------- # _get_workflow # --------------------------------------------------------------------------- -class TestGetWorkflow: +class TestGetWorkflow(_RealSessionTest): def test_debugger_fetches_draft(self, mocker: MockerFixture): draft_wf = _make_workflow() ws = MagicMock() ws.get_draft_workflow.return_value = draft_wf mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) - result = AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER, session=MagicMock()) + result = AppGenerateService._get_workflow( + _make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER, session=self.session + ) assert result is draft_wf ws.get_draft_workflow.assert_called_once() @@ -725,7 +734,7 @@ class TestGetWorkflow: mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) with pytest.raises(ValueError, match="Workflow not initialized"): - AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER, session=MagicMock()) + AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER, session=self.session) def test_non_debugger_fetches_published(self, mocker: MockerFixture): pub_wf = _make_workflow() @@ -734,7 +743,7 @@ class TestGetWorkflow: mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) result = AppGenerateService._get_workflow( - _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, session=MagicMock() + _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, session=self.session ) assert result is pub_wf ws.get_published_workflow.assert_called_once() @@ -745,7 +754,7 @@ class TestGetWorkflow: mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) with pytest.raises(ValueError, match="Workflow not published"): - AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, session=MagicMock()) + AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, session=self.session) def test_specific_workflow_id_valid_uuid(self, mocker: MockerFixture): valid_uuid = str(uuid.uuid4()) @@ -758,7 +767,7 @@ class TestGetWorkflow: _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, workflow_id=valid_uuid, - session=MagicMock(), + session=self.session, ) assert result is specific_wf ws.get_published_workflow_by_id.assert_called_once() @@ -772,7 +781,7 @@ class TestGetWorkflow: _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, workflow_id="not-a-uuid", - session=MagicMock(), + session=self.session, ) def test_specific_workflow_id_not_found(self, mocker: MockerFixture): @@ -786,14 +795,14 @@ class TestGetWorkflow: _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, workflow_id=valid_uuid, - session=MagicMock(), + session=self.session, ) # --------------------------------------------------------------------------- # generate_single_iteration # --------------------------------------------------------------------------- -class TestGenerateSingleIteration: +class TestGenerateSingleIteration(_RealSessionTest): def test_advanced_chat_mode(self, mocker: MockerFixture): workflow = _make_workflow() mocker.patch.object(AppGenerateService, "_get_workflow", return_value=workflow) @@ -806,7 +815,7 @@ class TestGenerateSingleIteration: return_value={"event": "iteration"}, ) app = _make_app(AppMode.ADVANCED_CHAT) - session = MagicMock() + session = self.session result = AppGenerateService.generate_single_iteration( app_model=app, user=_make_user(), @@ -830,7 +839,7 @@ class TestGenerateSingleIteration: return_value={"event": "wf-iteration"}, ) app = _make_app(AppMode.WORKFLOW) - session = MagicMock() + session = self.session result = AppGenerateService.generate_single_iteration( app_model=app, user=_make_user(), @@ -846,14 +855,14 @@ class TestGenerateSingleIteration: app = _make_app(AppMode.CHAT) with pytest.raises(ValueError, match="Invalid app mode"): AppGenerateService.generate_single_iteration( - app_model=app, user=_make_user(), node_id="n1", args={}, session=MagicMock() + app_model=app, user=_make_user(), node_id="n1", args={}, session=self.session ) # --------------------------------------------------------------------------- # generate_single_loop # --------------------------------------------------------------------------- -class TestGenerateSingleLoop: +class TestGenerateSingleLoop(_RealSessionTest): def test_advanced_chat_mode(self, mocker: MockerFixture): workflow = _make_workflow() mocker.patch.object(AppGenerateService, "_get_workflow", return_value=workflow) @@ -866,7 +875,7 @@ class TestGenerateSingleLoop: return_value={"event": "loop"}, ) app = _make_app(AppMode.ADVANCED_CHAT) - session = MagicMock() + session = self.session result = AppGenerateService.generate_single_loop( app_model=app, user=_make_user(), @@ -890,7 +899,7 @@ class TestGenerateSingleLoop: return_value={"event": "wf-loop"}, ) app = _make_app(AppMode.WORKFLOW) - session = MagicMock() + session = self.session result = AppGenerateService.generate_single_loop( app_model=app, user=_make_user(), @@ -906,20 +915,20 @@ class TestGenerateSingleLoop: app = _make_app(AppMode.COMPLETION) with pytest.raises(ValueError, match="Invalid app mode"): AppGenerateService.generate_single_loop( - app_model=app, user=_make_user(), node_id="n1", args=MagicMock(), session=MagicMock() + app_model=app, user=_make_user(), node_id="n1", args=MagicMock(), session=self.session ) # --------------------------------------------------------------------------- # generate_more_like_this # --------------------------------------------------------------------------- -class TestGenerateMoreLikeThis: +class TestGenerateMoreLikeThis(_RealSessionTest): def test_delegates_to_completion_generator(self, mocker: MockerFixture): gen_spy = mocker.patch( "services.app_generate_service.CompletionAppGenerator.generate_more_like_this", return_value={"result": "similar"}, ) - session = MagicMock() + session = self.session result = AppGenerateService.generate_more_like_this( app_model=_make_app(AppMode.COMPLETION), user=_make_user(),