rewrite_proxy.py110 lines · main
| 1 | #!/usr/bin/env python3 |
| 2 | """Phase-0 SuperTokens→Doltgres SQL rewrite proxy (spike only, not production).""" |
| 3 | import asyncio, re, struct, sys |
| 4 | LISTEN_PORT = int(sys.argv[1]) if len(sys.argv)>1 else 15432 |
| 5 | UPSTREAM_HOST = sys.argv[2] if len(sys.argv)>2 else "doltgres" |
| 6 | UPSTREAM_PORT = int(sys.argv[3]) if len(sys.argv)>3 else 5432 |
| 7 | stats = {"rewrites": 0, "queries": 0} |
| 8 | |
| 9 | def 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 | |
| 27 | def 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 | |
| 68 | async 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 | |
| 86 | async 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 | |
| 97 | async 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 | |
| 106 | async 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() |
| 110 | asyncio.run(main()) |