base_router.py 6.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214
  1. import functools
  2. import logging
  3. from abc import abstractmethod
  4. from typing import Callable
  5. from fastapi import APIRouter, Depends, HTTPException, Request, status
  6. from fastapi.responses import StreamingResponse
  7. from core.base import R2RException, manage_run
  8. logger = logging.getLogger()
  9. class BaseRouterV3:
  10. def __init__(self, providers, services, orchestration_provider, run_type):
  11. self.providers = providers
  12. self.services = services
  13. self.run_type = run_type
  14. self.orchestration_provider = orchestration_provider
  15. self.router = APIRouter()
  16. self.openapi_extras = self._load_openapi_extras()
  17. self._setup_routes()
  18. self._register_workflows()
  19. def get_router(self):
  20. return self.router
  21. def base_endpoint(self, func: Callable):
  22. @functools.wraps(func)
  23. async def wrapper(*args, **kwargs):
  24. async with manage_run(
  25. self.services["ingestion"].run_manager, func.__name__
  26. ) as run_id:
  27. auth_user = kwargs.get("auth_user")
  28. if auth_user:
  29. await self.services[
  30. "ingestion"
  31. ].run_manager.log_run_info( # TODO - this is a bit of a hack
  32. run_type=self.run_type,
  33. user=auth_user,
  34. )
  35. try:
  36. func_result = await func(*args, **kwargs)
  37. if (
  38. isinstance(func_result, tuple)
  39. and len(func_result) == 2
  40. ):
  41. results, outer_kwargs = func_result
  42. else:
  43. results, outer_kwargs = func_result, {}
  44. if isinstance(results, StreamingResponse):
  45. return results
  46. return {"results": results, **outer_kwargs}
  47. except R2RException:
  48. raise
  49. except Exception as e:
  50. logger.error(
  51. f"Error in base endpoint {func.__name__}() - \n\n{str(e)}",
  52. exc_info=True,
  53. )
  54. raise HTTPException(
  55. status_code=500,
  56. detail={
  57. "message": f"An error '{e}' occurred during {func.__name__}",
  58. "error": str(e),
  59. "error_type": type(e).__name__,
  60. },
  61. ) from e
  62. return wrapper
  63. @classmethod
  64. def build_router(cls, engine):
  65. return cls(engine).router
  66. def _register_workflows(self):
  67. pass
  68. def _load_openapi_extras(self):
  69. return {}
  70. @abstractmethod
  71. def _setup_routes(self):
  72. pass
  73. import functools
  74. import logging
  75. from abc import abstractmethod
  76. from typing import Callable
  77. from fastapi import APIRouter, Depends, HTTPException, Request
  78. from fastapi.responses import StreamingResponse
  79. from core.base import R2RException, manage_run
  80. logger = logging.getLogger()
  81. class BaseRouterV3:
  82. def __init__(self, providers, services, orchestration_provider, run_type):
  83. self.providers = providers
  84. self.services = services
  85. self.run_type = run_type
  86. self.orchestration_provider = orchestration_provider
  87. self.router = APIRouter()
  88. self.openapi_extras = self._load_openapi_extras()
  89. self.set_rate_limiting()
  90. self._setup_routes()
  91. self._register_workflows()
  92. def get_router(self):
  93. return self.router
  94. def base_endpoint(self, func: Callable):
  95. @functools.wraps(func)
  96. async def wrapper(*args, **kwargs):
  97. async with manage_run(
  98. self.services["ingestion"].run_manager, func.__name__
  99. ) as run_id:
  100. auth_user = kwargs.get("auth_user")
  101. if auth_user:
  102. await self.services["ingestion"].run_manager.log_run_info(
  103. run_type=self.run_type,
  104. user=auth_user,
  105. )
  106. try:
  107. func_result = await func(*args, **kwargs)
  108. if (
  109. isinstance(func_result, tuple)
  110. and len(func_result) == 2
  111. ):
  112. results, outer_kwargs = func_result
  113. else:
  114. results, outer_kwargs = func_result, {}
  115. if isinstance(results, StreamingResponse):
  116. return results
  117. return {"results": results, **outer_kwargs}
  118. except R2RException:
  119. raise
  120. except Exception as e:
  121. logger.error(
  122. f"Error in base endpoint {func.__name__}() - \n\n{str(e)}",
  123. exc_info=True,
  124. )
  125. raise HTTPException(
  126. status_code=500,
  127. detail={
  128. "message": f"An error '{e}' occurred during {func.__name__}",
  129. "error": str(e),
  130. "error_type": type(e).__name__,
  131. },
  132. ) from e
  133. return wrapper
  134. @classmethod
  135. def build_router(cls, engine):
  136. return cls(engine).router
  137. def _register_workflows(self):
  138. pass
  139. def _load_openapi_extras(self):
  140. return {}
  141. @abstractmethod
  142. def _setup_routes(self):
  143. pass
  144. def set_rate_limiting(self):
  145. """
  146. Set up a yield dependency for rate limiting and logging.
  147. """
  148. async def rate_limit_dependency(
  149. request: Request,
  150. auth_user=Depends(self.providers.auth.auth_wrapper),
  151. ):
  152. user_id = auth_user.id
  153. route = request.scope["path"]
  154. # Check the limits before proceeding
  155. try:
  156. await self.providers.database.limits_handler.check_limits(
  157. user_id, route
  158. )
  159. except ValueError as e:
  160. raise HTTPException(status_code=429, detail=str(e))
  161. request.state.user_id = user_id
  162. request.state.route = route
  163. print("in rate limit dependency....")
  164. # Yield to run the route
  165. try:
  166. yield
  167. finally:
  168. print("finally....")
  169. # After the route completes successfully, log the request
  170. await self.providers.database.limits_handler.log_request(
  171. user_id, route
  172. )
  173. self.rate_limit_dependency = rate_limit_dependency