Files
sub-provider/app/main.py
2026-04-20 14:16:09 +08:00

649 lines
24 KiB
Python

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,
),
)