config.py 2.2 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667
  1. """
  2. HealthAgent 核心配置模块
  3. """
  4. from dataclasses import dataclass, field
  5. from typing import Optional
  6. import os
  7. from dotenv import load_dotenv
  8. load_dotenv()
  9. # ========== LLM ==========
  10. @dataclass
  11. class LLMConfig:
  12. model_name: str = field(
  13. default_factory=lambda: os.getenv("OPENAI_MODEL_ID", "qwen-turbo")
  14. )
  15. api_key: Optional[str] = field(
  16. default_factory=lambda: os.getenv("OPENAI_API_KEY")
  17. )
  18. base_url: Optional[str] = field(
  19. default_factory=lambda: os.getenv("OPENAI_BASE_URL")
  20. )
  21. temperature: float = 0.7
  22. max_tokens: int = 2048
  23. timeout: int = 60
  24. # ========== Agent ==========
  25. @dataclass
  26. class AgentConfig:
  27. max_steps: int = 5
  28. timeout: int = 300
  29. history_limit: int = 50
  30. # ========== RAG ==========
  31. @dataclass
  32. class RAGConfig:
  33. enabled: bool = field(
  34. default_factory=lambda: os.getenv("RAG_ENABLED", "false").lower() in ("1", "true", "yes")
  35. )
  36. top_k: int = field(default_factory=lambda: int(os.getenv("RAG_TOP_K", "5")))
  37. milvus_uri: str = field(default_factory=lambda: os.getenv("MILVUS_URI", "http://127.0.0.1:19530"))
  38. milvus_token: Optional[str] = field(default_factory=lambda: os.getenv("MILVUS_TOKEN"))
  39. milvus_collection: str = field(default_factory=lambda: os.getenv("MILVUS_COLLECTION", "health_memory_chunks"))
  40. embedding_model: str = field(default_factory=lambda: os.getenv("EMBEDDING_MODEL", "text-embedding-v1"))
  41. embedding_api_key: Optional[str] = field(default_factory=lambda: os.getenv("EMBEDDING_API_KEY"))
  42. embedding_base_url: Optional[str] = field(default_factory=lambda: os.getenv("EMBEDDING_BASE_URL"))
  43. fallback_embedding_dim: int = field(default_factory=lambda: int(os.getenv("RAG_FALLBACK_EMBED_DIM", "64")))
  44. # ========== App ==========
  45. @dataclass
  46. class AppConfig:
  47. app_name: str = "HealthRecordAgent"
  48. debug: bool = False
  49. log_level: str = "INFO"
  50. # ========== Main ==========
  51. @dataclass
  52. class HealthAgentConfig:
  53. app: AppConfig = field(default_factory=AppConfig)
  54. llm: LLMConfig = field(default_factory=LLMConfig)
  55. agent: AgentConfig = field(default_factory=AgentConfig)
  56. rag: RAGConfig = field(default_factory=RAGConfig)
  57. # 全局配置
  58. _config = HealthAgentConfig()
  59. def get_config() -> HealthAgentConfig:
  60. return _config