np_app/service/account.py
2026-07-26 03:00:54 +08:00

539 lines
25 KiB
Python

import base64
import hashlib
import json
import os
import secrets
import sqlite3
import threading
import time
import zlib
from http.cookies import SimpleCookie
from pathlib import Path
from urllib.parse import urlencode
from urllib.request import Request, urlopen
def _json_blob(value):
if value is None:
return None
raw = json.dumps(value, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
return sqlite3.Binary(zlib.compress(raw, level=6))
def _blob_json(value):
if value is None:
return None
return json.loads(zlib.decompress(value).decode("utf-8"))
def _payload_summary(payload):
payload = payload or {}
workflow = payload.get("workflow", "analysis")
mode = payload.get("mode", "tube")
model = payload.get("model") or {}
strands = payload.get("strands") or []
if workflow == "design":
complexes = payload.get("design_complexes") or payload.get("design_targets") or []
tube_sizes = [int(row.get("max_size", 0) or 0) for row in (payload.get("design_tubes") or [])]
max_size = max(tube_sizes, default=int((payload.get("design") or {}).get("off_target_max_size", 0) or 0))
else:
complexes = [line for line in str(payload.get("complexes_text") or "").splitlines() if line.strip()]
max_size = int((payload.get("tube") or {}).get("max_size", 0) or 0)
return {
"workflow": workflow,
"mode": mode,
"material": model.get("material", "rna"),
"celsius": float(model.get("celsius", 37) or 37),
"sodium": float(model.get("sodium", 0) or 0),
"magnesium": float(model.get("magnesium", 0) or 0),
"max_size": max_size,
"strand_count": len(strands),
"complex_count": len(complexes),
"compute": list(payload.get("compute") or (["design"] if workflow == "design" else [])),
"trials": int((payload.get("design") or {}).get("trials", 0) or 0),
"stop_condition": float((payload.get("design") or {}).get("stop_condition", 0) or 0),
"max_time_seconds": int((payload.get("design") or {}).get("max_time_seconds", 0) or 0),
}
class AccountStore:
def __init__(self, path):
self.path = Path(path)
self._init_lock = threading.Lock()
self._initialized = False
def _connect(self):
self.path.parent.mkdir(parents=True, exist_ok=True)
connection = sqlite3.connect(self.path, timeout=30)
connection.row_factory = sqlite3.Row
connection.execute("PRAGMA busy_timeout = 30000")
connection.execute("PRAGMA foreign_keys = ON")
return connection
def initialize(self):
with self._init_lock:
if self._initialized:
return
with self._connect() as connection:
connection.execute("PRAGMA journal_mode = WAL")
connection.execute("PRAGMA synchronous = NORMAL")
connection.executescript(
"""
CREATE TABLE IF NOT EXISTS users (
user_id TEXT PRIMARY KEY,
username TEXT NOT NULL,
email TEXT,
display_name TEXT,
groups_json TEXT NOT NULL DEFAULT '[]',
created_at REAL NOT NULL,
last_seen_at REAL NOT NULL
);
CREATE TABLE IF NOT EXISTS jobs (
job_id TEXT PRIMARY KEY,
user_id TEXT NOT NULL REFERENCES users(user_id) ON DELETE CASCADE,
status TEXT NOT NULL,
source TEXT NOT NULL DEFAULT 'job',
created_at REAL NOT NULL,
updated_at REAL NOT NULL,
elapsed_seconds REAL,
workflow TEXT NOT NULL,
mode TEXT NOT NULL,
material TEXT NOT NULL,
celsius REAL NOT NULL,
sodium REAL NOT NULL,
magnesium REAL NOT NULL,
max_size INTEGER NOT NULL DEFAULT 0,
strand_count INTEGER NOT NULL DEFAULT 0,
complex_count INTEGER NOT NULL DEFAULT 0,
compute_json TEXT NOT NULL DEFAULT '[]',
trials INTEGER NOT NULL DEFAULT 0,
stop_condition REAL NOT NULL DEFAULT 0,
max_time_seconds INTEGER NOT NULL DEFAULT 0,
payload_blob BLOB NOT NULL,
result_blob BLOB,
error_blob BLOB,
stored_bytes INTEGER NOT NULL DEFAULT 0
);
CREATE INDEX IF NOT EXISTS jobs_user_created_idx ON jobs(user_id, created_at DESC);
CREATE INDEX IF NOT EXISTS jobs_user_status_idx ON jobs(user_id, status, created_at DESC);
CREATE TABLE IF NOT EXISTS shares (
share_id TEXT PRIMARY KEY,
job_id TEXT NOT NULL REFERENCES jobs(job_id) ON DELETE CASCADE,
user_id TEXT NOT NULL REFERENCES users(user_id) ON DELETE CASCADE,
active INTEGER NOT NULL DEFAULT 1,
created_at REAL NOT NULL,
expires_at REAL,
last_access_at REAL,
access_count INTEGER NOT NULL DEFAULT 0
);
CREATE INDEX IF NOT EXISTS shares_user_created_idx ON shares(user_id, created_at DESC);
"""
)
self._initialized = True
def upsert_user(self, user):
self.initialize()
now = time.time()
with self._connect() as connection:
connection.execute(
"""
INSERT INTO users(user_id, username, email, display_name, groups_json, created_at, last_seen_at)
VALUES(?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(user_id) DO UPDATE SET
username=excluded.username,
email=excluded.email,
display_name=excluded.display_name,
groups_json=excluded.groups_json,
last_seen_at=excluded.last_seen_at
""",
(
user["user_id"], user.get("username") or user["user_id"], user.get("email"),
user.get("display_name"), json.dumps(user.get("groups") or [], ensure_ascii=False), now, now,
),
)
def create_job(self, job_id, user, payload, status="queued", source="job", created_at=None, result=None, error=None):
self.upsert_user(user)
created_at = float(created_at or time.time())
summary = _payload_summary(payload)
payload_blob = _json_blob(payload)
result_blob = _json_blob(result)
error_blob = _json_blob(error)
stored_bytes = sum(len(item) for item in (payload_blob, result_blob, error_blob) if item is not None)
with self._connect() as connection:
connection.execute(
"""
INSERT OR IGNORE INTO jobs(
job_id, user_id, status, source, created_at, updated_at, workflow, mode, material,
celsius, sodium, magnesium, max_size, strand_count, complex_count, compute_json,
trials, stop_condition, max_time_seconds, payload_blob, result_blob, error_blob, stored_bytes
) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
job_id, user["user_id"], status, source, created_at, created_at,
summary["workflow"], summary["mode"], summary["material"], summary["celsius"],
summary["sodium"], summary["magnesium"], summary["max_size"], summary["strand_count"],
summary["complex_count"], json.dumps(summary["compute"]), summary["trials"],
summary["stop_condition"], summary["max_time_seconds"], payload_blob, result_blob,
error_blob, stored_bytes,
),
)
def update_job(self, job_id, status, *, result=None, error=None, elapsed_seconds=None):
self.initialize()
fields = ["status = ?", "updated_at = ?"]
values = [status, time.time()]
if result is not None:
fields.append("result_blob = ?")
values.append(_json_blob(result))
if error is not None:
fields.append("error_blob = ?")
values.append(_json_blob(error))
if elapsed_seconds is not None:
fields.append("elapsed_seconds = ?")
values.append(float(elapsed_seconds))
values.append(job_id)
with self._connect() as connection:
connection.execute(f"UPDATE jobs SET {', '.join(fields)} WHERE job_id = ?", values)
connection.execute(
"UPDATE jobs SET stored_bytes = length(payload_blob) + coalesce(length(result_blob), 0) + coalesce(length(error_blob), 0) WHERE job_id = ?",
(job_id,),
)
def owner_id(self, job_id):
self.initialize()
with self._connect() as connection:
row = connection.execute("SELECT user_id FROM jobs WHERE job_id = ?", (job_id,)).fetchone()
return row["user_id"] if row else None
def get_job(self, user_id, job_id, include_content=True):
self.initialize()
columns = "*" if include_content else "job_id,user_id,status,source,created_at,updated_at,elapsed_seconds,workflow,mode,material,celsius,sodium,magnesium,max_size,strand_count,complex_count,compute_json,trials,stop_condition,max_time_seconds,stored_bytes"
with self._connect() as connection:
row = connection.execute(f"SELECT {columns} FROM jobs WHERE job_id = ? AND user_id = ?", (job_id, user_id)).fetchone()
return self._job_row(row, include_content=include_content) if row else None
def list_jobs(self, user_id, filters=None):
self.initialize()
filters = filters or {}
limit = min(100, max(1, int(filters.get("limit", 30))))
offset = max(0, int(filters.get("offset", 0)))
clauses = ["user_id = ?"]
values = [user_id]
for field in ("status", "workflow", "mode", "material"):
value = str(filters.get(field) or "").strip()
if value:
clauses.append(f"{field} = ?")
values.append(value)
search = str(filters.get("q") or "").strip()
if search:
clauses.append("(job_id LIKE ? OR workflow LIKE ? OR mode LIKE ? OR material LIKE ?)")
values.extend([f"%{search}%"] * 4)
where = " AND ".join(clauses)
columns = "job_id,user_id,status,source,created_at,updated_at,elapsed_seconds,workflow,mode,material,celsius,sodium,magnesium,max_size,strand_count,complex_count,compute_json,trials,stop_condition,max_time_seconds,stored_bytes"
with self._connect() as connection:
total = connection.execute(f"SELECT count(*) AS count FROM jobs WHERE {where}", values).fetchone()["count"]
rows = connection.execute(
f"SELECT {columns} FROM jobs WHERE {where} ORDER BY created_at DESC LIMIT ? OFFSET ?",
[*values, limit, offset],
).fetchall()
return {"items": [self._job_row(row, include_content=False) for row in rows], "total": total, "limit": limit, "offset": offset}
def usage(self, user_id):
self.initialize()
with self._connect() as connection:
totals = connection.execute(
"""SELECT count(*) AS total_jobs,
coalesce(sum(CASE WHEN status='success' THEN 1 ELSE 0 END),0) AS success_jobs,
coalesce(sum(CASE WHEN status='error' THEN 1 ELSE 0 END),0) AS error_jobs,
coalesce(sum(CASE WHEN status IN ('queued','running','cancel_requested') THEN 1 ELSE 0 END),0) AS active_jobs,
coalesce(sum(elapsed_seconds),0) AS compute_seconds,
coalesce(sum(stored_bytes),0) AS stored_bytes FROM jobs WHERE user_id=?""",
(user_id,),
).fetchone()
shares = connection.execute("SELECT count(*) AS count FROM shares WHERE user_id=? AND active=1", (user_id,)).fetchone()["count"]
result = dict(totals)
result["active_shares"] = shares
return result
def import_history(self, user, entries):
imported = 0
for entry in entries:
payload = entry.get("payload") if isinstance(entry, dict) else None
if not isinstance(payload, dict):
continue
result = entry.get("result") if isinstance(entry.get("result"), dict) else None
raw_id = str(entry.get("id") or entry.get("job_id") or uuid_token())
job_id = f"import-{hashlib.sha256((user['user_id'] + ':' + raw_id).encode()).hexdigest()[:24]}"
created = entry.get("created_at") or entry.get("created_at_iso")
try:
created_at = float(created)
except (TypeError, ValueError):
try:
created_at = time.mktime(time.strptime(str(created).split(".")[0].replace("Z", ""), "%Y-%m-%dT%H:%M:%S"))
except (TypeError, ValueError):
created_at = time.time()
before = self.get_job(user["user_id"], job_id, include_content=False)
self.create_job(job_id, user, payload, status="success" if result else "input_only", source="import", created_at=created_at, result=result)
if before is None:
imported += 1
return imported
def delete_job(self, user_id, job_id):
self.initialize()
with self._connect() as connection:
row = connection.execute("SELECT status FROM jobs WHERE job_id=? AND user_id=?", (job_id, user_id)).fetchone()
if row is None:
return False
if row["status"] in {"queued", "running", "cancel_requested"}:
raise ValueError("Active jobs cannot be deleted.")
connection.execute("DELETE FROM jobs WHERE job_id=? AND user_id=?", (job_id, user_id))
return True
def create_share(self, user_id, job_id, expires_in=None):
job = self.get_job(user_id, job_id, include_content=False)
if job is None:
raise ValueError("Job not found.")
now = time.time()
expires_at = now + int(expires_in) if expires_in else None
share_id = uuid_token(18)
with self._connect() as connection:
connection.execute(
"INSERT INTO shares(share_id,job_id,user_id,active,created_at,expires_at) VALUES(?,?,?,?,?,?)",
(share_id, job_id, user_id, 1, now, expires_at),
)
return self.get_share_for_owner(user_id, share_id)
def share_metadata(self, share_id):
self.initialize()
now = time.time()
with self._connect() as connection:
row = connection.execute(
"""SELECT s.share_id, s.job_id, s.user_id, s.created_at, s.expires_at, s.active
FROM shares s WHERE s.share_id=? AND s.active=1 AND (s.expires_at IS NULL OR s.expires_at>?)""",
(share_id, now),
).fetchone()
return dict(row) if row else None
def record_share_access(self, share_id):
self.initialize()
now = time.time()
with self._connect() as connection:
connection.execute(
"UPDATE shares SET access_count=access_count+1,last_access_at=? WHERE share_id=?",
(now, share_id),
)
def list_shares(self, user_id):
self.initialize()
with self._connect() as connection:
rows = connection.execute(
"""SELECT s.*,j.status,j.workflow,j.mode,j.material,j.created_at AS job_created_at
FROM shares s JOIN jobs j ON j.job_id=s.job_id WHERE s.user_id=? ORDER BY s.created_at DESC""",
(user_id,),
).fetchall()
return [dict(row) for row in rows]
def get_share_for_owner(self, user_id, share_id):
self.initialize()
with self._connect() as connection:
row = connection.execute("SELECT * FROM shares WHERE share_id=? AND user_id=?", (share_id, user_id)).fetchone()
return dict(row) if row else None
def update_share(self, user_id, share_id, *, active=None, expires_in="unchanged"):
share = self.get_share_for_owner(user_id, share_id)
if share is None:
return None
fields = []
values = []
if active is not None:
fields.append("active=?")
values.append(1 if active else 0)
if expires_in != "unchanged":
fields.append("expires_at=?")
values.append(time.time() + int(expires_in) if expires_in else None)
if fields:
values.extend([share_id, user_id])
with self._connect() as connection:
connection.execute(f"UPDATE shares SET {', '.join(fields)} WHERE share_id=? AND user_id=?", values)
return self.get_share_for_owner(user_id, share_id)
def resolve_share(self, share_id, *, record_access=True):
self.initialize()
now = time.time()
with self._connect() as connection:
row = connection.execute(
"""SELECT
s.share_id, s.created_at AS share_created_at,
j.job_id, j.user_id, j.status, j.source, j.created_at, j.updated_at,
j.elapsed_seconds, j.workflow, j.mode, j.material, j.celsius, j.sodium,
j.magnesium, j.max_size, j.strand_count, j.complex_count, j.compute_json,
j.trials, j.stop_condition, j.max_time_seconds, j.payload_blob,
j.result_blob, j.error_blob, j.stored_bytes
FROM shares s JOIN jobs j ON j.job_id=s.job_id
WHERE s.share_id=? AND s.active=1 AND (s.expires_at IS NULL OR s.expires_at>?)""",
(share_id, now),
).fetchone()
if row is None:
return None
if record_access:
connection.execute("UPDATE shares SET access_count=access_count+1,last_access_at=? WHERE share_id=?", (now, share_id))
job = self._job_row(row, include_content=True)
return {
"id": share_id,
"share_id": share_id,
"job_id": job["job_id"],
"created_at": row["share_created_at"],
"status": job["status"],
"updated_at": job["updated_at"],
"elapsed_seconds": job["elapsed_seconds"],
"payload": job["payload"],
"result": job["result"],
"error": job["error"],
"result_summary": job["result_summary"],
}
@staticmethod
def _job_row(row, include_content):
item = dict(row)
item["compute"] = json.loads(item.pop("compute_json") or "[]")
item["result_summary"] = {
key: item[key] for key in (
"workflow", "mode", "material", "celsius", "sodium", "magnesium", "max_size",
"strand_count", "complex_count", "compute", "trials", "stop_condition", "max_time_seconds",
)
}
if include_content:
item["payload"] = _blob_json(item.pop("payload_blob"))
item["result"] = _blob_json(item.pop("result_blob"))
item["error"] = _blob_json(item.pop("error_blob"))
return item
def uuid_token(size=24):
return secrets.token_urlsafe(size)
class OIDCAuth:
def __init__(self, redis_factory):
self.redis_factory = redis_factory
self.issuer = os.environ.get("NP_OIDC_ISSUER", "https://auth.lihato.icu/application/o/nupack-account/").rstrip("/") + "/"
self.client_id = os.environ.get("NP_OIDC_CLIENT_ID", "np-replica-web")
self.redirect_uri = os.environ.get("NP_OIDC_REDIRECT_URI", "https://np.lihato.icu/auth/callback")
self.post_logout_uri = os.environ.get("NP_OIDC_POST_LOGOUT_URI", "https://np.lihato.icu/")
self.scopes = os.environ.get("NP_OIDC_SCOPES", "openid profile email").strip()
self.cookie_name = os.environ.get("NP_AUTH_COOKIE_NAME", "np_session")
self.session_ttl = int(os.environ.get("NP_AUTH_SESSION_TTL_SECONDS", str(7 * 86400)))
self.required = os.environ.get("NP_AUTH_REQUIRED", "1") != "0"
self._configuration = None
self._config_lock = threading.Lock()
def configuration(self):
with self._config_lock:
if self._configuration is None:
self._configuration = self._request_json(self.issuer + ".well-known/openid-configuration")
return self._configuration
@staticmethod
def _request_json(url, *, data=None, headers=None):
body = urlencode(data).encode("utf-8") if data is not None else None
request = Request(url, data=body, headers=headers or {})
with urlopen(request, timeout=15) as response:
return json.loads(response.read().decode("utf-8"))
def begin_login(self, next_path="/"):
if not next_path.startswith("/") or next_path.startswith("//"):
next_path = "/"
state = uuid_token(24)
verifier = uuid_token(48)
challenge = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
self.redis_factory().setex(
f"np_replica:oidc_state:{state}", 600,
json.dumps({"verifier": verifier, "next": next_path}),
)
config = self.configuration()
query = urlencode({
"client_id": self.client_id,
"response_type": "code",
"redirect_uri": self.redirect_uri,
"scope": self.scopes,
"state": state,
"code_challenge": challenge,
"code_challenge_method": "S256",
})
return f"{config['authorization_endpoint']}?{query}"
def complete_login(self, params):
state = (params.get("state") or [""])[0]
code = (params.get("code") or [""])[0]
if not state or not code:
raise ValueError("OIDC callback is missing code or state.")
key = f"np_replica:oidc_state:{state}"
client = self.redis_factory()
raw = client.get(key)
client.delete(key)
if raw is None:
raise ValueError("OIDC login state has expired.")
pending = json.loads(raw)
config = self.configuration()
tokens = self._request_json(config["token_endpoint"], data={
"grant_type": "authorization_code",
"client_id": self.client_id,
"code": code,
"redirect_uri": self.redirect_uri,
"code_verifier": pending["verifier"],
}, headers={"Content-Type": "application/x-www-form-urlencoded"})
claims = self._request_json(config["userinfo_endpoint"], headers={"Authorization": f"Bearer {tokens['access_token']}"})
subject = str(claims.get("sub") or "").strip()
if not subject:
raise ValueError("OIDC userinfo did not return a subject.")
user = {
"user_id": subject,
"username": claims.get("preferred_username") or claims.get("nickname") or claims.get("email") or subject,
"email": claims.get("email"),
"display_name": claims.get("name") or claims.get("preferred_username") or subject,
"groups": claims.get("groups") or [],
}
session_id = uuid_token(32)
client.setex(
f"np_replica:session:{session_id}", self.session_ttl,
json.dumps({"user": user, "id_token": tokens.get("id_token")}, ensure_ascii=False),
)
return session_id, user, pending.get("next") or "/"
def current_session(self, headers):
if not self.required:
user = {"user_id": "development", "username": "development", "email": None, "display_name": "Development", "groups": []}
return {"user": user, "session_id": None}
cookie = SimpleCookie()
cookie.load(headers.get("Cookie", ""))
morsel = cookie.get(self.cookie_name)
if morsel is None:
return None
session_id = morsel.value
raw = self.redis_factory().get(f"np_replica:session:{session_id}")
if raw is None:
return None
session = json.loads(raw)
session["session_id"] = session_id
return session
def logout_url(self, session):
config = self.configuration()
endpoint = config.get("end_session_endpoint")
if not endpoint:
return "/"
query = {"post_logout_redirect_uri": self.post_logout_uri}
if session and session.get("id_token"):
query["id_token_hint"] = session["id_token"]
return f"{endpoint}?{urlencode(query)}"
def delete_session(self, session):
if session and session.get("session_id"):
self.redis_factory().delete(f"np_replica:session:{session['session_id']}")
def cookie_header(self, session_id):
return f"{self.cookie_name}={session_id}; Path=/; Max-Age={self.session_ttl}; HttpOnly; Secure; SameSite=Lax"
def clear_cookie_header(self):
return f"{self.cookie_name}=; Path=/; Max-Age=0; HttpOnly; Secure; SameSite=Lax"