initial
This commit is contained in:
0
backend/app/api/__init__.py
Normal file
0
backend/app/api/__init__.py
Normal file
144
backend/app/api/admin.py
Normal file
144
backend/app/api/admin.py
Normal file
@@ -0,0 +1,144 @@
|
||||
"""Admin panel routes: settings, LLM logs, users."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||
|
||||
from app.core.settings_service import EDITABLE_SETTING_KEYS, get_all_settings, update_settings
|
||||
from app.db import get_db_dep
|
||||
from app.deps import require_admin
|
||||
from app.models import LlmCallLog, Setting, User
|
||||
from app.schemas import LlmLogOut, SettingsOut, SettingsUpdate
|
||||
|
||||
router = APIRouter(prefix="/api/admin", tags=["admin"])
|
||||
|
||||
|
||||
def _mask_secrets(values: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Mask sensitive api_key fields in outbound responses."""
|
||||
for k in ("llm.api_key", "embedding.api_key"):
|
||||
v = values.get(k)
|
||||
if isinstance(v, str) and v:
|
||||
values[k] = v[:4] + "***" + v[-4:] if len(v) > 8 else "***"
|
||||
# Never expose admin setup token via this endpoint
|
||||
values.pop("admin.setup_token", None)
|
||||
return values
|
||||
|
||||
|
||||
@router.get("/settings", response_model=SettingsOut)
|
||||
async def get_settings_endpoint(
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
_: User = Depends(require_admin),
|
||||
):
|
||||
values = await get_all_settings(db)
|
||||
values = _mask_secrets(values)
|
||||
return SettingsOut(values=values, editable_keys=sorted(EDITABLE_SETTING_KEYS.keys()))
|
||||
|
||||
|
||||
@router.put("/settings", response_model=SettingsOut)
|
||||
async def update_settings_endpoint(
|
||||
payload: SettingsUpdate,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
_: User = Depends(require_admin),
|
||||
):
|
||||
# Strip masked api_key fields unless the user typed a new value
|
||||
cleaned: Dict[str, Any] = {}
|
||||
for k, v in (payload.values or {}).items():
|
||||
if k in ("llm.api_key", "embedding.api_key") and isinstance(v, str) and "***" in v:
|
||||
continue
|
||||
cleaned[k] = v
|
||||
new_values = await update_settings(db, cleaned)
|
||||
# If embedding settings changed, drop the cached RAG client so the next
|
||||
# get_rag() call rebuilds it (and reconfigures Qdrant collections if dim changed).
|
||||
if any(k.startswith("embedding.") for k in cleaned):
|
||||
from app.core.rag import reset_rag
|
||||
await reset_rag()
|
||||
new_values = _mask_secrets(new_values)
|
||||
return SettingsOut(values=new_values, editable_keys=sorted(EDITABLE_SETTING_KEYS.keys()))
|
||||
|
||||
|
||||
@router.post("/embeddings/test")
|
||||
async def test_embeddings_endpoint(
|
||||
payload: Dict[str, Any] = Body(default={}),
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
_: User = Depends(require_admin),
|
||||
):
|
||||
"""Probe the currently configured embeddings endpoint.
|
||||
|
||||
Accepts an optional `overrides` dict with embedding.* keys (e.g. to test
|
||||
a new endpoint before saving). Returns: ok, provider, base_url, model,
|
||||
dim, sample_norm (or error).
|
||||
"""
|
||||
from app.core.rag import probe_embeddings
|
||||
settings_map = await get_all_settings(db)
|
||||
# Apply ad-hoc overrides (without saving) so the admin can try before save
|
||||
overrides = (payload or {}).get("overrides") or {}
|
||||
for k, v in overrides.items():
|
||||
if k in EDITABLE_SETTING_KEYS:
|
||||
settings_map[k] = v
|
||||
return await probe_embeddings(settings_map)
|
||||
|
||||
|
||||
@router.get("/llm-logs", response_model=List[LlmLogOut])
|
||||
async def list_llm_logs(
|
||||
limit: int = 50,
|
||||
offset: int = 0,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
_: User = Depends(require_admin),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(LlmCallLog).order_by(LlmCallLog.created_at.desc()).limit(min(limit, 200)).offset(offset)
|
||||
)
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
@router.get("/llm-logs/{log_id}")
|
||||
async def get_llm_log(
|
||||
log_id: str,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
_: User = Depends(require_admin),
|
||||
):
|
||||
from uuid import UUID
|
||||
result = await db.execute(select(LlmCallLog).where(LlmCallLog.id == UUID(log_id)))
|
||||
log = result.scalars().first()
|
||||
if not log:
|
||||
raise HTTPException(status_code=404, detail="log_not_found")
|
||||
return {
|
||||
"id": str(log.id),
|
||||
"purpose": log.purpose,
|
||||
"model": log.model,
|
||||
"base_url": log.base_url,
|
||||
"prompt_messages": log.prompt_messages,
|
||||
"tools": log.tools,
|
||||
"response_text": log.response_text,
|
||||
"tool_calls": log.tool_calls,
|
||||
"prompt_tokens": log.prompt_tokens,
|
||||
"completion_tokens": log.completion_tokens,
|
||||
"total_tokens": log.total_tokens,
|
||||
"latency_ms": log.latency_ms,
|
||||
"error": log.error,
|
||||
"created_at": log.created_at.isoformat() if log.created_at else None,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/users")
|
||||
async def list_users(
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
_: User = Depends(require_admin),
|
||||
):
|
||||
result = await db.execute(select(User).order_by(User.created_at.desc()))
|
||||
users = result.scalars().all()
|
||||
return [
|
||||
{
|
||||
"id": str(u.id),
|
||||
"email": u.email,
|
||||
"username": u.username,
|
||||
"is_admin": u.is_admin,
|
||||
"is_active": u.is_active,
|
||||
"created_at": u.created_at.isoformat() if u.created_at else None,
|
||||
}
|
||||
for u in users
|
||||
]
|
||||
83
backend/app/api/auth.py
Normal file
83
backend/app/api/auth.py
Normal file
@@ -0,0 +1,83 @@
|
||||
"""Authentication routes: register, login, me, admin setup."""
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
|
||||
from app.core.security import create_access_token, hash_password, verify_password
|
||||
from app.core.settings_service import get_setting
|
||||
from app.db import get_db_dep
|
||||
from app.deps import get_current_user
|
||||
from app.models import User
|
||||
from app.schemas import AdminSetupRequest, TokenOut, UserLogin, UserOut, UserRegister
|
||||
|
||||
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
||||
|
||||
|
||||
@router.post("/register", response_model=TokenOut, status_code=status.HTTP_201_CREATED)
|
||||
async def register(payload: UserRegister, db: AsyncSession = Depends(get_db_dep)):
|
||||
existing = await db.execute(select(User).where((User.email == payload.email) | (User.username == payload.username)))
|
||||
if existing.scalars().first():
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="user_already_exists")
|
||||
user = User(
|
||||
email=payload.email,
|
||||
username=payload.username,
|
||||
hashed_password=hash_password(payload.password),
|
||||
is_admin=False,
|
||||
)
|
||||
db.add(user)
|
||||
await db.commit()
|
||||
await db.refresh(user)
|
||||
token = create_access_token(subject=str(user.id), extra={"is_admin": user.is_admin})
|
||||
return TokenOut(access_token=token, user=UserOut.model_validate(user))
|
||||
|
||||
|
||||
@router.post("/login", response_model=TokenOut)
|
||||
async def login(payload: UserLogin, db: AsyncSession = Depends(get_db_dep)):
|
||||
result = await db.execute(select(User).where(User.email == payload.email))
|
||||
user = result.scalars().first()
|
||||
if not user or not verify_password(payload.password, user.hashed_password):
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid_credentials")
|
||||
if not user.is_active:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="user_disabled")
|
||||
token = create_access_token(subject=str(user.id), extra={"is_admin": user.is_admin})
|
||||
return TokenOut(access_token=token, user=UserOut.model_validate(user))
|
||||
|
||||
|
||||
@router.get("/me", response_model=UserOut)
|
||||
async def me(user: User = Depends(get_current_user)):
|
||||
return user
|
||||
|
||||
|
||||
@router.post("/admin-setup", response_model=TokenOut)
|
||||
async def admin_setup(payload: AdminSetupRequest, db: AsyncSession = Depends(get_db_dep)):
|
||||
"""One-time endpoint to create the first admin user using a setup token."""
|
||||
# Check if any admin already exists
|
||||
existing_admins = await db.execute(select(User).where(User.is_admin.is_(True)))
|
||||
if existing_admins.scalars().first() is not None:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="admin_already_exists")
|
||||
|
||||
# Validate setup token (from DB or env)
|
||||
db_token = await get_setting(db, "admin.setup_token", default=None)
|
||||
env_token = payload.token # what the user supplied
|
||||
if not db_token or db_token != env_token:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="invalid_setup_token")
|
||||
|
||||
# Check user collision
|
||||
existing = await db.execute(select(User).where((User.email == payload.email) | (User.username == payload.username)))
|
||||
if existing.scalars().first():
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="user_already_exists")
|
||||
|
||||
user = User(
|
||||
email=payload.email,
|
||||
username=payload.username,
|
||||
hashed_password=hash_password(payload.password),
|
||||
is_admin=True,
|
||||
)
|
||||
db.add(user)
|
||||
await db.commit()
|
||||
await db.refresh(user)
|
||||
token = create_access_token(subject=str(user.id), extra={"is_admin": user.is_admin})
|
||||
return TokenOut(access_token=token, user=UserOut.model_validate(user))
|
||||
65
backend/app/api/misc.py
Normal file
65
backend/app/api/misc.py
Normal file
@@ -0,0 +1,65 @@
|
||||
"""Glossary + Triggers routes."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from app.db import get_db_dep
|
||||
from app.deps import get_current_user
|
||||
from app.models import DeferredTrigger, GlossaryEntry, Session, User, World
|
||||
from app.schemas import GlossaryEntryOut, TriggerOut
|
||||
|
||||
router = APIRouter(prefix="/api", tags=["misc"])
|
||||
|
||||
|
||||
@router.get("/worlds/{world_id}/glossary", response_model=List[GlossaryEntryOut])
|
||||
async def list_glossary(
|
||||
world_id: UUID,
|
||||
kind: str | None = None,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
w_result = await db.execute(select(World).where(World.id == world_id))
|
||||
world = w_result.scalars().first()
|
||||
if not world:
|
||||
raise HTTPException(status_code=404, detail="world_not_found")
|
||||
if world.owner_id != user.id and not user.is_admin:
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
|
||||
q = select(GlossaryEntry).where(GlossaryEntry.world_id == world_id)
|
||||
if kind:
|
||||
q = q.where(GlossaryEntry.kind == kind)
|
||||
q = q.order_by(GlossaryEntry.created_at.desc())
|
||||
result = await db.execute(q)
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
@router.get("/sessions/{session_id}/triggers", response_model=List[TriggerOut])
|
||||
async def list_triggers(
|
||||
session_id: UUID,
|
||||
include_fired: bool = True,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
s_result = await db.execute(
|
||||
select(Session).join(World, Session.world_id == World.id).where(Session.id == session_id)
|
||||
)
|
||||
session = s_result.scalars().first()
|
||||
if not session:
|
||||
raise HTTPException(status_code=404, detail="session_not_found")
|
||||
w_result = await db.execute(select(World).where(World.id == session.world_id))
|
||||
world = w_result.scalars().first()
|
||||
if not world or (world.owner_id != user.id and not user.is_admin):
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
|
||||
q = select(DeferredTrigger).where(DeferredTrigger.session_id == session_id)
|
||||
if not include_fired:
|
||||
q = q.where(DeferredTrigger.fired.is_(False))
|
||||
q = q.order_by(DeferredTrigger.fire_at)
|
||||
result = await db.execute(q)
|
||||
return result.scalars().all()
|
||||
71
backend/app/api/presets.py
Normal file
71
backend/app/api/presets.py
Normal file
@@ -0,0 +1,71 @@
|
||||
"""Preset routes: list / get / create."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from app.db import get_db_dep
|
||||
from app.deps import get_current_user
|
||||
from app.models import Preset, User
|
||||
from app.schemas import PresetCreate, PresetOut
|
||||
|
||||
router = APIRouter(prefix="/api/presets", tags=["presets"])
|
||||
|
||||
|
||||
@router.get("", response_model=List[PresetOut])
|
||||
async def list_presets(
|
||||
language: str | None = None,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
"""List public presets + user's private ones, optionally filtered by language."""
|
||||
q = select(Preset).where(
|
||||
(Preset.is_public.is_(True)) | (Preset.author_id == user.id)
|
||||
)
|
||||
if language:
|
||||
q = q.where(Preset.language == language)
|
||||
q = q.order_by(Preset.is_builtin.desc(), Preset.created_at.desc())
|
||||
result = await db.execute(q)
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
@router.get("/{preset_id}", response_model=PresetOut)
|
||||
async def get_preset(
|
||||
preset_id: UUID,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
_: User = Depends(get_current_user),
|
||||
):
|
||||
result = await db.execute(select(Preset).where(Preset.id == preset_id))
|
||||
preset = result.scalars().first()
|
||||
if not preset:
|
||||
raise HTTPException(status_code=404, detail="preset_not_found")
|
||||
if not preset.is_public and preset.author_id != _.id:
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
return preset
|
||||
|
||||
|
||||
@router.post("", response_model=PresetOut, status_code=201)
|
||||
async def create_preset(
|
||||
payload: PresetCreate,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
preset = Preset(
|
||||
slug=payload.slug,
|
||||
title=payload.title,
|
||||
description=payload.description,
|
||||
language=payload.language,
|
||||
is_public=payload.is_public,
|
||||
is_builtin=False,
|
||||
payload=payload.payload,
|
||||
author_id=user.id,
|
||||
)
|
||||
db.add(preset)
|
||||
await db.commit()
|
||||
await db.refresh(preset)
|
||||
return preset
|
||||
162
backend/app/api/sessions.py
Normal file
162
backend/app/api/sessions.py
Normal file
@@ -0,0 +1,162 @@
|
||||
"""Sessions routes: list / create / get / messages / start iteration (SSE)."""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
from typing import List
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sse_starlette.sse import EventSourceResponse
|
||||
|
||||
from app.db import get_db_dep
|
||||
from app.deps import get_current_user
|
||||
from app.engine.orchestrator import run_iteration
|
||||
from app.models import Message, Session, User, World
|
||||
from app.schemas import IterationRequest, MessageOut, SessionCreate, SessionOut
|
||||
|
||||
router = APIRouter(prefix="/api/sessions", tags=["sessions"])
|
||||
|
||||
|
||||
@router.get("", response_model=List[SessionOut])
|
||||
async def list_sessions(
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(Session)
|
||||
.join(World, Session.world_id == World.id)
|
||||
.where(World.owner_id == user.id)
|
||||
.order_by(Session.last_played_at.desc().nullslast())
|
||||
)
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
@router.post("", response_model=SessionOut, status_code=201)
|
||||
async def create_session(
|
||||
payload: SessionCreate,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
# Verify world ownership
|
||||
result = await db.execute(select(World).where(World.id == payload.world_id))
|
||||
world = result.scalars().first()
|
||||
if not world:
|
||||
raise HTTPException(status_code=404, detail="world_not_found")
|
||||
if world.owner_id != user.id and not user.is_admin:
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
if world.status not in ("ready", "active"):
|
||||
raise HTTPException(status_code=400, detail=f"world_not_ready: status={world.status}")
|
||||
|
||||
session = Session(
|
||||
world_id=world.id,
|
||||
title=payload.title or f"Сессия в мире «{world.name}»",
|
||||
)
|
||||
db.add(session)
|
||||
# Mark world as active
|
||||
world.status = "active"
|
||||
await db.commit()
|
||||
await db.refresh(session)
|
||||
return session
|
||||
|
||||
|
||||
@router.get("/{session_id}", response_model=SessionOut)
|
||||
async def get_session(
|
||||
session_id: UUID,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(Session).join(World, Session.world_id == World.id).where(Session.id == session_id)
|
||||
)
|
||||
session = result.scalars().first()
|
||||
if not session:
|
||||
raise HTTPException(status_code=404, detail="session_not_found")
|
||||
# Verify ownership via world
|
||||
w_result = await db.execute(select(World).where(World.id == session.world_id))
|
||||
world = w_result.scalars().first()
|
||||
if not world or (world.owner_id != user.id and not user.is_admin):
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
return session
|
||||
|
||||
|
||||
@router.get("/{session_id}/messages", response_model=List[MessageOut])
|
||||
async def list_messages(
|
||||
session_id: UUID,
|
||||
include_hidden: bool = Query(False),
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
# Verify access
|
||||
result = await db.execute(
|
||||
select(Session).join(World, Session.world_id == World.id).where(Session.id == session_id)
|
||||
)
|
||||
session = result.scalars().first()
|
||||
if not session:
|
||||
raise HTTPException(status_code=404, detail="session_not_found")
|
||||
w_result = await db.execute(select(World).where(World.id == session.world_id))
|
||||
world = w_result.scalars().first()
|
||||
if not world or (world.owner_id != user.id and not user.is_admin):
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
|
||||
q = select(Message).where(Message.session_id == session_id).order_by(Message.seq)
|
||||
if not include_hidden:
|
||||
q = q.where(Message.hidden.is_(False))
|
||||
result = await db.execute(q)
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
@router.post("/{session_id}/iterate")
|
||||
async def iterate_session(
|
||||
session_id: UUID,
|
||||
payload: IterationRequest,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
"""SSE stream of the iteration."""
|
||||
# Verify access
|
||||
result = await db.execute(
|
||||
select(Session).join(World, Session.world_id == World.id).where(Session.id == session_id)
|
||||
)
|
||||
session = result.scalars().first()
|
||||
if not session:
|
||||
raise HTTPException(status_code=404, detail="session_not_found")
|
||||
w_result = await db.execute(select(World).where(World.id == session.world_id))
|
||||
world = w_result.scalars().first()
|
||||
if not world or (world.owner_id != user.id and not user.is_admin):
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
if payload.session_id != session_id:
|
||||
raise HTTPException(status_code=400, detail="session_id_mismatch")
|
||||
|
||||
async def event_generator():
|
||||
try:
|
||||
async for event in run_iteration(db=db, user_id=user.id, session_id=session_id, action_text=payload.action_text):
|
||||
yield {"event": event["type"], "data": json.dumps(event.get("data", {}), ensure_ascii=False, default=str)}
|
||||
except Exception as e:
|
||||
yield {"event": "error", "data": json.dumps({"message": str(e)}, ensure_ascii=False)}
|
||||
yield {"event": "done", "data": "{}"}
|
||||
|
||||
return EventSourceResponse(event_generator())
|
||||
|
||||
|
||||
@router.delete("/{session_id}", status_code=204)
|
||||
async def delete_session(
|
||||
session_id: UUID,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(Session).join(World, Session.world_id == World.id).where(Session.id == session_id)
|
||||
)
|
||||
session = result.scalars().first()
|
||||
if not session:
|
||||
raise HTTPException(status_code=404, detail="session_not_found")
|
||||
w_result = await db.execute(select(World).where(World.id == session.world_id))
|
||||
world = w_result.scalars().first()
|
||||
if not world or (world.owner_id != user.id and not user.is_admin):
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
await db.delete(session)
|
||||
await db.commit()
|
||||
161
backend/app/api/worlds.py
Normal file
161
backend/app/api/worlds.py
Normal file
@@ -0,0 +1,161 @@
|
||||
"""Worlds routes: CRUD + world builder flow."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from app.db import get_db_dep
|
||||
from app.deps import get_current_user
|
||||
from app.engine.world_builder import commit_world_builder, continue_world_builder, start_world_builder
|
||||
from app.models import User, World
|
||||
from app.schemas import (
|
||||
WorldBuilderCommit,
|
||||
WorldBuilderMessage,
|
||||
WorldBuilderReply,
|
||||
WorldBuilderStart,
|
||||
WorldCreate,
|
||||
WorldOut,
|
||||
WorldUpdate,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/api/worlds", tags=["worlds"])
|
||||
|
||||
|
||||
@router.get("", response_model=List[WorldOut])
|
||||
async def list_worlds(
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(World).where(World.owner_id == user.id).order_by(World.updated_at.desc())
|
||||
)
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
@router.get("/{world_id}", response_model=WorldOut)
|
||||
async def get_world(
|
||||
world_id: UUID,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
result = await db.execute(select(World).where(World.id == world_id))
|
||||
world = result.scalars().first()
|
||||
if not world:
|
||||
raise HTTPException(status_code=404, detail="world_not_found")
|
||||
if world.owner_id != user.id and not user.is_admin:
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
return world
|
||||
|
||||
|
||||
@router.post("", response_model=WorldOut, status_code=201)
|
||||
async def create_world(
|
||||
payload: WorldCreate,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
world = World(
|
||||
owner_id=user.id,
|
||||
name=payload.name,
|
||||
language=payload.language,
|
||||
definition={},
|
||||
state={},
|
||||
status="draft",
|
||||
preset_id=payload.preset_id,
|
||||
)
|
||||
db.add(world)
|
||||
await db.commit()
|
||||
await db.refresh(world)
|
||||
return world
|
||||
|
||||
|
||||
@router.patch("/{world_id}", response_model=WorldOut)
|
||||
async def update_world(
|
||||
world_id: UUID,
|
||||
payload: WorldUpdate,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
result = await db.execute(select(World).where(World.id == world_id))
|
||||
world = result.scalars().first()
|
||||
if not world:
|
||||
raise HTTPException(status_code=404, detail="world_not_found")
|
||||
if world.owner_id != user.id and not user.is_admin:
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
for field, value in payload.model_dump(exclude_unset=True).items():
|
||||
setattr(world, field, value)
|
||||
await db.commit()
|
||||
await db.refresh(world)
|
||||
return world
|
||||
|
||||
|
||||
@router.delete("/{world_id}", status_code=204)
|
||||
async def delete_world(
|
||||
world_id: UUID,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
result = await db.execute(select(World).where(World.id == world_id))
|
||||
world = result.scalars().first()
|
||||
if not world:
|
||||
raise HTTPException(status_code=404, detail="world_not_found")
|
||||
if world.owner_id != user.id and not user.is_admin:
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
await db.delete(world)
|
||||
await db.commit()
|
||||
|
||||
|
||||
# === World Builder flow ===
|
||||
|
||||
@router.post("/builder/start", response_model=WorldBuilderReply)
|
||||
async def builder_start(
|
||||
payload: WorldBuilderStart,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
try:
|
||||
return await start_world_builder(
|
||||
db=db,
|
||||
user=user,
|
||||
world_name=payload.world_name,
|
||||
language=payload.language,
|
||||
preset_id=payload.preset_id,
|
||||
setting_brief=payload.setting_brief,
|
||||
character_brief=payload.character_brief,
|
||||
rules_brief=payload.rules_brief,
|
||||
notes=payload.notes,
|
||||
)
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=f"builder_start_failed: {e}")
|
||||
|
||||
|
||||
@router.post("/builder/continue", response_model=WorldBuilderReply)
|
||||
async def builder_continue(
|
||||
payload: WorldBuilderMessage,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
try:
|
||||
return await continue_world_builder(db=db, user=user, session_id=payload.session_id, user_message=payload.message)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=f"builder_continue_failed: {e}")
|
||||
|
||||
|
||||
@router.post("/builder/commit", response_model=WorldOut)
|
||||
async def builder_commit(
|
||||
payload: WorldBuilderCommit,
|
||||
db: AsyncSession = Depends(get_db_dep),
|
||||
user: User = Depends(get_current_user),
|
||||
):
|
||||
try:
|
||||
return await commit_world_builder(db=db, user=user, session_id=payload.session_id, name=payload.name)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=f"builder_commit_failed: {e}")
|
||||
Reference in New Issue
Block a user