__init__.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268
  1. from typing import List, Optional, Any, Union, Dict
  2. from pydantic import BaseModel, ConfigDict, field_validator
  3. import re
  4. from datetime import datetime
  5. import json
  6. # --- User Schemas ---
  7. class UserBase(BaseModel):
  8. username: str
  9. email: Optional[str] = None
  10. role: Optional[str] = "user"
  11. class UserCreate(UserBase):
  12. password: str
  13. @field_validator("email")
  14. @classmethod
  15. def validate_email(cls, value: Optional[str]) -> Optional[str]:
  16. if value is None:
  17. return None
  18. normalized = value.strip().lower()
  19. if not re.fullmatch(r"[^\s@]+@[^\s@]+\.[^\s@]+", normalized):
  20. raise ValueError("请输入有效的邮箱地址")
  21. return normalized
  22. class UserResponse(UserBase):
  23. id: int
  24. created_at: datetime
  25. model_config = ConfigDict(from_attributes=True)
  26. class Token(BaseModel):
  27. access_token: str
  28. token_type: str
  29. class TokenData(BaseModel):
  30. username: Optional[str] = None
  31. # --- Persona Schemas ---
  32. class PersonaBase(BaseModel):
  33. name: str
  34. title: Optional[str] = None
  35. bio: Optional[str] = None
  36. theories: Optional[List[str]] = []
  37. stance: Optional[str] = None
  38. system_prompt: Optional[str] = None
  39. is_public: bool = False
  40. @field_validator('name')
  41. @classmethod
  42. def validate_name(cls, value: str) -> str:
  43. value = value.strip()
  44. if not value:
  45. raise ValueError('Persona name must not be blank')
  46. return value
  47. class PersonaCreate(PersonaBase):
  48. pass
  49. class PersonaUpdate(BaseModel):
  50. name: Optional[str] = None
  51. title: Optional[str] = None
  52. bio: Optional[str] = None
  53. theories: Optional[List[str]] = None
  54. stance: Optional[str] = None
  55. system_prompt: Optional[str] = None
  56. is_public: Optional[bool] = None
  57. @field_validator('name')
  58. @classmethod
  59. def validate_name(cls, value: Optional[str]) -> Optional[str]:
  60. if value is None:
  61. return value
  62. value = value.strip()
  63. if not value:
  64. raise ValueError('Persona name must not be blank')
  65. return value
  66. class PersonaResponse(PersonaBase):
  67. id: int
  68. owner_id: int
  69. created_at: datetime
  70. theories: Optional[Union[List[str], str]] = []
  71. model_config = ConfigDict(from_attributes=True)
  72. @field_validator('theories', mode='before')
  73. @classmethod
  74. def parse_theories(cls, v: Any) -> List[str]:
  75. if isinstance(v, str):
  76. try:
  77. parsed = json.loads(v)
  78. if isinstance(parsed, list):
  79. return parsed
  80. return []
  81. except json.JSONDecodeError:
  82. return []
  83. elif v is None:
  84. return []
  85. return v
  86. # --- Moderator Schemas ---
  87. class ModeratorBase(BaseModel):
  88. name: str
  89. title: Optional[str] = "主持人"
  90. bio: Optional[str] = None
  91. system_prompt: Optional[str] = None
  92. greeting_template: Optional[str] = None
  93. closing_template: Optional[str] = None
  94. summary_template: Optional[str] = None
  95. class ModeratorCreate(ModeratorBase):
  96. pass
  97. class ModeratorUpdate(ModeratorBase):
  98. pass
  99. class ModeratorResponse(ModeratorBase):
  100. id: int
  101. creator_id: int
  102. created_at: datetime
  103. model_config = ConfigDict(from_attributes=True)
  104. from .system_log import SystemLogCreate, SystemLogResponse
  105. # --- Forum Schemas ---
  106. class ForumBase(BaseModel):
  107. topic: str
  108. @field_validator('topic')
  109. @classmethod
  110. def validate_topic(cls, value: str) -> str:
  111. value = value.strip()
  112. if not value:
  113. raise ValueError('讨论主题不能为空')
  114. if len(value) > 200:
  115. raise ValueError('讨论主题不能超过 200 个字符')
  116. return value
  117. class ForumCreate(ForumBase):
  118. participant_ids: List[int]
  119. moderator_id: Optional[int] = None # Optional for backward compatibility (can use default)
  120. duration_minutes: int = 30
  121. @field_validator('participant_ids')
  122. @classmethod
  123. def validate_participants(cls, value: List[int]) -> List[int]:
  124. unique_ids = list(dict.fromkeys(value))
  125. if not unique_ids:
  126. raise ValueError('请至少选择一位智能体')
  127. if len(unique_ids) > 5:
  128. raise ValueError('每个论坛最多选择 5 位智能体')
  129. if any(persona_id <= 0 for persona_id in unique_ids):
  130. raise ValueError('智能体编号无效')
  131. return unique_ids
  132. @field_validator('duration_minutes')
  133. @classmethod
  134. def validate_duration(cls, value: int) -> int:
  135. if value < 1 or value > 120:
  136. raise ValueError('论坛时长必须在 1 到 120 分钟之间')
  137. return value
  138. class ForumParticipantResponse(BaseModel):
  139. persona_id: int
  140. thoughts_history: Optional[Union[List[Any], str]] = [] # Changed from List[str] to List[Any] to support dicts
  141. persona: Optional[PersonaResponse] = None
  142. model_config = ConfigDict(from_attributes=True)
  143. @field_validator('thoughts_history', mode='before')
  144. @classmethod
  145. def parse_thoughts_history(cls, v: Any) -> List[Any]:
  146. if isinstance(v, str):
  147. try:
  148. parsed = json.loads(v)
  149. if isinstance(parsed, list):
  150. return parsed
  151. # If it's a dict (single thought), wrap in list? Or return empty?
  152. # Based on log, it seems to be a list of dicts.
  153. return []
  154. except json.JSONDecodeError:
  155. return []
  156. elif isinstance(v, list):
  157. return v
  158. elif v is None:
  159. return []
  160. return [v] if v else []
  161. class ForumResponse(ForumBase):
  162. id: int
  163. creator_id: int
  164. moderator_id: Optional[int] = None
  165. status: str
  166. start_time: Optional[datetime] = None
  167. end_time: Optional[datetime] = None
  168. duration_minutes: Optional[int] = 30
  169. summary_history: Optional[Union[List[Any], str]] = [] # Changed to List[Any] for flexibility
  170. ablation_flags: Optional[Dict[str, bool]] = {}
  171. participants: Optional[List[ForumParticipantResponse]] = []
  172. moderator: Optional[ModeratorResponse] = None # Include moderator info
  173. model_config = ConfigDict(from_attributes=True)
  174. @field_validator('summary_history', mode='before')
  175. @classmethod
  176. def parse_summary_history(cls, v: Any) -> List[Any]:
  177. if isinstance(v, str):
  178. try:
  179. parsed = json.loads(v)
  180. if isinstance(parsed, list):
  181. return parsed
  182. return []
  183. except json.JSONDecodeError:
  184. return []
  185. elif isinstance(v, list):
  186. return v
  187. elif v is None:
  188. return []
  189. return [v] if v else []
  190. @field_validator('ablation_flags', mode='before')
  191. @classmethod
  192. def parse_ablation_flags(cls, v: Any) -> Dict[str, bool]:
  193. if isinstance(v, str):
  194. try:
  195. parsed = json.loads(v)
  196. return parsed if isinstance(parsed, dict) else {}
  197. except json.JSONDecodeError:
  198. return {}
  199. return v if isinstance(v, dict) else {}
  200. # --- Message Schemas ---
  201. class MessageBase(BaseModel):
  202. speaker_name: str
  203. content: str
  204. thought: Optional[str] = None # Added thought field
  205. turn_count: int = 0
  206. class MessageCreate(MessageBase):
  207. forum_id: int
  208. persona_id: Optional[int] = None
  209. moderator_id: Optional[int] = None
  210. class MessageResponse(MessageBase):
  211. id: int
  212. forum_id: int
  213. persona_id: Optional[int]
  214. moderator_id: Optional[int] = None
  215. timestamp: datetime
  216. thought: Optional[str] = None # Ensure it's in response
  217. model_config = ConfigDict(from_attributes=True)
  218. class TriggerAgentRequest(BaseModel):
  219. persona_id: Optional[int] = None
  220. class TriggerModeratorRequest(BaseModel):
  221. action: str = "auto" # auto, opening, summary, closing
  222. class GodGenerateRequest(BaseModel):
  223. prompt: str
  224. n: int = 1
  225. class ForumStartRequest(BaseModel):
  226. ablation_flags: Optional[Dict[str, bool]] = None