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.

deps.py 6.0 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194
  1. # api/deps.py
  2. import logging
  3. from contextlib import contextmanager
  4. from typing import Optional
  5. from fastapi import Depends, HTTPException, status
  6. from ..database import ConfigurationManager, DatabaseManager
  7. from ..teammanager import TeamManager
  8. from .config import settings
  9. from .managers.connection import WebSocketManager
  10. logger = logging.getLogger(__name__)
  11. # Global manager instances
  12. _db_manager: Optional[DatabaseManager] = None
  13. _websocket_manager: Optional[WebSocketManager] = None
  14. _team_manager: Optional[TeamManager] = None
  15. # Context manager for database sessions
  16. @contextmanager
  17. def get_db_context():
  18. """Provide a transactional scope around a series of operations."""
  19. if not _db_manager:
  20. raise HTTPException(
  21. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Database manager not initialized"
  22. )
  23. try:
  24. yield _db_manager
  25. except Exception as e:
  26. logger.error(f"Database operation failed: {str(e)}")
  27. raise HTTPException(
  28. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Database operation failed"
  29. ) from e
  30. # Dependency providers
  31. async def get_db() -> DatabaseManager:
  32. """Dependency provider for database manager"""
  33. if not _db_manager:
  34. raise HTTPException(
  35. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Database manager not initialized"
  36. )
  37. return _db_manager
  38. async def get_websocket_manager() -> WebSocketManager:
  39. """Dependency provider for connection manager"""
  40. if not _websocket_manager:
  41. raise HTTPException(
  42. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Connection manager not initialized"
  43. )
  44. return _websocket_manager
  45. async def get_team_manager() -> TeamManager:
  46. """Dependency provider for team manager"""
  47. if not _team_manager:
  48. raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Team manager not initialized")
  49. return _team_manager
  50. # Authentication dependency
  51. async def get_current_user(
  52. # Add your authentication logic here
  53. # For example: token: str = Depends(oauth2_scheme)
  54. ) -> str:
  55. """
  56. Dependency for getting the current authenticated user.
  57. Replace with your actual authentication logic.
  58. """
  59. # Implement your user authentication here
  60. return "user_id" # Replace with actual user identification
  61. # Manager initialization and cleanup
  62. async def init_managers(database_uri: str, config_dir: str, app_root: str) -> None:
  63. """Initialize all manager instances"""
  64. global _db_manager, _websocket_manager, _team_manager
  65. logger.info("Initializing managers...")
  66. try:
  67. # Initialize database manager
  68. _db_manager = DatabaseManager(engine_uri=database_uri, base_dir=app_root)
  69. _db_manager.initialize_database(auto_upgrade=settings.UPGRADE_DATABASE)
  70. # init default team config
  71. _team_config_manager = ConfigurationManager(db_manager=_db_manager)
  72. await _team_config_manager.import_directory(config_dir, settings.DEFAULT_USER_ID, check_exists=True)
  73. # Initialize connection manager
  74. _websocket_manager = WebSocketManager(db_manager=_db_manager)
  75. logger.info("Connection manager initialized")
  76. # Initialize team manager
  77. _team_manager = TeamManager()
  78. logger.info("Team manager initialized")
  79. except Exception as e:
  80. logger.error(f"Failed to initialize managers: {str(e)}")
  81. await cleanup_managers() # Cleanup any partially initialized managers
  82. raise
  83. async def cleanup_managers() -> None:
  84. """Cleanup and shutdown all manager instances"""
  85. global _db_manager, _websocket_manager, _team_manager
  86. logger.info("Cleaning up managers...")
  87. # Cleanup connection manager first to ensure all active connections are closed
  88. if _websocket_manager:
  89. try:
  90. await _websocket_manager.cleanup()
  91. except Exception as e:
  92. logger.error(f"Error cleaning up connection manager: {str(e)}")
  93. finally:
  94. _websocket_manager = None
  95. # TeamManager doesn't need explicit cleanup since WebSocketManager handles it
  96. _team_manager = None
  97. # Cleanup database manager last
  98. if _db_manager:
  99. try:
  100. await _db_manager.close()
  101. except Exception as e:
  102. logger.error(f"Error cleaning up database manager: {str(e)}")
  103. finally:
  104. _db_manager = None
  105. logger.info("All managers cleaned up")
  106. # Utility functions for dependency management
  107. def get_manager_status() -> dict:
  108. """Get the initialization status of all managers"""
  109. return {
  110. "database_manager": _db_manager is not None,
  111. "websocket_manager": _websocket_manager is not None,
  112. "team_manager": _team_manager is not None,
  113. }
  114. # Combined dependencies
  115. async def get_managers():
  116. """Get all managers in one dependency"""
  117. return {"db": await get_db(), "connection": await get_websocket_manager(), "team": await get_team_manager()}
  118. # Error handling for manager operations
  119. class ManagerOperationError(Exception):
  120. """Custom exception for manager operation errors"""
  121. def __init__(self, manager_name: str, operation: str, detail: str):
  122. self.manager_name = manager_name
  123. self.operation = operation
  124. self.detail = detail
  125. super().__init__(f"{manager_name} failed during {operation}: {detail}")
  126. # Dependency for requiring specific managers
  127. def require_managers(*manager_names: str):
  128. """Decorator to require specific managers for a route"""
  129. async def dependency():
  130. status = get_manager_status()
  131. missing = [name for name in manager_names if not status.get(f"{name}_manager")]
  132. if missing:
  133. raise HTTPException(
  134. status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
  135. detail=f"Required managers not available: {', '.join(missing)}",
  136. )
  137. return True
  138. return Depends(dependency)