M2: cut the hosted product away from the local one
94 files, +1,395 -6,578. Three files are new; twenty-four are gone. The milestone is subtraction, and what is left is the single-user local storyteller the specification describes. Removed in full: campaign scripting and its QuickJS sandbox; multi-user accounts, guest sessions, login, registration and the shared demo key; the visitor-analytics tables, dashboard and page beacon; the access log of sign-ins, addresses and devices; per-IP and per-user rate limiting and quotas; Render deployment config; Postgres and psycopg; cloud inference providers, the API-key field and the key encryption that existed to store it; session-cookie signing. None of it was hidden behind a flag — the routes are gone and answer 404. Two things were kept that the brief allowed keeping. The `users` table and its foreign keys stay as an internal ownership detail, because rewriting them out means a migration across most of the schema to delete a column that costs nothing; nothing creates a second user and no request carries an identity. Five inert tables and four inert columns stay for the same reason, so an M1 campaign database opens unchanged. The one addition is app/endpoints.py, which decides where a story may be sent. Loopback, RFC1918, link-local, unique-local and CGNAT — an explicit allowlist of networks, not a guess at what `ipaddress` means by "private", which calls the documentation ranges private and IPv6 loopback reserved. Every address a hostname resolves to must be in it, so a split answer does not squeak through, and the rule runs both when the endpoint is saved and before every outbound request, because a name that resolved to the LAN this morning can resolve elsewhere this afternoon. Known cloud hosts are named in the refusal so the error says why rather than looking like broken DNS. TLS is never traded against it: M1's shared trust context is intact on all four clients and there is no way to skip verification. The hardcoded 120-second model timeout is now a setting. That was not theoretical — on this GPU-less four-core host a cold load of qwen2.5:3b-instruct took 648.9 seconds to produce the first turn, while turns 2 to 5 of the same campaign took 3.6 to 13.1. Connect stays short at 10s so a wrong address still fails fast; the read timeout defaults to 300s and is bounded at 3600, because "wait longer" must stay a number. Two defects found while testing and fixed here. An unknown /api path fell through the SPA catch-all and came back as HTML with status 200, so a client asking for JSON parsed a web page instead of learning the route was gone. And AIDND_CORS_ORIGINS accepted "*", which on an unauthenticated loopback API would hand every page on the Internet a write handle on the campaign database; it now refuses to start. Verified rather than assumed. Offline, on a network with no route out and no DNS: five turns, retry with both takes retained, restart with an identical transcript digest, a failed model call leaving the accepted AI-turn count untouched, and a capture with zero non-loopback unicast packets. Against a real second machine on the LAN over HTTPS with a private CA: four turns, restart, and a capture showing 289 packets to the approved host, 344 loopback, zero anywhere else, zero DNS queries. Cloud and public endpoints refused with their reasons; no API key settable; every removed route 404. 604 backend tests pass, down from 648 by the fifteen retired with the subsystems they tested and up by the twenty-nine added for the endpoint policy and the removed surface. The scripting tests were not deleted: eight files used a JavaScript counter as instrumentation for the state snapshot and rollback machinery, which M2 does not touch, so the counter moved to the world-state engine and those tests still assert what they always did. Frontend lint and build are clean; the image builds, and its wheel-building stage is gone with quickjs. No M3 work. Undo is still destructive and there is still no Redo.
This commit is contained in:
@@ -1,168 +0,0 @@
|
||||
"""The access log: who arrived, when, and from where.
|
||||
|
||||
The deliberate opposite of analytics.py. That module counts and stores nothing
|
||||
that points at a person; this one records addresses, email addresses and
|
||||
devices, because an access log that cannot identify the access is not an access
|
||||
log. The two live in separate modules and separate tables on purpose, so that
|
||||
the anonymity of the counters is a property of the code rather than a convention
|
||||
someone has to remember.
|
||||
|
||||
Owner-only, and never shown to the people it records.
|
||||
|
||||
Four kinds of row:
|
||||
|
||||
- `session` A browser that has a session made a request. For a guest, this
|
||||
is their first visit.
|
||||
- `login` An existing account signed in.
|
||||
- `register` A guest upgraded to an account.
|
||||
- `login_failed` A password attempt that did not match, with the address tried.
|
||||
|
||||
Session rows are the only ones that need thinning. `/auth/me` runs on every page
|
||||
load, and one row per load would be noise rather than a log. A row is written
|
||||
when the day or the address changes for that user. That is the granularity a log
|
||||
is read at, such as seen on the 3rd from 1.2.3.4, and it still records someone
|
||||
moving networks during a day.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import threading
|
||||
|
||||
from sqlalchemy import desc, or_, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from . import analytics, models
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SESSION = "session"
|
||||
LOGIN = "login"
|
||||
REGISTER = "register"
|
||||
LOGIN_FAILED = "login_failed"
|
||||
|
||||
MAX_UA = 200
|
||||
|
||||
# user id -> (day, ip) of the last session row written for them. Process-local
|
||||
# like the rate limiter's windows, and for the same reason: this is a single
|
||||
# process, and the worst case after a restart is one redundant row per user.
|
||||
_last_session: dict[int, tuple[str, str]] = {}
|
||||
_guard = threading.Lock()
|
||||
_MAX_TRACKED = 10_000
|
||||
|
||||
|
||||
def _client_ip(request) -> str:
|
||||
# This import is deferred. `limits` imports `auth`, which the routers that
|
||||
# call this function import, so a module-level import here would create a
|
||||
# cycle. The spoof resistance lives in `limits` and must not be
|
||||
# reimplemented. A second, looser answer to which address belongs to the
|
||||
# client is how one of them ends up trusting a header it should not.
|
||||
from . import limits
|
||||
|
||||
return limits.client_ip(request)
|
||||
|
||||
|
||||
def describe(user: models.User) -> str:
|
||||
"""Returns how a user is named in the log.
|
||||
|
||||
A guest has no email, and their id is the only handle anyone has for them.
|
||||
The third case is a local install's implicit single user, who also has no
|
||||
email but is the operator rather than a visitor. Naming that user "Guest #1"
|
||||
would be wrong in the one row they are certain to read.
|
||||
"""
|
||||
if user.email:
|
||||
return user.email
|
||||
return f"Guest #{user.id}" if user.is_guest else f"Local user #{user.id}"
|
||||
|
||||
|
||||
def _country(request) -> str:
|
||||
"""Returns the edge's country header, or "" when there is none.
|
||||
|
||||
The blank differs from the counters' "(unknown)" label. A table column reads
|
||||
better as a dash than as a word, and an empty string is the correct value for
|
||||
a country that is not known.
|
||||
"""
|
||||
country = analytics.country_of(request.headers)
|
||||
return "" if country == analytics.UNKNOWN else country
|
||||
|
||||
|
||||
def record(
|
||||
db: Session,
|
||||
kind: str,
|
||||
request,
|
||||
*,
|
||||
user: models.User | None = None,
|
||||
who: str | None = None,
|
||||
) -> None:
|
||||
"""Writes one row.
|
||||
|
||||
This function never raises. The log observes sign-in rather than guarding it,
|
||||
and a logging failure must not lock anyone out.
|
||||
"""
|
||||
try:
|
||||
event = models.AccessEvent(
|
||||
kind=kind,
|
||||
user_id=user.id if user is not None else None,
|
||||
who=(who if who is not None else describe(user) if user else "")[:320],
|
||||
is_guest=bool(user.is_guest) if user is not None else False,
|
||||
ip=_client_ip(request)[:45],
|
||||
country=_country(request),
|
||||
device=analytics.device_of(request.headers.get("user-agent", "")),
|
||||
user_agent=(request.headers.get("user-agent") or "")[:MAX_UA],
|
||||
)
|
||||
db.add(event)
|
||||
db.commit()
|
||||
except Exception: # pragma: no cover - defensive
|
||||
db.rollback()
|
||||
logger.exception("Access log write failed; continuing.")
|
||||
|
||||
|
||||
def note_session(db: Session, user: models.User, request) -> None:
|
||||
"""Records that a session made a request, at most one row per day per address."""
|
||||
try:
|
||||
today = analytics._today()
|
||||
ip = _client_ip(request)
|
||||
with _guard:
|
||||
if _last_session.get(user.id) == (today, ip):
|
||||
return
|
||||
_last_session[user.id] = (today, ip)
|
||||
if len(_last_session) > _MAX_TRACKED:
|
||||
# Nothing here needs to persist. Clearing the map costs at most
|
||||
# one extra row per active user.
|
||||
_last_session.clear()
|
||||
_last_session[user.id] = (today, ip)
|
||||
except Exception: # pragma: no cover - defensive
|
||||
logger.exception("Access log session check failed; continuing.")
|
||||
return
|
||||
record(db, SESSION, request, user=user)
|
||||
|
||||
|
||||
def recent(
|
||||
db: Session,
|
||||
*,
|
||||
limit: int = 50,
|
||||
before_id: int | None = None,
|
||||
kind: str | None = None,
|
||||
query: str | None = None,
|
||||
) -> dict:
|
||||
"""Returns a page of the log, newest first.
|
||||
|
||||
The page is anchored on a row id rather than an offset, as the story pager
|
||||
is. Rows keep arriving while the log is read, and an offset would shift the
|
||||
page under whoever is reading it.
|
||||
"""
|
||||
statement = select(models.AccessEvent).order_by(desc(models.AccessEvent.id))
|
||||
if before_id is not None:
|
||||
statement = statement.where(models.AccessEvent.id < before_id)
|
||||
if kind:
|
||||
statement = statement.where(models.AccessEvent.kind == kind)
|
||||
if query:
|
||||
like = f"%{query.strip()}%"
|
||||
statement = statement.where(or_(
|
||||
models.AccessEvent.who.ilike(like),
|
||||
models.AccessEvent.ip.ilike(like),
|
||||
models.AccessEvent.country.ilike(like),
|
||||
))
|
||||
# Requesting one extra row reports whether more rows exist, without a
|
||||
# second COUNT over the whole table.
|
||||
rows = list(db.scalars(statement.limit(limit + 1)))
|
||||
has_more = len(rows) > limit
|
||||
return {"events": rows[:limit], "has_more": has_more}
|
||||
@@ -1,583 +0,0 @@
|
||||
"""Visit analytics for the hosted demo.
|
||||
|
||||
This is a small self-hosted counter that answers whether anyone visited and
|
||||
whether they played. It is built into the app rather than added with a
|
||||
third-party script, because the CSP in `main.py` allows scripts from 'self'
|
||||
only, ad blockers block the popular trackers, and none of those trackers can see
|
||||
what is worth knowing here: turns taken, demo-key spend, and which seeded
|
||||
scenario people pick.
|
||||
|
||||
Three rules shape the design:
|
||||
|
||||
1. It stores nothing personal. It records no IP addresses, no user agents, no
|
||||
user ids, and no title of anything a player wrote. A visitor appears only as
|
||||
an HMAC of their user id, which is one-way and salted with the app's secret
|
||||
key, so these tables cannot be joined back to an account even by someone
|
||||
holding the database. Story content never reaches this module. What one
|
||||
specific person did is unanswerable by design, and only totals are
|
||||
available.
|
||||
2. Egress is the budget. Neon bills for bytes leaving the database, and this
|
||||
project has already paid for forgetting that once. Counts are therefore
|
||||
aggregated in memory and flushed as UPSERTs, so a visit is a write and never
|
||||
a read, and every dashboard query is a GROUP BY that returns tens of rows
|
||||
rather than per-visit rows. A month of traffic costs a few kilobytes to read
|
||||
back.
|
||||
3. The numbers come from the server, not from the browser. The client reports
|
||||
one thing, which is the page that was viewed. Everything with meaning, such
|
||||
as a turn happening or an account being created, is recorded by the code that
|
||||
performs it, where a stranger cannot fake it and an extension cannot block
|
||||
it.
|
||||
|
||||
Storage is two tables, both bounded. `analytics_daily` holds one counter row per
|
||||
day, metric, and label, which is a few dozen rows a day.
|
||||
`analytics_visitor_days` holds one row per visitor per day carrying the funnel
|
||||
flags, which is what makes the funnel count people rather than clicks. It is the
|
||||
only table that grows with traffic, and cleanup ages it out.
|
||||
"""
|
||||
|
||||
import hmac
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
from datetime import timedelta
|
||||
from hashlib import sha256
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from sqlalchemy import case, func, or_, select
|
||||
from sqlalchemy.dialects.postgresql import insert as pg_insert
|
||||
from sqlalchemy.dialects.sqlite import insert as sqlite_insert
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from . import models, security
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------- Metrics ----------
|
||||
# `metric` is the family, and `label` is the bucket within it. One generic
|
||||
# counter table is better than a column per measurement, because adding a new
|
||||
# question later costs nothing rather than a migration.
|
||||
|
||||
M_PAGE = "pageview"
|
||||
M_EVENT = "event"
|
||||
M_REFERRER = "referrer"
|
||||
M_DEVICE = "device"
|
||||
M_COUNTRY = "country"
|
||||
M_SCENARIO = "scenario" # Which seeded or public scenario was played.
|
||||
M_ERROR = "error" # "<status> <route>" for a 4xx or 5xx on /api.
|
||||
|
||||
EV_SCENARIO_OPEN = "scenario_opened"
|
||||
EV_ADVENTURE = "adventure_created"
|
||||
EV_IMPORT = "adventure_imported"
|
||||
EV_TURN = "turn"
|
||||
EV_DEMO_TURN = "demo_turn" # A turn billed to the shared demo key.
|
||||
EV_TURN_ERROR = "turn_error"
|
||||
EV_SIGNUP = "signup"
|
||||
EV_LOGIN = "login"
|
||||
|
||||
# Events that are also funnel steps. Recording one sets a flag on the visitor's
|
||||
# row for the day, so the funnel counts distinct visitor-days rather than repeat
|
||||
# clicks. This name-to-column map is the whole definition of the funnel, and the
|
||||
# dashboard reads it back in this order.
|
||||
FUNNEL_FLAGS = {
|
||||
EV_SCENARIO_OPEN: "opened",
|
||||
EV_ADVENTURE: "created",
|
||||
EV_TURN: "played",
|
||||
EV_SIGNUP: "signed_up",
|
||||
}
|
||||
|
||||
OTHER = "(other)"
|
||||
NONE_LABEL = "(direct)"
|
||||
UNKNOWN = "(unknown)"
|
||||
|
||||
# ---------- Bounds ----------
|
||||
# These bounds exist so that a hostile visitor can add rows to these tables no
|
||||
# faster than an honest one. The only label a client can influence is the
|
||||
# referrer, and together these caps mean the worst it can do is fill one day's
|
||||
# referrer list and then be folded into "(other)".
|
||||
|
||||
MAX_LABEL_LEN = 80
|
||||
MAX_LABELS_PER_METRIC = 200 # Distinct labels per metric per day, then OTHER.
|
||||
MAX_PENDING = 4000 # Buffered entries before an inline flush.
|
||||
FLUSH_INTERVAL_SECONDS = 60
|
||||
|
||||
# How long the per-visitor-day rows are kept. The daily counters are small and
|
||||
# are kept indefinitely. These rows are the ones that scale with traffic. A
|
||||
# visitor whose last visit ages out counts as new again, which is an acceptable
|
||||
# trade at this horizon and keeps the table from being a permanent record of
|
||||
# anyone.
|
||||
RETENTION_DAYS = int(os.environ.get("AIDND_ANALYTICS_RETENTION_DAYS", "400") or 400)
|
||||
|
||||
_HOST_OK = re.compile(r"^[a-z0-9.-]+$")
|
||||
_COUNTRY_OK = re.compile(r"^[A-Z]{2}$")
|
||||
_NUMERIC_SEGMENT = re.compile(r"^\d+$")
|
||||
|
||||
# SPA routes, in the form the dashboard shows them. Any other path a client
|
||||
# reports becomes OTHER, so the page list cannot be filled with junk and cannot
|
||||
# record which adventure someone is reading.
|
||||
KNOWN_ROUTES = {
|
||||
"/", "/adventures", "/scenarios", "/scenarios/:id", "/play/:id",
|
||||
"/scripts", "/scripts/:id", "/settings", "/chat", "/analytics",
|
||||
}
|
||||
|
||||
# ---------- In-process buffer ----------
|
||||
# The deployment is a single process, which is the same assumption `limits.py`
|
||||
# makes, so a plain dict under a lock is the whole design. Losing up to a minute
|
||||
# of counts to a hard restart is acceptable for traffic numbers, and the flusher
|
||||
# also runs on shutdown. On Render's free tier the service is idle when it
|
||||
# sleeps, so the buffer it sleeps on is empty.
|
||||
|
||||
_counts: dict[tuple[str, str, str], int] = {}
|
||||
_visits: dict[tuple[str, str], set[str]] = {} # (day, visitor) -> flags.
|
||||
_labels_seen: dict[tuple[str, str], set[str]] = {} # (day, metric) -> labels.
|
||||
_guard = threading.Lock()
|
||||
|
||||
|
||||
def _today() -> str:
|
||||
return models.utcnow().date().isoformat()
|
||||
|
||||
|
||||
def record(metric: str, label: str = "", *, n: int = 1) -> None:
|
||||
"""Adds `n` to one counter.
|
||||
|
||||
This function never raises. Analytics must not fail a request that it is
|
||||
only observing.
|
||||
"""
|
||||
try:
|
||||
day = _today()
|
||||
label = (label or "").strip()[:MAX_LABEL_LEN]
|
||||
with _guard:
|
||||
seen = _labels_seen.setdefault((day, metric), set())
|
||||
if label not in seen:
|
||||
if len(seen) >= MAX_LABELS_PER_METRIC:
|
||||
label = OTHER
|
||||
else:
|
||||
seen.add(label)
|
||||
key = (day, metric, label)
|
||||
_counts[key] = _counts.get(key, 0) + n
|
||||
pending = len(_counts) + len(_visits)
|
||||
except Exception: # pragma: no cover - defensive
|
||||
logger.exception("Analytics counter failed; continuing.")
|
||||
return
|
||||
if pending >= MAX_PENDING:
|
||||
flush()
|
||||
|
||||
|
||||
def visitor_id(user: models.User) -> str:
|
||||
"""Returns a stable, one-way handle for one visitor.
|
||||
|
||||
The handle is an HMAC of the user id under the app's secret key. It is
|
||||
stable, so a returning visitor can be distinguished from a new one. It is
|
||||
one-way, so nothing in the analytics tables points back at an account. It is
|
||||
keyed, so a client cannot compute one and claim to be someone else. One
|
||||
consequence follows: rotating `AIDND_SECRET_KEY` makes every returning
|
||||
visitor look new.
|
||||
"""
|
||||
digest = hmac.new(security.SECRET_KEY, f"visitor:{user.id}".encode(), sha256)
|
||||
return digest.hexdigest()[:32]
|
||||
|
||||
|
||||
def record_visit(user: models.User | None, *, flag: str | None = None) -> None:
|
||||
"""Records that this visitor was here today, and optionally sets one funnel
|
||||
flag.
|
||||
|
||||
Without a user the call does nothing. A page loaded before a session exists
|
||||
still counts as a pageview, but not as a person.
|
||||
"""
|
||||
if user is None:
|
||||
return
|
||||
try:
|
||||
with _guard:
|
||||
flags = _visits.setdefault((_today(), visitor_id(user)), set())
|
||||
if flag:
|
||||
flags.add(flag)
|
||||
except Exception: # pragma: no cover - defensive
|
||||
logger.exception("Analytics visit failed; continuing.")
|
||||
|
||||
|
||||
def record_event(name: str, user: models.User | None = None) -> None:
|
||||
"""Records one event, and credits the visitor's day if it is a funnel step.
|
||||
|
||||
This is the whole interface the call sites use.
|
||||
"""
|
||||
record(M_EVENT, name)
|
||||
record_visit(user, flag=FUNNEL_FLAGS.get(name))
|
||||
|
||||
|
||||
# ---------- Normalizing what the browser reports ----------
|
||||
|
||||
def normalize_route(path: str) -> str:
|
||||
"""Reduces a client-reported path to one of `KNOWN_ROUTES`.
|
||||
|
||||
Numeric segments become ":id". That bounds the label count, and it keeps
|
||||
which adventure someone opened out of the statistics.
|
||||
"""
|
||||
path = (path or "/").split("?")[0].split("#")[0]
|
||||
if not path.startswith("/"):
|
||||
path = "/" + path
|
||||
if len(path) > 1:
|
||||
path = path.rstrip("/")
|
||||
parts = [":id" if _NUMERIC_SEGMENT.match(p) else p for p in path.split("/")]
|
||||
route = "/".join(parts) or "/"
|
||||
return route if route in KNOWN_ROUTES else OTHER
|
||||
|
||||
|
||||
def normalize_referrer(referrer: str, own_host: str = "") -> str:
|
||||
"""Returns the sending site as a bare host.
|
||||
|
||||
This app's own host means an internal navigation, which is not a referral.
|
||||
In that case the function returns "", which tells the caller to skip it.
|
||||
"""
|
||||
if not referrer:
|
||||
return NONE_LABEL
|
||||
host = (urlsplit(referrer).hostname or "").lower().lstrip(".")
|
||||
if not host or not _HOST_OK.match(host) or len(host) > MAX_LABEL_LEN:
|
||||
return OTHER
|
||||
if host == (own_host or "").lower() or host in ("localhost", "127.0.0.1"):
|
||||
return ""
|
||||
return host[4:] if host.startswith("www.") else host
|
||||
|
||||
|
||||
def api_route_label(scope: dict, status: int) -> str:
|
||||
"""Returns an error bucket such as "500 /api/adventures/{adventure_id}".
|
||||
|
||||
The label uses the route template, never the request path. That keeps one
|
||||
bucket per endpoint rather than one per adventure id. It also bounds the
|
||||
table: an unmatched path is chosen entirely by the caller, so labeling by it
|
||||
would let anyone create rows by requesting arbitrary paths.
|
||||
"""
|
||||
template = getattr(scope.get("route"), "path", None)
|
||||
return f"{status} {template}" if template else f"{status} (unmatched)"
|
||||
|
||||
|
||||
def device_of(user_agent: str) -> str:
|
||||
"""Returns "mobile", "tablet", or "desktop", and nothing more specific.
|
||||
|
||||
The user-agent string itself is never stored, because it is a fingerprint
|
||||
and the useful answer is one word.
|
||||
"""
|
||||
ua = (user_agent or "").lower()
|
||||
if not ua:
|
||||
return UNKNOWN
|
||||
if any(bot in ua for bot in ("bot", "crawler", "spider", "headless", "preview")):
|
||||
return "bot"
|
||||
if "ipad" in ua or "tablet" in ua or ("android" in ua and "mobile" not in ua):
|
||||
return "tablet"
|
||||
if any(m in ua for m in ("mobi", "iphone", "ipod", "android", "phone")):
|
||||
return "mobile"
|
||||
return "desktop"
|
||||
|
||||
|
||||
# Geo headers an edge network may add. Render fronts services with a CDN that
|
||||
# can set `cf-ipcountry`, and the others cost nothing to check. A value is
|
||||
# trusted only if it looks like an ISO code, because a client can send any
|
||||
# header, so the worst case is a wrong country rather than an unbounded label.
|
||||
_GEO_HEADERS = ("cf-ipcountry", "x-vercel-ip-country", "x-geo-country", "x-country-code")
|
||||
|
||||
|
||||
def country_of(headers) -> str:
|
||||
for name in _GEO_HEADERS:
|
||||
value = (headers.get(name) or "").strip().upper()
|
||||
if _COUNTRY_OK.match(value) and value != "XX":
|
||||
return value
|
||||
return UNKNOWN
|
||||
|
||||
|
||||
# ---------- Flushing ----------
|
||||
|
||||
def _insert(db: Session):
|
||||
return sqlite_insert if db.get_bind().dialect.name == "sqlite" else pg_insert
|
||||
|
||||
|
||||
def _drain() -> tuple[dict, dict]:
|
||||
with _guard:
|
||||
counts, visits = _counts.copy(), _visits.copy()
|
||||
_counts.clear()
|
||||
_visits.clear()
|
||||
# The label sets bound cardinality within one day, so drop the
|
||||
# previous day's rather than grow a map that never shrinks.
|
||||
today = _today()
|
||||
for key in [k for k in _labels_seen if k[0] != today]:
|
||||
del _labels_seen[key]
|
||||
return counts, visits
|
||||
|
||||
|
||||
def _restore(counts: dict, visits: dict) -> None:
|
||||
"""Returns a failed flush's work to the buffer, so the next flush retries it."""
|
||||
with _guard:
|
||||
for key, n in counts.items():
|
||||
_counts[key] = _counts.get(key, 0) + n
|
||||
for key, flags in visits.items():
|
||||
_visits.setdefault(key, set()).update(flags)
|
||||
|
||||
|
||||
def flush(db: Session | None = None) -> None:
|
||||
"""Writes the buffer out. This is safe to call from anywhere and never raises."""
|
||||
counts, visits = _drain()
|
||||
if not counts and not visits:
|
||||
return
|
||||
own_session = db is None
|
||||
if own_session:
|
||||
from .database import SessionLocal
|
||||
db = SessionLocal()
|
||||
try:
|
||||
_write_counts(db, counts)
|
||||
_write_visits(db, visits)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
_restore(counts, visits)
|
||||
logger.exception("Analytics flush failed; counts held for the next one.")
|
||||
finally:
|
||||
if own_session:
|
||||
db.close()
|
||||
|
||||
|
||||
def _write_counts(db: Session, counts: dict) -> None:
|
||||
if not counts:
|
||||
return
|
||||
table = models.AnalyticsDaily.__table__
|
||||
rows = [
|
||||
{"day": day, "metric": metric, "label": label, "hits": hits}
|
||||
for (day, metric, label), hits in counts.items()
|
||||
]
|
||||
stmt = _insert(db)(table).values(rows)
|
||||
db.execute(stmt.on_conflict_do_update(
|
||||
index_elements=["day", "metric", "label"],
|
||||
set_={"hits": table.c.hits + stmt.excluded.hits},
|
||||
))
|
||||
|
||||
|
||||
def _write_visits(db: Session, visits: dict) -> None:
|
||||
if not visits:
|
||||
return
|
||||
table = models.AnalyticsVisitorDay.__table__
|
||||
ids = {visitor for _, visitor in visits}
|
||||
# One indexed lookup decides new against returning for the whole batch. It
|
||||
# is the only read this module makes outside the dashboard, and it returns
|
||||
# short hashes for the visitors active right now, so the batch bounds it.
|
||||
known = set(db.scalars(
|
||||
select(models.AnalyticsVisitorDay.visitor)
|
||||
.where(models.AnalyticsVisitorDay.visitor.in_(ids))
|
||||
.distinct()
|
||||
))
|
||||
rows = [
|
||||
{
|
||||
"day": day,
|
||||
"visitor": visitor,
|
||||
"is_new": visitor not in known,
|
||||
**{column: column in flags for column in FUNNEL_FLAGS.values()},
|
||||
}
|
||||
for (day, visitor), flags in visits.items()
|
||||
]
|
||||
stmt = _insert(db)(table).values(rows)
|
||||
db.execute(stmt.on_conflict_do_update(
|
||||
index_elements=["day", "visitor"],
|
||||
# Flags only turn on, and `is_new` is absent on purpose. The first
|
||||
# write of a visitor's first day is what decided it.
|
||||
set_={
|
||||
column: or_(table.c[column], stmt.excluded[column])
|
||||
for column in FUNNEL_FLAGS.values()
|
||||
},
|
||||
))
|
||||
|
||||
|
||||
def purge_old_visitor_days(db: Session) -> int:
|
||||
"""Deletes visitor-day rows past the retention horizon.
|
||||
|
||||
The cleanup sweeper calls this. The daily counters are never purged, because
|
||||
they are aggregates, they are small, and this project keeps its history.
|
||||
"""
|
||||
if RETENTION_DAYS <= 0:
|
||||
return 0
|
||||
cutoff = (models.utcnow().date() - timedelta(days=RETENTION_DAYS)).isoformat()
|
||||
removed = db.query(models.AnalyticsVisitorDay).filter(
|
||||
models.AnalyticsVisitorDay.day < cutoff
|
||||
).delete(synchronize_session=False)
|
||||
db.commit()
|
||||
return removed or 0
|
||||
|
||||
|
||||
# ---------- Reading it back ----------
|
||||
# Every query below is an aggregate. The database does the counting and returns
|
||||
# tens of rows, however much traffic is behind them. No query here can return a
|
||||
# row that belongs to one visitor.
|
||||
|
||||
TOP_N = 12
|
||||
|
||||
|
||||
def _top(rows: list[dict], limit: int = TOP_N) -> list[dict]:
|
||||
return rows[:limit]
|
||||
|
||||
|
||||
def summary(db: Session, days: int = 30) -> dict:
|
||||
"""Returns everything the dashboard shows for the last `days` days, including
|
||||
today.
|
||||
|
||||
The function flushes first, so the numbers include the last minute.
|
||||
"""
|
||||
flush(db)
|
||||
today = models.utcnow().date()
|
||||
since = (today - timedelta(days=days - 1)).isoformat()
|
||||
daily = models.AnalyticsDaily
|
||||
visitor = models.AnalyticsVisitorDay
|
||||
|
||||
# 1. Every counter in the window, reduced to (metric, label) totals. The
|
||||
# page, referrer, country, device, scenario, and error tables all come
|
||||
# from this one pass rather than from a query each.
|
||||
by_metric: dict[str, list[dict]] = {}
|
||||
for metric, label, hits in db.execute(
|
||||
select(daily.metric, daily.label, func.sum(daily.hits))
|
||||
.where(daily.day >= since)
|
||||
.group_by(daily.metric, daily.label)
|
||||
):
|
||||
by_metric.setdefault(metric, []).append({"label": label, "hits": int(hits)})
|
||||
for rows in by_metric.values():
|
||||
rows.sort(key=lambda row: -row["hits"])
|
||||
events = {row["label"]: row["hits"] for row in by_metric.get(M_EVENT, [])}
|
||||
|
||||
# 2. The two per-day series the dashboard draws.
|
||||
pageviews_by_day = {
|
||||
day: int(hits)
|
||||
for day, hits in db.execute(
|
||||
select(daily.day, func.sum(daily.hits))
|
||||
.where(daily.day >= since, daily.metric == M_PAGE)
|
||||
.group_by(daily.day)
|
||||
)
|
||||
}
|
||||
turns_by_day = {
|
||||
day: int(hits)
|
||||
for day, hits in db.execute(
|
||||
select(daily.day, func.sum(daily.hits))
|
||||
.where(daily.day >= since, daily.metric == M_EVENT, daily.label == EV_TURN)
|
||||
.group_by(daily.day)
|
||||
)
|
||||
}
|
||||
|
||||
# 3. People, per day. There is one row per visitor per day, so COUNT(*) is
|
||||
# already the day's unique visitors and no DISTINCT is needed.
|
||||
visitors_by_day: dict[str, dict] = {}
|
||||
for day, total, fresh in db.execute(
|
||||
select(
|
||||
visitor.day,
|
||||
func.count(),
|
||||
func.sum(case((visitor.is_new, 1), else_=0)),
|
||||
)
|
||||
.where(visitor.day >= since)
|
||||
.group_by(visitor.day)
|
||||
):
|
||||
visitors_by_day[day] = {"visitors": int(total), "new": int(fresh or 0)}
|
||||
|
||||
# 4. The funnel over the whole window, counting each person once.
|
||||
# COUNT(DISTINCT CASE WHEN flag THEN visitor END) ignores the NULLs the
|
||||
# CASE leaves for everyone who did not reach that step.
|
||||
unique, unique_new, *reached = db.execute(
|
||||
select(
|
||||
func.count(func.distinct(visitor.visitor)),
|
||||
func.count(func.distinct(case((visitor.is_new, visitor.visitor)))),
|
||||
*[
|
||||
func.count(func.distinct(case((visitor.__table__.c[column], visitor.visitor))))
|
||||
for column in FUNNEL_FLAGS.values()
|
||||
],
|
||||
).where(visitor.day >= since)
|
||||
).one()
|
||||
|
||||
series = []
|
||||
for offset in range(days):
|
||||
day = (today - timedelta(days=days - 1 - offset)).isoformat()
|
||||
counted = visitors_by_day.get(day, {})
|
||||
series.append({
|
||||
"day": day,
|
||||
"visitors": counted.get("visitors", 0),
|
||||
"new": counted.get("new", 0),
|
||||
"pageviews": pageviews_by_day.get(day, 0),
|
||||
"turns": turns_by_day.get(day, 0),
|
||||
})
|
||||
|
||||
visits = sum(row["visitors"] for row in series)
|
||||
pageviews = sum(pageviews_by_day.values())
|
||||
turns = events.get(EV_TURN, 0)
|
||||
errors = by_metric.get(M_ERROR, [])
|
||||
return {
|
||||
"days": days,
|
||||
"since": since,
|
||||
"until": today.isoformat(),
|
||||
"generated_at": models.utcnow().isoformat(),
|
||||
"totals": {
|
||||
# `visitors` counts each person once for the window. `visits`
|
||||
# counts them once per day they returned, which is the closest
|
||||
# measure to "sessions" that does not track sessions.
|
||||
"visitors": int(unique),
|
||||
"new_visitors": int(unique_new),
|
||||
"visits": visits,
|
||||
"pageviews": pageviews,
|
||||
"turns": turns,
|
||||
"demo_turns": events.get(EV_DEMO_TURN, 0),
|
||||
"adventures": events.get(EV_ADVENTURE, 0),
|
||||
"signups": events.get(EV_SIGNUP, 0),
|
||||
"logins": events.get(EV_LOGIN, 0),
|
||||
"turn_errors": events.get(EV_TURN_ERROR, 0),
|
||||
"errors": sum(row["hits"] for row in errors),
|
||||
"turns_per_visit": round(turns / visits, 1) if visits else 0,
|
||||
"pages_per_visit": round(pageviews / visits, 1) if visits else 0,
|
||||
},
|
||||
"series": series,
|
||||
# Step 0 is everyone who arrived, so the drop-off between it and
|
||||
# "Opened a scenario" appears as a step like any other.
|
||||
"funnel": [{"step": "Visited", "count": int(unique)}] + [
|
||||
{"step": step, "count": int(count)}
|
||||
for step, count in zip(
|
||||
["Opened a scenario", "Started an adventure", "Played a turn", "Signed up"],
|
||||
reached,
|
||||
)
|
||||
],
|
||||
"pages": _top(by_metric.get(M_PAGE, [])),
|
||||
"referrers": _top(by_metric.get(M_REFERRER, [])),
|
||||
"countries": _top(by_metric.get(M_COUNTRY, [])),
|
||||
"devices": by_metric.get(M_DEVICE, []),
|
||||
"scenarios": _top(by_metric.get(M_SCENARIO, [])),
|
||||
"errors": _top(errors),
|
||||
"events": by_metric.get(M_EVENT, []),
|
||||
}
|
||||
|
||||
|
||||
# ---------- Background flusher ----------
|
||||
# This matches the start and stop pair in `cleanup`, so the lifespan in
|
||||
# `main.py` reads the same way for both. The interval bounds how much a hard
|
||||
# restart can lose.
|
||||
|
||||
async def _flush_loop() -> None:
|
||||
import asyncio
|
||||
|
||||
from starlette.concurrency import run_in_threadpool
|
||||
|
||||
while True:
|
||||
await asyncio.sleep(FLUSH_INTERVAL_SECONDS)
|
||||
# This is blocking database work, so keep it off the event loop, which
|
||||
# is also serving SSE turn streams.
|
||||
await run_in_threadpool(flush)
|
||||
|
||||
|
||||
def start_flusher():
|
||||
import asyncio
|
||||
|
||||
return asyncio.create_task(_flush_loop())
|
||||
|
||||
|
||||
async def stop_flusher(task) -> None:
|
||||
"""Cancels the loop and writes out whatever it was holding.
|
||||
|
||||
A deploy is the one restart that is both frequent and predictable, so it
|
||||
should not be what loses a minute of counts.
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
from starlette.concurrency import run_in_threadpool
|
||||
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
await run_in_threadpool(flush)
|
||||
@@ -41,11 +41,11 @@ from .context import lineage
|
||||
|
||||
# The slices of a context snapshot that belong to one attempt rather than to the
|
||||
# turn. They are the world-state delta the attempt proposed and what the engine
|
||||
# did with it, the script report, the model's literal reply, and the endpoint's
|
||||
# did with it, the model's literal reply, and the endpoint's
|
||||
# token accounting. Each attempt is its own API call, and a retry is the call
|
||||
# most likely to read the prompt back out of cache. Everything else in a snapshot
|
||||
# is the prompt, which is assembled once per turn.
|
||||
ATTEMPT_KEYS = ("world_state", "script", "raw_output", "usage")
|
||||
ATTEMPT_KEYS = ("world_state", "raw_output", "usage")
|
||||
|
||||
|
||||
# ------------------------------------------------------------------ reading
|
||||
@@ -152,7 +152,7 @@ def preceding(
|
||||
# ------------------------------------------------------------------ writing
|
||||
|
||||
def restore_state(adventure: models.Adventure, node: models.Action | None) -> None:
|
||||
"""Restores the script state and world state that `node` left behind.
|
||||
"""Restores the world state that `node` left behind.
|
||||
|
||||
A NULL snapshot means leave the live state as it is, never reset it. Rows
|
||||
written before SP4 that the migration could not derive an outcome for carry
|
||||
@@ -161,17 +161,16 @@ def restore_state(adventure: models.Adventure, node: models.Action | None) -> No
|
||||
"""
|
||||
if node is None:
|
||||
return
|
||||
if isinstance(node.state_after, dict):
|
||||
adventure.script_state = copy.deepcopy(node.state_after)
|
||||
if isinstance(node.world_state_after, dict):
|
||||
adventure.world_state = copy.deepcopy(node.world_state_after)
|
||||
|
||||
|
||||
def snapshot_outcome(adventure: models.Adventure, node: models.Action) -> None:
|
||||
"""Records on `node` the state of the adventure now that the node has played."""
|
||||
state = adventure.script_state if isinstance(adventure.script_state, dict) else {}
|
||||
world = adventure.world_state if isinstance(adventure.world_state, dict) else {}
|
||||
node.state_after = copy.deepcopy(state)
|
||||
# `state_after` held the scripting engine's shared state, which M2 removed.
|
||||
# The column stays for schema compatibility and is written empty.
|
||||
node.state_after = {}
|
||||
node.world_state_after = copy.deepcopy(world)
|
||||
|
||||
|
||||
|
||||
+34
-219
@@ -1,216 +1,44 @@
|
||||
"""Phase 8: user resolution, sessions, and the shared demo key.
|
||||
"""Resolving the one local user. **There is no authentication in this product.**
|
||||
|
||||
The `AIDND_MULTI_USER` environment variable selects one of two modes:
|
||||
The module keeps its name so the dependency every router already depends on
|
||||
keeps working, but nothing here authenticates anybody. The Adventure
|
||||
Storyteller is a single-user application that binds to loopback: whoever can
|
||||
reach the API is the person who started it, and there is nobody else to tell
|
||||
them apart from.
|
||||
|
||||
* Local mode, the default. Every request resolves to one automatically created
|
||||
local user. There are no cookies and no login UI, so a clone or a
|
||||
docker-compose run behaves like the single-user app from before Phase 8.
|
||||
* Multi-user mode, used for hosted deployments. Requests carry a signed session
|
||||
cookie. `GET /api/auth/me` creates a guest user on the first visit, and
|
||||
registering upgrades that guest in place so their data survives. A request
|
||||
without a valid session gets a 401, and the frontend re-establishes the
|
||||
session through `/me`.
|
||||
Upstream had two modes. `AIDND_MULTI_USER` selected a hosted deployment with
|
||||
signed session cookies, guest accounts, registration, login, a shared demo API
|
||||
key with a per-day cap, "power users", and an owner allowlist for the analytics
|
||||
dashboard. M2 removed all of it: this product has no hosted mode to protect, and
|
||||
every one of those surfaces was a way for the application to be reached by
|
||||
someone other than its owner.
|
||||
|
||||
The shared demo key, which is the fallback when a user brings no key of their
|
||||
own, is also configured here. A user whose settings hold no API key is routed to
|
||||
a server-funded endpoint with a model allowlist and a per-day turn cap.
|
||||
What is left is the local path that upstream already had. Every request
|
||||
resolves to one automatically created user row.
|
||||
|
||||
The `users` table and the `user_id` foreign keys on scenarios, adventures and
|
||||
settings stay. They are an **internal ownership detail**, not an account
|
||||
system: nothing creates a second user, nothing logs in, and no request carries
|
||||
an identity. They remain because rewriting them out would mean a migration
|
||||
across most of the schema to delete a column that costs nothing and keeps every
|
||||
existing M1 database readable.
|
||||
"""
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from datetime import timezone
|
||||
|
||||
from fastapi import Depends, HTTPException, Request
|
||||
from fastapi import Depends, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from . import models, security
|
||||
from . import models
|
||||
from .database import get_db
|
||||
|
||||
|
||||
def _env_flag(name: str) -> bool:
|
||||
return os.environ.get(name, "").strip().lower() in ("1", "true", "yes", "on")
|
||||
|
||||
|
||||
MULTI_USER = _env_flag("AIDND_MULTI_USER")
|
||||
|
||||
SESSION_COOKIE = "aidnd_session"
|
||||
# Secure cookies are on by default in multi-user mode, because a hosted
|
||||
# deployment serves HTTPS and browsers also accept Secure on http://localhost.
|
||||
# `AIDND_COOKIE_SECURE` overrides the default with 0 or 1. Use 0 when testing
|
||||
# multi-user mode over plain HTTP on a LAN address.
|
||||
_cookie_secure_env = os.environ.get("AIDND_COOKIE_SECURE", "").strip().lower()
|
||||
COOKIE_SECURE = (
|
||||
_cookie_secure_env in ("1", "true", "yes", "on")
|
||||
if _cookie_secure_env
|
||||
else MULTI_USER
|
||||
)
|
||||
COOKIE_MAX_AGE = 60 * 60 * 24 * 365
|
||||
|
||||
# ---------- Shared demo key (BYOK fallback) ----------
|
||||
|
||||
DEMO_API_KEY = os.environ.get("AIDND_DEMO_API_KEY", "").strip()
|
||||
DEMO_ENDPOINT_URL = (
|
||||
os.environ.get("AIDND_DEMO_ENDPOINT_URL", "").strip()
|
||||
or "https://openrouter.ai/api/v1"
|
||||
)
|
||||
DEMO_MODELS = [
|
||||
m.strip()
|
||||
for m in os.environ.get("AIDND_DEMO_MODELS", "").split(",")
|
||||
if m.strip()
|
||||
] or ["google/gemma-4-26b-a4b-it:free"]
|
||||
DEMO_TURNS_PER_DAY = int(os.environ.get("AIDND_DEMO_TURNS_PER_DAY", "20") or 20)
|
||||
|
||||
# Trusted testers, listed by email, who bypass the daily demo cap and take
|
||||
# unmetered turns on the shared demo key. The list is comma-separated, and the
|
||||
# match ignores case.
|
||||
POWER_USERS = {
|
||||
e.strip().lower()
|
||||
for e in os.environ.get("AIDND_POWER_USERS", "").split(",")
|
||||
if e.strip()
|
||||
}
|
||||
|
||||
# Who can see the visit analytics. This is a separate list from `POWER_USERS` on
|
||||
# purpose. A trusted tester gets unmetered turns and the AI Chat page, which is
|
||||
# not a reason to give them the site's traffic numbers. An empty list, which is
|
||||
# the default, means nobody sees the dashboard in a hosted deployment.
|
||||
ANALYTICS_EMAILS = {
|
||||
e.strip().lower()
|
||||
for e in os.environ.get("AIDND_ANALYTICS_EMAILS", "").split(",")
|
||||
if e.strip()
|
||||
}
|
||||
|
||||
DEMO_CAP_MESSAGE = (
|
||||
f"You've used all {DEMO_TURNS_PER_DAY} free demo turns for today. "
|
||||
"Add your own API key in Settings to keep playing (it resets tomorrow)."
|
||||
)
|
||||
|
||||
|
||||
def demo_enabled() -> bool:
|
||||
# The demo key is a hosted-deployment feature. A local install talks to
|
||||
# whatever endpoint Settings points at, even with no API key, such as
|
||||
# Ollama.
|
||||
return MULTI_USER and bool(DEMO_API_KEY)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProviderConfig:
|
||||
"""What the turn engine connects with, after the decision between a
|
||||
user-supplied key and the demo key.
|
||||
|
||||
Build one of these with `resolve_provider_config()`.
|
||||
"""
|
||||
|
||||
endpoint_url: str
|
||||
api_key: str
|
||||
model: str
|
||||
using_demo: bool
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
# A second guard around server-funded turns. `resolve_provider_config()`
|
||||
# already pins the model, and this makes the pin a property of the config
|
||||
# object too, so a later caller cannot construct an unpinned one. This
|
||||
# raise is unreachable by design. Reaching it means a new code path
|
||||
# bypassed the pinning, which is worth failing on rather than billing
|
||||
# for.
|
||||
#
|
||||
# The test is `using_demo`, not `api_key == DEMO_API_KEY`. Keying on the
|
||||
# key value looks stricter and is wrong. The demo key is an ordinary
|
||||
# OpenRouter key, so a user can legitimately paste that same key into
|
||||
# their own Settings. Every resolution then raised, which returned a 500
|
||||
# even from `GET /auth/me` and took the whole SPA down. `using_demo` is
|
||||
# what means the server is paying, and only the demo branch below sets
|
||||
# it.
|
||||
if self.using_demo and self.model not in DEMO_MODELS:
|
||||
raise ValueError(
|
||||
f"Refusing to use the shared demo key with non-whitelisted model {self.model!r}"
|
||||
)
|
||||
|
||||
|
||||
def resolve_provider_config(
|
||||
settings: models.Settings, *, model_override: str | None = None
|
||||
) -> ProviderConfig:
|
||||
"""Returns the user's own key when they have one, and the shared demo key
|
||||
otherwise.
|
||||
|
||||
The demo branch is the security-relevant one, and it is the only place the
|
||||
allowlist rule lives. Every caller has to come through this function rather
|
||||
than build a `ProviderConfig` itself. On the demo key:
|
||||
|
||||
* The model is pinned to `DEMO_MODELS`, so a caller-supplied override from
|
||||
the AI Chat page, or a hand-edited Settings row, cannot point a
|
||||
server-funded key at a paid model. An unrecognized model falls back to
|
||||
`DEMO_MODELS[0]`.
|
||||
* The endpoint is pinned to `DEMO_ENDPOINT_URL`, so the key cannot be
|
||||
redirected to a URL the user controls and captured there.
|
||||
|
||||
`model_override` is a per-request preference and never a grant. It is used
|
||||
verbatim with the user's own key, and on the demo key only when the model is
|
||||
on the allowlist.
|
||||
"""
|
||||
key = settings.api_key_plain
|
||||
requested = (model_override or "").strip() or settings.model
|
||||
if key or not demo_enabled():
|
||||
return ProviderConfig(settings.endpoint_url, key, requested, False)
|
||||
model = requested if requested in DEMO_MODELS else DEMO_MODELS[0]
|
||||
return ProviderConfig(DEMO_ENDPOINT_URL, DEMO_API_KEY, model, True)
|
||||
|
||||
|
||||
def _today() -> str:
|
||||
return models.utcnow().date().isoformat()
|
||||
|
||||
|
||||
def is_power_user(user: models.User) -> bool:
|
||||
"""Returns whether this user is a trusted tester.
|
||||
|
||||
A trusted tester gets unmetered demo turns, plus tooling that is not part of
|
||||
the game, such as the AI Chat scratchpad. A local install is always trusted,
|
||||
because it runs on the operator's own machine with their own API key. The
|
||||
provider debug log is local-only for the same reason.
|
||||
"""
|
||||
if not MULTI_USER:
|
||||
return True
|
||||
return bool(user.email) and user.email.lower() in POWER_USERS
|
||||
|
||||
|
||||
def is_owner(user: models.User) -> bool:
|
||||
"""Returns whether this user may see the visit analytics.
|
||||
|
||||
A local install always may, because it runs on the operator's own machine and
|
||||
shows their own visits. The provider debug log follows the same reasoning. A
|
||||
hosted deployment checks `AIDND_ANALYTICS_EMAILS`.
|
||||
"""
|
||||
if not MULTI_USER:
|
||||
return True
|
||||
return bool(user.email) and user.email.lower() in ANALYTICS_EMAILS
|
||||
|
||||
|
||||
def demo_turns_left(user: models.User) -> int:
|
||||
# A power user is never capped, so report the full cap and let the banner
|
||||
# read "N of N" rather than count down.
|
||||
if is_power_user(user):
|
||||
return DEMO_TURNS_PER_DAY
|
||||
used = user.demo_turns_used if user.demo_turns_date == _today() else 0
|
||||
return max(0, DEMO_TURNS_PER_DAY - used)
|
||||
|
||||
|
||||
def count_demo_turn(user: models.User) -> None:
|
||||
"""Records one demo turn. The caller's commit stores it."""
|
||||
if is_power_user(user):
|
||||
return # A power user's turns do not count against the cap.
|
||||
today = _today()
|
||||
if user.demo_turns_date != today:
|
||||
user.demo_turns_date = today
|
||||
user.demo_turns_used = 0
|
||||
user.demo_turns_used += 1
|
||||
|
||||
|
||||
# ---------- User resolution ----------
|
||||
|
||||
def local_user(db: Session) -> models.User:
|
||||
"""Returns the single implicit user used in local mode.
|
||||
"""Returns the single implicit user, creating it on first use.
|
||||
|
||||
A migration gives this user ownership of data written before Phase 8. On a
|
||||
fresh database the user is created on first use.
|
||||
A migration gives this user ownership of data written before per-user rows
|
||||
existed, so an older database resolves to the row that already owns its
|
||||
campaigns rather than to a fresh empty one.
|
||||
"""
|
||||
user = (
|
||||
db.query(models.User)
|
||||
@@ -237,27 +65,14 @@ def _touch(user: models.User, db: Session) -> None:
|
||||
db.commit()
|
||||
|
||||
|
||||
def resolve_session_user(request: Request, db: Session) -> models.User | None:
|
||||
token = request.cookies.get(SESSION_COOKIE)
|
||||
if not token:
|
||||
return None
|
||||
user_id = security.verify_session(token)
|
||||
if user_id is None:
|
||||
return None
|
||||
return db.get(models.User, user_id)
|
||||
def get_current_user(
|
||||
request: Request, db: Session = Depends(get_db)
|
||||
) -> models.User:
|
||||
"""The dependency every router uses. It always succeeds.
|
||||
|
||||
|
||||
def get_current_user(request: Request, db: Session = Depends(get_db)) -> models.User:
|
||||
"""The dependency every router uses to resolve the current user.
|
||||
|
||||
In multi-user mode a 401 means the frontend has to establish a session again
|
||||
through `GET /api/auth/me`.
|
||||
`request` is unused and kept so the signature stays a FastAPI dependency
|
||||
the routers can depend on unchanged.
|
||||
"""
|
||||
if not MULTI_USER:
|
||||
user = local_user(db)
|
||||
else:
|
||||
user = resolve_session_user(request, db)
|
||||
if user is None:
|
||||
raise HTTPException(401, "No session. Call GET /api/auth/me first.")
|
||||
user = local_user(db)
|
||||
_touch(user, db)
|
||||
return user
|
||||
|
||||
+4
-24
@@ -109,7 +109,6 @@ def export(db: Session, adventure: models.Adventure) -> dict:
|
||||
"pronouns": adventure.persona_pronouns,
|
||||
"desc": adventure.persona_desc,
|
||||
},
|
||||
"scriptState": adventure.script_state,
|
||||
"worldState": adventure.world_state,
|
||||
"autoSummarize": adventure.auto_summarize,
|
||||
"memoryBankEnabled": adventure.memory_bank_enabled,
|
||||
@@ -126,15 +125,6 @@ def export(db: Session, adventure: models.Adventure) -> dict:
|
||||
"entry": c.entry, "notes": c.notes}
|
||||
for c in adventure.story_cards
|
||||
],
|
||||
"scripts": [
|
||||
{
|
||||
"position": s.position, "enabled": s.enabled,
|
||||
"name": s.name, "description": s.description,
|
||||
"library": s.library_js, "input": s.input_js,
|
||||
"context": s.context_js, "output": s.output_js,
|
||||
}
|
||||
for s in adventure.scripts
|
||||
],
|
||||
"actions": [_exported_node(a, local) for a in nodes],
|
||||
}
|
||||
|
||||
@@ -654,7 +644,6 @@ def materialize(
|
||||
authors_note=str(payload.get("authorsNote") or ""),
|
||||
ai_instructions=str(payload.get("aiInstructions") or ""),
|
||||
story_summary=str(payload.get("storySummary") or ""),
|
||||
script_state=payload.get("scriptState") or {},
|
||||
world_state=payload.get("worldState") or {},
|
||||
auto_summarize=bool(payload.get("autoSummarize", False)),
|
||||
memory_bank_enabled=bool(payload.get("memoryBankEnabled", False)),
|
||||
@@ -674,19 +663,10 @@ def materialize(
|
||||
notes=str(card.get("notes") or ""),
|
||||
))
|
||||
|
||||
for i, item in enumerate(payload.get("scripts") or []):
|
||||
if isinstance(item, dict):
|
||||
db.add(models.AdventureScript(
|
||||
adventure_id=adventure.id,
|
||||
position=int(item.get("position", i)),
|
||||
enabled=bool(item.get("enabled", True)),
|
||||
name=str(item.get("name") or "Imported Script")[:schemas.NAME_MAX],
|
||||
description=str(item.get("description") or ""),
|
||||
library_js=str(item.get("library") or ""),
|
||||
input_js=str(item.get("input") or ""),
|
||||
context_js=str(item.get("context") or ""),
|
||||
output_js=str(item.get("output") or ""),
|
||||
))
|
||||
# A bundle exported before M2 may carry "scripts" and "scriptState".
|
||||
# Campaign scripting is gone, so both are ignored rather than rejected: the
|
||||
# story, its tree, its cards and its memories still import intact, which is
|
||||
# what the bundle is for.
|
||||
|
||||
write(db, adventure, story)
|
||||
db.flush()
|
||||
|
||||
@@ -1,178 +0,0 @@
|
||||
"""Retention policy for throwaway guest accounts.
|
||||
|
||||
In multi-user mode every first visit creates a `users` row through
|
||||
`GET /api/auth/me`, so a public demo accumulates one account per visitor. Most
|
||||
of those visitors never return, and each one leaves behind whatever scenarios,
|
||||
adventures, actions, and memories they generated. This module deletes guests
|
||||
that have been inactive for `AIDND_GUEST_RETENTION_DAYS`, which defaults to 5,
|
||||
along with everything they made.
|
||||
|
||||
Why this is safe to run unattended:
|
||||
|
||||
- Only rows with `is_guest` AND `email IS NULL` are ever touched, and both
|
||||
clauses are checked rather than either alone. Registering upgrades the row
|
||||
in place (is_guest -> False), so a guest who signs up keeps everything;
|
||||
local mode's implicit single user is also is_guest=False.
|
||||
- Idle time is `COALESCE(last_seen_at, created_at)`. `auth._touch` writes
|
||||
`last_seen_at` at most once an hour, and a guest created by `/auth/me` has
|
||||
NULL there until its second request, so `created_at` is the correct floor for
|
||||
a new visitor. Without the coalesce, those rows look arbitrarily old.
|
||||
- Nothing a guest owns is reachable by anyone else. `is_public` is an
|
||||
output-only field, as `schemas.ScenarioBase` shows, so the only shared
|
||||
scenarios are the seeded ones, which have a NULL `user_id` and are outside
|
||||
this filter. Deleting a guest cannot remove content from another user.
|
||||
|
||||
The sweep uses one Core DELETE rather than an ORM cascade. `db.delete(user)`
|
||||
would SELECT every adventure, action, memory, and story card into Python only to
|
||||
delete them, which on Neon is the egress pattern that has already cost this
|
||||
project once. Every foreign key from `users` downward is ON DELETE CASCADE, from
|
||||
users to scenarios, adventures, scripts, and settings, and from those to actions,
|
||||
memories, and cards, so the database deletes the whole graph in one statement and
|
||||
returns a row count.
|
||||
|
||||
The scan gets no index. The sweep runs a few times a day against a table holding
|
||||
at most a few thousand rows, which does not justify a migration and the schema
|
||||
surface it adds.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from sqlalchemy import delete, func
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.concurrency import run_in_threadpool
|
||||
|
||||
from . import analytics, auth, models
|
||||
from .database import SessionLocal
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _int_env(name: str, default: int) -> int:
|
||||
try:
|
||||
return int(os.environ.get(name, "").strip() or default)
|
||||
except ValueError:
|
||||
logger.warning("%s is not an integer; using %d.", name, default)
|
||||
return default
|
||||
|
||||
|
||||
# Days of inactivity before a guest account is deleted. A value of 0 or less
|
||||
# disables the policy, for a deployment that keeps everything.
|
||||
RETENTION_DAYS = _int_env("AIDND_GUEST_RETENTION_DAYS", 5)
|
||||
|
||||
# How often a long-lived process re-checks. Hours, not minutes: nothing here is
|
||||
# time-critical, and on Render's free tier the service sleeps and cold-starts
|
||||
# often enough that the startup sweep does most of the work by itself.
|
||||
SWEEP_INTERVAL_SECONDS = _int_env("AIDND_CLEANUP_INTERVAL_HOURS", 6) * 3600
|
||||
|
||||
|
||||
def enabled() -> bool:
|
||||
"""Guests only exist in multi-user mode, so local runs skip the sweep
|
||||
rather than pointing a DELETE at a database that has nothing to collect."""
|
||||
return auth.MULTI_USER and RETENTION_DAYS > 0
|
||||
|
||||
|
||||
def anything_to_sweep() -> bool:
|
||||
"""Whether the periodic task is worth starting at all. The two jobs it runs
|
||||
are independent: a deployment can keep every guest forever and still want
|
||||
its analytics rows aged out, and vice versa."""
|
||||
return enabled() or analytics.RETENTION_DAYS > 0
|
||||
|
||||
|
||||
def delete_stale_guests(db: Session, *, now: datetime | None = None) -> int:
|
||||
"""Delete guests idle for RETENTION_DAYS or more. Returns the row count.
|
||||
|
||||
The caller owns error handling; `sweep` is the safe wrapper.
|
||||
"""
|
||||
if RETENTION_DAYS <= 0:
|
||||
return 0
|
||||
# Stored timestamps are UTC without a timezone on both backends. SQLite
|
||||
# drops the timezone, and the Postgres columns are TIMESTAMP WITHOUT TIME
|
||||
# ZONE with the session pinned to UTC in `database.py`. Match that, so the
|
||||
# comparison does not depend on how a dialect renders a value that carries a
|
||||
# timezone.
|
||||
reference = now or models.utcnow()
|
||||
cutoff = reference.replace(tzinfo=None) - timedelta(days=RETENTION_DAYS)
|
||||
|
||||
stmt = (
|
||||
delete(models.User)
|
||||
.where(
|
||||
models.User.is_guest.is_(True),
|
||||
models.User.email.is_(None),
|
||||
func.coalesce(models.User.last_seen_at, models.User.created_at) < cutoff,
|
||||
)
|
||||
# Without this option, the "auto" strategy cannot evaluate coalesce in
|
||||
# Python and falls back to fetching every matching primary key first.
|
||||
# That is a second round trip for no benefit, because this session holds
|
||||
# no User objects to synchronize.
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
removed = db.execute(stmt).rowcount or 0
|
||||
db.commit()
|
||||
return removed
|
||||
|
||||
|
||||
def sweep() -> int:
|
||||
"""One pass, with its own session. Never raises: a failed cleanup must not
|
||||
be able to take the app down (same rule as seeding). Returns the guest
|
||||
count, which is the number worth logging about."""
|
||||
if not anything_to_sweep():
|
||||
return 0
|
||||
db = SessionLocal()
|
||||
try:
|
||||
# Ages out the per-visitor analytics rows, on its own terms: it is not
|
||||
# about guests, and it must still happen on a deployment that has
|
||||
# chosen to keep every account it ever minted.
|
||||
aged = analytics.purge_old_visitor_days(db)
|
||||
if aged:
|
||||
logger.info("Aged out %d analytics visitor-day row(s).", aged)
|
||||
removed = delete_stale_guests(db) if enabled() else 0
|
||||
if removed:
|
||||
logger.info(
|
||||
"Cleaned up %d guest account(s) idle for %d+ days.",
|
||||
removed,
|
||||
RETENTION_DAYS,
|
||||
)
|
||||
return removed
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Guest cleanup failed; continuing without it.")
|
||||
return 0
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
async def _sweep_loop() -> None:
|
||||
while True:
|
||||
# Blocking DB work: keep it off the event loop, which is also serving
|
||||
# SSE turn streams.
|
||||
await run_in_threadpool(sweep)
|
||||
await asyncio.sleep(SWEEP_INTERVAL_SECONDS)
|
||||
|
||||
|
||||
def start_sweeper() -> asyncio.Task | None:
|
||||
"""Kick off the periodic sweep; None when there is nothing to sweep."""
|
||||
if not enabled():
|
||||
logger.info("Guest cleanup disabled (multi_user=%s, retention_days=%d).",
|
||||
auth.MULTI_USER, RETENTION_DAYS)
|
||||
else:
|
||||
logger.info(
|
||||
"Guest cleanup on: deleting guests idle %d+ days, every %d hour(s).",
|
||||
RETENTION_DAYS,
|
||||
SWEEP_INTERVAL_SECONDS // 3600,
|
||||
)
|
||||
if not anything_to_sweep():
|
||||
return None
|
||||
return asyncio.create_task(_sweep_loop())
|
||||
|
||||
|
||||
async def stop_sweeper(task: asyncio.Task | None) -> None:
|
||||
if task is None:
|
||||
return
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
+22
-43
@@ -4,10 +4,17 @@ from pathlib import Path
|
||||
from sqlalchemy import create_engine, event
|
||||
from sqlalchemy.orm import DeclarativeBase, sessionmaker
|
||||
|
||||
# AIDND_DB_PATH lets deployments (Docker volume, hosted disk) relocate the
|
||||
# SQLite database; default stays backend/data.db for local runs. The parent
|
||||
# directory also hosts the auto-generated secret.key (see security.py), so
|
||||
# DB_PATH stays defined even when Postgres is in use.
|
||||
# The one database. It is SQLite, on this machine, in a file.
|
||||
#
|
||||
# Upstream could also point at a server database — `AIDND_DATABASE_URL` or the
|
||||
# platform-conventional `DATABASE_URL`, normalised onto psycopg3, with
|
||||
# pre-ping for a serverless Postgres that suspends when idle. That existed for
|
||||
# a hosted deployment. M2 removed it along with the deployment: a local
|
||||
# single-user storyteller has one reader, and a network database would be one
|
||||
# more thing that has to be running, and one more place the story lives.
|
||||
#
|
||||
# `AIDND_DB_PATH` stays. It is how the Docker image puts the database on a
|
||||
# volume, and how a test points at a throwaway file.
|
||||
_env_db_path = os.environ.get("AIDND_DB_PATH")
|
||||
DB_PATH = (
|
||||
Path(_env_db_path).resolve()
|
||||
@@ -15,48 +22,20 @@ DB_PATH = (
|
||||
else Path(__file__).resolve().parent.parent / "data.db"
|
||||
)
|
||||
|
||||
# `AIDND_DATABASE_URL`, or the conventional `DATABASE_URL`, switches the app to
|
||||
# a server database. Any SQLAlchemy URL works, and hosted deploys use Postgres,
|
||||
# which Phase 9 settled on Neon for. If neither variable is set, the app uses
|
||||
# SQLite.
|
||||
DATABASE_URL = (
|
||||
os.environ.get("AIDND_DATABASE_URL", "").strip()
|
||||
or os.environ.get("DATABASE_URL", "").strip()
|
||||
DB_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
engine = create_engine(
|
||||
f"sqlite:///{DB_PATH}",
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
|
||||
def _normalize_url(url: str) -> str:
|
||||
"""Map the postgres:// / postgresql:// schemes hosts hand out to the
|
||||
psycopg3 driver installed in requirements.txt."""
|
||||
for prefix in ("postgres://", "postgresql://"):
|
||||
if url.startswith(prefix):
|
||||
return "postgresql+psycopg://" + url[len(prefix):]
|
||||
return url
|
||||
|
||||
|
||||
if DATABASE_URL:
|
||||
engine = create_engine(
|
||||
_normalize_url(DATABASE_URL),
|
||||
# Serverless Postgres (Neon) suspends idle databases; pre-ping
|
||||
# replaces silently-dead pooled connections instead of erroring.
|
||||
pool_pre_ping=True,
|
||||
# Store/read naive UTC like SQLite does, regardless of server default.
|
||||
connect_args={"options": "-c timezone=UTC"},
|
||||
)
|
||||
else:
|
||||
DB_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
engine = create_engine(
|
||||
f"sqlite:///{DB_PATH}",
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
@event.listens_for(engine, "connect")
|
||||
def _enable_sqlite_foreign_keys(dbapi_connection, _record):
|
||||
# SQLite ships with foreign keys OFF per connection; without this every
|
||||
# ondelete=CASCADE/SET NULL in models.py is silently ignored.
|
||||
cursor = dbapi_connection.cursor()
|
||||
cursor.execute("PRAGMA foreign_keys=ON")
|
||||
cursor.close()
|
||||
@event.listens_for(engine, "connect")
|
||||
def _enable_sqlite_foreign_keys(dbapi_connection, _record):
|
||||
# SQLite ships with foreign keys OFF per connection; without this every
|
||||
# ondelete=CASCADE/SET NULL in models.py is silently ignored.
|
||||
cursor = dbapi_connection.cursor()
|
||||
cursor.execute("PRAGMA foreign_keys=ON")
|
||||
cursor.close()
|
||||
|
||||
|
||||
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
"""Which inference endpoints this product is willing to talk to.
|
||||
|
||||
The Adventure Storyteller sends the player's prose, the assembled context, the
|
||||
retrieved memories, and the embedding inputs to whatever address the model
|
||||
endpoint names. That makes the endpoint the single most consequential setting
|
||||
in the application: point it somewhere else and the whole campaign goes there.
|
||||
|
||||
The v1 rule (`planning/DECISIONS/002-ollama-only-v1.md`, ADR 004) is that
|
||||
inference runs on user-controlled local infrastructure. Two deployments are
|
||||
supported and no third is:
|
||||
|
||||
* **same-host** — Ollama on loopback, the default;
|
||||
* **explicitly configured trusted LAN** — Ollama on another machine the user
|
||||
controls, named by them, reached over HTTP or over HTTPS with a certificate
|
||||
their machine trusts.
|
||||
|
||||
Everything on the public Internet is refused. Not discouraged in the UI, not
|
||||
absent from a dropdown — refused, here, on the way out, so that a hand-edited
|
||||
database row or a hostname that starts resolving somewhere new cannot quietly
|
||||
turn a local install into an exfiltration path.
|
||||
|
||||
## How the line is drawn
|
||||
|
||||
By **address**, not by name, and against an explicit allowlist of networks:
|
||||
loopback, the three RFC1918 ranges, link-local, IPv6 unique-local, and
|
||||
carrier-grade NAT — the last of which is what a mesh VPN such as Tailscale
|
||||
hands out and is as user-controlled as a LAN.
|
||||
|
||||
Every address the host resolves to must be in one of them. One address outside
|
||||
is enough to refuse the endpoint, so a name resolving to both a private and a
|
||||
public address does not squeak through.
|
||||
|
||||
The networks are spelled out rather than inferred from `ipaddress`'s own
|
||||
classifications, which do not mean what this rule needs: `is_private` is true
|
||||
of the documentation ranges and of `0.0.0.0/8`, and `is_reserved` is true of
|
||||
IPv6 loopback — so a rule written around it refuses `http://[::1]:11434/v1`,
|
||||
which is an ordinary same-host Ollama. Naming the networks keeps the policy
|
||||
readable and makes anything unnamed refused by default.
|
||||
|
||||
Checking addresses rather than hostnames is what makes the rule hard to talk
|
||||
around. A cloud provider cannot be reached by spelling its name differently,
|
||||
and `localhost.` or a DNS entry pointing at a public host is judged on where it
|
||||
actually goes.
|
||||
|
||||
## What this is not
|
||||
|
||||
It is not a general network-policy framework, and there is nothing to configure.
|
||||
There is one predicate, and it is applied in two places: when the endpoint is
|
||||
saved, so the user gets a clear error immediately, and again before every
|
||||
outbound request, because a name that resolved to `192.168.1.50` this morning
|
||||
can resolve to something else this afternoon.
|
||||
|
||||
TLS is a separate matter and is never traded against this one. See
|
||||
`tlstrust.py`: certificates are verified in full, and no endpoint — however
|
||||
private its address — may skip that.
|
||||
"""
|
||||
|
||||
import ipaddress
|
||||
import socket
|
||||
from urllib.parse import urlparse
|
||||
|
||||
#: The networks an inference endpoint may live on. Anything else is refused.
|
||||
ALLOWED_NETWORKS = tuple(
|
||||
ipaddress.ip_network(cidr)
|
||||
for cidr in (
|
||||
"127.0.0.0/8", # this machine
|
||||
"10.0.0.0/8", # RFC1918
|
||||
"172.16.0.0/12", # RFC1918
|
||||
"192.168.0.0/16", # RFC1918
|
||||
"169.254.0.0/16", # link-local
|
||||
"100.64.0.0/10", # carrier-grade NAT, which mesh VPNs use
|
||||
"::1/128", # this machine, v6
|
||||
"fc00::/7", # unique-local, v6
|
||||
"fe80::/10", # link-local, v6
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _is_local(ip) -> bool:
|
||||
return any(ip in net for net in ALLOWED_NETWORKS)
|
||||
|
||||
#: Hosts that are only ever a cloud inference service. The address rule below
|
||||
#: already refuses every one of them, because they all resolve to public
|
||||
#: addresses; this list exists solely so the error says *why* rather than
|
||||
#: leaving the user to wonder whether their DNS is broken.
|
||||
CLOUD_HOSTS = (
|
||||
"openrouter.ai",
|
||||
"api.openai.com",
|
||||
"api.anthropic.com",
|
||||
"api.groq.com",
|
||||
"api.mistral.ai",
|
||||
"api.together.xyz",
|
||||
"api.deepseek.com",
|
||||
"generativelanguage.googleapis.com",
|
||||
"api.cohere.ai",
|
||||
"api.perplexity.ai",
|
||||
)
|
||||
|
||||
_CLOUD_REASON = (
|
||||
"this build talks to Ollama on your own machine or on your own network, "
|
||||
"and has no cloud provider support"
|
||||
)
|
||||
|
||||
|
||||
def _cloud_host(host: str) -> bool:
|
||||
host = host.lower().rstrip(".")
|
||||
return any(host == h or host.endswith("." + h) for h in CLOUD_HOSTS)
|
||||
|
||||
|
||||
def rejection_reason(url: str) -> str | None:
|
||||
"""Why this endpoint may not be used, or None if it may.
|
||||
|
||||
The string is shown to the user, so it says what to do rather than what
|
||||
went wrong internally.
|
||||
"""
|
||||
parsed = urlparse((url or "").strip())
|
||||
if parsed.scheme not in ("http", "https"):
|
||||
return "the endpoint URL must start with http:// or https://"
|
||||
host = parsed.hostname
|
||||
if not host:
|
||||
return "the endpoint URL has no host"
|
||||
if _cloud_host(host):
|
||||
return f"{host} is a cloud inference service — {_CLOUD_REASON}"
|
||||
|
||||
port = parsed.port or (443 if parsed.scheme == "https" else 80)
|
||||
try:
|
||||
infos = socket.getaddrinfo(host, port, type=socket.SOCK_STREAM)
|
||||
except socket.gaierror:
|
||||
return (
|
||||
f"the host {host!r} could not be resolved — check the address, and "
|
||||
"that the machine running Ollama is reachable from here"
|
||||
)
|
||||
|
||||
for info in infos:
|
||||
try:
|
||||
ip = ipaddress.ip_address(info[4][0])
|
||||
except ValueError:
|
||||
return "the endpoint host resolved to an address that could not be read"
|
||||
if _is_local(ip):
|
||||
continue
|
||||
if ip.is_global:
|
||||
return (
|
||||
f"{host} resolves to {ip}, which is a public Internet address — "
|
||||
f"{_CLOUD_REASON}. Use Ollama on this machine "
|
||||
"(http://127.0.0.1:11434/v1) or on a machine on your own network"
|
||||
)
|
||||
return (
|
||||
f"{host} resolves to {ip}, which is not on this machine and not on "
|
||||
"your own network. Use http://127.0.0.1:11434/v1, or the address of "
|
||||
"a machine on your network"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def check(url: str) -> None:
|
||||
"""Raises `EndpointRejected` if this endpoint is outside the policy."""
|
||||
reason = rejection_reason(url)
|
||||
if reason is not None:
|
||||
raise EndpointRejected(reason)
|
||||
|
||||
|
||||
class EndpointRejected(Exception):
|
||||
"""The configured endpoint is not one this product will send a story to."""
|
||||
|
||||
|
||||
def is_loopback(url: str) -> bool:
|
||||
"""Whether this endpoint is on this machine. Used for reporting, not for
|
||||
gating: a trusted-LAN endpoint is equally allowed."""
|
||||
host = urlparse((url or "").strip()).hostname
|
||||
if not host:
|
||||
return False
|
||||
try:
|
||||
infos = socket.getaddrinfo(host, None, type=socket.SOCK_STREAM)
|
||||
except socket.gaierror:
|
||||
return False
|
||||
try:
|
||||
return all(ipaddress.ip_address(i[4][0]).is_loopback for i in infos)
|
||||
except ValueError:
|
||||
return False
|
||||
+30
-194
@@ -1,183 +1,35 @@
|
||||
"""Phase 9: abuse guards for hosted, multi-user deployments.
|
||||
"""Resource bounds on what a single request or a single story may cost.
|
||||
|
||||
Rate limits and row caps do nothing in local mode, because a single local player
|
||||
should never be throttled by their own app. The values are hardcoded on purpose.
|
||||
They are generous enough that a legitimate player never notices them, and tight
|
||||
enough that a hostile visitor cannot exhaust the demo key, saturate the CPU, or
|
||||
fill the database.
|
||||
Upstream carried three things here, and only one of them belongs in a local
|
||||
single-user product. Per-IP and per-user **rate limiting**, the login-attempt
|
||||
throttle, and the per-user **quotas** were hosted-service policy: they existed
|
||||
to stop a hostile visitor exhausting a shared demo key or filling a shared
|
||||
database. M2 removed all of it. There are no visitors, and throttling the one
|
||||
person who started the application would be a bug rather than a guard.
|
||||
|
||||
What is left is defensive programming, and it applies whatever the deployment:
|
||||
|
||||
* a ceiling on the **request body**, so a malformed or hostile payload cannot
|
||||
be read into memory before anything looks at it;
|
||||
* ceilings on how large **one adventure** may grow, in actions, memories,
|
||||
story cards and branches. These bound storage and the cost of the queries
|
||||
that walk them. They are per-story, not per-user: nothing here counts how
|
||||
many campaigns a person may have.
|
||||
|
||||
An import is checked against the same per-adventure ceilings that live creation
|
||||
uses, so a bundle cannot carry a story past a limit that play could not reach.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from collections import defaultdict, deque
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from . import auth, models
|
||||
from . import models
|
||||
|
||||
# ---------- Rate limiting ----------
|
||||
# Fixed windows per scope and caller. The windows live in memory, which is
|
||||
# enough for the single-process deployment this app targets. The worst case
|
||||
# after a restart is a brief extra allowance.
|
||||
# ---------- Per-story row caps ----------
|
||||
|
||||
# Maps a scope to (max requests, window seconds).
|
||||
RATE_LIMITS: dict[str, tuple[int, int]] = {
|
||||
"turn": (10, 60), # AI turn generation. The demo key also has a daily cap.
|
||||
"chat": (30, 60), # The AI Chat scratchpad, for power users.
|
||||
"script-test": (30, 60), # Sandboxed, but each run costs up to 2s of CPU.
|
||||
"connection-test": (10, 60), # Outbound HTTP to a user-supplied URL.
|
||||
"import": (30, 60), # Large writes.
|
||||
"auth": (10, 300), # Register and login attempts, per IP.
|
||||
"guest": (30, 300), # New guest users, per IP. Each one is a database row.
|
||||
# Pageview beacons. The limit is generous, because a real reader clicking
|
||||
# around a SPA sends a handful a minute, and it is low enough that nobody
|
||||
# can inflate the traffic numbers faster than by reloading the page.
|
||||
"analytics": (120, 60),
|
||||
}
|
||||
|
||||
_windows: dict[tuple[str, str], deque] = defaultdict(deque)
|
||||
_windows_guard = threading.Lock()
|
||||
|
||||
|
||||
# How many proxy hops sit between the app and the real client. On Render, and on
|
||||
# most platforms, that is one, because the platform's edge appends the connecting
|
||||
# IP to the right of `X-Forwarded-For`. A client can prepend any value on the
|
||||
# left, but it cannot push a value past the edge's own append, so the trustworthy
|
||||
# client IP is the entry that many places from the right rather than uvicorn's
|
||||
# leftmost choice. Trusting the leftmost entry let anyone rotate
|
||||
# `X-Forwarded-For` to get a fresh rate-limit bucket per request and bypass the
|
||||
# auth and guest limits. If the deployment adds more hops, set
|
||||
# `AIDND_TRUSTED_PROXY_HOPS`.
|
||||
TRUSTED_PROXY_HOPS = max(1, int(os.environ.get("AIDND_TRUSTED_PROXY_HOPS", "1") or 1))
|
||||
|
||||
|
||||
def client_ip(request: Request) -> str:
|
||||
"""Returns the real client IP, resisting a spoofed `X-Forwarded-For`.
|
||||
|
||||
The function reads the hop the trusted edge appended, which is the rightmost
|
||||
entry minus any extra trusted hops. If no forwarded header is present, which
|
||||
happens locally, in development, and on a direct connection, it falls back to
|
||||
the socket peer.
|
||||
|
||||
The function is public because the access log needs the same answer. Two
|
||||
functions that each decide which address belongs to the caller is how one of
|
||||
them ends up trusting a header it should not.
|
||||
"""
|
||||
forwarded = request.headers.get("x-forwarded-for")
|
||||
if forwarded:
|
||||
parts = [p.strip() for p in forwarded.split(",") if p.strip()]
|
||||
if parts:
|
||||
return parts[-min(TRUSTED_PROXY_HOPS, len(parts))]
|
||||
return request.client.host if request.client else "unknown"
|
||||
|
||||
|
||||
def rate_limit(scope: str, request: Request, user: models.User | None = None) -> None:
|
||||
"""Raises a 429 when the caller exceeds the scope's window.
|
||||
|
||||
The window is keyed per user when a user is known, because an account
|
||||
survives an IP change, and per IP otherwise.
|
||||
"""
|
||||
if not auth.MULTI_USER:
|
||||
return
|
||||
limit, window_seconds = RATE_LIMITS[scope]
|
||||
key = (scope, f"u{user.id}" if user else f"ip{client_ip(request)}")
|
||||
now = time.time()
|
||||
with _windows_guard:
|
||||
window = _windows[key]
|
||||
while window and window[0] < now - window_seconds:
|
||||
window.popleft()
|
||||
if len(window) >= limit:
|
||||
raise HTTPException(
|
||||
429, "You're doing that too fast — wait a minute and try again."
|
||||
)
|
||||
window.append(now)
|
||||
if len(_windows) > 10_000:
|
||||
_prune(now)
|
||||
|
||||
|
||||
# ---------- Per-account login throttle ----------
|
||||
# This is defense in depth next to the per-IP `auth` limit. A botnet dilutes
|
||||
# that limit, because many real source IPs each get their own bucket, so it
|
||||
# cannot by itself stop a distributed guessing run against one account. This cap
|
||||
# keys on the target email rather than on the caller, so guessing one account's
|
||||
# password stays expensive however many addresses the guesses come from.
|
||||
#
|
||||
# Only failures count, and a correct password clears the record. The window
|
||||
# slides over a short period rather than locking the account, so a user who
|
||||
# mistypes a few times recovers within minutes. The trade-off is that an
|
||||
# attacker can keep a known account throttled, which is an inconvenience and is
|
||||
# preferable to letting the account be brute-forced.
|
||||
LOGIN_FAIL_LIMIT = 8 # Failed attempts per account.
|
||||
LOGIN_FAIL_WINDOW = 900 # The window in seconds, which is 15 minutes.
|
||||
|
||||
_login_fails: dict[str, deque] = defaultdict(deque)
|
||||
_login_guard = threading.Lock()
|
||||
|
||||
|
||||
def check_login_allowed(email: str) -> None:
|
||||
"""Raises a 429 when an account has too many recent failed logins.
|
||||
|
||||
Call this before verifying the password, so that a guess never reaches the
|
||||
hash.
|
||||
"""
|
||||
if not auth.MULTI_USER:
|
||||
return
|
||||
now = time.time()
|
||||
with _login_guard:
|
||||
window = _login_fails[email]
|
||||
while window and window[0] < now - LOGIN_FAIL_WINDOW:
|
||||
window.popleft()
|
||||
if len(window) >= LOGIN_FAIL_LIMIT:
|
||||
raise HTTPException(
|
||||
429,
|
||||
"Too many failed sign-in attempts for this account — "
|
||||
"wait a few minutes and try again.",
|
||||
)
|
||||
|
||||
|
||||
def note_login_failure(email: str) -> None:
|
||||
"""Records one failed attempt against `email`."""
|
||||
if not auth.MULTI_USER:
|
||||
return
|
||||
now = time.time()
|
||||
with _login_guard:
|
||||
_login_fails[email].append(now)
|
||||
if len(_login_fails) > 10_000: # Bound the map against a flood of unique emails.
|
||||
stale = [
|
||||
key for key, window in _login_fails.items()
|
||||
if not window or window[-1] < now - LOGIN_FAIL_WINDOW
|
||||
]
|
||||
for key in stale:
|
||||
del _login_fails[key]
|
||||
|
||||
|
||||
def note_login_success(email: str) -> None:
|
||||
"""Clears the account's failure record after a correct password."""
|
||||
with _login_guard:
|
||||
_login_fails.pop(email, None)
|
||||
|
||||
|
||||
def _prune(now: float) -> None:
|
||||
"""Drops callers whose whole window has expired, so the per-IP dict stays bounded.
|
||||
|
||||
Call this with the guard held.
|
||||
"""
|
||||
longest = max(seconds for _, seconds in RATE_LIMITS.values())
|
||||
stale = [key for key, window in _windows.items()
|
||||
if not window or window[-1] < now - longest]
|
||||
for key in stale:
|
||||
del _windows[key]
|
||||
|
||||
|
||||
# ---------- Per-user row caps ----------
|
||||
|
||||
MAX_ADVENTURES_PER_USER = 100
|
||||
MAX_SCENARIOS_PER_USER = 200
|
||||
MAX_SCRIPTS_PER_USER = 200
|
||||
MAX_STORY_CARDS_PER_OWNER = 200 # Per scenario or per adventure.
|
||||
MAX_MEMORIES_PER_ADVENTURE = 1000
|
||||
MAX_ACTIONS_PER_ADVENTURE = 5000
|
||||
@@ -200,28 +52,14 @@ def check_row_cap(
|
||||
) -> None:
|
||||
"""Raises a 409 when creating one more row of `kind` would exceed its cap.
|
||||
|
||||
The caller has already checked ownership of the scenario or adventure passed
|
||||
in.
|
||||
Only per-story kinds are capped. `adventures` and `scenarios` were per-user
|
||||
quotas and are no longer checked; the callers still pass them, and they are
|
||||
accepted and ignored so that adding a cap back is a change here rather than
|
||||
at every call site.
|
||||
"""
|
||||
if not auth.MULTI_USER:
|
||||
if kind in ("adventures", "scenarios"):
|
||||
return
|
||||
if kind == "adventures":
|
||||
count = _count(db, models.Adventure, models.Adventure.user_id == user.id)
|
||||
cap, subject, hint = (
|
||||
MAX_ADVENTURES_PER_USER, "adventures",
|
||||
"delete one you no longer play to make room",
|
||||
)
|
||||
elif kind == "scenarios":
|
||||
count = _count(db, models.Scenario, models.Scenario.user_id == user.id)
|
||||
cap, subject, hint = (
|
||||
MAX_SCENARIOS_PER_USER, "scenarios", "delete one to make room"
|
||||
)
|
||||
elif kind == "scripts":
|
||||
count = _count(db, models.Script, models.Script.user_id == user.id)
|
||||
cap, subject, hint = (
|
||||
MAX_SCRIPTS_PER_USER, "scripts", "delete one to make room"
|
||||
)
|
||||
elif kind == "story_cards":
|
||||
if kind == "story_cards":
|
||||
owner_filter = (
|
||||
models.StoryCard.scenario_id == scenario_id
|
||||
if scenario_id is not None
|
||||
@@ -273,8 +111,6 @@ def check_bundle_lists(**lists) -> None:
|
||||
The keyword arguments are `story_cards`, `memories`, `actions`, and
|
||||
`branches`.
|
||||
"""
|
||||
if not auth.MULTI_USER:
|
||||
return
|
||||
for name, value in lists.items():
|
||||
cap = _BUNDLE_LIST_CAPS[name]
|
||||
if isinstance(value, list) and len(value) > cap:
|
||||
@@ -286,8 +122,8 @@ def check_bundle_lists(**lists) -> None:
|
||||
|
||||
# ---------- Request body size ----------
|
||||
# The limit is generous enough for the largest legitimate payload, which is an
|
||||
# adventure export holding thousands of actions. It applies in every mode, and no
|
||||
# honest request approaches it.
|
||||
# adventure export holding thousands of actions. No honest request approaches
|
||||
# it.
|
||||
|
||||
MAX_BODY_BYTES = 2 * 1024 * 1024
|
||||
MAX_IMPORT_BODY_BYTES = 20 * 1024 * 1024
|
||||
|
||||
+31
-68
@@ -1,6 +1,5 @@
|
||||
import mimetypes
|
||||
import os
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import FastAPI
|
||||
@@ -8,15 +7,10 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from starlette.exceptions import HTTPException as StarletteHTTPException
|
||||
|
||||
from . import analytics, cleanup
|
||||
from .auth import MULTI_USER
|
||||
from .database import engine
|
||||
from .limits import BodySizeLimitMiddleware
|
||||
from .migrations import bootstrap
|
||||
from .routers import (
|
||||
adventures, analytics as analytics_router, auth, chat, debug, scenarios, scripts,
|
||||
settings, story_cards,
|
||||
)
|
||||
from .routers import adventures, chat, debug, scenarios, settings, story_cards
|
||||
from .seed import seed_public_scenarios
|
||||
|
||||
bootstrap(engine)
|
||||
@@ -24,37 +18,32 @@ seed_public_scenarios(engine)
|
||||
|
||||
# Production serves the SPA same-origin, so CORS only matters for the Vite dev
|
||||
# server; AIDND_CORS_ORIGINS overrides for any other cross-origin setup.
|
||||
#
|
||||
# A wildcard is refused rather than honoured. This API is unauthenticated by
|
||||
# design and bound to loopback, so its only protection from a page the user
|
||||
# happens to have open in another tab is the same-origin policy. `*` would hand
|
||||
# every site on the Internet a write handle on the local campaign database. If the
|
||||
# value is wrong the app refuses to start, because a permissive CORS policy that
|
||||
# nobody notices is worse than one that fails loudly.
|
||||
CORS_ORIGINS = [
|
||||
o.strip()
|
||||
for o in os.environ.get("AIDND_CORS_ORIGINS", "").split(",")
|
||||
if o.strip()
|
||||
] or ["http://localhost:5173", "http://127.0.0.1:5173"]
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(_app: FastAPI):
|
||||
# Sweeps once on boot, then on an interval. Booting is the reliable
|
||||
# trigger on Render's free tier, where the service sleeps after ~15
|
||||
# minutes and a long-running timer rarely gets to fire.
|
||||
sweeper = cleanup.start_sweeper()
|
||||
# Visit counters are buffered in memory and written in batches; this is
|
||||
# what turns them into rows, and stop_flusher writes out the last batch so
|
||||
# a deploy doesn't drop it.
|
||||
flusher = analytics.start_flusher()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
await cleanup.stop_sweeper(sweeper)
|
||||
await analytics.stop_flusher(flusher)
|
||||
if any(o == "*" or o.strip() == "*" for o in CORS_ORIGINS):
|
||||
raise RuntimeError(
|
||||
"AIDND_CORS_ORIGINS must not contain '*'. The storyteller API is "
|
||||
"unauthenticated and loopback-bound; a wildcard origin would let any "
|
||||
"web page read and rewrite every campaign. List the exact origins "
|
||||
"instead."
|
||||
)
|
||||
|
||||
|
||||
# The interactive API docs stay local-only: in multi-user mode they just hand
|
||||
# strangers a map of the API surface.
|
||||
app = FastAPI(
|
||||
title="AI D&D",
|
||||
docs_url=None if MULTI_USER else "/docs",
|
||||
title="Adventure Storyteller",
|
||||
docs_url="/docs",
|
||||
redoc_url=None,
|
||||
openapi_url=None if MULTI_USER else "/openapi.json",
|
||||
lifespan=lifespan,
|
||||
openapi_url="/openapi.json",
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
@@ -117,50 +106,14 @@ class SecurityHeadersMiddleware:
|
||||
await self.app(scope, receive, send_with_headers)
|
||||
|
||||
|
||||
class ApiErrorMiddleware:
|
||||
"""Counts failed API responses for the analytics dashboard.
|
||||
|
||||
This is middleware rather than an exception handler, because it observes
|
||||
what the client received. A 429 from a rate limiter, a 404 from routing, and
|
||||
a 500 from a handler that never returned all reach it the same way. It is
|
||||
pure ASGI for the same reason as the headers above: an SSE turn must not be
|
||||
buffered on its way out. It watches `/api` only, because a 404 on the SPA
|
||||
mount is a page load rather than a fault.
|
||||
"""
|
||||
|
||||
def __init__(self, app):
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope, receive, send):
|
||||
if scope["type"] != "http" or not scope.get("path", "").startswith("/api"):
|
||||
return await self.app(scope, receive, send)
|
||||
|
||||
async def send_counting(message):
|
||||
if message["type"] == "http.response.start" and message["status"] >= 400:
|
||||
# The router has already put the matched route on the scope by
|
||||
# the time a response starts, so the label can name the
|
||||
# endpoint rather than the caller's path.
|
||||
analytics.record(
|
||||
analytics.M_ERROR,
|
||||
analytics.api_route_label(scope, message["status"]),
|
||||
)
|
||||
await send(message)
|
||||
|
||||
await self.app(scope, receive, send_counting)
|
||||
|
||||
|
||||
app.add_middleware(ApiErrorMiddleware)
|
||||
app.add_middleware(SecurityHeadersMiddleware)
|
||||
|
||||
app.include_router(auth.router)
|
||||
app.include_router(scenarios.router)
|
||||
app.include_router(adventures.router)
|
||||
app.include_router(story_cards.router)
|
||||
app.include_router(scripts.router)
|
||||
app.include_router(settings.router)
|
||||
app.include_router(chat.router)
|
||||
app.include_router(debug.router)
|
||||
app.include_router(analytics_router.router)
|
||||
|
||||
|
||||
@app.get("/api/health")
|
||||
@@ -171,7 +124,8 @@ def health():
|
||||
# In production, serve the built frontend (frontend/dist) as static files.
|
||||
class SPAStaticFiles(StaticFiles):
|
||||
"""Serve index.html for unknown paths so client-side routes (/play/3)
|
||||
survive a page reload. API routes are matched before this mount."""
|
||||
survive a page reload. API routes are matched before this mount, and an
|
||||
unmatched one 404s rather than falling through to the page."""
|
||||
|
||||
async def get_response(self, path, scope):
|
||||
try:
|
||||
@@ -179,11 +133,20 @@ class SPAStaticFiles(StaticFiles):
|
||||
except StarletteHTTPException as exc:
|
||||
if exc.status_code != 404:
|
||||
raise
|
||||
return await super().get_response("index.html", scope)
|
||||
return await self._fallback(path, scope)
|
||||
if response.status_code == 404:
|
||||
return await super().get_response("index.html", scope)
|
||||
return await self._fallback(path, scope)
|
||||
return response
|
||||
|
||||
async def _fallback(self, path, scope):
|
||||
# The mount is a catch-all, so an /api path no router claims — a typo,
|
||||
# or an endpoint this build removed — used to come back as the SPA's
|
||||
# HTML with status 200, and a client asking for JSON parsed a web page
|
||||
# instead of seeing that the route is not there.
|
||||
if path == "api" or path.startswith("api/"):
|
||||
raise StarletteHTTPException(status_code=404)
|
||||
return await super().get_response("index.html", scope)
|
||||
|
||||
|
||||
# Python's mimetypes table has no entry for woff2 on a slim Debian image, so
|
||||
# StaticFiles served the self-hosted fonts as application/octet-stream. Browsers
|
||||
|
||||
@@ -338,6 +338,10 @@ MIGRATIONS: list[tuple[int, str | dict[str, str]]] = [
|
||||
(74, "ALTER TABLE adventures ADD COLUMN persona_name VARCHAR(80) NOT NULL DEFAULT ''"),
|
||||
(75, "ALTER TABLE adventures ADD COLUMN persona_pronouns VARCHAR(40) NOT NULL DEFAULT ''"),
|
||||
(76, "ALTER TABLE adventures ADD COLUMN persona_desc TEXT NOT NULL DEFAULT ''"),
|
||||
# M2: how long to wait for the model. Upstream hardcoded 120s in the HTTP
|
||||
# client, which a cold model load on a CPU-only machine can exceed. The
|
||||
# default matches `providers.openai_compatible.DEFAULT_READ_TIMEOUT`.
|
||||
(77, "ALTER TABLE settings ADD COLUMN model_timeout_seconds INTEGER NOT NULL DEFAULT 300"),
|
||||
]
|
||||
|
||||
LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1)
|
||||
@@ -1106,23 +1110,3 @@ def bootstrap(engine: Engine, through: int = LATEST_VERSION) -> None:
|
||||
_backfill_parents(conn)
|
||||
current = version
|
||||
_set_version(conn, current)
|
||||
_encrypt_plaintext_api_keys(conn)
|
||||
|
||||
|
||||
def _encrypt_plaintext_api_keys(conn) -> None:
|
||||
"""Encrypts API keys saved before encryption at rest existed (Phase 8).
|
||||
|
||||
Those keys are stored in plain text, so this pass wraps them in Fernet. Plain
|
||||
SQL cannot do it. The pass runs on every start, and it matches no rows once
|
||||
every row carries the `enc:` prefix.
|
||||
"""
|
||||
from . import security # Deferred: security derives its key from DB_PATH setup.
|
||||
|
||||
rows = conn.execute(text(
|
||||
"SELECT id, api_key FROM settings WHERE api_key != '' AND api_key NOT LIKE 'enc:%'"
|
||||
)).all()
|
||||
for row_id, plain in rows:
|
||||
conn.execute(
|
||||
text("UPDATE settings SET api_key = :key WHERE id = :id"),
|
||||
{"key": security.encrypt_secret(plain), "id": row_id},
|
||||
)
|
||||
|
||||
+29
-179
@@ -1,8 +1,8 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import (
|
||||
JSON, Boolean, Column, DateTime, Float, ForeignKey, Index, Integer, LargeBinary,
|
||||
String, Table, Text, UniqueConstraint, event,
|
||||
JSON, Boolean, DateTime, Float, ForeignKey, Index, Integer, LargeBinary,
|
||||
String, Text, event,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column, relationship
|
||||
|
||||
@@ -37,19 +37,12 @@ class User(Base):
|
||||
is_guest: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow)
|
||||
last_seen_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
|
||||
# Shared demo key usage (resets when the UTC date changes).
|
||||
# Was the shared demo key's per-day tally. M2 removed the demo key; these
|
||||
# columns stay so existing databases open unchanged and are never written.
|
||||
demo_turns_used: Mapped[int] = mapped_column(Integer, default=0)
|
||||
demo_turns_date: Mapped[str] = mapped_column(String(10), default="")
|
||||
|
||||
|
||||
scenario_scripts = Table(
|
||||
"scenario_scripts",
|
||||
Base.metadata,
|
||||
Column("scenario_id", ForeignKey("scenarios.id", ondelete="CASCADE"), primary_key=True),
|
||||
Column("script_id", ForeignKey("scripts.id", ondelete="CASCADE"), primary_key=True),
|
||||
)
|
||||
|
||||
|
||||
class Scenario(Base):
|
||||
__tablename__ = "scenarios"
|
||||
|
||||
@@ -89,7 +82,6 @@ class Scenario(Base):
|
||||
back_populates="scenario", cascade="all, delete-orphan"
|
||||
)
|
||||
adventures: Mapped[list["Adventure"]] = relationship(back_populates="scenario")
|
||||
scripts: Mapped[list["Script"]] = relationship(secondary=scenario_scripts)
|
||||
|
||||
|
||||
class Adventure(Base):
|
||||
@@ -121,6 +113,9 @@ class Adventure(Base):
|
||||
persona_name: Mapped[str] = mapped_column(String(80), default="")
|
||||
persona_pronouns: Mapped[str] = mapped_column(String(40), default="")
|
||||
persona_desc: Mapped[str] = mapped_column(Text, default="")
|
||||
# Was the campaign scripting engine's shared `state` object. M2 removed
|
||||
# scripting; the column stays so existing databases open unchanged, and it
|
||||
# is never written with anything but an empty dict.
|
||||
script_state: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||
# Phase 12: live RPG world state (world/player/npc stats + milestones),
|
||||
# instantiated from the scenario's stat_schema. Empty when there's no RPG layer.
|
||||
@@ -177,11 +172,6 @@ class Adventure(Base):
|
||||
cascade="all, delete-orphan",
|
||||
order_by="Action.id",
|
||||
)
|
||||
scripts: Mapped[list["AdventureScript"]] = relationship(
|
||||
back_populates="adventure",
|
||||
cascade="all, delete-orphan",
|
||||
order_by="AdventureScript.position",
|
||||
)
|
||||
memories: Mapped[list["Memory"]] = relationship(
|
||||
back_populates="adventure",
|
||||
cascade="all, delete-orphan",
|
||||
@@ -419,9 +409,10 @@ class Action(Base):
|
||||
# and for re-attaching the emit block when replaying history to the model.
|
||||
# Mirrors the active variant, same as text/reasoning/context_snapshot.
|
||||
world_delta: Mapped[dict | None] = mapped_column(JSON, nullable=True)
|
||||
# Phase 14, SP4: the shared script state and the RPG world state as they
|
||||
# stood after this node was played. These columns record the node's outcome
|
||||
# rather than its starting position.
|
||||
# Phase 14, SP4: the RPG world state as it stood after this node was
|
||||
# played. These columns record the node's outcome rather than its starting
|
||||
# position. (`state_after` held the scripting engine's state, which M2
|
||||
# removed; it is now always written empty.)
|
||||
#
|
||||
# Two operations need this outcome, and neither can use a snapshot taken
|
||||
# before the turn. Switching between siblings must restore the state that
|
||||
@@ -508,52 +499,6 @@ class Action(Base):
|
||||
return out
|
||||
|
||||
|
||||
class Script(Base):
|
||||
__tablename__ = "scripts"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
user_id: Mapped[int | None] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), nullable=True
|
||||
)
|
||||
name: Mapped[str] = mapped_column(String(200), default="Untitled Script")
|
||||
description: Mapped[str] = mapped_column(Text, default="")
|
||||
library_js: Mapped[str] = mapped_column(Text, default="")
|
||||
input_js: Mapped[str] = mapped_column(Text, default="")
|
||||
context_js: Mapped[str] = mapped_column(Text, default="")
|
||||
output_js: Mapped[str] = mapped_column(Text, default="")
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow, onupdate=utcnow)
|
||||
|
||||
|
||||
class AdventureScript(Base):
|
||||
"""A script copied into an adventure at creation, so library edits don't
|
||||
change running adventures unless the player explicitly re-syncs it from
|
||||
`source_script_id`. `state` lives on Adventure.script_state (one shared
|
||||
state per adventure, as in AI Dungeon)."""
|
||||
|
||||
__tablename__ = "adventure_scripts"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
adventure_id: Mapped[int] = mapped_column(ForeignKey("adventures.id", ondelete="CASCADE"))
|
||||
# The library Script that this copy was made from, which lets the player
|
||||
# re-sync it on demand. The value is NULL for legacy copies that predate
|
||||
# this column, and for demo-derived copies whose source the player does not
|
||||
# own. Those copies fall back to matching by name, or cannot be synced.
|
||||
source_script_id: Mapped[int | None] = mapped_column(
|
||||
ForeignKey("scripts.id", ondelete="SET NULL"), nullable=True
|
||||
)
|
||||
position: Mapped[int] = mapped_column(Integer, default=0)
|
||||
enabled: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
name: Mapped[str] = mapped_column(String(200), default="Untitled Script")
|
||||
description: Mapped[str] = mapped_column(Text, default="")
|
||||
library_js: Mapped[str] = mapped_column(Text, default="")
|
||||
input_js: Mapped[str] = mapped_column(Text, default="")
|
||||
context_js: Mapped[str] = mapped_column(Text, default="")
|
||||
output_js: Mapped[str] = mapped_column(Text, default="")
|
||||
|
||||
adventure: Mapped[Adventure] = relationship(back_populates="scripts")
|
||||
|
||||
|
||||
class Settings(Base):
|
||||
__tablename__ = "settings"
|
||||
|
||||
@@ -564,8 +509,10 @@ class Settings(Base):
|
||||
ForeignKey("users.id", ondelete="CASCADE"), nullable=True, unique=True
|
||||
)
|
||||
endpoint_url: Mapped[str] = mapped_column(String(500), default="http://localhost:11434/v1")
|
||||
# Encrypted at rest with Fernet, which produces a value that starts with
|
||||
# "enc:". See security.py. To read the key, use `api_key_plain`.
|
||||
# Was a cloud provider's API key, encrypted at rest. Ollama does not use
|
||||
# one and M2 removed cloud providers, so nothing reads or writes this now.
|
||||
# The column stays so existing databases open unchanged; an old value is
|
||||
# left where it is rather than migrated or decrypted.
|
||||
api_key: Mapped[str] = mapped_column(String(500), default="")
|
||||
model: Mapped[str] = mapped_column(String(200), default="")
|
||||
api_mode: Mapped[str] = mapped_column(String(20), default="chat") # chat|completion
|
||||
@@ -579,6 +526,11 @@ class Settings(Base):
|
||||
# max_output_tokens so story output keeps its full budget.
|
||||
reasoning_max_tokens: Mapped[int] = mapped_column(Integer, default=0)
|
||||
context_token_budget: Mapped[int] = mapped_column(Integer, default=16384)
|
||||
# How long to wait for the model, in seconds, before giving up on a turn.
|
||||
# A cold load of a mid-sized model on a CPU-only machine can take minutes,
|
||||
# while the same turn takes seconds once the model is resident. See
|
||||
# `providers.openai_compatible.DEFAULT_READ_TIMEOUT`.
|
||||
model_timeout_seconds: Mapped[int] = mapped_column(Integer, default=300)
|
||||
narrator_prompt: Mapped[str] = mapped_column(
|
||||
Text,
|
||||
default=(
|
||||
@@ -599,120 +551,18 @@ class Settings(Base):
|
||||
memory_bank_capacity: Mapped[int] = mapped_column(Integer, default=80)
|
||||
memory_top_k: Mapped[int] = mapped_column(Integer, default=5)
|
||||
|
||||
@property
|
||||
def has_api_key(self) -> bool:
|
||||
return bool(self.api_key)
|
||||
|
||||
@property
|
||||
def api_key_plain(self) -> str:
|
||||
from . import security # local import: models is imported before security
|
||||
|
||||
return security.decrypt_secret(self.api_key)
|
||||
|
||||
|
||||
# ---------- Visit analytics (see analytics.py) ----------
|
||||
# Two intentionally simple tables. Neither can hold text that a player wrote,
|
||||
# and neither can be joined back to a `users` row, because the visitor column
|
||||
# holds an HMAC and has no foreign key. When guest cleanup deletes an account,
|
||||
# the history that account contributed remains intact and anonymous.
|
||||
|
||||
|
||||
class AnalyticsDaily(Base):
|
||||
"""One counter: how many times `label` happened within `metric` on `day`.
|
||||
|
||||
The table stores a generic triple of metric, label, and hits rather than one
|
||||
column per statistic. Measuring something new therefore costs a constant
|
||||
rather than a migration. The only writer is an UPSERT that runs from a
|
||||
buffer. See `analytics.flush`.
|
||||
"""
|
||||
|
||||
__tablename__ = "analytics_daily"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
day: Mapped[str] = mapped_column(String(10), index=True) # YYYY-MM-DD, UTC
|
||||
metric: Mapped[str] = mapped_column(String(32))
|
||||
label: Mapped[str] = mapped_column(String(80), default="")
|
||||
hits: Mapped[int] = mapped_column(Integer, default=0)
|
||||
|
||||
# The upsert target: one row per bucket per day, created or incremented.
|
||||
__table_args__ = (
|
||||
UniqueConstraint("day", "metric", "label", name="uq_analytics_daily_bucket"),
|
||||
)
|
||||
|
||||
|
||||
class AnalyticsVisitorDay(Base):
|
||||
"""One visitor, one day, and which funnel steps they reached on it.
|
||||
|
||||
This table exists so that the funnel counts people rather than clicks. A
|
||||
player who starts six adventures counts as one person who started an
|
||||
adventure. `is_new` is set when the visitor has no earlier row, which is why
|
||||
the visitor column also has an index of its own.
|
||||
"""
|
||||
|
||||
__tablename__ = "analytics_visitor_days"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
day: Mapped[str] = mapped_column(String(10))
|
||||
# HMAC of the user id under the app secret; not reversible, not a key.
|
||||
visitor: Mapped[str] = mapped_column(String(32))
|
||||
is_new: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
opened: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
created: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
played: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
signed_up: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint("day", "visitor", name="uq_analytics_visitor_day"),
|
||||
Index("ix_analytics_visitor", "visitor"),
|
||||
)
|
||||
|
||||
|
||||
class AccessEvent(Base):
|
||||
"""One sign-in, registration, failed attempt, or session first-seen.
|
||||
|
||||
This table is the counterpart to the two above, and it is kept separate from
|
||||
them on purpose. It identifies people by design, recording address, email,
|
||||
and device. Keeping it in its own table and its own module means that the
|
||||
structure enforces the anonymity of the counters rather than a convention.
|
||||
|
||||
`user_id` is a plain integer with no foreign key. An access log that
|
||||
disappeared when the account did would not serve its purpose, and guest
|
||||
cleanup deletes accounts on a schedule. `who` and `is_guest` are snapshots
|
||||
for the same reason, so a row still reads correctly after the account is
|
||||
gone.
|
||||
"""
|
||||
|
||||
__tablename__ = "access_events"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
at: Mapped[datetime] = mapped_column(DateTime, default=utcnow, index=True)
|
||||
# session | login | register | login_failed
|
||||
kind: Mapped[str] = mapped_column(String(16))
|
||||
user_id: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
# The email for a registered account, or a label such as "Guest #12"
|
||||
# otherwise. For a failed sign-in, this holds the address that was tried,
|
||||
# which is the reason the row exists.
|
||||
who: Mapped[str] = mapped_column(String(320), default="")
|
||||
is_guest: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
ip: Mapped[str] = mapped_column(String(45), default="") # 45 = max IPv6
|
||||
country: Mapped[str] = mapped_column(String(16), default="")
|
||||
device: Mapped[str] = mapped_column(String(16), default="")
|
||||
user_agent: Mapped[str] = mapped_column(String(200), default="")
|
||||
|
||||
|
||||
# Phase 14: the fallback under `tree.place_action`.
|
||||
# The hosted visitor dashboard's two counter tables and the access log that
|
||||
# recorded sign-ins, addresses and devices used to be mapped here. M2 removed
|
||||
# the hosted deployment they served.
|
||||
#
|
||||
# Since SP2, reads filter on `branch_id` and `depth`. A node written without
|
||||
# them is invisible to every page, every context build, and every memory pass.
|
||||
# The failure is silent, because nothing raises an error. Every current writer
|
||||
# places its nodes explicitly, but relying on that would also mean relying on
|
||||
# every fixture, script, and test written from now on. The session therefore
|
||||
# enforces the rule as rows travel to the database.
|
||||
#
|
||||
# This listener is registered here rather than in tree.py so that importing the
|
||||
# models is enough to enable it. The invariant belongs to the rows, not to the
|
||||
# module that usually writes them. The import sits inside the callback because
|
||||
# tree.py imports this module.
|
||||
# The tables are left in the database rather than dropped: they are inert,
|
||||
# nothing reads or writes them, and a destructive migration would risk an
|
||||
# existing campaign database for tidiness alone. They are not product
|
||||
# functionality.
|
||||
|
||||
|
||||
@event.listens_for(Session, "before_flush")
|
||||
def _place_new_nodes_on_the_tree(session, flush_context, instances):
|
||||
from . import tree
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
"""SSRF guard for the one place the server makes an outbound request to a
|
||||
user-supplied address: the BYOK `endpoint_url` (connection test + turns/chat).
|
||||
|
||||
Without this guard, a hosted user could point `endpoint_url` at an internal
|
||||
service or at the cloud metadata endpoint, 169.254.169.254, and have the server
|
||||
fetch it. The connection test even returns part of the response. The guard
|
||||
therefore refuses any URL that resolves to a non-public address.
|
||||
|
||||
The guard does nothing in local mode. A local install talking to
|
||||
http://localhost:11434, which is Ollama, is the intended case. The guard applies
|
||||
only to a hosted, multi-user deployment, where the endpoint comes from an
|
||||
untrusted visitor.
|
||||
"""
|
||||
|
||||
import ipaddress
|
||||
import socket
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from . import auth
|
||||
|
||||
|
||||
def endpoint_block_reason(url: str) -> str | None:
|
||||
"""A human-readable reason this URL must NOT be fetched server-side, or None
|
||||
if it's allowed. Resolves the host and rejects it if any resulting address
|
||||
is non-public (private, loopback, link-local/metadata, reserved, …).
|
||||
|
||||
Checking at request time (not just on save) is deliberate: it resists a DNS
|
||||
record that flips to a private IP after the value was stored.
|
||||
"""
|
||||
if not auth.MULTI_USER:
|
||||
return None
|
||||
parsed = urlparse(url)
|
||||
if parsed.scheme not in ("http", "https"):
|
||||
return "the endpoint URL must start with http:// or https://"
|
||||
host = parsed.hostname
|
||||
if not host:
|
||||
return "the endpoint URL has no host"
|
||||
port = parsed.port or (443 if parsed.scheme == "https" else 80)
|
||||
try:
|
||||
infos = socket.getaddrinfo(host, port, type=socket.SOCK_STREAM)
|
||||
except socket.gaierror:
|
||||
return "the endpoint host could not be resolved"
|
||||
for info in infos:
|
||||
try:
|
||||
ip = ipaddress.ip_address(info[4][0])
|
||||
except ValueError:
|
||||
return "the endpoint host resolved to an unrecognized address"
|
||||
# is_global is the strict allowlist: private/loopback/link-local/CGNAT
|
||||
# all report False, so this one check covers the metadata IP too.
|
||||
if not ip.is_global or ip.is_multicast or ip.is_reserved:
|
||||
return "the endpoint URL resolves to a non-public address"
|
||||
return None
|
||||
@@ -3,32 +3,27 @@ from typing import AsyncIterator
|
||||
|
||||
import httpx
|
||||
|
||||
from .. import debuglog, netguard, tlstrust
|
||||
from .. import debuglog, endpoints, tlstrust
|
||||
from .base import PromptParts, Provider, ProviderError
|
||||
|
||||
# Appended after the story text in chat mode, so a chat-tuned model continues
|
||||
# the prose rather than replying conversationally.
|
||||
CHAT_CONTINUE_HINT = "\n\n[Continue the story directly. Output only story text.]"
|
||||
|
||||
# OpenRouter serves one model from whichever upstream is available, and every
|
||||
# upstream holds its own prompt cache, so a request routed somewhere new starts
|
||||
# with a cold cache however stable the prompt is. Naming a preferred upstream
|
||||
# makes routing deterministic, which is what allows a cache hit at all.
|
||||
#
|
||||
# `allow_fallbacks` stays at its default of true on purpose, because this is a
|
||||
# preference rather than a restriction. If the named upstream is down, the
|
||||
# request still goes elsewhere and only misses the cache, which is the behavior
|
||||
# without this setting.
|
||||
#
|
||||
# This is a list rather than a value derived from the model slug. The vendor half
|
||||
# of a slug is usually the provider slug, such as "deepseek/..." mapping to
|
||||
# "deepseek", which was verified against /api/v1/providers, but not reliably.
|
||||
# Google's models are served by "google-ai-studio" and "google-vertex", and there
|
||||
# is no "google". Look a vendor up on the model's Providers tab before adding it
|
||||
# here. A slug that does not exist is a routing preference that, at best, does
|
||||
# nothing.
|
||||
_OPENROUTER_HOST = "openrouter.ai"
|
||||
_PREFERRED_UPSTREAM = {"deepseek": "deepseek"}
|
||||
# A machine that is not listening refuses in milliseconds, so a slow connect
|
||||
# means the wrong address rather than a busy model.
|
||||
CONNECT_TIMEOUT = 10.0
|
||||
|
||||
# How long to wait for generation when Settings names no value. Upstream
|
||||
# hardcoded 120s, and M1 measured a *cold* load of a 3B model on a GPU-less
|
||||
# four-core host exceeding it three times while the same turn took 6-9 seconds
|
||||
# once the model was resident. 300s covers a cold start on modest hardware and
|
||||
# is still a number: a wedged endpoint fails rather than hanging forever.
|
||||
DEFAULT_READ_TIMEOUT = 300.0
|
||||
|
||||
# Embeddings are short and never cold-load a large model.
|
||||
EMBED_READ_TIMEOUT = 60.0
|
||||
|
||||
|
||||
|
||||
# Completion endpoints have no roles, so a chat has to be flattened into one
|
||||
@@ -44,74 +39,47 @@ def flatten_messages(messages: list[dict]) -> str:
|
||||
|
||||
|
||||
class OpenAICompatibleProvider(Provider):
|
||||
"""Adapter for any /v1-style endpoint.
|
||||
"""Adapter for Ollama's OpenAI-compatible `/v1` API.
|
||||
|
||||
This covers Ollama, LM Studio, OpenAI, OpenRouter, vLLM, and Groq, among
|
||||
others.
|
||||
The protocol is OpenAI's, which is what the module is named for; the
|
||||
product speaks it to Ollama and to nothing else. `endpoints.py` decides
|
||||
which addresses may be reached, and every request re-checks — the shape of
|
||||
the wire format is not the same thing as permission to use it.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
endpoint_url: str,
|
||||
api_key: str,
|
||||
model: str,
|
||||
api_mode: str = "chat",
|
||||
reasoning_max_tokens: int = 0,
|
||||
read_timeout: float | None = None,
|
||||
):
|
||||
self.base_url = endpoint_url.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.api_mode = api_mode # Either "chat" or "completion".
|
||||
# The thinking budget for reasoning models, on top of `max_tokens`. A
|
||||
# value of 0 means the `reasoning` parameter is not sent, because an
|
||||
# endpoint that does not know the field may reject it. A negative value
|
||||
# asks the endpoint to turn reasoning off.
|
||||
self.reasoning_max_tokens = reasoning_max_tokens
|
||||
# How long to wait for the model, in seconds. Cold-loading a model on a
|
||||
# CPU-only machine can take minutes, and a fixed short timeout reports
|
||||
# that as a failure. See `DEFAULT_READ_TIMEOUT`.
|
||||
self.read_timeout = read_timeout or DEFAULT_READ_TIMEOUT
|
||||
# The token accounting from the last call, when the endpoint reported
|
||||
# any. It holds the prompt and completion counts, plus, on OpenRouter,
|
||||
# `prompt_tokens_details.cached_tokens`, which is the number of prompt
|
||||
# tokens read from cache rather than billed in full. Every request method
|
||||
# writes it, so a caller reads it after the call it made. One provider is
|
||||
# built per request.
|
||||
# any. Every request method writes it, so a caller reads it after the
|
||||
# call it made. One provider is built per request.
|
||||
self.last_usage: dict | None = None
|
||||
|
||||
def _headers(self) -> dict:
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
return headers
|
||||
# No Authorization header: Ollama does not use one, and this build has
|
||||
# no cloud provider to carry a key for.
|
||||
return {"Content-Type": "application/json"}
|
||||
|
||||
def _apply_reasoning_budget(self, body: dict) -> None:
|
||||
"""Gives reasoning models their own thinking budget, in the OpenRouter style.
|
||||
def _timeout(self, seconds: float | None = None) -> httpx.Timeout:
|
||||
"""Short to connect, patient to read.
|
||||
|
||||
The method raises `max_tokens`, so the output keeps its full budget.
|
||||
|
||||
A negative budget does the opposite. It sends `effort: "none"` to turn
|
||||
reasoning off on a model that reasons by default, such as DeepSeek V4
|
||||
Flash. That differs from `exclude: true`, which still reasons and still
|
||||
bills for it while hiding the trace. Zero still means send nothing, so an
|
||||
endpoint that rejects unknown fields, such as Ollama, keeps working.
|
||||
A machine that is not listening says so in milliseconds, so a slow
|
||||
connect is a wrong address rather than a busy model and should fail
|
||||
fast. Generation is the opposite: the first token can be minutes away
|
||||
while a model loads.
|
||||
"""
|
||||
if self.api_mode != "chat":
|
||||
return
|
||||
if self.reasoning_max_tokens < 0:
|
||||
body["reasoning"] = {"effort": "none"}
|
||||
elif self.reasoning_max_tokens > 0:
|
||||
body["reasoning"] = {"max_tokens": self.reasoning_max_tokens}
|
||||
body["max_tokens"] += self.reasoning_max_tokens
|
||||
|
||||
def _apply_provider_routing(self, body: dict) -> None:
|
||||
"""Prefers one upstream on OpenRouter, so the prompt cache stays warm.
|
||||
|
||||
The method does nothing anywhere else. `provider` is an OpenRouter
|
||||
extension, and Ollama and similar servers reject fields they do not know.
|
||||
The `reasoning` parameter above is written around the same constraint.
|
||||
"""
|
||||
if _OPENROUTER_HOST not in self.base_url:
|
||||
return
|
||||
upstream = _PREFERRED_UPSTREAM.get(self.model.split("/", 1)[0].lower())
|
||||
if upstream:
|
||||
body["provider"] = {"order": [upstream]}
|
||||
return httpx.Timeout(seconds or self.read_timeout, connect=CONNECT_TIMEOUT)
|
||||
|
||||
def _record_usage(self, payload: dict) -> None:
|
||||
"""Records the endpoint's own token accounting, if it reported any.
|
||||
@@ -147,8 +115,6 @@ class OpenAICompatibleProvider(Provider):
|
||||
"max_tokens": max_tokens,
|
||||
"stream": True,
|
||||
}
|
||||
self._apply_reasoning_budget(body)
|
||||
self._apply_provider_routing(body)
|
||||
return url, body
|
||||
|
||||
@staticmethod
|
||||
@@ -227,8 +193,6 @@ class OpenAICompatibleProvider(Provider):
|
||||
"max_tokens": max_tokens,
|
||||
"stream": True,
|
||||
}
|
||||
self._apply_reasoning_budget(body)
|
||||
self._apply_provider_routing(body)
|
||||
async for event in self._stream(url, body):
|
||||
yield event
|
||||
|
||||
@@ -238,17 +202,17 @@ class OpenAICompatibleProvider(Provider):
|
||||
The method POSTs a streaming request, yields `("text", chunk)` and
|
||||
`("reasoning", chunk)` pairs, and logs the exchange.
|
||||
"""
|
||||
# SSRF guard for hosted mode. A user-supplied `endpoint_url` must not
|
||||
# point at an internal or metadata address. This does nothing for a
|
||||
# local install.
|
||||
reason = netguard.endpoint_block_reason(url)
|
||||
# Re-checked on every request, not only when the endpoint was saved: a
|
||||
# hostname that resolved to a LAN address yesterday can resolve
|
||||
# somewhere else today, and a database row can be edited by hand.
|
||||
reason = endpoints.rejection_reason(url)
|
||||
if reason:
|
||||
raise ProviderError(f"This endpoint can't be used — {reason}.")
|
||||
log = debuglog.start_entry(url, self.model, body)
|
||||
received: list[str] = []
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(120, connect=10), verify=tlstrust.ssl_context()
|
||||
timeout=self._timeout(), verify=tlstrust.ssl_context()
|
||||
) as client:
|
||||
async with client.stream("POST", url, json=body, headers=self._headers()) as resp:
|
||||
if resp.status_code != 200:
|
||||
@@ -351,13 +315,16 @@ class OpenAICompatibleProvider(Provider):
|
||||
"max_tokens": max_tokens,
|
||||
"stream": False,
|
||||
}
|
||||
self._apply_reasoning_budget(body)
|
||||
self._apply_provider_routing(body)
|
||||
|
||||
# Same check as `_stream`: every outbound request re-tests the
|
||||
# endpoint, so no path reaches an address the policy refuses.
|
||||
reason = endpoints.rejection_reason(url)
|
||||
if reason:
|
||||
raise ProviderError(f"This endpoint can't be used — {reason}.")
|
||||
log = debuglog.start_entry(url, self.model, body)
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(120, connect=10), verify=tlstrust.ssl_context()
|
||||
timeout=self._timeout(), verify=tlstrust.ssl_context()
|
||||
) as client:
|
||||
resp = await client.post(url, json=body, headers=self._headers())
|
||||
except httpx.HTTPError as exc:
|
||||
@@ -383,10 +350,15 @@ class OpenAICompatibleProvider(Provider):
|
||||
raise ProviderError("No embedding model configured — set one in Settings.")
|
||||
url = f"{self.base_url}/embeddings"
|
||||
body = {"model": self.model, "input": texts}
|
||||
# Same check as `_stream`: every outbound request re-tests the
|
||||
# endpoint, so no path reaches an address the policy refuses.
|
||||
reason = endpoints.rejection_reason(url)
|
||||
if reason:
|
||||
raise ProviderError(f"This endpoint can't be used — {reason}.")
|
||||
log = debuglog.start_entry(url, self.model, body)
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(60, connect=10), verify=tlstrust.ssl_context()
|
||||
timeout=self._timeout(EMBED_READ_TIMEOUT), verify=tlstrust.ssl_context()
|
||||
) as client:
|
||||
resp = await client.post(url, json=body, headers=self._headers())
|
||||
except httpx.HTTPError as exc:
|
||||
|
||||
@@ -32,7 +32,6 @@ from . import ( # noqa: F401
|
||||
takes,
|
||||
branches,
|
||||
bundle_io,
|
||||
scripts,
|
||||
refresh,
|
||||
insights,
|
||||
memories,
|
||||
|
||||
@@ -7,7 +7,7 @@ only check ownership and hand the work over.
|
||||
from fastapi import Body, Depends, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ... import analytics, bundle, limits, models, schemas
|
||||
from ... import bundle, limits, models, schemas
|
||||
from ...database import get_db
|
||||
|
||||
from .deps import CurrentUser, current_adventure, router
|
||||
@@ -34,7 +34,6 @@ def import_adventure(
|
||||
user: models.User = CurrentUser,
|
||||
):
|
||||
version = bundle.check_format(payload)
|
||||
limits.rate_limit("import", request, user)
|
||||
limits.check_row_cap("adventures", db, user)
|
||||
limits.check_bundle_lists(
|
||||
story_cards=payload.get("storyCards"),
|
||||
@@ -66,5 +65,4 @@ def import_adventure(
|
||||
# This is not a funnel step. A returning player imports a bundle, so it
|
||||
# says nothing about how far a first-time visitor got. It is counted anyway,
|
||||
# because it is the clearest evidence that anyone uses the export format.
|
||||
analytics.record_event(analytics.EV_IMPORT, user)
|
||||
return adventure
|
||||
|
||||
@@ -9,7 +9,7 @@ from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy.orm.attributes import set_committed_value
|
||||
|
||||
from ... import analytics, attempts, images, limits, memorybank, models, schemas, tree, worldstate
|
||||
from ... import attempts, images, limits, memorybank, models, schemas, tree, worldstate
|
||||
from ...database import get_db
|
||||
|
||||
from .deps import CurrentUser, current_adventure, router
|
||||
@@ -195,20 +195,6 @@ def create_adventure(
|
||||
if scenario:
|
||||
for ref, spec in scenario_card_specs(scenario, values).items():
|
||||
db.add(models.StoryCard(adventure_id=adventure.id, source_ref=ref, **spec))
|
||||
for position, script in enumerate(scenario.scripts):
|
||||
db.add(
|
||||
models.AdventureScript(
|
||||
adventure_id=adventure.id,
|
||||
source_script_id=script.id,
|
||||
position=position,
|
||||
name=script.name,
|
||||
description=script.description,
|
||||
library_js=script.library_js,
|
||||
input_js=script.input_js,
|
||||
context_js=script.context_js,
|
||||
output_js=script.output_js,
|
||||
)
|
||||
)
|
||||
if scenario.prompt.strip():
|
||||
opening = models.Action(
|
||||
adventure_id=adventure.id,
|
||||
@@ -223,12 +209,6 @@ def create_adventure(
|
||||
|
||||
db.commit()
|
||||
db.refresh(adventure)
|
||||
analytics.record_event(analytics.EV_ADVENTURE, user)
|
||||
# Track which shared scenarios players pick. This is the only content this
|
||||
# module records, and it records only public scenarios. A player's own
|
||||
# scenario titles stay private.
|
||||
if scenario is not None and scenario.is_public:
|
||||
analytics.record(analytics.M_SCENARIO, scenario.title)
|
||||
return adventure
|
||||
|
||||
|
||||
@@ -262,20 +242,6 @@ def get_adventure(
|
||||
return out
|
||||
|
||||
|
||||
@router.get("/{adventure_id}/script-state")
|
||||
def get_script_state(
|
||||
adventure: models.Adventure = Depends(current_adventure),
|
||||
):
|
||||
"""Returns the scripting `state` object.
|
||||
|
||||
The object holds every variable that scripts read and write through
|
||||
`state.x`, persisted after each hook. It stays `{}` until a script sets a
|
||||
variable.
|
||||
"""
|
||||
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: models.Adventure = Depends(current_adventure),
|
||||
|
||||
@@ -7,7 +7,7 @@ returns the prompt a turn was actually generated from. Neither writes anything.
|
||||
from fastapi import Depends, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ... import auth, memorybank, models
|
||||
from ... import memorybank, models
|
||||
from ...context import build_context
|
||||
from ...database import get_db
|
||||
from ..settings import get_settings
|
||||
@@ -23,14 +23,7 @@ async def dry_run_context(
|
||||
):
|
||||
"""Returns what the app would send to the AI if the player continued now."""
|
||||
settings = get_settings(db, user)
|
||||
if auth.resolve_provider_config(settings).using_demo:
|
||||
memories = (
|
||||
{"used": [], "error": "Memory bank is unavailable on the shared demo key."}
|
||||
if adventure.memory_bank_enabled
|
||||
else None
|
||||
)
|
||||
else:
|
||||
memories = await memorybank.retrieve_memories(adventure, settings, update_stats=False)
|
||||
memories = await memorybank.retrieve_memories(adventure, settings, update_stats=False)
|
||||
_, _, report = build_context(adventure, settings, memories)
|
||||
return report
|
||||
|
||||
|
||||
@@ -1,116 +0,0 @@
|
||||
"""The per-adventure copies of library scripts.
|
||||
|
||||
An adventure snapshots a library `Script` when it starts, so editing the library
|
||||
does not change a story in progress. These endpoints report whether a snapshot
|
||||
has fallen behind its library original, and copy the original over on request.
|
||||
"""
|
||||
|
||||
from fastapi import Depends, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ... import models, schemas
|
||||
from ...database import get_db
|
||||
|
||||
from .deps import CurrentUser, current_adventure, router
|
||||
|
||||
|
||||
# Fields that are copied from a library Script into its adventure-script
|
||||
# snapshot, and compared to decide whether a copy is out of date.
|
||||
SYNC_FIELDS = ("name", "description", "library_js", "input_js", "context_js", "output_js")
|
||||
|
||||
|
||||
def resolve_library_script(
|
||||
adv_script: models.AdventureScript, db: Session, user: models.User
|
||||
) -> models.Script | None:
|
||||
"""Returns the library Script an adventure script can re-sync from.
|
||||
|
||||
The result is the script this copy was made from. For a legacy copy with no
|
||||
link, it is one of the player's own scripts with the same name. Only the
|
||||
player's own scripts are considered, so a copy derived from a demo scenario
|
||||
has nothing to sync to.
|
||||
"""
|
||||
if adv_script.source_script_id is not None:
|
||||
script = db.get(models.Script, adv_script.source_script_id)
|
||||
if script is not None and script.user_id == user.id:
|
||||
return script
|
||||
return (
|
||||
db.query(models.Script)
|
||||
.filter(models.Script.user_id == user.id, models.Script.name == adv_script.name)
|
||||
.order_by(models.Script.updated_at.desc())
|
||||
.first()
|
||||
)
|
||||
|
||||
|
||||
def _mark_out_of_date(
|
||||
adv_script: models.AdventureScript, db: Session, user: models.User
|
||||
) -> models.AdventureScript:
|
||||
"""Attaches a transient `out_of_date` flag, which `AdventureScriptOut` reads.
|
||||
|
||||
The flag is `True` or `False` when a syncable library version exists, and
|
||||
`None` when none exists.
|
||||
"""
|
||||
library = resolve_library_script(adv_script, db, user)
|
||||
adv_script.out_of_date = (
|
||||
None if library is None
|
||||
else any(getattr(adv_script, f) != getattr(library, f) for f in SYNC_FIELDS)
|
||||
)
|
||||
return adv_script
|
||||
|
||||
|
||||
@router.get("/{adventure_id}/scripts", response_model=list[schemas.AdventureScriptOut])
|
||||
def list_adventure_scripts(
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = CurrentUser,
|
||||
adventure: models.Adventure = Depends(current_adventure),
|
||||
):
|
||||
return [_mark_out_of_date(s, db, user) for s in adventure.scripts]
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{adventure_id}/scripts/{adv_script_id}/sync",
|
||||
response_model=schemas.AdventureScriptOut,
|
||||
)
|
||||
def sync_adventure_script(
|
||||
adventure_id: int,
|
||||
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.
|
||||
"""
|
||||
script = db.get(models.AdventureScript, adv_script_id)
|
||||
if script is None or script.adventure_id != adventure_id:
|
||||
raise HTTPException(404, "Script not found")
|
||||
library = resolve_library_script(script, db, user)
|
||||
if library is None:
|
||||
raise HTTPException(404, "No library script to sync from")
|
||||
for field in SYNC_FIELDS:
|
||||
setattr(script, field, getattr(library, field))
|
||||
# Store the link, so that a name-matched legacy copy syncs by id next
|
||||
# time.
|
||||
script.source_script_id = library.id
|
||||
db.commit()
|
||||
db.refresh(script)
|
||||
return _mark_out_of_date(script, db, user)
|
||||
|
||||
|
||||
@router.patch(
|
||||
"/{adventure_id}/scripts/{adv_script_id}", response_model=schemas.AdventureScriptOut
|
||||
)
|
||||
def update_adventure_script(
|
||||
adventure_id: int,
|
||||
adv_script_id: int,
|
||||
payload: schemas.AdventureScriptUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
adventure: models.Adventure = Depends(current_adventure),
|
||||
):
|
||||
script = db.get(models.AdventureScript, adv_script_id)
|
||||
if script is None or script.adventure_id != adventure_id:
|
||||
raise HTTPException(404, "Script not found")
|
||||
for field, value in payload.model_dump(exclude_unset=True).items():
|
||||
setattr(script, field, value)
|
||||
db.commit()
|
||||
return script
|
||||
@@ -15,7 +15,6 @@ from ... import attempts, limits, memorybank, models, schemas, tree
|
||||
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
|
||||
@@ -39,8 +38,6 @@ def retry_action(
|
||||
attempt is stored as a sibling at the same coordinate. No text the AI wrote
|
||||
is rewritten or deleted.
|
||||
"""
|
||||
limits.rate_limit("turn", request, user)
|
||||
turns.check_demo_cap(db, user)
|
||||
turns.acquire_turn_lock(adventure_id)
|
||||
last_ai = None
|
||||
try:
|
||||
@@ -62,9 +59,7 @@ def retry_action(
|
||||
return StreamingResponse(
|
||||
turns.with_turn_lock(
|
||||
adventure_id,
|
||||
turns.generate_turn(
|
||||
adventure, db, ScriptPipeline(adventure, db), user, retry_of=last_ai
|
||||
),
|
||||
turns.generate_turn(adventure, db, user, retry_of=last_ai),
|
||||
),
|
||||
media_type="text/event-stream",
|
||||
headers=SSE_HEADERS,
|
||||
@@ -254,9 +249,7 @@ 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.
|
||||
"""
|
||||
limits.rate_limit("turn", request, user)
|
||||
limits.check_row_cap("actions", db, user, adventure=adventure)
|
||||
turns.check_demo_cap(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")
|
||||
@@ -292,9 +285,7 @@ def add_take(
|
||||
if action.type == "ai":
|
||||
# There is no player action to write. The action this turn answers is
|
||||
# already on the path, borrowed from the line being left.
|
||||
stream = turns.generate_turn(
|
||||
adventure, db, ScriptPipeline(adventure, db), user, retry_of=retry_of
|
||||
)
|
||||
stream = turns.generate_turn(adventure, db, user, retry_of=retry_of)
|
||||
else:
|
||||
stream = turns.run_player_turn(
|
||||
adventure,
|
||||
|
||||
@@ -3,8 +3,8 @@
|
||||
Everything a test needs to intercept lives here, and other modules reach it as
|
||||
`turns.<name>` rather than importing it by value. That matters twice. The turn
|
||||
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.
|
||||
`OpenAICompatibleProvider` or `generate_turn` patches this module, which every
|
||||
caller reads through.
|
||||
"""
|
||||
import threading
|
||||
|
||||
@@ -13,12 +13,11 @@ from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ... import (
|
||||
analytics, attempts, auth, limits, memorybank, models, schemas, tree, worldstate,
|
||||
attempts, limits, memorybank, models, schemas, tree, worldstate,
|
||||
)
|
||||
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
|
||||
|
||||
@@ -115,14 +114,11 @@ def action_json(action: models.Action, db: Session | None = None) -> dict:
|
||||
async def generate_turn(
|
||||
adventure: models.Adventure,
|
||||
db: Session,
|
||||
pipeline: ScriptPipeline,
|
||||
user: models.User,
|
||||
retry_of: models.Action | None = None,
|
||||
):
|
||||
"""Streams the AI continuation as SSE, then stores the result.
|
||||
|
||||
The continuation passes through the `context` and `output` script hooks.
|
||||
|
||||
If `retry_of` is set, the result is stored as a sibling of that AI action, at
|
||||
the same turn and the same coordinate, and the discarded attempt stays where
|
||||
it was written. Before calling, the caller must roll the adventure back to
|
||||
@@ -132,15 +128,15 @@ async def generate_turn(
|
||||
"""
|
||||
saved = False
|
||||
try:
|
||||
async for event in _generate_turn(adventure, db, pipeline, user, retry_of):
|
||||
async for event in _generate_turn(adventure, db, user, retry_of):
|
||||
if event is _SAVED:
|
||||
saved = True
|
||||
continue
|
||||
yield event
|
||||
finally:
|
||||
if retry_of is not None and not saved:
|
||||
# The turn failed with a provider error, an empty reply, a script
|
||||
# stop, or a disconnected client. No sibling was written, so the
|
||||
# The turn failed with a provider error, an empty reply, or a
|
||||
# disconnected client. No sibling was written, so the
|
||||
# attempt on screen is still the live one. Restore the state it
|
||||
# produced.
|
||||
attempts.restore_state(adventure, retry_of)
|
||||
@@ -155,55 +151,26 @@ _SAVED = object()
|
||||
async def _generate_turn(
|
||||
adventure: models.Adventure,
|
||||
db: Session,
|
||||
pipeline: ScriptPipeline,
|
||||
user: models.User,
|
||||
retry_of: models.Action | None = None,
|
||||
):
|
||||
settings = get_settings(db, user)
|
||||
cfg = auth.resolve_provider_config(settings)
|
||||
# On a retry, the attempt being replaced is still the live node of its turn,
|
||||
# because it stays live until a replacement exists. Filter it out of the
|
||||
# context. Otherwise the model reads the attempt it is replacing as
|
||||
# established story and writes a sequel to it.
|
||||
replacing_id = retry_of.id if retry_of is not None else None
|
||||
if cfg.using_demo:
|
||||
# The server-funded key makes no embedding or summarization calls, so
|
||||
# memory retrieval is skipped. If the bank is on, return a note.
|
||||
memories = (
|
||||
{"used": [], "error": "Memory bank is unavailable on the shared demo key — add your own API key in Settings."}
|
||||
if adventure.memory_bank_enabled
|
||||
else None
|
||||
)
|
||||
else:
|
||||
memories = await memorybank.retrieve_memories(
|
||||
adventure, settings, update_stats=True, exclude_action_id=replacing_id
|
||||
)
|
||||
memories = await memorybank.retrieve_memories(
|
||||
adventure, settings, update_stats=True, exclude_action_id=replacing_id
|
||||
)
|
||||
system_text, story_text, snapshot = build_context(
|
||||
adventure, settings, memories, exclude_action_id=replacing_id
|
||||
)
|
||||
|
||||
# onModelContext: scripts read, and can rewrite, the whole assembled
|
||||
# context.
|
||||
combined = f"{system_text}\n\n{story_text}" if system_text else story_text
|
||||
modified, stop = pipeline.run("context", combined)
|
||||
if stop:
|
||||
yield sse({"type": "stopped", "script": pipeline.report()})
|
||||
return
|
||||
context_changed = modified != combined
|
||||
parts = (
|
||||
PromptParts(system="", story=modified)
|
||||
if context_changed
|
||||
else PromptParts(system=system_text, story=story_text)
|
||||
)
|
||||
snapshot["script"] = pipeline.report() | {
|
||||
"context_changed": context_changed,
|
||||
"context_before": combined if context_changed else None,
|
||||
"context_after": modified if context_changed else None,
|
||||
}
|
||||
parts = PromptParts(system=system_text, story=story_text)
|
||||
|
||||
provider = OpenAICompatibleProvider(
|
||||
cfg.endpoint_url, cfg.api_key, cfg.model, settings.api_mode,
|
||||
settings.reasoning_max_tokens,
|
||||
settings.endpoint_url, settings.model, settings.api_mode
|
||||
)
|
||||
chunks: list[str] = []
|
||||
reasoning_chunks: list[str] = []
|
||||
@@ -239,13 +206,6 @@ async def _generate_turn(
|
||||
yield turn_error(detail)
|
||||
return
|
||||
|
||||
# onOutput
|
||||
text, _ = pipeline.run("output", text)
|
||||
if not text.strip():
|
||||
yield turn_error("A script's output modifier returned empty text.")
|
||||
return
|
||||
snapshot["script"] = snapshot["script"] | pipeline.report()
|
||||
|
||||
# RPG world state (Phase 12): read the AI's state delta out of the reply,
|
||||
# apply it through the engine, and strip the block from the displayed text.
|
||||
#
|
||||
@@ -304,39 +264,15 @@ async def _generate_turn(
|
||||
tree.place_action(db, adventure, ai_action)
|
||||
db.add(ai_action)
|
||||
adventure.updated_at = models.utcnow()
|
||||
if cfg.using_demo:
|
||||
# Successful demo turns count against the daily cap, which the endpoint
|
||||
# checks before the turn starts. A failed provider call above returns
|
||||
# before this line.
|
||||
auth.count_demo_turn(user)
|
||||
db.commit()
|
||||
# Count the turn here, after every path on which it could still have failed,
|
||||
# so the number means "stories advanced" rather than "requests attempted".
|
||||
# The demo tally counts those same turns as spend on the server-funded key.
|
||||
analytics.record_event(analytics.EV_TURN, user)
|
||||
if cfg.using_demo:
|
||||
analytics.record(analytics.M_EVENT, analytics.EV_DEMO_TURN)
|
||||
db.refresh(ai_action)
|
||||
yield _SAVED
|
||||
yield sse({"type": "done", "action": action_json(ai_action, db), "script": pipeline.report()})
|
||||
yield sse({"type": "done", "action": action_json(ai_action, db)})
|
||||
# Phase 6: schedule summarization and embedding without waiting for them.
|
||||
# The task opens its own database session. It is skipped on the demo key,
|
||||
# because background AI calls are unmetered spend on the server-funded
|
||||
# key.
|
||||
if not cfg.using_demo:
|
||||
memorybank.schedule_post_turn(adventure)
|
||||
# The task opens its own database session.
|
||||
memorybank.schedule_post_turn(adventure)
|
||||
|
||||
|
||||
def check_demo_cap(db: Session, user: models.User) -> None:
|
||||
"""Checks the demo cap before a turn starts.
|
||||
|
||||
Checking first avoids storing a capped player's input and then leaving it
|
||||
without a reply.
|
||||
"""
|
||||
settings = get_settings(db, user)
|
||||
if auth.resolve_provider_config(settings).using_demo and auth.demo_turns_left(user) <= 0:
|
||||
raise HTTPException(429, auth.DEMO_CAP_MESSAGE)
|
||||
|
||||
|
||||
async def run_player_turn(
|
||||
adventure: models.Adventure,
|
||||
@@ -353,29 +289,20 @@ async def run_player_turn(
|
||||
formatted, and a plain edit puts that same text in the box and writes it back
|
||||
verbatim. Formatting it a second time produces `> You > You ...`.
|
||||
"""
|
||||
pipeline = ScriptPipeline(adventure, db)
|
||||
|
||||
# An empty do, say, or story action behaves as a continue.
|
||||
if payload.type != "continue" and payload.text.strip():
|
||||
# onInput reads the formatted text, as in AI Dungeon: "> You ...".
|
||||
formatted = (
|
||||
payload.text.strip() if preformatted
|
||||
else format_player_input(payload.type, payload.text)
|
||||
)
|
||||
modified, stop = pipeline.run("input", formatted)
|
||||
if not modified.strip():
|
||||
yield turn_error("A script's input modifier returned empty text.",
|
||||
script=pipeline.report())
|
||||
return
|
||||
player_action = models.Action(
|
||||
adventure_id=adventure.id,
|
||||
depth=next_depth(adventure),
|
||||
type=payload.type,
|
||||
text=modified,
|
||||
text=formatted,
|
||||
)
|
||||
# The state after the input hook has run. The node leaves this state
|
||||
# behind. The AI turn after it starts here, and a retry of that turn
|
||||
# rolls back to here.
|
||||
# The state this node leaves behind. The AI turn after it starts here,
|
||||
# and a retry of that turn rolls back to here.
|
||||
attempts.snapshot_outcome(adventure, player_action)
|
||||
tree.place_action(db, adventure, player_action)
|
||||
db.add(player_action)
|
||||
@@ -387,12 +314,8 @@ async def run_player_turn(
|
||||
# that was just saved.
|
||||
db.expire(adventure, ["actions"])
|
||||
yield sse({"type": "player", "action": action_json(player_action, db)})
|
||||
if stop:
|
||||
# If onInput returns `{ stop: true }`, skip the AI call.
|
||||
yield sse({"type": "stopped", "script": pipeline.report()})
|
||||
return
|
||||
|
||||
async for event in generate_turn(adventure, db, pipeline, user):
|
||||
async for event in generate_turn(adventure, db, user):
|
||||
yield event
|
||||
|
||||
|
||||
@@ -405,9 +328,7 @@ def create_action(
|
||||
user: models.User = CurrentUser,
|
||||
adventure: models.Adventure = Depends(current_adventure),
|
||||
):
|
||||
limits.rate_limit("turn", request, user)
|
||||
limits.check_row_cap("actions", db, user, adventure=adventure)
|
||||
check_demo_cap(db, user)
|
||||
acquire_turn_lock(adventure_id)
|
||||
try:
|
||||
_move_to_after(db, adventure, payload.after_id)
|
||||
|
||||
@@ -1,140 +0,0 @@
|
||||
"""Visit analytics: one endpoint the browser writes to, one the owner reads.
|
||||
|
||||
The split matters. `/collect` is public and accepts one fact, which page was
|
||||
viewed, because anything a stranger can POST is a number a stranger can invent.
|
||||
Everything the dashboard relies on, meaning turns, adventures, sign-ups, demo
|
||||
spend, and errors, is recorded on the server by the code that performs it, so
|
||||
those counts are as trustworthy as the app itself.
|
||||
|
||||
The two reading endpoints are owner-only and 404 for everyone else, the same
|
||||
way the AI Chat router does: a feature nobody else can use is better off not
|
||||
appearing to exist. `/summary` serves the anonymous counters (analytics.py,
|
||||
which stores nothing that points at a person) and `/access` serves the access
|
||||
log (accesslog.py, which identifies people on purpose).
|
||||
"""
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import accesslog, analytics, auth, limits, models
|
||||
from ..database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/analytics", tags=["analytics"])
|
||||
|
||||
|
||||
def owner(
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
) -> models.User:
|
||||
"""Gates the reading half. It returns 404 rather than 403. See the module
|
||||
docstring."""
|
||||
if not auth.is_owner(user):
|
||||
raise HTTPException(404, "Not found")
|
||||
return user
|
||||
|
||||
|
||||
Owner = Depends(owner)
|
||||
|
||||
|
||||
class Pageview(BaseModel):
|
||||
"""What the SPA reports on a page load or a route change.
|
||||
|
||||
`first` marks a real page load rather than a client-side navigation. The
|
||||
facts that describe a visit rather than a view, which are where it came from,
|
||||
on what kind of device, and from which country, are recorded only on a page
|
||||
load, so a visitor who clicks through five pages is still one referral.
|
||||
"""
|
||||
|
||||
path: str = Field("", max_length=300)
|
||||
referrer: str = Field("", max_length=500)
|
||||
first: bool = False
|
||||
|
||||
|
||||
@router.post("/collect", status_code=204)
|
||||
def collect(
|
||||
payload: Pageview,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Response:
|
||||
"""Record one pageview. Always 204, even when nothing was counted: the
|
||||
browser has no business knowing whether it was."""
|
||||
limits.rate_limit("analytics", request)
|
||||
# Resolved by hand rather than through get_current_user: a pageview that
|
||||
# arrives before /auth/me has minted a session should still be counted as a
|
||||
# view, not turned into a 401 the SPA has to handle.
|
||||
user = (
|
||||
auth.resolve_session_user(request, db)
|
||||
if auth.MULTI_USER
|
||||
else auth.local_user(db)
|
||||
)
|
||||
# The operator's own clicks are not traffic. This applies only in
|
||||
# multi-user mode. Locally every user is the owner, and excluding them would
|
||||
# leave the dashboard empty on the machine the app is developed on.
|
||||
if auth.MULTI_USER and user is not None and auth.is_owner(user):
|
||||
return Response(status_code=204)
|
||||
|
||||
analytics.record(analytics.M_PAGE, analytics.normalize_route(payload.path))
|
||||
analytics.record_visit(user)
|
||||
if payload.first:
|
||||
referrer = analytics.normalize_referrer(
|
||||
payload.referrer, request.url.hostname or ""
|
||||
)
|
||||
if referrer: # "" means same-origin, which is not a referral
|
||||
analytics.record(analytics.M_REFERRER, referrer)
|
||||
analytics.record(
|
||||
analytics.M_DEVICE,
|
||||
analytics.device_of(request.headers.get("user-agent", "")),
|
||||
)
|
||||
analytics.record(analytics.M_COUNTRY, analytics.country_of(request.headers))
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
@router.get("/summary")
|
||||
def summary(
|
||||
days: int = Query(30, ge=1, le=365),
|
||||
db: Session = Depends(get_db),
|
||||
_user: models.User = Owner,
|
||||
) -> dict:
|
||||
"""Returns the whole dashboard in one aggregate response.
|
||||
|
||||
The response is a few kilobytes however much traffic is behind it.
|
||||
"""
|
||||
return analytics.summary(db, days)
|
||||
|
||||
|
||||
@router.get("/access")
|
||||
def access_log(
|
||||
limit: int = Query(50, ge=1, le=200),
|
||||
before_id: int | None = Query(None),
|
||||
kind: str | None = Query(None),
|
||||
q: str | None = Query(None, max_length=120),
|
||||
db: Session = Depends(get_db),
|
||||
_user: models.User = Owner,
|
||||
) -> dict:
|
||||
"""A page of the access log, newest first.
|
||||
|
||||
Unlike `/summary`, this returns rows about people, which is what it is for.
|
||||
It is therefore behind the same owner gate, it is paged rather than returned
|
||||
in full, and the people it describes have no endpoint that reaches it.
|
||||
"""
|
||||
page = accesslog.recent(
|
||||
db, limit=limit, before_id=before_id, kind=kind, query=q
|
||||
)
|
||||
return {
|
||||
"events": [
|
||||
{
|
||||
"id": event.id,
|
||||
"at": event.at.isoformat(),
|
||||
"kind": event.kind,
|
||||
"who": event.who,
|
||||
"is_guest": event.is_guest,
|
||||
"ip": event.ip,
|
||||
"country": event.country,
|
||||
"device": event.device,
|
||||
"user_agent": event.user_agent,
|
||||
}
|
||||
for event in page["events"]
|
||||
],
|
||||
"has_more": page["has_more"],
|
||||
}
|
||||
@@ -1,158 +0,0 @@
|
||||
import re
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import (accesslog, analytics, auth, cleanup, limits, models, schemas,
|
||||
security, starter)
|
||||
from ..database import get_db
|
||||
from .settings import get_settings
|
||||
|
||||
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
||||
|
||||
EMAIL_RE = re.compile(r"^[^@\s]+@[^@\s]+\.[^@\s]+$")
|
||||
|
||||
|
||||
def _set_session_cookie(response: Response, user_id: int) -> None:
|
||||
response.set_cookie(
|
||||
auth.SESSION_COOKIE,
|
||||
security.sign_session(user_id),
|
||||
max_age=auth.COOKIE_MAX_AGE,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
secure=auth.COOKIE_SECURE,
|
||||
path="/",
|
||||
)
|
||||
|
||||
|
||||
def me_payload(user: models.User, db: Session) -> dict:
|
||||
settings = get_settings(db, user)
|
||||
cfg = auth.resolve_provider_config(settings)
|
||||
return {
|
||||
"multi_user": auth.MULTI_USER,
|
||||
"id": user.id,
|
||||
"email": user.email,
|
||||
"is_guest": user.is_guest,
|
||||
# Trusted testers: unmetered demo turns, plus the AI Chat scratchpad.
|
||||
"power_user": auth.is_power_user(user),
|
||||
# Separate allowlist: shows the visit-analytics page and its nav link.
|
||||
"analytics": auth.is_owner(user),
|
||||
# How long an idle guest is kept before cleanup deletes it (None when
|
||||
# the policy is off). Served rather than hardcoded in the UI so the
|
||||
# number a guest is shown is the number actually enforced.
|
||||
"guest_retention_days": cleanup.RETENTION_DAYS if cleanup.enabled() else None,
|
||||
"demo": {
|
||||
"enabled": auth.demo_enabled(),
|
||||
"using_demo": cfg.using_demo,
|
||||
"model": cfg.model if cfg.using_demo else None,
|
||||
"turns_per_day": auth.DEMO_TURNS_PER_DAY,
|
||||
"turns_left": auth.demo_turns_left(user) if auth.demo_enabled() else None,
|
||||
"models": auth.DEMO_MODELS if auth.demo_enabled() else [],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@router.get("/me")
|
||||
def me(request: Request, response: Response, db: Session = Depends(get_db)):
|
||||
"""Returns the current user.
|
||||
|
||||
In multi-user mode this also establishes the session. If the cookie is
|
||||
missing or invalid, the endpoint creates a guest user and sets a cookie. The
|
||||
frontend calls it on load and after any 401.
|
||||
"""
|
||||
if not auth.MULTI_USER:
|
||||
user = auth.local_user(db)
|
||||
else:
|
||||
user = auth.resolve_session_user(request, db)
|
||||
if user is None:
|
||||
# Each new guest is a database row, so cap how fast one IP can
|
||||
# create them.
|
||||
limits.rate_limit("guest", request)
|
||||
user = models.User(is_guest=True)
|
||||
db.add(user)
|
||||
db.commit()
|
||||
# The guest is committed first, so a failure while copying the
|
||||
# starter adventure still leaves them with an account.
|
||||
starter.give(db, user)
|
||||
db.commit()
|
||||
_set_session_cookie(response, user.id)
|
||||
# This endpoint is the SPA's bootstrap call, so it is where a session first
|
||||
# shows itself; accesslog thins the rows down to one per day per address.
|
||||
accesslog.note_session(db, user, request)
|
||||
return me_payload(user, db)
|
||||
|
||||
|
||||
@router.post("/register")
|
||||
def register(
|
||||
payload: schemas.AuthCredentials,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""Upgrades the current guest in place.
|
||||
|
||||
The `user_id` does not change, so every adventure, scenario, script, and
|
||||
setting they created as a guest is kept.
|
||||
"""
|
||||
if not auth.MULTI_USER:
|
||||
raise HTTPException(400, "Accounts are disabled in local mode.")
|
||||
limits.rate_limit("auth", request)
|
||||
email = payload.email.strip().lower()
|
||||
if not EMAIL_RE.match(email):
|
||||
raise HTTPException(422, "Enter a valid email address.")
|
||||
if len(payload.password) < 8:
|
||||
raise HTTPException(422, "Password must be at least 8 characters.")
|
||||
if not user.is_guest:
|
||||
raise HTTPException(400, "This session is already registered.")
|
||||
if db.query(models.User).filter(models.User.email == email).first():
|
||||
raise HTTPException(409, "An account with this email already exists — log in instead.")
|
||||
user.email = email
|
||||
user.password_hash = security.hash_password(payload.password)
|
||||
user.is_guest = False
|
||||
db.commit()
|
||||
analytics.record_event(analytics.EV_SIGNUP, user)
|
||||
accesslog.record(db, accesslog.REGISTER, request, user=user)
|
||||
return me_payload(user, db)
|
||||
|
||||
|
||||
@router.post("/login")
|
||||
def login(
|
||||
payload: schemas.AuthCredentials,
|
||||
request: Request,
|
||||
response: Response,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""Point this browser's session at an existing account. Any current guest
|
||||
session is simply abandoned (its data stays under the guest user)."""
|
||||
if not auth.MULTI_USER:
|
||||
raise HTTPException(400, "Accounts are disabled in local mode.")
|
||||
limits.rate_limit("auth", request)
|
||||
email = payload.email.strip().lower()
|
||||
# Per-account throttle: stops distributed guessing against one email even
|
||||
# when the per-IP limit above is diluted across many source addresses.
|
||||
limits.check_login_allowed(email)
|
||||
user = db.query(models.User).filter(models.User.email == email).first()
|
||||
if (
|
||||
user is None
|
||||
or not user.password_hash
|
||||
or not security.verify_password(payload.password, user.password_hash)
|
||||
):
|
||||
limits.note_login_failure(email)
|
||||
# Logged with the address that was tried, not the account that owns it:
|
||||
# a guessing run against an address that has no account is exactly the
|
||||
# thing worth being able to see.
|
||||
accesslog.record(db, accesslog.LOGIN_FAILED, request, who=email)
|
||||
raise HTTPException(401, "Incorrect email or password.")
|
||||
limits.note_login_success(email)
|
||||
_set_session_cookie(response, user.id)
|
||||
analytics.record_event(analytics.EV_LOGIN, user)
|
||||
accesslog.record(db, accesslog.LOGIN, request, user=user)
|
||||
return me_payload(user, db)
|
||||
|
||||
|
||||
@router.post("/logout")
|
||||
def logout(response: Response):
|
||||
if not auth.MULTI_USER:
|
||||
raise HTTPException(400, "Accounts are disabled in local mode.")
|
||||
response.delete_cookie(auth.SESSION_COOKIE, path="/")
|
||||
return {"ok": True}
|
||||
+36
-79
@@ -1,22 +1,23 @@
|
||||
"""AI Chat: a plain scratchpad for talking to a model directly.
|
||||
"""AI Chat: a plain scratchpad for talking to the configured model directly.
|
||||
|
||||
Power users reach it, which means the `AIDND_POWER_USERS` email allowlist. It is
|
||||
deliberately thin. It adds no story context, no scripts, and no world state, and
|
||||
it persists nothing. The conversation lives in the browser and is posted in full
|
||||
on each turn. It exists for testing models, prompts, and endpoints without
|
||||
starting an adventure.
|
||||
Deliberately thin. It adds no story context and no world state, and it persists
|
||||
nothing. The conversation lives in the browser and is posted in full on each
|
||||
turn. It exists for checking a model, a prompt, or an endpoint without starting
|
||||
an adventure — which is exactly the kind of thing a local single-user install
|
||||
wants a page for.
|
||||
|
||||
Model choice is free when the user brought their own API key. On the shared demo
|
||||
key the model stays pinned to the `AIDND_DEMO_MODELS` allowlist, exactly as it is
|
||||
for turns. The server funds that key, so this page must not let it reach paid
|
||||
models.
|
||||
Upstream gated this behind a "power user" email allowlist and pinned the model
|
||||
when a shared demo key was in play. M2 removed both: there is one local user,
|
||||
who owns the endpoint, and there is no server-funded key to protect. The model
|
||||
this page talks to is the one in Settings, or one the user names per request —
|
||||
either way it is their own Ollama.
|
||||
"""
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import auth, limits, models, schemas
|
||||
from .. import auth, models, schemas
|
||||
from ..database import get_db
|
||||
from ..providers import OpenAICompatibleProvider, ProviderError
|
||||
from ..sse import SSE_HEADERS, sse
|
||||
@@ -25,80 +26,40 @@ from .settings import get_settings, list_endpoint_models
|
||||
router = APIRouter(prefix="/api/chat", tags=["chat"])
|
||||
|
||||
|
||||
def power_user(
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
) -> models.User:
|
||||
"""Gate for the whole router. 404 rather than 403 so the feature simply
|
||||
doesn't appear to exist for everyone else."""
|
||||
if not auth.is_power_user(user):
|
||||
raise HTTPException(404, "Not found")
|
||||
return user
|
||||
|
||||
|
||||
PowerUser = Depends(power_user)
|
||||
|
||||
|
||||
def _resolve_model(
|
||||
settings: models.Settings, requested: str | None
|
||||
) -> tuple[auth.ProviderConfig, str | None]:
|
||||
"""Returns the provider config for this chat, plus a note when the requested
|
||||
model was not used.
|
||||
|
||||
The pinning rule lives in `resolve_provider_config`. This function only
|
||||
reports the substitution that call made, so one place decides what the demo
|
||||
key may talk to.
|
||||
"""
|
||||
cfg = auth.resolve_provider_config(settings, model_override=requested)
|
||||
wanted = (requested or "").strip()
|
||||
if wanted and wanted != cfg.model:
|
||||
return cfg, (
|
||||
f"'{wanted}' isn't available on the shared demo key — using "
|
||||
f"{cfg.model}. Add your own API key in Settings to use any model."
|
||||
)
|
||||
return cfg, None
|
||||
|
||||
|
||||
@router.get("/config")
|
||||
async def chat_config(
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = PowerUser,
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""Returns what this page can talk to.
|
||||
|
||||
The response holds the resolved endpoint and model, whether model choice is
|
||||
pinned to the demo allowlist, and the endpoint's model listing. The listing
|
||||
is best effort, and an unreachable endpoint returns an empty list.
|
||||
The model listing is best effort: an unreachable endpoint returns an empty
|
||||
list and the reason, rather than failing the page.
|
||||
"""
|
||||
settings = get_settings(db, user)
|
||||
cfg = auth.resolve_provider_config(settings)
|
||||
listing = await list_endpoint_models(cfg)
|
||||
listing = await list_endpoint_models(settings.endpoint_url)
|
||||
return {
|
||||
"endpoint_url": cfg.endpoint_url,
|
||||
"model": cfg.model,
|
||||
"using_demo": cfg.using_demo,
|
||||
"endpoint_url": settings.endpoint_url,
|
||||
"model": settings.model,
|
||||
"api_mode": settings.api_mode,
|
||||
"temperature": settings.temperature,
|
||||
"max_tokens": settings.max_output_tokens,
|
||||
# On the demo key the whitelist IS the list of choices; otherwise it's
|
||||
# whatever the endpoint advertises (suggestions, not a restriction).
|
||||
"models": auth.DEMO_MODELS if cfg.using_demo else listing.get("models", []),
|
||||
# Suggestions from the endpoint, not a restriction.
|
||||
"models": listing.get("models", []),
|
||||
"models_error": None if listing.get("ok") else listing.get("detail"),
|
||||
}
|
||||
|
||||
|
||||
async def run_chat(cfg: auth.ProviderConfig, settings: models.Settings, payload: schemas.ChatRequest,
|
||||
note: str | None, db: Session, user: models.User):
|
||||
async def run_chat(
|
||||
settings: models.Settings, model: str, payload: schemas.ChatRequest
|
||||
):
|
||||
"""Streams the reply as SSE, using the turn stream's event shape.
|
||||
|
||||
The generator emits `reasoning` and `chunk` events while generating and then
|
||||
a `done` event, so the frontend reuses the same code.
|
||||
"""
|
||||
if note:
|
||||
yield sse({"type": "note", "detail": note})
|
||||
provider = OpenAICompatibleProvider(
|
||||
cfg.endpoint_url, cfg.api_key, cfg.model, settings.api_mode,
|
||||
settings.reasoning_max_tokens,
|
||||
settings.endpoint_url, model, settings.api_mode
|
||||
)
|
||||
messages = [m.model_dump() for m in payload.messages]
|
||||
chunks: list[str] = []
|
||||
@@ -106,7 +67,11 @@ async def run_chat(cfg: auth.ProviderConfig, settings: models.Settings, payload:
|
||||
try:
|
||||
async for kind, chunk in provider.chat(
|
||||
messages,
|
||||
temperature=payload.temperature if payload.temperature is not None else settings.temperature,
|
||||
temperature=(
|
||||
payload.temperature
|
||||
if payload.temperature is not None
|
||||
else settings.temperature
|
||||
),
|
||||
max_tokens=payload.max_tokens or settings.max_output_tokens,
|
||||
):
|
||||
if kind == "reasoning":
|
||||
@@ -123,33 +88,26 @@ async def run_chat(cfg: auth.ProviderConfig, settings: models.Settings, payload:
|
||||
if not text:
|
||||
detail = (
|
||||
"The model used its entire token budget on reasoning and returned no "
|
||||
"reply — raise max tokens, cap the reasoning budget in Settings, or "
|
||||
"use a non-reasoning model."
|
||||
"reply — raise max tokens or use a non-reasoning model."
|
||||
if reasoning_chunks
|
||||
else "The AI returned an empty response."
|
||||
)
|
||||
yield sse({"type": "error", "detail": detail})
|
||||
return
|
||||
|
||||
if cfg.using_demo:
|
||||
# Unmetered for power users (count_demo_turn is a no-op for them), but
|
||||
# keep the call so the accounting stays right if the gate ever widens.
|
||||
auth.count_demo_turn(user)
|
||||
db.commit()
|
||||
yield sse({
|
||||
"type": "done",
|
||||
"text": text,
|
||||
"reasoning": "".join(reasoning_chunks).strip() or None,
|
||||
"model": cfg.model,
|
||||
"model": model,
|
||||
})
|
||||
|
||||
|
||||
@router.post("/stream")
|
||||
def chat_stream(
|
||||
payload: schemas.ChatRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = PowerUser,
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
total = sum(len(m.content) for m in payload.messages)
|
||||
if total > schemas.CHAT_TOTAL_MAX:
|
||||
@@ -157,13 +115,12 @@ def chat_stream(
|
||||
413, f"This conversation is too long to send ({total:,} characters) — "
|
||||
"clear it or start a new one."
|
||||
)
|
||||
limits.rate_limit("chat", request, user)
|
||||
settings = get_settings(db, user)
|
||||
cfg, note = _resolve_model(settings, payload.model)
|
||||
if not cfg.model:
|
||||
model = (payload.model or "").strip() or settings.model
|
||||
if not model:
|
||||
raise HTTPException(400, "No model configured — set one in Settings or pick one here.")
|
||||
return StreamingResponse(
|
||||
run_chat(cfg, settings, payload, note, db, user),
|
||||
run_chat(settings, model, payload),
|
||||
media_type="text/event-stream",
|
||||
headers=SSE_HEADERS,
|
||||
)
|
||||
|
||||
@@ -1,19 +1,16 @@
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi import APIRouter
|
||||
|
||||
from .. import auth, debuglog
|
||||
from .. import debuglog
|
||||
|
||||
router = APIRouter(prefix="/api/debug", tags=["debug"])
|
||||
|
||||
|
||||
@router.get("/requests")
|
||||
def recent_requests():
|
||||
"""Most-recent-first log of provider requests/responses (no API keys).
|
||||
"""Most-recent-first log of provider requests and responses.
|
||||
|
||||
The log is a single process-wide ring buffer with no per-user attribution,
|
||||
so in multi-user mode, which is how a hosted deployment runs, it would expose
|
||||
other players' prompts. It is disabled there and available on a local
|
||||
install.
|
||||
A single process-wide ring buffer. It holds the prompts this install sent
|
||||
to its own Ollama, which is exactly what the person running it needs to
|
||||
diagnose a turn, and there is nobody else it could expose them to.
|
||||
"""
|
||||
if auth.MULTI_USER:
|
||||
raise HTTPException(403, "The debug log is only available on local installs.")
|
||||
return debuglog.recent()
|
||||
|
||||
@@ -3,7 +3,7 @@ from fastapi.responses import Response
|
||||
from sqlalchemy import or_
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import analytics, auth, images, limits, models, schemas
|
||||
from .. import auth, images, limits, models, schemas
|
||||
from ..database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/scenarios", tags=["scenarios"])
|
||||
@@ -54,12 +54,6 @@ def get_scenario(
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
scenario = get_scenario_or_404(scenario_id, db, user)
|
||||
# A funnel step, recorded for shared scenarios only. Opening one is the
|
||||
# first sign that a visitor is interested, and someone editing their own
|
||||
# scenario is already past this point. Their titles are theirs rather than a
|
||||
# statistic.
|
||||
if scenario.is_public:
|
||||
analytics.record_event(analytics.EV_SCENARIO_OPEN, user)
|
||||
return scenario
|
||||
|
||||
|
||||
@@ -96,18 +90,8 @@ def update_scenario(
|
||||
):
|
||||
scenario = get_scenario_or_404(scenario_id, db, user, edit=True)
|
||||
data = payload.model_dump(exclude_unset=True)
|
||||
script_ids = data.pop("script_ids", None)
|
||||
for field, value in data.items():
|
||||
setattr(scenario, field, value)
|
||||
if script_ids is not None:
|
||||
scripts = (
|
||||
db.query(models.Script)
|
||||
.filter(models.Script.id.in_(script_ids), models.Script.user_id == user.id)
|
||||
.all()
|
||||
)
|
||||
if len(scripts) != len(set(script_ids)):
|
||||
raise HTTPException(404, "One or more scripts not found")
|
||||
scenario.scripts = sorted(scripts, key=lambda s: script_ids.index(s.id))
|
||||
db.commit()
|
||||
return scenario
|
||||
|
||||
@@ -148,13 +132,6 @@ def export_scenario(
|
||||
{"type": c.type, "name": c.name, "keys": c.keys, "entry": c.entry, "notes": c.notes}
|
||||
for c in s.story_cards
|
||||
],
|
||||
"scripts": [
|
||||
{
|
||||
"name": sc.name, "description": sc.description, "library": sc.library_js,
|
||||
"input": sc.input_js, "context": sc.context_js, "output": sc.output_js,
|
||||
}
|
||||
for sc in s.scripts
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@@ -185,7 +162,6 @@ def import_scenario(
|
||||
):
|
||||
"""Accepts our export format and AI Dungeon scenario exports best-effort;
|
||||
reports any keys it didn't understand."""
|
||||
limits.rate_limit("import", request, user)
|
||||
limits.check_row_cap("scenarios", db, user)
|
||||
fields: dict = {}
|
||||
unmapped: list[str] = []
|
||||
@@ -250,21 +226,6 @@ def import_scenario(
|
||||
)
|
||||
)
|
||||
|
||||
for item in bundle.get("scripts") or []:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
script = models.Script(
|
||||
user_id=user.id,
|
||||
name=str(item.get("name") or "Imported Script")[:schemas.NAME_MAX],
|
||||
description=str(item.get("description") or ""),
|
||||
library_js=str(item.get("library") or item.get("sharedLibrary") or ""),
|
||||
input_js=str(item.get("input") or item.get("onInput") or ""),
|
||||
context_js=str(item.get("context") or item.get("onModelContext") or ""),
|
||||
output_js=str(item.get("output") or item.get("onOutput") or ""),
|
||||
)
|
||||
db.add(script)
|
||||
db.flush()
|
||||
scenario.scripts.append(script)
|
||||
|
||||
db.commit()
|
||||
out = schemas.ScenarioOut.model_validate(scenario).model_dump(mode="json")
|
||||
|
||||
@@ -1,159 +0,0 @@
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import auth, limits, models, schemas
|
||||
from ..database import get_db
|
||||
from ..scripting import run_hook
|
||||
|
||||
router = APIRouter(prefix="/api/scripts", tags=["scripts"])
|
||||
|
||||
HOOK_FIELDS = {"input": "input_js", "context": "context_js", "output": "output_js"}
|
||||
|
||||
|
||||
def get_script_or_404(script_id: int, db: Session, user: models.User) -> models.Script:
|
||||
script = db.get(models.Script, script_id)
|
||||
if script is None or script.user_id != user.id:
|
||||
raise HTTPException(404, "Script not found")
|
||||
return script
|
||||
|
||||
|
||||
@router.get("", response_model=list[schemas.ScriptOut])
|
||||
def list_scripts(
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
return (
|
||||
db.query(models.Script)
|
||||
.filter(models.Script.user_id == user.id)
|
||||
.order_by(models.Script.updated_at.desc())
|
||||
.all()
|
||||
)
|
||||
|
||||
|
||||
@router.post("", response_model=schemas.ScriptOut, status_code=201)
|
||||
def create_script(
|
||||
payload: schemas.ScriptCreate,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
limits.check_row_cap("scripts", db, user)
|
||||
script = models.Script(**payload.model_dump(), user_id=user.id)
|
||||
db.add(script)
|
||||
db.commit()
|
||||
return script
|
||||
|
||||
|
||||
@router.get("/{script_id}", response_model=schemas.ScriptOut)
|
||||
def get_script(
|
||||
script_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
return get_script_or_404(script_id, db, user)
|
||||
|
||||
|
||||
@router.patch("/{script_id}", response_model=schemas.ScriptOut)
|
||||
def update_script(
|
||||
script_id: int,
|
||||
payload: schemas.ScriptUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
script = get_script_or_404(script_id, db, user)
|
||||
for field, value in payload.model_dump(exclude_unset=True).items():
|
||||
setattr(script, field, value)
|
||||
db.commit()
|
||||
return script
|
||||
|
||||
|
||||
@router.delete("/{script_id}", status_code=204)
|
||||
def delete_script(
|
||||
script_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
db.delete(get_script_or_404(script_id, db, user))
|
||||
db.commit()
|
||||
|
||||
|
||||
@router.post("/{script_id}/test")
|
||||
def test_script(
|
||||
script_id: int,
|
||||
payload: schemas.ScriptTestRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""Runs one hook against sample text, making no AI call and storing nothing."""
|
||||
script = get_script_or_404(script_id, db, user)
|
||||
limits.rate_limit("script-test", request, user)
|
||||
result = run_hook(
|
||||
script.library_js,
|
||||
getattr(script, HOOK_FIELDS[payload.hook]),
|
||||
payload.text,
|
||||
payload.state,
|
||||
history=[],
|
||||
story_cards=[],
|
||||
info={"actionCount": 0, "characterNames": [], "memoryLength": 0, "maxChars": 0},
|
||||
)
|
||||
return {
|
||||
"text": result.text,
|
||||
"stop": result.stop,
|
||||
"state": result.state,
|
||||
"storyCards": result.story_cards,
|
||||
"logs": result.logs,
|
||||
"error": result.error,
|
||||
}
|
||||
|
||||
|
||||
# ---------- Import / Export ----------
|
||||
|
||||
@router.get("/{script_id}/export")
|
||||
def export_script(
|
||||
script_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""JSON bundle matching how AI Dungeon scripts circulate."""
|
||||
script = get_script_or_404(script_id, db, user)
|
||||
return {
|
||||
"name": script.name,
|
||||
"description": script.description,
|
||||
"library": script.library_js,
|
||||
"input": script.input_js,
|
||||
"context": script.context_js,
|
||||
"output": script.output_js,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/import", response_model=schemas.ScriptOut, status_code=201)
|
||||
def import_script(
|
||||
request: Request,
|
||||
bundle: dict = Body(...),
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""Accepts our export bundle; tolerates *_js key names too."""
|
||||
limits.rate_limit("import", request, user)
|
||||
limits.check_row_cap("scripts", db, user)
|
||||
def pick(*keys: str) -> str:
|
||||
for key in keys:
|
||||
value = bundle.get(key)
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return ""
|
||||
|
||||
script = models.Script(
|
||||
user_id=user.id,
|
||||
# A raw-dict import bypasses the schemas, so truncate to the VARCHAR
|
||||
# width.
|
||||
name=(pick("name") or "Imported Script")[:schemas.NAME_MAX],
|
||||
description=pick("description"),
|
||||
library_js=pick("library", "library_js", "sharedLibrary"),
|
||||
input_js=pick("input", "input_js", "onInput"),
|
||||
context_js=pick("context", "context_js", "onModelContext"),
|
||||
output_js=pick("output", "output_js", "onOutput"),
|
||||
)
|
||||
db.add(script)
|
||||
db.commit()
|
||||
return script
|
||||
@@ -1,20 +1,33 @@
|
||||
"""The model settings, and the connection test that tells you why they don't work.
|
||||
|
||||
There is one settings row, belonging to the one local user. It describes an
|
||||
Ollama: where it is, which model to narrate with, which to embed with, and how
|
||||
long to wait for it.
|
||||
|
||||
Upstream let this row name any OpenAI-compatible endpoint and carry an
|
||||
encrypted API key for it. M2 narrowed both: `endpoints.py` decides which
|
||||
addresses may be named, and there is no key field, because Ollama does not use
|
||||
one and this build has no cloud provider to carry a key for.
|
||||
"""
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.concurrency import run_in_threadpool
|
||||
|
||||
from .. import auth, limits, models, netguard, schemas, security, tlstrust
|
||||
from .. import auth, endpoints, models, schemas, tlstrust
|
||||
from ..database import get_db
|
||||
from ..providers.openai_compatible import CONNECT_TIMEOUT
|
||||
|
||||
router = APIRouter(prefix="/api/settings", tags=["settings"])
|
||||
|
||||
#: The connection test is a listing, not a generation, so it never waits on a
|
||||
#: model load and does not need the turn engine's patience.
|
||||
TEST_TIMEOUT = 15.0
|
||||
|
||||
|
||||
def get_settings(db: Session, user: models.User) -> models.Settings:
|
||||
"""Returns the user's settings row, creating it on first access.
|
||||
|
||||
Phase 8 made settings per user rather than global. They cover the endpoint,
|
||||
the key, the models, and the memory configuration.
|
||||
"""
|
||||
"""Returns the settings row, creating it on first access."""
|
||||
settings = (
|
||||
db.query(models.Settings).filter(models.Settings.user_id == user.id).first()
|
||||
)
|
||||
@@ -34,16 +47,24 @@ def read_settings(
|
||||
|
||||
|
||||
@router.put("", response_model=schemas.SettingsOut)
|
||||
def update_settings(
|
||||
async def update_settings(
|
||||
payload: schemas.SettingsUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
settings = get_settings(db, user)
|
||||
fields = payload.model_dump(exclude_unset=True)
|
||||
# Write-only API key: absent = unchanged, "" = cleared, else encrypted.
|
||||
if "api_key" in fields:
|
||||
fields["api_key"] = security.encrypt_secret(fields["api_key"].strip())
|
||||
|
||||
if "endpoint_url" in fields:
|
||||
# Refused here so the user finds out while they are looking at the
|
||||
# field, rather than on their next turn. The provider re-checks before
|
||||
# every request regardless; this is the friendly half of the same rule.
|
||||
reason = await run_in_threadpool(
|
||||
endpoints.rejection_reason, fields["endpoint_url"]
|
||||
)
|
||||
if reason is not None:
|
||||
raise HTTPException(400, f"That endpoint can't be used — {reason}.")
|
||||
|
||||
embedding_model_changed = (
|
||||
"embedding_model" in fields
|
||||
and fields["embedding_model"] != settings.embedding_model
|
||||
@@ -53,8 +74,6 @@ def update_settings(
|
||||
if embedding_model_changed:
|
||||
# Vectors from the old model have a different dimensionality/space;
|
||||
# clear them so the post-turn task re-embeds with the new model.
|
||||
# This covers only this user's adventures, because settings are per
|
||||
# user now.
|
||||
#
|
||||
# Both columns, and the flag. This is the one place that clears vectors
|
||||
# in bulk rather than through memorybank.set_vector, and when the
|
||||
@@ -80,28 +99,61 @@ def update_settings(
|
||||
return settings
|
||||
|
||||
|
||||
async def list_endpoint_models(cfg: auth.ProviderConfig) -> dict:
|
||||
"""Fetches the endpoint's /models listing.
|
||||
async def list_endpoint_models(endpoint_url: str) -> dict:
|
||||
"""Fetches the endpoint's `/models` listing, and doubles as the connection test.
|
||||
|
||||
The call also serves as a connectivity check, so a failure returns
|
||||
`{"ok": False, "detail": ...}` rather than raising.
|
||||
Returns `{"ok": False, "detail": ...}` rather than raising, because every
|
||||
caller wants to show the reason rather than fail the page.
|
||||
|
||||
The failure cases are told apart on purpose. "Ollama isn't running", "that
|
||||
address isn't allowed", "the certificate doesn't verify" and "it answered,
|
||||
but with an error" need four different things done about them, and a single
|
||||
"connection failed" leaves the user guessing which they have.
|
||||
"""
|
||||
# SSRF guard. Never probe a non-public address the user supplied.
|
||||
reason = await run_in_threadpool(netguard.endpoint_block_reason, cfg.endpoint_url)
|
||||
if reason:
|
||||
return {"ok": False, "detail": f"Can't reach that endpoint — {reason}."}
|
||||
url = cfg.endpoint_url.rstrip("/") + "/models"
|
||||
headers = {}
|
||||
if cfg.api_key:
|
||||
headers["Authorization"] = f"Bearer {cfg.api_key}"
|
||||
reason = await run_in_threadpool(endpoints.rejection_reason, endpoint_url)
|
||||
if reason is not None:
|
||||
return {
|
||||
"ok": False, "kind": "rejected",
|
||||
"detail": f"That endpoint can't be used — {reason}.",
|
||||
}
|
||||
|
||||
url = endpoint_url.rstrip("/") + "/models"
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10, verify=tlstrust.ssl_context()) as client:
|
||||
resp = await client.get(url, headers=headers)
|
||||
async with httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(TEST_TIMEOUT, connect=CONNECT_TIMEOUT),
|
||||
verify=tlstrust.ssl_context(),
|
||||
) as client:
|
||||
resp = await client.get(url)
|
||||
except httpx.ConnectError as exc:
|
||||
# A TLS failure arrives as a ConnectError too, and it needs a different
|
||||
# answer from "nothing is listening": install the CA, don't start Ollama.
|
||||
if "CERTIFICATE_VERIFY" in str(exc).upper() or "SSL" in str(exc).upper():
|
||||
return {
|
||||
"ok": False, "kind": "tls",
|
||||
"detail": (
|
||||
"The endpoint's TLS certificate could not be verified. If it "
|
||||
"uses a private or self-signed CA, install that CA on this "
|
||||
"machine so the system trusts it. Certificate checking is "
|
||||
"not optional."
|
||||
),
|
||||
}
|
||||
return {
|
||||
"ok": False, "kind": "unreachable",
|
||||
"detail": f"Could not connect to {endpoint_url} — is Ollama running there?",
|
||||
}
|
||||
except httpx.TimeoutException:
|
||||
return {
|
||||
"ok": False, "kind": "timeout",
|
||||
"detail": f"{endpoint_url} did not answer within {TEST_TIMEOUT:.0f}s.",
|
||||
}
|
||||
except httpx.HTTPError as exc:
|
||||
return {"ok": False, "detail": f"Connection failed: {exc}"}
|
||||
return {"ok": False, "kind": "error", "detail": f"Connection failed: {exc}"}
|
||||
|
||||
if resp.status_code != 200:
|
||||
return {"ok": False, "detail": f"HTTP {resp.status_code}: {resp.text[:300]}"}
|
||||
return {
|
||||
"ok": False, "kind": "http",
|
||||
"detail": f"HTTP {resp.status_code}: {resp.text[:300]}",
|
||||
}
|
||||
|
||||
models_available: list[str] = []
|
||||
try:
|
||||
@@ -115,15 +167,19 @@ async def list_endpoint_models(cfg: auth.ProviderConfig) -> dict:
|
||||
|
||||
@router.post("/test")
|
||||
async def test_connection(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
user: models.User = Depends(auth.get_current_user),
|
||||
):
|
||||
"""Runs a cheap connectivity check against whatever the turn engine would use.
|
||||
|
||||
That includes the shared demo endpoint, when the user has no key of their
|
||||
own.
|
||||
"""
|
||||
limits.rate_limit("connection-test", request, user)
|
||||
"""Checks the endpoint the turn engine would use, and lists its models."""
|
||||
settings = get_settings(db, user)
|
||||
return await list_endpoint_models(auth.resolve_provider_config(settings))
|
||||
result = await list_endpoint_models(settings.endpoint_url)
|
||||
if result.get("ok") and settings.model and settings.model not in result["models"]:
|
||||
# Reachable, but pointed at a model that is not installed there — the
|
||||
# commonest way for a correct endpoint to still fail every turn.
|
||||
return result | {
|
||||
"warning": (
|
||||
f"{settings.endpoint_url} is reachable, but has no model named "
|
||||
f"'{settings.model}'. Pull it there, or pick one from the list."
|
||||
)
|
||||
}
|
||||
return result
|
||||
|
||||
@@ -159,7 +159,6 @@ def import_story_cards(
|
||||
raise HTTPException(422, 'Expected a "cards" array of story cards.')
|
||||
cards_in = [c for c in cards_in if isinstance(c, dict)]
|
||||
|
||||
limits.rate_limit("import", request, user)
|
||||
limits.check_bundle_lists(story_cards=cards_in)
|
||||
existing = len(owner.story_cards)
|
||||
if auth.MULTI_USER and existing + len(cards_in) > limits.MAX_STORY_CARDS_PER_OWNER:
|
||||
|
||||
+5
-76
@@ -14,7 +14,6 @@ NAME_MAX = 200 # Titles and names. VARCHAR(200).
|
||||
TAGS_MAX = 500 # VARCHAR(500).
|
||||
CARD_TYPE_MAX = 100 # VARCHAR(100).
|
||||
PROSE_MAX = 50_000 # Memory, author's note, prompts, entries, and notes.
|
||||
SCRIPT_MAX = 200_000 # One JavaScript source file.
|
||||
ACTION_MAX = 20_000 # One player action.
|
||||
MEMORY_TEXT_MAX = 5_000
|
||||
# A scenario cover image, stored inline as a base64 data URI. A 400x300 WebP at
|
||||
@@ -31,7 +30,6 @@ Name = Annotated[str, Field(max_length=NAME_MAX)]
|
||||
Tags = Annotated[str, Field(max_length=TAGS_MAX)]
|
||||
CardType = Annotated[str, Field(max_length=CARD_TYPE_MAX)]
|
||||
Prose = Annotated[str, Field(max_length=PROSE_MAX)]
|
||||
ScriptSource = Annotated[str, Field(max_length=SCRIPT_MAX)]
|
||||
ActionText = Annotated[str, Field(max_length=ACTION_MAX)]
|
||||
Image = Annotated[str, Field(max_length=IMAGE_MAX)]
|
||||
Icon = Annotated[str, Field(max_length=ICON_MAX)]
|
||||
@@ -106,7 +104,6 @@ class ScenarioUpdate(BaseModel):
|
||||
image: Image | None = None
|
||||
icon: Icon | None = None
|
||||
stat_schema: dict | None = None
|
||||
script_ids: list[int] | None = None
|
||||
|
||||
|
||||
class ScenarioOut(ORMModel, ScenarioBase):
|
||||
@@ -115,7 +112,6 @@ class ScenarioOut(ORMModel, ScenarioBase):
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
story_cards: list[StoryCardOut] = []
|
||||
scripts: list["ScriptOut"] = []
|
||||
|
||||
|
||||
class ScenarioListItem(ORMModel):
|
||||
@@ -374,67 +370,6 @@ class AdventureListItem(ORMModel):
|
||||
icon: str = ""
|
||||
|
||||
|
||||
# ---------- Scripts ----------
|
||||
|
||||
class ScriptBase(BaseModel):
|
||||
name: Name = "Untitled Script"
|
||||
description: Prose = ""
|
||||
library_js: ScriptSource = ""
|
||||
input_js: ScriptSource = ""
|
||||
context_js: ScriptSource = ""
|
||||
output_js: ScriptSource = ""
|
||||
|
||||
|
||||
class ScriptCreate(ScriptBase):
|
||||
pass
|
||||
|
||||
|
||||
class ScriptUpdate(BaseModel):
|
||||
name: Name | None = None
|
||||
description: Prose | None = None
|
||||
library_js: ScriptSource | None = None
|
||||
input_js: ScriptSource | None = None
|
||||
context_js: ScriptSource | None = None
|
||||
output_js: ScriptSource | None = None
|
||||
|
||||
|
||||
class ScriptOut(ORMModel, ScriptBase):
|
||||
id: int
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class ScriptTestRequest(BaseModel):
|
||||
hook: Literal["input", "context", "output"]
|
||||
text: Prose = ""
|
||||
state: dict = {}
|
||||
|
||||
|
||||
class AdventureScriptOut(ORMModel):
|
||||
id: int
|
||||
adventure_id: int
|
||||
position: int
|
||||
enabled: bool
|
||||
name: str
|
||||
description: str
|
||||
library_js: str
|
||||
input_js: str
|
||||
context_js: str
|
||||
output_js: str
|
||||
# The router sets this field, which is not stored. It is `True` when a
|
||||
# syncable library version exists whose code differs from this copy, and
|
||||
# `None` when there is nothing to sync from.
|
||||
out_of_date: bool | None = None
|
||||
|
||||
|
||||
class AdventureScriptUpdate(BaseModel):
|
||||
enabled: bool | None = None
|
||||
library_js: ScriptSource | None = None
|
||||
input_js: ScriptSource | None = None
|
||||
context_js: ScriptSource | None = None
|
||||
output_js: ScriptSource | None = None
|
||||
|
||||
|
||||
# ---------- Auth (Phase 8) ----------
|
||||
|
||||
class AuthCredentials(BaseModel):
|
||||
@@ -448,15 +383,12 @@ class AuthCredentials(BaseModel):
|
||||
|
||||
class SettingsOut(ORMModel):
|
||||
endpoint_url: str
|
||||
# The key itself is never returned. It is encrypted at rest and
|
||||
# write-only.
|
||||
has_api_key: bool
|
||||
model: str
|
||||
api_mode: str
|
||||
temperature: float
|
||||
max_output_tokens: int
|
||||
reasoning_max_tokens: int
|
||||
context_token_budget: int
|
||||
model_timeout_seconds: int
|
||||
narrator_prompt: str
|
||||
summary_model: str
|
||||
embedding_model: str
|
||||
@@ -491,18 +423,15 @@ class ChatRequest(BaseModel):
|
||||
|
||||
class SettingsUpdate(BaseModel):
|
||||
endpoint_url: Annotated[str, Field(max_length=500)] | None = None # VARCHAR(500).
|
||||
# Encryption expands the stored value by about four thirds into the same
|
||||
# VARCHAR(500), so 256 plaintext characters is the largest safe input. The
|
||||
# stored form is "enc:" plus Fernet plus base64.
|
||||
api_key: Annotated[str, Field(max_length=256)] | None = None
|
||||
model: Name | None = None
|
||||
api_mode: Annotated[str, Field(max_length=20)] | None = None
|
||||
temperature: Annotated[float, Field(ge=0, le=5)] | None = None
|
||||
max_output_tokens: Annotated[int, Field(ge=1, le=100_000)] | None = None
|
||||
# A value of -1 turns reasoning off explicitly, which sends
|
||||
# `reasoning: {effort: none}`. A value of 0 sends nothing.
|
||||
reasoning_max_tokens: Annotated[int, Field(ge=-1, le=100_000)] | None = None
|
||||
context_token_budget: Annotated[int, Field(ge=256, le=200_000)] | None = None
|
||||
# Seconds to wait for the model. The floor is high enough that a normal
|
||||
# turn cannot trip it; the ceiling exists so that "wait longer" stays a
|
||||
# number rather than becoming "wait forever".
|
||||
model_timeout_seconds: Annotated[int, Field(ge=30, le=3600)] | None = None
|
||||
narrator_prompt: Prose | None = None
|
||||
summary_model: Name | None = None
|
||||
embedding_model: Name | None = None
|
||||
|
||||
@@ -1,4 +0,0 @@
|
||||
from .engine import HookResult, run_hook
|
||||
from .pipeline import ScriptPipeline
|
||||
|
||||
__all__ = ["HookResult", "ScriptPipeline", "run_hook"]
|
||||
@@ -1,146 +0,0 @@
|
||||
"""AI Dungeon-compatible script execution in an embedded QuickJS sandbox.
|
||||
|
||||
Each hook run is fully isolated (fresh Context), capped at 16 MB memory and
|
||||
2 seconds CPU, with no filesystem/network/process access (QuickJS has none by
|
||||
default). Scripts follow the AI Dungeon contract: define a `modifier(text)`
|
||||
and call it as the last line; its return value `{ text, stop }` is the result.
|
||||
"""
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import quickjs
|
||||
|
||||
MEMORY_LIMIT = 16 * 1024 * 1024
|
||||
TIME_LIMIT_SECONDS = 2
|
||||
HISTORY_WINDOW = 100 # recent actions exposed as `history`
|
||||
|
||||
# Globals per the official docs: text, state, history, storyCards, info,
|
||||
# log/console.log, story card functions, plus legacy worldInfo aliases.
|
||||
PRELUDE = """
|
||||
"use strict";
|
||||
var __logs = [];
|
||||
var state = __DATA__.state;
|
||||
var text = __DATA__.text;
|
||||
var history = __DATA__.history;
|
||||
var storyCards = __DATA__.storyCards;
|
||||
var info = __DATA__.info;
|
||||
|
||||
function log(msg) {
|
||||
__logs.push(typeof msg === "string" ? msg : JSON.stringify(msg));
|
||||
}
|
||||
var console = { log: log };
|
||||
|
||||
// Returns the new card's index, or false if a card with those keys exists —
|
||||
// matching real AI Dungeon. Note index 0 is falsy; that quirk is upstream's.
|
||||
function addStoryCard(keys, entry, type) {
|
||||
for (var i = 0; i < storyCards.length; i++) {
|
||||
if (storyCards[i].keys === keys) return false;
|
||||
}
|
||||
storyCards.push({ id: null, keys: keys || "", entry: entry || "", type: type || "" });
|
||||
return storyCards.length - 1;
|
||||
}
|
||||
function updateStoryCard(index, keys, entry, type) {
|
||||
var card = storyCards[index];
|
||||
if (!card) throw new Error("Story card not found");
|
||||
card.keys = keys;
|
||||
card.entry = entry;
|
||||
card.type = type;
|
||||
}
|
||||
function removeStoryCard(index) {
|
||||
if (!storyCards[index]) throw new Error("Story card not found");
|
||||
storyCards.splice(index, 1);
|
||||
}
|
||||
|
||||
// Legacy aliases used by older AI Dungeon scripts.
|
||||
var worldInfo = storyCards;
|
||||
var worldEntries = storyCards;
|
||||
function addWorldEntry(keys, entry) { return addStoryCard(keys, entry, ""); }
|
||||
function updateWorldEntry(index, keys, entry) {
|
||||
var card = storyCards[index];
|
||||
if (!card) throw new Error("World entry not found");
|
||||
card.keys = keys;
|
||||
card.entry = entry;
|
||||
}
|
||||
function removeWorldEntry(index) { return removeStoryCard(index); }
|
||||
"""
|
||||
|
||||
COLLECT = """
|
||||
JSON.stringify({
|
||||
result: (typeof __result === "undefined" || __result === null) ? null : __result,
|
||||
state: state,
|
||||
storyCards: storyCards,
|
||||
logs: __logs
|
||||
})
|
||||
"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class HookResult:
|
||||
text: str
|
||||
stop: bool = False
|
||||
state: dict = field(default_factory=dict)
|
||||
story_cards: list = field(default_factory=list)
|
||||
logs: list = field(default_factory=list)
|
||||
error: str | None = None
|
||||
|
||||
|
||||
def run_hook(
|
||||
library_js: str,
|
||||
hook_js: str,
|
||||
text: str,
|
||||
state: dict,
|
||||
history: list[dict],
|
||||
story_cards: list[dict],
|
||||
info: dict,
|
||||
) -> HookResult:
|
||||
"""Run one modifier hook. This function never raises. Failures return as
|
||||
`.error` with text, state, and cards unchanged, so a bad script cannot
|
||||
break a turn."""
|
||||
unchanged = HookResult(text=text, state=state, story_cards=story_cards)
|
||||
source = f"{library_js}\n;\n{hook_js}" if library_js.strip() else hook_js
|
||||
if not source.strip():
|
||||
return unchanged
|
||||
|
||||
data = {
|
||||
"state": state,
|
||||
"text": text,
|
||||
"history": history[-HISTORY_WINDOW:],
|
||||
"storyCards": story_cards,
|
||||
"info": info,
|
||||
}
|
||||
try:
|
||||
ctx = quickjs.Context()
|
||||
ctx.set_memory_limit(MEMORY_LIMIT)
|
||||
ctx.set_time_limit(TIME_LIMIT_SECONDS)
|
||||
ctx.eval(f"var __DATA__ = {json.dumps(data)};")
|
||||
ctx.eval(PRELUDE)
|
||||
ctx.eval(f"var __SRC__ = {json.dumps(source)};")
|
||||
# Indirect eval keeps the script in global scope, so `modifier(text)` as the
|
||||
# script's final expression statement becomes the completion value.
|
||||
ctx.eval("var __result = (0, eval)(__SRC__);")
|
||||
collected = json.loads(ctx.eval(COLLECT))
|
||||
except quickjs.JSException as exc:
|
||||
unchanged.error = f"Script error: {exc}"
|
||||
return unchanged
|
||||
except Exception as exc: # memory limit, invalid JSON state, engine faults
|
||||
unchanged.error = f"Script execution failed: {exc}"
|
||||
return unchanged
|
||||
|
||||
result = collected.get("result")
|
||||
new_text, stop = text, False
|
||||
if isinstance(result, dict):
|
||||
if isinstance(result.get("text"), str):
|
||||
new_text = result["text"]
|
||||
stop = bool(result.get("stop"))
|
||||
elif isinstance(result, str):
|
||||
new_text = result
|
||||
|
||||
new_state = collected.get("state")
|
||||
return HookResult(
|
||||
text=new_text,
|
||||
stop=stop,
|
||||
state=new_state if isinstance(new_state, dict) else {},
|
||||
story_cards=collected.get("storyCards") or [],
|
||||
logs=collected.get("logs") or [],
|
||||
)
|
||||
@@ -1,112 +0,0 @@
|
||||
"""Runs an adventure's enabled scripts through a turn's hook points, applying
|
||||
state and story-card mutations back to the database after each hook."""
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .. import models
|
||||
from ..context import history as context_history
|
||||
from .engine import run_hook
|
||||
|
||||
MAX_STORY_CARDS = 5000 # AI Dungeon's per-adventure sanity cap
|
||||
|
||||
|
||||
class ScriptPipeline:
|
||||
def __init__(self, adventure: models.Adventure, db: Session):
|
||||
self.adventure = adventure
|
||||
self.db = db
|
||||
self.logs: list[str] = []
|
||||
self.errors: list[str] = []
|
||||
|
||||
@property
|
||||
def message(self) -> str | None:
|
||||
state = self.adventure.script_state
|
||||
msg = state.get("message") if isinstance(state, dict) else None
|
||||
return msg if isinstance(msg, str) and msg.strip() else None
|
||||
|
||||
def _history(self) -> list[dict]:
|
||||
# Read the path rather than `adventure.actions`. That collection holds
|
||||
# every branch's actions, and this is the documented history API a user
|
||||
# script reads. Giving a script the siblings of the turn it is running on
|
||||
# would be the same bug as building a prompt from them, and visible to
|
||||
# the user.
|
||||
#
|
||||
# `story_actions` also drops rows with blank text, which
|
||||
# `adventure.actions` kept, so this array is shorter than it was for an
|
||||
# adventure that has any such rows. `info.actionCount` counts the same
|
||||
# way. That is intended. A row with no text is this app's bookkeeping, it
|
||||
# has no counterpart in the AI Dungeon history a ported script was
|
||||
# written against, and the prompt has never included one. A script keyed
|
||||
# on every N actions lands on different turns than it did before phase
|
||||
# 14, and no reading of this is compatible with both.
|
||||
return [
|
||||
{"text": a.text, "rawText": a.text, "type": a.type}
|
||||
for a in context_history.story_actions(self.adventure)
|
||||
]
|
||||
|
||||
def _cards(self) -> list[dict]:
|
||||
return [
|
||||
{"id": c.id, "keys": c.keys, "entry": c.entry, "type": c.type}
|
||||
for c in self.adventure.story_cards
|
||||
]
|
||||
|
||||
def _info(self) -> dict:
|
||||
return {
|
||||
"actionCount": context_history.count(self.adventure),
|
||||
"characterNames": [],
|
||||
"memoryLength": len(self.adventure.memory),
|
||||
"maxChars": 0,
|
||||
}
|
||||
|
||||
def _apply_cards(self, returned: list) -> None:
|
||||
existing = {c.id: c for c in self.adventure.story_cards}
|
||||
seen_ids = set()
|
||||
added = 0
|
||||
for item in returned:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
card_id = item.get("id")
|
||||
keys = str(item.get("keys") or "")
|
||||
entry = str(item.get("entry") or "")
|
||||
card_type = str(item.get("type") or "")
|
||||
if card_id in existing:
|
||||
seen_ids.add(card_id)
|
||||
card = existing[card_id]
|
||||
card.keys, card.entry, card.type = keys, entry, card_type
|
||||
elif len(existing) + added < MAX_STORY_CARDS:
|
||||
self.db.add(
|
||||
models.StoryCard(
|
||||
adventure_id=self.adventure.id,
|
||||
keys=keys, entry=entry, type=card_type,
|
||||
)
|
||||
)
|
||||
added += 1
|
||||
for card_id, card in existing.items():
|
||||
if card_id not in seen_ids:
|
||||
self.db.delete(card)
|
||||
|
||||
def run(self, hook: str, text: str) -> tuple[str, bool]:
|
||||
"""Chain `hook` across all enabled scripts. Returns (text, stop)."""
|
||||
state = self.adventure.script_state if isinstance(self.adventure.script_state, dict) else {}
|
||||
for script in self.adventure.scripts:
|
||||
hook_js = getattr(script, f"{hook}_js")
|
||||
if not script.enabled or not hook_js.strip():
|
||||
continue
|
||||
result = run_hook(
|
||||
script.library_js, hook_js, text, state,
|
||||
self._history(), self._cards(), self._info(),
|
||||
)
|
||||
if result.error:
|
||||
self.errors.append(f"{script.name} ({hook}): {result.error}")
|
||||
continue # a broken script never breaks the turn
|
||||
self.logs.extend(f"[{script.name}/{hook}] {line}" for line in result.logs)
|
||||
self._apply_cards(result.story_cards)
|
||||
state = result.state
|
||||
self.adventure.script_state = state
|
||||
self.db.commit()
|
||||
text = result.text
|
||||
if result.stop:
|
||||
return text, True
|
||||
return text, False
|
||||
|
||||
def report(self) -> dict:
|
||||
return {"logs": self.logs, "errors": self.errors, "message": self.message}
|
||||
@@ -1,131 +0,0 @@
|
||||
"""Phase 8: secrets and crypto primitives for optional accounts.
|
||||
|
||||
Everything derives from one server-side secret:
|
||||
|
||||
* Session cookies are HMAC-signed with it.
|
||||
* Stored LLM API keys are Fernet-encrypted with a key derived from it.
|
||||
|
||||
The secret comes from `AIDND_SECRET_KEY`, or it is generated once into
|
||||
`secret.key` next to the database, so a local install and a Docker volume work
|
||||
with no configuration. Losing that file logs everyone out and makes the stored
|
||||
API keys unreadable, and users then re-enter them. A multi-user deployment has
|
||||
to set the environment variable, because a hosted filesystem is ephemeral and a
|
||||
`secret.key` regenerated on every deploy would log out every user each time.
|
||||
|
||||
Passwords use `hashlib.scrypt`, which is in the standard library and backed by
|
||||
OpenSSL, so this needs no separate hashing dependency.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import os
|
||||
import secrets
|
||||
|
||||
from cryptography.fernet import Fernet, InvalidToken
|
||||
|
||||
from .database import DB_PATH
|
||||
|
||||
_SECRET_FILE = DB_PATH.parent / "secret.key"
|
||||
|
||||
|
||||
def _load_secret() -> bytes:
|
||||
env = os.environ.get("AIDND_SECRET_KEY", "").strip()
|
||||
if env:
|
||||
return env.encode()
|
||||
# Same flag parse as auth.MULTI_USER (auth imports this module, so it
|
||||
# can't be imported from there).
|
||||
if os.environ.get("AIDND_MULTI_USER", "").strip().lower() in ("1", "true", "yes", "on"):
|
||||
raise RuntimeError(
|
||||
"AIDND_SECRET_KEY must be set when AIDND_MULTI_USER is on: an "
|
||||
"auto-generated secret.key on an ephemeral hosted filesystem would "
|
||||
"rotate on every deploy, logging out every user and orphaning "
|
||||
"their stored API keys. Generate one with: "
|
||||
"python -c \"import secrets; print(secrets.token_urlsafe(48))\""
|
||||
)
|
||||
if _SECRET_FILE.exists():
|
||||
return _SECRET_FILE.read_bytes().strip()
|
||||
secret = secrets.token_urlsafe(48).encode()
|
||||
_SECRET_FILE.write_bytes(secret)
|
||||
return secret
|
||||
|
||||
|
||||
SECRET_KEY = _load_secret()
|
||||
_fernet = Fernet(base64.urlsafe_b64encode(hashlib.sha256(SECRET_KEY).digest()))
|
||||
|
||||
|
||||
# ---------- Password hashing (scrypt) ----------
|
||||
|
||||
_SCRYPT_N, _SCRYPT_R, _SCRYPT_P = 2**14, 8, 1
|
||||
|
||||
|
||||
def hash_password(password: str) -> str:
|
||||
salt = secrets.token_bytes(16)
|
||||
key = hashlib.scrypt(
|
||||
password.encode(), salt=salt, n=_SCRYPT_N, r=_SCRYPT_R, p=_SCRYPT_P
|
||||
)
|
||||
return f"scrypt${_SCRYPT_N}${_SCRYPT_R}${_SCRYPT_P}${salt.hex()}${key.hex()}"
|
||||
|
||||
|
||||
def verify_password(password: str, stored: str) -> bool:
|
||||
try:
|
||||
scheme, n, r, p, salt_hex, key_hex = stored.split("$")
|
||||
if scheme != "scrypt":
|
||||
return False
|
||||
key = hashlib.scrypt(
|
||||
password.encode(), salt=bytes.fromhex(salt_hex),
|
||||
n=int(n), r=int(r), p=int(p),
|
||||
)
|
||||
return hmac.compare_digest(key, bytes.fromhex(key_hex))
|
||||
except (ValueError, AttributeError):
|
||||
return False
|
||||
|
||||
|
||||
# ---------- Session tokens ----------
|
||||
# The token is "v1.<user_id>.<hmac>". It does not expire, because a long-lived
|
||||
# guest session is what this is for.
|
||||
|
||||
def sign_session(user_id: int) -> str:
|
||||
payload = f"v1.{user_id}"
|
||||
sig = hmac.new(SECRET_KEY, payload.encode(), hashlib.sha256).hexdigest()
|
||||
return f"{payload}.{sig}"
|
||||
|
||||
|
||||
def verify_session(token: str) -> int | None:
|
||||
try:
|
||||
version, user_id, sig = token.split(".")
|
||||
if version != "v1":
|
||||
return None
|
||||
payload = f"{version}.{user_id}"
|
||||
expected = hmac.new(SECRET_KEY, payload.encode(), hashlib.sha256).hexdigest()
|
||||
if not hmac.compare_digest(sig, expected):
|
||||
return None
|
||||
return int(user_id)
|
||||
except (ValueError, AttributeError):
|
||||
return None
|
||||
|
||||
|
||||
# ---------- API-key encryption at rest ----------
|
||||
# Stored values carry an "enc:" prefix so plaintext keys from pre-Phase-8
|
||||
# databases can be recognized and migrated.
|
||||
|
||||
ENC_PREFIX = "enc:"
|
||||
|
||||
|
||||
def encrypt_secret(plain: str) -> str:
|
||||
if not plain:
|
||||
return ""
|
||||
return ENC_PREFIX + _fernet.encrypt(plain.encode()).decode()
|
||||
|
||||
|
||||
def decrypt_secret(stored: str) -> str:
|
||||
"""Returns the plaintext key. Tolerates legacy plaintext values (returned
|
||||
as-is) and undecryptable tokens (secret rotated → treated as unset)."""
|
||||
if not stored:
|
||||
return ""
|
||||
if not stored.startswith(ENC_PREFIX):
|
||||
return stored
|
||||
try:
|
||||
return _fernet.decrypt(stored[len(ENC_PREFIX):].encode()).decode()
|
||||
except (InvalidToken, ValueError):
|
||||
return ""
|
||||
+4
-41
@@ -4,8 +4,7 @@ Every JSON file in ``seed_data/`` describes one demo scenario in the same
|
||||
model-native shape the export endpoint produces. Seeded scenarios have a NULL
|
||||
owner and ``is_public=True``, so every visitor (including guests) sees them and
|
||||
can start an adventure from them, while nobody can edit them. Starting an
|
||||
adventure copies the scenario's story cards and scripts into the adventure, so
|
||||
the seeded scripts run for guests too.
|
||||
adventure copies the scenario's story cards into the adventure.
|
||||
|
||||
Seed files are the source of truth for demo content: a scenario is inserted if
|
||||
missing, reconciled in place when a seed file's content changes, and deleted
|
||||
@@ -13,7 +12,7 @@ when no file claims its title any more, so an edit ships on the next deploy.
|
||||
Rename a seed by changing its `title` and listing the old one under
|
||||
`previous_titles`, which moves the rename onto the existing row. When a seed already matches, nothing is written, so
|
||||
this stays cheap to run on every boot. An adventure already started from a demo
|
||||
keeps its own copied cards and scripts and is unchanged. Only a new adventure
|
||||
keeps its own copied cards and is unchanged. Only a new adventure
|
||||
picks up the updated content.
|
||||
"""
|
||||
|
||||
@@ -35,7 +34,6 @@ SEED_DIR = Path(__file__).resolve().parent / "seed_data"
|
||||
_SCALARS = ("title", "description", "prompt", "memory", "authors_note", "ai_instructions",
|
||||
"tags", "image", "icon")
|
||||
_CARD_FIELDS = ("type", "name", "keys", "entry", "notes")
|
||||
_SCRIPT_FIELDS = ("name", "library_js", "input_js", "context_js", "output_js")
|
||||
|
||||
|
||||
def seed_public_scenarios(engine: Engine) -> None:
|
||||
@@ -98,7 +96,7 @@ def _sweep_unclaimed(db, claimed: set[str]) -> int:
|
||||
nothing anybody created can be reached from here.
|
||||
|
||||
An adventure started from a deleted demo survives. `adventures.scenario_id`
|
||||
is `ON DELETE SET NULL`, so the story, its cards, and its scripts are its
|
||||
is `ON DELETE SET NULL`, so the story and its cards are its
|
||||
own copies and stay; the adventure loses the cover art it inherited.
|
||||
|
||||
The caller skips this when a seed file failed to parse. A file that cannot
|
||||
@@ -117,11 +115,6 @@ def _sweep_unclaimed(db, claimed: set[str]) -> int:
|
||||
)
|
||||
for scenario in stale:
|
||||
logger.info("Removing seeded scenario %r; no seed file claims it.", scenario.title)
|
||||
# The scripts are joined through a secondary table, so nothing cascades
|
||||
# to them. They have a NULL owner and no other reader.
|
||||
for script in list(scenario.scripts):
|
||||
db.delete(script)
|
||||
scenario.scripts = []
|
||||
db.delete(scenario)
|
||||
return len(stale)
|
||||
|
||||
@@ -130,10 +123,6 @@ def _card_tuple(source, get) -> tuple:
|
||||
return tuple(get(source, f) for f in _CARD_FIELDS)
|
||||
|
||||
|
||||
def _script_tuple(source, get) -> tuple:
|
||||
return tuple(get(source, f) for f in _SCRIPT_FIELDS)
|
||||
|
||||
|
||||
def find_seeded(db, title: str) -> models.Scenario | None:
|
||||
"""Returns the seeded scenario with this exact title, if there is one."""
|
||||
return (
|
||||
@@ -178,14 +167,7 @@ def _matches(scenario: models.Scenario, data: dict) -> bool:
|
||||
_card_tuple(c, lambda o, f: o.get(f, ""))
|
||||
for c in (data.get("story_cards") or []) if isinstance(c, dict)
|
||||
)
|
||||
if have_cards != want_cards:
|
||||
return False
|
||||
have_scripts = sorted(_script_tuple(s, lambda o, f: getattr(o, f)) for s in scenario.scripts)
|
||||
want_scripts = sorted(
|
||||
_script_tuple(s, lambda o, f: (o.get(f, "") or ("Script" if f == "name" else "")))
|
||||
for s in (data.get("scripts") or []) if isinstance(s, dict)
|
||||
)
|
||||
return have_scripts == want_scripts
|
||||
return have_cards == want_cards
|
||||
|
||||
|
||||
def _insert_scenario(db, data: dict) -> None:
|
||||
@@ -203,9 +185,6 @@ def _update_scenario(db, scenario: models.Scenario, data: dict) -> None:
|
||||
# adventure foreign keys that point at it, intact.
|
||||
for card in list(scenario.story_cards):
|
||||
db.delete(card)
|
||||
for script in list(scenario.scripts):
|
||||
db.delete(script)
|
||||
scenario.scripts = []
|
||||
db.flush()
|
||||
_populate_children(db, scenario, data)
|
||||
|
||||
@@ -231,19 +210,3 @@ def _populate_children(db, scenario: models.Scenario, data: dict) -> None:
|
||||
notes=card.get("notes", ""),
|
||||
)
|
||||
)
|
||||
|
||||
for item in data.get("scripts") or []:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
script = models.Script(
|
||||
user_id=None,
|
||||
name=item.get("name", "Script"),
|
||||
description=item.get("description", ""),
|
||||
library_js=item.get("library_js", ""),
|
||||
input_js=item.get("input_js", ""),
|
||||
context_js=item.get("context_js", ""),
|
||||
output_js=item.get("output_js", ""),
|
||||
)
|
||||
db.add(script)
|
||||
db.flush()
|
||||
scenario.scripts.append(script)
|
||||
|
||||
+5
-9
@@ -6,11 +6,9 @@ frames, so the format lives here rather than in either one.
|
||||
"""
|
||||
import json
|
||||
|
||||
from . import analytics
|
||||
|
||||
# `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.
|
||||
# makes an nginx-style reverse proxy flush each event immediately rather than
|
||||
# buffer it, which matters if anyone puts one in front of the app.
|
||||
SSE_HEADERS = {"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}
|
||||
|
||||
|
||||
@@ -20,11 +18,9 @@ def sse(obj: dict) -> str:
|
||||
|
||||
|
||||
def turn_error(detail: str, **extra) -> str:
|
||||
"""Returns an SSE error for a turn that could not be produced, and counts it.
|
||||
"""Returns an SSE error for a turn that could not be produced.
|
||||
|
||||
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.
|
||||
A failed turn is still an HTTP 200 response, because the error is reported
|
||||
inside the stream the client is already reading.
|
||||
"""
|
||||
analytics.record(analytics.M_EVENT, analytics.EV_TURN_ERROR)
|
||||
return sse({"type": "error", "detail": detail, **extra})
|
||||
|
||||
Reference in New Issue
Block a user