Remove three kinds of duplication in the backend

Stage 2, items 1, 4, and 5 of `plan/17-refactor.md`.

**One path resolver in `worldstate`.** `apply_delta` and `apply_override` routed
`flags.<name>`, `milestones.<id>`, `world.<stat>`, `player.<stat>`, and
`npc.<id>.<stat>` with parallel code, about 100 lines each. `_resolve` now says
what a path points at and returns either a target or the rejection to report.
Each function keeps its own write rule, because the rules genuinely differ: an
override sets a number rather than adding to it, ignores `cooldown`,
`max_delta_per_turn`, and the rule that a counter only counts up, and can un-set
a milestone.

A differential check ran both implementations over 3960 payloads: twenty paths,
fourteen values, three starting states, plus every three-path combination. The
results are identical except that 674 rejections from `apply_override` now carry
a `fix` string. `apply_delta` already worded those, and the world-state editor
renders them, so an override that names an unknown flag now explains itself the
way a delta does.

**`sse`, `SSE_HEADERS`, and `turn_error` move to `app/sse.py`.** Two routers
stream, and `chat.py` had to import from `routers.adventures` to reach them.

**`get_adventure_or_404` becomes the `current_adventure` dependency.** All 32
handlers repeated the call as their first statement. The ownership check now
reads in the signature and runs before the body. FastAPI caches a dependency for
one request, so the handler's `db` is the session the adventure came from.

The generated OpenAPI document is byte-identical except on `rename_branch`,
where `branch_id` is now listed before `adventure_id`, because that handler no
longer names `adventure_id` itself. Parameter order in the document is
cosmetic.

Six tests in `test_state_revert.py` call `undo_turn` and `retry_action`
directly rather than over HTTP. They pass the adventure they already hold
instead of an id.

