diff --git a/deploy/convert_server.py b/deploy/convert_server.py index f5515ca..bfc6520 100755 --- a/deploy/convert_server.py +++ b/deploy/convert_server.py @@ -163,7 +163,11 @@ class Handler(BaseHTTPRequestHandler): doc_id = parsed.path.split("/")[-1] supps = query("SELECT id, contract_id FROM supplements WHERE document_id = %s", (doc_id,)) for s in (supps or []): - execute("DELETE FROM spec_current WHERE contract_id = %s", (s["contract_id"],)) + execute( + """DELETE FROM spec_current WHERE contract_id = %s + AND last_event_id IN (SELECT id FROM spec_events WHERE supplement_id = %s)""", + (s["contract_id"], s["id"]), + ) execute("DELETE FROM spec_events WHERE supplement_id = %s", (s["id"],)) execute("DELETE FROM supplements WHERE id = %s", (s["id"],)) execute("DELETE FROM documents WHERE id = %s", (doc_id,)) @@ -177,6 +181,10 @@ class Handler(BaseHTTPRequestHandler): length = int(self.headers.get("Content-Length", 0)) body = json.loads(self.rfile.read(length)) if length > 0 else {} keep_ids = set(body.get("keep_ids", [])) + # Prevent DoS: too many IDs + if len(keep_ids) > 1000: + self._json({"ok": False, "error": "too many keep_ids (max 1000)"}, 400) + return docs = query("SELECT id FROM documents", ()) deleted = 0 for d in (docs or []): @@ -184,7 +192,11 @@ class Handler(BaseHTTPRequestHandler): continue supps = query("SELECT id, contract_id FROM supplements WHERE document_id = %s", (d["id"],)) for s in (supps or []): - execute("DELETE FROM spec_current WHERE contract_id = %s", (s["contract_id"],)) + execute( + """DELETE FROM spec_current WHERE contract_id = %s + AND last_event_id IN (SELECT id FROM spec_events WHERE supplement_id = %s)""", + (s["contract_id"], s["id"]), + ) execute("DELETE FROM spec_events WHERE supplement_id = %s", (s["id"],)) execute("DELETE FROM supplements WHERE id = %s", (s["id"],)) execute("DELETE FROM documents WHERE id = %s", (d["id"],)) diff --git a/deploy/db/supplements.py b/deploy/db/supplements.py index 33652b1..41fa40e 100644 --- a/deploy/db/supplements.py +++ b/deploy/db/supplements.py @@ -28,25 +28,31 @@ def get(supp_id): def delete_by_document(contract_id, filename): - """Delete supplement+document by contract+filename (cascade: spec_events first).""" - row = query( + """Delete ALL supplements+documents by contract+filename (cascade: spec_events first). + Handles duplicates from previously failed uploads.""" + rows = query( """SELECT s.id as sid, s.document_id FROM supplements s JOIN documents d ON d.id = s.document_id WHERE s.contract_id = %s AND d.filename = %s""", (contract_id, filename), ) - if row: - r = row[0] - # 1. Delete spec_current for this contract (FK to spec_events) - execute("DELETE FROM spec_current WHERE contract_id = %s", (contract_id,)) + if not rows: + return False + + for r in rows: + # 1. Delete spec_current rows referencing this supplement's events + execute( + """DELETE FROM spec_current WHERE contract_id = %s + AND last_event_id IN (SELECT id FROM spec_events WHERE supplement_id = %s)""", + (contract_id, r["sid"]), + ) # 2. Delete spec_events referencing this supplement execute("DELETE FROM spec_events WHERE supplement_id = %s", (r["sid"],)) # 3. Delete supplement execute("DELETE FROM supplements WHERE id = %s", (r["sid"],)) # 4. Delete document execute("DELETE FROM documents WHERE id = %s", (r["document_id"],)) - return True - return False + return True def delete(supp_id):