connection.go 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558
  1. package pgserver
  2. import (
  3. "bufio"
  4. "bytes"
  5. "encoding/binary"
  6. "fmt"
  7. "io"
  8. "log"
  9. "net"
  10. "strings"
  11. "github.com/danfragoso/pizzasql-next/pkg/executor"
  12. "github.com/danfragoso/pizzasql-next/pkg/lexer"
  13. "github.com/danfragoso/pizzasql-next/pkg/parser"
  14. "github.com/danfragoso/pizzasql-next/pkg/storage"
  15. )
  16. // Connection represents a client connection
  17. type Connection struct {
  18. conn net.Conn
  19. reader *bufio.Reader
  20. writer *bufio.Writer
  21. executor *executor.Executor
  22. schema *storage.SchemaManager
  23. dbManager *storage.DatabaseManager
  24. database string
  25. params map[string]string
  26. txStatus byte
  27. quiet bool // Disable query logging
  28. }
  29. // NewConnection creates a new connection handler
  30. func NewConnection(conn net.Conn, dbManager *storage.DatabaseManager, quiet bool) *Connection {
  31. return &Connection{
  32. conn: conn,
  33. reader: bufio.NewReader(conn),
  34. writer: bufio.NewWriter(conn),
  35. dbManager: dbManager,
  36. params: make(map[string]string),
  37. txStatus: TxStatusIdle,
  38. quiet: quiet,
  39. }
  40. }
  41. // Handle processes the connection
  42. func (c *Connection) Handle() error {
  43. defer c.conn.Close()
  44. // First, check for SSL request (sent before startup message)
  45. // SSL request is 8 bytes: length(4) + code(4) where code = 80877103
  46. firstBytes := make([]byte, 8)
  47. n, err := io.ReadFull(c.reader, firstBytes)
  48. if err != nil {
  49. return fmt.Errorf("failed to read initial bytes: %w", err)
  50. }
  51. // Check if it's an SSL request (code 80877103 = 0x04D2162F)
  52. if n == 8 {
  53. length := binary.BigEndian.Uint32(firstBytes[0:4])
  54. code := binary.BigEndian.Uint32(firstBytes[4:8])
  55. if length == 8 && code == 80877103 {
  56. // SSL request - we don't support SSL, send 'N'
  57. if !c.quiet {
  58. log.Printf("Client requested SSL, sending rejection")
  59. }
  60. if _, err := c.conn.Write([]byte{'N'}); err != nil {
  61. return fmt.Errorf("failed to send SSL rejection: %w", err)
  62. }
  63. // Now read the actual startup message
  64. } else {
  65. // Not SSL request, this is part of startup message
  66. // We need to prepend these bytes back for ReadStartupMessage
  67. // Create a multi-reader that first reads our buffered bytes, then continues with the reader
  68. c.reader = bufio.NewReader(io.MultiReader(bytes.NewReader(firstBytes), c.reader))
  69. }
  70. }
  71. // Read startup message
  72. if !c.quiet {
  73. log.Printf("Reading startup message...")
  74. }
  75. params, err := ReadStartupMessage(c.reader)
  76. if err != nil {
  77. return fmt.Errorf("failed to read startup message: %w", err)
  78. }
  79. c.params = params
  80. if !c.quiet {
  81. log.Printf("Startup params: %+v", params)
  82. }
  83. // Get database name from params (default to "pizzasql")
  84. dbName := params["database"]
  85. if dbName == "" {
  86. dbName = "pizzasql"
  87. }
  88. c.database = dbName
  89. if !c.quiet {
  90. log.Printf("New connection: user=%s database=%s", params["user"], dbName)
  91. }
  92. // Initialize database
  93. if err := c.initDatabase(dbName); err != nil {
  94. c.sendError("FATAL", ErrCodeConnectionFailure, fmt.Sprintf("Failed to initialize database: %v", err))
  95. return err
  96. }
  97. // Send authentication OK (no auth for now)
  98. if err := c.sendAuthenticationOk(); err != nil {
  99. return err
  100. }
  101. // Send parameter status messages
  102. if err := c.sendParameterStatus("server_version", "14.0 (PizzaSQL)"); err != nil {
  103. return err
  104. }
  105. if err := c.sendParameterStatus("server_encoding", "UTF8"); err != nil {
  106. return err
  107. }
  108. if err := c.sendParameterStatus("client_encoding", "UTF8"); err != nil {
  109. return err
  110. }
  111. if err := c.sendParameterStatus("DateStyle", "ISO, MDY"); err != nil {
  112. return err
  113. }
  114. if err := c.sendParameterStatus("TimeZone", "UTC"); err != nil {
  115. return err
  116. }
  117. // Send backend key data (for cancellation - we don't implement this yet)
  118. if err := c.sendBackendKeyData(12345, 67890); err != nil {
  119. return err
  120. }
  121. // Send ready for query
  122. if err := c.sendReadyForQuery(); err != nil {
  123. return err
  124. }
  125. // Message loop
  126. for {
  127. msg, err := ReadMessage(c.reader)
  128. if err != nil {
  129. if err == io.EOF {
  130. log.Printf("Connection closed by client")
  131. return nil
  132. }
  133. return fmt.Errorf("failed to read message: %w", err)
  134. }
  135. if err := c.handleMessage(msg); err != nil {
  136. if err == io.EOF {
  137. // Normal termination
  138. log.Printf("Connection closed normally")
  139. return nil
  140. }
  141. log.Printf("Error handling message: %v", err)
  142. return err
  143. }
  144. }
  145. }
  146. // initDatabase initializes the database connection
  147. func (c *Connection) initDatabase(dbName string) error {
  148. db, err := c.dbManager.GetDatabase(dbName)
  149. if err != nil {
  150. return err
  151. }
  152. c.schema = db.Schema
  153. c.executor = executor.New(db.Schema, db.Table)
  154. c.executor.SyncCatalog()
  155. return nil
  156. }
  157. // handleMessage processes a client message
  158. func (c *Connection) handleMessage(msg *Message) error {
  159. switch msg.Type {
  160. case MsgQuery:
  161. return c.handleQuery(msg)
  162. case MsgTerminate:
  163. log.Printf("Client requested termination")
  164. return io.EOF
  165. case MsgParse, MsgBind, MsgDescribe, MsgExecute, MsgSync, MsgClose:
  166. // Extended query protocol - not implemented yet
  167. c.sendError("ERROR", ErrCodeFeatureNotSupported, "Extended query protocol not yet supported")
  168. return c.sendReadyForQuery()
  169. default:
  170. log.Printf("Unknown message type: %c (%d)", msg.Type, msg.Type)
  171. c.sendError("ERROR", ErrCodeProtocolViolation, fmt.Sprintf("Unknown message type: %c", msg.Type))
  172. return c.sendReadyForQuery()
  173. }
  174. }
  175. // handleQuery processes a simple query
  176. func (c *Connection) handleQuery(msg *Message) error {
  177. // Parse query string (null-terminated)
  178. sql := string(msg.Data[:len(msg.Data)-1])
  179. if !c.quiet {
  180. log.Printf("Query: %s", sql)
  181. }
  182. // Handle empty query
  183. sqlTrimmed := strings.TrimSpace(sql)
  184. if sqlTrimmed == "" || sqlTrimmed == ";" {
  185. if err := c.sendEmptyQueryResponse(); err != nil {
  186. return err
  187. }
  188. return c.sendReadyForQuery()
  189. }
  190. // Handle special PostgreSQL system queries that drivers send
  191. sqlUpper := strings.ToUpper(strings.TrimSpace(sql))
  192. // lib/pq and other drivers query these for connection validation
  193. if strings.Contains(sqlUpper, "SELECT VERSION()") {
  194. // Return a fake PostgreSQL version
  195. return c.handleVersionQuery()
  196. }
  197. if strings.Contains(sqlUpper, "SELECT CURRENT_USER") {
  198. // Return the current user
  199. return c.handleCurrentUserQuery()
  200. }
  201. if strings.Contains(sqlUpper, "SHOW") && (strings.Contains(sqlUpper, "SERVER_VERSION") ||
  202. strings.Contains(sqlUpper, "SERVER_ENCODING") ||
  203. strings.Contains(sqlUpper, "CLIENT_ENCODING")) {
  204. // Handle SHOW commands
  205. return c.handleShowCommand(sqlUpper)
  206. }
  207. // Execute query
  208. l := lexer.New(sql)
  209. p := parser.New(l)
  210. stmt, err := p.Parse()
  211. if err != nil {
  212. c.sendError("ERROR", ErrCodeSyntaxError, fmt.Sprintf("Syntax error: %v", err))
  213. return c.sendReadyForQuery()
  214. }
  215. result, err := c.executor.Execute(stmt)
  216. if err != nil {
  217. c.sendError("ERROR", ErrCodeInternalError, fmt.Sprintf("Execution error: %v", err))
  218. return c.sendReadyForQuery()
  219. }
  220. // Send result based on statement type
  221. if err := c.sendResult(result, stmt); err != nil {
  222. return err
  223. }
  224. return c.sendReadyForQuery()
  225. }
  226. // sendResult sends query results
  227. func (c *Connection) sendResult(result *executor.Result, stmt parser.Statement) error {
  228. // For SELECT statements, send row description and data rows
  229. if _, isSelect := stmt.(*parser.SelectStmt); isSelect && len(result.Columns) > 0 {
  230. // Send row description
  231. if err := c.sendRowDescription(result.Columns, result.ColumnTypes); err != nil {
  232. return err
  233. }
  234. // Send data rows
  235. for _, row := range result.Rows {
  236. if err := c.sendDataRow(row, result.Columns); err != nil {
  237. return err
  238. }
  239. }
  240. // Send command complete
  241. tag := fmt.Sprintf("SELECT %d", len(result.Rows))
  242. return c.sendCommandComplete(tag)
  243. }
  244. // For other statements, just send command complete
  245. tag := c.getCommandTag(stmt, result)
  246. return c.sendCommandComplete(tag)
  247. }
  248. // getCommandTag returns the command completion tag
  249. func (c *Connection) getCommandTag(stmt parser.Statement, result *executor.Result) string {
  250. switch stmt.(type) {
  251. case *parser.CreateTableStmt:
  252. return "CREATE TABLE"
  253. case *parser.DropTableStmt:
  254. return "DROP TABLE"
  255. case *parser.CreateIndexStmt:
  256. return "CREATE INDEX"
  257. case *parser.DropIndexStmt:
  258. return "DROP INDEX"
  259. case *parser.InsertStmt:
  260. return fmt.Sprintf("INSERT 0 %d", result.RowsAffected)
  261. case *parser.UpdateStmt:
  262. return fmt.Sprintf("UPDATE %d", result.RowsAffected)
  263. case *parser.DeleteStmt:
  264. return fmt.Sprintf("DELETE %d", result.RowsAffected)
  265. case *parser.BeginStmt:
  266. c.txStatus = TxStatusInBlock
  267. return "BEGIN"
  268. case *parser.CommitStmt:
  269. c.txStatus = TxStatusIdle
  270. return "COMMIT"
  271. case *parser.RollbackStmt:
  272. c.txStatus = TxStatusIdle
  273. return "ROLLBACK"
  274. default:
  275. return "OK"
  276. }
  277. }
  278. // sendAuthenticationOk sends authentication OK message
  279. func (c *Connection) sendAuthenticationOk() error {
  280. mb := NewMessageBuilder()
  281. mb.WriteInt32(0) // Auth OK
  282. return c.writeMessage(MsgAuthenticationOk, mb.Bytes())
  283. }
  284. // sendParameterStatus sends a parameter status message
  285. func (c *Connection) sendParameterStatus(name, value string) error {
  286. mb := NewMessageBuilder()
  287. mb.WriteString(name)
  288. mb.WriteString(value)
  289. return c.writeMessage(MsgParameterStatus, mb.Bytes())
  290. }
  291. // sendBackendKeyData sends backend key data
  292. func (c *Connection) sendBackendKeyData(processID, secretKey int32) error {
  293. mb := NewMessageBuilder()
  294. mb.WriteInt32(processID)
  295. mb.WriteInt32(secretKey)
  296. return c.writeMessage(MsgBackendKeyData, mb.Bytes())
  297. }
  298. // sendReadyForQuery sends ready for query message
  299. func (c *Connection) sendReadyForQuery() error {
  300. mb := NewMessageBuilder()
  301. mb.WriteByte(c.txStatus)
  302. return c.writeMessage(MsgReadyForQuery, mb.Bytes())
  303. }
  304. // sendEmptyQueryResponse sends empty query response
  305. func (c *Connection) sendEmptyQueryResponse() error {
  306. return c.writeMessage(MsgEmptyQueryResponse, []byte{})
  307. }
  308. // sendRowDescription sends row description (column metadata)
  309. func (c *Connection) sendRowDescription(columns []string, columnTypes []string) error {
  310. mb := NewMessageBuilder()
  311. mb.WriteInt16(int16(len(columns)))
  312. for i, col := range columns {
  313. colType := ""
  314. if i < len(columnTypes) {
  315. colType = columnTypes[i]
  316. }
  317. mb.WriteString(col)
  318. mb.WriteInt32(0) // table OID
  319. mb.WriteInt16(0) // column attribute number
  320. mb.WriteInt32(c.getOIDForType(colType)) // type OID
  321. mb.WriteInt16(c.getTypeSizeForType(colType)) // type size
  322. mb.WriteInt32(-1) // type modifier
  323. mb.WriteInt16(0) // format code (text)
  324. }
  325. return c.writeMessage(MsgRowDescription, mb.Bytes())
  326. }
  327. // sendDataRow sends a data row
  328. func (c *Connection) sendDataRow(row []interface{}, columns []string) error {
  329. mb := NewMessageBuilder()
  330. mb.WriteInt16(int16(len(row)))
  331. for _, value := range row {
  332. if value == nil {
  333. mb.WriteInt32(-1) // NULL indicator
  334. continue
  335. }
  336. // Convert value to string
  337. strValue := c.valueToString(value)
  338. mb.WriteInt32(int32(len(strValue)))
  339. mb.WriteBytes([]byte(strValue))
  340. }
  341. return c.writeMessage(MsgDataRow, mb.Bytes())
  342. }
  343. // sendCommandComplete sends command complete message
  344. func (c *Connection) sendCommandComplete(tag string) error {
  345. mb := NewMessageBuilder()
  346. mb.WriteString(tag)
  347. return c.writeMessage(MsgCommandComplete, mb.Bytes())
  348. }
  349. // sendError sends an error response
  350. func (c *Connection) sendError(severity, code, message string) error {
  351. mb := NewMessageBuilder()
  352. mb.WriteByte(ErrorFieldSeverity)
  353. mb.WriteString(severity)
  354. mb.WriteByte(ErrorFieldCode)
  355. mb.WriteString(code)
  356. mb.WriteByte(ErrorFieldMessage)
  357. mb.WriteString(message)
  358. mb.WriteByte(0) // Terminator
  359. return c.writeMessage(MsgErrorResponse, mb.Bytes())
  360. }
  361. // writeMessage writes a message to the connection
  362. func (c *Connection) writeMessage(msgType byte, data []byte) error {
  363. if !c.quiet {
  364. log.Printf("Sending message type=%c length=%d", msgType, len(data)+4)
  365. }
  366. if err := WriteMessage(c.writer, msgType, data); err != nil {
  367. return err
  368. }
  369. return c.writer.Flush()
  370. }
  371. // getOIDForType returns PostgreSQL OID for type
  372. func (c *Connection) getOIDForType(typeName string) int32 {
  373. switch strings.ToUpper(typeName) {
  374. case "INTEGER", "INT":
  375. return 23 // INT4OID
  376. case "TEXT", "VARCHAR", "CHAR":
  377. return 25 // TEXTOID
  378. case "REAL", "FLOAT":
  379. return 700 // FLOAT4OID
  380. case "DOUBLE":
  381. return 701 // FLOAT8OID
  382. case "BOOLEAN", "BOOL":
  383. return 16 // BOOLOID
  384. case "BLOB":
  385. return 17 // BYTEAOID
  386. default:
  387. return 25 // Default to TEXT
  388. }
  389. }
  390. // getTypeSizeForType returns type size
  391. func (c *Connection) getTypeSizeForType(typeName string) int16 {
  392. switch strings.ToUpper(typeName) {
  393. case "INTEGER", "INT":
  394. return 4
  395. case "REAL", "FLOAT":
  396. return 4
  397. case "DOUBLE":
  398. return 8
  399. case "BOOLEAN", "BOOL":
  400. return 1
  401. default:
  402. return -1 // Variable length
  403. }
  404. }
  405. // valueToString converts a value to string
  406. func (c *Connection) valueToString(value interface{}) string {
  407. if value == nil {
  408. return ""
  409. }
  410. return fmt.Sprintf("%v", value)
  411. }
  412. // handleVersionQuery handles SELECT version()
  413. func (c *Connection) handleVersionQuery() error {
  414. columns := []string{"version"}
  415. columnTypes := []string{"TEXT"}
  416. if err := c.sendRowDescription(columns, columnTypes); err != nil {
  417. return err
  418. }
  419. row := []interface{}{"PostgreSQL 14.0 (PizzaSQL)"}
  420. if err := c.sendDataRow(row, columns); err != nil {
  421. return err
  422. }
  423. if err := c.sendCommandComplete("SELECT 1"); err != nil {
  424. return err
  425. }
  426. return c.sendReadyForQuery()
  427. }
  428. // handleCurrentUserQuery handles SELECT current_user
  429. func (c *Connection) handleCurrentUserQuery() error {
  430. columns := []string{"current_user"}
  431. columnTypes := []string{"TEXT"}
  432. if err := c.sendRowDescription(columns, columnTypes); err != nil {
  433. return err
  434. }
  435. user := c.params["user"]
  436. if user == "" {
  437. user = "pizzasql"
  438. }
  439. row := []interface{}{user}
  440. if err := c.sendDataRow(row, columns); err != nil {
  441. return err
  442. }
  443. if err := c.sendCommandComplete("SELECT 1"); err != nil {
  444. return err
  445. }
  446. return c.sendReadyForQuery()
  447. }
  448. // handleShowCommand handles SHOW commands
  449. func (c *Connection) handleShowCommand(sqlUpper string) error {
  450. var value string
  451. var name string
  452. if strings.Contains(sqlUpper, "SERVER_VERSION") {
  453. name = "server_version"
  454. value = "14.0"
  455. } else if strings.Contains(sqlUpper, "SERVER_ENCODING") {
  456. name = "server_encoding"
  457. value = "UTF8"
  458. } else if strings.Contains(sqlUpper, "CLIENT_ENCODING") {
  459. name = "client_encoding"
  460. value = "UTF8"
  461. } else {
  462. // Unknown SHOW command
  463. c.sendError("ERROR", ErrCodeFeatureNotSupported, "SHOW command not supported")
  464. return c.sendReadyForQuery()
  465. }
  466. columns := []string{name}
  467. columnTypes := []string{"TEXT"}
  468. if err := c.sendRowDescription(columns, columnTypes); err != nil {
  469. return err
  470. }
  471. row := []interface{}{value}
  472. if err := c.sendDataRow(row, columns); err != nil {
  473. return err
  474. }
  475. if err := c.sendCommandComplete("SHOW"); err != nil {
  476. return err
  477. }
  478. return c.sendReadyForQuery()
  479. }