Cleanup of origins handling. Added site_url and site_path such that these can be determined reliably, and we print it in the startbox.

This commit is contained in:
2025-12-06 03:39:05 +00:00
parent df5c176bcd
commit a1b73711e6
6 changed files with 157 additions and 89 deletions
+18
View File
@@ -1,3 +1,4 @@
from dataclasses import dataclass
from datetime import timedelta from datetime import timedelta
# Shared configuration constants for session management. # Shared configuration constants for session management.
@@ -5,3 +6,20 @@ SESSION_LIFETIME = timedelta(hours=24)
# Lifetime for reset links created by admins # Lifetime for reset links created by admins
RESET_LIFETIME = timedelta(days=14) RESET_LIFETIME = timedelta(days=14)
@dataclass
class PaskiaConfig:
"""Runtime configuration for the Paskia authentication server."""
rp_id: str
rp_name: str | None
origins: list[str] | None
auth_host: str | None
site_url: str # Base URL without trailing path (e.g. https://example.com)
site_path: str # Path to auth UI: "/" if auth_host, else "/auth/"
# Listen address (one of host:port or uds)
host: str | None = None
port: int | None = None
uds: str | None = None
devmode: bool = False
+61 -21
View File
@@ -7,6 +7,8 @@ from urllib.parse import urlparse
import uvicorn import uvicorn
from paskia.util.hostutil import normalize_origin
DEFAULT_HOST = "localhost" DEFAULT_HOST = "localhost"
DEFAULT_SERVE_PORT = 4401 DEFAULT_SERVE_PORT = 4401
@@ -172,53 +174,91 @@ def main():
else: else:
host = port = uds = all_ifaces = None # type: ignore host = port = uds = all_ifaces = None # type: ignore
# Collect origins and handle auth_host # Collect and normalize origins, handle auth_host
origins = getattr(args, "origins", None) or [] origins = [normalize_origin(o) for o in (getattr(args, "origins", None) or [])]
if args.auth_host: if args.auth_host:
# Normalize auth_host with scheme # Normalize auth_host with scheme
if "://" not in args.auth_host: if "://" not in args.auth_host:
args.auth_host = f"https://{args.auth_host}" args.auth_host = f"https://{args.auth_host}"
validate_auth_host(args.auth_host, args.rp_id) validate_auth_host(args.auth_host, args.rp_id)
from paskia.util import hostutil as _hostutil # local import
_hostutil.reload_config()
# If origins are configured, ensure auth_host is included at top # If origins are configured, ensure auth_host is included at top
if origins: if origins:
# Insert auth_host at the beginning (Passkey.__init__ will normalize/dedupe) # Insert auth_host at the beginning (Passkey.__init__ will dedupe)
origins.insert(0, args.auth_host) origins.insert(0, args.auth_host)
# Compute site_url and site_path for reset links
# Priority: auth_host > first origin with localhost > http://localhost:port
if args.auth_host:
site_url = args.auth_host.rstrip("/")
site_path = "/"
elif origins:
# Find localhost origin if rp_id is localhost, else use first origin
localhost_origin = (
next((o for o in origins if "://localhost" in o), None)
if args.rp_id == "localhost"
else None
)
site_url = (localhost_origin or origins[0]).rstrip("/")
site_path = "/auth/"
elif args.rp_id == "localhost" and port:
# Dev mode: use http with port
site_url = f"http://localhost:{port}"
site_path = "/auth/"
else:
site_url = f"https://{args.rp_id}"
site_path = "/auth/"
# Build runtime configuration
from paskia.config import PaskiaConfig
config = PaskiaConfig(
rp_id=args.rp_id,
rp_name=args.rp_name or None,
origins=origins or None,
auth_host=args.auth_host or None,
site_url=site_url,
site_path=site_path,
host=host,
port=port,
uds=uds,
)
# Export configuration via single JSON env variable for worker processes # Export configuration via single JSON env variable for worker processes
# (PASKIA_DEVMODE is kept separate as it's externally defined)
# All keys are always present; None is used where no value is configured
import json import json
config = { config_json = {
"rp_id": args.rp_id, "rp_id": config.rp_id,
"rp_name": args.rp_name or None, "rp_name": config.rp_name,
"origins": origins or None, "origins": config.origins,
"auth_host": args.auth_host or None, "auth_host": config.auth_host,
"site_url": config.site_url,
"site_path": config.site_path,
} }
os.environ["PASKIA_CONFIG"] = json.dumps(config) os.environ["PASKIA_CONFIG"] = json.dumps(config_json)
# One-time initialization + bootstrap before starting any server processes. # Initialize globals (without bootstrap yet)
# Lifespan in worker processes will call globals.init with bootstrap disabled.
from paskia import globals as _globals # local import from paskia import globals as _globals # local import
asyncio.run( asyncio.run(
_globals.init( _globals.init(
rp_id=config["rp_id"], rp_id=config.rp_id,
rp_name=config["rp_name"], rp_name=config.rp_name,
origins=config["origins"], origins=config.origins,
bootstrap=True, bootstrap=False,
) )
) )
# Print startup configuration # Print startup configuration
from paskia.util import startupbox from paskia.util import startupbox
startupbox.print_startup_config(_globals.passkey.instance, args, host, port, uds) startupbox.print_startup_config(config)
# Bootstrap after startup box is printed
from paskia.bootstrap import bootstrap_if_needed
asyncio.run(bootstrap_if_needed())
# Handle recover-admin command (no server start) # Handle recover-admin command (no server start)
if args.command == "reset": if args.command == "reset":
+21 -1
View File
@@ -1,15 +1,19 @@
import logging import logging
import os import os
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from pathlib import Path
from fastapi import FastAPI, HTTPException, Request, Response from fastapi import FastAPI, HTTPException, Request, Response
from fastapi.responses import RedirectResponse from fastapi.responses import FileResponse, RedirectResponse
from fastapi.staticfiles import StaticFiles from fastapi.staticfiles import StaticFiles
from paskia.fastapi import admin, api, auth_host, ws from paskia.fastapi import admin, api, auth_host, ws
from paskia.fastapi.session import AUTH_COOKIE from paskia.fastapi.session import AUTH_COOKIE
from paskia.util import frontend, hostutil, passphrase from paskia.util import frontend, hostutil, passphrase
# Path to examples/index.html when running from source tree
_EXAMPLES_DIR = Path(__file__).parent.parent.parent / "examples"
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): # pragma: no cover - startup path async def lifespan(app: FastAPI): # pragma: no cover - startup path
@@ -92,6 +96,22 @@ async def admin_root(request: Request, auth=AUTH_COOKIE):
return await admin.adminapp(request, auth) # Delegated to admin app return await admin.adminapp(request, auth) # Delegated to admin app
@app.get("/auth/examples/", include_in_schema=False)
async def examples_page():
"""Serve examples/index.html when running from source tree.
This provides a simple test page for API mode authentication flows
without depending on the Vue frontend build.
"""
index_file = _EXAMPLES_DIR / "index.html"
if not index_file.is_file():
raise HTTPException(
status_code=404,
detail="Examples not available (not running from source tree)",
)
return FileResponse(index_file, media_type="text/html")
# Note: this catch-all handler must be the last route defined # Note: this catch-all handler must be the last route defined
@app.get("/{reset}") @app.get("/{reset}")
@app.get("/auth/{reset}") @app.get("/auth/{reset}")
+11 -14
View File
@@ -67,34 +67,31 @@ class Passkey:
self.rp_name = rp_name or rp_id self.rp_name = rp_name or rp_id
self.allowed_origins: set[str] | None = None self.allowed_origins: set[str] | None = None
if origins: if origins:
# Normalize and deduplicate origins into a set for O(1) lookups # Validate and deduplicate origins into a set for O(1) lookups
self.allowed_origins = { for o in origins:
self._normalize_and_validate_origin(o, rp_id) for o in origins self._validate_origin(o, rp_id)
} self.allowed_origins = set(origins)
self.supported_pub_key_algs = supported_pub_key_algs or [ self.supported_pub_key_algs = supported_pub_key_algs or [
COSEAlgorithmIdentifier.EDDSA, COSEAlgorithmIdentifier.EDDSA,
COSEAlgorithmIdentifier.ECDSA_SHA_256, COSEAlgorithmIdentifier.ECDSA_SHA_256,
COSEAlgorithmIdentifier.RSASSA_PKCS1_v1_5_SHA_256, COSEAlgorithmIdentifier.RSASSA_PKCS1_v1_5_SHA_256,
] ]
def _normalize_and_validate_origin(self, origin: str, rp_id: str) -> str: def _validate_origin(self, origin: str, rp_id: str) -> None:
"""Normalize and validate an origin URL against the rp_id.""" """Validate an origin URL against the rp_id."""
if "://" not in origin:
origin = f"https://{origin}"
hostname = urlparse(origin).hostname hostname = urlparse(origin).hostname
if not hostname: if not hostname:
raise ValueError(f"Invalid origin URL: no hostname found in '{origin}'") raise ValueError(f"Invalid origin URL: no hostname found in '{origin}'")
if hostname == rp_id or hostname.endswith(f".{rp_id}"): if hostname == rp_id or hostname.endswith(f".{rp_id}"):
return origin return
raise ValueError( raise ValueError(
f"Origin domain '{hostname}' must be the same as or a subdomain of rp_id '{rp_id}'" f"Origin domain '{hostname}' must be the same as or a subdomain of rp_id '{rp_id}'"
) )
def validate_origin(self, origin: str) -> str: def validate_origin(self, origin: str) -> str:
"""Validate that origin is allowed and return the normalized form. """Validate that origin is allowed and return it.
Args: Args:
origin: The origin URL to validate (from WebSocket request header) origin: The origin URL to validate (from WebSocket request header)
@@ -106,10 +103,10 @@ class Passkey:
ValueError: If origin is not in the allowed list (when origins are configured) ValueError: If origin is not in the allowed list (when origins are configured)
or if origin is not a valid subdomain of rp_id or if origin is not a valid subdomain of rp_id
""" """
normalized = self._normalize_and_validate_origin(origin, self.rp_id) self._validate_origin(origin, self.rp_id)
if self.allowed_origins is not None and normalized not in self.allowed_origins: if self.allowed_origins is not None and origin not in self.allowed_origins:
raise ValueError(f"Origin '{origin}' is not in the allowed origins list") raise ValueError(f"Origin '{origin}' is not in the allowed origins list")
return normalized return origin
### Registration Methods ### ### Registration Methods ###
+27 -36
View File
@@ -3,38 +3,31 @@
import json import json
import os import os
from functools import lru_cache from functools import lru_cache
from urllib.parse import urlparse, urlsplit from urllib.parse import urlsplit
from paskia.globals import passkey as global_passkey
@lru_cache(maxsize=1) @lru_cache(maxsize=1)
def _load_config() -> tuple[str, str] | None: def _load_config() -> dict:
"""Load auth_host from PASKIA_CONFIG JSON. """Load PASKIA_CONFIG JSON."""
Returns (scheme, netloc) tuple if configured, None otherwise.
"""
config_json = os.getenv("PASKIA_CONFIG") config_json = os.getenv("PASKIA_CONFIG")
if not config_json: if not config_json:
return None return {}
config = json.loads(config_json) return json.loads(config_json)
raw = config["auth_host"] # Always present, may be None
if not raw:
return None
parsed = urlparse(raw if "://" in raw else f"//{raw}")
netloc = parsed.netloc or parsed.path
if not netloc:
return None
return (parsed.scheme or "https", netloc.strip("/"))
def configured_auth_host() -> str | None:
cfg = _load_config()
return cfg[1] if cfg else None
def is_root_mode() -> bool: def is_root_mode() -> bool:
return _load_config() is not None return _load_config().get("auth_host") is not None
def configured_auth_host() -> str | None:
"""Return configured auth_host netloc, or None."""
auth_host = _load_config().get("auth_host")
if not auth_host:
return None
from urllib.parse import urlparse
parsed = urlparse(auth_host if "://" in auth_host else f"//{auth_host}")
return parsed.netloc or parsed.path or None
def ui_base_path() -> str: def ui_base_path() -> str:
@@ -42,25 +35,23 @@ def ui_base_path() -> str:
def auth_site_base_url() -> str: def auth_site_base_url() -> str:
"""Return the base URL for the auth site UI. """Return the base URL for the auth site UI (computed at startup)."""
If auth_host is configured (root mode), returns its URL.
Otherwise, constructs URL from rp_id with /auth/ path.
"""
cfg = _load_config() cfg = _load_config()
if cfg: return cfg.get("site_url", "https://localhost") + cfg.get("site_path", "/auth/")
scheme, netloc = cfg
return f"{scheme}://{netloc}/"
# Not in root mode: use rp_id with /auth/ path
rp_id = global_passkey.instance.rp_id
return f"https://{rp_id}/auth/"
def reset_link_url(token: str) -> str: def reset_link_url(token: str) -> str:
"""Generate a reset link URL for the given token."""
return f"{auth_site_base_url()}{token}" return f"{auth_site_base_url()}{token}"
def normalize_origin(origin: str) -> str:
"""Normalize an origin URL by adding https:// if no scheme is present."""
if "://" not in origin:
return f"https://{origin}"
return origin
def reload_config() -> None: def reload_config() -> None:
_load_config.cache_clear() _load_config.cache_clear()
+19 -17
View File
@@ -1,6 +1,13 @@
"""Startup configuration box formatting utilities.""" """Startup configuration box formatting utilities."""
import os import os
from sys import stderr
from typing import TYPE_CHECKING
from paskia._version import __version__
if TYPE_CHECKING:
from paskia.config import PaskiaConfig
BOX_WIDTH = 60 # Inner width (excluding box chars) BOX_WIDTH = 60 # Inner width (excluding box chars)
@@ -20,23 +27,18 @@ def bottom() -> str:
return "" + "" * (BOX_WIDTH + 2) + "\n" return "" + "" * (BOX_WIDTH + 2) + "\n"
def print_startup_config(passkey_instance, args, host, port, uds) -> None: def print_startup_config(config: "PaskiaConfig") -> None:
"""Print server configuration on startup.""" """Print server configuration on startup."""
from sys import stderr
from paskia._version import __version__
lines = [top()] lines = [top()]
lines.append(line(" ▄▄▄▄▄")) lines.append(line(" ▄▄▄▄▄"))
lines.append(line("█ █ Paskia " + __version__)) lines.append(line("█ █ Paskia " + __version__))
lines.append(line("█ █▄▄▄▄▄▄▄▄▄▄▄▄")) lines.append(line("█ █▄▄▄▄▄▄▄▄▄▄▄▄"))
lines.append(line("█ █▀▀▀▀█▀▀█▀▀█")) lines.append(line("█ █▀▀▀▀█▀▀█▀▀█ " + config.site_url + config.site_path))
lines.append(line(" ▀▀▀▀▀")) lines.append(line(" ▀▀▀▀▀"))
# Format auth host section # Format auth host section
auth_host = getattr(args, "auth_host", None) if config.auth_host:
if auth_host: lines.append(line(f"Auth Host: {config.auth_host}"))
lines.append(line(f"Auth Host: {auth_host}"))
# Show frontend URL if in dev mode # Show frontend URL if in dev mode
devmode = os.environ.get("PASKIA_DEVMODE") devmode = os.environ.get("PASKIA_DEVMODE")
@@ -44,24 +46,24 @@ def print_startup_config(passkey_instance, args, host, port, uds) -> None:
lines.append(line(f"Dev Frontend: {devmode}")) lines.append(line(f"Dev Frontend: {devmode}"))
# Format listen address with scheme # Format listen address with scheme
if uds: if config.uds:
listen = f"unix:{uds}" listen = f"unix:{config.uds}"
elif host: elif config.host:
listen = f"http://{host}:{port}" listen = f"http://{config.host}:{config.port}"
else: else:
listen = f"http://0.0.0.0:{port} + [::]:{port}" listen = f"http://0.0.0.0:{config.port} + [::]:{config.port}"
lines.append(line(f"Backend: {listen}")) lines.append(line(f"Backend: {listen}"))
# Relying Party line (omit name if same as id) # Relying Party line (omit name if same as id)
rp_id = passkey_instance.rp_id rp_id = config.rp_id
rp_name = passkey_instance.rp_name rp_name = config.rp_name
if rp_name and rp_name != rp_id: if rp_name and rp_name != rp_id:
lines.append(line(f"Relying Party: {rp_id} ({rp_name})")) lines.append(line(f"Relying Party: {rp_id} ({rp_name})"))
else: else:
lines.append(line(f"Relying Party: {rp_id}")) lines.append(line(f"Relying Party: {rp_id}"))
# Format origins section # Format origins section
allowed = passkey_instance.allowed_origins allowed = config.origins
if allowed: if allowed:
lines.append(line("Permitted Origins:")) lines.append(line("Permitted Origins:"))
for origin in sorted(allowed): for origin in sorted(allowed):