cosmo/backend/app/api/rocket.py

223 lines
7.9 KiB
Python

"""Rocket simulator configuration APIs."""
from datetime import datetime
from typing import Annotated
from fastapi import APIRouter, Depends, HTTPException, Query, status
from pydantic import BaseModel, ConfigDict, Field, StringConstraints, field_validator
from sqlalchemy import func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from app.models.db import RocketConfig, User
from app.services.auth_deps import require_admin
router = APIRouter(prefix="/rockets", tags=["rockets"])
RocketCode = Annotated[str, StringConstraints(pattern=r"^[a-z0-9][a-z0-9-]*$")]
class RocketStage(BaseModel):
name: str = Field(min_length=1, max_length=100)
dry_mass_kg: float = Field(gt=0)
fuel_mass_kg: float = Field(gt=0)
max_thrust_n: float = Field(gt=0)
specific_impulse_s: float = Field(gt=0)
engine_count: int = Field(gt=0)
length_ratio: float = Field(gt=0, le=1)
class RocketBase(BaseModel):
code: RocketCode
name: str = Field(min_length=1, max_length=100)
name_zh: str | None = Field(default=None, max_length=100)
manufacturer: str | None = Field(default=None, max_length=100)
country: str | None = Field(default=None, max_length=100)
launch_site_name: str = Field(default="Equatorial Launch Site", min_length=1, max_length=120)
launch_latitude_deg: float = Field(default=0, ge=-90, le=90)
launch_longitude_deg: float = Field(default=0, ge=-180, le=180)
description: str | None = None
color: str = Field(default="#f8fafc", pattern=r"^#[0-9a-fA-F]{6}$")
height_m: float = Field(gt=0)
diameter_m: float = Field(gt=0)
payload_mass_kg: float = Field(ge=0, default=0)
drag_coefficient: float = Field(gt=0, le=2, default=0.4)
reference_area_m2: float = Field(gt=0)
target_orbit_km: float = Field(gt=0, default=200)
target_velocity_mps: float = Field(gt=0, default=7800)
separation_delay_seconds: float = Field(ge=0, le=30, default=2)
second_stage_ignition_delay_seconds: float = Field(ge=0, le=30, default=1)
stage_1: RocketStage
stage_2: RocketStage
is_active: bool = True
sort_order: int = 0
@field_validator("code", mode="before")
@classmethod
def normalize_code(cls, value: str) -> str:
return value.strip().lower()
class RocketCreate(RocketBase):
pass
class RocketUpdate(BaseModel):
code: RocketCode | None = None
name: str | None = Field(default=None, min_length=1, max_length=100)
name_zh: str | None = Field(default=None, max_length=100)
manufacturer: str | None = Field(default=None, max_length=100)
country: str | None = Field(default=None, max_length=100)
launch_site_name: str | None = Field(default=None, min_length=1, max_length=120)
launch_latitude_deg: float | None = Field(default=None, ge=-90, le=90)
launch_longitude_deg: float | None = Field(default=None, ge=-180, le=180)
description: str | None = None
color: str | None = Field(default=None, pattern=r"^#[0-9a-fA-F]{6}$")
height_m: float | None = Field(default=None, gt=0)
diameter_m: float | None = Field(default=None, gt=0)
payload_mass_kg: float | None = Field(default=None, ge=0)
drag_coefficient: float | None = Field(default=None, gt=0, le=2)
reference_area_m2: float | None = Field(default=None, gt=0)
target_orbit_km: float | None = Field(default=None, gt=0)
target_velocity_mps: float | None = Field(default=None, gt=0)
separation_delay_seconds: float | None = Field(default=None, ge=0, le=30)
second_stage_ignition_delay_seconds: float | None = Field(default=None, ge=0, le=30)
stage_1: RocketStage | None = None
stage_2: RocketStage | None = None
is_active: bool | None = None
sort_order: int | None = None
@field_validator("code", mode="before")
@classmethod
def normalize_code(cls, value: str | None) -> str | None:
return value.strip().lower() if value else value
class RocketResponse(RocketBase):
id: int
created_at: datetime | None = None
updated_at: datetime | None = None
model_config = ConfigDict(from_attributes=True)
def rocket_values(data: RocketCreate | RocketUpdate, *, exclude_unset: bool = False) -> dict:
values = data.model_dump(exclude_unset=exclude_unset)
for key in ("stage_1", "stage_2"):
if key in values and values[key] is not None:
values[key] = dict(values[key])
return values
async def ensure_unique_code(db: AsyncSession, code: str, exclude_id: int | None = None) -> None:
query = select(RocketConfig.id).where(RocketConfig.code == code)
if exclude_id is not None:
query = query.where(RocketConfig.id != exclude_id)
if (await db.execute(query)).scalar_one_or_none() is not None:
raise HTTPException(status_code=400, detail="Rocket code already exists")
@router.get("", response_model=list[RocketResponse])
async def list_active_rockets(db: AsyncSession = Depends(get_db)):
result = await db.execute(
select(RocketConfig)
.where(RocketConfig.is_active.is_(True))
.order_by(RocketConfig.sort_order, RocketConfig.id)
)
return result.scalars().all()
@router.get("/admin")
async def list_rockets_for_admin(
skip: int = Query(0, ge=0),
limit: int = Query(20, ge=1, le=100),
search: str | None = None,
db: AsyncSession = Depends(get_db),
_: User = Depends(require_admin),
):
filters = []
if search:
keyword = f"%{search.strip()}%"
filters.append(or_(
RocketConfig.code.ilike(keyword),
RocketConfig.name.ilike(keyword),
RocketConfig.name_zh.ilike(keyword),
))
total = await db.scalar(select(func.count(RocketConfig.id)).where(*filters))
result = await db.execute(
select(RocketConfig)
.where(*filters)
.order_by(RocketConfig.sort_order, RocketConfig.id)
.offset(skip)
.limit(limit)
)
return {
"rockets": [RocketResponse.model_validate(item) for item in result.scalars().all()],
"total": total or 0,
}
@router.get("/{code}", response_model=RocketResponse)
async def get_active_rocket(code: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(
select(RocketConfig).where(
RocketConfig.code == code.lower(),
RocketConfig.is_active.is_(True),
)
)
rocket = result.scalar_one_or_none()
if not rocket:
raise HTTPException(status_code=404, detail="Rocket not found")
return rocket
@router.post("/admin", response_model=RocketResponse, status_code=status.HTTP_201_CREATED)
async def create_rocket(
data: RocketCreate,
db: AsyncSession = Depends(get_db),
_: User = Depends(require_admin),
):
await ensure_unique_code(db, data.code)
rocket = RocketConfig(**rocket_values(data))
db.add(rocket)
await db.commit()
await db.refresh(rocket)
return rocket
@router.put("/admin/{rocket_id}", response_model=RocketResponse)
async def update_rocket(
rocket_id: int,
data: RocketUpdate,
db: AsyncSession = Depends(get_db),
_: User = Depends(require_admin),
):
rocket = await db.get(RocketConfig, rocket_id)
if not rocket:
raise HTTPException(status_code=404, detail="Rocket not found")
values = rocket_values(data, exclude_unset=True)
if not values:
raise HTTPException(status_code=400, detail="No fields to update")
if "code" in values:
await ensure_unique_code(db, values["code"], rocket_id)
for key, value in values.items():
setattr(rocket, key, value)
await db.commit()
await db.refresh(rocket)
return rocket
@router.delete("/admin/{rocket_id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_rocket(
rocket_id: int,
db: AsyncSession = Depends(get_db),
_: User = Depends(require_admin),
):
rocket = await db.get(RocketConfig, rocket_id)
if not rocket:
raise HTTPException(status_code=404, detail="Rocket not found")
await db.delete(rocket)
await db.commit()