현재 작업 전체 반영

This commit is contained in:
Yun Chan 2026-06-27 16:08:41 +09:00
parent 5560638e54
commit c0dddab594
85 changed files with 11322 additions and 539 deletions

View file

@ -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."""

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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)}

View file

@ -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,

View file

@ -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)