repositories.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356
  1. from datetime import datetime
  2. from sqlalchemy import func, select
  3. from sqlalchemy.orm import Session
  4. from core_domain import (
  5. AgentRunStatus,
  6. AgentStatus,
  7. AgentToolInvocationStatus,
  8. AgentVersionStatus)
  9. from core_shared import JSONValue
  10. from app.db.models import AgentDefinition, AgentRun, AgentToolInvocation, AgentVersion
  11. class AgentDefinitionRepository:
  12. def __init__(self, db: Session) -> None:
  13. self.db = db
  14. def create(
  15. self,
  16. *,
  17. code: str,
  18. name: str,
  19. description: str | None,
  20. agent_type: str,
  21. owner_user_id: str | None,
  22. metadata_json: dict[str, JSONValue] | None) -> AgentDefinition:
  23. entity = AgentDefinition(
  24. code=code,
  25. name=name,
  26. description=description,
  27. agent_type=agent_type,
  28. owner_user_id=owner_user_id,
  29. metadata_json=metadata_json)
  30. self.db.add(entity)
  31. self.db.commit()
  32. self.db.refresh(entity)
  33. return entity
  34. def list_all(self) -> list[AgentDefinition]:
  35. stmt = (
  36. select(AgentDefinition)
  37. .order_by(AgentDefinition.created_time.desc())
  38. )
  39. return list(self.db.scalars(stmt))
  40. def get_by_id(self, *, agent_id: str) -> AgentDefinition | None:
  41. stmt = (
  42. select(AgentDefinition)
  43. .where(AgentDefinition.id == agent_id)
  44. )
  45. return self.db.scalar(stmt)
  46. def update(
  47. self,
  48. *,
  49. agent_id: str,
  50. name: str | None,
  51. description: str | None,
  52. metadata_json: dict[str, JSONValue] | None) -> AgentDefinition | None:
  53. entity = self.get_by_id(agent_id=agent_id)
  54. if entity is None:
  55. return None
  56. if name is not None:
  57. entity.name = name
  58. if description is not None:
  59. entity.description = description
  60. if metadata_json is not None:
  61. entity.metadata_json = metadata_json
  62. self.db.commit()
  63. self.db.refresh(entity)
  64. return entity
  65. def update_status(
  66. self,
  67. *,
  68. agent_id: str,
  69. status: AgentStatus) -> AgentDefinition | None:
  70. entity = self.get_by_id(agent_id=agent_id)
  71. if entity is None:
  72. return None
  73. entity.status = status
  74. self.db.commit()
  75. self.db.refresh(entity)
  76. return entity
  77. class AgentVersionRepository:
  78. def __init__(self, db: Session) -> None:
  79. self.db = db
  80. def create(
  81. self,
  82. *,
  83. agent_id: str,
  84. status: AgentVersionStatus,
  85. role: str,
  86. goal: str | None,
  87. system_prompt: str,
  88. model_config_json: dict[str, JSONValue],
  89. memory_policy_json: dict[str, JSONValue],
  90. tool_refs_json: list[dict[str, JSONValue]],
  91. skill_refs_json: list[dict[str, JSONValue]]) -> AgentVersion:
  92. version_no = self._next_version_no(agent_id)
  93. entity = AgentVersion(
  94. agent_id=agent_id,
  95. version_no=version_no,
  96. status=status,
  97. role=role,
  98. goal=goal,
  99. system_prompt=system_prompt,
  100. model_config_json=model_config_json,
  101. memory_policy_json=memory_policy_json,
  102. tool_refs_json=tool_refs_json,
  103. skill_refs_json=skill_refs_json,
  104. published_time=datetime.utcnow() if status == "published" else None)
  105. self.db.add(entity)
  106. self.db.commit()
  107. self.db.refresh(entity)
  108. return entity
  109. def list_by_agent(self, *, agent_id: str) -> list[AgentVersion]:
  110. stmt = (
  111. select(AgentVersion)
  112. .where(AgentVersion.agent_id == agent_id)
  113. .order_by(AgentVersion.version_no.desc())
  114. )
  115. return list(self.db.scalars(stmt))
  116. def get_by_id(self, *, agent_version_id: str) -> AgentVersion | None:
  117. stmt = (
  118. select(AgentVersion)
  119. .where(AgentVersion.id == agent_version_id)
  120. )
  121. return self.db.scalar(stmt)
  122. def get_latest_published(self, *, agent_id: str) -> AgentVersion | None:
  123. stmt = (
  124. select(AgentVersion)
  125. .where(AgentVersion.agent_id == agent_id)
  126. .where(AgentVersion.status == "published")
  127. .order_by(AgentVersion.version_no.desc())
  128. .limit(1)
  129. )
  130. return self.db.scalar(stmt)
  131. def get_latest_by_agent(self, *, agent_id: str) -> AgentVersion | None:
  132. stmt = (
  133. select(AgentVersion)
  134. .where(AgentVersion.agent_id == agent_id)
  135. .order_by(AgentVersion.version_no.desc())
  136. .limit(1)
  137. )
  138. return self.db.scalar(stmt)
  139. def _next_version_no(self, agent_id: str) -> int:
  140. stmt = select(func.max(AgentVersion.version_no)).where(AgentVersion.agent_id == agent_id)
  141. current_max = self.db.scalar(stmt)
  142. return (current_max or 0) + 1
  143. class AgentRunRepository:
  144. def __init__(self, db: Session) -> None:
  145. self.db = db
  146. def create(
  147. self,
  148. *,
  149. agent_id: str,
  150. agent_version_id: str,
  151. session_id: str | None,
  152. input_text: str | None,
  153. input_json: dict[str, JSONValue] | None) -> AgentRun:
  154. now = datetime.utcnow()
  155. entity = AgentRun(
  156. agent_id=agent_id,
  157. agent_version_id=agent_version_id,
  158. session_id=session_id,
  159. input_text=input_text,
  160. input_json=input_json,
  161. status="queued",
  162. queued_time=now)
  163. self.db.add(entity)
  164. self.db.commit()
  165. self.db.refresh(entity)
  166. return entity
  167. def list_by_scope(
  168. self,
  169. *,
  170. agent_id: str | None = None,
  171. session_id: str | None = None) -> list[AgentRun]:
  172. stmt = select(AgentRun)
  173. if agent_id is not None:
  174. stmt = stmt.where(AgentRun.agent_id == agent_id)
  175. if session_id is not None:
  176. stmt = stmt.where(AgentRun.session_id == session_id)
  177. stmt = stmt.order_by(AgentRun.created_time.desc())
  178. return list(self.db.scalars(stmt))
  179. def get_by_id(self, *, agent_run_id: str) -> AgentRun | None:
  180. stmt = (
  181. select(AgentRun)
  182. .where(AgentRun.id == agent_run_id)
  183. )
  184. return self.db.scalar(stmt)
  185. def claim_next_queued(
  186. self,
  187. *,
  188. worker_key: str,
  189. lease_expire_time: datetime) -> AgentRun | None:
  190. stmt = (
  191. select(AgentRun)
  192. .where(AgentRun.status == "queued")
  193. .order_by(AgentRun.created_time.asc())
  194. .with_for_update(skip_locked=True)
  195. .limit(1)
  196. )
  197. entity = self.db.scalar(stmt)
  198. if entity is None:
  199. return None
  200. now = datetime.utcnow()
  201. entity.status = "running"
  202. entity.worker_key = worker_key
  203. entity.started_time = entity.started_time or now
  204. entity.lease_expire_time = lease_expire_time
  205. self.db.commit()
  206. self.db.refresh(entity)
  207. return entity
  208. def release_expired_leases(self, *, now_time: datetime, max_items: int = 100) -> int:
  209. stmt = (
  210. select(AgentRun)
  211. .where(AgentRun.status == "running")
  212. .where(AgentRun.lease_expire_time.is_not(None))
  213. .where(AgentRun.lease_expire_time <= now_time)
  214. .order_by(AgentRun.lease_expire_time.asc())
  215. .limit(max_items)
  216. )
  217. entities = list(self.db.scalars(stmt))
  218. for entity in entities:
  219. entity.status = "queued"
  220. entity.worker_key = None
  221. entity.lease_expire_time = None
  222. entity.queued_time = now_time
  223. entity.started_time = None
  224. entity.finished_time = None
  225. if entities:
  226. self.db.commit()
  227. return len(entities)
  228. def update_status(
  229. self,
  230. *,
  231. agent_run_id: str,
  232. status: AgentRunStatus,
  233. worker_key: str | None = None,
  234. output_text: str | None = None,
  235. output_json: dict[str, JSONValue] | None = None,
  236. error_code: str | None = None,
  237. error_message: str | None = None) -> AgentRun | None:
  238. entity = self.db.get(AgentRun, agent_run_id)
  239. if entity is None:
  240. return None
  241. now = datetime.utcnow()
  242. entity.status = status
  243. entity.worker_key = worker_key
  244. entity.output_text = output_text
  245. entity.output_json = output_json
  246. entity.error_code = error_code
  247. entity.error_message = error_message
  248. if status == "running" and entity.started_time is None:
  249. entity.started_time = now
  250. if status in {"completed", "failed", "cancelled"}:
  251. entity.finished_time = now
  252. entity.lease_expire_time = None
  253. self.db.commit()
  254. self.db.refresh(entity)
  255. return entity
  256. class AgentToolInvocationRepository:
  257. def __init__(self, db: Session) -> None:
  258. self.db = db
  259. def create(
  260. self,
  261. *,
  262. agent_run_id: str,
  263. agent_id: str,
  264. agent_version_id: str,
  265. tool_code: str | None,
  266. tool_binding_id: str | None,
  267. status: AgentToolInvocationStatus,
  268. reason: str | None = None,
  269. input_json: dict[str, JSONValue] | None = None) -> AgentToolInvocation:
  270. entity = AgentToolInvocation(
  271. agent_run_id=agent_run_id,
  272. agent_id=agent_id,
  273. agent_version_id=agent_version_id,
  274. tool_code=tool_code,
  275. tool_binding_id=tool_binding_id,
  276. status=status,
  277. reason=reason,
  278. input_json=input_json or {})
  279. self.db.add(entity)
  280. self.db.commit()
  281. self.db.refresh(entity)
  282. return entity
  283. def list_by_run(
  284. self,
  285. *,
  286. agent_run_id: str) -> list[AgentToolInvocation]:
  287. stmt = (
  288. select(AgentToolInvocation)
  289. .where(AgentToolInvocation.agent_run_id == agent_run_id)
  290. .order_by(AgentToolInvocation.created_time.asc())
  291. )
  292. return list(self.db.scalars(stmt))
  293. def update_status(
  294. self,
  295. *,
  296. invocation_id: str,
  297. status: AgentToolInvocationStatus,
  298. reason: str | None = None,
  299. output_text: str | None = None,
  300. output_json: dict[str, JSONValue] | None = None,
  301. error_message: str | None = None) -> AgentToolInvocation | None:
  302. entity = self.db.get(AgentToolInvocation, invocation_id)
  303. if entity is None:
  304. return None
  305. now = datetime.utcnow()
  306. entity.status = status
  307. entity.reason = reason
  308. entity.output_text = output_text
  309. entity.output_json = output_json
  310. entity.error_message = error_message
  311. if status == "running" and entity.started_time is None:
  312. entity.started_time = now
  313. if status in {"completed", "failed", "skipped"}:
  314. if entity.started_time is None:
  315. entity.started_time = now
  316. entity.finished_time = now
  317. self.db.commit()
  318. self.db.refresh(entity)
  319. return entity