| import logging |
| import os |
| from typing import Optional |
| from datetime import datetime, UTC |
|
|
| from fastapi import Depends, HTTPException, Request, status |
| from fastapi.security import APIKeyHeader |
|
|
| from src.db.init_db import get_session |
| from src.db.schemas.models import User as DBUser, AgentTemplate, UserTemplatePreference |
| from src.schemas.user_schema import User |
| from src.utils.logger import Logger |
|
|
| logger = Logger("user_manager", see_time=True, console_log=False) |
|
|
| |
| API_KEY_NAME = "X-API-Key" |
| api_key_header = APIKeyHeader(name=API_KEY_NAME, auto_error=False) |
|
|
|
|
| async def get_current_user( |
| request: Request, |
| api_key: Optional[str] = Depends(api_key_header) |
| ) -> Optional[User]: |
| """ |
| Dependency to get the current authenticated user. |
| Returns None if no user is authenticated. |
| """ |
| |
| |
| |
| |
|
|
| if not api_key or not isinstance(api_key, str): |
| |
| api_key = request.headers.get(API_KEY_NAME) or request.query_params.get("api_key") |
|
|
| |
| if not api_key: |
| return None |
| |
| try: |
| |
| |
| session = get_session() |
| |
| try: |
| |
| |
| try: |
| |
| if isinstance(api_key, str): |
| user_id = int(api_key) |
| db_user = session.query(DBUser).filter(DBUser.user_id == user_id).first() |
| else: |
| |
| logger.log_message("API key is not a string", level=logging.ERROR) |
| return None |
| except ValueError: |
| |
| logger.log_message(f"API key is not a number: {api_key}", level=logging.ERROR) |
| db_user = session.query(DBUser).filter(DBUser.username == api_key).first() |
| |
| if not db_user: |
| logger.log_message("User not found", level=logging.ERROR) |
| return None |
| |
| return User( |
| user_id=db_user.user_id, |
| username=db_user.username, |
| email=db_user.email |
| ) |
| |
| finally: |
| session.close() |
| |
| except Exception as e: |
| logger.log_message(f"Error authenticating user: {str(e)}", level=logging.ERROR) |
| return None |
|
|
| |
| def create_user(username: str, email: str) -> User: |
| """Create a new user in the database""" |
| session = get_session() |
| try: |
| |
| existing_user = session.query(DBUser).filter(DBUser.email == email).first() |
| if existing_user: |
| return User( |
| user_id=existing_user.user_id, |
| username=existing_user.username, |
| email=existing_user.email |
| ) |
| |
| |
| new_user = DBUser( |
| username=username, |
| email=email |
| ) |
| session.add(new_user) |
| session.commit() |
| session.refresh(new_user) |
| |
| |
| _enable_default_agents_for_user(new_user.user_id, session) |
| |
| return User( |
| user_id=new_user.user_id, |
| username=new_user.username, |
| email=new_user.email |
| ) |
| |
| except Exception as e: |
| session.rollback() |
| logger.log_message(f"Error creating user: {str(e)}", logging.ERROR) |
| raise |
| |
| finally: |
| session.close() |
|
|
| def get_user_by_email(email: str) -> Optional[User]: |
| """Get a user by email""" |
| session = get_session() |
| try: |
| user = session.query(DBUser).filter(DBUser.email == email).first() |
| if user is None: |
| return None |
| return User( |
| user_id=user.user_id, |
| username=user.username, |
| email=user.email |
| ) |
| except Exception as e: |
| logger.log_message(f"Error getting user by email: {str(e)}", logging.ERROR) |
| return None |
| finally: |
| session.close() |
|
|
| def _enable_default_agents_for_user(user_id: int, session): |
| """Enable default agents for a new user""" |
| try: |
| |
| default_agent_names = [ |
| "preprocessing_agent", |
| "statistical_analytics_agent", |
| "sk_learn_agent", |
| "data_viz_agent" |
| ] |
| |
| |
| default_agents = session.query(AgentTemplate).filter( |
| AgentTemplate.template_name.in_(default_agent_names), |
| AgentTemplate.is_active == True |
| ).all() |
| |
| |
| for agent in default_agents: |
| |
| existing_pref = session.query(UserTemplatePreference).filter( |
| UserTemplatePreference.user_id == user_id, |
| UserTemplatePreference.template_id == agent.template_id |
| ).first() |
| |
| if not existing_pref: |
| |
| new_pref = UserTemplatePreference( |
| user_id=user_id, |
| template_id=agent.template_id, |
| is_enabled=True, |
| usage_count=0, |
| created_at=datetime.now(UTC), |
| updated_at=datetime.now(UTC) |
| ) |
| session.add(new_pref) |
| |
| session.commit() |
| logger.log_message(f"Enabled {len(default_agents)} default agents for user {user_id}", level=logging.INFO) |
| |
| except Exception as e: |
| session.rollback() |
| logger.log_message(f"Error enabling default agents for user {user_id}: {str(e)}", level=logging.ERROR) |
| raise |
|
|