services.py 25 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617
  1. from dataclasses import dataclass
  2. from datetime import datetime, timedelta
  3. from core_events import EventPublishContract, EventServiceClient, EventServiceClientError
  4. from core_domain import AgentRunContract, TeamMemberContract
  5. from core_shared import JSONValue, try_build_redis_client
  6. from core_shared.task_queue import TaskQueuePublisher
  7. from app.bootstrap.settings import TeamServiceSettings
  8. from app.db.models import TeamDefinition, TeamRun, TeamVersion
  9. from app.domain.repositories import (
  10. TeamDefinitionRepository,
  11. TeamRunRepository,
  12. TeamVersionRepository)
  13. from app.infrastructure.agent_client import AgentServiceClient, AgentServiceClientError
  14. from app.schemas.team import (
  15. TeamConfigCreateRequestDto,
  16. TeamConfigUpdateRequestDto,
  17. TeamCreateRequest,
  18. TeamCreateRequestDto,
  19. TeamDeleteRequestDto,
  20. TeamRunCreateRequest,
  21. TeamRunCreateRequestDto,
  22. TeamRunExecuteRequest,
  23. TeamRunStatusUpdateRequestDto,
  24. TeamRunStatusUpdateRequest,
  25. TeamStatusUpdateRequest,
  26. TeamUpdateRequestDto,
  27. TeamVersionCreateRequest)
  28. @dataclass(frozen=True)
  29. class TeamMemberRunResult:
  30. member: TeamMemberContract
  31. run: AgentRunContract
  32. class TeamApplicationService:
  33. def __init__(
  34. self,
  35. *,
  36. team_repository: TeamDefinitionRepository,
  37. team_version_repository: TeamVersionRepository,
  38. team_run_repository: TeamRunRepository,
  39. agent_client: AgentServiceClient | None = None,
  40. event_client: EventServiceClient | None = None,
  41. task_queue_publisher: TaskQueuePublisher | None = None) -> None:
  42. self.team_repository = team_repository
  43. self.team_version_repository = team_version_repository
  44. self.team_run_repository = team_run_repository
  45. self.agent_client = agent_client
  46. self.event_client = event_client
  47. self.task_queue_publisher = task_queue_publisher
  48. def create_team(self, payload: TeamCreateRequest) -> TeamDefinition:
  49. return self.team_repository.create(
  50. code=payload.code,
  51. name=payload.name,
  52. description=payload.description,
  53. team_type=payload.team_type,
  54. owner_user_id=payload.owner_user_id,
  55. metadata_json=payload.metadata_json)
  56. def create_team_from_contract(self, payload: TeamCreateRequestDto) -> TeamDefinition:
  57. return self.create_team(
  58. TeamCreateRequest(
  59. code=self._build_team_code(payload.name),
  60. name=payload.name,
  61. description=payload.description,
  62. team_type=payload.teamType,
  63. owner_user_id=payload.ownerUserId,
  64. metadata_json=payload.metadata))
  65. def list_teams(self) -> list[TeamDefinition]:
  66. return self.team_repository.list_all()
  67. def get_team(self, *, team_id: str) -> TeamDefinition | None:
  68. return self.team_repository.get_by_id(team_id=team_id)
  69. def update_team_from_contract(self, payload: TeamUpdateRequestDto) -> TeamDefinition | None:
  70. entity = self.team_repository.get_by_id(team_id=payload.teamId)
  71. if entity is None:
  72. return None
  73. if payload.name is not None:
  74. entity.name = payload.name
  75. entity.code = self._build_team_code(payload.name)
  76. if payload.description is not None:
  77. entity.description = payload.description
  78. if payload.teamType is not None:
  79. entity.team_type = payload.teamType
  80. if payload.status is not None:
  81. entity.status = payload.status
  82. if payload.ownerUserId is not None:
  83. entity.owner_user_id = payload.ownerUserId
  84. if payload.metadata is not None:
  85. entity.metadata_json = payload.metadata
  86. return self.team_repository.save(entity)
  87. def delete_team_from_contract(self, payload: TeamDeleteRequestDto) -> bool:
  88. entity = self.team_repository.get_by_id(team_id=payload.teamId)
  89. if entity is None:
  90. return False
  91. self.team_repository.delete(entity)
  92. return True
  93. def update_team_status(
  94. self,
  95. *,
  96. team_id: str,
  97. payload: TeamStatusUpdateRequest) -> TeamDefinition | None:
  98. return self.team_repository.update_status(
  99. team_id=team_id,
  100. status=payload.status)
  101. def create_team_version(self, payload: TeamVersionCreateRequest) -> TeamVersion:
  102. team = self.team_repository.get_by_id(team_id=payload.team_id)
  103. if team is None:
  104. raise ValueError(f"team not found: {payload.team_id}")
  105. if not payload.member_refs:
  106. raise ValueError("team version requires at least one member")
  107. return self.team_version_repository.create(
  108. team_id=payload.team_id,
  109. status=payload.status,
  110. coordination_mode=payload.coordination_mode,
  111. objective=payload.objective,
  112. member_refs_json=[item.model_dump(mode="json") for item in payload.member_refs],
  113. policy_json=payload.policy_json)
  114. def create_team_config_from_contract(self, payload: TeamConfigCreateRequestDto) -> TeamVersion:
  115. return self.create_team_version(
  116. TeamVersionCreateRequest(
  117. team_id=payload.teamId,
  118. status=payload.status,
  119. coordination_mode=payload.coordinationMode,
  120. objective=payload.objective,
  121. member_refs=self._normalize_member_refs(payload.memberRefs),
  122. policy_json=payload.policy))
  123. def list_team_versions(self, *, team_id: str) -> list[TeamVersion]:
  124. return self.team_version_repository.list_by_team(team_id=team_id)
  125. def list_team_configs(self, *, team_id: str | None = None) -> list[TeamVersion]:
  126. if team_id is not None:
  127. return self.team_version_repository.list_by_team(team_id=team_id)
  128. return self.team_version_repository.list_all()
  129. def get_team_config(self, *, config_id: str) -> TeamVersion | None:
  130. return self.team_version_repository.get_by_id(team_version_id=config_id)
  131. def update_team_config_from_contract(self, payload: TeamConfigUpdateRequestDto) -> TeamVersion | None:
  132. entity = self.team_version_repository.get_by_id(team_version_id=payload.configId)
  133. if entity is None:
  134. return None
  135. if payload.status is not None:
  136. entity.status = payload.status
  137. entity.published_time = datetime.utcnow() if payload.status == "published" else entity.published_time
  138. if payload.coordinationMode is not None:
  139. entity.coordination_mode = payload.coordinationMode
  140. if payload.objective is not None:
  141. entity.objective = payload.objective
  142. if payload.memberRefs is not None:
  143. normalized = self._normalize_member_refs(payload.memberRefs)
  144. if not normalized:
  145. raise ValueError("team config requires at least one member")
  146. entity.member_refs_json = [item.model_dump(mode="json") for item in normalized]
  147. if payload.policy is not None:
  148. entity.policy_json = payload.policy
  149. return self.team_version_repository.save(entity)
  150. def delete_team_config(self, *, config_id: str) -> bool:
  151. entity = self.team_version_repository.get_by_id(team_version_id=config_id)
  152. if entity is None:
  153. return False
  154. self.team_version_repository.delete(entity)
  155. return True
  156. def create_team_run(self, payload: TeamRunCreateRequest) -> TeamRun:
  157. team_version = self._resolve_team_version(
  158. team_id=payload.team_id,
  159. team_version_id=payload.team_version_id)
  160. if team_version is None:
  161. raise ValueError("published team version not found")
  162. team_run = self.team_run_repository.create(
  163. team_id=payload.team_id,
  164. team_version_id=team_version.id,
  165. session_id=payload.session_id,
  166. input_text=payload.input_text,
  167. input_json=payload.input_json)
  168. self._publish_event(
  169. event_type="team.run.created",
  170. team_run=team_run,
  171. payload_json={"team_run_id": team_run.id, "status": team_run.status})
  172. if self.task_queue_publisher is not None:
  173. self.task_queue_publisher.publish_team_run(
  174. team_run_id=team_run.id)
  175. return team_run
  176. def create_team_run_from_contract(self, payload: TeamRunCreateRequestDto) -> TeamRun:
  177. return self.create_team_run(
  178. TeamRunCreateRequest(
  179. team_id=payload.teamId,
  180. team_version_id=payload.teamConfigId,
  181. session_id=payload.sessionId,
  182. input_text=payload.inputText,
  183. input_json=payload.inputJson))
  184. def list_team_runs(
  185. self,
  186. *,
  187. team_id: str | None = None,
  188. session_id: str | None = None) -> list[TeamRun]:
  189. return self.team_run_repository.list_by_scope(
  190. team_id=team_id,
  191. session_id=session_id)
  192. def get_team_run(self, *, team_run_id: str) -> TeamRun | None:
  193. return self.team_run_repository.get_by_id(team_run_id=team_run_id)
  194. def delete_team_run(self, *, team_run_id: str) -> bool:
  195. entity = self.team_run_repository.get_by_id(team_run_id=team_run_id)
  196. if entity is None:
  197. return False
  198. self.team_run_repository.delete(entity)
  199. return True
  200. def update_team_run_status(
  201. self,
  202. *,
  203. team_run_id: str,
  204. payload: TeamRunStatusUpdateRequest) -> TeamRun | None:
  205. entity = self.team_run_repository.get_by_id(
  206. team_run_id=team_run_id)
  207. if entity is None:
  208. return None
  209. return self.team_run_repository.update_status(
  210. team_run_id=team_run_id,
  211. status=payload.status,
  212. worker_key=payload.worker_key,
  213. output_text=payload.output_text,
  214. output_json=payload.output_json,
  215. error_code=payload.error_code,
  216. error_message=payload.error_message)
  217. def update_team_run_status_from_contract(
  218. self,
  219. payload: TeamRunStatusUpdateRequestDto) -> TeamRun | None:
  220. return self.update_team_run_status(
  221. team_run_id=payload.teamRunId,
  222. payload=TeamRunStatusUpdateRequest(
  223. status=payload.status,
  224. worker_key=payload.workerKey,
  225. output_text=payload.outputText,
  226. output_json=payload.outputJson,
  227. error_code=payload.errorCode,
  228. error_message=payload.errorMessage))
  229. def execute_team_run(
  230. self,
  231. *,
  232. team_run_id: str,
  233. payload: TeamRunExecuteRequest) -> TeamRun | None:
  234. team_run = self.team_run_repository.get_by_id(
  235. team_run_id=team_run_id)
  236. if team_run is None:
  237. return None
  238. team_version = self.team_version_repository.get_by_id(
  239. team_version_id=team_run.team_version_id)
  240. if team_version is None:
  241. failed_run = self.team_run_repository.update_status(
  242. team_run_id=team_run.id,
  243. status="failed",
  244. worker_key=payload.worker_key,
  245. error_code="team_version_missing",
  246. error_message=f"team version not found: {team_run.team_version_id}")
  247. if failed_run is not None:
  248. self._publish_event(
  249. event_type="team.run.failed",
  250. team_run=failed_run,
  251. payload_json={
  252. "team_run_id": failed_run.id,
  253. "status": failed_run.status,
  254. "error_code": "team_version_missing",
  255. })
  256. return failed_run
  257. running_run = self.team_run_repository.update_status(
  258. team_run_id=team_run.id,
  259. status="running",
  260. worker_key=payload.worker_key)
  261. if running_run is None:
  262. return None
  263. members = self._read_team_members(team_version)
  264. if not members:
  265. return self.team_run_repository.update_status(
  266. team_run_id=team_run.id,
  267. status="failed",
  268. worker_key=payload.worker_key,
  269. error_code="team_members_missing",
  270. error_message="team version has no valid members")
  271. try:
  272. member_results = self._execute_members(
  273. team_run=team_run,
  274. team_version=team_version,
  275. members=members,
  276. worker_key=payload.worker_key,
  277. dry_run=payload.dry_run)
  278. except AgentServiceClientError as exc:
  279. return self.team_run_repository.update_status(
  280. team_run_id=team_run.id,
  281. status="failed",
  282. worker_key=payload.worker_key,
  283. error_code="agent_service_error",
  284. error_message=str(exc))
  285. failed_results = [item for item in member_results if item.run.status != "completed"]
  286. output_text = self._build_team_output_text(
  287. team_version=team_version,
  288. member_results=member_results)
  289. output_json: dict[str, JSONValue] = {
  290. "dry_run": payload.dry_run,
  291. "coordination_mode": team_version.coordination_mode,
  292. "team_version_id": team_version.id,
  293. "member_run_count": len(member_results),
  294. "member_results": [
  295. self._member_result_to_json(item) for item in member_results
  296. ],
  297. }
  298. if failed_results:
  299. failed_run = self.team_run_repository.update_status(
  300. team_run_id=team_run.id,
  301. status="failed",
  302. worker_key=payload.worker_key,
  303. output_text=output_text,
  304. output_json=output_json,
  305. error_code="member_run_failed",
  306. error_message=f"{len(failed_results)} member run(s) failed")
  307. if failed_run is not None:
  308. self._publish_event(
  309. event_type="team.run.failed",
  310. team_run=failed_run,
  311. payload_json={
  312. "team_run_id": failed_run.id,
  313. "status": failed_run.status,
  314. "failed_member_count": len(failed_results),
  315. })
  316. return failed_run
  317. completed_run = self.team_run_repository.update_status(
  318. team_run_id=team_run.id,
  319. status="completed",
  320. worker_key=payload.worker_key,
  321. output_text=output_text,
  322. output_json=output_json)
  323. if completed_run is not None:
  324. self._publish_event(
  325. event_type="team.run.completed",
  326. team_run=completed_run,
  327. payload_json={
  328. "team_run_id": completed_run.id,
  329. "status": completed_run.status,
  330. "member_run_count": len(member_results),
  331. })
  332. return completed_run
  333. def execute_next_claimed_team_run(
  334. self,
  335. *,
  336. worker_key: str,
  337. lease_seconds: int,
  338. stale_running_seconds: int,
  339. dry_run: bool) -> tuple[TeamRun, int] | None:
  340. released_lease_count = self.team_run_repository.release_expired_leases(
  341. now_time=datetime.utcnow(),
  342. stale_running_seconds=stale_running_seconds)
  343. claimed_team_run = self.team_run_repository.claim_next_queued(
  344. worker_key=worker_key,
  345. lease_expire_time=datetime.utcnow() + timedelta(seconds=lease_seconds))
  346. if claimed_team_run is None:
  347. return None
  348. result = self.execute_team_run(
  349. team_run_id=claimed_team_run.id,
  350. payload=TeamRunExecuteRequest(
  351. worker_key=worker_key,
  352. dry_run=dry_run))
  353. if result is None:
  354. return None
  355. return result, released_lease_count
  356. def _resolve_team_version(
  357. self,
  358. *,
  359. team_id: str,
  360. team_version_id: str | None) -> TeamVersion | None:
  361. if team_version_id is not None:
  362. return self.team_version_repository.get_by_id(
  363. team_version_id=team_version_id)
  364. return self.team_version_repository.get_latest_published(
  365. team_id=team_id)
  366. def _execute_members(
  367. self,
  368. *,
  369. team_run: TeamRun,
  370. team_version: TeamVersion,
  371. members: list[TeamMemberContract],
  372. worker_key: str | None,
  373. dry_run: bool) -> list[TeamMemberRunResult]:
  374. if self.agent_client is None:
  375. raise AgentServiceClientError("agent service client is not configured")
  376. member_results: list[TeamMemberRunResult] = []
  377. prior_outputs: list[dict[str, JSONValue]] = []
  378. for member in self._order_members(members):
  379. member_input_json = self._build_member_input_json(
  380. team_run=team_run,
  381. team_version=team_version,
  382. member=member,
  383. prior_outputs=prior_outputs)
  384. created_run = self.agent_client.create_agent_run(
  385. agent_id=member.agent_id,
  386. agent_version_id=member.agent_version_id,
  387. session_id=team_run.session_id,
  388. input_text=self._build_member_input_text(
  389. team_run=team_run,
  390. team_version=team_version,
  391. member=member),
  392. input_json=member_input_json)
  393. executed_run = self.agent_client.execute_agent_run(
  394. agent_run_id=created_run.id,
  395. worker_key=worker_key,
  396. dry_run=dry_run)
  397. member_results.append(TeamMemberRunResult(member=member, run=executed_run))
  398. prior_outputs.append(
  399. {
  400. "member_key": member.member_key,
  401. "member_role": member.role,
  402. "member_name": member.name,
  403. "agent_id": member.agent_id,
  404. "agent_run_id": executed_run.id,
  405. "status": executed_run.status,
  406. "output_text": executed_run.output_text,
  407. "output_json": executed_run.output_json or {},
  408. }
  409. )
  410. return member_results
  411. def _read_team_members(self, team_version: TeamVersion) -> list[TeamMemberContract]:
  412. members: list[TeamMemberContract] = []
  413. for item in team_version.member_refs_json:
  414. try:
  415. members.append(TeamMemberContract.model_validate(item))
  416. except ValueError:
  417. continue
  418. return members
  419. def _order_members(self, members: list[TeamMemberContract]) -> list[TeamMemberContract]:
  420. role_priority = {
  421. "planner": 0,
  422. "supervisor": 1,
  423. "specialist": 2,
  424. "executor": 3,
  425. "reviewer": 4,
  426. }
  427. return sorted(members, key=lambda item: role_priority.get(item.role, 10))
  428. def _build_member_input_text(
  429. self,
  430. *,
  431. team_run: TeamRun,
  432. team_version: TeamVersion,
  433. member: TeamMemberContract) -> str:
  434. lines = [
  435. f"Team objective: {team_version.objective or 'No objective provided.'}",
  436. f"Member role: {member.role}",
  437. ]
  438. if member.responsibility:
  439. lines.append(f"Responsibility: {member.responsibility}")
  440. if team_run.input_text:
  441. lines.append(f"User task: {team_run.input_text}")
  442. return "\n".join(lines)
  443. def _build_member_input_json(
  444. self,
  445. *,
  446. team_run: TeamRun,
  447. team_version: TeamVersion,
  448. member: TeamMemberContract,
  449. prior_outputs: list[dict[str, JSONValue]]) -> dict[str, JSONValue]:
  450. input_json: dict[str, JSONValue] = dict(team_run.input_json or {})
  451. input_json.update(
  452. {
  453. "team_id": team_run.team_id,
  454. "team_run_id": team_run.id,
  455. "team_version_id": team_version.id,
  456. "team_objective": team_version.objective,
  457. "member_key": member.member_key,
  458. "member_role": member.role,
  459. "member_responsibility": member.responsibility,
  460. "prior_member_outputs": prior_outputs,
  461. }
  462. )
  463. configured_input = member.config_json.get("input_json")
  464. if isinstance(configured_input, dict):
  465. input_json.update(
  466. {str(item_key): item_value for item_key, item_value in configured_input.items()}
  467. )
  468. return input_json
  469. def _build_team_output_text(
  470. self,
  471. *,
  472. team_version: TeamVersion,
  473. member_results: list[TeamMemberRunResult]) -> str:
  474. lines = [
  475. f"Team objective: {team_version.objective or 'No objective provided.'}",
  476. f"Coordination mode: {team_version.coordination_mode}",
  477. "Team conversation:",
  478. ]
  479. for index, item in enumerate(member_results, start=1):
  480. output_text = item.run.output_text or item.run.error_message or ""
  481. speaker = item.member.name or item.member.member_key
  482. lines.append(
  483. f"{index}. {speaker} ({item.member.role}) "
  484. f"status={item.run.status}: {output_text}")
  485. return "\n".join(lines)
  486. def _member_result_to_json(self, result: TeamMemberRunResult) -> dict[str, JSONValue]:
  487. return {
  488. "member_key": result.member.member_key,
  489. "member_role": result.member.role,
  490. "member_name": result.member.name,
  491. "member_responsibility": result.member.responsibility,
  492. "agent_run_id": result.run.id,
  493. "agent_id": result.run.agent_id,
  494. "agent_version_id": result.run.agent_version_id,
  495. "status": result.run.status,
  496. "output_text": result.run.output_text,
  497. "output_json": result.run.output_json or {},
  498. "error_code": result.run.error_code,
  499. "error_message": result.run.error_message,
  500. }
  501. def _publish_event(
  502. self,
  503. *,
  504. event_type: str,
  505. team_run: TeamRun,
  506. payload_json: dict[str, JSONValue]) -> None:
  507. if self.event_client is None:
  508. return
  509. try:
  510. self.event_client.publish_event(
  511. EventPublishContract(
  512. event_type=event_type,
  513. source_service="team-service",
  514. aggregate_type="team_run",
  515. aggregate_id=team_run.id,
  516. correlation_id=team_run.session_id,
  517. payload_json={
  518. **payload_json,
  519. "team_id": team_run.team_id,
  520. "team_version_id": team_run.team_version_id,
  521. })
  522. )
  523. except EventServiceClientError:
  524. return
  525. def _build_team_code(self, name: str) -> str:
  526. base = "".join(
  527. char.lower() if char.isalnum() else "_"
  528. for char in name
  529. ).strip("_") or "team"
  530. return base[:64]
  531. def _normalize_member_refs(self, member_refs: list[dict[str, JSONValue]]) -> list[TeamMemberContract]:
  532. members: list[TeamMemberContract] = []
  533. for index, item in enumerate(member_refs, start=1):
  534. role = item.get("role")
  535. normalized_role = "executor" if role == "worker" else role
  536. member = {
  537. **item,
  538. "member_key": item.get("member_key") or item.get("memberKey") or f"member_{index}",
  539. "agent_id": item.get("agent_id") or item.get("agentId"),
  540. "agent_version_id": item.get("agent_version_id") or item.get("agentVersionId"),
  541. "role": normalized_role or "specialist",
  542. "config_json": item.get("config_json") or item.get("configJson") or {},
  543. }
  544. members.append(TeamMemberContract.model_validate(member))
  545. return members
  546. def build_team_application_service(
  547. *,
  548. team_repository: TeamDefinitionRepository,
  549. team_version_repository: TeamVersionRepository,
  550. team_run_repository: TeamRunRepository,
  551. settings: TeamServiceSettings) -> TeamApplicationService:
  552. redis_client = try_build_redis_client(settings.redis_url)
  553. return TeamApplicationService(
  554. team_repository=team_repository,
  555. team_version_repository=team_version_repository,
  556. team_run_repository=team_run_repository,
  557. agent_client=AgentServiceClient(
  558. base_url=settings.agent_service_url,
  559. timeout_seconds=settings.agent_service_timeout_seconds),
  560. event_client=EventServiceClient(
  561. base_url=settings.event_service_url,
  562. timeout_seconds=settings.event_service_timeout_seconds),
  563. task_queue_publisher=(
  564. TaskQueuePublisher(client=redis_client) if redis_client is not None else None
  565. ))