from __future__ import annotations import hmac import logging from urllib.parse import quote from fastapi import FastAPI, Form, HTTPException, Query, Request from fastapi.responses import HTMLResponse, RedirectResponse, Response from fastapi.templating import Jinja2Templates from starlette.middleware.sessions import SessionMiddleware from app.config import get_settings from app.models import RuleConfig, SourceConfig, SourceSnapshot from app.services.bundle_cache import build_bundle_cache_key, load_bundle_cache, save_bundle_cache from app.services.config_store import ( delete_source, get_profile_default_source_keys, list_profile_groups, list_profile_rule_bindings, list_profile_source_bindings, list_profiles, list_sources, load_profile_app_config_from_db, replace_profile_groups, save_profile, save_source, update_profile_source_bindings, update_profile_rule_bindings, ) from app.services.loader import load_app_config from app.services.profiles import build_bundle_profile, build_thin_profile, dump_yaml from app.services.rules import load_rule_text from app.services.subscriptions import ( build_merged_provider_document, build_provider_document, build_source_snapshots, dump_provider_yaml, get_first_quota, ) settings = get_settings() logger = logging.getLogger(__name__) app = FastAPI(title=settings.app_name) app.add_middleware( SessionMiddleware, secret_key=settings.admin_session_secret, session_cookie="sub_provider_admin", max_age=settings.admin_session_max_age, same_site="lax", https_only=False, ) app_config = load_app_config(settings.sources_file) PUBLIC_PREFIX = "/" + (app_config.public_path or settings.public_path).strip("/") templates = Jinja2Templates(directory=str(settings.config_dir.parent / "templates")) @app.get("/healthz") async def healthz() -> dict[str, str]: return {"status": "ok"} def _base_url(request: Request) -> str: if settings.public_base_url: return settings.public_base_url.rstrip("/") return str(request.base_url).rstrip("/") def _admin_auth_enabled() -> bool: return bool((settings.admin_token or "").strip()) def _is_admin_authenticated(request: Request) -> bool: return not _admin_auth_enabled() or request.session.get("admin_authenticated") is True def _set_flash(request: Request, message: str, level: str = "info") -> None: request.session["_flash"] = {"message": message, "level": level} def _pop_flash(request: Request): return request.session.pop("_flash", None) def _admin_context(request: Request, *, active: str, **extra): return { "request": request, "active": active, "flash": _pop_flash(request), "auth_enabled": _admin_auth_enabled(), "authenticated": _is_admin_authenticated(request), **extra, } def _admin_guard(request: Request) -> RedirectResponse | None: if _is_admin_authenticated(request): return None next_path = quote(str(request.url.path)) return _redirect(f"/admin/login?next={next_path}") def _current_app_config(profile_key: str | None = None): if profile_key: profile_config = load_profile_app_config_from_db(profile_key) if profile_config is not None: return profile_config return load_app_config(settings.sources_file) def _resolve_sources(sources: str | None, profile_key: str | None = None) -> list[tuple[str, SourceConfig]]: current_config = _current_app_config(profile_key=profile_key) enabled = [(name, src) for name, src in current_config.sources.items() if src.enabled and str(src.url).strip()] if not sources: if profile_key: default_keys = get_profile_default_source_keys(profile_key) if default_keys: selected = [(name, current_config.sources[name]) for name in default_keys if name in current_config.sources] logger.info("resolve_sources profile default: profile=%s selected=%s", profile_key, [name for name, _ in selected]) return selected logger.info("resolve_sources default: selected=%s", [name for name, _ in enabled]) return enabled names = [item.strip() for item in sources.split(",") if item.strip()] selected: list[tuple[str, SourceConfig]] = [] seen: set[str] = set() for name in names: if name in seen: continue source = current_config.sources.get(name) if source is None or not source.enabled or not str(source.url).strip(): raise HTTPException(status_code=404, detail=f"source not found or disabled: {name}") selected.append((name, source)) seen.add(name) if not selected: raise HTTPException(status_code=400, detail="no sources selected") logger.info("resolve_sources explicit: requested=%s selected=%s", sources, [name for name, _ in selected]) return selected def _rule_path(rule: RuleConfig): if not rule.file: raise HTTPException(status_code=404, detail="rule file not available") path = (settings.rules_dir / rule.file).resolve() if not path.is_file() or settings.rules_dir.resolve() not in path.parents: raise HTTPException(status_code=404, detail="rule file missing") return path async def _build_quota_headers(source_items: list[tuple[str, SourceConfig]]) -> dict[str, str]: headers: dict[str, str] = {} quota = await get_first_quota(source_items) if quota and not quota.is_empty(): headers["Subscription-Userinfo"] = quota.to_header_value() return headers def _quota_headers_from_snapshots(snapshots: list[SourceSnapshot]) -> dict[str, str]: if not snapshots: return {} quota = snapshots[0].quota if quota and not quota.is_empty(): return {"Subscription-Userinfo": quota.to_header_value()} return {} def _yaml_response(content: str, request: Request, headers: dict[str, str] | None = None, filename: str | None = None) -> Response: final_headers = { "Content-Type": "text/yaml; charset=utf-8", "Cache-Control": "no-store", } if headers: final_headers.update(headers) if filename: final_headers["Content-Disposition"] = f"attachment; filename*=UTF-8''{filename}" body = "" if request.method == "HEAD" else content return Response(content=body, media_type="text/yaml; charset=utf-8", headers=final_headers) @app.api_route(PUBLIC_PREFIX + "/providers/merged.yaml", methods=["GET", "HEAD"]) async def merged_provider(request: Request, sources: str | None = Query(default=None)) -> Response: source_items = _resolve_sources(sources) try: document = await build_merged_provider_document(source_items) except Exception as exc: # noqa: BLE001 logger.exception("merged_provider failed: sources=%s", [name for name, _ in source_items]) raise HTTPException(status_code=502, detail=f"failed to build merged provider: {exc}") from exc content = dump_provider_yaml(document) headers = await _build_quota_headers(source_items) return _yaml_response(content, request, headers=headers, filename="merged.yaml") @app.api_route(PUBLIC_PREFIX + "/providers/{name}.yaml", methods=["GET", "HEAD"]) async def provider(name: str, request: Request) -> Response: current_config = _current_app_config() source = current_config.sources.get(name) if source is None or not source.enabled: raise HTTPException(status_code=404, detail="provider not found") try: document = await build_provider_document(name, source) except Exception as exc: # noqa: BLE001 logger.exception("provider failed: source=%s", name) raise HTTPException(status_code=502, detail=f"failed to build provider: {exc}") from exc content = dump_provider_yaml(document) headers = await _build_quota_headers([(name, source)]) return _yaml_response(content, request, headers=headers, filename=f"{name}.yaml") @app.api_route(PUBLIC_PREFIX + "/rules/{name}.yaml", methods=["GET", "HEAD"]) async def rule_file(name: str, request: Request) -> Response: current_config = _current_app_config() rule = current_config.rules.get(name) if rule is None: raise HTTPException(status_code=404, detail="rule not found") content = load_rule_text(_rule_path(rule)) return _yaml_response(content, request, filename=f"{name}.yaml") @app.api_route(PUBLIC_PREFIX + "/clients/{client_type}.yaml", methods=["GET", "HEAD"]) async def client_profile(client_type: str, request: Request, sources: str | None = Query(default=None)) -> Response: current_config = _current_app_config(profile_key=client_type) client = current_config.clients.get(client_type) if client is None: raise HTTPException(status_code=404, detail="client config not found") source_items = _resolve_sources(sources, profile_key=client_type) content = dump_yaml( build_thin_profile( client_type=client_type, app_config=current_config, selected_source_names=[name for name, _ in source_items], base_url=_base_url(request), public_path=(current_config.public_path or settings.public_path).strip("/"), ) ) headers = {"profile-update-interval": str(client.provider_interval)} headers.update(await _build_quota_headers(source_items)) return _yaml_response(content, request, headers=headers, filename=f"{client_type}.yaml") @app.api_route(PUBLIC_PREFIX + "/bundle/{client_type}.yaml", methods=["GET", "HEAD"]) async def bundle_profile( client_type: str, request: Request, sources: str | None = Query(default=None), force_refresh: bool = Query(default=False), ) -> Response: current_config = _current_app_config(profile_key=client_type) client = current_config.clients.get(client_type) if client is None: raise HTTPException(status_code=404, detail="client config not found") source_items = _resolve_sources(sources, profile_key=client_type) cache_key = build_bundle_cache_key(client_type=client_type, source_names=[name for name, _ in source_items]) if not force_refresh: cached = load_bundle_cache( cache_dir=settings.bundle_cache_dir, cache_key=cache_key, ttl_seconds=settings.bundle_cache_ttl_seconds, ) if cached is not None: logger.info("bundle cache hit: client=%s sources=%s", client_type, [name for name, _ in source_items]) headers = { "profile-update-interval": str(client.provider_interval), "X-Sub-Provider-Bundle-Cache": "HIT", } headers.update(cached.headers) return _yaml_response(content=cached.content, request=request, headers=headers, filename=f"bundle-{client_type}.yaml") try: snapshots = await build_source_snapshots(source_items) except Exception as exc: # noqa: BLE001 logger.exception("bundle_profile failed: client=%s sources=%s", client_type, [name for name, _ in source_items]) raise HTTPException(status_code=502, detail=f"failed to build bundle: {exc}") from exc content = dump_yaml( build_bundle_profile( client_type=client_type, app_config=current_config, snapshots=snapshots, ) ) headers = { "profile-update-interval": str(client.provider_interval), "X-Sub-Provider-Bundle-Cache": "BYPASS" if force_refresh else "MISS", } headers.update(_quota_headers_from_snapshots(snapshots)) save_bundle_cache( cache_dir=settings.bundle_cache_dir, cache_key=cache_key, content=content, headers={key: value for key, value in headers.items() if key != "X-Sub-Provider-Bundle-Cache"}, ) return _yaml_response(content, request, headers=headers, filename=f"bundle-{client_type}.yaml") def _redirect(url: str) -> RedirectResponse: return RedirectResponse(url=url, status_code=303) @app.get("/admin/login", response_class=HTMLResponse) async def admin_login(request: Request, next: str = Query(default="/admin/sources")) -> HTMLResponse: if _is_admin_authenticated(request): return _redirect(next) return templates.TemplateResponse( request, "login.html", _admin_context(request, active="", next=next), ) @app.post("/admin/login") async def admin_login_submit( request: Request, token: str = Form(default=""), next: str = Form(default="/admin/sources"), ) -> Response: if not _admin_auth_enabled(): request.session["admin_authenticated"] = True return _redirect(next) if hmac.compare_digest((settings.admin_token or "").strip(), token.strip()): request.session["admin_authenticated"] = True _set_flash(request, "登录成功", "success") return _redirect(next) _set_flash(request, "Token 不正确", "error") return _redirect(f"/admin/login?next={quote(next)}") @app.post("/admin/logout") async def admin_logout(request: Request) -> Response: request.session.clear() return _redirect("/admin/login") @app.get("/admin", response_class=HTMLResponse) async def admin_index(request: Request) -> Response: redirect = _admin_guard(request) if redirect: return redirect return _redirect("/admin/sources") @app.get("/admin/sources", response_class=HTMLResponse) async def admin_sources(request: Request) -> HTMLResponse: redirect = _admin_guard(request) if redirect: return redirect return templates.TemplateResponse( request, "sources.html", _admin_context( request, active="sources", sources=list_sources(), profiles=list_profiles(), ), ) @app.post("/admin/sources") async def admin_save_source( request: Request, key: str = Form(...), url: str = Form(...), kind: str = Form(default="auto"), display_name: str = Form(default=""), prefix: str = Form(default=""), suffix: str = Form(default=""), include_regex: str = Form(default=""), exclude_regex: str = Form(default=""), cache_ttl_seconds: str = Form(default=""), enabled: str | None = Form(default=None), ) -> Response: redirect = _admin_guard(request) if redirect: return redirect save_source( key=key.strip(), enabled=enabled is not None, kind=kind, url=url.strip(), display_name=display_name.strip() or None, headers={}, include_regex=include_regex.strip() or None, exclude_regex=exclude_regex.strip() or None, prefix=prefix, suffix=suffix, cache_ttl_seconds=int(cache_ttl_seconds) if cache_ttl_seconds.strip() else None, ) _set_flash(request, f"Source {key.strip()} 已保存", "success") return _redirect("/admin/sources") @app.post("/admin/sources/{key}/delete") async def admin_delete_source(request: Request, key: str) -> Response: redirect = _admin_guard(request) if redirect: return redirect delete_source(key) _set_flash(request, f"Source {key} 已删除", "success") return _redirect("/admin/sources") @app.get("/admin/profiles", response_class=HTMLResponse) async def admin_profiles(request: Request) -> HTMLResponse: redirect = _admin_guard(request) if redirect: return redirect return templates.TemplateResponse( request, "profiles.html", _admin_context( request, active="profiles", profiles=list_profiles(), ), ) @app.post("/admin/profiles") async def admin_save_profile( request: Request, key: str = Form(...), title: str = Form(...), provider_interval: int = Form(default=21600), rule_interval: int = Form(default=86400), test_url: str = Form(...), test_interval: int = Form(default=300), main_policy: str = Form(...), source_policy: str = Form(...), mixed_auto_policy: str = Form(...), manual_policy: str = Form(...), direct_policy: str = Form(...), mode: str = Form(default="rule"), allow_lan: str | None = Form(default=None), ipv6: str | None = Form(default=None), mixed_port: str = Form(default=""), socks_port: str = Form(default=""), log_level: str = Form(default="info"), ) -> Response: redirect = _admin_guard(request) if redirect: return redirect save_profile( key=key.strip(), title=title.strip(), provider_interval=provider_interval, rule_interval=rule_interval, test_url=test_url.strip(), test_interval=test_interval, main_policy=main_policy.strip(), source_policy=source_policy.strip(), mixed_auto_policy=mixed_auto_policy.strip(), manual_policy=manual_policy.strip(), direct_policy=direct_policy.strip(), mode=mode.strip(), allow_lan=allow_lan is not None, ipv6=ipv6 is not None, mixed_port=int(mixed_port) if mixed_port.strip() else None, socks_port=int(socks_port) if socks_port.strip() else None, log_level=log_level.strip() or None, ) _set_flash(request, f"Profile {key.strip()} 已保存", "success") return _redirect("/admin/profiles") @app.get("/admin/profiles/{profile_key}/sources", response_class=HTMLResponse) async def admin_profile_sources(request: Request, profile_key: str) -> HTMLResponse: redirect = _admin_guard(request) if redirect: return redirect return templates.TemplateResponse( request, "profile_sources.html", _admin_context( request, active="profiles", profile_key=profile_key, bindings=list_profile_source_bindings(profile_key), ), ) @app.post("/admin/profiles/{profile_key}/sources") async def admin_update_profile_sources(request: Request, profile_key: str) -> Response: redirect = _admin_guard(request) if redirect: return redirect form = await request.form() keys = form.getlist("source_key") rows: list[dict] = [] for index, key in enumerate(keys): rows.append( { "key": key, "enabled": form.get(f"enabled_{key}") is not None, "order_index": form.get(f"order_index_{key}", str(index)), } ) update_profile_source_bindings(profile_key, rows) _set_flash(request, f"Profile {profile_key} 的默认源已更新", "success") return _redirect(f"/admin/profiles/{profile_key}/sources") @app.get("/admin/profiles/{profile_key}/rules", response_class=HTMLResponse) async def admin_profile_rules(request: Request, profile_key: str) -> HTMLResponse: redirect = _admin_guard(request) if redirect: return redirect return templates.TemplateResponse( request, "rules.html", _admin_context( request, active="profiles", profile_key=profile_key, bindings=list_profile_rule_bindings(profile_key), ), ) @app.get("/admin/profiles/{profile_key}/groups", response_class=HTMLResponse) async def admin_profile_groups(request: Request, profile_key: str) -> HTMLResponse: redirect = _admin_guard(request) if redirect: return redirect return templates.TemplateResponse( request, "groups.html", _admin_context( request, active="profiles", profile_key=profile_key, groups=list_profile_groups(profile_key), ), ) @app.post("/admin/profiles/{profile_key}/groups") async def admin_update_profile_groups(request: Request, profile_key: str) -> Response: redirect = _admin_guard(request) if redirect: return redirect form = await request.form() row_ids = form.getlist("group_id") rows: list[dict] = [] for index, group_id in enumerate(row_ids): group_key = str(index) name = str(form.get(f"group_name_{group_key}", "")).strip() if not str(name).strip(): continue proxies_raw = str(form.get(f"proxies_{group_key}", "")).strip() if form.get(f"delete_{group_key}") is not None: continue rows.append( { "id": int(group_id) if str(group_id).strip() else None, "name": str(name).strip(), "group_kind": str(form.get(f"group_kind_{group_key}", "policy")).strip(), "type": str(form.get(f"type_{group_key}", "select")).strip(), "order_index": int(str(form.get(f"order_index_{group_key}", index))), "proxies": [item.strip() for item in proxies_raw.splitlines() if item.strip()], "filter_regex": str(form.get(f"filter_regex_{group_key}", "")).strip(), "tolerance": int(str(form.get(f"tolerance_{group_key}", "")).strip()) if str(form.get(f"tolerance_{group_key}", "")).strip() else None, "url": str(form.get(f"url_{group_key}", "")).strip(), "interval": int(str(form.get(f"interval_{group_key}", "")).strip()) if str(form.get(f"interval_{group_key}", "")).strip() else None, "enabled": form.get(f"enabled_{group_key}") is not None, } ) new_name = str(form.get("new_group_name", "")).strip() if new_name: new_proxies_raw = str(form.get("new_group_proxies", "")).strip() rows.append( { "id": None, "name": new_name, "group_kind": str(form.get("new_group_kind", "policy")).strip(), "type": str(form.get("new_group_type", "select")).strip(), "order_index": int(str(form.get("new_group_order_index", len(rows))).strip() or len(rows)), "proxies": [item.strip() for item in new_proxies_raw.splitlines() if item.strip()], "filter_regex": str(form.get("new_group_filter_regex", "")).strip(), "tolerance": int(str(form.get("new_group_tolerance", "")).strip()) if str(form.get("new_group_tolerance", "")).strip() else None, "url": str(form.get("new_group_url", "")).strip(), "interval": int(str(form.get("new_group_interval", "")).strip()) if str(form.get("new_group_interval", "")).strip() else None, "enabled": form.get("new_group_enabled") is not None, } ) replace_profile_groups(profile_key, rows) _set_flash(request, f"Profile {profile_key} 的策略组已更新", "success") return _redirect(f"/admin/profiles/{profile_key}/groups") @app.post("/admin/profiles/{profile_key}/rules") async def admin_update_profile_rules(request: Request, profile_key: str) -> Response: redirect = _admin_guard(request) if redirect: return redirect form = await request.form() keys = form.getlist("rule_key") rows: list[dict] = [] for index, key in enumerate(keys): rows.append( { "key": key, "enabled": form.get(f"enabled_{key}") is not None, "order_index": form.get(f"order_index_{key}", str(index)), "policy": form.get(f"policy_{key}", ""), } ) update_profile_rule_bindings(profile_key, rows) _set_flash(request, f"Profile {profile_key} 的规则绑定已更新", "success") return _redirect(f"/admin/profiles/{profile_key}/rules") @app.get("/admin/profiles/{profile_key}/preview", response_class=HTMLResponse) async def admin_profile_preview( request: Request, profile_key: str, sources: str | None = Query(default=None), ) -> HTMLResponse: redirect = _admin_guard(request) if redirect: return redirect profile_config = _current_app_config(profile_key=profile_key) source_items = _resolve_sources(sources, profile_key=profile_key) selected_names = [name for name, _ in source_items] snapshots = await build_source_snapshots([(name, profile_config.sources[name]) for name in selected_names]) yaml_text = dump_yaml( build_bundle_profile( client_type=profile_key, app_config=profile_config, snapshots=snapshots, ) ) return templates.TemplateResponse( request, "preview.html", _admin_context( request, active="profiles", profile_key=profile_key, yaml_text=yaml_text, available_sources=list(profile_config.sources.keys()), selected_sources=selected_names, ), )