app refactored

This commit is contained in:
2026-03-18 02:39:33 +08:00
parent 6605224e30
commit db5aefedeb
4 changed files with 467 additions and 92 deletions
+264 -41
View File
@@ -4,24 +4,43 @@ URL Shortener — stdlib-only, Python 3.8+
Usage: python urlshort.py <config.json>
API (all routes are prefixed with base_path, e.g. /s):
GET / Health check
GET / Frontend (or health check if no static/)
GET /api/health Health check
POST /api/shorten Create a short URL (API key required)
GET /api/urls List all short URLs (API key required)
GET /api/urls/<code> Get info for a code (no API key required)
GET /api/lookup?url=<url> Look up by original URL (no API key required)
DELETE /api/urls/<code> Delete a short URL (API key required)
GET /<code> Redirect to original URL
"""
import json
import logging
import os
import random
import sqlite3
import string
import sys
import time
import collections
import mimetypes
import re
import threading
from http.server import BaseHTTPRequestHandler, HTTPServer
from urllib.parse import parse_qs, urlparse
# ---------------------------------------------------------------------------
# Logging — outputs to console (stdout); Docker captures it automatically.
# ---------------------------------------------------------------------------
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
stream=sys.stdout,
)
log = logging.getLogger("urlshort")
# ---------------------------------------------------------------------------
# Config
@@ -31,19 +50,26 @@ def load_config(path: str) -> dict:
with open(path, "r", encoding="utf-8") as fh:
cfg = json.load(fh)
for key in ("base_url", "short_length", "api_key"):
for key in ("base_url", "api_key"):
if key not in cfg:
raise ValueError(f"Missing required config key: '{key}'")
# Backward compat: old 'short_length' → 'min_short_length'
if "short_length" in cfg and "min_short_length" not in cfg:
cfg["min_short_length"] = cfg["short_length"]
cfg.setdefault("host", "0.0.0.0")
cfg.setdefault("port", 8080)
cfg.setdefault("db_path", "data/urlshort.db")
cfg.setdefault("retention_days", 0)
cfg.setdefault("min_short_length", 6)
cfg.setdefault("max_short_length", 32)
cfg.setdefault("max_url_length", 2048)
cfg.setdefault("max_retention_days", 3650)
cfg.setdefault("rate_limit_requests", 60)
cfg.setdefault("rate_limit_window", 60)
# Derive base_path from the path component of base_url.
# e.g. "http://example.com/s" → "/s"
# "http://example.com/s/" → "/s"
# "http://example.com" → ""
raw = urlparse(cfg["base_url"]).path.strip("/")
cfg["base_path"] = f"/{raw}" if raw else ""
@@ -118,6 +144,53 @@ def is_valid_url(url: str) -> bool:
return False
# ---------------------------------------------------------------------------
# Validation helpers
# ---------------------------------------------------------------------------
_CODE_RE = re.compile(r'^[A-Za-z0-9]+$')
def is_valid_code(code: str, max_length: int = 32) -> bool:
"""Short codes must be alphanumeric and within length limits."""
return bool(code) and len(code) <= max_length and bool(_CODE_RE.match(code))
# ---------------------------------------------------------------------------
# Static file serving
# ---------------------------------------------------------------------------
STATIC_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "static")
# ---------------------------------------------------------------------------
# Rate limiter (sliding window, per IP)
# ---------------------------------------------------------------------------
class RateLimiter:
"""Simple in-memory sliding-window rate limiter."""
def __init__(self, max_requests: int = 60, window: int = 60):
self.max_requests = max_requests
self.window = window
self._hits: dict = collections.defaultdict(list)
self._lock = threading.Lock()
def is_allowed(self, ip: str) -> bool:
now = time.time()
cutoff = now - self.window
with self._lock:
hits = self._hits[ip]
self._hits[ip] = hits = [t for t in hits if t > cutoff]
if len(hits) >= self.max_requests:
return False
hits.append(now)
return True
_rate_limiter = RateLimiter()
# ---------------------------------------------------------------------------
# HTTP Handler
# ---------------------------------------------------------------------------
@@ -132,6 +205,20 @@ class Handler(BaseHTTPRequestHandler):
# Routing helpers
# ------------------------------------------------------------------
def _client_ip(self) -> str:
"""Get the real client IP, checking reverse-proxy headers first.
nginx adds X-Real-IP and X-Forwarded-For via proxy_set_header.
Without these headers, falls back to the TCP connection source.
"""
ip = self.headers.get("X-Real-IP", "").strip()
if ip:
return ip
xff = self.headers.get("X-Forwarded-For", "").strip()
if xff:
return xff.split(",")[0].strip()
return self.client_address[0]
def _local_path(self):
"""Return the request path with base_path prefix stripped.
@@ -155,11 +242,7 @@ class Handler(BaseHTTPRequestHandler):
# ------------------------------------------------------------------
def log_message(self, fmt, *args): # noqa: N802 stdlib override
sys.stderr.write(
f"[{self.log_date_time_string()}] {self.address_string()} - "
+ (fmt % args)
+ "\n"
)
log.info("%s - %s", self._client_ip(), fmt % args)
# ------------------------------------------------------------------
# Low-level response helpers
@@ -195,6 +278,20 @@ class Handler(BaseHTTPRequestHandler):
self.send_header("Content-Length", "0")
self.end_headers()
# ------------------------------------------------------------------
# CORS & Security headers (injected into every response)
# ------------------------------------------------------------------
def end_headers(self):
self.send_header("Access-Control-Allow-Origin", "*")
self.send_header("Access-Control-Allow-Methods", "GET, POST, DELETE, OPTIONS")
self.send_header("Access-Control-Allow-Headers", "Content-Type")
self.send_header("Access-Control-Max-Age", "86400")
self.send_header("X-Content-Type-Options", "nosniff")
self.send_header("X-Frame-Options", "DENY")
self.send_header("X-XSS-Protection", "1; mode=block")
super().end_headers()
# ------------------------------------------------------------------
# Auth — purely api_key based (query param or JSON body)
# ------------------------------------------------------------------
@@ -232,12 +329,31 @@ class Handler(BaseHTTPRequestHandler):
def do_GET(self): # noqa: N802
cleanup_expired(self.cfg["db_path"])
if not _rate_limiter.is_allowed(self._client_ip()):
self._error(429, "Too many requests")
return
# Redirect bare base path to base path + / for correct relative URLs
raw = urlparse(self.path).path
base = self.cfg.get("base_path", "")
if base and raw == base:
self._redirect(base + "/")
return
path = self._local_path()
if path is None:
self._send_empty(404)
return
if path == "/":
# Serve frontend if static/index.html exists, else health check
index_path = os.path.join(STATIC_DIR, "index.html")
if os.path.isfile(index_path):
self._serve_static("index.html")
else:
self._send_json(200, {"status": "ok", "service": "url-shortener"})
elif path == "/api/health":
self._send_json(200, {"status": "ok", "service": "url-shortener"})
elif path == "/api/urls":
@@ -247,6 +363,13 @@ class Handler(BaseHTTPRequestHandler):
code = path[len("/api/urls/"):]
self._handle_get_url(code)
elif path == "/api/lookup":
self._handle_lookup()
elif path == "/static" or path.startswith("/static/"):
rel = path[len("/static"):].lstrip("/") or "index.html"
self._serve_static(rel)
else:
code = path.lstrip("/")
self._handle_redirect(code)
@@ -254,6 +377,10 @@ class Handler(BaseHTTPRequestHandler):
def do_POST(self): # noqa: N802
cleanup_expired(self.cfg["db_path"])
if not _rate_limiter.is_allowed(self._client_ip()):
self._error(429, "Too many requests")
return
path = self._local_path()
if path is None:
self._send_empty(404)
@@ -267,6 +394,10 @@ class Handler(BaseHTTPRequestHandler):
def do_DELETE(self): # noqa: N802
cleanup_expired(self.cfg["db_path"])
if not _rate_limiter.is_allowed(self._client_ip()):
self._error(429, "Too many requests")
return
path = self._local_path()
if path is None:
self._send_empty(404)
@@ -278,6 +409,14 @@ class Handler(BaseHTTPRequestHandler):
else:
self._send_empty(404)
def do_OPTIONS(self): # noqa: N802
"""Handle CORS preflight requests."""
self.send_response(200)
self.send_header("Content-Length", "0")
self.send_header("Allow", "GET, POST, DELETE, OPTIONS")
self.end_headers()
self.wfile.flush()
# ------------------------------------------------------------------
# Handlers
# ------------------------------------------------------------------
@@ -303,19 +442,10 @@ class Handler(BaseHTTPRequestHandler):
if not is_valid_url(original_url):
self._error(400, "Invalid URL — must start with http:// or https://")
return
# short_length: query param, then body, then config default
raw_sl = qs.get("short_length", [None])[0]
if raw_sl is None:
raw_sl = body.get("short_length")
if raw_sl is not None:
try:
short_length = int(raw_sl)
except (ValueError, TypeError):
self._error(400, "Invalid short_length")
return
else:
short_length = self.cfg["short_length"]
max_url = self.cfg["max_url_length"]
if len(original_url) > max_url:
self._error(400, f"URL too long (max {max_url} characters)")
return
# retention_days: query param, then body, then config default
raw_rd = qs.get("retention_days", [None])[0]
@@ -330,8 +460,13 @@ class Handler(BaseHTTPRequestHandler):
else:
retention_days = self.cfg.get("retention_days", 0)
max_ret = self.cfg["max_retention_days"]
if not (0 <= retention_days <= max_ret):
self._error(400, f"retention_days must be 0{max_ret}")
return
with db_connect(self.cfg["db_path"]) as conn:
short_code = self._unique_code(conn, short_length)
short_code = self._unique_code(conn)
if short_code is None:
self._error(500, "Could not generate a unique short code — try again")
return
@@ -367,11 +502,16 @@ class Handler(BaseHTTPRequestHandler):
"original_url": r["original_url"],
"created_at": r["created_at"],
"visit_count": r["visit_count"],
"retention_days": r["retention_days"],
})
self._send_json(200, {"count": len(urls), "urls": urls})
def _handle_get_url(self, code: str) -> None:
# No API key required
if not is_valid_code(code, self.cfg["max_short_length"]):
self._send_empty(404)
return
with db_connect(self.cfg["db_path"]) as conn:
row = conn.execute(
"SELECT * FROM urls WHERE short_code = ?", (code,)
@@ -388,6 +528,7 @@ class Handler(BaseHTTPRequestHandler):
"original_url": row["original_url"],
"created_at": row["created_at"],
"visit_count": row["visit_count"],
"retention_days": row["retention_days"],
})
def _handle_delete_url(self, code: str) -> None:
@@ -395,6 +536,9 @@ class Handler(BaseHTTPRequestHandler):
if not self._check_api_key():
self._send_empty(403)
return
if not is_valid_code(code, self.cfg["max_short_length"]):
self._send_empty(404)
return
with db_connect(self.cfg["db_path"]) as conn:
row = conn.execute(
@@ -411,6 +555,10 @@ class Handler(BaseHTTPRequestHandler):
self._send_empty(204)
def _handle_redirect(self, code: str) -> None:
if not is_valid_code(code, self.cfg["max_short_length"]):
self._send_empty(404)
return
with db_connect(self.cfg["db_path"]) as conn:
row = conn.execute(
"SELECT original_url FROM urls WHERE short_code = ?", (code,)
@@ -426,20 +574,85 @@ class Handler(BaseHTTPRequestHandler):
self._redirect(row["original_url"])
def _handle_lookup(self) -> None:
"""Look up a URL by its original URL. No API key required."""
qs = parse_qs(urlparse(self.path).query)
url = qs.get("url", [""])[0].strip()
if not url:
self._error(400, "Missing required parameter: url")
return
max_url = self.cfg["max_url_length"]
if len(url) > max_url:
self._error(400, f"URL too long (max {max_url} characters)")
return
with db_connect(self.cfg["db_path"]) as conn:
row = conn.execute(
"SELECT * FROM urls WHERE original_url = ? ORDER BY created_at DESC LIMIT 1",
(url,),
).fetchone()
if row is None:
self._send_empty(404)
return
base = self.cfg["base_url"].rstrip("/")
self._send_json(200, {
"short_code": row["short_code"],
"short_url": f"{base}/{row['short_code']}",
"original_url": row["original_url"],
"created_at": row["created_at"],
"visit_count": row["visit_count"],
"retention_days": row["retention_days"],
})
def _serve_static(self, rel_path: str) -> None:
"""Serve a static file from STATIC_DIR (for local dev; nginx in prod)."""
if not rel_path:
rel_path = "index.html"
# Prevent directory traversal
safe = os.path.normpath(rel_path)
if safe.startswith("..") or os.path.isabs(safe):
self._send_empty(403)
return
fpath = os.path.join(STATIC_DIR, safe)
if not os.path.isfile(fpath):
self._send_empty(404)
return
mime, _ = mimetypes.guess_type(fpath)
if not mime:
mime = "application/octet-stream"
with open(fpath, "rb") as f:
data = f.read()
self.send_response(200)
self.send_header("Content-Type", mime)
self.send_header("Content-Length", str(len(data)))
self.send_header("Cache-Control", "public, max-age=3600")
self.end_headers()
self.wfile.write(data)
# ------------------------------------------------------------------
# Internal utils
# ------------------------------------------------------------------
def _unique_code(self, conn: sqlite3.Connection, length: int = None, attempts: int = 10):
if length is None:
length = self.cfg["short_length"]
for _ in range(attempts):
code = generate_code(length)
exists = conn.execute(
"SELECT 1 FROM urls WHERE short_code = ?", (code,)
).fetchone()
if not exists:
return code
def _unique_code(self, conn: sqlite3.Connection, attempts_per_length: int = 10):
"""Generate a unique short code with progressive length increment.
Starts at min_short_length and tries `attempts_per_length` random
codes at each length. On exhaustion, increments the length by 1
and repeats, up to max_short_length. Returns None only when the
entire range is exhausted (extremely unlikely).
"""
min_len = self.cfg["min_short_length"]
max_len = self.cfg["max_short_length"]
for length in range(min_len, max_len + 1):
for _ in range(attempts_per_length):
code = generate_code(length)
exists = conn.execute(
"SELECT 1 FROM urls WHERE short_code = ?", (code,)
).fetchone()
if not exists:
return code
return None
@@ -448,29 +661,39 @@ class Handler(BaseHTTPRequestHandler):
# ---------------------------------------------------------------------------
def main() -> None:
global _rate_limiter
if len(sys.argv) != 2:
print(f"Usage: {sys.argv[0]} <config.json>", file=sys.stderr)
log.error("Usage: %s <config.json>", sys.argv[0])
sys.exit(1)
cfg = load_config(sys.argv[1])
init_db(cfg["db_path"])
cleanup_expired(cfg["db_path"])
# Re-initialize rate limiter with config values
_rate_limiter = RateLimiter(
max_requests=cfg["rate_limit_requests"],
window=cfg["rate_limit_window"],
)
Handler.cfg = cfg
host, port = cfg["host"], int(cfg["port"])
server = HTTPServer((host, port), Handler)
print(f"URL Shortener listening on http://{host}:{port}", file=sys.stderr)
print(f"Base URL : {cfg['base_url']}", file=sys.stderr)
print(f"Base path : {cfg['base_path'] or '/'}", file=sys.stderr)
print(f"DB path : {cfg['db_path']}", file=sys.stderr)
print(f"Retention days : {cfg['retention_days']}", file=sys.stderr)
log.info("URL Shortener listening on http://%s:%s", host, port)
log.info("Base URL : %s", cfg["base_url"])
log.info("Base path : %s", cfg["base_path"] or "/")
log.info("DB path : %s", cfg["db_path"])
log.info("Short codes : %s%s chars", cfg["min_short_length"], cfg["max_short_length"])
log.info("Retention days : %s", cfg["retention_days"])
log.info("Rate limit : %s req/%ss per IP", cfg["rate_limit_requests"], cfg["rate_limit_window"])
try:
server.serve_forever()
except KeyboardInterrupt:
print("\nShutting down.", file=sys.stderr)
log.info("Shutting down.")
server.shutdown()