kv.go 20 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846
  1. package storage
  2. import (
  3. "bufio"
  4. "errors"
  5. "fmt"
  6. "hash/crc32"
  7. "io"
  8. "net"
  9. "strings"
  10. "sync"
  11. "time"
  12. )
  13. const (
  14. headerSize = 32
  15. headerMagic = "PKBF"
  16. headerVersion = 1
  17. opPing = 1
  18. opStatus = 2
  19. opGet = 3
  20. opPut = 4
  21. opDelete = 5
  22. opExists = 6
  23. opMultiGet = 7
  24. opBatchWrite = 8
  25. opScanOpen = 9
  26. opScanNext = 10
  27. opScanClose = 11
  28. batchPut = 1
  29. batchDelete = 2
  30. statusOK = 0
  31. statusNotFound = 1
  32. statusError = 2
  33. maxKeySize = 1024 * 1024
  34. maxValueSize = 64 * 1024 * 1024
  35. maxTransactionSize = 64 * 1024 * 1024
  36. maxOperations = 65535
  37. maxFrameSize = maxKeySize + maxValueSize + 1024
  38. scanPageSize = 1024
  39. existsPipelineSize = 128
  40. )
  41. var crc32cTable = crc32.MakeTable(crc32.Castagnoli)
  42. var (
  43. ErrKeyNotFound = errors.New("key not found")
  44. ErrProtocol = errors.New("pkbfi protocol error")
  45. )
  46. func crc32c(p []byte) uint32 {
  47. return crc32.Checksum(p, crc32cTable)
  48. }
  49. func putU16(b []byte, v uint16) {
  50. b[0] = byte(v)
  51. b[1] = byte(v >> 8)
  52. }
  53. func putU32(b []byte, v uint32) {
  54. b[0] = byte(v)
  55. b[1] = byte(v >> 8)
  56. b[2] = byte(v >> 16)
  57. b[3] = byte(v >> 24)
  58. }
  59. func putU64(b []byte, v uint64) {
  60. b[0] = byte(v)
  61. b[1] = byte(v >> 8)
  62. b[2] = byte(v >> 16)
  63. b[3] = byte(v >> 24)
  64. b[4] = byte(v >> 32)
  65. b[5] = byte(v >> 40)
  66. b[6] = byte(v >> 48)
  67. b[7] = byte(v >> 56)
  68. }
  69. func getU16(b []byte) uint16 {
  70. return uint16(b[0]) | uint16(b[1])<<8
  71. }
  72. func getU32(b []byte) uint32 {
  73. return uint32(b[0]) | uint32(b[1])<<8 | uint32(b[2])<<16 | uint32(b[3])<<24
  74. }
  75. func getU64(b []byte) uint64 {
  76. return uint64(b[0]) | uint64(b[1])<<8 | uint64(b[2])<<16 | uint64(b[3])<<24 |
  77. uint64(b[4])<<32 | uint64(b[5])<<40 | uint64(b[6])<<48 | uint64(b[7])<<56
  78. }
  79. func encodeFrame(opcode, flags uint16, requestID uint64, payload []byte) []byte {
  80. frame := make([]byte, headerSize+len(payload))
  81. copy(frame[0:4], headerMagic)
  82. putU16(frame[4:6], headerVersion)
  83. putU16(frame[6:8], 0)
  84. putU16(frame[8:10], opcode)
  85. putU16(frame[10:12], flags)
  86. putU64(frame[12:20], requestID)
  87. putU32(frame[20:24], uint32(len(payload)))
  88. putU32(frame[24:28], crc32c(payload))
  89. putU32(frame[28:32], 0)
  90. putU32(frame[28:32], crc32c(frame[0:32]))
  91. copy(frame[32:], payload)
  92. return frame
  93. }
  94. func readFrame(r *bufio.Reader) (uint16, uint16, uint64, []byte, error) {
  95. var header [headerSize]byte
  96. if _, err := io.ReadFull(r, header[:]); err != nil {
  97. return 0, 0, 0, nil, err
  98. }
  99. if string(header[0:4]) != headerMagic {
  100. return 0, 0, 0, nil, fmt.Errorf("%w: invalid magic", ErrProtocol)
  101. }
  102. if getU16(header[4:6]) != headerVersion {
  103. return 0, 0, 0, nil, fmt.Errorf("%w: incompatible version", ErrProtocol)
  104. }
  105. payloadLen := getU32(header[20:24])
  106. if payloadLen > maxFrameSize {
  107. return 0, 0, 0, nil, fmt.Errorf("%w: frame too large", ErrProtocol)
  108. }
  109. headerCRC := getU32(header[28:32])
  110. var headerCopy [headerSize]byte
  111. copy(headerCopy[:], header[:])
  112. putU32(headerCopy[28:32], 0)
  113. if crc32c(headerCopy[:]) != headerCRC {
  114. return 0, 0, 0, nil, fmt.Errorf("%w: header checksum mismatch", ErrProtocol)
  115. }
  116. payload := make([]byte, payloadLen)
  117. if _, err := io.ReadFull(r, payload); err != nil {
  118. return 0, 0, 0, nil, err
  119. }
  120. if crc32c(payload) != getU32(header[24:28]) {
  121. return 0, 0, 0, nil, fmt.Errorf("%w: payload checksum mismatch", ErrProtocol)
  122. }
  123. return getU16(header[8:10]), getU16(header[10:12]), getU64(header[12:20]), payload, nil
  124. }
  125. func encodeResponse(opcode uint16, requestID uint64, body []byte) []byte {
  126. return encodeFrame(opcode|0x8000, 1, requestID, body)
  127. }
  128. func oneKeyPayload(key []byte) []byte {
  129. payload := make([]byte, 4+len(key))
  130. putU32(payload[0:4], uint32(len(key)))
  131. copy(payload[4:], key)
  132. return payload
  133. }
  134. func validateKey(key []byte) error {
  135. if len(key) > maxKeySize {
  136. return fmt.Errorf("pkbfi: key exceeds %d bytes", maxKeySize)
  137. }
  138. return nil
  139. }
  140. func validateValue(value []byte) error {
  141. if len(value) > maxValueSize {
  142. return fmt.Errorf("pkbfi: value exceeds %d bytes", maxValueSize)
  143. }
  144. return nil
  145. }
  146. func parseOneKey(payload []byte) ([]byte, bool) {
  147. if len(payload) < 4 {
  148. return nil, false
  149. }
  150. length := getU32(payload[0:4])
  151. if length > maxKeySize || uint64(4)+uint64(length) != uint64(len(payload)) {
  152. return nil, false
  153. }
  154. return payload[4:], true
  155. }
  156. func errorBody(message string) []byte {
  157. body := make([]byte, 2+len(message))
  158. putU16(body[0:2], statusError)
  159. copy(body[2:], message)
  160. return body
  161. }
  162. type KVClient struct {
  163. conn net.Conn
  164. reader *bufio.Reader
  165. writer *bufio.Writer
  166. mu sync.Mutex
  167. nextID uint64
  168. requestTimeout time.Duration
  169. lastUsed time.Time
  170. }
  171. type KVResult struct {
  172. Value []byte
  173. LSN uint64
  174. Found bool
  175. }
  176. type KVEntry struct {
  177. Key []byte
  178. Value []byte
  179. LSN uint64
  180. }
  181. type BatchOp struct {
  182. Op byte
  183. Key []byte
  184. Value []byte
  185. }
  186. type ScanCursor struct {
  187. client *KVClient
  188. id uint64
  189. }
  190. func NewKVClient(addr string) (*KVClient, error) {
  191. network, target := parseAddr(addr)
  192. conn, err := net.Dial(network, target)
  193. if err != nil {
  194. return nil, fmt.Errorf("failed to connect to PizzaKV: %w", err)
  195. }
  196. return &KVClient{
  197. conn: conn,
  198. reader: bufio.NewReader(conn),
  199. writer: bufio.NewWriter(conn),
  200. nextID: 1,
  201. lastUsed: time.Now(),
  202. }, nil
  203. }
  204. func parseAddr(addr string) (string, string) {
  205. if strings.HasPrefix(addr, "unix:") {
  206. return "unix", strings.TrimPrefix(addr, "unix:")
  207. }
  208. return "tcp", addr
  209. }
  210. func (c *KVClient) Close() error {
  211. c.mu.Lock()
  212. defer c.mu.Unlock()
  213. if c.conn != nil {
  214. err := c.conn.Close()
  215. c.conn = nil
  216. return err
  217. }
  218. return nil
  219. }
  220. func (c *KVClient) SetDeadline(t time.Time) error {
  221. if c.conn == nil {
  222. return nil
  223. }
  224. return c.conn.SetDeadline(t)
  225. }
  226. func (c *KVClient) writeFrame(opcode, flags uint16, requestID uint64, payload []byte) error {
  227. frame := encodeFrame(opcode, flags, requestID, payload)
  228. if _, err := c.writer.Write(frame); err != nil {
  229. return err
  230. }
  231. return c.writer.Flush()
  232. }
  233. func (c *KVClient) request(opcode uint16, payload []byte) (uint16, []byte, error) {
  234. if c.requestTimeout > 0 {
  235. if err := c.conn.SetDeadline(time.Now().Add(c.requestTimeout)); err != nil {
  236. return 0, nil, err
  237. }
  238. }
  239. requestID := c.nextID
  240. c.nextID++
  241. if err := c.writeFrame(opcode, 0, requestID, payload); err != nil {
  242. return 0, nil, err
  243. }
  244. respOpcode, respFlags, respID, body, err := readFrame(c.reader)
  245. if err != nil {
  246. return 0, nil, err
  247. }
  248. if respOpcode != opcode|0x8000 {
  249. return 0, nil, fmt.Errorf("%w: unexpected response opcode %d", ErrProtocol, respOpcode)
  250. }
  251. if respFlags != 1 {
  252. return 0, nil, fmt.Errorf("%w: unexpected response flags %d", ErrProtocol, respFlags)
  253. }
  254. if respID != requestID {
  255. return 0, nil, fmt.Errorf("%w: response id %d does not match request %d", ErrProtocol, respID, requestID)
  256. }
  257. c.lastUsed = time.Now()
  258. if len(body) < 2 {
  259. return 0, nil, fmt.Errorf("%w: response too short", ErrProtocol)
  260. }
  261. status := getU16(body[0:2])
  262. if status == statusError {
  263. return status, nil, fmt.Errorf("pkbfi server error: %s", body[2:])
  264. }
  265. return status, body[2:], nil
  266. }
  267. func (c *KVClient) Put(key, value []byte) (uint64, error) {
  268. if err := validateKey(key); err != nil {
  269. return 0, err
  270. }
  271. if err := validateValue(value); err != nil {
  272. return 0, err
  273. }
  274. c.mu.Lock()
  275. defer c.mu.Unlock()
  276. payload := make([]byte, 8+len(key)+len(value))
  277. putU32(payload[0:4], uint32(len(key)))
  278. putU32(payload[4:8], uint32(len(value)))
  279. copy(payload[8:], key)
  280. copy(payload[8+len(key):], value)
  281. status, body, err := c.request(opPut, payload)
  282. if err != nil {
  283. return 0, err
  284. }
  285. if status != statusOK || len(body) != 8 {
  286. return 0, fmt.Errorf("%w: malformed put response", ErrProtocol)
  287. }
  288. return getU64(body[0:8]), nil
  289. }
  290. func (c *KVClient) Get(key []byte) (KVResult, error) {
  291. if err := validateKey(key); err != nil {
  292. return KVResult{}, err
  293. }
  294. c.mu.Lock()
  295. defer c.mu.Unlock()
  296. status, body, err := c.request(opGet, oneKeyPayload(key))
  297. if err != nil {
  298. return KVResult{}, err
  299. }
  300. if status == statusNotFound {
  301. return KVResult{}, ErrKeyNotFound
  302. }
  303. if status != statusOK || len(body) < 12 {
  304. return KVResult{}, fmt.Errorf("%w: malformed get response", ErrProtocol)
  305. }
  306. lsn := getU64(body[0:8])
  307. valueLen := getU32(body[8:12])
  308. if valueLen > maxValueSize || uint64(len(body)) != 12+uint64(valueLen) {
  309. return KVResult{}, fmt.Errorf("%w: malformed get value length", ErrProtocol)
  310. }
  311. return KVResult{Value: body[12:], LSN: lsn, Found: true}, nil
  312. }
  313. func (c *KVClient) Del(key []byte) (bool, error) {
  314. if err := validateKey(key); err != nil {
  315. return false, err
  316. }
  317. c.mu.Lock()
  318. defer c.mu.Unlock()
  319. status, body, err := c.request(opDelete, oneKeyPayload(key))
  320. if err != nil {
  321. return false, err
  322. }
  323. if status != statusOK || len(body) != 1 {
  324. return false, fmt.Errorf("%w: malformed delete response", ErrProtocol)
  325. }
  326. return body[0] != 0, nil
  327. }
  328. func (c *KVClient) Exists(key []byte) (bool, error) {
  329. if err := validateKey(key); err != nil {
  330. return false, err
  331. }
  332. c.mu.Lock()
  333. defer c.mu.Unlock()
  334. status, body, err := c.request(opExists, oneKeyPayload(key))
  335. if err != nil {
  336. return false, err
  337. }
  338. if status != statusOK || len(body) != 1 {
  339. return false, fmt.Errorf("%w: malformed exists response", ErrProtocol)
  340. }
  341. return body[0] != 0, nil
  342. }
  343. func (c *KVClient) ExistsMany(keys [][]byte) ([]bool, error) {
  344. if len(keys) > maxOperations {
  345. return nil, fmt.Errorf("pkbfi: too many keys")
  346. }
  347. for _, key := range keys {
  348. if err := validateKey(key); err != nil {
  349. return nil, err
  350. }
  351. }
  352. c.mu.Lock()
  353. defer c.mu.Unlock()
  354. results := make([]bool, len(keys))
  355. for start := 0; start < len(keys); start += existsPipelineSize {
  356. end := start + existsPipelineSize
  357. if end > len(keys) {
  358. end = len(keys)
  359. }
  360. ids := make([]uint64, end-start)
  361. if c.requestTimeout > 0 {
  362. if err := c.conn.SetDeadline(time.Now().Add(c.requestTimeout)); err != nil {
  363. return nil, err
  364. }
  365. }
  366. for i, key := range keys[start:end] {
  367. ids[i] = c.nextID
  368. c.nextID++
  369. if _, err := c.writer.Write(encodeFrame(opExists, 0, ids[i], oneKeyPayload(key))); err != nil {
  370. return nil, err
  371. }
  372. }
  373. if err := c.writer.Flush(); err != nil {
  374. return nil, err
  375. }
  376. for i, requestID := range ids {
  377. opcode, flags, responseID, body, err := readFrame(c.reader)
  378. if err != nil {
  379. return nil, err
  380. }
  381. if opcode != opExists|0x8000 || flags != 1 || responseID != requestID {
  382. return nil, fmt.Errorf("%w: malformed exists response frame", ErrProtocol)
  383. }
  384. if len(body) < 2 {
  385. return nil, fmt.Errorf("%w: response too short", ErrProtocol)
  386. }
  387. status := getU16(body[0:2])
  388. if status == statusError {
  389. return nil, fmt.Errorf("pkbfi server error: %s", body[2:])
  390. }
  391. if status != statusOK || len(body) != 3 {
  392. return nil, fmt.Errorf("%w: malformed exists response", ErrProtocol)
  393. }
  394. results[start+i] = body[2] != 0
  395. c.lastUsed = time.Now()
  396. }
  397. }
  398. return results, nil
  399. }
  400. func (c *KVClient) MultiGet(keys [][]byte) ([]KVResult, error) {
  401. if len(keys) > maxOperations {
  402. return nil, fmt.Errorf("pkbfi: too many keys")
  403. }
  404. payloadSize := 4
  405. for _, key := range keys {
  406. if err := validateKey(key); err != nil {
  407. return nil, err
  408. }
  409. payloadSize += 4 + len(key)
  410. if payloadSize > maxFrameSize {
  411. return nil, fmt.Errorf("pkbfi: multi_get request exceeds frame limit")
  412. }
  413. }
  414. c.mu.Lock()
  415. defer c.mu.Unlock()
  416. payload := make([]byte, 4, payloadSize)
  417. putU32(payload[0:4], uint32(len(keys)))
  418. for _, key := range keys {
  419. var length [4]byte
  420. putU32(length[:], uint32(len(key)))
  421. payload = append(payload, length[:]...)
  422. payload = append(payload, key...)
  423. }
  424. status, body, err := c.request(opMultiGet, payload)
  425. if err != nil {
  426. return nil, err
  427. }
  428. if status != statusOK || len(body) < 4 {
  429. return nil, fmt.Errorf("%w: malformed multi_get response", ErrProtocol)
  430. }
  431. count := getU32(body[0:4])
  432. if count != uint32(len(keys)) {
  433. return nil, fmt.Errorf("%w: multi_get count mismatch", ErrProtocol)
  434. }
  435. results := make([]KVResult, count)
  436. pos := 4
  437. for i := uint32(0); i < count; i++ {
  438. if len(body)-pos < 16 {
  439. return nil, fmt.Errorf("%w: truncated multi_get entry", ErrProtocol)
  440. }
  441. present := body[pos] != 0
  442. valueLen := getU32(body[pos+4 : pos+8])
  443. lsn := getU64(body[pos+8 : pos+16])
  444. pos += 16
  445. results[i] = KVResult{LSN: lsn, Found: present}
  446. if present {
  447. if valueLen > maxValueSize || len(body)-pos < int(valueLen) {
  448. return nil, fmt.Errorf("%w: multi_get value length", ErrProtocol)
  449. }
  450. results[i].Value = body[pos : pos+int(valueLen)]
  451. pos += int(valueLen)
  452. }
  453. }
  454. if pos != len(body) {
  455. return nil, fmt.Errorf("%w: multi_get trailing bytes", ErrProtocol)
  456. }
  457. return results, nil
  458. }
  459. func (c *KVClient) BatchWrite(ops []BatchOp, metadata []byte) (uint64, error) {
  460. if len(ops) == 0 || len(ops) > maxOperations {
  461. return 0, fmt.Errorf("pkbfi: invalid operation count")
  462. }
  463. if len(metadata) > maxTransactionSize-8 {
  464. return 0, fmt.Errorf("pkbfi: batch metadata exceeds transaction limit")
  465. }
  466. payloadSize := 8 + len(metadata)
  467. for _, op := range ops {
  468. if op.Op != batchPut && op.Op != batchDelete {
  469. return 0, fmt.Errorf("pkbfi: invalid batch opcode %d", op.Op)
  470. }
  471. if err := validateKey(op.Key); err != nil {
  472. return 0, err
  473. }
  474. if err := validateValue(op.Value); err != nil {
  475. return 0, err
  476. }
  477. if op.Op == batchDelete && len(op.Value) != 0 {
  478. return 0, fmt.Errorf("pkbfi: delete operation with value")
  479. }
  480. payloadSize += 12 + len(op.Key) + len(op.Value)
  481. if payloadSize > maxTransactionSize {
  482. return 0, fmt.Errorf("pkbfi: batch exceeds transaction limit")
  483. }
  484. }
  485. c.mu.Lock()
  486. defer c.mu.Unlock()
  487. payload := make([]byte, 8, payloadSize)
  488. putU32(payload[0:4], uint32(len(ops)))
  489. putU32(payload[4:8], uint32(len(metadata)))
  490. payload = append(payload, metadata...)
  491. for _, op := range ops {
  492. var header [12]byte
  493. header[0] = op.Op
  494. putU32(header[4:8], uint32(len(op.Key)))
  495. putU32(header[8:12], uint32(len(op.Value)))
  496. payload = append(payload, header[:]...)
  497. payload = append(payload, op.Key...)
  498. payload = append(payload, op.Value...)
  499. }
  500. status, body, err := c.request(opBatchWrite, payload)
  501. if err != nil {
  502. return 0, err
  503. }
  504. if status != statusOK || len(body) != 8 {
  505. return 0, fmt.Errorf("%w: malformed batch_write response", ErrProtocol)
  506. }
  507. return getU64(body[0:8]), nil
  508. }
  509. func (c *KVClient) Scan(prefix []byte) (*ScanCursor, error) {
  510. return c.openScan(prefix, true, scanPageSize)
  511. }
  512. func (c *KVClient) ScanWithLimit(prefix []byte, pageSize uint32) (*ScanCursor, error) {
  513. return c.openScan(prefix, true, pageSize)
  514. }
  515. // ScanKeys opens a key-only scan: the server omits values from the returned
  516. // pages. It is used when only the key set (or its size) is needed, such as the
  517. // COUNT(*) fast path or bulk deletion, so a full-table scan does not pull row
  518. // values across the wire.
  519. func (c *KVClient) ScanKeys(prefix []byte) (*ScanCursor, error) {
  520. return c.openScan(prefix, false, scanPageSize)
  521. }
  522. func (c *KVClient) openScan(prefix []byte, includeValues bool, pageSize uint32) (*ScanCursor, error) {
  523. if err := validateKey(prefix); err != nil {
  524. return nil, err
  525. }
  526. if pageSize == 0 || pageSize > 4096 {
  527. return nil, fmt.Errorf("pkbfi: scan page size must be between 1 and 4096")
  528. }
  529. c.mu.Lock()
  530. defer c.mu.Unlock()
  531. payload := make([]byte, 12+len(prefix))
  532. if includeValues {
  533. payload[0] = 1
  534. }
  535. putU32(payload[4:8], pageSize)
  536. putU32(payload[8:12], uint32(len(prefix)))
  537. copy(payload[12:], prefix)
  538. status, body, err := c.request(opScanOpen, payload)
  539. if err != nil {
  540. return nil, err
  541. }
  542. if status != statusOK || len(body) != 8 {
  543. return nil, fmt.Errorf("%w: malformed scan_open response", ErrProtocol)
  544. }
  545. return &ScanCursor{client: c, id: getU64(body[0:8])}, nil
  546. }
  547. func (s *ScanCursor) Next() ([]KVEntry, bool, error) {
  548. s.client.mu.Lock()
  549. defer s.client.mu.Unlock()
  550. var payload [12]byte
  551. putU64(payload[0:8], s.id)
  552. putU32(payload[8:12], 0)
  553. status, body, err := s.client.request(opScanNext, payload[:])
  554. if err != nil {
  555. return nil, false, err
  556. }
  557. if status != statusOK || len(body) < 8 {
  558. return nil, false, fmt.Errorf("%w: malformed scan_next response", ErrProtocol)
  559. }
  560. done := body[0] != 0
  561. count := getU32(body[4:8])
  562. entries := make([]KVEntry, 0, count)
  563. pos := 8
  564. for i := uint32(0); i < count; i++ {
  565. if len(body)-pos < 16 {
  566. return nil, false, fmt.Errorf("%w: truncated scan entry", ErrProtocol)
  567. }
  568. keyLen := getU32(body[pos : pos+4])
  569. valueLen := getU32(body[pos+4 : pos+8])
  570. lsn := getU64(body[pos+8 : pos+16])
  571. pos += 16
  572. if keyLen > maxKeySize || valueLen > maxValueSize || len(body)-pos < int(keyLen)+int(valueLen) {
  573. return nil, false, fmt.Errorf("%w: scan entry length", ErrProtocol)
  574. }
  575. key := body[pos : pos+int(keyLen)]
  576. pos += int(keyLen)
  577. value := body[pos : pos+int(valueLen)]
  578. pos += int(valueLen)
  579. entries = append(entries, KVEntry{Key: key, Value: value, LSN: lsn})
  580. }
  581. if pos != len(body) {
  582. return nil, false, fmt.Errorf("%w: scan trailing bytes", ErrProtocol)
  583. }
  584. return entries, done, nil
  585. }
  586. func (s *ScanCursor) Close() error {
  587. s.client.mu.Lock()
  588. defer s.client.mu.Unlock()
  589. var payload [8]byte
  590. putU64(payload[0:8], s.id)
  591. status, body, err := s.client.request(opScanClose, payload[:])
  592. if err != nil {
  593. return err
  594. }
  595. if status != statusOK || len(body) != 1 {
  596. return fmt.Errorf("%w: malformed scan_close response", ErrProtocol)
  597. }
  598. return nil
  599. }
  600. func (c *KVClient) Write(key, value string) error {
  601. _, err := c.Put([]byte(key), []byte(value))
  602. return err
  603. }
  604. func (c *KVClient) Read(key string) (string, error) {
  605. res, err := c.Get([]byte(key))
  606. if err != nil {
  607. return "", err
  608. }
  609. return string(res.Value), nil
  610. }
  611. func (c *KVClient) Delete(key string) error {
  612. _, err := c.Del([]byte(key))
  613. return err
  614. }
  615. func (c *KVClient) Reads(prefix string) ([]string, error) {
  616. scan, err := c.Scan([]byte(prefix))
  617. if err != nil {
  618. return nil, err
  619. }
  620. defer scan.Close()
  621. values := make([]string, 0)
  622. for {
  623. entries, done, err := scan.Next()
  624. if err != nil {
  625. return nil, err
  626. }
  627. for _, entry := range entries {
  628. values = append(values, string(entry.Value))
  629. }
  630. if done {
  631. return values, nil
  632. }
  633. }
  634. }
  635. func (c *KVClient) IsAlive() bool {
  636. c.mu.Lock()
  637. defer c.mu.Unlock()
  638. if c.conn == nil {
  639. return false
  640. }
  641. c.conn.SetDeadline(time.Now().Add(500 * time.Millisecond))
  642. defer c.conn.SetDeadline(time.Time{})
  643. _, _, err := c.request(opPing, nil)
  644. return err == nil
  645. }
  646. type KVPool struct {
  647. addr string
  648. pool chan *KVClient
  649. size int
  650. timeout time.Duration
  651. mu sync.Mutex
  652. closed bool
  653. }
  654. func (p *KVPool) replacementClient() (*KVClient, error) {
  655. client, err := NewKVClient(p.addr)
  656. if err == nil {
  657. client.requestTimeout = p.timeout
  658. return client, nil
  659. }
  660. p.mu.Lock()
  661. if !p.closed {
  662. select {
  663. case p.pool <- nil:
  664. default:
  665. }
  666. }
  667. p.mu.Unlock()
  668. return nil, err
  669. }
  670. func NewKVPool(addr string, size int, timeout time.Duration) (*KVPool, error) {
  671. p := &KVPool{
  672. addr: addr,
  673. pool: make(chan *KVClient, size),
  674. size: size,
  675. timeout: timeout,
  676. }
  677. for i := 0; i < size; i++ {
  678. client, err := NewKVClient(addr)
  679. if err != nil {
  680. p.Close()
  681. return nil, fmt.Errorf("failed to create connection pool: %w", err)
  682. }
  683. p.pool <- client
  684. }
  685. return p, nil
  686. }
  687. func (p *KVPool) Get() (*KVClient, error) {
  688. p.mu.Lock()
  689. if p.closed {
  690. p.mu.Unlock()
  691. return nil, fmt.Errorf("pool is closed")
  692. }
  693. p.mu.Unlock()
  694. select {
  695. case client := <-p.pool:
  696. if client != nil && client.conn != nil {
  697. if client.lastUsed.IsZero() {
  698. client.lastUsed = time.Now()
  699. } else if time.Since(client.lastUsed) >= 20*time.Second && !client.IsAlive() {
  700. client.Close()
  701. return p.replacementClient()
  702. }
  703. client.requestTimeout = p.timeout
  704. return client, nil
  705. }
  706. return p.replacementClient()
  707. case <-time.After(30 * time.Second):
  708. return nil, fmt.Errorf("kv pool timeout: no connection available after 30s")
  709. }
  710. }
  711. func (p *KVPool) Put(client *KVClient) {
  712. if client == nil {
  713. return
  714. }
  715. p.mu.Lock()
  716. if p.closed {
  717. p.mu.Unlock()
  718. client.Close()
  719. return
  720. }
  721. p.mu.Unlock()
  722. client.requestTimeout = 0
  723. client.SetDeadline(time.Time{})
  724. select {
  725. case p.pool <- client:
  726. default:
  727. client.Close()
  728. }
  729. }
  730. func (p *KVPool) Close() error {
  731. p.mu.Lock()
  732. if p.closed {
  733. p.mu.Unlock()
  734. return nil
  735. }
  736. p.closed = true
  737. p.mu.Unlock()
  738. close(p.pool)
  739. for client := range p.pool {
  740. if client != nil {
  741. client.Close()
  742. }
  743. }
  744. return nil
  745. }
  746. func (p *KVPool) WithClient(fn func(*KVClient) error) error {
  747. client, err := p.Get()
  748. if err != nil {
  749. return err
  750. }
  751. if err := fn(client); err != nil {
  752. if !isConnectionError(err) {
  753. p.Put(client)
  754. return err
  755. }
  756. client.Close()
  757. p.mu.Lock()
  758. if !p.closed {
  759. select {
  760. case p.pool <- nil:
  761. default:
  762. }
  763. }
  764. p.mu.Unlock()
  765. return err
  766. }
  767. p.Put(client)
  768. return nil
  769. }
  770. func isConnectionError(err error) bool {
  771. var netErr net.Error
  772. if errors.As(err, &netErr) {
  773. return true
  774. }
  775. if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
  776. return true
  777. }
  778. return errors.Is(err, ErrProtocol)
  779. }