refactor: agent draft (#37356)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
zyssyz123
2026-06-12 03:46:21 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent 72faca2592
commit e32a732812
17 changed files with 554 additions and 75 deletions
+32 -9
View File
@@ -2,7 +2,7 @@ import logging
import uuid
from typing import Any
from sqlalchemy import func, select
from sqlalchemy import func, or_, select
from sqlalchemy.exc import IntegrityError
from extensions.ext_database import db
@@ -74,10 +74,15 @@ class AgentComposerService:
return cls._empty_workflow_state(app_id=app_id, workflow_id=workflow.id, node_id=node_id)
agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id)
version_id = (
agent.active_config_snapshot_id
if agent and binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT
else binding.current_snapshot_id
)
version = cls._get_version_if_present(
tenant_id=tenant_id,
agent_id=agent.id if agent else None,
version_id=binding.current_snapshot_id,
version_id=version_id,
)
return cls._serialize_workflow_state(binding=binding, agent=agent, version=version)
@@ -129,10 +134,15 @@ class AgentComposerService:
db.session.commit()
agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id)
version_id = (
agent.active_config_snapshot_id
if agent and binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT
else binding.current_snapshot_id
)
version = cls._get_version_if_present(
tenant_id=tenant_id,
agent_id=agent.id if agent else None,
version_id=binding.current_snapshot_id,
version_id=version_id,
)
state = cls._serialize_workflow_state(binding=binding, agent=agent, version=version)
state["validation"] = cls.collect_validation_findings(tenant_id=tenant_id, payload=payload)
@@ -489,11 +499,26 @@ class AgentComposerService:
@classmethod
def calculate_impact(cls, *, tenant_id: str, current_snapshot_id: str) -> dict[str, Any]:
snapshot = db.session.scalar(
select(AgentConfigSnapshot)
.where(
AgentConfigSnapshot.tenant_id == tenant_id,
AgentConfigSnapshot.id == current_snapshot_id,
)
.limit(1)
)
agent_id = snapshot.agent_id if snapshot else None
predicates = [WorkflowAgentNodeBinding.current_snapshot_id == current_snapshot_id]
if agent_id:
predicates.append(
(WorkflowAgentNodeBinding.agent_id == agent_id)
& (WorkflowAgentNodeBinding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT)
)
bindings = list(
db.session.scalars(
select(WorkflowAgentNodeBinding).where(
WorkflowAgentNodeBinding.tenant_id == tenant_id,
WorkflowAgentNodeBinding.current_snapshot_id == current_snapshot_id,
or_(*predicates),
)
).all()
)
@@ -1003,7 +1028,7 @@ class AgentComposerService:
"id": binding.id,
"binding_type": binding.binding_type.value,
"agent_id": binding.agent_id,
"current_snapshot_id": binding.current_snapshot_id,
"current_snapshot_id": version.id if version else binding.current_snapshot_id,
"workflow_id": binding.workflow_id,
"node_id": binding.node_id,
},
@@ -1022,10 +1047,8 @@ class AgentComposerService:
# this is the same list (so callers don't need to special-case).
"effective_declared_outputs": cls._serialize_effective_outputs(cls._declared_outputs_from_binding(binding)),
"save_options": save_options,
"impact_summary": cls.calculate_impact(
tenant_id=binding.tenant_id, current_snapshot_id=binding.current_snapshot_id
)
if binding.current_snapshot_id
"impact_summary": cls.calculate_impact(tenant_id=binding.tenant_id, current_snapshot_id=version.id)
if version
else None,
}
+93 -38
View File
@@ -18,6 +18,7 @@ from models.agent import (
WorkflowAgentNodeBinding,
)
from models.agent_config_entities import AgentSoulConfig
from models.enums import AppStatus
from models.model import App
from models.workflow import Workflow
from services.agent.composer_validator import ComposerConfigValidator
@@ -37,6 +38,7 @@ class AgentReferencingWorkflow(TypedDict):
app_name: str
app_mode: str
workflow_id: str
workflow_version: str
node_ids: list[str]
@@ -45,11 +47,17 @@ class AgentRosterService:
self._session = session
@staticmethod
def serialize_agent(agent: Agent, active_version: AgentConfigSnapshot | None = None) -> dict[str, Any]:
def serialize_agent(
agent: Agent,
active_version: AgentConfigSnapshot | None = None,
published_references: list[AgentReferencingWorkflow] | None = None,
) -> dict[str, Any]:
published_references = published_references or []
return {
"id": agent.id,
"name": agent.name,
"description": agent.description,
"role": agent.role or "",
"icon_type": agent.icon_type.value if agent.icon_type else None,
"icon": agent.icon,
"icon_background": agent.icon_background,
@@ -68,6 +76,9 @@ class AgentRosterService:
"archived_at": to_timestamp(agent.archived_at),
"created_at": to_timestamp(agent.created_at),
"updated_at": to_timestamp(agent.updated_at),
"published_reference_count": len(published_references),
"published_node_reference_count": sum(len(item["node_ids"]) for item in published_references),
"published_references": published_references,
}
@staticmethod
@@ -104,13 +115,23 @@ class AgentRosterService:
versions_by_id = self._load_versions_by_id(
[agent.active_config_snapshot_id for agent in agents if agent.active_config_snapshot_id]
)
published_references_by_agent_id = self._load_published_references_by_agent_id(
tenant_id=tenant_id,
agent_ids=[agent.id for agent in agents],
)
data = []
for agent in agents:
active_version = (
versions_by_id.get(agent.active_config_snapshot_id) if agent.active_config_snapshot_id else None
)
data.append(self.serialize_agent(agent, active_version))
data.append(
self.serialize_agent(
agent,
active_version,
published_references_by_agent_id.get(agent.id, []),
)
)
return {
"data": data,
@@ -170,6 +191,7 @@ class AgentRosterService:
tenant_id=tenant_id,
name=payload.name,
description=payload.description,
role=payload.role,
icon_type=payload.icon_type,
icon=payload.icon,
icon_background=payload.icon_background,
@@ -241,6 +263,7 @@ class AgentRosterService:
tenant_id=tenant_id,
name=name,
description=description,
role="",
icon_type=icon_type,
icon=icon,
icon_background=icon_background,
@@ -306,48 +329,18 @@ class AgentRosterService:
if agent is None:
return []
bindings = self._session.scalars(
select(WorkflowAgentNodeBinding).where(
WorkflowAgentNodeBinding.tenant_id == tenant_id,
WorkflowAgentNodeBinding.agent_id == agent.id,
WorkflowAgentNodeBinding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT,
)
).all()
if not bindings:
return []
# Collapse the per-version / per-node rows into one entry per workflow app.
node_ids_by_workflow: dict[tuple[str, str], set[str]] = {}
for binding in bindings:
node_ids_by_workflow.setdefault((binding.app_id, binding.workflow_id), set()).add(binding.node_id)
referenced_app_ids = {workflow_app_id for workflow_app_id, _ in node_ids_by_workflow}
apps = {app.id: app for app in self._session.scalars(select(App).where(App.id.in_(referenced_app_ids))).all()}
result: list[AgentReferencingWorkflow] = []
for (workflow_app_id, workflow_id), node_ids in node_ids_by_workflow.items():
app = apps.get(workflow_app_id)
if app is None:
# Orphaned binding (workflow app deleted): skip rather than 500.
continue
result.append(
AgentReferencingWorkflow(
app_id=workflow_app_id,
app_name=app.name,
app_mode=str(app.mode),
workflow_id=workflow_id,
node_ids=sorted(node_ids),
)
)
result.sort(key=lambda item: item["app_name"].lower())
return result
return self._load_published_references_by_agent_id(tenant_id=tenant_id, agent_ids=[agent.id]).get(agent.id, [])
def get_roster_agent_detail(self, *, tenant_id: str, agent_id: str) -> dict[str, Any]:
agent = self._get_agent(tenant_id=tenant_id, agent_id=agent_id, roster_only=True)
active_version = self._get_version(
tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id
)
return self.serialize_agent(agent, active_version)
published_references_by_agent_id = self._load_published_references_by_agent_id(
tenant_id=tenant_id,
agent_ids=[agent.id],
)
return self.serialize_agent(agent, active_version, published_references_by_agent_id.get(agent.id, []))
def update_roster_agent(
self, *, tenant_id: str, agent_id: str, account_id: str, payload: RosterAgentUpdatePayload
@@ -450,6 +443,68 @@ class AgentRosterService:
raise AgentVersionNotFoundError()
return version
def _load_published_references_by_agent_id(
self, *, tenant_id: str, agent_ids: list[str]
) -> dict[str, list[AgentReferencingWorkflow]]:
if not agent_ids:
return {}
bindings = list(
self._session.scalars(
select(WorkflowAgentNodeBinding).where(
WorkflowAgentNodeBinding.tenant_id == tenant_id,
WorkflowAgentNodeBinding.agent_id.in_(agent_ids),
WorkflowAgentNodeBinding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT,
WorkflowAgentNodeBinding.workflow_version != Workflow.VERSION_DRAFT,
)
).all()
)
if not bindings:
return {}
app_ids = {binding.app_id for binding in bindings}
apps = {
app.id: app
for app in self._session.scalars(
select(App).where(
App.tenant_id == tenant_id,
App.id.in_(app_ids),
App.status == AppStatus.NORMAL,
)
).all()
}
grouped: dict[str, dict[tuple[str, str], AgentReferencingWorkflow]] = {}
for binding in bindings:
if not binding.agent_id:
continue
app = apps.get(binding.app_id)
if app is None or app.workflow_id != binding.workflow_id:
continue
by_workflow = grouped.setdefault(binding.agent_id, {})
key = (binding.app_id, binding.workflow_id)
item = by_workflow.setdefault(
key,
AgentReferencingWorkflow(
app_id=binding.app_id,
app_name=app.name,
app_mode=str(app.mode),
workflow_id=binding.workflow_id,
workflow_version=binding.workflow_version,
node_ids=[],
),
)
item["node_ids"].append(binding.node_id)
result: dict[str, list[AgentReferencingWorkflow]] = {}
for agent_id, by_workflow in grouped.items():
references = list(by_workflow.values())
for reference in references:
reference["node_ids"] = sorted(set(reference["node_ids"]))
references.sort(key=lambda item: (item["app_name"].lower(), item["workflow_id"]))
result[agent_id] = references
return result
def _load_versions_by_id(self, version_ids: list[str]) -> dict[str, AgentConfigSnapshot]:
if not version_ids:
return {}
+120 -2
View File
@@ -1,10 +1,13 @@
from __future__ import annotations
from collections.abc import Mapping
from typing import Any
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.workflow.nodes.agent_v2.validators import WorkflowAgentNodeValidator
from models.agent import WorkflowAgentNodeBinding
from models.agent import Agent, AgentScope, AgentStatus, WorkflowAgentBindingType, WorkflowAgentNodeBinding
from models.agent_config_entities import WorkflowNodeJobConfig
from models.workflow import Workflow
@@ -12,6 +15,9 @@ from models.workflow import Workflow
class WorkflowAgentPublishService:
"""Validate and freeze Workflow Agent v2 bindings during workflow publish."""
_DRAFT_WORKFLOW_VERSION = Workflow.VERSION_DRAFT
_AGENT_BINDING_KEY = "agent_binding"
@classmethod
def validate_agent_nodes_for_publish(cls, *, session: Session, draft_workflow: Workflow) -> None:
WorkflowAgentNodeValidator.validate_published_workflow(session=session, workflow=draft_workflow)
@@ -20,6 +26,100 @@ class WorkflowAgentPublishService:
def validate_agent_nodes_for_draft_sync(cls, *, session: Session, draft_workflow: Workflow) -> None:
WorkflowAgentNodeValidator.validate_draft_workflow(session=session, workflow=draft_workflow)
@classmethod
def sync_roster_agent_bindings_for_draft(
cls,
*,
session: Session,
draft_workflow: Workflow,
account_id: str,
) -> None:
agent_nodes = dict(WorkflowAgentNodeValidator.iter_agent_v2_nodes(draft_workflow.graph_dict))
existing_bindings = list(
session.scalars(
select(WorkflowAgentNodeBinding).where(
WorkflowAgentNodeBinding.tenant_id == draft_workflow.tenant_id,
WorkflowAgentNodeBinding.app_id == draft_workflow.app_id,
WorkflowAgentNodeBinding.workflow_id == draft_workflow.id,
WorkflowAgentNodeBinding.workflow_version == cls._DRAFT_WORKFLOW_VERSION,
)
).all()
)
existing_by_node_id = {binding.node_id: binding for binding in existing_bindings}
for binding in existing_bindings:
if binding.node_id not in agent_nodes:
session.delete(binding)
for node_id, node_data in agent_nodes.items():
binding_payload = node_data.get(cls._AGENT_BINDING_KEY)
if binding_payload is None:
continue
if not isinstance(binding_payload, Mapping):
raise ValueError(f"Workflow Agent node {node_id} has invalid agent_binding.")
cls._sync_roster_agent_binding_for_node(
session=session,
draft_workflow=draft_workflow,
node_id=node_id,
node_binding=binding_payload,
existing_binding=existing_by_node_id.get(node_id),
account_id=account_id,
)
session.flush()
@classmethod
def _sync_roster_agent_binding_for_node(
cls,
*,
session: Session,
draft_workflow: Workflow,
node_id: str,
node_binding: Mapping[str, Any],
existing_binding: WorkflowAgentNodeBinding | None,
account_id: str,
) -> None:
binding_type = node_binding.get("binding_type")
if binding_type != WorkflowAgentBindingType.ROSTER_AGENT.value:
raise ValueError(f"Workflow Agent node {node_id} only supports roster_agent graph binding.")
agent_id = node_binding.get("agent_id")
if not isinstance(agent_id, str) or not agent_id:
raise ValueError(f"Workflow Agent node {node_id} roster_agent binding requires agent_id.")
agent = session.scalar(
select(Agent)
.where(
Agent.tenant_id == draft_workflow.tenant_id,
Agent.id == agent_id,
Agent.scope == AgentScope.ROSTER,
Agent.status == AgentStatus.ACTIVE,
)
.limit(1)
)
if agent is None:
raise ValueError(f"Workflow Agent node {node_id} references an unavailable roster agent.")
if not agent.active_config_snapshot_id:
raise ValueError(f"Workflow Agent node {node_id} roster agent has no active config snapshot.")
binding = existing_binding
if binding is None:
binding = WorkflowAgentNodeBinding(
tenant_id=draft_workflow.tenant_id,
app_id=draft_workflow.app_id,
workflow_id=draft_workflow.id,
workflow_version=cls._DRAFT_WORKFLOW_VERSION,
node_id=node_id,
node_job_config=WorkflowNodeJobConfig(),
created_by=account_id,
)
session.add(binding)
elif not binding.node_job_config:
binding.node_job_config = WorkflowNodeJobConfig()
binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT
binding.agent_id = agent.id
binding.current_snapshot_id = agent.active_config_snapshot_id
binding.updated_by = account_id
@classmethod
def copy_agent_node_bindings_to_published(
cls,
@@ -43,8 +143,26 @@ class WorkflowAgentPublishService:
WorkflowAgentNodeBinding.node_id.in_(node_ids),
)
).all()
if not bindings:
return
agents_by_id = {
agent.id: agent
for agent in session.scalars(
select(Agent).where(
Agent.tenant_id == draft_workflow.tenant_id,
Agent.id.in_({binding.agent_id for binding in bindings if binding.agent_id}),
)
).all()
}
for binding in bindings:
agent = agents_by_id.get(binding.agent_id) if binding.agent_id else None
current_snapshot_id = (
agent.active_config_snapshot_id
if agent is not None and binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT
else binding.current_snapshot_id
)
copied = WorkflowAgentNodeBinding(
tenant_id=binding.tenant_id,
app_id=binding.app_id,
@@ -53,7 +171,7 @@ class WorkflowAgentPublishService:
node_id=binding.node_id,
binding_type=binding.binding_type,
agent_id=binding.agent_id,
current_snapshot_id=binding.current_snapshot_id,
current_snapshot_id=current_snapshot_id,
node_job_config=WorkflowNodeJobConfig.model_validate(binding.node_job_config_dict),
created_by=binding.created_by,
updated_by=binding.updated_by,
+2
View File
@@ -61,6 +61,7 @@ class ComposerSavePayload(BaseModel):
class RosterAgentCreatePayload(BaseModel):
name: str = Field(min_length=1, max_length=255)
description: str = ""
role: str = Field(default="", max_length=255)
icon_type: AgentIconType | None = None
icon: str | None = Field(default=None, max_length=255)
icon_background: str | None = Field(default=None, max_length=255)
@@ -71,6 +72,7 @@ class RosterAgentCreatePayload(BaseModel):
class RosterAgentUpdatePayload(BaseModel):
name: str | None = Field(default=None, min_length=1, max_length=255)
description: str | None = None
role: str | None = Field(default=None, max_length=255)
icon_type: AgentIconType | None = None
icon: str | None = Field(default=None, max_length=255)
icon_background: str | None = Field(default=None, max_length=255)
+6
View File
@@ -322,6 +322,12 @@ class WorkflowService:
from services.agent.workflow_publish_service import WorkflowAgentPublishService
db.session.flush()
WorkflowAgentPublishService.sync_roster_agent_bindings_for_draft(
session=cast(Session, db.session),
draft_workflow=workflow,
account_id=account.id,
)
WorkflowAgentPublishService.validate_agent_nodes_for_draft_sync(
session=cast(Session, db.session),
draft_workflow=workflow,