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 4.3 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105
  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, Union
  5. from uuid import UUID, uuid4
  6. from autogen_core import ComponentModel
  7. from pydantic import ConfigDict
  8. from sqlalchemy import ForeignKey, Integer
  9. from sqlmodel import JSON, Column, DateTime, Field, SQLModel, func
  10. from .types import MessageConfig, MessageMeta, TeamResult
  11. class Team(SQLModel, table=True):
  12. __table_args__ = {"sqlite_autoincrement": True}
  13. id: Optional[int] = Field(default=None, primary_key=True)
  14. created_at: datetime = Field(
  15. default_factory=datetime.now,
  16. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  17. ) # pylint: disable=not-callable
  18. updated_at: datetime = Field(
  19. default_factory=datetime.now,
  20. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  21. ) # pylint: disable=not-callable
  22. user_id: Optional[str] = None
  23. version: Optional[str] = "0.0.1"
  24. component: Union[ComponentModel, dict] = Field(sa_column=Column(JSON))
  25. class Message(SQLModel, table=True):
  26. __table_args__ = {"sqlite_autoincrement": True}
  27. id: Optional[int] = Field(default=None, primary_key=True)
  28. created_at: datetime = Field(
  29. default_factory=datetime.now,
  30. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  31. ) # pylint: disable=not-callable
  32. updated_at: datetime = Field(
  33. default_factory=datetime.now,
  34. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  35. ) # pylint: disable=not-callable
  36. user_id: Optional[str] = None
  37. version: Optional[str] = "0.0.1"
  38. config: Union[MessageConfig, dict] = Field(default_factory=MessageConfig, sa_column=Column(JSON))
  39. session_id: Optional[int] = Field(
  40. default=None, sa_column=Column(Integer, ForeignKey("session.id", ondelete="CASCADE"))
  41. )
  42. run_id: Optional[UUID] = Field(default=None, foreign_key="run.id")
  43. message_meta: Optional[Union[MessageMeta, dict]] = Field(default={}, sa_column=Column(JSON))
  44. class Session(SQLModel, table=True):
  45. __table_args__ = {"sqlite_autoincrement": True}
  46. id: Optional[int] = Field(default=None, primary_key=True)
  47. created_at: datetime = Field(
  48. default_factory=datetime.now,
  49. sa_column=Column(DateTime(timezone=True), server_default=func.now()),
  50. ) # pylint: disable=not-callable
  51. updated_at: datetime = Field(
  52. default_factory=datetime.now,
  53. sa_column=Column(DateTime(timezone=True), onupdate=func.now()),
  54. ) # pylint: disable=not-callable
  55. user_id: Optional[str] = None
  56. version: Optional[str] = "0.0.1"
  57. team_id: Optional[int] = Field(default=None, sa_column=Column(Integer, ForeignKey("team.id", ondelete="CASCADE")))
  58. name: Optional[str] = None
  59. class RunStatus(str, Enum):
  60. CREATED = "created"
  61. ACTIVE = "active"
  62. COMPLETE = "complete"
  63. ERROR = "error"
  64. STOPPED = "stopped"
  65. class Run(SQLModel, table=True):
  66. """Represents a single execution run within a session"""
  67. __table_args__ = {"sqlite_autoincrement": True}
  68. id: UUID = Field(default_factory=uuid4, primary_key=True, index=True)
  69. created_at: datetime = Field(
  70. default_factory=datetime.now, sa_column=Column(DateTime(timezone=True), server_default=func.now())
  71. )
  72. updated_at: datetime = Field(
  73. default_factory=datetime.now, sa_column=Column(DateTime(timezone=True), onupdate=func.now())
  74. )
  75. session_id: Optional[int] = Field(
  76. default=None, sa_column=Column(Integer, ForeignKey("session.id", ondelete="CASCADE"), nullable=False)
  77. )
  78. status: RunStatus = Field(default=RunStatus.CREATED)
  79. # Store the original user task
  80. task: Union[MessageConfig, dict] = Field(default_factory=MessageConfig, sa_column=Column(JSON))
  81. # Store TeamResult which contains TaskResult
  82. team_result: Union[TeamResult, dict] = Field(default=None, sa_column=Column(JSON))
  83. error_message: Optional[str] = None
  84. version: Optional[str] = "0.0.1"
  85. messages: Union[List[Message], List[dict]] = Field(default_factory=list, sa_column=Column(JSON))
  86. model_config = ConfigDict(json_encoders={UUID: str, datetime: lambda v: v.isoformat()})