2
0

kv_test.go 23 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929
  1. package storage
  2. import (
  3. "bufio"
  4. "bytes"
  5. "errors"
  6. "fmt"
  7. "net"
  8. "strings"
  9. "testing"
  10. "time"
  11. )
  12. func getResponseBody(value []byte, lsn uint64) []byte {
  13. body := make([]byte, 14+len(value))
  14. putU16(body[0:2], statusOK)
  15. putU64(body[2:10], lsn)
  16. putU32(body[10:14], uint32(len(value)))
  17. copy(body[14:], value)
  18. return body
  19. }
  20. func pipeClient(conn net.Conn) *KVClient {
  21. return &KVClient{
  22. conn: conn,
  23. reader: bufio.NewReader(conn),
  24. writer: bufio.NewWriter(conn),
  25. nextID: 1,
  26. lastUsed: time.Now(),
  27. }
  28. }
  29. func TestCRC32C(t *testing.T) {
  30. if got := crc32c([]byte("123456789")); got != 0xe3069283 {
  31. t.Fatalf("crc32c = %#x, want 0xe3069283", got)
  32. }
  33. }
  34. func TestEncodeReadFrameRoundTrip(t *testing.T) {
  35. payload := []byte{0x00, 0x01, 0xfe, '\n', '\r'}
  36. frame := encodeFrame(opPut, 0, 42, payload)
  37. r := bufio.NewReader(bytes.NewReader(frame))
  38. opcode, flags, requestID, body, err := readFrame(r)
  39. if err != nil {
  40. t.Fatalf("readFrame: %v", err)
  41. }
  42. if opcode != opPut || flags != 0 || requestID != 42 {
  43. t.Fatalf("opcode=%d flags=%d requestID=%d", opcode, flags, requestID)
  44. }
  45. if !bytes.Equal(body, payload) {
  46. t.Fatalf("body = %x, want %x", body, payload)
  47. }
  48. }
  49. func TestReadFrameRejectsOversizedFrame(t *testing.T) {
  50. var header [headerSize]byte
  51. copy(header[0:4], headerMagic)
  52. putU16(header[4:6], headerVersion)
  53. putU32(header[20:24], maxFrameSize+1)
  54. r := bufio.NewReader(strings.NewReader(string(header[:])))
  55. if _, _, _, _, err := readFrame(r); !errors.Is(err, ErrProtocol) {
  56. t.Fatalf("err = %v, want ErrProtocol", err)
  57. }
  58. }
  59. func TestClientRejectsOversizedRequests(t *testing.T) {
  60. c := &KVClient{}
  61. largeKey := make([]byte, maxKeySize+1)
  62. largeValue := make([]byte, maxValueSize+1)
  63. if _, err := c.Put(largeKey, nil); err == nil {
  64. t.Fatal("Put accepted an oversized key")
  65. }
  66. if _, err := c.Put(nil, largeValue); err == nil {
  67. t.Fatal("Put accepted an oversized value")
  68. }
  69. if _, err := c.Get(largeKey); err == nil {
  70. t.Fatal("Get accepted an oversized key")
  71. }
  72. if _, err := c.MultiGet([][]byte{largeKey}); err == nil {
  73. t.Fatal("MultiGet accepted an oversized key")
  74. }
  75. if _, err := c.BatchWrite([]BatchOp{{Op: batchPut, Key: []byte("k"), Value: largeValue}}, nil); err == nil {
  76. t.Fatal("BatchWrite accepted an oversized value")
  77. }
  78. if _, err := c.Scan(largeKey); err == nil {
  79. t.Fatal("Scan accepted an oversized prefix")
  80. }
  81. if _, err := c.ScanWithLimit(nil, 0); err == nil {
  82. t.Fatal("ScanWithLimit accepted a zero page size")
  83. }
  84. if _, err := c.ScanWithLimit(nil, 4097); err == nil {
  85. t.Fatal("ScanWithLimit accepted an oversized page size")
  86. }
  87. }
  88. func TestPutGetBinaryRoundTrip(t *testing.T) {
  89. kv := newTestKVServer(t)
  90. defer kv.close()
  91. c := kv.client()
  92. defer c.Close()
  93. key := []byte{0x00, 0x01, 0x02, 'k', '\n'}
  94. value := []byte{0xff, 0x00, '\r', '\n', 'v', 0x80}
  95. lsn, err := c.Put(key, value)
  96. if err != nil {
  97. t.Fatalf("put: %v", err)
  98. }
  99. if lsn == 0 {
  100. t.Fatalf("put returned zero lsn")
  101. }
  102. res, err := c.Get(key)
  103. if err != nil {
  104. t.Fatalf("get: %v", err)
  105. }
  106. if !res.Found {
  107. t.Fatalf("expected found")
  108. }
  109. if res.LSN != lsn {
  110. t.Fatalf("lsn = %d, want %d", res.LSN, lsn)
  111. }
  112. if !bytes.Equal(res.Value, value) {
  113. t.Fatalf("value = %x, want %x", res.Value, value)
  114. }
  115. }
  116. func TestGetNotFound(t *testing.T) {
  117. kv := newTestKVServer(t)
  118. defer kv.close()
  119. c := kv.client()
  120. defer c.Close()
  121. if _, err := c.Get([]byte("missing")); err != ErrKeyNotFound {
  122. t.Fatalf("err = %v, want ErrKeyNotFound", err)
  123. }
  124. if _, err := c.Read("missing"); err != ErrKeyNotFound {
  125. t.Fatalf("read err = %v, want ErrKeyNotFound", err)
  126. }
  127. }
  128. func TestDeleteAndExists(t *testing.T) {
  129. kv := newTestKVServer(t)
  130. defer kv.close()
  131. c := kv.client()
  132. defer c.Close()
  133. if _, err := c.Put([]byte("k"), []byte("v")); err != nil {
  134. t.Fatalf("put: %v", err)
  135. }
  136. found, err := c.Exists([]byte("k"))
  137. if err != nil || !found {
  138. t.Fatalf("exists = %v, %v", found, err)
  139. }
  140. deleted, err := c.Del([]byte("k"))
  141. if err != nil || !deleted {
  142. t.Fatalf("del = %v, %v", deleted, err)
  143. }
  144. found, err = c.Exists([]byte("k"))
  145. if err != nil || found {
  146. t.Fatalf("exists after delete = %v, %v", found, err)
  147. }
  148. deleted, err = c.Del([]byte("k"))
  149. if err != nil || deleted {
  150. t.Fatalf("del missing = %v, %v", deleted, err)
  151. }
  152. }
  153. func TestExistsManyPipelinesRequests(t *testing.T) {
  154. kv := newTestKVServer(t)
  155. defer kv.close()
  156. ln, err := net.Listen("tcp", "127.0.0.1:0")
  157. if err != nil {
  158. t.Fatal(err)
  159. }
  160. defer ln.Close()
  161. go func() {
  162. conn, err := ln.Accept()
  163. if err == nil {
  164. kv.handle(conn)
  165. }
  166. }()
  167. c, err := NewKVClient(ln.Addr().String())
  168. if err != nil {
  169. t.Fatal(err)
  170. }
  171. defer c.Close()
  172. keys := make([][]byte, 300)
  173. for i := range keys {
  174. keys[i] = []byte(fmt.Sprintf("key:%03d", i))
  175. if i%2 == 0 {
  176. if _, err := c.Put(keys[i], []byte("v")); err != nil {
  177. t.Fatal(err)
  178. }
  179. }
  180. }
  181. exists, err := c.ExistsMany(keys)
  182. if err != nil {
  183. t.Fatal(err)
  184. }
  185. for i, found := range exists {
  186. if found != (i%2 == 0) {
  187. t.Fatalf("exists[%d]=%v", i, found)
  188. }
  189. }
  190. }
  191. func TestMultiGet(t *testing.T) {
  192. kv := newTestKVServer(t)
  193. defer kv.close()
  194. c := kv.client()
  195. defer c.Close()
  196. if _, err := c.Put([]byte("a"), []byte("1")); err != nil {
  197. t.Fatalf("put a: %v", err)
  198. }
  199. if _, err := c.Put([]byte("b"), []byte{0x00, 0x02}); err != nil {
  200. t.Fatalf("put b: %v", err)
  201. }
  202. results, err := c.MultiGet([][]byte{[]byte("a"), []byte("missing"), []byte("b")})
  203. if err != nil {
  204. t.Fatalf("multi_get: %v", err)
  205. }
  206. if len(results) != 3 {
  207. t.Fatalf("len = %d", len(results))
  208. }
  209. if !results[0].Found || string(results[0].Value) != "1" {
  210. t.Fatalf("results[0] = %+v", results[0])
  211. }
  212. if results[1].Found {
  213. t.Fatalf("results[1] should be missing")
  214. }
  215. if !results[2].Found || !bytes.Equal(results[2].Value, []byte{0x00, 0x02}) {
  216. t.Fatalf("results[2] = %+v", results[2])
  217. }
  218. }
  219. func TestBatchWrite(t *testing.T) {
  220. kv := newTestKVServer(t)
  221. defer kv.close()
  222. c := kv.client()
  223. defer c.Close()
  224. if _, err := c.Put([]byte("d"), []byte("old")); err != nil {
  225. t.Fatalf("put d: %v", err)
  226. }
  227. lsn, err := c.BatchWrite([]BatchOp{
  228. {Op: batchPut, Key: []byte("a"), Value: []byte("1")},
  229. {Op: batchPut, Key: []byte("b"), Value: []byte("2")},
  230. {Op: batchDelete, Key: []byte("d")},
  231. }, []byte("meta"))
  232. if err != nil {
  233. t.Fatalf("batch_write: %v", err)
  234. }
  235. if lsn == 0 {
  236. t.Fatalf("zero lsn")
  237. }
  238. ra, err := c.Get([]byte("a"))
  239. if err != nil || !ra.Found || string(ra.Value) != "1" {
  240. t.Fatalf("a = %+v, %v", ra, err)
  241. }
  242. rb, err := c.Get([]byte("b"))
  243. if err != nil || !rb.Found || string(rb.Value) != "2" {
  244. t.Fatalf("b = %+v, %v", rb, err)
  245. }
  246. rd, err := c.Get([]byte("d"))
  247. if err != ErrKeyNotFound {
  248. t.Fatalf("d err = %v, want ErrKeyNotFound", err)
  249. }
  250. if rd.Found {
  251. t.Fatalf("d should be deleted")
  252. }
  253. }
  254. func TestScanPagination(t *testing.T) {
  255. kv := newTestKVServer(t)
  256. defer kv.close()
  257. kv.maxScanPage = 2
  258. c := kv.client()
  259. defer c.Close()
  260. expected := make(map[string]string)
  261. for i := 0; i < 5; i++ {
  262. key := fmt.Sprintf("pre:%02d", i)
  263. value := fmt.Sprintf("v%d", i)
  264. if _, err := c.Put([]byte(key), []byte(value)); err != nil {
  265. t.Fatalf("put %s: %v", key, err)
  266. }
  267. expected[key] = value
  268. }
  269. scan, err := c.Scan([]byte("pre:"))
  270. if err != nil {
  271. t.Fatalf("scan: %v", err)
  272. }
  273. defer scan.Close()
  274. var entries []KVEntry
  275. pages := 0
  276. for {
  277. batch, done, err := scan.Next()
  278. if err != nil {
  279. t.Fatalf("scan next: %v", err)
  280. }
  281. pages++
  282. entries = append(entries, batch...)
  283. if done {
  284. break
  285. }
  286. }
  287. if pages < 3 {
  288. t.Fatalf("expected pagination across multiple pages, got %d", pages)
  289. }
  290. if len(entries) != 5 {
  291. t.Fatalf("entries = %d, want 5", len(entries))
  292. }
  293. for _, e := range entries {
  294. if want := expected[string(e.Key)]; string(e.Value) != want {
  295. t.Fatalf("key %q value = %q, want %q", e.Key, e.Value, want)
  296. }
  297. if e.LSN == 0 {
  298. t.Fatalf("key %q has zero lsn", e.Key)
  299. }
  300. }
  301. }
  302. func TestScanBinaryValues(t *testing.T) {
  303. kv := newTestKVServer(t)
  304. defer kv.close()
  305. c := kv.client()
  306. defer c.Close()
  307. key := []byte{0x00, 0x01, 'p'}
  308. value := []byte{0xff, 0x00, '\n', '\r'}
  309. if _, err := c.Put(key, value); err != nil {
  310. t.Fatalf("put: %v", err)
  311. }
  312. scan, err := c.Scan([]byte{0x00})
  313. if err != nil {
  314. t.Fatalf("scan: %v", err)
  315. }
  316. defer scan.Close()
  317. batch, done, err := scan.Next()
  318. if err != nil {
  319. t.Fatalf("next: %v", err)
  320. }
  321. if !done || len(batch) != 1 {
  322. t.Fatalf("done=%v len=%d", done, len(batch))
  323. }
  324. if !bytes.Equal(batch[0].Key, key) {
  325. t.Fatalf("key = %x, want %x", batch[0].Key, key)
  326. }
  327. if !bytes.Equal(batch[0].Value, value) {
  328. t.Fatalf("value = %x, want %x", batch[0].Value, value)
  329. }
  330. }
  331. func TestReadsNoDelimiterAssumptions(t *testing.T) {
  332. kv := newTestKVServer(t)
  333. defer kv.close()
  334. c := kv.client()
  335. defer c.Close()
  336. values := []string{"a\nb", "c\r\nd", "", "e\rf\ng"}
  337. for i, v := range values {
  338. if _, err := c.Put([]byte(fmt.Sprintf("p:%d", i)), []byte(v)); err != nil {
  339. t.Fatalf("put %d: %v", i, err)
  340. }
  341. }
  342. got, err := c.Reads("p:")
  343. if err != nil {
  344. t.Fatalf("reads: %v", err)
  345. }
  346. if len(got) != len(values) {
  347. t.Fatalf("reads returned %d values, want %d", len(got), len(values))
  348. }
  349. for i, v := range values {
  350. if got[i] != v {
  351. t.Fatalf("got[%d] = %q, want %q", i, got[i], v)
  352. }
  353. }
  354. }
  355. func TestHeaderCRCError(t *testing.T) {
  356. clientConn, serverConn := net.Pipe()
  357. defer clientConn.Close()
  358. defer serverConn.Close()
  359. c := pipeClient(clientConn)
  360. go func() {
  361. r := bufio.NewReader(serverConn)
  362. _, _, requestID, _, err := readFrame(r)
  363. if err != nil {
  364. return
  365. }
  366. resp := encodeResponse(opGet, requestID, getResponseBody([]byte("v"), 1))
  367. resp[6] ^= 0xff
  368. serverConn.Write(resp)
  369. }()
  370. if _, err := c.Read("k"); !errors.Is(err, ErrProtocol) {
  371. t.Fatalf("err = %v, want ErrProtocol", err)
  372. }
  373. }
  374. func TestPayloadCRCError(t *testing.T) {
  375. clientConn, serverConn := net.Pipe()
  376. defer clientConn.Close()
  377. defer serverConn.Close()
  378. c := pipeClient(clientConn)
  379. go func() {
  380. r := bufio.NewReader(serverConn)
  381. _, _, requestID, _, err := readFrame(r)
  382. if err != nil {
  383. return
  384. }
  385. resp := encodeResponse(opGet, requestID, getResponseBody([]byte("v"), 1))
  386. resp[headerSize] ^= 0xff
  387. serverConn.Write(resp)
  388. }()
  389. if _, err := c.Read("k"); !errors.Is(err, ErrProtocol) {
  390. t.Fatalf("err = %v, want ErrProtocol", err)
  391. }
  392. }
  393. func TestServerErrorStatus(t *testing.T) {
  394. clientConn, serverConn := net.Pipe()
  395. defer clientConn.Close()
  396. defer serverConn.Close()
  397. c := pipeClient(clientConn)
  398. go func() {
  399. r := bufio.NewReader(serverConn)
  400. _, _, requestID, _, err := readFrame(r)
  401. if err != nil {
  402. return
  403. }
  404. serverConn.Write(encodeResponse(opGet, requestID, errorBody("Boom")))
  405. }()
  406. _, err := c.Read("k")
  407. if err == nil {
  408. t.Fatalf("expected error")
  409. }
  410. if errors.Is(err, ErrProtocol) {
  411. t.Fatalf("server error should not be ErrProtocol: %v", err)
  412. }
  413. if !strings.Contains(err.Error(), "Boom") {
  414. t.Fatalf("err = %v", err)
  415. }
  416. }
  417. func TestRequestIDMismatch(t *testing.T) {
  418. clientConn, serverConn := net.Pipe()
  419. defer clientConn.Close()
  420. defer serverConn.Close()
  421. c := pipeClient(clientConn)
  422. go func() {
  423. r := bufio.NewReader(serverConn)
  424. _, _, requestID, _, err := readFrame(r)
  425. if err != nil {
  426. return
  427. }
  428. serverConn.Write(encodeResponse(opGet, requestID+1, getResponseBody([]byte("v"), 1)))
  429. }()
  430. if _, err := c.Read("k"); !errors.Is(err, ErrProtocol) {
  431. t.Fatalf("err = %v, want ErrProtocol", err)
  432. }
  433. }
  434. func TestOpcodeMismatch(t *testing.T) {
  435. clientConn, serverConn := net.Pipe()
  436. defer clientConn.Close()
  437. defer serverConn.Close()
  438. c := pipeClient(clientConn)
  439. go func() {
  440. r := bufio.NewReader(serverConn)
  441. _, _, requestID, _, err := readFrame(r)
  442. if err != nil {
  443. return
  444. }
  445. serverConn.Write(encodeResponse(opPut, requestID, getResponseBody([]byte("v"), 1)))
  446. }()
  447. if _, err := c.Read("k"); !errors.Is(err, ErrProtocol) {
  448. t.Fatalf("err = %v, want ErrProtocol", err)
  449. }
  450. }
  451. func TestPartialReads(t *testing.T) {
  452. clientConn, serverConn := net.Pipe()
  453. defer clientConn.Close()
  454. defer serverConn.Close()
  455. c := pipeClient(clientConn)
  456. go func() {
  457. r := bufio.NewReader(serverConn)
  458. _, _, requestID, _, err := readFrame(r)
  459. if err != nil {
  460. return
  461. }
  462. resp := encodeResponse(opGet, requestID, getResponseBody([]byte("hello"), 7))
  463. for _, b := range resp {
  464. if _, err := serverConn.Write([]byte{b}); err != nil {
  465. return
  466. }
  467. }
  468. }()
  469. res, err := c.Get([]byte("k"))
  470. if err != nil {
  471. t.Fatalf("get: %v", err)
  472. }
  473. if !res.Found || string(res.Value) != "hello" || res.LSN != 7 {
  474. t.Fatalf("res = %+v", res)
  475. }
  476. }
  477. func TestMonotonicRequestIDs(t *testing.T) {
  478. clientConn, serverConn := net.Pipe()
  479. defer clientConn.Close()
  480. defer serverConn.Close()
  481. c := pipeClient(clientConn)
  482. ids := make(chan uint64, 4)
  483. go func() {
  484. r := bufio.NewReader(serverConn)
  485. for i := 0; i < 4; i++ {
  486. _, _, requestID, _, err := readFrame(r)
  487. if err != nil {
  488. return
  489. }
  490. ids <- requestID
  491. serverConn.Write(encodeResponse(opPing, requestID, []byte{0, 0}))
  492. }
  493. }()
  494. for i := 0; i < 4; i++ {
  495. if _, _, err := c.request(opPing, nil); err != nil {
  496. t.Fatalf("ping %d: %v", i, err)
  497. }
  498. }
  499. var prev uint64
  500. for i := 0; i < 4; i++ {
  501. id := <-ids
  502. if i > 0 && id <= prev {
  503. t.Fatalf("request id not monotonic: %d then %d", prev, id)
  504. }
  505. prev = id
  506. }
  507. }
  508. func TestPoolReconnectsAfterBrokenConnection(t *testing.T) {
  509. s := newTestKVServer(t)
  510. ln, err := net.Listen("tcp", "127.0.0.1:0")
  511. if err != nil {
  512. t.Fatalf("listen: %v", err)
  513. }
  514. t.Cleanup(func() { ln.Close() })
  515. accepted := make(chan net.Conn, 8)
  516. go func() {
  517. for {
  518. conn, err := ln.Accept()
  519. if err != nil {
  520. return
  521. }
  522. accepted <- conn
  523. go s.handle(conn)
  524. }
  525. }()
  526. pool, err := NewKVPool(ln.Addr().String(), 1, 2*time.Second)
  527. if err != nil {
  528. t.Fatalf("pool: %v", err)
  529. }
  530. defer pool.Close()
  531. if err := pool.WithClient(func(c *KVClient) error { return c.Write("k", "v") }); err != nil {
  532. t.Fatalf("first write: %v", err)
  533. }
  534. if err := pool.WithClient(func(c *KVClient) error {
  535. v, err := c.Read("k")
  536. if err != nil {
  537. return err
  538. }
  539. if v != "v" {
  540. return fmt.Errorf("read = %q", v)
  541. }
  542. return nil
  543. }); err != nil {
  544. t.Fatalf("first read: %v", err)
  545. }
  546. first := <-accepted
  547. first.Close()
  548. if err := pool.WithClient(func(c *KVClient) error { return c.Write("k", "v") }); err == nil {
  549. t.Fatalf("expected write on broken connection to fail")
  550. }
  551. if err := pool.WithClient(func(c *KVClient) error {
  552. v, err := c.Read("k")
  553. if err != nil {
  554. return err
  555. }
  556. if v != "v" {
  557. return fmt.Errorf("read = %q", v)
  558. }
  559. return nil
  560. }); err != nil {
  561. t.Fatalf("read after reconnect: %v", err)
  562. }
  563. client, err := pool.Get()
  564. if err != nil {
  565. t.Fatal(err)
  566. }
  567. client.lastUsed = time.Now().Add(-31 * time.Second)
  568. pool.Put(client)
  569. second := <-accepted
  570. second.Close()
  571. if err := pool.WithClient(func(c *KVClient) error {
  572. v, err := c.Read("k")
  573. if err != nil {
  574. return err
  575. }
  576. if v != "v" {
  577. return fmt.Errorf("read = %q", v)
  578. }
  579. return nil
  580. }); err != nil {
  581. t.Fatalf("stale idle connection was not replaced before use: %v", err)
  582. }
  583. }
  584. func compareBatchResponseBody(committed bool, lsn uint64) []byte {
  585. body := make([]byte, 18)
  586. putU16(body[0:2], statusOK)
  587. if committed {
  588. body[2] = 1
  589. }
  590. putU64(body[10:18], lsn)
  591. return body
  592. }
  593. func parseComparePayload(payload []byte) ([]CompareCheck, []BatchOp, bool) {
  594. if len(payload) < 16 {
  595. return nil, nil, false
  596. }
  597. checkCount := getU32(payload[0:4])
  598. opCount := getU32(payload[4:8])
  599. metadataLen := getU32(payload[8:12])
  600. pos := 16
  601. checks := make([]CompareCheck, 0, checkCount)
  602. for i := uint32(0); i < checkCount; i++ {
  603. if len(payload)-pos < 16 {
  604. return nil, nil, false
  605. }
  606. keyLen := getU32(payload[pos : pos+4])
  607. lsn := getU64(payload[pos+8 : pos+16])
  608. pos += 16
  609. if keyLen > maxKeySize || len(payload)-pos < int(keyLen) {
  610. return nil, nil, false
  611. }
  612. checks = append(checks, CompareCheck{Key: payload[pos : pos+int(keyLen)], LSN: lsn})
  613. pos += int(keyLen)
  614. }
  615. if len(payload)-pos < int(metadataLen) {
  616. return nil, nil, false
  617. }
  618. pos += int(metadataLen)
  619. ops := make([]BatchOp, 0, opCount)
  620. for i := uint32(0); i < opCount; i++ {
  621. if len(payload)-pos < 12 {
  622. return nil, nil, false
  623. }
  624. op := payload[pos]
  625. keyLen := getU32(payload[pos+4 : pos+8])
  626. valueLen := getU32(payload[pos+8 : pos+12])
  627. pos += 12
  628. if keyLen > maxKeySize || valueLen > maxValueSize || len(payload)-pos < int(keyLen)+int(valueLen) {
  629. return nil, nil, false
  630. }
  631. ops = append(ops, BatchOp{
  632. Op: op,
  633. Key: payload[pos : pos+int(keyLen)],
  634. Value: payload[pos+int(keyLen) : pos+int(keyLen)+int(valueLen)],
  635. })
  636. pos += int(keyLen) + int(valueLen)
  637. }
  638. if pos != len(payload) {
  639. return nil, nil, false
  640. }
  641. return checks, ops, true
  642. }
  643. func TestCompareBatchWriteClientSemantics(t *testing.T) {
  644. clientConn, serverConn := net.Pipe()
  645. defer clientConn.Close()
  646. c := pipeClient(clientConn)
  647. state := make(map[string]uint64)
  648. var nextLSN uint64
  649. go func() {
  650. defer serverConn.Close()
  651. r := bufio.NewReader(serverConn)
  652. for i := 0; i < 3; i++ {
  653. opcode, _, requestID, payload, err := readFrame(r)
  654. if err != nil {
  655. return
  656. }
  657. if opcode != opCompareBatch {
  658. serverConn.Write(encodeResponse(opcode, requestID, errorBody("UnknownOpcode")))
  659. continue
  660. }
  661. checks, ops, ok := parseComparePayload(payload)
  662. if !ok {
  663. serverConn.Write(encodeResponse(opcode, requestID, errorBody("InvalidPayload")))
  664. continue
  665. }
  666. committed := true
  667. for _, check := range checks {
  668. current, found := state[string(check.Key)]
  669. if check.LSN == 0 {
  670. if found {
  671. committed = false
  672. }
  673. } else if !found || current != check.LSN {
  674. committed = false
  675. }
  676. }
  677. responseLSN := uint64(0)
  678. if committed {
  679. nextLSN++
  680. responseLSN = nextLSN
  681. for _, op := range ops {
  682. if op.Op == batchPut {
  683. state[string(op.Key)] = nextLSN
  684. } else {
  685. delete(state, string(op.Key))
  686. }
  687. }
  688. }
  689. serverConn.Write(encodeResponse(opcode, requestID, compareBatchResponseBody(committed, responseLSN)))
  690. }
  691. }()
  692. lsn, committed, err := c.CompareBatchWrite(
  693. []CompareCheck{{Key: []byte("k"), LSN: 0}},
  694. []BatchOp{{Op: batchPut, Key: []byte("k"), Value: []byte("v")}},
  695. []byte("meta"),
  696. )
  697. if err != nil {
  698. t.Fatalf("absent check: %v", err)
  699. }
  700. if !committed || lsn != 1 {
  701. t.Fatalf("absent check committed=%v lsn=%d, want committed lsn=1", committed, lsn)
  702. }
  703. lsn, committed, err = c.CompareBatchWrite(
  704. []CompareCheck{{Key: []byte("k"), LSN: 1}},
  705. []BatchOp{{Op: batchPut, Key: []byte("k"), Value: []byte("v2")}},
  706. nil,
  707. )
  708. if err != nil {
  709. t.Fatalf("matching lsn: %v", err)
  710. }
  711. if !committed || lsn != 2 {
  712. t.Fatalf("matching lsn committed=%v lsn=%d, want committed lsn=2", committed, lsn)
  713. }
  714. lsn, committed, err = c.CompareBatchWrite(
  715. []CompareCheck{{Key: []byte("k"), LSN: 1}},
  716. []BatchOp{{Op: batchPut, Key: []byte("k"), Value: []byte("v3")}},
  717. nil,
  718. )
  719. if err != nil {
  720. t.Fatalf("stale lsn: %v", err)
  721. }
  722. if committed || lsn != 0 {
  723. t.Fatalf("stale lsn committed=%v lsn=%d, want conflict committed=false lsn=0", committed, lsn)
  724. }
  725. }
  726. func TestCompareBatchWriteWireFormat(t *testing.T) {
  727. clientConn, serverConn := net.Pipe()
  728. defer clientConn.Close()
  729. c := pipeClient(clientConn)
  730. var captured []byte
  731. go func() {
  732. defer serverConn.Close()
  733. r := bufio.NewReader(serverConn)
  734. opcode, _, requestID, payload, err := readFrame(r)
  735. if err != nil {
  736. return
  737. }
  738. captured = payload
  739. serverConn.Write(encodeResponse(opcode, requestID, compareBatchResponseBody(true, 7)))
  740. }()
  741. checks := []CompareCheck{
  742. {Key: []byte("a"), LSN: 0},
  743. {Key: []byte("bb"), LSN: 42},
  744. }
  745. ops := []BatchOp{
  746. {Op: batchPut, Key: []byte("x"), Value: []byte("yy")},
  747. {Op: batchDelete, Key: []byte("z")},
  748. }
  749. if _, _, err := c.CompareBatchWrite(checks, ops, []byte("m")); err != nil {
  750. t.Fatalf("compare batch: %v", err)
  751. }
  752. if !bytes.Equal(captured[0:4], []byte{2, 0, 0, 0}) {
  753. t.Fatalf("check count = %v", captured[0:4])
  754. }
  755. if !bytes.Equal(captured[4:8], []byte{2, 0, 0, 0}) {
  756. t.Fatalf("op count = %v", captured[4:8])
  757. }
  758. if !bytes.Equal(captured[8:12], []byte{1, 0, 0, 0}) {
  759. t.Fatalf("metadata len = %v", captured[8:12])
  760. }
  761. if !bytes.Equal(captured[12:16], []byte{0, 0, 0, 0}) {
  762. t.Fatalf("reserved = %v", captured[12:16])
  763. }
  764. pos := 16
  765. if !bytes.Equal(captured[pos:pos+4], []byte{1, 0, 0, 0}) {
  766. t.Fatalf("check0 key len = %v", captured[pos:pos+4])
  767. }
  768. if getU64(captured[pos+8:pos+16]) != 0 {
  769. t.Fatalf("check0 lsn = %d", getU64(captured[pos+8:pos+16]))
  770. }
  771. if !bytes.Equal(captured[pos+16:pos+17], []byte("a")) {
  772. t.Fatalf("check0 key = %q", captured[pos+16:pos+17])
  773. }
  774. pos += 17
  775. if !bytes.Equal(captured[pos:pos+4], []byte{2, 0, 0, 0}) {
  776. t.Fatalf("check1 key len = %v", captured[pos:pos+4])
  777. }
  778. if getU64(captured[pos+8:pos+16]) != 42 {
  779. t.Fatalf("check1 lsn = %d", getU64(captured[pos+8:pos+16]))
  780. }
  781. if !bytes.Equal(captured[pos+16:pos+18], []byte("bb")) {
  782. t.Fatalf("check1 key = %q", captured[pos+16:pos+18])
  783. }
  784. pos += 18
  785. if !bytes.Equal(captured[pos:pos+1], []byte("m")) {
  786. t.Fatalf("metadata = %q", captured[pos:pos+1])
  787. }
  788. pos++
  789. if captured[pos] != batchPut {
  790. t.Fatalf("op0 opcode = %d", captured[pos])
  791. }
  792. if !bytes.Equal(captured[pos+4:pos+8], []byte{1, 0, 0, 0}) {
  793. t.Fatalf("op0 key len = %v", captured[pos+4:pos+8])
  794. }
  795. if !bytes.Equal(captured[pos+8:pos+12], []byte{2, 0, 0, 0}) {
  796. t.Fatalf("op0 value len = %v", captured[pos+8:pos+12])
  797. }
  798. pos += 12
  799. if !bytes.Equal(captured[pos:pos+3], []byte("xyy")) {
  800. t.Fatalf("op0 body = %q", captured[pos:pos+3])
  801. }
  802. pos += 3
  803. if captured[pos] != batchDelete {
  804. t.Fatalf("op1 opcode = %d", captured[pos])
  805. }
  806. if !bytes.Equal(captured[pos+4:pos+8], []byte{1, 0, 0, 0}) {
  807. t.Fatalf("op1 key len = %v", captured[pos+4:pos+8])
  808. }
  809. if !bytes.Equal(captured[pos+8:pos+12], []byte{0, 0, 0, 0}) {
  810. t.Fatalf("op1 value len = %v", captured[pos+8:pos+12])
  811. }
  812. pos += 12
  813. if !bytes.Equal(captured[pos:pos+1], []byte("z")) {
  814. t.Fatalf("op1 key = %q", captured[pos:pos+1])
  815. }
  816. pos++
  817. if pos != len(captured) {
  818. t.Fatalf("payload trailing bytes: got len %d want %d", len(captured), pos)
  819. }
  820. }
  821. func TestCompareBatchWriteRejectsInvalidInput(t *testing.T) {
  822. c := &KVClient{}
  823. largeKey := make([]byte, maxKeySize+1)
  824. if _, _, err := c.CompareBatchWrite(nil, nil, nil); err == nil {
  825. t.Fatal("accepted empty ops")
  826. }
  827. if _, _, err := c.CompareBatchWrite([]CompareCheck{{Key: largeKey}}, []BatchOp{{Op: batchPut, Key: []byte("k")}}, nil); err == nil {
  828. t.Fatal("accepted oversized check key")
  829. }
  830. if _, _, err := c.CompareBatchWrite(nil, []BatchOp{{Op: 99, Key: []byte("k")}}, nil); err == nil {
  831. t.Fatal("accepted invalid batch opcode")
  832. }
  833. if _, _, err := c.CompareBatchWrite(nil, []BatchOp{{Op: batchDelete, Key: []byte("k"), Value: []byte("v")}}, nil); err == nil {
  834. t.Fatal("accepted delete with value")
  835. }
  836. }
  837. func TestCompareBatchWriteMalformedResponse(t *testing.T) {
  838. clientConn, serverConn := net.Pipe()
  839. defer clientConn.Close()
  840. c := pipeClient(clientConn)
  841. go func() {
  842. defer serverConn.Close()
  843. r := bufio.NewReader(serverConn)
  844. opcode, _, requestID, _, err := readFrame(r)
  845. if err != nil {
  846. return
  847. }
  848. serverConn.Write(encodeResponse(opcode, requestID, []byte{0, 0, 1}))
  849. }()
  850. if _, _, err := c.CompareBatchWrite(
  851. []CompareCheck{{Key: []byte("k"), LSN: 0}},
  852. []BatchOp{{Op: batchPut, Key: []byte("k"), Value: []byte("v")}},
  853. nil,
  854. ); !errors.Is(err, ErrProtocol) {
  855. t.Fatalf("err = %v, want ErrProtocol", err)
  856. }
  857. }