You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

db.py 11 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281
  1. # defines how core data types in autogenstudio are serialized and stored in the database
  2. from datetime import datetime
  3. from enum import Enum
  4. from typing import List, Optional, Tuple, Type, Union
  5. from uuid import UUID, uuid4
  6. from loguru import logger
  7. from pydantic import BaseModel
  8. from sqlalchemy import ForeignKey, Integer, UniqueConstraint
  9. from sqlmodel import JSON, Column, DateTime, Field, Relationship, SQLModel, func
  10. from .types import AgentConfig, MessageConfig, MessageMeta, ModelConfig, TeamConfig, TeamResult, ToolConfig
  11. # added for python3.11 and sqlmodel 0.0.22 incompatibility
  12. if hasattr(SQLModel, "model_config"):
  13. SQLModel.model_config["protected_namespaces"] = ()
  14. elif hasattr(SQLModel, "Config"):
  15. class CustomSQLModel(SQLModel):
  16. class Config:
  17. protected_namespaces = ()
  18. SQLModel = CustomSQLModel
  19. else:
  20. logger.warning("Unable to set protected_namespaces.")
  21. # pylint: disable=protected-access
  22. class ComponentTypes(Enum):
  23. TEAM = "team"
  24. AGENT = "agent"
  25. MODEL = "model"
  26. TOOL = "tool"
  27. @property
  28. def model_class(self) -> Type[SQLModel]:
  29. return {
  30. ComponentTypes.TEAM: Team,
  31. ComponentTypes.AGENT: Agent,
  32. ComponentTypes.MODEL: Model,
  33. ComponentTypes.TOOL: Tool,
  34. }[self]
  35. class LinkTypes(Enum):
  36. AGENT_MODEL = "agent_model"
  37. AGENT_TOOL = "agent_tool"
  38. TEAM_AGENT = "team_agent"
  39. @property
  40. # type: ignore
  41. def link_config(self) -> Tuple[Type[SQLModel], Type[SQLModel], Type[SQLModel]]:
  42. return {
  43. LinkTypes.AGENT_MODEL: (Agent, Model, AgentModelLink),
  44. LinkTypes.AGENT_TOOL: (Agent, Tool, AgentToolLink),
  45. LinkTypes.TEAM_AGENT: (Team, Agent, TeamAgentLink),
  46. }[self]
  47. @property
  48. def primary_class(self) -> Type[SQLModel]: # type: ignore
  49. return self.link_config[0]
  50. @property
  51. def secondary_class(self) -> Type[SQLModel]: # type: ignore
  52. return self.link_config[1]
  53. @property
  54. def link_table(self) -> Type[SQLModel]: # type: ignore
  55. return self.link_config[2]
  56. # link models
  57. class AgentToolLink(SQLModel, table=True):
  58. __table_args__ = (
  59. UniqueConstraint("agent_id", "sequence", name="unique_agent_tool_sequence"),
  60. {"sqlite_autoincrement": True},
  61. )
  62. agent_id: int = Field(default=None, primary_key=True, foreign_key="agent.id")
  63. tool_id: int = Field(default=None, primary_key=True, foreign_key="tool.id")
  64. sequence: Optional[int] = Field(default=0, primary_key=True)
  65. class AgentModelLink(SQLModel, table=True):
  66. __table_args__ = (
  67. UniqueConstraint("agent_id", "sequence", name="unique_agent_tool_sequence"),
  68. {"sqlite_autoincrement": True},
  69. )
  70. agent_id: int = Field(default=None, primary_key=True, foreign_key="agent.id")
  71. model_id: int = Field(default=None, primary_key=True, foreign_key="model.id")
  72. sequence: Optional[int] = Field(default=0, primary_key=True)
  73. class TeamAgentLink(SQLModel, table=True):
  74. __table_args__ = (
  75. UniqueConstraint("agent_id", "sequence", name="unique_agent_tool_sequence"),
  76. {"sqlite_autoincrement": True},
  77. )
  78. team_id: int = Field(default=None, primary_key=True, foreign_key="team.id")
  79. agent_id: int = Field(default=None, primary_key=True, foreign_key="agent.id")
  80. sequence: Optional[int] = Field(default=0, primary_key=True)
  81. # database models
  82. class Tool(SQLModel, table=True):
  83. __table_args__ = {"sqlite_autoincrement": True}
  84. id: Optional[int] = Field(default=None, primary_key=True)
  85. created_at: datetime = Field(
  86. default_factory=datetime.now,
  87. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  88. ) # pylint: disable=not-callable
  89. updated_at: datetime = Field(
  90. default_factory=datetime.now,
  91. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  92. ) # pylint: disable=not-callable
  93. user_id: Optional[str] = None
  94. version: Optional[str] = "0.0.1"
  95. config: Union[ToolConfig, dict] = Field(sa_column=Column(JSON))
  96. agents: List["Agent"] = Relationship(back_populates="tools", link_model=AgentToolLink)
  97. class Model(SQLModel, table=True):
  98. __table_args__ = {"sqlite_autoincrement": True}
  99. id: Optional[int] = Field(default=None, primary_key=True)
  100. created_at: datetime = Field(
  101. default_factory=datetime.now,
  102. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  103. ) # pylint: disable=not-callable
  104. updated_at: datetime = Field(
  105. default_factory=datetime.now,
  106. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  107. ) # pylint: disable=not-callable
  108. user_id: Optional[str] = None
  109. version: Optional[str] = "0.0.1"
  110. config: Union[ModelConfig, dict] = Field(sa_column=Column(JSON))
  111. agents: List["Agent"] = Relationship(back_populates="models", link_model=AgentModelLink)
  112. class Team(SQLModel, table=True):
  113. __table_args__ = {"sqlite_autoincrement": True}
  114. id: Optional[int] = Field(default=None, primary_key=True)
  115. created_at: datetime = Field(
  116. default_factory=datetime.now,
  117. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  118. ) # pylint: disable=not-callable
  119. updated_at: datetime = Field(
  120. default_factory=datetime.now,
  121. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  122. ) # pylint: disable=not-callable
  123. user_id: Optional[str] = None
  124. version: Optional[str] = "0.0.1"
  125. config: Union[TeamConfig, dict] = Field(sa_column=Column(JSON))
  126. agents: List["Agent"] = Relationship(back_populates="teams", link_model=TeamAgentLink)
  127. class Agent(SQLModel, table=True):
  128. __table_args__ = {"sqlite_autoincrement": True}
  129. id: Optional[int] = Field(default=None, primary_key=True)
  130. created_at: datetime = Field(
  131. default_factory=datetime.now,
  132. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  133. ) # pylint: disable=not-callable
  134. updated_at: datetime = Field(
  135. default_factory=datetime.now,
  136. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  137. ) # pylint: disable=not-callable
  138. user_id: Optional[str] = None
  139. version: Optional[str] = "0.0.1"
  140. config: Union[AgentConfig, dict] = Field(sa_column=Column(JSON))
  141. tools: List[Tool] = Relationship(back_populates="agents", link_model=AgentToolLink)
  142. models: List[Model] = Relationship(back_populates="agents", link_model=AgentModelLink)
  143. teams: List[Team] = Relationship(back_populates="agents", link_model=TeamAgentLink)
  144. class Message(SQLModel, table=True):
  145. __table_args__ = {"sqlite_autoincrement": True}
  146. id: Optional[int] = Field(default=None, primary_key=True)
  147. created_at: datetime = Field(
  148. default_factory=datetime.now,
  149. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  150. ) # pylint: disable=not-callable
  151. updated_at: datetime = Field(
  152. default_factory=datetime.now,
  153. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  154. ) # pylint: disable=not-callable
  155. user_id: Optional[str] = None
  156. version: Optional[str] = "0.0.1"
  157. config: Union[MessageConfig, dict] = Field(default_factory=MessageConfig, sa_column=Column(JSON))
  158. session_id: Optional[int] = Field(
  159. default=None, sa_column=Column(Integer, ForeignKey("session.id", ondelete="CASCADE"))
  160. )
  161. run_id: Optional[UUID] = Field(default=None, foreign_key="run.id")
  162. message_meta: Optional[Union[MessageMeta, dict]] = Field(default={}, sa_column=Column(JSON))
  163. class Session(SQLModel, table=True):
  164. __table_args__ = {"sqlite_autoincrement": True}
  165. id: Optional[int] = Field(default=None, primary_key=True)
  166. created_at: datetime = Field(
  167. default_factory=datetime.now,
  168. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  169. ) # pylint: disable=not-callable
  170. updated_at: datetime = Field(
  171. default_factory=datetime.now,
  172. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  173. ) # pylint: disable=not-callable
  174. user_id: Optional[str] = None
  175. version: Optional[str] = "0.0.1"
  176. team_id: Optional[int] = Field(default=None, sa_column=Column(Integer, ForeignKey("team.id", ondelete="CASCADE")))
  177. name: Optional[str] = None
  178. class RunStatus(str, Enum):
  179. CREATED = "created"
  180. ACTIVE = "active"
  181. COMPLETE = "complete"
  182. ERROR = "error"
  183. STOPPED = "stopped"
  184. class Run(SQLModel, table=True):
  185. """Represents a single execution run within a session"""
  186. __table_args__ = {"sqlite_autoincrement": True}
  187. id: UUID = Field(default_factory=uuid4, primary_key=True, index=True)
  188. created_at: datetime = Field(
  189. default_factory=datetime.now, sa_column=Column(DateTime(timezone=True), server_default=func.now())
  190. )
  191. updated_at: datetime = Field(
  192. default_factory=datetime.now, sa_column=Column(DateTime(timezone=True), onupdate=func.now())
  193. )
  194. session_id: Optional[int] = Field(
  195. default=None, sa_column=Column(Integer, ForeignKey("session.id", ondelete="CASCADE"), nullable=False)
  196. )
  197. status: RunStatus = Field(default=RunStatus.CREATED)
  198. # Store the original user task
  199. task: Union[MessageConfig, dict] = Field(default_factory=MessageConfig, sa_column=Column(JSON))
  200. # Store TeamResult which contains TaskResult
  201. team_result: Union[TeamResult, dict] = Field(default=None, sa_column=Column(JSON))
  202. error_message: Optional[str] = None
  203. version: Optional[str] = "0.0.1"
  204. messages: Union[List[Message], List[dict]] = Field(default_factory=list, sa_column=Column(JSON))
  205. class Config:
  206. json_encoders = {UUID: str, datetime: lambda v: v.isoformat()}
  207. class GalleryConfig(SQLModel, table=False):
  208. id: UUID = Field(default_factory=uuid4, primary_key=True, index=True)
  209. title: Optional[str] = None
  210. description: Optional[str] = None
  211. run: Run
  212. team: TeamConfig = None
  213. tags: Optional[List[str]] = None
  214. visibility: str = "public" # public, private, shared
  215. class Config:
  216. json_encoders = {UUID: str, datetime: lambda v: v.isoformat()}
  217. class Gallery(SQLModel, table=True):
  218. __table_args__ = {"sqlite_autoincrement": True}
  219. id: Optional[int] = Field(default=None, primary_key=True)
  220. created_at: datetime = Field(
  221. default_factory=datetime.now,
  222. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  223. )
  224. updated_at: datetime = Field(
  225. default_factory=datetime.now,
  226. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  227. )
  228. user_id: Optional[str] = None
  229. version: Optional[str] = "0.0.1"
  230. config: Union[GalleryConfig, dict] = Field(default_factory=GalleryConfig, sa_column=Column(JSON))