rewrite_proxy.py 4.5 KB

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