549 tests pass.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014Dix4oGV3njgWRdu7P9t6r
This commit is contained in:
parththakkar106
2026-08-29 02:04:02 +05:30
co-authored by Claude Opus 5
parent 0623dd8b78
commit 2c57b1ceab
16 changed files with 258 additions and 261 deletions
+1 -3
View File
@@ -42,17 +42,15 @@ from ... import limits # noqa: F401 `adventures.limits` is patched by tests.
from .crud import SNIPPET_MAX, _snippet
from .paging import ACTION_PAGE
from .takes import retry_action, undo_turn
from .turns import SSE_HEADERS, sse, world_delta_of
from .turns import world_delta_of
__all__ = [
"ACTION_PAGE",
"SNIPPET_MAX",
"SSE_HEADERS",
"_snippet",
"limits",
"retry_action",
"router",
"sse",
"undo_turn",
"world_delta_of",
]
+4 -8
View File
@@ -11,18 +11,17 @@ from sqlalchemy.orm import Session
from ... import models, schemas, tree
from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
from .nodes import delete_turn
from .paging import ACTION_PAGE, action_window, annotate_takes
@router.get("/{adventure_id}/actions", response_model=schemas.ActionPage)
def list_actions(
adventure_id: int,
before_id: int | None = None,
limit: int = ACTION_PAGE,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Returns a page of the story, working backwards from the newest action.
@@ -30,7 +29,6 @@ def list_actions(
asks for what comes before it. Omit `before_id` for the newest window. See
`action_window` for why this anchors on a row rather than an offset.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
limit = max(1, min(limit, ACTION_PAGE * 4))
actions, total, has_more = action_window(
db, adventure, before_id=before_id, limit=limit
@@ -51,9 +49,8 @@ def update_action(
action_id: int,
payload: schemas.ActionUpdate,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found")
@@ -70,9 +67,8 @@ def delete_action(
adventure_id: int,
action_id: int,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
adventure = get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found")
+6 -10
View File
@@ -18,14 +18,15 @@ from ...context import lineage
from ...database import get_db
from . import turns
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
from .nodes import db_tip
from .paging import current_window
@router.get("/{adventure_id}/branches", response_model=list[schemas.BranchOut])
def list_branches(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
):
"""Returns every branch of the adventure and where each one leaves its parent.
@@ -35,7 +36,6 @@ def list_branches(
`actions`, never one query per branch, so a view of a hundred forks does not
cost a hundred round trips.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
branches = (
db.query(models.Branch)
.filter(models.Branch.adventure_id == adventure.id)
@@ -95,11 +95,10 @@ def get_branch_or_404(
"/{adventure_id}/branches/{branch_id}", response_model=schemas.BranchOut
)
def rename_branch(
adventure_id: int,
branch_id: int,
payload: schemas.BranchRename,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Names a branch, or clears the name to leave it unnamed.
@@ -107,7 +106,6 @@ def rename_branch(
name anyone chose, and storing one gives the client an empty label to draw
instead of the fork depth.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
branch = get_branch_or_404(adventure, branch_id, db)
name = (payload.name or "").strip()
branch.name = name or None
@@ -145,7 +143,7 @@ def delete_branch(
adventure_id: int,
branch_id: int,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Deletes a branch and everything forked from it.
@@ -163,7 +161,6 @@ def delete_branch(
cascade on `branches.parent_branch_id`, so the delete is a single statement
however deep the subtree is.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
branch = get_branch_or_404(adventure, branch_id, db)
if branch.parent_branch_id is None:
raise HTTPException(
@@ -238,7 +235,7 @@ def switch_branch(
adventure_id: int,
branch_id: int,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Reads and plays a different branch of the story.
@@ -249,7 +246,6 @@ def switch_branch(
another branch's numbers, including the world-state cooldown clock inside the
snapshot.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
branch = db.get(models.Branch, branch_id)
if branch is None or branch.adventure_id != adventure.id:
raise HTTPException(404, "Branch not found")
+3 -3
View File
@@ -10,19 +10,19 @@ from sqlalchemy.orm import Session
from ... import analytics, bundle, limits, models, schemas
from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
@router.get("/{adventure_id}/export")
def export_adventure(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
db: Session = Depends(get_db),
adv: models.Adventure = Depends(current_adventure),
):
"""Returns a full backup: plot components, story cards, scripts, state, and tree.
`app/bundle.py` owns the format, in both of its versions. A backup outlives
the schema, so no call site decides anything about its shape.
"""
adv = get_adventure_or_404(adventure_id, db, user)
return bundle.export(db, adv)
+10 -15
View File
@@ -12,7 +12,7 @@ from sqlalchemy.orm.attributes import set_committed_value
from ... import analytics, attempts, images, limits, memorybank, models, schemas, tree, worldstate
from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
from .paging import action_window, annotate_takes
from .scenario_text import fill_placeholders, scenario_card_specs
@@ -229,7 +229,8 @@ def create_adventure(
@router.get("/{adventure_id}", response_model=schemas.AdventureOut)
def get_adventure(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
):
"""Returns the adventure and the newest window of its story.
@@ -238,7 +239,6 @@ def get_adventure(
exist above. `GET /{id}/actions` serves the older pages as the reader
scrolls up.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
actions, total, _ = action_window(db, adventure)
# Annotate before handing over the window. This path serializes through the
# relationship rather than building `ActionOut` itself, so the pager numbers
@@ -259,7 +259,7 @@ def get_adventure(
@router.get("/{adventure_id}/script-state")
def get_script_state(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
adventure: models.Adventure = Depends(current_adventure),
):
"""Returns the scripting `state` object.
@@ -267,21 +267,19 @@ def get_script_state(
`state.x`, persisted after each hook. It stays `{}` until a script sets a
variable.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
state = adventure.script_state if isinstance(adventure.script_state, dict) else {}
return {"state": state}
@router.get("/{adventure_id}/world-state")
def get_world_state(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
adventure: models.Adventure = Depends(current_adventure),
):
"""Returns the live RPG world state and the scenario's `stat_schema`.
The play view uses both to render the character sheet and the milestones.
`schema` is null when the adventure has no RPG layer.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
schema = adventure.scenario.stat_schema if adventure.scenario else None
state = adventure.world_state if isinstance(adventure.world_state, dict) else {}
return {
@@ -292,10 +290,9 @@ def get_world_state(
@router.put("/{adventure_id}/world-state")
def override_world_state(
adventure_id: int,
overrides: dict = Body(...),
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Edits the live RPG values directly, as a manual correction rather than a turn.
@@ -303,7 +300,6 @@ def override_world_state(
`milestones.y` to their new absolute values. The endpoint rejects unknown
paths and wrong types one at a time, and applies the rest.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
schema = adventure.scenario.stat_schema if adventure.scenario else None
if not worldstate.has_schema(schema):
raise HTTPException(400, "This adventure has no RPG world-state layer")
@@ -316,12 +312,10 @@ def override_world_state(
@router.patch("/{adventure_id}", response_model=schemas.AdventureOut)
def update_adventure(
adventure_id: int,
payload: schemas.AdventureUpdate,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
adventure = get_adventure_or_404(adventure_id, db, user)
for field, value in payload.model_dump(exclude_unset=True).items():
setattr(adventure, field, value)
db.commit()
@@ -330,9 +324,10 @@ def update_adventure(
@router.delete("/{adventure_id}", status_code=204)
def delete_adventure(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
adventure_id: int,
db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
):
adventure = get_adventure_or_404(adventure_id, db, user)
db.delete(adventure)
db.commit()
# No later request reads this adventure's vectors, so drop them now. The
+17
View File
@@ -9,6 +9,7 @@ from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from ... import auth, models
from ...database import get_db
router = APIRouter(prefix="/api/adventures", tags=["adventures"])
@@ -23,3 +24,19 @@ def get_adventure_or_404(
if adventure is None or adventure.user_id != user.id:
raise HTTPException(404, "Adventure not found")
return adventure
def current_adventure(
adventure_id: int,
db: Session = Depends(get_db),
user: models.User = Depends(auth.get_current_user),
) -> models.Adventure:
"""Resolves the `{adventure_id}` in the path to the caller's own adventure.
Declare this in a handler's signature. The ownership check then reads in the
signature, where you look for it, and it runs before the handler body.
FastAPI caches a dependency for the length of one request, so the handler's
own `db` is the same session this one loaded the adventure from.
"""
return get_adventure_or_404(adventure_id, db, user)
+5 -5
View File
@@ -12,15 +12,16 @@ from ...context import build_context
from ...database import get_db
from ..settings import get_settings
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
@router.get("/{adventure_id}/context")
async def dry_run_context(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Returns what the app would send to the AI if the player continued now."""
adventure = get_adventure_or_404(adventure_id, db, user)
settings = get_settings(db, user)
if auth.resolve_provider_config(settings).using_demo:
memories = (
@@ -39,9 +40,8 @@ def action_context(
adventure_id: int,
action_id: int,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found")
+7 -9
View File
@@ -11,7 +11,7 @@ from ... import limits, memorybank, models, schemas, tree
from ...context import lineage
from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
# The columns `schemas.MemoryOut` renders. `embedded` is a real column and
@@ -32,9 +32,10 @@ MEMORY_LIST_COLUMNS = (
@router.get("/{adventure_id}/memories", response_model=list[schemas.MemoryOut])
def list_memories(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
adventure_id: int,
db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
):
adventure = get_adventure_or_404(adventure_id, db, user)
# Name the columns in a query rather than walk `adventure.memories`.
# Retrieval used to walk the relationship, which is why a turn cost
# megabytes: a relationship load returns whole entities, so it reads
@@ -63,13 +64,12 @@ def list_memories(
@router.post("/{adventure_id}/memories", response_model=schemas.MemoryOut, status_code=201)
def create_memory(
adventure_id: int,
payload: schemas.MemoryCreate,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Adds a memory manually. The next post-turn pass embeds it."""
adventure = get_adventure_or_404(adventure_id, db, user)
limits.check_row_cap("memories", db, user, adventure=adventure)
if not payload.text.strip():
raise HTTPException(400, "Memory text cannot be empty")
@@ -88,9 +88,8 @@ def update_memory(
memory_id: int,
payload: schemas.MemoryUpdate,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
get_adventure_or_404(adventure_id, db, user)
memory = db.get(models.Memory, memory_id)
if memory is None or memory.adventure_id != adventure_id:
raise HTTPException(404, "Memory not found")
@@ -108,9 +107,8 @@ def delete_memory(
adventure_id: int,
memory_id: int,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
get_adventure_or_404(adventure_id, db, user)
memory = db.get(models.Memory, memory_id)
if memory is None or memory.adventure_id != adventure_id:
raise HTTPException(404, "Memory not found")
+5 -4
View File
@@ -14,7 +14,7 @@ from ... import models, schemas, worldstate
from ...database import get_db
from . import turns
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
from .scenario_text import (
CARD_FIELDS, SCENARIO_TEXT_FIELDS, fill_placeholders, scenario_card_specs,
scenario_placeholder_names,
@@ -113,10 +113,11 @@ def plan_refresh(
@router.get("/{adventure_id}/refresh", response_model=schemas.RefreshPlan)
def preview_refresh(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Returns what "Update from scenario" would change, for the confirm dialog."""
adventure = get_adventure_or_404(adventure_id, db, user)
scenario = resolve_source_scenario(adventure, db, user)
if scenario is None:
raise HTTPException(404, "No scenario to update from")
@@ -137,6 +138,7 @@ def refresh_from_scenario(
payload: schemas.AdventureRefresh = Body(default=schemas.AdventureRefresh()),
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Copies the scenario's current plot text, story cards, and stat schema over
this adventure's copy.
@@ -150,7 +152,6 @@ def refresh_from_scenario(
cards; and, through `worldstate.reconcile`, the live value of every stat the
schema still defines.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
scenario = resolve_source_scenario(adventure, db, user)
if scenario is None:
raise HTTPException(404, "No scenario to update from")
+6 -6
View File
@@ -11,7 +11,7 @@ from sqlalchemy.orm import Session
from ... import models, schemas
from ...database import get_db
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
# Fields that are copied from a library Script into its adventure-script
@@ -59,9 +59,10 @@ def _mark_out_of_date(
@router.get("/{adventure_id}/scripts", response_model=list[schemas.AdventureScriptOut])
def list_adventure_scripts(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
adventure = get_adventure_or_404(adventure_id, db, user)
return [_mark_out_of_date(s, db, user) for s in adventure.scripts]
@@ -74,12 +75,12 @@ def sync_adventure_script(
adv_script_id: int,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Overwrites this copy's code with the latest from its library script.
`enabled`, `position`, and the adventure's shared `script_state` are kept.
"""
get_adventure_or_404(adventure_id, db, user)
script = db.get(models.AdventureScript, adv_script_id)
if script is None or script.adventure_id != adventure_id:
raise HTTPException(404, "Script not found")
@@ -104,9 +105,8 @@ def update_adventure_script(
adv_script_id: int,
payload: schemas.AdventureScriptUpdate,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
get_adventure_or_404(adventure_id, db, user)
script = db.get(models.AdventureScript, adv_script_id)
if script is None or script.adventure_id != adventure_id:
raise HTTPException(404, "Script not found")
+10 -12
View File
@@ -16,12 +16,12 @@ from ...context import cursors
from ...context import lineage
from ...database import get_db
from ...scripting import ScriptPipeline
from ...sse import SSE_HEADERS
from . import turns
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
from .nodes import delete_turn, last_action, stand_on
from .paging import action_window, annotate_takes, current_window
from .turns import SSE_HEADERS
@router.post("/{adventure_id}/retry")
@@ -30,6 +30,7 @@ def retry_action(
request: Request,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Regenerates the last AI action and keeps the discarded attempt.
@@ -38,7 +39,6 @@ def retry_action(
attempt is stored as a sibling at the same coordinate. No text the AI wrote
is rewritten or deleted.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
limits.rate_limit("turn", request, user)
turns.check_demo_cap(db, user)
turns.acquire_turn_lock(adventure_id)
@@ -79,7 +79,7 @@ def list_variants(
adventure_id: int,
action_id: int,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Returns every attempt made for one AI turn.
@@ -90,7 +90,6 @@ def list_variants(
Switching changes which row the story tells, and a client that holds an id it
received a moment ago still has to be able to ask about the same turn.
"""
get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found")
@@ -119,7 +118,7 @@ def select_variant(
action_id: int,
payload: schemas.VariantSelect,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Makes an earlier attempt live again and restores the state it produced.
@@ -130,7 +129,6 @@ def select_variant(
would leave the story contradicting itself. The attempts of earlier turns
stay readable through `list_variants`.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found")
@@ -172,7 +170,7 @@ def fork_from_attempt(
adventure_id: int,
action_id: int,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Continues the story from this attempt, forking a branch if one is needed.
@@ -184,7 +182,6 @@ def fork_from_attempt(
* The story has moved past its turn, so the endpoint forks. The attempt gets
a branch of its own, and the line it leaves keeps every turn it has.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
action = db.get(models.Action, action_id)
if action is None or action.adventure_id != adventure_id:
raise HTTPException(404, "Action not found")
@@ -236,6 +233,7 @@ def add_take(
request: Request,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
"""Plays a turn again, whoever wrote it.
@@ -256,7 +254,6 @@ def add_take(
turn, so the new attempt is written at the same depth under the same parent,
and the line it leaves is unchanged. No node below is copied.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
limits.rate_limit("turn", request, user)
limits.check_row_cap("actions", db, user, adventure=adventure)
turns.check_demo_cap(db, user)
@@ -317,7 +314,9 @@ def add_take(
@router.post("/{adventure_id}/undo", response_model=schemas.ActionPage)
def undo_turn(
adventure_id: int, db: Session = Depends(get_db), user: models.User = CurrentUser
adventure_id: int,
db: Session = Depends(get_db),
adventure: models.Adventure = Depends(current_adventure),
):
"""Deletes the last turn: the trailing AI action and its player action, if any.
@@ -325,7 +324,6 @@ def undo_turn(
ran, and it prunes any memory that summarized the removed actions. The turn
lock prevents an undo while a turn is still generating.
"""
adventure = get_adventure_or_404(adventure_id, db, user)
turns.acquire_turn_lock(adventure_id)
try:
# Only the last turn is removed, so fetch the two actions it can
+3 -24
View File
@@ -6,7 +6,6 @@ lock guards one set only while one module owns it. And a test that replaces
`OpenAICompatibleProvider`, `generate_turn`, or `check_demo_cap` patches this
module, which every caller reads through.
"""
import json
import threading
from fastapi import Depends, HTTPException, Request
@@ -20,9 +19,10 @@ from ...context import build_context, cursors
from ...database import get_db
from ...providers import OpenAICompatibleProvider, PromptParts, ProviderError
from ...scripting import ScriptPipeline
from ...sse import SSE_HEADERS, sse, turn_error
from ..settings import get_settings
from .deps import CurrentUser, get_adventure_or_404, router
from .deps import CurrentUser, current_adventure, router
from .nodes import _move_to_after, next_depth, next_index
from .paging import annotate_takes
@@ -98,27 +98,6 @@ def format_player_input(action_type: str, text: str) -> str:
return text # The "story" type is appended as raw text.
def sse(obj: dict) -> str:
return f"data: {json.dumps(obj)}\n\n"
def turn_error(detail: str, **extra) -> str:
"""Returns an SSE error for a turn that could not be produced, and counts it.
A failed turn is still an HTTP 200 response, so the middleware's status-code
tally cannot see it. This metric exists so that a demo whose model refuses
every request does not report as healthy.
"""
analytics.record(analytics.M_EVENT, analytics.EV_TURN_ERROR)
return sse({"type": "error", "detail": detail, **extra})
# `no-cache` stops an intermediary from caching the stream. `X-Accel-Buffering`
# makes nginx-style reverse proxies, which hosted deploys use, flush each event
# immediately rather than buffer it.
SSE_HEADERS = {"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}
def action_json(action: models.Action, db: Session | None = None) -> dict:
"""Serializes one action for the wire.
@@ -425,8 +404,8 @@ def create_action(
request: Request,
db: Session = Depends(get_db),
user: models.User = CurrentUser,
adventure: models.Adventure = Depends(current_adventure),
):
adventure = get_adventure_or_404(adventure_id, db, user)
limits.rate_limit("turn", request, user)
limits.check_row_cap("actions", db, user, adventure=adventure)
check_demo_cap(db, user)
+1 -1
View File
@@ -19,7 +19,7 @@ from sqlalchemy.orm import Session
from .. import auth, limits, models, schemas
from ..database import get_db
from ..providers import OpenAICompatibleProvider, ProviderError
from .adventures import SSE_HEADERS, sse
from ..sse import SSE_HEADERS, sse
from .settings import get_settings, list_endpoint_models
router = APIRouter(prefix="/api/chat", tags=["chat"])