223 lines
7.9 KiB
Python
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()
|