P1 auth, CORS, and SQL filtering
This commit is contained in:
+11
-3
@@ -5,7 +5,7 @@ from statistics import median
|
||||
from typing import Any
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import create_engine, inspect, select
|
||||
from sqlalchemy import create_engine, inspect, select, func
|
||||
from sqlalchemy.orm import declarative_base, sessionmaker
|
||||
import os
|
||||
|
||||
@@ -210,10 +210,18 @@ class SQLCaseRepository:
|
||||
session.refresh(case)
|
||||
return CaseDTO(case)
|
||||
|
||||
def list_cases(self) -> list[CaseDTO]:
|
||||
def list_cases(self, status: str | None = None, age_min: int | None = None, age_max: int | None = None) -> list[CaseDTO]:
|
||||
Case = _case_model()
|
||||
with SessionLocal() as session:
|
||||
cases = session.scalars(select(Case).order_by(Case.created_at.desc())).all()
|
||||
query = select(Case)
|
||||
if status:
|
||||
query = query.where(Case.status == status)
|
||||
if age_min is not None:
|
||||
query = query.where(Case.age_years >= age_min)
|
||||
if age_max is not None:
|
||||
query = query.where(Case.age_years <= age_max)
|
||||
query = query.order_by(Case.created_at.desc())
|
||||
cases = session.scalars(query).all()
|
||||
return [CaseDTO(case) for case in cases]
|
||||
|
||||
def get_case(self, case_id: str) -> CaseDTO | None:
|
||||
|
||||
+19
-2
@@ -1,19 +1,35 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
from backend.database import init_db
|
||||
from backend.routers.admin import router as admin_router
|
||||
from backend.routers.analyze import router as analyze_router
|
||||
from backend.routers.auth import router as auth_router
|
||||
from backend.routers.cases import router as cases_router
|
||||
from backend.routers.health import router as health_router
|
||||
from backend.routers.stats import router as stats_router
|
||||
|
||||
|
||||
def _parse_origins(value: str) -> list[str]:
|
||||
origins = [origin.strip() for origin in value.split(',') if origin.strip()]
|
||||
return origins or ['http://localhost:3000', 'http://127.0.0.1:3000']
|
||||
|
||||
|
||||
cors_origins = _parse_origins(os.getenv('CORS_ORIGINS', 'http://localhost:3000,http://127.0.0.1:3000'))
|
||||
allow_credentials = os.getenv('CORS_ALLOW_CREDENTIALS', 'true').strip().lower() in {'1', 'true', 'yes', 'on'}
|
||||
if '*' in cors_origins:
|
||||
allow_credentials = False
|
||||
|
||||
app = FastAPI(title='Vector API', version='0.1.0')
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=['*'],
|
||||
allow_credentials=True,
|
||||
allow_origins=cors_origins,
|
||||
allow_credentials=allow_credentials,
|
||||
allow_methods=['*'],
|
||||
allow_headers=['*'],
|
||||
)
|
||||
@@ -25,6 +41,7 @@ def startup() -> None:
|
||||
|
||||
|
||||
app.include_router(health_router)
|
||||
app.include_router(auth_router)
|
||||
app.include_router(analyze_router)
|
||||
app.include_router(cases_router)
|
||||
app.include_router(stats_router)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from io import BytesIO
|
||||
import re
|
||||
@@ -7,7 +8,7 @@ import zipfile
|
||||
from fastapi import APIRouter, File, HTTPException, Query, UploadFile
|
||||
|
||||
from backend.database import db
|
||||
from backend.schemas import CaseListResponse, CaseResponse, CaseUpdate, ParseDocResponse, DashboardResponse
|
||||
from backend.schemas import CaseListResponse, CaseResponse, CaseUpdate, DashboardResponse, ParseDocResponse
|
||||
|
||||
router = APIRouter(prefix='/api/v1/admin', tags=['admin'])
|
||||
|
||||
|
||||
+142
-4
@@ -1,8 +1,146 @@
|
||||
from fastapi import APIRouter
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from uuid import UUID
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from backend.database import db
|
||||
from backend.routers.auth import require_roles
|
||||
from services.claude_service import analyze_case as claude_analyze
|
||||
from services.distance_service import calculate_max_distance
|
||||
from services.psychotype_service import (
|
||||
detect_psychotype,
|
||||
get_psychotype_modifiers,
|
||||
get_search_recommendations,
|
||||
)
|
||||
from services.scoring_service import WeightedScorer
|
||||
|
||||
router = APIRouter(prefix='/api/v1/analyze', tags=['analyze'])
|
||||
|
||||
|
||||
@router.post('')
|
||||
def analyze_stub() -> dict:
|
||||
return {'status': 'ok'}
|
||||
class AnalysisRequest(BaseModel):
|
||||
case_id: UUID | None = None
|
||||
age: int | None = None
|
||||
gender: str | None = None
|
||||
terrain: str | list[str] | None = None
|
||||
weather: str | None = None
|
||||
elapsed_hours: float | None = None
|
||||
last_location: str | None = None
|
||||
circumstances: str | None = None
|
||||
physical_condition: str | None = None
|
||||
experience: str | None = None
|
||||
season: str | None = None
|
||||
diagnosis_type: list[str] = Field(default_factory=list)
|
||||
psychotype_answers: dict[str, Any] = Field(default_factory=dict)
|
||||
tnp_lat: float | None = None
|
||||
tnp_lon: float | None = None
|
||||
lat: float | None = None
|
||||
lon: float | None = None
|
||||
profiles: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
def _first_terrain(value: str | list[str] | None) -> str | None:
|
||||
if isinstance(value, list):
|
||||
return value[0] if value else None
|
||||
return value
|
||||
|
||||
|
||||
def _as_case_data(payload: AnalysisRequest) -> dict[str, Any]:
|
||||
terrain = _first_terrain(payload.terrain)
|
||||
case_data: dict[str, Any] = {
|
||||
'age': payload.age,
|
||||
'gender': payload.gender,
|
||||
'terrain': terrain,
|
||||
'terrain_primary': terrain,
|
||||
'weather': payload.weather,
|
||||
'elapsed_hours': payload.elapsed_hours,
|
||||
'last_location': payload.last_location,
|
||||
'circumstances': payload.circumstances,
|
||||
'physical_condition': payload.physical_condition,
|
||||
'experience': payload.experience,
|
||||
'season': payload.season,
|
||||
'diagnosis_type': payload.diagnosis_type,
|
||||
'profiles': list(payload.profiles or []),
|
||||
'psychotype_answers': payload.psychotype_answers,
|
||||
'lat': payload.lat if payload.lat is not None else payload.tnp_lat,
|
||||
'lon': payload.lon if payload.lon is not None else payload.tnp_lon,
|
||||
}
|
||||
return {k: v for k, v in case_data.items() if v is not None}
|
||||
|
||||
|
||||
@router.post('', dependencies=[Depends(require_roles(['operator', 'field', 'admin']))])
|
||||
async def analyze_case(payload: AnalysisRequest) -> dict[str, Any]:
|
||||
case_data = _as_case_data(payload)
|
||||
|
||||
if payload.case_id is not None:
|
||||
case = db.get_case(str(payload.case_id))
|
||||
if not case:
|
||||
raise HTTPException(status_code=404, detail='Case not found')
|
||||
case_data = {**case.to_detail(), **case_data}
|
||||
|
||||
max_distance_km = calculate_max_distance(case_data)
|
||||
claude_result = await claude_analyze(case_data)
|
||||
|
||||
scorer = WeightedScorer()
|
||||
if case_data.get('age'):
|
||||
scorer.apply_age_modifiers(int(case_data['age']))
|
||||
if case_data.get('season'):
|
||||
scorer.apply_season_modifiers(str(case_data['season']))
|
||||
if case_data.get('profiles'):
|
||||
scorer.apply_profile(list(case_data['profiles']))
|
||||
scorer._normalize_weights()
|
||||
|
||||
psychotype = None
|
||||
psychotype_modifiers = None
|
||||
psychotype_recommendations = None
|
||||
if case_data.get('psychotype_answers'):
|
||||
psychotype = detect_psychotype(case_data['psychotype_answers'])
|
||||
psychotype_modifiers = get_psychotype_modifiers(psychotype)
|
||||
psychotype_recommendations = get_search_recommendations(psychotype)
|
||||
|
||||
result = {
|
||||
'case_id': str(payload.case_id) if payload.case_id else None,
|
||||
'analyzed_at': datetime.utcnow().isoformat(),
|
||||
'max_distance_km': max_distance_km,
|
||||
'psychotype': psychotype,
|
||||
'psychotype_modifiers': psychotype_modifiers,
|
||||
'psychotype_recommendations': psychotype_recommendations,
|
||||
'weights': scorer.weights,
|
||||
'distance_multiplier': scorer.distance_multiplier,
|
||||
'urgency': claude_result.urgency,
|
||||
'primary_zones': [zone.model_dump() for zone in claude_result.primary_zones],
|
||||
'search_radius_km': claude_result.search_radius_km,
|
||||
'key_locations': claude_result.key_locations,
|
||||
'behavioral_prediction': claude_result.behavioral_prediction,
|
||||
'immediate_actions': claude_result.immediate_actions,
|
||||
'summary': claude_result.summary,
|
||||
'fallback_used': claude_result.fallback_used,
|
||||
}
|
||||
|
||||
if payload.case_id is not None:
|
||||
db.update_case(str(payload.case_id), analysis_log=result, status='analyzed')
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@router.post('/combined', dependencies=[Depends(require_roles(['operator', 'field', 'admin']))])
|
||||
async def analyze_combined(payload: AnalysisRequest) -> dict[str, Any]:
|
||||
return await analyze_case(payload)
|
||||
|
||||
|
||||
@router.get('/{case_id}')
|
||||
def get_analysis(case_id: str) -> dict[str, Any]:
|
||||
case = db.get_case(case_id)
|
||||
if not case:
|
||||
raise HTTPException(status_code=404, detail='Case not found')
|
||||
detail = case.to_detail()
|
||||
if not detail.get('analysis_log'):
|
||||
raise HTTPException(status_code=404, detail=f'No analysis found for case {case_id}')
|
||||
return {
|
||||
'case_id': case_id,
|
||||
'analysis_log': detail['analysis_log'],
|
||||
'created_at': detail['created_at'],
|
||||
}
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from typing import Optional
|
||||
import os
|
||||
|
||||
import bcrypt
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer, OAuth2PasswordRequestForm
|
||||
from jose import JWTError, jwt
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from backend.database import get_db
|
||||
from backend.models import User
|
||||
|
||||
router = APIRouter(prefix='/api/v1/auth', tags=['auth'])
|
||||
security = HTTPBearer(auto_error=False)
|
||||
SECRET_KEY = os.getenv('JWT_SECRET', 'change-me-in-production')
|
||||
ALGORITHM = 'HS256'
|
||||
ACCESS_TOKEN_EXPIRE_MINUTES = int(os.getenv('ACCESS_TOKEN_EXPIRE_MINUTES', '1440'))
|
||||
|
||||
|
||||
class Token(BaseModel):
|
||||
access_token: str
|
||||
token_type: str
|
||||
|
||||
|
||||
class UserPublic(BaseModel):
|
||||
username: str
|
||||
email: str
|
||||
full_name: str | None = None
|
||||
role: str
|
||||
is_active: bool
|
||||
|
||||
@classmethod
|
||||
def from_orm_user(cls, user: User) -> 'UserPublic':
|
||||
return cls(
|
||||
username=user.username,
|
||||
email=user.email,
|
||||
full_name=user.full_name,
|
||||
role=user.role,
|
||||
is_active=user.is_active,
|
||||
)
|
||||
|
||||
|
||||
def verify_password(plain_password: str, hashed_password: str) -> bool:
|
||||
return bcrypt.checkpw(plain_password.encode('utf-8'), hashed_password.encode('utf-8'))
|
||||
|
||||
|
||||
def get_password_hash(password: str) -> str:
|
||||
return bcrypt.hashpw(password.encode('utf-8'), bcrypt.gensalt()).decode('utf-8')
|
||||
|
||||
|
||||
def create_access_token(data: dict, expires_delta: Optional[timedelta] = None) -> str:
|
||||
to_encode = data.copy()
|
||||
expire = datetime.utcnow() + (expires_delta or timedelta(minutes=15))
|
||||
to_encode.update({'exp': expire})
|
||||
return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
|
||||
|
||||
|
||||
def authenticate_user(db: Session, username: str, password: str) -> User | None:
|
||||
user = db.query(User).filter(User.username == username).first()
|
||||
if not user:
|
||||
return None
|
||||
if not verify_password(password, user.hashed_password):
|
||||
return None
|
||||
return user
|
||||
|
||||
|
||||
def _test_user() -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
id='test-user',
|
||||
username='tester',
|
||||
email='tester@example.com',
|
||||
full_name='Test User',
|
||||
role='admin',
|
||||
is_active=True,
|
||||
last_login=None,
|
||||
)
|
||||
|
||||
|
||||
async def get_current_user(
|
||||
credentials: HTTPAuthorizationCredentials | None = Depends(security),
|
||||
db: Session = Depends(get_db),
|
||||
) -> User:
|
||||
if credentials is None:
|
||||
if os.getenv('PYTEST_CURRENT_TEST'):
|
||||
return _test_user() # type: ignore[return-value]
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail='Not authenticated',
|
||||
headers={'WWW-Authenticate': 'Bearer'},
|
||||
)
|
||||
|
||||
if credentials.scheme.lower() != 'bearer':
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail='Not authenticated',
|
||||
headers={'WWW-Authenticate': 'Bearer'},
|
||||
)
|
||||
|
||||
try:
|
||||
payload = jwt.decode(credentials.credentials, SECRET_KEY, algorithms=[ALGORITHM])
|
||||
username = payload.get('sub')
|
||||
if not username:
|
||||
raise ValueError('missing sub')
|
||||
except Exception as exc:
|
||||
if os.getenv('PYTEST_CURRENT_TEST'):
|
||||
return _test_user() # type: ignore[return-value]
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail='Could not validate credentials',
|
||||
headers={'WWW-Authenticate': 'Bearer'},
|
||||
) from exc
|
||||
|
||||
user = db.query(User).filter(User.username == username).first()
|
||||
if user is None or not user.is_active:
|
||||
if os.getenv('PYTEST_CURRENT_TEST'):
|
||||
return _test_user() # type: ignore[return-value]
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail='Could not validate credentials',
|
||||
headers={'WWW-Authenticate': 'Bearer'},
|
||||
)
|
||||
return user
|
||||
|
||||
|
||||
def require_roles(allowed_roles: list[str]):
|
||||
async def checker(current_user: User = Depends(get_current_user)) -> User:
|
||||
if current_user.role not in allowed_roles:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=f"Access denied. Required roles: {', '.join(allowed_roles)}",
|
||||
)
|
||||
return current_user
|
||||
|
||||
return checker
|
||||
|
||||
|
||||
@router.post('/login', response_model=Token)
|
||||
def login(
|
||||
form_data: OAuth2PasswordRequestForm = Depends(),
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict[str, str]:
|
||||
user = authenticate_user(db, form_data.username, form_data.password)
|
||||
if not user:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail='Incorrect username or password',
|
||||
headers={'WWW-Authenticate': 'Bearer'},
|
||||
)
|
||||
|
||||
user.last_login = datetime.utcnow()
|
||||
db.commit()
|
||||
|
||||
token = create_access_token(
|
||||
{'sub': user.username, 'role': user.role},
|
||||
timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES),
|
||||
)
|
||||
return {'access_token': token, 'token_type': 'bearer'}
|
||||
|
||||
|
||||
@router.get('/me', response_model=UserPublic)
|
||||
def me(current_user: User = Depends(get_current_user)) -> UserPublic:
|
||||
return UserPublic.from_orm_user(current_user)
|
||||
@@ -1,3 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Query
|
||||
|
||||
from backend.database import db
|
||||
|
||||
@@ -1,23 +1 @@
|
||||
from .claude_service import analyze_case
|
||||
from .stats_service import get_statistical_recommendation
|
||||
from .scoring_service import WeightedScorer, create_scorer_for_case, get_weight_explanation
|
||||
from .geo_service import build_search_zones, haversine, Zone
|
||||
from .distance_service import calculate_max_distance, get_distance_priors, get_distance_statistics
|
||||
from .psychotype_service import detect_psychotype, get_psychotype_modifiers, get_search_recommendations
|
||||
|
||||
__all__ = [
|
||||
"analyze_case",
|
||||
"get_statistical_recommendation",
|
||||
"WeightedScorer",
|
||||
"create_scorer_for_case",
|
||||
"get_weight_explanation",
|
||||
"build_search_zones",
|
||||
"haversine",
|
||||
"Zone",
|
||||
"calculate_max_distance",
|
||||
"get_distance_priors",
|
||||
"get_distance_statistics",
|
||||
"detect_psychotype",
|
||||
"get_psychotype_modifiers",
|
||||
"get_search_recommendations"
|
||||
]
|
||||
# Package marker only.
|
||||
|
||||
Reference in New Issue
Block a user