365 lines
12 KiB
Python
365 lines
12 KiB
Python
#!/usr/bin/env python3
|
||
"""
|
||
URL Shortener — stdlib-only, Python 3.8+
|
||
Usage: python urlshort.py <config.json>
|
||
|
||
API:
|
||
GET / 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 (API key required)
|
||
DELETE /api/urls/<code> Delete a short URL (API key required)
|
||
GET /<code> Redirect to original URL
|
||
"""
|
||
|
||
import json
|
||
import os
|
||
import random
|
||
import sqlite3
|
||
import string
|
||
import sys
|
||
import time
|
||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||
from urllib.parse import parse_qs, urlparse
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Config
|
||
# ---------------------------------------------------------------------------
|
||
|
||
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"):
|
||
if key not in cfg:
|
||
raise ValueError(f"Missing required config key: '{key}'")
|
||
|
||
cfg.setdefault("host", "0.0.0.0")
|
||
cfg.setdefault("port", 8080)
|
||
cfg.setdefault("db_path", "data/urlshort.db")
|
||
return cfg
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Database
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def init_db(db_path: str) -> None:
|
||
parent = os.path.dirname(db_path)
|
||
if parent:
|
||
os.makedirs(parent, exist_ok=True)
|
||
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.execute("""
|
||
CREATE TABLE IF NOT EXISTS urls (
|
||
short_code TEXT PRIMARY KEY,
|
||
original_url TEXT NOT NULL,
|
||
created_at INTEGER NOT NULL,
|
||
visit_count INTEGER NOT NULL DEFAULT 0
|
||
)
|
||
""")
|
||
conn.commit()
|
||
|
||
|
||
def db_connect(db_path: str) -> sqlite3.Connection:
|
||
conn = sqlite3.connect(db_path)
|
||
conn.row_factory = sqlite3.Row
|
||
return conn
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Helpers
|
||
# ---------------------------------------------------------------------------
|
||
|
||
_CHARS = string.ascii_letters + string.digits
|
||
|
||
|
||
def generate_code(length: int) -> str:
|
||
return "".join(random.choices(_CHARS, k=length))
|
||
|
||
|
||
def is_valid_url(url: str) -> bool:
|
||
try:
|
||
p = urlparse(url)
|
||
return p.scheme in ("http", "https") and bool(p.netloc)
|
||
except Exception:
|
||
return False
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# HTTP Handler
|
||
# ---------------------------------------------------------------------------
|
||
|
||
class Handler(BaseHTTPRequestHandler):
|
||
"""Single handler that serves the whole URL shortener API."""
|
||
|
||
# Injected by main() before the server starts.
|
||
cfg: dict = {}
|
||
|
||
# ------------------------------------------------------------------
|
||
# Logging
|
||
# ------------------------------------------------------------------
|
||
|
||
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"
|
||
)
|
||
|
||
# ------------------------------------------------------------------
|
||
# Low-level response helpers
|
||
# ------------------------------------------------------------------
|
||
|
||
def _send_json(self, status: int, payload: object) -> None:
|
||
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
|
||
self.send_response(status)
|
||
self.send_header("Content-Type", "application/json; charset=utf-8")
|
||
self.send_header("Content-Length", str(len(body)))
|
||
self.end_headers()
|
||
self.wfile.write(body)
|
||
|
||
def _error(self, status: int, message: str) -> None:
|
||
self._send_json(status, {"error": message})
|
||
|
||
def _redirect(self, location: str) -> None:
|
||
self.send_response(302)
|
||
self.send_header("Location", location)
|
||
self.send_header("Content-Length", "0")
|
||
self.end_headers()
|
||
|
||
# ------------------------------------------------------------------
|
||
# Auth
|
||
# ------------------------------------------------------------------
|
||
|
||
def _authenticated(self) -> bool:
|
||
"""Accept key via Authorization: Bearer <key> header or ?api_key=."""
|
||
auth = self.headers.get("Authorization", "")
|
||
if auth.lower().startswith("bearer "):
|
||
return auth[7:].strip() == self.cfg["api_key"]
|
||
|
||
qs = parse_qs(urlparse(self.path).query)
|
||
return qs.get("api_key", [""])[0] == self.cfg["api_key"]
|
||
|
||
# ------------------------------------------------------------------
|
||
# Body
|
||
# ------------------------------------------------------------------
|
||
|
||
def _read_json(self):
|
||
length = int(self.headers.get("Content-Length", 0))
|
||
if length == 0:
|
||
return {}
|
||
raw = self.rfile.read(length)
|
||
try:
|
||
return json.loads(raw)
|
||
except json.JSONDecodeError:
|
||
return None # signals parse failure to caller
|
||
|
||
# ------------------------------------------------------------------
|
||
# Route dispatch
|
||
# ------------------------------------------------------------------
|
||
|
||
def do_GET(self): # noqa: N802
|
||
path = urlparse(self.path).path.rstrip("/") or "/"
|
||
|
||
if path == "/":
|
||
self._send_json(200, {"status": "ok", "service": "url-shortener"})
|
||
|
||
elif path == "/api/urls":
|
||
self._handle_list_urls()
|
||
|
||
elif path.startswith("/api/urls/"):
|
||
code = path[len("/api/urls/"):]
|
||
self._handle_get_url(code)
|
||
|
||
else:
|
||
code = path.lstrip("/")
|
||
self._handle_redirect(code)
|
||
|
||
def do_POST(self): # noqa: N802
|
||
path = urlparse(self.path).path.rstrip("/")
|
||
|
||
if path == "/api/shorten":
|
||
self._handle_shorten()
|
||
else:
|
||
self._error(404, "Not found")
|
||
|
||
def do_DELETE(self): # noqa: N802
|
||
path = urlparse(self.path).path.rstrip("/")
|
||
|
||
if path.startswith("/api/urls/"):
|
||
code = path[len("/api/urls/"):]
|
||
self._handle_delete_url(code)
|
||
else:
|
||
self._error(404, "Not found")
|
||
|
||
# ------------------------------------------------------------------
|
||
# Handlers
|
||
# ------------------------------------------------------------------
|
||
|
||
def _handle_shorten(self) -> None:
|
||
if not self._authenticated():
|
||
self._error(401, "Unauthorized")
|
||
return
|
||
|
||
body = self._read_json()
|
||
if body is None:
|
||
self._error(400, "Invalid JSON body")
|
||
return
|
||
|
||
original_url = str(body.get("url", "")).strip()
|
||
if not original_url:
|
||
self._error(400, "Missing required field: url")
|
||
return
|
||
if not is_valid_url(original_url):
|
||
self._error(400, "Invalid URL — must start with http:// or https://")
|
||
return
|
||
|
||
custom_code = str(body.get("custom_code", "")).strip()
|
||
|
||
with db_connect(self.cfg["db_path"]) as conn:
|
||
if custom_code:
|
||
row = conn.execute(
|
||
"SELECT short_code FROM urls WHERE short_code = ?", (custom_code,)
|
||
).fetchone()
|
||
if row:
|
||
self._error(409, "Custom code already in use")
|
||
return
|
||
short_code = custom_code
|
||
else:
|
||
short_code = self._unique_code(conn)
|
||
if short_code is None:
|
||
self._error(500, "Could not generate a unique short code — try again")
|
||
return
|
||
|
||
created_at = int(time.time())
|
||
conn.execute(
|
||
"INSERT INTO urls (short_code, original_url, created_at) VALUES (?, ?, ?)",
|
||
(short_code, original_url, created_at),
|
||
)
|
||
conn.commit()
|
||
|
||
short_url = f"{self.cfg['base_url'].rstrip('/')}/{short_code}"
|
||
self._send_json(201, {
|
||
"short_code": short_code,
|
||
"short_url": short_url,
|
||
"original_url": original_url,
|
||
"created_at": created_at,
|
||
})
|
||
|
||
def _handle_list_urls(self) -> None:
|
||
if not self._authenticated():
|
||
self._error(401, "Unauthorized")
|
||
return
|
||
|
||
with db_connect(self.cfg["db_path"]) as conn:
|
||
rows = conn.execute(
|
||
"SELECT * FROM urls ORDER BY created_at DESC"
|
||
).fetchall()
|
||
|
||
base = self.cfg["base_url"].rstrip("/")
|
||
urls = [
|
||
{**dict(r), "short_url": f"{base}/{r['short_code']}"}
|
||
for r in rows
|
||
]
|
||
self._send_json(200, {"count": len(urls), "urls": urls})
|
||
|
||
def _handle_get_url(self, code: str) -> None:
|
||
if not self._authenticated():
|
||
self._error(401, "Unauthorized")
|
||
return
|
||
|
||
with db_connect(self.cfg["db_path"]) as conn:
|
||
row = conn.execute(
|
||
"SELECT * FROM urls WHERE short_code = ?", (code,)
|
||
).fetchone()
|
||
|
||
if row is None:
|
||
self._error(404, "Short code not found")
|
||
return
|
||
|
||
base = self.cfg["base_url"].rstrip("/")
|
||
self._send_json(200, {**dict(row), "short_url": f"{base}/{code}"})
|
||
|
||
def _handle_delete_url(self, code: str) -> None:
|
||
if not self._authenticated():
|
||
self._error(401, "Unauthorized")
|
||
return
|
||
|
||
with db_connect(self.cfg["db_path"]) as conn:
|
||
row = conn.execute(
|
||
"SELECT short_code FROM urls WHERE short_code = ?", (code,)
|
||
).fetchone()
|
||
if row is None:
|
||
self._error(404, "Short code not found")
|
||
return
|
||
conn.execute("DELETE FROM urls WHERE short_code = ?", (code,))
|
||
conn.commit()
|
||
|
||
self._send_json(200, {"message": f"Deleted '{code}'"})
|
||
|
||
def _handle_redirect(self, code: str) -> None:
|
||
with db_connect(self.cfg["db_path"]) as conn:
|
||
row = conn.execute(
|
||
"SELECT original_url FROM urls WHERE short_code = ?", (code,)
|
||
).fetchone()
|
||
if row is None:
|
||
self._error(404, "Short code not found")
|
||
return
|
||
conn.execute(
|
||
"UPDATE urls SET visit_count = visit_count + 1 WHERE short_code = ?",
|
||
(code,),
|
||
)
|
||
conn.commit()
|
||
|
||
self._redirect(row["original_url"])
|
||
|
||
# ------------------------------------------------------------------
|
||
# Internal utils
|
||
# ------------------------------------------------------------------
|
||
|
||
def _unique_code(self, conn: sqlite3.Connection, attempts: int = 10):
|
||
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
|
||
return None
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Entry point
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def main() -> None:
|
||
if len(sys.argv) != 2:
|
||
print(f"Usage: {sys.argv[0]} <config.json>", file=sys.stderr)
|
||
sys.exit(1)
|
||
|
||
cfg = load_config(sys.argv[1])
|
||
init_db(cfg["db_path"])
|
||
|
||
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"DB path : {cfg['db_path']}", file=sys.stderr)
|
||
|
||
try:
|
||
server.serve_forever()
|
||
except KeyboardInterrupt:
|
||
print("\nShutting down.", file=sys.stderr)
|
||
server.shutdown()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|
||
|