Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 13 additions & 21 deletions app/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,11 +36,10 @@
from app.core.utils.log import LogConfig
from app.dependencies import (
disconnect_state,
get_app_state,
get_db,
get_notification_manager,
get_redis_client,
init_app_state,
init_state,
)
from app.module import all_modules, module_list
from app.types.exceptions import (
Expand Down Expand Up @@ -494,11 +493,11 @@ async def init_lifespan(
) -> LifespanState:
hyperion_error_logger.info("Startup: Initializing application")

# We get `init_app_state` as a dependency, as tests
# We get `init_state` as a dependency, as tests
# should override it to provide their own state
state: LifespanState = await app.dependency_overrides.get(
init_app_state,
init_app_state,
await app.dependency_overrides.get(
init_state,
init_state,
)(
app=app,
settings=settings,
Expand All @@ -508,7 +507,7 @@ async def init_lifespan(
redis_client: Redis | None = app.dependency_overrides.get(
get_redis_client,
get_redis_client,
)(state=state)
)()

# Initialization steps should only be run once across all workers
# We use Redis locks to ensure that the initialization steps are only run once
Expand Down Expand Up @@ -544,14 +543,14 @@ async def init_lifespan(
)

get_db_dependency: Callable[
[LifespanState],
[],
AsyncGenerator[AsyncSession, None],
] = app.dependency_overrides.get(
get_db,
get_db,
)
# We need to run the factories only once across all the workers
async for db in get_db_dependency(state):
async for db in get_db_dependency():
await initialization.use_lock_for_workers(
run_factories,
"run_factories",
Expand All @@ -562,7 +561,7 @@ async def init_lifespan(
settings=settings,
hyperion_error_logger=hyperion_error_logger,
)
async for db in get_db_dependency(state):
async for db in get_db_dependency():
await initialization.use_lock_for_workers(
init_google_API,
"init_google_API",
Expand All @@ -573,11 +572,11 @@ async def init_lifespan(
settings=settings,
)

async for db in get_db_dependency(state):
async for db in get_db_dependency():
notification_manager = app.dependency_overrides.get(
get_notification_manager,
get_notification_manager,
)(state)
)()
await initialization.use_lock_for_workers(
initialize_notification_topics,
"initialize_notification_topics",
Expand All @@ -589,7 +588,7 @@ async def init_lifespan(
notification_manager=notification_manager,
)

return state
return LifespanState()


# We wrap the application in a function to be able to pass the settings and drop_db parameters
Expand Down Expand Up @@ -620,7 +619,6 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[LifespanState, None]:
disconnect_state,
disconnect_state,
)(
state=state,
hyperion_error_logger=hyperion_error_logger,
)

Expand All @@ -644,10 +642,6 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[LifespanState, None]:
calypsso = get_calypsso_app()
app.mount("/calypsso", calypsso, "Calypsso")

get_app_state_dependency = app.dependency_overrides.get(
get_app_state,
get_app_state,
)
get_redis_client_dependency = app.dependency_overrides.get(
get_redis_client,
get_redis_client,
Expand Down Expand Up @@ -683,9 +677,7 @@ async def logging_middleware(
port = request.client.port
client_address = f"{ip_address}:{port}"

redis_client: redis.Redis | None = get_redis_client_dependency(
state=get_app_state_dependency(request),
)
redis_client: redis.Redis | None = get_redis_client_dependency()

# We test the ip address with the redis limiter
process = True
Expand Down
4 changes: 2 additions & 2 deletions app/core/notification/endpoints_notification.py
Original file line number Diff line number Diff line change
Expand Up @@ -299,7 +299,7 @@ async def send_test_future_notification(
message=message,
defer_date=datetime.now(UTC) + timedelta(seconds=10),
scheduler=scheduler,
job_id="testtt",
job_id="send_test_future_notification",
)


Expand Down Expand Up @@ -350,7 +350,7 @@ async def send_test_future_notification_topic(
topic_id=notification_test_topic.id,
message=message,
defer_date=datetime.now(UTC) + timedelta(seconds=10),
job_id="test26",
job_id="notification_test_future",
scheduler=scheduler,
)

Expand Down
75 changes: 37 additions & 38 deletions app/dependencies.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ async def get_users(db: AsyncSession = Depends(get_db)):
from app.utils.auth import auth_utils
from app.utils.communication.notifications import NotificationManager, NotificationTool
from app.utils.state import (
LifespanState,
GlobalState,
RuntimeLifespanState,
disconnect_redis_client,
disconnect_scheduler,
Expand All @@ -61,28 +61,31 @@ async def get_users(db: AsyncSession = Depends(get_db)):
hyperion_access_logger = logging.getLogger("hyperion.access")
hyperion_error_logger = logging.getLogger("hyperion.error")

GLOBAL_STATE: GlobalState

async def init_app_state(

async def init_state(
app: FastAPI,
settings: Settings,
hyperion_error_logger: logging.Logger,
) -> LifespanState:
) -> None:
"""
Initialize the state of the application. This dependency should be used at the start of the application lifespan.
Initialize the global state for the project. This dependency should be called at the start of the application lifespan.

This methode should be called as a dependency, and test may override it to provide their own state.
```python
state = app.dependency_overrides.get(
init_app_state,
init_app_state,
app.dependency_overrides.get(
init_state,
init_state,
)(
app=app,
settings=settings,
hyperion_error_logger=hyperion_error_logger,
)
state = cast("LifespanState", state)
```
"""
global GLOBAL_STATE

engine = init_engine(settings=settings)

SessionLocal = init_SessionLocal(engine)
Expand All @@ -94,7 +97,7 @@ async def init_app_state(

scheduler = await init_scheduler(
settings=settings,
app=app,
_dependency_overrides=app.dependency_overrides,
)

ws_manager = await init_websocket_connection_manager(
Expand All @@ -112,7 +115,7 @@ async def init_app_state(

mail_templates = init_mail_templates(settings=settings)

return LifespanState(
GLOBAL_STATE = GlobalState(
engine=engine,
SessionLocal=SessionLocal,
redis_client=redis_client,
Expand All @@ -126,17 +129,17 @@ async def init_app_state(


async def disconnect_state(
state: LifespanState,
hyperion_error_logger: logging.Logger,
) -> None:
"""
Disconnect items requiring it. This dependency should be used at the end of the application lifespan.

This methode should be called as a dependency as test may need to run additional steps
This methode should be called as a dependency as tests may need to run additional steps
"""
disconnect_redis_client(state["redis_client"])
await disconnect_scheduler(state["scheduler"])
await disconnect_websocket_connection_manager(state["ws_manager"])

disconnect_redis_client(GLOBAL_STATE["redis_client"])
await disconnect_scheduler(GLOBAL_STATE["scheduler"])
await disconnect_websocket_connection_manager(GLOBAL_STATE["ws_manager"])

hyperion_error_logger.info("Application state disconnected successfully.")

Expand Down Expand Up @@ -179,7 +182,7 @@ def get_settings() -> Settings:
return construct_prod_settings()


async def get_db(state: AppState) -> AsyncGenerator[AsyncSession, None]:
async def get_db() -> AsyncGenerator[AsyncSession, None]:
"""
Return a database session that will be automatically committed and closed after usage.

Expand All @@ -202,7 +205,7 @@ async def get_db(state: AppState) -> AsyncGenerator[AsyncSession, None]:
# Add objects that may be rolled back in case of an error here
```
"""
async with state["SessionLocal"]() as db:
async with GLOBAL_STATE["SessionLocal"]() as db:
try:
yield db
except HTTPException:
Expand All @@ -217,42 +220,42 @@ async def get_db(state: AppState) -> AsyncGenerator[AsyncSession, None]:
await db.close()


async def get_unsafe_db(state: AppState) -> AsyncGenerator[AsyncSession, None]:
async def get_unsafe_db() -> AsyncGenerator[AsyncSession, None]:
"""
Return a database session but don't close it automatically

It should only be used for really specific cases where `get_db` will not work
"""

async with state["SessionLocal"]() as db:
async with GLOBAL_STATE["SessionLocal"]() as db:
yield db


def get_redis_client(state: AppState) -> redis.Redis | None:
def get_redis_client() -> redis.Redis | None:
"""
Dependency that returns the redis client

If the redis client is not available, it will return None.
"""
return state["redis_client"]
return GLOBAL_STATE["redis_client"]


def get_scheduler(state: AppState) -> Scheduler:
return state["scheduler"]
def get_scheduler() -> Scheduler:
return GLOBAL_STATE["scheduler"]


def get_websocket_connection_manager(state: AppState) -> WebsocketConnectionManager:
return state["ws_manager"]
def get_websocket_connection_manager() -> WebsocketConnectionManager:
return GLOBAL_STATE["ws_manager"]


def get_notification_manager(state: AppState) -> NotificationManager:
def get_notification_manager() -> NotificationManager:
"""
Dependency that returns the notification manager.
This dependency provide a low level tool allowing to use notification manager internal methods.

If you want to send a notification, prefer `get_notification_tool` dependency.
"""
return state["notification_manager"]
return GLOBAL_STATE["notification_manager"]


def get_notification_tool(
Expand All @@ -271,22 +274,20 @@ def get_notification_tool(
)


def get_drive_file_manager(state: AppState) -> DriveFileManager:
def get_drive_file_manager() -> DriveFileManager:
"""
Dependency that returns the drive file manager.
"""

return state["drive_file_manager"]
return GLOBAL_STATE["drive_file_manager"]


@lru_cache
def get_payment_tool(
name: HelloAssoConfigName,
) -> Callable[[AppState], PaymentTool]:
def get_payment_tool(
state: AppState,
) -> PaymentTool:
payment_tools = state["payment_tools"]
) -> Callable[[], PaymentTool]:
def get_payment_tool() -> PaymentTool:
payment_tools = GLOBAL_STATE["payment_tools"]
if name not in payment_tools:
hyperion_error_logger.warning(
f"HelloAsso API credentials are not set for {name.value}, payment won't be available",
Expand All @@ -298,14 +299,12 @@ def get_payment_tool(
return get_payment_tool


def get_mail_templates(
state: AppState,
) -> calypsso.MailTemplates:
def get_mail_templates() -> calypsso.MailTemplates:
"""
Dependency that returns the mail templates manager.
"""

return state["mail_templates"]
return GLOBAL_STATE["mail_templates"]


def get_token_data(
Expand Down
Loading