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.1 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193
  1. # api/deps.py
  2. import logging
  3. from contextlib import contextmanager
  4. from pathlib import Path
  5. from typing import Optional
  6. from fastapi import Depends, HTTPException, status
  7. from ..database import DatabaseManager
  8. from ..teammanager import TeamManager
  9. from .config import settings
  10. from .managers.connection import WebSocketManager
  11. logger = logging.getLogger(__name__)
  12. # Global manager instances
  13. _db_manager: Optional[DatabaseManager] = None
  14. _websocket_manager: Optional[WebSocketManager] = None
  15. _team_manager: Optional[TeamManager] = None
  16. # Context manager for database sessions
  17. @contextmanager
  18. def get_db_context():
  19. """Provide a transactional scope around a series of operations."""
  20. if not _db_manager:
  21. raise HTTPException(
  22. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Database manager not initialized"
  23. )
  24. try:
  25. yield _db_manager
  26. except Exception as e:
  27. logger.error(f"Database operation failed: {str(e)}")
  28. raise HTTPException(
  29. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Database operation failed"
  30. ) from e
  31. # Dependency providers
  32. async def get_db() -> DatabaseManager:
  33. """Dependency provider for database manager"""
  34. if not _db_manager:
  35. raise HTTPException(
  36. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Database manager not initialized"
  37. )
  38. return _db_manager
  39. async def get_websocket_manager() -> WebSocketManager:
  40. """Dependency provider for connection manager"""
  41. if not _websocket_manager:
  42. raise HTTPException(
  43. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Connection manager not initialized"
  44. )
  45. return _websocket_manager
  46. async def get_team_manager() -> TeamManager:
  47. """Dependency provider for team manager"""
  48. if not _team_manager:
  49. raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Team manager not initialized")
  50. return _team_manager
  51. # Authentication dependency
  52. async def get_current_user(
  53. # Add your authentication logic here
  54. # For example: token: str = Depends(oauth2_scheme)
  55. ) -> str:
  56. """
  57. Dependency for getting the current authenticated user.
  58. Replace with your actual authentication logic.
  59. """
  60. # Implement your user authentication here
  61. return "user_id" # Replace with actual user identification
  62. # Manager initialization and cleanup
  63. async def init_managers(database_uri: str, config_dir: str | Path, app_root: str | Path) -> None:
  64. """Initialize all manager instances"""
  65. global _db_manager, _websocket_manager, _team_manager
  66. logger.info("Initializing managers...")
  67. try:
  68. # Initialize database manager
  69. _db_manager = DatabaseManager(engine_uri=database_uri, base_dir=app_root)
  70. _db_manager.initialize_database(auto_upgrade=settings.UPGRADE_DATABASE)
  71. # init default team config
  72. await _db_manager.import_teams_from_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. manager_status = get_manager_status() # Different name
  131. missing = [name for name in manager_names if not manager_status.get(f"{name}_manager")]
  132. if missing:
  133. raise HTTPException(
  134. status_code=status.HTTP_503_SERVICE_UNAVAILABLE, # Now this refers to the imported module
  135. detail=f"Required managers not available: {', '.join(missing)}",
  136. )
  137. return True
  138. return Depends(dependency)