539 lines
25 KiB
Python
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"
|