mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor: agent draft (#37356)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
autofix-ci[bot]
parent
72faca2592
commit
e32a732812
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user