137 lines
3.6 KiB
Python
137 lines
3.6 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""JWT utilities for OAuth access tokens."""
|
|
|
|
import datetime
|
|
import os
|
|
import secrets
|
|
import uuid
|
|
from typing import Optional
|
|
|
|
import jwt
|
|
|
|
from .db import get_conn
|
|
|
|
JWT_ALG = "HS256"
|
|
ACCESS_TOKEN_HOURS = 1
|
|
REFRESH_TOKEN_DAYS = 30
|
|
|
|
|
|
def jwt_secret() -> str:
|
|
secret = os.environ.get("MCP_JWT_SECRET", "").strip()
|
|
if not secret:
|
|
secret = secrets.token_urlsafe(48)
|
|
os.environ["MCP_JWT_SECRET"] = secret
|
|
return secret
|
|
|
|
|
|
def base_url() -> str:
|
|
return os.environ.get("MCP_BASE_URL", "https://mcp.loogle.it").rstrip("/")
|
|
|
|
|
|
def scopes_for_user(user: dict, requested: str) -> str:
|
|
parts = [s for s in requested.split() if s]
|
|
allowed = {
|
|
"context:read",
|
|
"context:write",
|
|
"knowledge:read",
|
|
"knowledge:write",
|
|
"gitea:read",
|
|
"gitea:write",
|
|
"home:read",
|
|
"irrigation:read",
|
|
"turni:read",
|
|
}
|
|
if user.get("is_admin"):
|
|
allowed.add("admin")
|
|
filtered = [s for s in parts if s in allowed]
|
|
if not filtered:
|
|
filtered = [
|
|
"context:read",
|
|
"context:write",
|
|
"knowledge:read",
|
|
"knowledge:write",
|
|
"gitea:read",
|
|
"gitea:write",
|
|
"home:read",
|
|
"irrigation:read",
|
|
"turni:read",
|
|
]
|
|
return " ".join(filtered)
|
|
|
|
|
|
def create_access_token(user: dict, scope: str, client_id: str) -> tuple[str, str]:
|
|
jti = str(uuid.uuid4())
|
|
now = datetime.datetime.utcnow()
|
|
payload = {
|
|
"iss": base_url(),
|
|
"sub": user["username"],
|
|
"uid": user["id"],
|
|
"scope": scope,
|
|
"client_id": client_id,
|
|
"jti": jti,
|
|
"iat": now,
|
|
"exp": now + datetime.timedelta(hours=ACCESS_TOKEN_HOURS),
|
|
}
|
|
token = jwt.encode(payload, jwt_secret(), algorithm=JWT_ALG)
|
|
return token, jti
|
|
|
|
|
|
def create_refresh_token(user_id: int, scope: str, client_id: str) -> str:
|
|
token = secrets.token_urlsafe(48)
|
|
expires = (
|
|
datetime.datetime.utcnow() + datetime.timedelta(days=REFRESH_TOKEN_DAYS)
|
|
).strftime("%Y-%m-%d %H:%M:%S")
|
|
get_conn().execute(
|
|
"INSERT INTO refresh_tokens(token,user_id,client_id,scope,expires_at) VALUES (?,?,?,?,?)",
|
|
(token, user_id, client_id, scope, expires),
|
|
)
|
|
get_conn().commit()
|
|
return token
|
|
|
|
|
|
def decode_access_token(token: str) -> Optional[dict]:
|
|
try:
|
|
payload = jwt.decode(token, jwt_secret(), algorithms=[JWT_ALG], issuer=base_url())
|
|
except jwt.PyJWTError:
|
|
return None
|
|
jti = payload.get("jti")
|
|
if not jti:
|
|
return None
|
|
row = get_conn().execute(
|
|
"SELECT 1 FROM revoked_jtis WHERE jti=? AND expires_at > datetime('now')",
|
|
(jti,),
|
|
).fetchone()
|
|
if row:
|
|
return None
|
|
return payload
|
|
|
|
|
|
def revoke_jti(jti: str, expires_at: datetime.datetime) -> None:
|
|
get_conn().execute(
|
|
"INSERT OR IGNORE INTO revoked_jtis(jti,expires_at) VALUES (?,?)",
|
|
(jti, expires_at.strftime("%Y-%m-%d %H:%M:%S")),
|
|
)
|
|
get_conn().commit()
|
|
|
|
|
|
def revoke_refresh_token(token: str) -> None:
|
|
get_conn().execute("UPDATE refresh_tokens SET revoked=1 WHERE token=?", (token,))
|
|
get_conn().commit()
|
|
|
|
|
|
def consume_refresh_token(token: str) -> Optional[dict]:
|
|
row = get_conn().execute(
|
|
"SELECT * FROM refresh_tokens WHERE token=? AND revoked=0 AND expires_at > datetime('now')",
|
|
(token,),
|
|
).fetchone()
|
|
if not row:
|
|
return None
|
|
return dict(row)
|
|
|
|
|
|
def has_scope(claims: dict, scope: str) -> bool:
|
|
scopes = set((claims.get("scope") or "").split())
|
|
if "admin" in scopes:
|
|
return True
|
|
return scope in scopes
|