rewrite_proxy.py110 lines · main
1#!/usr/bin/env python3
2"""Phase-0 SuperTokens→Doltgres SQL rewrite proxy (spike only, not production)."""
3import asyncio, re, struct, sys
4LISTEN_PORT = int(sys.argv[1]) if len(sys.argv)>1 else 15432
5UPSTREAM_HOST = sys.argv[2] if len(sys.argv)>2 else "doltgres"
6UPSTREAM_PORT = int(sys.argv[3]) if len(sys.argv)>3 else 5432
7stats = {"rewrites": 0, "queries": 0}
8
9def rewrite_sql(sql: str) -> str:
10 orig = sql
11 sql = re.sub(
12 r"SET\s+SESSION\s+CHARACTERISTICS\s+AS\s+TRANSACTION\s+ISOLATION\s+LEVEL\s+READ\s+COMMITTED\s*;?",
13 "SET default_transaction_isolation TO 'read committed'", sql, flags=re.I)
14 sql = re.sub(r"CONSTRAINT\s+[A-Za-z0-9_]+(\s+UNIQUE\b)", r"\1", sql, flags=re.I)
15 sql = re.sub(r"CONSTRAINT\s+[A-Za-z0-9_]+(\s+CHECK\b)", r"\1", sql, flags=re.I)
16 sql = re.sub(r"\s+PARTITION\s+BY\s+RANGE\s*\([^)]*\)", "", sql, flags=re.I)
17 sql = re.sub(r"\s+PARTITION\s+BY\s+LIST\s*\([^)]*\)", "", sql, flags=re.I)
18 sql = re.sub(r"\s+PARTITION\s+BY\s+HASH\s*\([^)]*\)", "", sql, flags=re.I)
19 if re.search(r"\bPARTITION\s+OF\b", sql, re.I):
20 sql = "SELECT 1"
21 sql = re.sub(r"\s+USING\s+brin\b", "", sql, flags=re.I)
22 sql = re.sub(r"\bDROP\s+(TABLE|INDEX|VIEW)\s+(.+?)\s+CASCADE\b", r"DROP \1 \2", sql, flags=re.I)
23 if sql != orig:
24 stats["rewrites"] += 1
25 return sql
26
27def process_client_buffer(buf: bytearray):
28 out = bytearray(); i = 0
29 while True:
30 if len(buf) - i < 5: break
31 mtype = buf[i]
32 if mtype == 0:
33 if len(buf)-i < 4: break
34 (length,) = struct.unpack_from("!I", buf, i)
35 if length < 4 or length > 10_000_000:
36 out.extend(buf[i:]); return out, bytearray()
37 if len(buf)-i < length: break
38 out.extend(buf[i:i+length]); i += length; continue
39 (length,) = struct.unpack_from("!I", buf, i+1)
40 total = 1 + length
41 if length < 4 or total > 10_000_000:
42 out.extend(buf[i:]); return out, bytearray()
43 if len(buf)-i < total: break
44 msg = bytes(buf[i:i+total])
45 if mtype == ord('Q'):
46 payload = msg[5:]
47 if payload.endswith(b'\x00'):
48 sql = payload[:-1].decode('utf-8','replace')
49 stats['queries'] += 1
50 new_sql = rewrite_sql(sql)
51 if new_sql != sql:
52 new_payload = new_sql.encode() + b'\x00'
53 msg = bytes([ord('Q')]) + struct.pack('!I', 4+len(new_payload)) + new_payload
54 elif mtype == ord('P'):
55 body = msg[5:]
56 try:
57 z1 = body.index(b'\x00'); name = body[:z1+1]; rest = body[z1+1:]
58 z2 = rest.index(b'\x00'); query = rest[:z2].decode('utf-8','replace'); tail = rest[z2:]
59 stats['queries'] += 1
60 new_q = rewrite_sql(query)
61 if new_q != query:
62 new_body = name + new_q.encode() + tail
63 msg = bytes([ord('P')]) + struct.pack('!I', 4+len(new_body)) + new_body
64 except ValueError: pass
65 out.extend(msg); i += total
66 return out, bytearray(buf[i:])
67
68async def pipe_c2s(reader, writer):
69 buf = bytearray()
70 try:
71 while True:
72 chunk = await reader.read(65536)
73 if not chunk: break
74 buf.extend(chunk)
75 to_send, buf = process_client_buffer(buf)
76 if to_send:
77 writer.write(to_send); await writer.drain()
78 except Exception: pass
79 finally:
80 if buf:
81 try: writer.write(buf); await writer.drain()
82 except: pass
83 try: writer.close(); await writer.wait_closed()
84 except: pass
85
86async def pipe_s2c(reader, writer):
87 try:
88 while True:
89 chunk = await reader.read(65536)
90 if not chunk: break
91 writer.write(chunk); await writer.drain()
92 except Exception: pass
93 finally:
94 try: writer.close(); await writer.wait_closed()
95 except: pass
96
97async def handle(cr, cw):
98 try:
99 ur, uw = await asyncio.open_connection(UPSTREAM_HOST, UPSTREAM_PORT)
100 except Exception as e:
101 print(f"[err] {e}", flush=True); cw.close(); return
102 t1=asyncio.create_task(pipe_c2s(cr,uw)); t2=asyncio.create_task(pipe_s2c(ur,cw))
103 await asyncio.wait([t1,t2], return_when=asyncio.FIRST_COMPLETED)
104 t1.cancel(); t2.cancel()
105
106async def main():
107 s = await asyncio.start_server(handle, '0.0.0.0', LISTEN_PORT)
108 print(f"[listen] :{LISTEN_PORT} -> {UPSTREAM_HOST}:{UPSTREAM_PORT}", flush=True)
109 async with s: await s.serve_forever()
110asyncio.run(main())