현재 작업 전체 반영
This commit is contained in:
parent
5560638e54
commit
c0dddab594
85 changed files with 11322 additions and 539 deletions
|
|
@ -4,9 +4,10 @@ from __future__ import annotations
|
|||
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..auth_sessions import (
|
||||
|
|
@ -23,11 +24,13 @@ from ..engine_client import engine_client
|
|||
from ..runtime_policy import require_runtime_fallback_allowed
|
||||
from ..services.voice import voice_service
|
||||
from ..services import rag
|
||||
from ..store import store
|
||||
|
||||
router = APIRouter(prefix="/admin", tags=["admin"])
|
||||
|
||||
AdminPrincipal = Annotated[Principal, Depends(require_role(Role.ADMIN))]
|
||||
HealthStatus = Literal["ok", "degraded", "down"]
|
||||
UsageBudgetStatus = Literal["disabled", "ok", "warn", "exceeded"]
|
||||
|
||||
|
||||
class AdminServiceHealth(BaseModel):
|
||||
|
|
@ -46,6 +49,36 @@ class AdminHealthResponse(BaseModel):
|
|||
services: list[AdminServiceHealth]
|
||||
|
||||
|
||||
class AdminUsageBreakdown(BaseModel):
|
||||
provider: str
|
||||
model: str
|
||||
turns: int
|
||||
tokens_in: int
|
||||
tokens_out: int
|
||||
cost_usd: float
|
||||
|
||||
|
||||
class AdminUsageBudget(BaseModel):
|
||||
limit_usd: float
|
||||
used_ratio: float
|
||||
remaining_usd: float | None
|
||||
status: UsageBudgetStatus
|
||||
|
||||
|
||||
class AdminUsageResponse(BaseModel):
|
||||
source: Literal["database", "server_session_registry"]
|
||||
durable: bool
|
||||
window_days: int
|
||||
generated_at: float
|
||||
total_turns: int
|
||||
metered_turns: int
|
||||
tokens_in: int
|
||||
tokens_out: int
|
||||
cost_usd: float
|
||||
budget: AdminUsageBudget
|
||||
by_provider: list[AdminUsageBreakdown]
|
||||
|
||||
|
||||
class AdminEngineConfigResponse(BaseModel):
|
||||
engine_mode: str
|
||||
engine_url: str
|
||||
|
|
@ -120,6 +153,47 @@ def _workload_load(count: int, expected_capacity: int) -> float:
|
|||
return _clamp01(count / expected_capacity)
|
||||
|
||||
|
||||
def _decimal_to_float(value: object) -> float:
|
||||
if value is None:
|
||||
return 0.0
|
||||
if isinstance(value, Decimal):
|
||||
return float(value)
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def _safe_usage_int(value: object) -> int:
|
||||
try:
|
||||
return int(value or 0)
|
||||
except (TypeError, ValueError):
|
||||
return 0
|
||||
|
||||
|
||||
def _usage_budget(cost_usd: float) -> AdminUsageBudget:
|
||||
limit = max(0.0, float(settings.admin_usage_budget_usd or 0.0))
|
||||
if limit <= 0:
|
||||
return AdminUsageBudget(
|
||||
limit_usd=0.0,
|
||||
used_ratio=0.0,
|
||||
remaining_usd=None,
|
||||
status="disabled",
|
||||
)
|
||||
used_ratio = max(0.0, cost_usd / limit)
|
||||
status_value: UsageBudgetStatus = "ok"
|
||||
if used_ratio >= 1.0:
|
||||
status_value = "exceeded"
|
||||
elif used_ratio >= 0.8:
|
||||
status_value = "warn"
|
||||
return AdminUsageBudget(
|
||||
limit_usd=round(limit, 6),
|
||||
used_ratio=round(used_ratio, 4),
|
||||
remaining_usd=round(max(0.0, limit - cost_usd), 6),
|
||||
status=status_value,
|
||||
)
|
||||
|
||||
|
||||
async def _runtime_health_metrics(*, db_ok: bool) -> RuntimeHealthMetrics:
|
||||
metrics = RuntimeHealthMetrics()
|
||||
|
||||
|
|
@ -175,6 +249,157 @@ async def _runtime_health_metrics(*, db_ok: bool) -> RuntimeHealthMetrics:
|
|||
return metrics
|
||||
|
||||
|
||||
async def _usage_from_database(window_days: int) -> AdminUsageResponse:
|
||||
async with acquire(role="admin") as conn:
|
||||
total_row = await conn.fetchrow(
|
||||
"""
|
||||
SELECT
|
||||
COUNT(*) FILTER (WHERE speaker = 'client') AS total_turns,
|
||||
COUNT(*) FILTER (
|
||||
WHERE speaker = 'client'
|
||||
AND (
|
||||
llm_provider IS NOT NULL OR model IS NOT NULL
|
||||
OR tokens_in IS NOT NULL OR tokens_out IS NOT NULL
|
||||
OR cost_usd IS NOT NULL
|
||||
)
|
||||
) AS metered_turns,
|
||||
COALESCE(SUM(tokens_in) FILTER (WHERE speaker = 'client'), 0)::bigint AS tokens_in,
|
||||
COALESCE(SUM(tokens_out) FILTER (WHERE speaker = 'client'), 0)::bigint AS tokens_out,
|
||||
COALESCE(SUM(cost_usd) FILTER (WHERE speaker = 'client'), 0)::numeric AS cost_usd
|
||||
FROM app.turns
|
||||
WHERE created_at >= now() - ($1::int * interval '1 day')
|
||||
""",
|
||||
window_days,
|
||||
)
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT
|
||||
COALESCE(llm_provider, 'unknown') AS provider,
|
||||
COALESCE(model, 'unknown') AS model,
|
||||
COUNT(*) AS turns,
|
||||
COALESCE(SUM(tokens_in), 0)::bigint AS tokens_in,
|
||||
COALESCE(SUM(tokens_out), 0)::bigint AS tokens_out,
|
||||
COALESCE(SUM(cost_usd), 0)::numeric AS cost_usd
|
||||
FROM app.turns
|
||||
WHERE created_at >= now() - ($1::int * interval '1 day')
|
||||
AND speaker = 'client'
|
||||
AND (
|
||||
llm_provider IS NOT NULL OR model IS NOT NULL
|
||||
OR tokens_in IS NOT NULL OR tokens_out IS NOT NULL
|
||||
OR cost_usd IS NOT NULL
|
||||
)
|
||||
GROUP BY 1, 2
|
||||
ORDER BY cost_usd DESC, tokens_in + tokens_out DESC, turns DESC
|
||||
LIMIT 12
|
||||
""",
|
||||
window_days,
|
||||
)
|
||||
|
||||
total_cost = round(_decimal_to_float(total_row["cost_usd"] if total_row else 0), 6)
|
||||
return AdminUsageResponse(
|
||||
source="database",
|
||||
durable=True,
|
||||
window_days=window_days,
|
||||
generated_at=time.time(),
|
||||
total_turns=_safe_usage_int(total_row["total_turns"] if total_row else 0),
|
||||
metered_turns=_safe_usage_int(total_row["metered_turns"] if total_row else 0),
|
||||
tokens_in=_safe_usage_int(total_row["tokens_in"] if total_row else 0),
|
||||
tokens_out=_safe_usage_int(total_row["tokens_out"] if total_row else 0),
|
||||
cost_usd=total_cost,
|
||||
budget=_usage_budget(total_cost),
|
||||
by_provider=[
|
||||
AdminUsageBreakdown(
|
||||
provider=str(row["provider"] or "unknown"),
|
||||
model=str(row["model"] or "unknown"),
|
||||
turns=_safe_usage_int(row["turns"]),
|
||||
tokens_in=_safe_usage_int(row["tokens_in"]),
|
||||
tokens_out=_safe_usage_int(row["tokens_out"]),
|
||||
cost_usd=round(_decimal_to_float(row["cost_usd"]), 6),
|
||||
)
|
||||
for row in rows
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _usage_from_runtime_store(window_days: int) -> AdminUsageResponse:
|
||||
window_start = time.time() - (window_days * 86400)
|
||||
total_turns = 0
|
||||
metered_turns = 0
|
||||
tokens_in = 0
|
||||
tokens_out = 0
|
||||
cost_usd = 0.0
|
||||
buckets: dict[tuple[str, str], dict[str, int | float]] = {}
|
||||
|
||||
for sess in store.list():
|
||||
for turn in getattr(sess, "turns", []) or []:
|
||||
if getattr(turn, "speaker", "") != "client":
|
||||
continue
|
||||
if float(getattr(turn, "created_at", 0.0) or 0.0) < window_start:
|
||||
continue
|
||||
total_turns += 1
|
||||
provider = str(getattr(turn, "llm_provider", None) or "unknown")
|
||||
model = str(getattr(turn, "model", None) or "unknown")
|
||||
turn_tokens_in = _safe_usage_int(getattr(turn, "tokens_in", 0))
|
||||
turn_tokens_out = _safe_usage_int(getattr(turn, "tokens_out", 0))
|
||||
turn_cost = _decimal_to_float(getattr(turn, "cost_usd", 0.0))
|
||||
is_metered = (
|
||||
provider != "unknown"
|
||||
or model != "unknown"
|
||||
or turn_tokens_in > 0
|
||||
or turn_tokens_out > 0
|
||||
or turn_cost > 0
|
||||
)
|
||||
if not is_metered:
|
||||
continue
|
||||
metered_turns += 1
|
||||
tokens_in += turn_tokens_in
|
||||
tokens_out += turn_tokens_out
|
||||
cost_usd += turn_cost
|
||||
key = (provider, model)
|
||||
bucket = buckets.setdefault(
|
||||
key,
|
||||
{"turns": 0, "tokens_in": 0, "tokens_out": 0, "cost_usd": 0.0},
|
||||
)
|
||||
bucket["turns"] = int(bucket["turns"]) + 1
|
||||
bucket["tokens_in"] = int(bucket["tokens_in"]) + turn_tokens_in
|
||||
bucket["tokens_out"] = int(bucket["tokens_out"]) + turn_tokens_out
|
||||
bucket["cost_usd"] = float(bucket["cost_usd"]) + turn_cost
|
||||
|
||||
by_provider = [
|
||||
AdminUsageBreakdown(
|
||||
provider=provider,
|
||||
model=model,
|
||||
turns=int(values["turns"]),
|
||||
tokens_in=int(values["tokens_in"]),
|
||||
tokens_out=int(values["tokens_out"]),
|
||||
cost_usd=round(float(values["cost_usd"]), 6),
|
||||
)
|
||||
for (provider, model), values in sorted(
|
||||
buckets.items(),
|
||||
key=lambda item: (
|
||||
-float(item[1]["cost_usd"]),
|
||||
-(int(item[1]["tokens_in"]) + int(item[1]["tokens_out"])),
|
||||
-int(item[1]["turns"]),
|
||||
),
|
||||
)[:12]
|
||||
]
|
||||
|
||||
total_cost = round(cost_usd, 6)
|
||||
return AdminUsageResponse(
|
||||
source="server_session_registry",
|
||||
durable=False,
|
||||
window_days=window_days,
|
||||
generated_at=time.time(),
|
||||
total_turns=total_turns,
|
||||
metered_turns=metered_turns,
|
||||
tokens_in=tokens_in,
|
||||
tokens_out=tokens_out,
|
||||
cost_usd=total_cost,
|
||||
budget=_usage_budget(total_cost),
|
||||
by_provider=by_provider,
|
||||
)
|
||||
|
||||
|
||||
class AdminUserCreate(BaseModel):
|
||||
email: str = Field(..., min_length=3, max_length=254)
|
||||
display_name: str = Field(..., min_length=1, max_length=80)
|
||||
|
|
@ -407,6 +632,19 @@ async def admin_health(principal: AdminPrincipal) -> AdminHealthResponse:
|
|||
)
|
||||
|
||||
|
||||
@router.get("/usage", response_model=AdminUsageResponse)
|
||||
async def admin_usage(
|
||||
principal: AdminPrincipal,
|
||||
window_days: Annotated[int, Query(ge=1, le=90)] = 7,
|
||||
) -> AdminUsageResponse:
|
||||
"""Return AI token/cost usage from persisted turns or dev fallback state."""
|
||||
try:
|
||||
return await _usage_from_database(window_days)
|
||||
except Exception:
|
||||
require_runtime_fallback_allowed("admin usage")
|
||||
return _usage_from_runtime_store(window_days)
|
||||
|
||||
|
||||
@router.get("/engine-config", response_model=AdminEngineConfigResponse)
|
||||
async def get_engine_config(principal: AdminPrincipal) -> AdminEngineConfigResponse:
|
||||
"""Return the current admin-managed engine settings."""
|
||||
|
|
|
|||
|
|
@ -11,6 +11,9 @@ from __future__ import annotations
|
|||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import secrets
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -34,11 +37,14 @@ from ..saml import (
|
|||
)
|
||||
|
||||
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
GOOGLE_AUTHORIZE_URL = "https://accounts.google.com/o/oauth2/v2/auth"
|
||||
GOOGLE_TOKEN_URL = "https://oauth2.googleapis.com/token"
|
||||
GOOGLE_TOKENINFO_URL = "https://oauth2.googleapis.com/tokeninfo"
|
||||
OAUTH_STATE_TTL_SECONDS = 10 * 60
|
||||
OAUTH_STATE_COOKIE_NAME = "__Host-vignette_oauth_state"
|
||||
DEV_OAUTH_STATE_COOKIE_NAME = "vignette_oauth_state"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
|
|
@ -201,6 +207,54 @@ def _role_for_saml_identity(identity: SamlIdentity) -> Role:
|
|||
return _role_for_email(identity.email)
|
||||
|
||||
|
||||
def _split_cohort_values(value: str | None) -> list[str]:
|
||||
if not value:
|
||||
return []
|
||||
return [item.strip() for item in value.split(",") if item.strip()]
|
||||
|
||||
|
||||
def _append_unique(items: list[str], values: list[str]) -> None:
|
||||
seen = {item.lower() for item in items}
|
||||
for value in values:
|
||||
key = value.lower()
|
||||
if key and key not in seen:
|
||||
items.append(value)
|
||||
seen.add(key)
|
||||
|
||||
|
||||
def _configured_cohort_ids(
|
||||
*,
|
||||
email: str,
|
||||
hosted_domain: str | None = None,
|
||||
claim_hint: str | None = None,
|
||||
) -> list[str]:
|
||||
normalized_email = _normalize_email(email)
|
||||
domain = _email_domain(normalized_email)
|
||||
hd = _normalize_domain(hosted_domain)
|
||||
email_map = {
|
||||
_normalize_email(key): value
|
||||
for key, value in settings.auth_email_cohort_map.items()
|
||||
if _normalize_email(key)
|
||||
}
|
||||
domain_map = {
|
||||
_normalize_domain(key): value
|
||||
for key, value in settings.auth_domain_cohort_map.items()
|
||||
if _normalize_domain(key)
|
||||
}
|
||||
cohorts: list[str] = []
|
||||
_append_unique(cohorts, _split_cohort_values(email_map.get(normalized_email)))
|
||||
_append_unique(cohorts, _split_cohort_values(domain_map.get(domain)))
|
||||
if hd and hd != domain:
|
||||
_append_unique(cohorts, _split_cohort_values(domain_map.get(hd)))
|
||||
_append_unique(cohorts, _split_cohort_values(claim_hint))
|
||||
return cohorts
|
||||
|
||||
|
||||
def _provider_external_id(provider: str, subject: str | None, email: str) -> str:
|
||||
value = (subject or "").strip() or _normalize_email(email)
|
||||
return f"{provider}:{value.lower()}"
|
||||
|
||||
|
||||
def _safe_next_path(next_path: str | None) -> str:
|
||||
if not next_path or not next_path.startswith("/") or next_path.startswith("//"):
|
||||
return "/"
|
||||
|
|
@ -254,7 +308,12 @@ def _is_dev_login_allowed_origin(origin: str) -> bool:
|
|||
|
||||
def _configured_frontend_origins() -> list[str]:
|
||||
origins: list[str] = []
|
||||
for value in [settings.frontend_base_url, *settings.frontend_origin_map.values(), *settings.cors_origins]:
|
||||
for value in [
|
||||
settings.frontend_base_url,
|
||||
*settings.frontend_origin_map.values(),
|
||||
*settings.cors_origins,
|
||||
*settings.auth_dev_login_extra_origins,
|
||||
]:
|
||||
origin = _url_origin(value)
|
||||
if origin and origin not in origins:
|
||||
origins.append(origin)
|
||||
|
|
@ -276,6 +335,17 @@ def _frontend_origin_for_request(request: Request | None = None) -> str:
|
|||
hostname = host.rsplit(":", 1)[0].lower() if host else ""
|
||||
if mapped_origin := _frontend_origin_map().get(hostname):
|
||||
return mapped_origin
|
||||
forwarded_proto = (
|
||||
request.headers.get("x-forwarded-proto", "").split(",", 1)[0].strip().lower()
|
||||
)
|
||||
if forwarded_host and forwarded_proto in {"http", "https"}:
|
||||
forwarded_origin = _url_origin(f"{forwarded_proto}://{host}")
|
||||
if (
|
||||
settings.environment == "dev"
|
||||
and forwarded_origin
|
||||
and _is_dev_login_allowed_origin(forwarded_origin)
|
||||
):
|
||||
return forwarded_origin
|
||||
if hostname in {"localhost", "127.0.0.1", "::1"}:
|
||||
return fallback
|
||||
|
||||
|
|
@ -299,6 +369,93 @@ def _pkce_challenge(verifier: str) -> str:
|
|||
return base64.urlsafe_b64encode(digest).decode("ascii").rstrip("=")
|
||||
|
||||
|
||||
def _b64url_encode(value: bytes) -> str:
|
||||
return base64.urlsafe_b64encode(value).decode("ascii").rstrip("=")
|
||||
|
||||
|
||||
def _b64url_decode(value: str) -> bytes:
|
||||
padded = value + ("=" * (-len(value) % 4))
|
||||
return base64.urlsafe_b64decode(padded.encode("ascii"))
|
||||
|
||||
|
||||
def _oauth_state_signature(payload: str) -> str:
|
||||
digest = hmac.new(
|
||||
settings.session_secret.encode("utf-8"),
|
||||
f"oauth-state:{payload}".encode("utf-8"),
|
||||
hashlib.sha256,
|
||||
).digest()
|
||||
return _b64url_encode(digest)
|
||||
|
||||
|
||||
def _oauth_code_verifier_for_state(state: str) -> str:
|
||||
digest = hmac.new(
|
||||
settings.session_secret.encode("utf-8"),
|
||||
f"oauth-pkce:{state}".encode("utf-8"),
|
||||
hashlib.sha256,
|
||||
).digest()
|
||||
return _b64url_encode(digest)
|
||||
|
||||
|
||||
def _build_oauth_state(next_path: str | None) -> tuple[str, OAuthState]:
|
||||
created_at = time.time()
|
||||
safe_next = _safe_next_path(next_path)
|
||||
payload = {
|
||||
"iat": created_at,
|
||||
"next": safe_next,
|
||||
"nonce": secrets.token_urlsafe(24),
|
||||
}
|
||||
payload_blob = _b64url_encode(
|
||||
json.dumps(payload, separators=(",", ":"), sort_keys=True).encode("utf-8")
|
||||
)
|
||||
state = f"{payload_blob}.{_oauth_state_signature(payload_blob)}"
|
||||
return state, OAuthState(
|
||||
code_verifier=_oauth_code_verifier_for_state(state),
|
||||
next_path=safe_next,
|
||||
created_at=created_at,
|
||||
)
|
||||
|
||||
|
||||
def _oauth_state_from_signed_token(state: str | None) -> OAuthState | None:
|
||||
if not state or "." not in state:
|
||||
return None
|
||||
payload_blob, signature = state.rsplit(".", 1)
|
||||
if not payload_blob or not signature:
|
||||
return None
|
||||
if not hmac.compare_digest(signature, _oauth_state_signature(payload_blob)):
|
||||
return None
|
||||
|
||||
try:
|
||||
payload = json.loads(_b64url_decode(payload_blob).decode("utf-8"))
|
||||
created_at = float(payload.get("iat", 0))
|
||||
except (ValueError, TypeError, json.JSONDecodeError):
|
||||
return None
|
||||
|
||||
if created_at <= 0 or created_at < time.time() - OAUTH_STATE_TTL_SECONDS:
|
||||
return None
|
||||
if not isinstance(payload.get("nonce"), str):
|
||||
return None
|
||||
|
||||
next_path = payload.get("next")
|
||||
if not isinstance(next_path, str):
|
||||
return None
|
||||
return OAuthState(
|
||||
code_verifier=_oauth_code_verifier_for_state(state),
|
||||
next_path=_safe_next_path(next_path),
|
||||
created_at=created_at,
|
||||
)
|
||||
|
||||
|
||||
def _oauth_state_for_callback(state: str | None, cookie_state: str | None) -> OAuthState | None:
|
||||
if not state:
|
||||
return None
|
||||
if cookie_state != state:
|
||||
return None
|
||||
stored = _oauth_states.pop(state, None)
|
||||
if stored is not None:
|
||||
return stored
|
||||
return _oauth_state_from_signed_token(state)
|
||||
|
||||
|
||||
def _prune_oauth_states() -> None:
|
||||
cutoff = time.time() - OAUTH_STATE_TTL_SECONDS
|
||||
stale = [key for key, value in _oauth_states.items() if value.created_at < cutoff]
|
||||
|
|
@ -342,6 +499,46 @@ def _set_session_cookie(response: Response, sid: str) -> None:
|
|||
)
|
||||
|
||||
|
||||
def _set_oauth_state_cookie(response: Response, state: str) -> None:
|
||||
response.set_cookie(
|
||||
key=OAUTH_STATE_COOKIE_NAME,
|
||||
value=state,
|
||||
max_age=OAUTH_STATE_TTL_SECONDS,
|
||||
httponly=True,
|
||||
secure=True,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
if settings.environment == "dev":
|
||||
response.set_cookie(
|
||||
key=DEV_OAUTH_STATE_COOKIE_NAME,
|
||||
value=state,
|
||||
max_age=OAUTH_STATE_TTL_SECONDS,
|
||||
httponly=True,
|
||||
secure=False,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
|
||||
|
||||
def _delete_oauth_state_cookie(response: Response) -> None:
|
||||
response.delete_cookie(
|
||||
OAUTH_STATE_COOKIE_NAME,
|
||||
httponly=True,
|
||||
secure=True,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
if settings.environment == "dev":
|
||||
response.delete_cookie(
|
||||
DEV_OAUTH_STATE_COOKIE_NAME,
|
||||
httponly=True,
|
||||
secure=False,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
|
||||
|
||||
def _delete_session_cookie(response: Response) -> None:
|
||||
response.delete_cookie(
|
||||
settings.cookie_name,
|
||||
|
|
@ -375,13 +572,28 @@ def _frontend_login_redirect(reason: str, request: Request) -> RedirectResponse:
|
|||
return RedirectResponse(f"{base_url}/login?{urlencode({'oauth': reason})}", status_code=302)
|
||||
|
||||
|
||||
def _oauth_callback_error(reason: str, request: Request) -> RedirectResponse:
|
||||
response = _frontend_login_redirect(reason, request)
|
||||
_delete_oauth_state_cookie(response)
|
||||
return response
|
||||
|
||||
|
||||
def _log_oauth_callback_failure(request: Request, reason: str, **fields: object) -> None:
|
||||
"""Log OAuth callback failures without authorization codes, tokens, or raw user IDs."""
|
||||
host = request.headers.get("x-forwarded-host") or request.headers.get("host")
|
||||
logger.warning(
|
||||
"google_oauth_callback_failed reason=%s host=%s forwarded_proto=%s details=%s",
|
||||
reason,
|
||||
host,
|
||||
request.headers.get("x-forwarded-proto"),
|
||||
fields,
|
||||
)
|
||||
|
||||
|
||||
def _dev_login_available(request: Request) -> bool:
|
||||
if settings.environment != "dev" or not settings.auth_dev_login_enabled:
|
||||
return False
|
||||
|
||||
if _dev_login_extra_origins():
|
||||
return True
|
||||
|
||||
saw_browser_origin = False
|
||||
for header_name in ("origin", "referer"):
|
||||
origin = _url_origin(request.headers.get(header_name))
|
||||
|
|
@ -389,16 +601,28 @@ def _dev_login_available(request: Request) -> bool:
|
|||
saw_browser_origin = True
|
||||
if not _is_dev_login_allowed_origin(origin):
|
||||
return False
|
||||
|
||||
if not saw_browser_origin and _dev_login_extra_origins():
|
||||
if saw_browser_origin:
|
||||
return True
|
||||
|
||||
forwarded_host = request.headers.get("x-forwarded-host")
|
||||
host = (forwarded_host or request.headers.get("host") or "").split(",", 1)[0].strip()
|
||||
origin = _url_origin(f"http://{host}") if host else None
|
||||
forwarded_proto = (
|
||||
request.headers.get("x-forwarded-proto", "").split(",", 1)[0].strip().lower()
|
||||
)
|
||||
scheme = forwarded_proto if forwarded_host and forwarded_proto in {"http", "https"} else "http"
|
||||
origin = _url_origin(f"{scheme}://{host}") if host else None
|
||||
return bool(origin and _is_dev_login_allowed_origin(origin))
|
||||
|
||||
|
||||
def _dev_oauth_redirect_unavailable(request: Request) -> bool:
|
||||
redirect_origin = _url_origin(settings.oauth_redirect_uri)
|
||||
return bool(
|
||||
_dev_login_available(request)
|
||||
and redirect_origin
|
||||
and not _is_local_origin(redirect_origin)
|
||||
)
|
||||
|
||||
|
||||
@router.get("/config", response_model=AuthConfigResponse)
|
||||
async def auth_config(request: Request) -> AuthConfigResponse:
|
||||
"""Return non-secret login configuration for the browser login screen."""
|
||||
|
|
@ -449,15 +673,12 @@ async def login(
|
|||
return _frontend_login_redirect("unsupported_provider", request)
|
||||
if not _google_configured():
|
||||
return _frontend_login_redirect("not_configured", request)
|
||||
if _dev_oauth_redirect_unavailable(request):
|
||||
return _frontend_login_redirect("local_oauth_unavailable", request)
|
||||
|
||||
_prune_oauth_states()
|
||||
state = secrets.token_urlsafe(32)
|
||||
verifier = secrets.token_urlsafe(64)
|
||||
_oauth_states[state] = OAuthState(
|
||||
code_verifier=verifier,
|
||||
next_path=_safe_next_path(next),
|
||||
created_at=time.time(),
|
||||
)
|
||||
state, stored_state = _build_oauth_state(next)
|
||||
_oauth_states[state] = stored_state
|
||||
|
||||
params = {
|
||||
"client_id": settings.oauth_google_client_id,
|
||||
|
|
@ -465,11 +686,13 @@ async def login(
|
|||
"response_type": "code",
|
||||
"scope": "openid email profile",
|
||||
"state": state,
|
||||
"code_challenge": _pkce_challenge(verifier),
|
||||
"code_challenge": _pkce_challenge(stored_state.code_verifier),
|
||||
"code_challenge_method": "S256",
|
||||
"prompt": "select_account",
|
||||
}
|
||||
return RedirectResponse(f"{GOOGLE_AUTHORIZE_URL}?{urlencode(params)}", status_code=302)
|
||||
response = RedirectResponse(f"{GOOGLE_AUTHORIZE_URL}?{urlencode(params)}", status_code=302)
|
||||
_set_oauth_state_cookie(response, state)
|
||||
return response
|
||||
|
||||
|
||||
@router.get("/callback")
|
||||
|
|
@ -477,15 +700,45 @@ async def callback(
|
|||
request: Request,
|
||||
code: Annotated[Optional[str], Query()] = None,
|
||||
state: Annotated[Optional[str], Query()] = None,
|
||||
error: Annotated[Optional[str], Query()] = None,
|
||||
error_description: Annotated[Optional[str], Query()] = None,
|
||||
oauth_state_cookie: Annotated[Optional[str], Cookie(alias=OAUTH_STATE_COOKIE_NAME)] = None,
|
||||
dev_oauth_state_cookie: Annotated[Optional[str], Cookie(alias=DEV_OAUTH_STATE_COOKIE_NAME)] = None,
|
||||
) -> RedirectResponse:
|
||||
"""Exchange Google auth code, validate identity, and issue a BFF cookie."""
|
||||
if error:
|
||||
reason = "access_denied" if error == "access_denied" else "provider_error"
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
reason,
|
||||
provider_error=error,
|
||||
has_state=bool(state),
|
||||
has_error_description=bool(error_description),
|
||||
)
|
||||
return _oauth_callback_error(reason, request)
|
||||
|
||||
if not code or not state:
|
||||
return _frontend_login_redirect("missing_callback", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"missing_callback",
|
||||
has_code=bool(code),
|
||||
has_state=bool(state),
|
||||
)
|
||||
return _oauth_callback_error("missing_callback", request)
|
||||
|
||||
_prune_oauth_states()
|
||||
stored = _oauth_states.pop(state, None)
|
||||
cookie_state = oauth_state_cookie or (
|
||||
dev_oauth_state_cookie if settings.environment == "dev" else None
|
||||
)
|
||||
stored = _oauth_state_for_callback(state, cookie_state)
|
||||
if stored is None:
|
||||
return _frontend_login_redirect("invalid_state", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"invalid_state",
|
||||
has_cookie=bool(cookie_state),
|
||||
state_in_memory=state in _oauth_states if state else False,
|
||||
)
|
||||
return _oauth_callback_error("invalid_state", request)
|
||||
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
token_res = await client.post(
|
||||
|
|
@ -501,22 +754,45 @@ async def callback(
|
|||
headers={"Accept": "application/json"},
|
||||
)
|
||||
if token_res.status_code >= 400:
|
||||
return _frontend_login_redirect("token_exchange_failed", request)
|
||||
token_error: object
|
||||
try:
|
||||
token_body = token_res.json()
|
||||
token_error = {
|
||||
"error": token_body.get("error"),
|
||||
"error_description": token_body.get("error_description"),
|
||||
}
|
||||
except Exception:
|
||||
token_error = "non_json_error"
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"token_exchange_failed",
|
||||
status_code=token_res.status_code,
|
||||
token_error=token_error,
|
||||
)
|
||||
return _oauth_callback_error("token_exchange_failed", request)
|
||||
token_payload = token_res.json()
|
||||
id_token = token_payload.get("id_token")
|
||||
if not isinstance(id_token, str) or not id_token:
|
||||
return _frontend_login_redirect("id_token_missing", request)
|
||||
_log_oauth_callback_failure(request, "id_token_missing")
|
||||
return _oauth_callback_error("id_token_missing", request)
|
||||
|
||||
info_res = await client.get(GOOGLE_TOKENINFO_URL, params={"id_token": id_token})
|
||||
if info_res.status_code >= 400:
|
||||
return _frontend_login_redirect("id_token_invalid", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"id_token_invalid",
|
||||
status_code=info_res.status_code,
|
||||
)
|
||||
return _oauth_callback_error("id_token_invalid", request)
|
||||
claims = info_res.json()
|
||||
|
||||
if claims.get("aud") != settings.oauth_google_client_id:
|
||||
return _frontend_login_redirect("audience_mismatch", request)
|
||||
_log_oauth_callback_failure(request, "audience_mismatch")
|
||||
return _oauth_callback_error("audience_mismatch", request)
|
||||
issuer = claims.get("iss")
|
||||
if issuer not in {"accounts.google.com", "https://accounts.google.com"}:
|
||||
return _frontend_login_redirect("issuer_mismatch", request)
|
||||
_log_oauth_callback_failure(request, "issuer_mismatch", issuer=issuer)
|
||||
return _oauth_callback_error("issuer_mismatch", request)
|
||||
|
||||
try:
|
||||
email = validate_google_identity_domain(
|
||||
|
|
@ -525,21 +801,39 @@ async def callback(
|
|||
hosted_domain=claims.get("hd"),
|
||||
)
|
||||
except HTTPException:
|
||||
return _frontend_login_redirect("domain_not_allowed", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"domain_not_allowed",
|
||||
email_domain=_email_domain(str(claims.get("email") or "")),
|
||||
hosted_domain=_normalize_domain(str(claims.get("hd") or "")),
|
||||
)
|
||||
return _oauth_callback_error("domain_not_allowed", request)
|
||||
role = _role_for_email(email)
|
||||
display_name = str(claims.get("name") or email)
|
||||
cohort_ids = _configured_cohort_ids(
|
||||
email=email,
|
||||
hosted_domain=str(claims.get("hd") or ""),
|
||||
)
|
||||
external_id = _provider_external_id("google", str(claims.get("sub") or ""), email)
|
||||
try:
|
||||
sid, _ = await create_session(
|
||||
email=email,
|
||||
display_name=display_name,
|
||||
role=role.value,
|
||||
cohort_ids=[],
|
||||
cohort_ids=cohort_ids,
|
||||
external_id=external_id,
|
||||
)
|
||||
except InactiveUserError as exc:
|
||||
return _frontend_login_redirect("inactive_user", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"inactive_user",
|
||||
email_domain=_email_domain(email),
|
||||
)
|
||||
return _oauth_callback_error("inactive_user", request)
|
||||
|
||||
response = RedirectResponse(_frontend_url(stored.next_path, request), status_code=302)
|
||||
_set_session_cookie(response, sid)
|
||||
_delete_oauth_state_cookie(response)
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -580,12 +874,15 @@ async def saml_acs(request: Request) -> RedirectResponse:
|
|||
return _frontend_login_redirect("saml_assertion_invalid", request)
|
||||
|
||||
role = _role_for_saml_identity(identity)
|
||||
cohort_ids = _configured_cohort_ids(email=email, claim_hint=identity.cohort_hint)
|
||||
external_id = _provider_external_id("saml", identity.subject, email)
|
||||
try:
|
||||
sid, _ = await create_session(
|
||||
email=email,
|
||||
display_name=identity.display_name or email,
|
||||
role=role.value,
|
||||
cohort_ids=[],
|
||||
cohort_ids=cohort_ids,
|
||||
external_id=external_id,
|
||||
)
|
||||
except InactiveUserError:
|
||||
return _frontend_login_redirect("inactive_user", request)
|
||||
|
|
@ -615,7 +912,8 @@ async def dev_login(request: Request, body: DevLoginRequest, response: Response)
|
|||
email=email,
|
||||
display_name=body.display_name or email,
|
||||
role=body.role,
|
||||
cohort_ids=[],
|
||||
cohort_ids=_configured_cohort_ids(email=email),
|
||||
external_id=_provider_external_id("dev", email, email),
|
||||
)
|
||||
except InactiveUserError as exc:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail="user is inactive") from exc
|
||||
|
|
|
|||
|
|
@ -115,6 +115,7 @@ async def reevaluate_session(
|
|||
technique_codes=technique_codes,
|
||||
theory_mode=_theory_mode_of(sess),
|
||||
scope=body.scope if body.scope in ("session_end", "stage_transition") else "session_end",
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
)
|
||||
except EngineError as e:
|
||||
raise HTTPException(status.HTTP_503_SERVICE_UNAVAILABLE, detail=f"engine unavailable: {e}")
|
||||
|
|
@ -184,7 +185,12 @@ async def reevaluate_turn(
|
|||
recent_turns=recent,
|
||||
)
|
||||
|
||||
result = await evaluator.evaluate_turn(ctx, client_reply, engine=engine_client)
|
||||
result = await evaluator.evaluate_turn(
|
||||
ctx,
|
||||
client_reply,
|
||||
engine=engine_client,
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
)
|
||||
if result.error and result.error.startswith("engine_error"):
|
||||
raise HTTPException(status.HTTP_503_SERVICE_UNAVAILABLE, detail=result.error)
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -5,20 +5,26 @@ from __future__ import annotations
|
|||
from typing import Annotated, Any, Literal
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Response, status
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ..deps import CurrentPrincipal, Principal, Role, require_role
|
||||
from ..persona_repository import (
|
||||
CatalogPersona,
|
||||
PersonaDraftRecord,
|
||||
PersonaReviewAction,
|
||||
PersonaReviewItem,
|
||||
create_persona_draft,
|
||||
get_persona_draft_record,
|
||||
list_catalog_personas,
|
||||
list_persona_review_queue,
|
||||
update_persona_draft,
|
||||
update_persona_review_status,
|
||||
)
|
||||
from ..services.persona import PersonaCard
|
||||
|
||||
router = APIRouter(prefix="/personas", tags=["personas"])
|
||||
TeacherOrAdmin = Annotated[Principal, Depends(require_role(Role.TEACHER, Role.ADMIN))]
|
||||
JSON_OBJECT_FIELD = {"additionalProperties": True}
|
||||
|
||||
|
||||
class PersonaSummary(BaseModel):
|
||||
|
|
@ -51,6 +57,37 @@ class PersonaReviewDecisionRequest(BaseModel):
|
|||
action: PersonaReviewAction
|
||||
|
||||
|
||||
class PersonaDraftPayload(BaseModel):
|
||||
code: str = Field(min_length=1, max_length=24)
|
||||
display_name: str = Field(min_length=1, max_length=80)
|
||||
difficulty: Literal["easy", "moderate", "hard"]
|
||||
theory_target: list[str] = Field(default_factory=list)
|
||||
demographics: dict[str, Any] = Field(default_factory=dict, json_schema_extra=JSON_OBJECT_FIELD)
|
||||
presenting: dict[str, Any] = Field(default_factory=dict, json_schema_extra=JSON_OBJECT_FIELD)
|
||||
history: dict[str, Any] = Field(default_factory=dict, json_schema_extra=JSON_OBJECT_FIELD)
|
||||
big5: dict[str, float] = Field(default_factory=dict)
|
||||
resistance: dict[str, float] = Field(default_factory=dict)
|
||||
speech_style: dict[str, Any] = Field(default_factory=dict, json_schema_extra=JSON_OBJECT_FIELD)
|
||||
affect_baseline: dict[str, float] = Field(default_factory=dict)
|
||||
ccd: dict[str, Any] = Field(default_factory=dict, json_schema_extra=JSON_OBJECT_FIELD)
|
||||
dsm5_dimensional: dict[str, Any] = Field(default_factory=dict, json_schema_extra=JSON_OBJECT_FIELD)
|
||||
source_provenance: str = Field(default="", max_length=240)
|
||||
is_synthetic: bool = True
|
||||
submit_for_review: bool = False
|
||||
|
||||
|
||||
class PersonaDraftDetail(PersonaReviewSummary):
|
||||
demographics: dict[str, Any] = Field(json_schema_extra=JSON_OBJECT_FIELD)
|
||||
presenting: dict[str, Any] = Field(json_schema_extra=JSON_OBJECT_FIELD)
|
||||
history: dict[str, Any] = Field(json_schema_extra=JSON_OBJECT_FIELD)
|
||||
big5: dict[str, float]
|
||||
resistance: dict[str, float]
|
||||
speech_style: dict[str, Any] = Field(json_schema_extra=JSON_OBJECT_FIELD)
|
||||
affect_baseline: dict[str, float]
|
||||
ccd: dict[str, Any] = Field(json_schema_extra=JSON_OBJECT_FIELD)
|
||||
dsm5_dimensional: dict[str, Any] = Field(json_schema_extra=JSON_OBJECT_FIELD)
|
||||
|
||||
|
||||
def _first_text_value(data: dict[str, Any]) -> str:
|
||||
for value in data.values():
|
||||
if isinstance(value, str) and value.strip():
|
||||
|
|
@ -88,6 +125,49 @@ def _review_summary(entry: PersonaReviewItem) -> PersonaReviewSummary:
|
|||
)
|
||||
|
||||
|
||||
def _draft_detail(entry: PersonaDraftRecord) -> PersonaDraftDetail:
|
||||
card = entry.card
|
||||
return PersonaDraftDetail(
|
||||
**_review_summary(entry.review).model_dump(),
|
||||
demographics=card.demographics,
|
||||
presenting=card.presenting,
|
||||
history=card.history,
|
||||
big5=card.big5,
|
||||
resistance=card.resistance,
|
||||
speech_style=card.speech_style,
|
||||
affect_baseline=card.affect_baseline,
|
||||
ccd=card.ccd,
|
||||
dsm5_dimensional=card.dsm5_dimensional,
|
||||
)
|
||||
|
||||
|
||||
def _card_from_draft_payload(request: PersonaDraftPayload) -> PersonaCard:
|
||||
code = request.code.strip().upper()
|
||||
display_name = request.display_name.strip()
|
||||
if not code:
|
||||
raise HTTPException(status.HTTP_422_UNPROCESSABLE_ENTITY, detail="persona code is required")
|
||||
if not display_name:
|
||||
raise HTTPException(status.HTTP_422_UNPROCESSABLE_ENTITY, detail="display_name is required")
|
||||
theory_target = [value.strip().lower() for value in request.theory_target if value.strip()]
|
||||
return PersonaCard(
|
||||
code=code,
|
||||
display_name=display_name,
|
||||
difficulty=request.difficulty,
|
||||
theory_target=theory_target,
|
||||
demographics=request.demographics,
|
||||
presenting=request.presenting,
|
||||
history=request.history,
|
||||
big5=request.big5,
|
||||
resistance=request.resistance,
|
||||
speech_style=request.speech_style,
|
||||
affect_baseline=request.affect_baseline,
|
||||
ccd=request.ccd,
|
||||
dsm5_dimensional=request.dsm5_dimensional,
|
||||
source_provenance=request.source_provenance.strip(),
|
||||
is_synthetic=request.is_synthetic,
|
||||
)
|
||||
|
||||
|
||||
def _ensure_teacher_or_admin(principal: Principal) -> None:
|
||||
if principal.role not in {Role.TEACHER, Role.ADMIN}:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail="only teachers and admins can review personas")
|
||||
|
|
@ -127,6 +207,79 @@ async def list_persona_reviews(principal: TeacherOrAdmin) -> list[PersonaReviewS
|
|||
return [_review_summary(entry) for entry in queue]
|
||||
|
||||
|
||||
@router.post("/drafts", response_model=PersonaReviewSummary, status_code=status.HTTP_201_CREATED)
|
||||
async def create_persona_draft_route(
|
||||
request: PersonaDraftPayload,
|
||||
principal: TeacherOrAdmin,
|
||||
) -> PersonaReviewSummary:
|
||||
"""Create a draft persona card version for faculty review."""
|
||||
_ensure_teacher_or_admin(principal)
|
||||
try:
|
||||
created = await create_persona_draft(
|
||||
card=_card_from_draft_payload(request),
|
||||
author_id=principal.user_id,
|
||||
role=principal.role.value,
|
||||
submit_for_review=request.submit_for_review,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=str(exc)) from exc
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="persona draft database unavailable",
|
||||
) from exc
|
||||
return _review_summary(created)
|
||||
|
||||
|
||||
@router.get("/drafts/{persona_id}", response_model=PersonaDraftDetail)
|
||||
async def get_persona_draft_route(
|
||||
persona_id: str,
|
||||
principal: TeacherOrAdmin,
|
||||
) -> PersonaDraftDetail:
|
||||
"""Return a draft/review persona card for editing."""
|
||||
_ensure_teacher_or_admin(principal)
|
||||
try:
|
||||
record = await get_persona_draft_record(persona_id=persona_id, role=principal.role.value)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=str(exc)) from exc
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="persona draft database unavailable",
|
||||
) from exc
|
||||
if record is None:
|
||||
raise HTTPException(status.HTTP_404_NOT_FOUND, detail="persona draft not found")
|
||||
return _draft_detail(record)
|
||||
|
||||
|
||||
@router.put("/drafts/{persona_id}", response_model=PersonaReviewSummary)
|
||||
async def update_persona_draft_route(
|
||||
persona_id: str,
|
||||
request: PersonaDraftPayload,
|
||||
principal: TeacherOrAdmin,
|
||||
) -> PersonaReviewSummary:
|
||||
"""Update a draft/review persona card and optionally submit it for review."""
|
||||
_ensure_teacher_or_admin(principal)
|
||||
try:
|
||||
updated = await update_persona_draft(
|
||||
persona_id=persona_id,
|
||||
card=_card_from_draft_payload(request),
|
||||
author_id=principal.user_id,
|
||||
role=principal.role.value,
|
||||
submit_for_review=request.submit_for_review,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=str(exc)) from exc
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="persona draft update database unavailable",
|
||||
) from exc
|
||||
if updated is None:
|
||||
raise HTTPException(status.HTTP_404_NOT_FOUND, detail="persona draft not found or locked")
|
||||
return _review_summary(updated)
|
||||
|
||||
|
||||
@router.post("/review/{persona_id}", response_model=PersonaReviewSummary)
|
||||
async def decide_persona_review(
|
||||
persona_id: str,
|
||||
|
|
|
|||
|
|
@ -18,18 +18,19 @@ from fastapi import APIRouter, HTTPException, status
|
|||
from pydantic import BaseModel, Field
|
||||
from sse_starlette.sse import EventSourceResponse
|
||||
|
||||
from .. import db, session_persistence
|
||||
from .. import db, session_persistence, turn_runtime
|
||||
from ..config import settings
|
||||
from ..deps import CurrentPrincipal, Principal, Role
|
||||
from ..engine_client import EngineError, engine_client
|
||||
from ..persona_repository import get_catalog_persona
|
||||
from ..runtime_policy import require_runtime_fallback_allowed, runtime_fallback_allowed
|
||||
from ..runtime_policy import require_runtime_fallback_allowed
|
||||
from ..services import evaluator, memory, orchestrator, rag, state_machine
|
||||
from ..store import InProcSession, TurnRecord, store
|
||||
|
||||
router = APIRouter(prefix="/sessions", tags=["sessions"])
|
||||
|
||||
TheoryMode = Literal["humanistic", "cbt", "integrative"]
|
||||
EndStateValue = str | int | float | bool | None | dict[str, float]
|
||||
|
||||
|
||||
class SessionStartRequest(BaseModel):
|
||||
|
|
@ -51,6 +52,12 @@ class TurnRequest(BaseModel):
|
|||
text: str = Field(..., min_length=1)
|
||||
|
||||
|
||||
class CrisisResourceResponse(BaseModel):
|
||||
title: str
|
||||
number: str
|
||||
message: str
|
||||
|
||||
|
||||
class TurnResponse(BaseModel):
|
||||
turn_seq: int
|
||||
stage: str
|
||||
|
|
@ -58,13 +65,15 @@ class TurnResponse(BaseModel):
|
|||
client_reply: Optional[str] = None
|
||||
safety_flagged: bool = False
|
||||
crisis_kind: str = "none"
|
||||
crisis_resource: Optional[CrisisResourceResponse] = None
|
||||
conversation_stopped: bool = False
|
||||
|
||||
|
||||
class SessionEndResponse(BaseModel):
|
||||
session_id: str
|
||||
session_no: int
|
||||
digest_pending: bool
|
||||
end_state: dict
|
||||
end_state: dict[str, EndStateValue]
|
||||
|
||||
|
||||
class LearnerSessionSummary(BaseModel):
|
||||
|
|
@ -121,6 +130,12 @@ class ReviewTechnique(BaseModel):
|
|||
label: str
|
||||
|
||||
|
||||
class ReviewNonverbalEvent(BaseModel):
|
||||
kind: Literal["audio", "silence", "pace", "barge_in"]
|
||||
label: str
|
||||
detail: str
|
||||
|
||||
|
||||
class ReviewNote(BaseModel):
|
||||
author: str
|
||||
tone: str
|
||||
|
|
@ -136,6 +151,7 @@ class ReviewTurn(BaseModel):
|
|||
who: str
|
||||
text: str
|
||||
techniques: list[ReviewTechnique] = Field(default_factory=list)
|
||||
nonverbal: list[ReviewNonverbalEvent] = Field(default_factory=list)
|
||||
note: Optional[ReviewNote] = None
|
||||
|
||||
|
||||
|
|
@ -164,6 +180,34 @@ class ReviewPoint(BaseModel):
|
|||
jumpTo: Optional[str] = None
|
||||
|
||||
|
||||
class ReviewWorksheetEvidence(BaseModel):
|
||||
turnId: str
|
||||
speaker: Literal["learner", "client"]
|
||||
quote: str
|
||||
|
||||
|
||||
class ReviewWorksheetItem(BaseModel):
|
||||
key: str
|
||||
label: str
|
||||
value: Optional[str] = None
|
||||
evidence: list[ReviewWorksheetEvidence] = Field(default_factory=list)
|
||||
confidence: Literal["none", "low", "medium"] = "none"
|
||||
emptyReason: Optional[str] = None
|
||||
|
||||
|
||||
class ReviewWorksheetSection(BaseModel):
|
||||
key: str
|
||||
title: str
|
||||
items: list[ReviewWorksheetItem] = Field(default_factory=list)
|
||||
|
||||
|
||||
class ReviewCaseWorksheet(BaseModel):
|
||||
status: Literal["empty", "draft_from_transcript"] = "empty"
|
||||
generatedBy: str = "rule-based transcript extractor"
|
||||
sections: list[ReviewWorksheetSection] = Field(default_factory=list)
|
||||
limitations: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class SessionReviewResponse(BaseModel):
|
||||
session_id: str
|
||||
client: ReviewClient
|
||||
|
|
@ -184,6 +228,7 @@ class SessionReviewResponse(BaseModel):
|
|||
rubric: list[ReviewRubricRow] = Field(default_factory=list)
|
||||
goodMoments: list[ReviewPoint] = Field(default_factory=list)
|
||||
growthPoints: list[ReviewPoint] = Field(default_factory=list)
|
||||
caseWorksheet: ReviewCaseWorksheet = Field(default_factory=ReviewCaseWorksheet)
|
||||
nextLine: Optional[str] = None
|
||||
clientFeedback: Optional[str] = None
|
||||
audioUrl: Optional[str] = None
|
||||
|
|
@ -337,6 +382,27 @@ async def _build_start_recall(*, case_id: str, card) -> memory.RecallContext:
|
|||
)
|
||||
|
||||
|
||||
async def _build_seed_recall(*, case_id: str | None) -> memory.RecallContext:
|
||||
if not case_id:
|
||||
return memory.build_recall_context()
|
||||
try:
|
||||
db.get_pool()
|
||||
except RuntimeError:
|
||||
return memory.build_recall_context()
|
||||
prev_summary = await _load_prev_case_summary(case_id)
|
||||
pinned = list((prev_summary or {}).get("pinned_facts") or [])
|
||||
return memory.build_recall_context(prev_summary=prev_summary, pinned_facts=pinned)
|
||||
|
||||
|
||||
async def ensure_recall_context(sess: InProcSession) -> memory.RecallContext:
|
||||
cached = _RECALL_CACHE.get(sess.session_id)
|
||||
if cached is not None:
|
||||
return cached
|
||||
recall = await _build_seed_recall(case_id=sess.case_id)
|
||||
_RECALL_CACHE[sess.session_id] = recall
|
||||
return recall
|
||||
|
||||
|
||||
async def _warm_rag_caches(session_id: str, case_id: str, card) -> None:
|
||||
"""RAG 회상·KB 행동단서를 **백그라운드**로 산출해 캐시한다(요청 경로 비차단).
|
||||
|
||||
|
|
@ -362,13 +428,7 @@ _PHASE_KEY_BY_LABEL = {
|
|||
|
||||
|
||||
def _stage_label(stage: object) -> str:
|
||||
name = getattr(stage, "name", "")
|
||||
return {
|
||||
"RAPPORT": "라포",
|
||||
"EXPLORE": "탐색",
|
||||
"INTERVENE": "개입",
|
||||
"CLOSE": "정리",
|
||||
}.get(name, str(getattr(stage, "value", stage)))
|
||||
return turn_runtime.stage_label(stage)
|
||||
|
||||
|
||||
def _ensure_learner(principal: Principal) -> None:
|
||||
|
|
@ -383,54 +443,22 @@ async def _load_session_or_404(
|
|||
allow_ended: bool = False,
|
||||
include_turn_evaluation: bool = False,
|
||||
) -> InProcSession:
|
||||
sess = await session_persistence.load_session(
|
||||
sess, err = await turn_runtime.load_owned_session(
|
||||
session_id,
|
||||
principal,
|
||||
allow_ended=True,
|
||||
allow_ended=allow_ended,
|
||||
include_turn_evaluation=include_turn_evaluation,
|
||||
)
|
||||
if sess is not None:
|
||||
store.put(sess)
|
||||
elif runtime_fallback_allowed():
|
||||
sess = store.get(session_id)
|
||||
if sess is None:
|
||||
if err == turn_runtime.SessionAccessError.NOT_FOUND:
|
||||
raise HTTPException(status.HTTP_404_NOT_FOUND, detail="session not found")
|
||||
if sess.learner_id != principal.user_id:
|
||||
if err == turn_runtime.SessionAccessError.FORBIDDEN:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail="session does not belong to user")
|
||||
if sess.ended and not allow_ended:
|
||||
if err == turn_runtime.SessionAccessError.ENDED:
|
||||
raise HTTPException(status.HTTP_409_CONFLICT, detail="session already ended")
|
||||
assert sess is not None
|
||||
return sess
|
||||
|
||||
|
||||
async def _append_session_turn(sess: InProcSession, turn: TurnRecord) -> None:
|
||||
if await session_persistence.append_turn(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
turn=turn,
|
||||
):
|
||||
sess.turns.append(turn)
|
||||
store.put(sess)
|
||||
return
|
||||
require_runtime_fallback_allowed("session turn append")
|
||||
store.append_turn(sess.session_id, turn)
|
||||
|
||||
|
||||
async def _update_session_state(
|
||||
sess: InProcSession,
|
||||
state: state_machine.SessionState,
|
||||
) -> None:
|
||||
if await session_persistence.update_state(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
state=state,
|
||||
):
|
||||
sess.state = state
|
||||
store.put(sess)
|
||||
return
|
||||
require_runtime_fallback_allowed("session state update")
|
||||
store.update_state(sess.session_id, state)
|
||||
|
||||
|
||||
async def _end_persisted_session(sess: InProcSession, carry: memory.CarryOver) -> None:
|
||||
if await session_persistence.end_session(sess, carry):
|
||||
sess.ended = True
|
||||
|
|
@ -639,6 +667,178 @@ def _latest_client_feedback(turns: list[ReviewTurn]) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def _worksheet_evidence(turn: ReviewTurn) -> ReviewWorksheetEvidence:
|
||||
return ReviewWorksheetEvidence(
|
||||
turnId=turn.id,
|
||||
speaker=turn.speaker,
|
||||
quote=_clip_text(turn.text, 120),
|
||||
)
|
||||
|
||||
|
||||
def _worksheet_item(
|
||||
*,
|
||||
key: str,
|
||||
label: str,
|
||||
turns: list[ReviewTurn],
|
||||
keywords: list[str],
|
||||
preferred_speaker: Literal["learner", "client"] | None = None,
|
||||
fallback_turn: ReviewTurn | None = None,
|
||||
) -> ReviewWorksheetItem:
|
||||
lowered_keywords = [keyword.lower() for keyword in keywords if keyword]
|
||||
candidates = turns
|
||||
if preferred_speaker:
|
||||
preferred = [turn for turn in turns if turn.speaker == preferred_speaker]
|
||||
candidates = preferred + [turn for turn in turns if turn.speaker != preferred_speaker]
|
||||
|
||||
for turn in candidates:
|
||||
text = _compact_text(turn.text)
|
||||
lower_text = text.lower()
|
||||
if lowered_keywords and any(keyword in lower_text for keyword in lowered_keywords):
|
||||
return ReviewWorksheetItem(
|
||||
key=key,
|
||||
label=label,
|
||||
value=_clip_text(text, 140),
|
||||
evidence=[_worksheet_evidence(turn)],
|
||||
confidence="medium",
|
||||
)
|
||||
|
||||
if fallback_turn is not None:
|
||||
return ReviewWorksheetItem(
|
||||
key=key,
|
||||
label=label,
|
||||
value=_clip_text(fallback_turn.text, 140),
|
||||
evidence=[_worksheet_evidence(fallback_turn)],
|
||||
confidence="low",
|
||||
)
|
||||
|
||||
return ReviewWorksheetItem(
|
||||
key=key,
|
||||
label=label,
|
||||
value=None,
|
||||
evidence=[],
|
||||
confidence="none",
|
||||
emptyReason="저장된 축어록에서 명시 근거를 찾지 못했습니다.",
|
||||
)
|
||||
|
||||
|
||||
def _worksheet_section(
|
||||
key: str,
|
||||
title: str,
|
||||
specs: list[tuple[str, str, list[str], Literal["learner", "client"] | None]],
|
||||
turns: list[ReviewTurn],
|
||||
fallback_client: ReviewTurn | None,
|
||||
fallback_learner: ReviewTurn | None,
|
||||
) -> ReviewWorksheetSection:
|
||||
items: list[ReviewWorksheetItem] = []
|
||||
for item_key, label, keywords, speaker in specs:
|
||||
fallback = fallback_client if speaker == "client" else fallback_learner if speaker == "learner" else None
|
||||
items.append(
|
||||
_worksheet_item(
|
||||
key=item_key,
|
||||
label=label,
|
||||
turns=turns,
|
||||
keywords=keywords,
|
||||
preferred_speaker=speaker,
|
||||
fallback_turn=fallback if item_key in {"presenting_complaint", "first_goal"} else None,
|
||||
)
|
||||
)
|
||||
return ReviewWorksheetSection(key=key, title=title, items=items)
|
||||
|
||||
|
||||
def _case_worksheet_from_turns(turns: list[ReviewTurn]) -> ReviewCaseWorksheet:
|
||||
if not turns:
|
||||
return ReviewCaseWorksheet(
|
||||
status="empty",
|
||||
sections=[],
|
||||
limitations=["저장된 축어록이 없어 사례개념화 워크시트를 생성하지 않았습니다."],
|
||||
)
|
||||
|
||||
fallback_client = next((turn for turn in turns if turn.speaker == "client"), None)
|
||||
fallback_learner = next((turn for turn in turns if turn.speaker == "learner"), None)
|
||||
section_specs: list[
|
||||
tuple[str, str, list[tuple[str, str, list[str], Literal["learner", "client"] | None]]]
|
||||
] = [
|
||||
(
|
||||
"exploration_11",
|
||||
"탐색 11항목",
|
||||
[
|
||||
("presenting_complaint", "주호소", ["힘들", "문제", "걱정", "불안", "우울", "스트레스", "관계"], "client"),
|
||||
("trigger_context", "계기·상황", ["언제", "상황", "최근", "계기", "때"], "client"),
|
||||
("emotion", "정서", ["불안", "우울", "화", "슬프", "답답", "무섭", "외롭", "걱정"], "client"),
|
||||
("cognition", "생각", ["생각", "느낌", "해야", "못", "실패", "의미"], "client"),
|
||||
("behavior", "행동", ["피하", "잠", "먹", "울", "말", "연락", "공부", "멈"], "client"),
|
||||
("body", "신체·수면", ["잠", "식욕", "몸", "두통", "심장", "숨", "피곤"], "client"),
|
||||
("relationship", "관계", ["친구", "가족", "부모", "엄마", "아빠", "교수", "사람", "관계"], "client"),
|
||||
("resources", "자원", ["도움", "지지", "친구", "상담", "선생님", "가족"], "client"),
|
||||
("risk", "위험 신호", ["죽", "자살", "해치", "사라지고", "끝내", "위험"], "client"),
|
||||
("motivation", "변화동기", ["원", "바라", "변화", "해보고", "싶"], None),
|
||||
("first_goal", "상담 목표 초안", ["목표", "계획", "다음", "해볼", "원하"], "learner"),
|
||||
],
|
||||
),
|
||||
(
|
||||
"five_domains",
|
||||
"호소 5영역",
|
||||
[
|
||||
("domain_emotion", "정서", ["불안", "우울", "화", "슬프", "답답", "외롭"], "client"),
|
||||
("domain_cognition", "인지", ["생각", "걱정", "실패", "못", "의미"], "client"),
|
||||
("domain_behavior", "행동", ["피하", "연락", "공부", "잠", "멈"], "client"),
|
||||
("domain_relationship", "대인관계", ["친구", "가족", "사람", "관계", "부모"], "client"),
|
||||
("domain_body", "신체", ["잠", "식욕", "몸", "두통", "피곤", "숨"], "client"),
|
||||
],
|
||||
),
|
||||
(
|
||||
"cognitive_triad_emotions",
|
||||
"인지삼제·1/2차 감정",
|
||||
[
|
||||
("triad_self", "자기", ["나는", "내가", "나 자신", "스스로"], "client"),
|
||||
("triad_world", "타인·세계", ["사람", "세상", "학교", "가족", "친구"], "client"),
|
||||
("triad_future", "미래", ["앞으로", "미래", "계속", "나중"], "client"),
|
||||
("primary_emotion", "1차 감정", ["불안", "슬프", "무섭", "외롭", "걱정"], "client"),
|
||||
("secondary_emotion", "2차 감정", ["화", "짜증", "수치", "죄책", "부끄"], "client"),
|
||||
],
|
||||
),
|
||||
(
|
||||
"protective_barrier_quadrants",
|
||||
"보호·방해 4사분면",
|
||||
[
|
||||
("internal_protective", "내적 보호요인", ["해보고", "버텼", "노력", "원", "견뎠"], None),
|
||||
("internal_barrier", "내적 방해요인", ["못", "두려", "불안", "회피", "걱정"], "client"),
|
||||
("external_protective", "외적 보호요인", ["친구", "가족", "상담", "교수", "도움"], "client"),
|
||||
("external_barrier", "외적 방해요인", ["갈등", "압박", "비난", "스트레스", "혼자"], "client"),
|
||||
],
|
||||
),
|
||||
(
|
||||
"biopsychosocial_goals",
|
||||
"생물·심리·사회 목표",
|
||||
[
|
||||
("bio_goal", "생물", ["잠", "식사", "운동", "몸", "피곤"], "client"),
|
||||
("psy_goal", "심리", ["생각", "감정", "불안", "연습", "조절"], None),
|
||||
("social_goal", "사회", ["관계", "대화", "연락", "도움", "친구"], None),
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
sections = [
|
||||
_worksheet_section(
|
||||
key,
|
||||
title,
|
||||
specs,
|
||||
turns,
|
||||
fallback_client,
|
||||
fallback_learner,
|
||||
)
|
||||
for key, title, specs in section_specs
|
||||
]
|
||||
return ReviewCaseWorksheet(
|
||||
status="draft_from_transcript",
|
||||
sections=sections,
|
||||
limitations=[
|
||||
"저장된 축어록에서 키워드 근거를 추출한 1차 초안입니다.",
|
||||
"임상팀 루브릭, 교수자 검수, 학습자 수정 입력 전에는 확정 사례개념화로 보지 않습니다.",
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _evaluation_payload(record: dict[str, object] | None) -> dict[str, object]:
|
||||
if not record:
|
||||
return {}
|
||||
|
|
@ -697,35 +897,88 @@ def _review_note_from_turn_eval(ev: dict[str, object] | None) -> Optional[Review
|
|||
return None
|
||||
|
||||
|
||||
async def _record_safety_event(sess: InProcSession, ctx, result) -> None:
|
||||
"""위기 escalate 시 app.safety_events 적재(교수자 감사·알림 레코드). C2.
|
||||
def _seconds_label(milliseconds: int) -> str:
|
||||
seconds = max(0, milliseconds) / 1000.0
|
||||
if seconds >= 10:
|
||||
return f"{seconds:.0f}초"
|
||||
return f"{seconds:.1f}초"
|
||||
|
||||
비차단: DB 미가용(degraded)·FK 미충족(in-memory 세션) 시 graceful skip — 상담 루프를
|
||||
절대 막지 않는다. 실시간 교수자 push 알림은 후속(이 레코드가 1차 알림원).
|
||||
"""
|
||||
crisis = getattr(ctx, "crisis", None)
|
||||
if crisis is None or not getattr(crisis, "escalate", False):
|
||||
return
|
||||
kind = getattr(crisis.kind, "value", None) or str(getattr(crisis, "kind", "crisis"))
|
||||
try:
|
||||
async with db.acquire() as conn:
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO app.safety_events
|
||||
(session_id, trigger_type, ko_risk_level, escalated, detail)
|
||||
VALUES ($1::uuid, $2, $3, TRUE, $4::jsonb)
|
||||
""",
|
||||
sess.session_id,
|
||||
kind,
|
||||
int(getattr(crisis, "risk_level", 0) or 0),
|
||||
json.dumps({
|
||||
"matched": list(getattr(crisis, "matched", []) or []),
|
||||
"stage": getattr(result, "stage", None),
|
||||
"turn_seq": getattr(result, "turn_seq", None),
|
||||
}),
|
||||
|
||||
def _review_nonverbal_events(turn: TurnRecord) -> list[ReviewNonverbalEvent]:
|
||||
events: list[ReviewNonverbalEvent] = []
|
||||
if turn.silence_ms is not None and turn.silence_ms >= 1000:
|
||||
events.append(
|
||||
ReviewNonverbalEvent(
|
||||
kind="silence",
|
||||
label="침묵",
|
||||
detail=_seconds_label(turn.silence_ms),
|
||||
)
|
||||
)
|
||||
if turn.speech_rate is not None:
|
||||
events.append(
|
||||
ReviewNonverbalEvent(
|
||||
kind="pace",
|
||||
label="발화 속도",
|
||||
detail=f"분당 {turn.speech_rate:.0f}자",
|
||||
)
|
||||
)
|
||||
if turn.barge_in is True:
|
||||
events.append(
|
||||
ReviewNonverbalEvent(
|
||||
kind="barge_in",
|
||||
label="끼어듦",
|
||||
detail="내담자 발화 중 시작",
|
||||
)
|
||||
)
|
||||
if turn.audio_ref:
|
||||
events.append(
|
||||
ReviewNonverbalEvent(
|
||||
kind="audio",
|
||||
label="음성 입력",
|
||||
detail="음성으로 기록됨",
|
||||
)
|
||||
)
|
||||
return events
|
||||
|
||||
|
||||
async def _evaluate_stream_turn(ctx: orchestrator.TurnContext, final_reply: str) -> Optional[dict]:
|
||||
"""stream 경로 완료 후 fast-loop 평가를 계산한다. 실패는 턴 저장을 막지 않는다."""
|
||||
if not final_reply:
|
||||
return None
|
||||
try:
|
||||
hook = evaluator.make_eval_hook(
|
||||
engine_client,
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
)
|
||||
return await hook(ctx, final_reply)
|
||||
except Exception:
|
||||
pass # 비차단(R5): 적재 실패가 위기 대응/상담을 막지 않음.
|
||||
return None
|
||||
|
||||
|
||||
def _stream_result_from_done(
|
||||
ctx: orchestrator.TurnContext,
|
||||
final_reply: str,
|
||||
data: dict[str, object],
|
||||
evaluation: Optional[dict],
|
||||
) -> orchestrator.TurnResult:
|
||||
assert ctx.state_after is not None
|
||||
return orchestrator.TurnResult(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
stage=_stage_label(ctx.state_after.stage),
|
||||
effective_openness=ctx.state_after.effective_openness,
|
||||
client_reply=final_reply or None,
|
||||
safety_flagged=bool(data.get("safety_flagged")),
|
||||
state_after=ctx.state_after,
|
||||
evaluation=evaluation,
|
||||
crisis_kind=ctx.crisis.kind.value if ctx.crisis else "none",
|
||||
crisis_resource=data.get("crisis_resource") if isinstance(data.get("crisis_resource"), dict) else None,
|
||||
conversation_stopped=bool(data.get("conversation_stopped")),
|
||||
llm_provider=str(data.get("llm_provider") or "") or None,
|
||||
model=str(data.get("model") or "") or None,
|
||||
tokens_in=int(data.get("tokens_in") or 0),
|
||||
tokens_out=int(data.get("tokens_out") or 0),
|
||||
cost_usd=float(data.get("cost_usd") or 0.0),
|
||||
)
|
||||
|
||||
|
||||
def _learner_visible_turns(sess: InProcSession) -> list[TurnRecord]:
|
||||
|
|
@ -752,6 +1005,7 @@ async def _generate_and_save_session_evaluation(sess: InProcSession) -> None:
|
|||
technique_codes=[],
|
||||
theory_mode=sess.theory_mode,
|
||||
scope="session_end",
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
),
|
||||
timeout=min(float(settings.engine_timeout), 45.0),
|
||||
)
|
||||
|
|
@ -907,7 +1161,12 @@ async def start_session(
|
|||
raise HTTPException(status.HTTP_404_NOT_FOUND, detail=f"unknown persona {body.persona_code}")
|
||||
card = catalog_persona.card
|
||||
|
||||
recall = memory.build_recall_context()
|
||||
case_context = await session_persistence.get_case_context(
|
||||
learner_id=principal.user_id,
|
||||
persona_id=catalog_persona.persona_id,
|
||||
)
|
||||
recall = await _build_seed_recall(case_id=case_context.case_id if case_context else None)
|
||||
session_no = (case_context.last_session_no + 1) if case_context else 1
|
||||
st = state_machine.init_state(
|
||||
params=card.openness_params(),
|
||||
carry=recall.carry,
|
||||
|
|
@ -919,10 +1178,11 @@ async def start_session(
|
|||
card=card,
|
||||
theory_mode=body.theory_mode,
|
||||
state=st,
|
||||
session_no=1,
|
||||
session_no=session_no,
|
||||
carry_rapport=carry_rapport,
|
||||
persona_id=catalog_persona.persona_id,
|
||||
persona_version=catalog_persona.version,
|
||||
case_id=case_context.case_id if case_context else None,
|
||||
)
|
||||
degraded = catalog_persona.degraded or sess is None
|
||||
if sess is None:
|
||||
|
|
@ -932,7 +1192,7 @@ async def start_session(
|
|||
persona=card,
|
||||
theory_mode=body.theory_mode,
|
||||
state=st,
|
||||
session_no=1,
|
||||
session_no=session_no,
|
||||
carry_rapport=carry_rapport,
|
||||
)
|
||||
else:
|
||||
|
|
@ -1007,6 +1267,7 @@ async def get_session_review(
|
|||
who="학습자" if speaker == "learner" else client_name,
|
||||
text=turn.text_masked,
|
||||
techniques=_review_techniques_from_turn_eval(turn_eval),
|
||||
nonverbal=_review_nonverbal_events(turn) if speaker == "learner" else [],
|
||||
note=_review_note_from_turn_eval(turn_eval),
|
||||
)
|
||||
)
|
||||
|
|
@ -1085,6 +1346,7 @@ async def get_session_review(
|
|||
rubric=rubric,
|
||||
goodMoments=good_moments,
|
||||
growthPoints=growth_points,
|
||||
caseWorksheet=_case_worksheet_from_turns(turns),
|
||||
nextLine=next_line,
|
||||
clientFeedback=client_feedback,
|
||||
audioUrl=None,
|
||||
|
|
@ -1103,7 +1365,7 @@ async def submit_turn(
|
|||
"""Submit one trainee utterance and return the generated client reply."""
|
||||
_ensure_learner(principal)
|
||||
sess = await _load_session_or_404(session_id, principal)
|
||||
recall = _RECALL_CACHE.get(session_id) or memory.RecallContext()
|
||||
recall = await ensure_recall_context(sess)
|
||||
kb_cues = _KB_CUES_CACHE.get(session_id) or [] # 비차단: warm 전이면 빈 단서(graceful)
|
||||
|
||||
ctx = orchestrator.prepare_turn(
|
||||
|
|
@ -1124,7 +1386,11 @@ async def submit_turn(
|
|||
result = await orchestrator.run_turn_generate(
|
||||
ctx,
|
||||
engine_client,
|
||||
eval_hook=evaluator.make_eval_hook(engine_client),
|
||||
eval_hook=evaluator.make_eval_hook(
|
||||
engine_client,
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
),
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
)
|
||||
except EngineError as exc:
|
||||
raise HTTPException(
|
||||
|
|
@ -1132,37 +1398,13 @@ async def submit_turn(
|
|||
detail=f"engine unavailable: {exc}",
|
||||
) from exc
|
||||
|
||||
# 턴별 fast-loop 평가는 학습자(상담자) 발화에 부착(기법 태깅·적절성·의도이탈).
|
||||
await _append_session_turn(
|
||||
await turn_runtime.record_completed_turn(
|
||||
sess,
|
||||
TurnRecord(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
speaker="counselor",
|
||||
stage=_stage_label(ctx.state_after.stage),
|
||||
text=body.text,
|
||||
text_masked=ctx.learner_text_masked,
|
||||
evaluation=result.evaluation,
|
||||
),
|
||||
ctx,
|
||||
result,
|
||||
context_prefix="session",
|
||||
)
|
||||
|
||||
if result.client_reply:
|
||||
await _append_session_turn(
|
||||
sess,
|
||||
TurnRecord(
|
||||
turn_seq=result.turn_seq,
|
||||
speaker="client",
|
||||
stage=_stage_label(result.state_after.stage),
|
||||
text=result.client_reply,
|
||||
text_masked=result.client_reply,
|
||||
llm_provider=result.llm_provider,
|
||||
model=result.model,
|
||||
tokens_in=result.tokens_in,
|
||||
tokens_out=result.tokens_out,
|
||||
cost_usd=result.cost_usd,
|
||||
),
|
||||
)
|
||||
await _update_session_state(sess, result.state_after)
|
||||
await _record_safety_event(sess, ctx, result) # C2: 위기 escalate 시 safety_events 적재(비차단)
|
||||
await turn_runtime.record_safety_event(sess, ctx, result)
|
||||
|
||||
return TurnResponse(
|
||||
turn_seq=result.turn_seq,
|
||||
|
|
@ -1171,6 +1413,8 @@ async def submit_turn(
|
|||
client_reply=result.client_reply,
|
||||
safety_flagged=result.safety_flagged,
|
||||
crisis_kind=result.crisis_kind,
|
||||
crisis_resource=result.crisis_resource,
|
||||
conversation_stopped=result.conversation_stopped,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1183,7 +1427,7 @@ async def stream_turn(
|
|||
"""Stream a generated client reply for one trainee utterance."""
|
||||
_ensure_learner(principal)
|
||||
sess = await _load_session_or_404(session_id, principal)
|
||||
recall = _RECALL_CACHE.get(session_id) or memory.RecallContext()
|
||||
recall = await ensure_recall_context(sess)
|
||||
kb_cues = _KB_CUES_CACHE.get(session_id) or [] # 비차단: warm 전이면 빈 단서(graceful)
|
||||
|
||||
ctx = orchestrator.prepare_turn(
|
||||
|
|
@ -1204,40 +1448,26 @@ async def stream_turn(
|
|||
last_beat = asyncio.get_running_loop().time()
|
||||
final_reply = ""
|
||||
try:
|
||||
async for ev in orchestrator.run_turn_stream(ctx, engine_client):
|
||||
async for ev in orchestrator.run_turn_stream(
|
||||
ctx,
|
||||
engine_client,
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
):
|
||||
if ev.event == "token":
|
||||
text = str(ev.data.get("text", ""))
|
||||
final_reply += text
|
||||
yield {"event": "token", "data": text}
|
||||
elif ev.event == "done":
|
||||
data = {**ev.data, "stage": _stage_label(ctx.state_after.stage)}
|
||||
await _append_session_turn(
|
||||
evaluation = await _evaluate_stream_turn(ctx, final_reply)
|
||||
result = _stream_result_from_done(ctx, final_reply, data, evaluation)
|
||||
await turn_runtime.record_completed_turn(
|
||||
sess,
|
||||
TurnRecord(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
speaker="counselor",
|
||||
stage=_stage_label(ctx.state_after.stage),
|
||||
text=body.text,
|
||||
text_masked=ctx.learner_text_masked,
|
||||
),
|
||||
ctx,
|
||||
result,
|
||||
context_prefix="session",
|
||||
)
|
||||
await _update_session_state(sess, ctx.state_after)
|
||||
if final_reply:
|
||||
await _append_session_turn(
|
||||
sess,
|
||||
TurnRecord(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
speaker="client",
|
||||
stage=_stage_label(ctx.state_after.stage),
|
||||
text=final_reply,
|
||||
text_masked=final_reply,
|
||||
llm_provider=str(ev.data.get("llm_provider") or ""),
|
||||
model=str(ev.data.get("model") or ""),
|
||||
tokens_in=int(ev.data.get("tokens_in") or 0),
|
||||
tokens_out=int(ev.data.get("tokens_out") or 0),
|
||||
cost_usd=float(ev.data.get("cost_usd") or 0.0),
|
||||
),
|
||||
)
|
||||
await turn_runtime.record_safety_event(sess, ctx, result)
|
||||
yield {"event": "done", "data": json.dumps(data, ensure_ascii=False)}
|
||||
else:
|
||||
yield {"event": ev.event, "data": json.dumps(ev.data, ensure_ascii=False)}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Annotated
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from pydantic import BaseModel, Field
|
||||
|
|
@ -34,17 +34,70 @@ class TeacherSessionSummary(BaseModel):
|
|||
ended_at: str | None = None
|
||||
|
||||
|
||||
class TeacherGrowthPoint(BaseModel):
|
||||
session_id: str
|
||||
session_no: int
|
||||
persona_code: str
|
||||
stage: str
|
||||
started_at: str
|
||||
ended_at: str | None = None
|
||||
score: float | None = None
|
||||
rapport: float | None = None
|
||||
technique_count: int = 0
|
||||
watch_count: int = 0
|
||||
|
||||
|
||||
class TeacherLearnerGrowth(BaseModel):
|
||||
learner_id: str
|
||||
learner_label: str
|
||||
sessions: int
|
||||
ended_sessions: int
|
||||
latest_at: str
|
||||
first_score: float | None = None
|
||||
latest_score: float | None = None
|
||||
score_delta: float | None = None
|
||||
avg_score: float | None = None
|
||||
avg_rapport: float | None = None
|
||||
trend: str = "insufficient"
|
||||
top_techniques: list[str] = Field(default_factory=list)
|
||||
points: list[TeacherGrowthPoint] = Field(default_factory=list)
|
||||
|
||||
|
||||
class TeacherSafetyAlert(BaseModel):
|
||||
id: str
|
||||
session_id: str
|
||||
learner_id: str
|
||||
learner_label: str
|
||||
persona_code: str
|
||||
session_no: int
|
||||
trigger_type: str
|
||||
ko_risk_level: int
|
||||
escalated: bool
|
||||
created_at: str
|
||||
resource_title: str = "자살예방상담전화 109"
|
||||
resource_number: str = "109"
|
||||
|
||||
|
||||
class TeacherDashboardResponse(BaseModel):
|
||||
source: str = "in_memory"
|
||||
cohort_label: str = "현재 학습 기록"
|
||||
total_learners: int
|
||||
active_sessions: int
|
||||
ended_sessions: int
|
||||
safety_alerts: list[TeacherSafetyAlert] = Field(default_factory=list)
|
||||
learner_growth: list[TeacherLearnerGrowth] = Field(default_factory=list)
|
||||
pending_reviews: list[TeacherSessionSummary] = Field(default_factory=list)
|
||||
recent_sessions: list[TeacherSessionSummary] = Field(default_factory=list)
|
||||
message: str
|
||||
|
||||
|
||||
_APPROPRIATENESS_SCORE = {
|
||||
"neg": 0.0,
|
||||
"neutral": 0.5,
|
||||
"pos": 1.0,
|
||||
}
|
||||
|
||||
|
||||
def _iso(ts: float | None) -> str | None:
|
||||
if ts is None:
|
||||
return None
|
||||
|
|
@ -56,6 +109,149 @@ def _learner_label(learner_id: str) -> str:
|
|||
return f"학습자 {suffix}"
|
||||
|
||||
|
||||
def _safe_float(value: object) -> float | None:
|
||||
try:
|
||||
return float(value) # type: ignore[arg-type]
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _avg(values: list[float]) -> float | None:
|
||||
if not values:
|
||||
return None
|
||||
return round(sum(values) / len(values), 3)
|
||||
|
||||
|
||||
def _turn_eval(turn: Any) -> dict[str, Any] | None:
|
||||
ev = getattr(turn, "evaluation", None)
|
||||
return ev if isinstance(ev, dict) else None
|
||||
|
||||
|
||||
def _turn_score(ev: dict[str, Any]) -> float | None:
|
||||
raw = str(ev.get("appropriateness") or "").strip().lower()
|
||||
return _APPROPRIATENESS_SCORE.get(raw)
|
||||
|
||||
|
||||
def _turn_rapport(ev: dict[str, Any]) -> float | None:
|
||||
value = _safe_float(ev.get("rapport_signal"))
|
||||
if value is None:
|
||||
return None
|
||||
return max(-1.0, min(1.0, value))
|
||||
|
||||
|
||||
def _turn_techniques(ev: dict[str, Any]) -> list[str]:
|
||||
raw = ev.get("techniques")
|
||||
if not isinstance(raw, list):
|
||||
return []
|
||||
labels: list[str] = []
|
||||
for item in raw:
|
||||
if isinstance(item, dict):
|
||||
label = item.get("label") or item.get("name") or item.get("id")
|
||||
else:
|
||||
label = item
|
||||
if label:
|
||||
labels.append(str(label))
|
||||
return labels
|
||||
|
||||
|
||||
def _session_growth_point(sess: InProcSession) -> TeacherGrowthPoint:
|
||||
scores: list[float] = []
|
||||
rapports: list[float] = []
|
||||
technique_count = 0
|
||||
watch_count = 0
|
||||
for turn in sess.turns:
|
||||
if turn.speaker != "counselor":
|
||||
continue
|
||||
ev = _turn_eval(turn)
|
||||
if ev is None:
|
||||
continue
|
||||
score = _turn_score(ev)
|
||||
if score is not None:
|
||||
scores.append(score)
|
||||
if score < 1.0:
|
||||
watch_count += 1
|
||||
rapport = _turn_rapport(ev)
|
||||
if rapport is not None:
|
||||
rapports.append(rapport)
|
||||
technique_count += len(_turn_techniques(ev))
|
||||
return TeacherGrowthPoint(
|
||||
session_id=sess.session_id,
|
||||
session_no=sess.session_no,
|
||||
persona_code=sess.persona_code,
|
||||
stage=sess.state.stage.value,
|
||||
started_at=_iso(sess.created_at) or "",
|
||||
ended_at=_iso(sess.ended_at),
|
||||
score=_avg(scores),
|
||||
rapport=_avg(rapports),
|
||||
technique_count=technique_count,
|
||||
watch_count=watch_count,
|
||||
)
|
||||
|
||||
|
||||
def _build_learner_growth(sessions: list[InProcSession]) -> list[TeacherLearnerGrowth]:
|
||||
grouped: dict[str, list[InProcSession]] = {}
|
||||
for sess in sessions:
|
||||
grouped.setdefault(sess.learner_id, []).append(sess)
|
||||
|
||||
result: list[TeacherLearnerGrowth] = []
|
||||
for learner_id, learner_sessions in grouped.items():
|
||||
ordered = sorted(learner_sessions, key=lambda sess: sess.created_at)
|
||||
points = [_session_growth_point(sess) for sess in ordered]
|
||||
scored = [point for point in points if point.score is not None]
|
||||
rapport_values = [point.rapport for point in points if point.rapport is not None]
|
||||
technique_counts: dict[str, int] = {}
|
||||
for sess in ordered:
|
||||
for turn in sess.turns:
|
||||
if turn.speaker != "counselor":
|
||||
continue
|
||||
ev = _turn_eval(turn)
|
||||
if ev is None:
|
||||
continue
|
||||
for label in _turn_techniques(ev):
|
||||
technique_counts[label] = technique_counts.get(label, 0) + 1
|
||||
|
||||
first_score = scored[0].score if scored else None
|
||||
latest_score = scored[-1].score if scored else None
|
||||
score_delta: float | None = None
|
||||
trend = "insufficient"
|
||||
if first_score is not None and latest_score is not None:
|
||||
score_delta = round(latest_score - first_score, 3)
|
||||
if len(scored) >= 2:
|
||||
if score_delta >= 0.1:
|
||||
trend = "up"
|
||||
elif score_delta <= -0.1:
|
||||
trend = "down"
|
||||
else:
|
||||
trend = "flat"
|
||||
|
||||
latest_session = ordered[-1]
|
||||
top_techniques = [
|
||||
label
|
||||
for label, _count in sorted(
|
||||
technique_counts.items(),
|
||||
key=lambda item: (-item[1], item[0]),
|
||||
)[:3]
|
||||
]
|
||||
result.append(
|
||||
TeacherLearnerGrowth(
|
||||
learner_id=learner_id,
|
||||
learner_label=_learner_label(learner_id),
|
||||
sessions=len(ordered),
|
||||
ended_sessions=sum(1 for sess in ordered if sess.ended),
|
||||
latest_at=_iso(latest_session.ended_at or latest_session.created_at) or "",
|
||||
first_score=first_score,
|
||||
latest_score=latest_score,
|
||||
score_delta=score_delta,
|
||||
avg_score=_avg([point.score for point in scored if point.score is not None]),
|
||||
avg_rapport=_avg([value for value in rapport_values if value is not None]),
|
||||
trend=trend,
|
||||
top_techniques=top_techniques,
|
||||
points=points[-6:],
|
||||
)
|
||||
)
|
||||
return sorted(result, key=lambda item: item.latest_at, reverse=True)[:12]
|
||||
|
||||
|
||||
def _summary(sess: InProcSession) -> TeacherSessionSummary:
|
||||
learner_turns = sum(1 for turn in sess.turns if turn.speaker == "counselor")
|
||||
client_turns = sum(1 for turn in sess.turns if turn.speaker == "client")
|
||||
|
|
@ -79,13 +275,45 @@ def _summary(sess: InProcSession) -> TeacherSessionSummary:
|
|||
@router.get("/dashboard", response_model=TeacherDashboardResponse)
|
||||
async def teacher_dashboard(principal: TeacherPrincipal) -> TeacherDashboardResponse:
|
||||
"""Return teacher-visible dashboard data from real sessions only."""
|
||||
sessions, durable = await session_persistence.list_sessions(principal)
|
||||
sessions, durable = await session_persistence.list_sessions(
|
||||
principal,
|
||||
include_turn_evaluation=True,
|
||||
)
|
||||
if not durable:
|
||||
require_runtime_fallback_allowed("teacher dashboard")
|
||||
sessions = sorted(store.list(), key=lambda sess: sess.created_at, reverse=True)
|
||||
summaries = [_summary(sess) for sess in sessions]
|
||||
pending_reviews = [item for item in summaries if item.status == "ended"]
|
||||
learners = {sess.learner_id for sess in sessions}
|
||||
learner_growth = _build_learner_growth(sessions)
|
||||
safety_alerts: list[TeacherSafetyAlert] = []
|
||||
if durable:
|
||||
raw_alerts, alerts_durable = await session_persistence.list_safety_alerts(principal)
|
||||
if alerts_durable:
|
||||
safety_alerts = [
|
||||
TeacherSafetyAlert(
|
||||
id=str(item.get("id") or ""),
|
||||
session_id=str(item.get("session_id") or ""),
|
||||
learner_id=str(item.get("learner_id") or ""),
|
||||
learner_label=str(item.get("learner_label") or "학습자"),
|
||||
persona_code=str(item.get("persona_code") or ""),
|
||||
session_no=int(item.get("session_no") or 0),
|
||||
trigger_type=str(item.get("trigger_type") or "crisis"),
|
||||
ko_risk_level=int(item.get("ko_risk_level") or 0),
|
||||
escalated=bool(item.get("escalated")),
|
||||
created_at=str(item.get("created_at") or ""),
|
||||
resource_title=str(
|
||||
(item.get("detail") or {}).get("crisis_resource", {}).get(
|
||||
"title",
|
||||
"자살예방상담전화 109",
|
||||
)
|
||||
),
|
||||
resource_number=str(
|
||||
(item.get("detail") or {}).get("crisis_resource", {}).get("number", "109")
|
||||
),
|
||||
)
|
||||
for item in raw_alerts
|
||||
]
|
||||
|
||||
if sessions:
|
||||
message = "현재 기록된 실제 학습 세션만 표시합니다."
|
||||
|
|
@ -97,6 +325,8 @@ async def teacher_dashboard(principal: TeacherPrincipal) -> TeacherDashboardResp
|
|||
total_learners=len(learners),
|
||||
active_sessions=sum(1 for sess in sessions if not sess.ended),
|
||||
ended_sessions=sum(1 for sess in sessions if sess.ended),
|
||||
safety_alerts=safety_alerts,
|
||||
learner_growth=learner_growth,
|
||||
pending_reviews=pending_reviews[:20],
|
||||
recent_sessions=summaries[:20],
|
||||
message=message,
|
||||
|
|
|
|||
|
|
@ -22,14 +22,14 @@ from fastapi import APIRouter, WebSocket, WebSocketDisconnect
|
|||
from fastapi.responses import JSONResponse
|
||||
from starlette.websockets import WebSocketState
|
||||
|
||||
from .. import session_persistence
|
||||
from .. import session_persistence, turn_runtime
|
||||
from ..auth_sessions import get_session
|
||||
from ..config import settings
|
||||
from ..deps import Principal, Role
|
||||
from ..engine_client import EngineError, engine_client
|
||||
from ..persona_repository import get_catalog_persona
|
||||
from ..runtime_policy import require_runtime_fallback_allowed, runtime_fallback_allowed
|
||||
from ..services import evaluator, memory, orchestrator, state_machine
|
||||
from ..runtime_policy import require_runtime_fallback_allowed
|
||||
from ..services import evaluator, orchestrator, state_machine
|
||||
from ..services import voice as voice_svc
|
||||
from ..services.voice import VoicePreset, VoiceUnavailable, resolve_voice, voice_service
|
||||
from ..store import InProcSession, TurnRecord, store
|
||||
|
|
@ -292,7 +292,10 @@ async def _run_turn_and_speak(
|
|||
await _safe_send_json(websocket, {"type": "state", "state": "idle"})
|
||||
return
|
||||
|
||||
recall = memory.RecallContext()
|
||||
from . import sessions as session_routes
|
||||
|
||||
recall = await session_routes.ensure_recall_context(sess)
|
||||
kb_cues = session_routes._KB_CUES_CACHE.get(session_id) or []
|
||||
ctx = orchestrator.prepare_turn(
|
||||
session_id=session_id,
|
||||
case_id=sess.case_id,
|
||||
|
|
@ -302,6 +305,7 @@ async def _run_turn_and_speak(
|
|||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(visible_to="client"),
|
||||
kb_behavior_cues=kb_cues,
|
||||
theory_mode=sess.theory_mode,
|
||||
)
|
||||
assert ctx.state_after is not None
|
||||
|
|
@ -311,7 +315,11 @@ async def _run_turn_and_speak(
|
|||
result = await orchestrator.run_turn_generate(
|
||||
ctx,
|
||||
engine_client,
|
||||
eval_hook=evaluator.make_eval_hook(engine_client),
|
||||
eval_hook=evaluator.make_eval_hook(
|
||||
engine_client,
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
),
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
)
|
||||
except EngineError as e:
|
||||
await _safe_send_json(websocket, {"type": "error", "detail": f"engine unavailable: {e}"})
|
||||
|
|
@ -321,12 +329,15 @@ async def _run_turn_and_speak(
|
|||
reply = result.client_reply or ""
|
||||
# Persist only after the client reply has been generated. A failed AI turn
|
||||
# must not leave a learner-only transcript in review or history.
|
||||
await _append_voice_turn(
|
||||
await turn_runtime.record_completed_turn(
|
||||
sess,
|
||||
TurnRecord(
|
||||
ctx,
|
||||
result,
|
||||
context_prefix="voice session",
|
||||
counselor_turn=TurnRecord(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
speaker="counselor",
|
||||
stage=ctx.state_after.stage.value,
|
||||
stage=turn_runtime.stage_label(ctx.state_after.stage),
|
||||
text=learner_text,
|
||||
text_masked=ctx.learner_text_masked,
|
||||
audio_ref=audio_ref,
|
||||
|
|
@ -336,24 +347,7 @@ async def _run_turn_and_speak(
|
|||
evaluation=result.evaluation,
|
||||
),
|
||||
)
|
||||
if reply:
|
||||
# Persist the generated client reply before TTS playback.
|
||||
await _append_voice_turn(
|
||||
sess,
|
||||
TurnRecord(
|
||||
turn_seq=result.turn_seq,
|
||||
speaker="client",
|
||||
stage=result.stage,
|
||||
text=reply,
|
||||
text_masked=reply,
|
||||
llm_provider=result.llm_provider,
|
||||
model=result.model,
|
||||
tokens_in=result.tokens_in,
|
||||
tokens_out=result.tokens_out,
|
||||
cost_usd=result.cost_usd,
|
||||
),
|
||||
)
|
||||
await _update_voice_state(sess, result.state_after)
|
||||
await turn_runtime.record_safety_event(sess, ctx, result)
|
||||
|
||||
# Send the final client text before audio playback.
|
||||
await _safe_send_json(
|
||||
|
|
@ -367,6 +361,8 @@ async def _run_turn_and_speak(
|
|||
"turn_seq": result.turn_seq,
|
||||
"safety_flagged": result.safety_flagged,
|
||||
"crisis_kind": result.crisis_kind,
|
||||
"crisis_resource": result.crisis_resource,
|
||||
"conversation_stopped": result.conversation_stopped,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -403,49 +399,17 @@ async def _load_voice_session(
|
|||
session_id: str,
|
||||
principal: Principal,
|
||||
) -> tuple[InProcSession | None, str | None]:
|
||||
sess = await session_persistence.load_session(session_id, principal, allow_ended=True)
|
||||
if sess is not None:
|
||||
store.put(sess)
|
||||
elif runtime_fallback_allowed():
|
||||
sess = store.get(session_id)
|
||||
if sess is None:
|
||||
sess, err = await turn_runtime.load_owned_session(session_id, principal)
|
||||
if err == turn_runtime.SessionAccessError.NOT_FOUND:
|
||||
return None, f"unknown session {session_id}"
|
||||
if sess.learner_id != principal.user_id:
|
||||
if err == turn_runtime.SessionAccessError.FORBIDDEN:
|
||||
return None, "session does not belong to user"
|
||||
if sess.ended:
|
||||
if err == turn_runtime.SessionAccessError.ENDED:
|
||||
return None, "session already ended"
|
||||
assert sess is not None
|
||||
return sess, None
|
||||
|
||||
|
||||
async def _append_voice_turn(sess: InProcSession, turn: TurnRecord) -> None:
|
||||
if await session_persistence.append_turn(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
turn=turn,
|
||||
):
|
||||
sess.turns.append(turn)
|
||||
store.put(sess)
|
||||
return
|
||||
require_runtime_fallback_allowed("voice session turn append")
|
||||
store.append_turn(sess.session_id, turn)
|
||||
|
||||
|
||||
async def _update_voice_state(
|
||||
sess: InProcSession,
|
||||
state: state_machine.SessionState,
|
||||
) -> None:
|
||||
if await session_persistence.update_state(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
state=state,
|
||||
):
|
||||
sess.state = state
|
||||
store.put(sess)
|
||||
return
|
||||
require_runtime_fallback_allowed("voice session state update")
|
||||
store.update_state(sess.session_id, state)
|
||||
|
||||
|
||||
async def _principal_from_websocket(websocket: WebSocket) -> Principal | None:
|
||||
"""Restore the same server-side browser session used by REST routes."""
|
||||
raw_cookie = websocket.cookies.get(settings.cookie_name)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue