From 1d9d74cfde17ef0341a23e56cf72aa382da297ba Mon Sep 17 00:00:00 2001 From: Emmanuel Di Pretoro Date: Tue, 4 Aug 2026 21:39:01 +0200 Subject: [PATCH] Adding the challenge 0402 --- 0402/assert.go | 9 + 0402/cell.go | 57 ++++++ 0402/cell_test.go | 27 +++ 0402/kv.go | 127 +++++++++++++ 0402/kv_entry.go | 63 +++++++ 0402/kv_test.go | 212 ++++++++++++++++++++++ 0402/log.go | 40 ++++ 0402/os_other.go | 11 ++ 0402/os_unix.go | 34 ++++ 0402/row.go | 95 ++++++++++ 0402/row_test.go | 38 ++++ 0402/sql_parser.go | 392 ++++++++++++++++++++++++++++++++++++++++ 0402/sql_parser_test.go | 110 +++++++++++ 0402/table.go | 267 +++++++++++++++++++++++++++ 0402/table_test.go | 126 +++++++++++++ 15 files changed, 1608 insertions(+) create mode 100644 0402/assert.go create mode 100644 0402/cell.go create mode 100644 0402/cell_test.go create mode 100644 0402/kv.go create mode 100644 0402/kv_entry.go create mode 100644 0402/kv_test.go create mode 100644 0402/log.go create mode 100644 0402/os_other.go create mode 100644 0402/os_unix.go create mode 100644 0402/row.go create mode 100644 0402/row_test.go create mode 100644 0402/sql_parser.go create mode 100644 0402/sql_parser_test.go create mode 100644 0402/table.go create mode 100644 0402/table_test.go diff --git a/0402/assert.go b/0402/assert.go new file mode 100644 index 0000000..010bc17 --- /dev/null +++ b/0402/assert.go @@ -0,0 +1,9 @@ +package db0402 + +func check(cond bool) { + if !cond { + panic("assertion failure") + } +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/cell.go b/0402/cell.go new file mode 100644 index 0000000..83f07e1 --- /dev/null +++ b/0402/cell.go @@ -0,0 +1,57 @@ +package db0402 + +import ( + "encoding/binary" + "errors" + "slices" +) + +type CellType uint8 + +const ( + TypeI64 CellType = 1 + TypeStr CellType = 2 +) + +type Cell struct { + Type CellType + I64 int64 + Str []byte +} + +func (cell *Cell) Encode(toAppend []byte) []byte { + switch cell.Type { + case TypeI64: + return binary.LittleEndian.AppendUint64(toAppend, uint64(cell.I64)) + case TypeStr: + toAppend = binary.LittleEndian.AppendUint32(toAppend, uint32(len(cell.Str))) + return append(toAppend, cell.Str...) + default: + panic("unreachable") + } +} + +func (cell *Cell) Decode(data []byte) (rest []byte, err error) { + switch cell.Type { + case TypeI64: + if len(data) < 8 { + return data, errors.New("expect more data") + } + cell.I64 = int64(binary.LittleEndian.Uint64(data[0:8])) + return data[8:], nil + case TypeStr: + if len(data) < 4 { + return data, errors.New("expect more data") + } + size := int(binary.LittleEndian.Uint32(data[0:4])) + if len(data) < 4+size { + return data, errors.New("expect more data") + } + cell.Str = slices.Clone(data[4 : 4+size]) + return data[4+size:], nil + default: + panic("unreachable") + } +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/cell_test.go b/0402/cell_test.go new file mode 100644 index 0000000..41f0d70 --- /dev/null +++ b/0402/cell_test.go @@ -0,0 +1,27 @@ +package db0402 + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestTableCell(t *testing.T) { + cell := Cell{Type: TypeI64, I64: -2} + data := []byte{0xfe, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff} + assert.Equal(t, data, cell.Encode(nil)) + decoded := Cell{Type: TypeI64} + rest, err := decoded.Decode(data) + assert.True(t, len(rest) == 0 && err == nil) + assert.Equal(t, cell, decoded) + + cell = Cell{Type: TypeStr, Str: []byte("asdf")} + data = []byte{4, 0, 0, 0, 'a', 's', 'd', 'f'} + assert.Equal(t, data, cell.Encode(nil)) + decoded = Cell{Type: TypeStr} + rest, err = decoded.Decode(data) + assert.True(t, len(rest) == 0 && err == nil) + assert.Equal(t, cell, decoded) +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/kv.go b/0402/kv.go new file mode 100644 index 0000000..7fff804 --- /dev/null +++ b/0402/kv.go @@ -0,0 +1,127 @@ +package db0402 + +import ( + "bytes" + "slices" +) + +type KV struct { + log Log + keys [][]byte + vals [][]byte +} + +func (kv *KV) Open() error { + if err := kv.log.Open(); err != nil { + return err + } + + entries := []Entry{} + for { + ent := Entry{} + eof, err := kv.log.Read(&ent) + if err != nil { + return err + } else if eof { + break + } + entries = append(entries, ent) + } + + slices.SortStableFunc(entries, func(a, b Entry) int { + return bytes.Compare(a.key, b.key) + }) + kv.keys, kv.vals = kv.keys[:0], kv.vals[:0] + for _, ent := range entries { + n := len(kv.keys) + if n > 0 && bytes.Equal(kv.keys[n-1], ent.key) { + kv.keys, kv.vals = kv.keys[:n-1], kv.vals[:n-1] + } + if !ent.deleted { + kv.keys = append(kv.keys, ent.key) + kv.vals = append(kv.vals, ent.val) + } + } + return nil +} + +func (kv *KV) Close() error { return kv.log.Close() } + +func (kv *KV) Get(key []byte) (val []byte, ok bool, err error) { + if idx, ok := slices.BinarySearchFunc(kv.keys, key, bytes.Compare); ok { + return kv.vals[idx], true, nil + } + return nil, false, nil +} + +type UpdateMode int + +const ( + ModeUpsert UpdateMode = 0 // insert or update + ModeInsert UpdateMode = 1 // insert new + ModeUpdate UpdateMode = 2 // update existing +) + +func (kv *KV) SetEx(key []byte, val []byte, mode UpdateMode) (updated bool, err error) { + idx, exist := slices.BinarySearchFunc(kv.keys, key, bytes.Compare) + switch mode { + case ModeUpsert: + updated = !exist || !bytes.Equal(kv.vals[idx], val) + case ModeInsert: + updated = !exist + case ModeUpdate: + updated = exist && !bytes.Equal(kv.vals[idx], val) + default: + panic("unreachable") + } + if updated { + if err = kv.log.Write(&Entry{key: key, val: val}); err != nil { + return false, err + } + if exist { + kv.vals[idx] = val + } else { + kv.keys = slices.Insert(kv.keys, idx, key) + kv.vals = slices.Insert(kv.vals, idx, val) + } + } + return +} + +func (kv *KV) Set(key []byte, val []byte) (updated bool, err error) { + return kv.SetEx(key, val, ModeUpsert) +} + +func (kv *KV) Del(key []byte) (deleted bool, err error) { + if idx, ok := slices.BinarySearchFunc(kv.keys, key, bytes.Compare); ok { + if err = kv.log.Write(&Entry{key: key, deleted: true}); err != nil { + return false, err + } + kv.keys = slices.Delete(kv.keys, idx, idx+1) + kv.vals = slices.Delete(kv.vals, idx, idx+1) + return true, nil + } + return false, nil +} + +type KVIterator struct { + keys [][]byte + vals [][]byte + pos int +} + +func (kv *KV) Seek(key []byte) (*KVIterator, error) + +func (iter *KVIterator) Valid() bool { + return 0 <= iter.pos && iter.pos < len(iter.keys) +} + +func (iter *KVIterator) Key() []byte + +func (iter *KVIterator) Val() []byte + +func (iter *KVIterator) Next() error + +func (iter *KVIterator) Prev() error + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/kv_entry.go b/0402/kv_entry.go new file mode 100644 index 0000000..dd55512 --- /dev/null +++ b/0402/kv_entry.go @@ -0,0 +1,63 @@ +package db0402 + +import ( + "encoding/binary" + "errors" + "hash/crc32" + "io" +) + +type Entry struct { + key []byte + val []byte + deleted bool +} + +func (ent *Entry) Encode() []byte { + data := make([]byte, 4+4+4+1+len(ent.key)+len(ent.val)) + binary.LittleEndian.PutUint32(data[4:8], uint32(len(ent.key))) + copy(data[4+4+4+1:], ent.key) + if ent.deleted { + data[4+4+4] = 1 + } else { + binary.LittleEndian.PutUint32(data[8:12], uint32(len(ent.val))) + copy(data[4+4+4+1+len(ent.key):], ent.val) + } + binary.LittleEndian.PutUint32(data[0:4], crc32.ChecksumIEEE(data[4:])) + return data +} + +var ErrBadSum = errors.New("bad checksum") + +func (ent *Entry) Decode(r io.Reader) error { + var header [4 + 4 + 4 + 1]byte + if _, err := io.ReadFull(r, header[:]); err != nil { + return err + } + klen := int(binary.LittleEndian.Uint32(header[4:8])) + vlen := int(binary.LittleEndian.Uint32(header[8:12])) + deleted := header[4+4+4] + + data := make([]byte, klen+vlen) + if _, err := io.ReadFull(r, data); err != nil { + return err + } + + h := crc32.NewIEEE() + h.Write(header[4:]) + h.Write(data) + if h.Sum32() != binary.LittleEndian.Uint32(header[0:4]) { + return ErrBadSum + } + + ent.key = data[:klen] + if deleted != 0 { + ent.deleted = true + } else { + ent.deleted = false + ent.val = data[klen:] + } + return nil +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/kv_test.go b/0402/kv_test.go new file mode 100644 index 0000000..ca11357 --- /dev/null +++ b/0402/kv_test.go @@ -0,0 +1,212 @@ +package db0402 + +import ( + "bytes" + "os" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestKVBasic(t *testing.T) { + kv := KV{} + kv.log.FileName = ".test_db" + defer os.Remove(kv.log.FileName) + + os.Remove(kv.log.FileName) + err := kv.Open() + assert.Nil(t, err) + defer kv.Close() + + updated, err := kv.Set([]byte("k1"), []byte("v1")) + assert.True(t, updated && err == nil) + + val, ok, err := kv.Get([]byte("k1")) + assert.True(t, string(val) == "v1" && ok && err == nil) + + _, ok, err = kv.Get([]byte("xxx")) + assert.True(t, !ok && err == nil) + + updated, err = kv.Del([]byte("xxx")) + assert.True(t, !updated && err == nil) + + updated, err = kv.Del([]byte("k1")) + assert.True(t, updated && err == nil) + + _, ok, err = kv.Get([]byte("xxx")) + assert.True(t, !ok && err == nil) + + updated, err = kv.Set([]byte("k2"), []byte("v2")) + assert.True(t, updated && err == nil) + + // reopen + kv.Close() + err = kv.Open() + assert.Nil(t, err) + + _, ok, err = kv.Get([]byte("k1")) + assert.True(t, !ok && err == nil) + val, ok, err = kv.Get([]byte("k2")) + assert.True(t, string(val) == "v2" && ok && err == nil) +} + +func TestKVUpdateMode(t *testing.T) { + kv := KV{} + kv.log.FileName = ".test_db" + defer os.Remove(kv.log.FileName) + + os.Remove(kv.log.FileName) + err := kv.Open() + assert.Nil(t, err) + defer kv.Close() + + updated, err := kv.SetEx([]byte("k1"), []byte("v1"), ModeUpdate) + assert.True(t, !updated && err == nil) + + updated, err = kv.SetEx([]byte("k1"), []byte("v1"), ModeUpdate) + assert.True(t, !updated && err == nil) + + updated, err = kv.SetEx([]byte("k1"), []byte("v1"), ModeInsert) + assert.True(t, updated && err == nil) + + updated, err = kv.SetEx([]byte("k1"), []byte("xx"), ModeInsert) + assert.True(t, !updated && err == nil) + + updated, err = kv.SetEx([]byte("k1"), []byte("yy"), ModeUpdate) + assert.True(t, updated && err == nil) + + updated, err = kv.SetEx([]byte("k1"), []byte("zz"), ModeUpsert) + assert.True(t, updated && err == nil) + + updated, err = kv.SetEx([]byte("k2"), []byte("tt"), ModeUpsert) + assert.True(t, updated && err == nil) +} + +func TestKVRecovery(t *testing.T) { + kv := KV{} + kv.log.FileName = ".test_db" + defer os.Remove(kv.log.FileName) + + prepare := func() { + os.Remove(kv.log.FileName) + + err := kv.Open() + assert.Nil(t, err) + defer kv.Close() + + updated, err := kv.Set([]byte("k1"), []byte("v1")) + assert.True(t, updated && err == nil) + updated, err = kv.Set([]byte("k2"), []byte("v2")) + assert.True(t, updated && err == nil) + } + + prepare() + // simulate truncated log + fp, _ := os.OpenFile(kv.log.FileName, os.O_RDWR, 0o644) + st, _ := fp.Stat() + fp.Truncate(st.Size() - 1) + fp.Close() + // reopen + err := kv.Open() + assert.Nil(t, err) + // test + val, ok, err := kv.Get([]byte("k1")) + assert.True(t, string(val) == "v1" && ok && err == nil) + _, ok, err = kv.Get([]byte("k2")) // bad + assert.True(t, !ok && err == nil) + kv.Close() + + prepare() + // simulate bad checksum + fp, _ = os.OpenFile(kv.log.FileName, os.O_RDWR, 0o644) + st, _ = fp.Stat() + fp.WriteAt([]byte{0}, st.Size()-1) + fp.Close() + // reopen + err = kv.Open() + assert.Nil(t, err) + // test + val, ok, err = kv.Get([]byte("k1")) + assert.True(t, string(val) == "v1" && ok && err == nil) + _, ok, err = kv.Get([]byte("k2")) // bad + assert.True(t, !ok && err == nil) + kv.Close() +} + +func TestEntryEncode(t *testing.T) { + ent := Entry{key: []byte("k1"), val: []byte("xxx")} + data := []byte{0xe9, 0xec, 0x4d, 0x9e, 2, 0, 0, 0, 3, 0, 0, 0, 0, 'k', '1', 'x', 'x', 'x'} + + assert.Equal(t, data, ent.Encode()) + + decoded := Entry{} + err := decoded.Decode(bytes.NewBuffer(data)) + assert.Nil(t, err) + assert.Equal(t, ent, decoded) + + ent = Entry{key: []byte("k1"), deleted: true} + data = []byte{0x4c, 0xd0, 0xfe, 0xe5, 2, 0, 0, 0, 0, 0, 0, 0, 1, 'k', '1'} + + assert.Equal(t, data, ent.Encode()) + + decoded = Entry{} + err = decoded.Decode(bytes.NewBuffer(data)) + assert.Nil(t, err) + assert.Equal(t, ent, decoded) +} + +func TestKVSeek(t *testing.T) { + kv := KV{} + kv.log.FileName = ".test_db" + defer os.Remove(kv.log.FileName) + + os.Remove(kv.log.FileName) + err := kv.Open() + assert.Nil(t, err) + defer kv.Close() + + keys := []string{"c", "e", "g"} + vals := []string{"3", "5", "7"} + for i := range keys { + _, _ = kv.Set([]byte(keys[i]), []byte(vals[i])) + } + + iter, err := kv.Seek([]byte("a")) + require.Nil(t, err) + for i := range keys { + assert.True(t, iter.Valid()) + assert.Equal(t, []byte(keys[i]), iter.Key()) + assert.Equal(t, []byte(vals[i]), iter.Val()) + err = iter.Next() + require.Nil(t, err) + } + assert.False(t, iter.Valid()) + + err = iter.Prev() + require.Nil(t, err) + for i := len(keys) - 1; i >= 0; i-- { + assert.True(t, iter.Valid()) + assert.Equal(t, []byte(keys[i]), iter.Key()) + assert.Equal(t, []byte(vals[i]), iter.Val()) + err = iter.Prev() + require.Nil(t, err) + } + assert.False(t, iter.Valid()) + + iter, err = kv.Seek([]byte("f")) + require.Nil(t, err) + assert.True(t, iter.Valid()) + assert.Equal(t, []byte("g"), iter.Key()) + + iter, err = kv.Seek([]byte("g")) + require.Nil(t, err) + assert.True(t, iter.Valid()) + assert.Equal(t, []byte("g"), iter.Key()) + + iter, err = kv.Seek([]byte("h")) + require.Nil(t, err) + assert.False(t, iter.Valid()) +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/log.go b/0402/log.go new file mode 100644 index 0000000..f871344 --- /dev/null +++ b/0402/log.go @@ -0,0 +1,40 @@ +package db0402 + +import ( + "io" + "os" +) + +type Log struct { + FileName string + fp *os.File +} + +func (log *Log) Open() (err error) { + log.fp, err = createFileSync(log.FileName) + return err +} + +func (log *Log) Close() error { + return log.fp.Close() +} + +func (log *Log) Write(ent *Entry) error { + if _, err := log.fp.Write(ent.Encode()); err != nil { + return err + } + return log.fp.Sync() // fsync +} + +func (log *Log) Read(ent *Entry) (eof bool, err error) { + err = ent.Decode(log.fp) + if err == io.EOF || err == io.ErrUnexpectedEOF || err == ErrBadSum { + return true, nil + } else if err != nil { + return false, err + } else { + return false, nil + } +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/os_other.go b/0402/os_other.go new file mode 100644 index 0000000..e5386f1 --- /dev/null +++ b/0402/os_other.go @@ -0,0 +1,11 @@ +//go:build !unix + +package db0402 + +import "os" + +func createFileSync(file string) (*os.File, error) { + return os.OpenFile(file, os.O_RDWR|os.O_CREATE, 0o644) +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/os_unix.go b/0402/os_unix.go new file mode 100644 index 0000000..b789ddc --- /dev/null +++ b/0402/os_unix.go @@ -0,0 +1,34 @@ +//go:build unix + +package db0402 + +import ( + "os" + "path" + "syscall" +) + +// open or create a file and fsync the directory +func createFileSync(file string) (*os.File, error) { + fp, err := os.OpenFile(file, os.O_RDWR|os.O_CREATE, 0o644) + if err != nil { + return nil, err + } + if err = syncDir(path.Base(file)); err != nil { + _ = fp.Close() + return nil, err + } + return fp, err +} + +func syncDir(file string) error { + flags := os.O_RDONLY | syscall.O_DIRECTORY + dirfd, err := syscall.Open(path.Dir(file), flags, 0o644) + if err != nil { + return err + } + defer syscall.Close(dirfd) + return syscall.Fsync(dirfd) +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/row.go b/0402/row.go new file mode 100644 index 0000000..2fa2070 --- /dev/null +++ b/0402/row.go @@ -0,0 +1,95 @@ +package db0402 + +import ( + "errors" + "slices" +) + +type Schema struct { + Table string + Cols []Column + PKey []int // indexes of primary key columns +} + +type Column struct { + Name string + Type CellType +} + +type Row []Cell + +func (schema *Schema) NewRow() Row { + return make(Row, len(schema.Cols)) +} + +func (row Row) EncodeKey(schema *Schema) (key []byte) { + key = append(key, []byte(schema.Table)...) + key = append(key, 0x00) + check(len(row) == len(schema.Cols)) + for idx, value := range row { + if slices.Contains(schema.PKey, idx) { + check(value.Type == schema.Cols[idx].Type) + key = row[idx].Encode(key) + } + } + return key +} + +func (row Row) EncodeVal(schema *Schema) (val []byte) { + check(len(row) == len(schema.Cols)) + for idx, value := range row { + if !slices.Contains(schema.PKey, idx) { + check(value.Type == schema.Cols[idx].Type) + val = row[idx].Encode(val) + } + } + return val +} + +func (row Row) DecodeKey(schema *Schema, key []byte) (err error) { + check(len(row) == len(schema.Cols)) + + if len(key) < len(schema.Table)+1 { + return errors.New("bad key") + } + if string(key[:len(schema.Table)+1]) != schema.Table+"\x00" { + return errors.New("bad key") + } + key = key[len(schema.Table)+1:] + + for idx, col := range schema.Cols { + if !slices.Contains(schema.PKey, idx) { + continue + } + row[idx] = Cell{Type: col.Type} + if key, err = row[idx].Decode(key); err != nil { + return err + } + } + + if len(key) != 0 { + return errors.New("trailing garbage") + } + return nil +} + +func (row Row) DecodeVal(schema *Schema, val []byte) (err error) { + check(len(row) == len(schema.Cols)) + + for idx, col := range schema.Cols { + if slices.Contains(schema.PKey, idx) { + continue + } + row[idx] = Cell{Type: col.Type} + if val, err = row[idx].Decode(val); err != nil { + return err + } + } + + if len(val) != 0 { + return errors.New("trailing garbage") + } + return nil +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/row_test.go b/0402/row_test.go new file mode 100644 index 0000000..76630d9 --- /dev/null +++ b/0402/row_test.go @@ -0,0 +1,38 @@ +package db0402 + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestRowEncode(t *testing.T) { + schema := &Schema{ + Table: "link", + Cols: []Column{ + {Name: "time", Type: TypeI64}, + {Name: "src", Type: TypeStr}, + {Name: "dst", Type: TypeStr}, + }, + PKey: []int{1, 2}, // (src, dst) + } + + row := Row{ + Cell{Type: TypeI64, I64: 123}, + Cell{Type: TypeStr, Str: []byte("a")}, + Cell{Type: TypeStr, Str: []byte("b")}, + } + key := []byte{'l', 'i', 'n', 'k', 0, 1, 0, 0, 0, 'a', 1, 0, 0, 0, 'b'} + val := []byte{123, 0, 0, 0, 0, 0, 0, 0} + assert.Equal(t, key, row.EncodeKey(schema)) + assert.Equal(t, val, row.EncodeVal(schema)) + + decoded := schema.NewRow() + err := decoded.DecodeKey(schema, key) + assert.Nil(t, err) + err = decoded.DecodeVal(schema, val) + assert.Nil(t, err) + assert.Equal(t, row, decoded) +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/sql_parser.go b/0402/sql_parser.go new file mode 100644 index 0000000..12e1a50 --- /dev/null +++ b/0402/sql_parser.go @@ -0,0 +1,392 @@ +package db0402 + +import ( + "errors" + "strconv" + "strings" +) + +type Parser struct { + buf string + pos int +} + +func NewParser(s string) Parser { + return Parser{buf: s, pos: 0} +} + +type StmtSelect struct { + table string + cols []string + keys []NamedCell +} + +type NamedCell struct { + column string + value Cell +} + +type StmtCreatTable struct { + table string + cols []Column + pkey []string +} + +type StmtInsert struct { + table string + value []Cell +} + +type StmtUpdate struct { + table string + keys []NamedCell + value []NamedCell +} + +type StmtDelete struct { + table string + keys []NamedCell +} + +func isSpace(ch byte) bool { + switch ch { + case '\t', '\n', '\v', '\f', '\r', ' ': + return true + } + return false +} +func isAlpha(ch byte) bool { + return 'a' <= (ch|32) && (ch|32) <= 'z' +} +func isDigit(ch byte) bool { + return '0' <= ch && ch <= '9' +} +func isNameStart(ch byte) bool { + return isAlpha(ch) || ch == '_' +} +func isNameContinue(ch byte) bool { + return isAlpha(ch) || isDigit(ch) || ch == '_' +} +func isSeparator(ch byte) bool { + return ch < 128 && !isNameContinue(ch) +} + +func (p *Parser) skipSpaces() { + for p.pos < len(p.buf) && isSpace(p.buf[p.pos]) { + p.pos += 1 + } +} + +func (p *Parser) tryKeyword(kws ...string) bool { + save := p.pos + for _, kw := range kws { + p.skipSpaces() + if !(p.pos+len(kw) <= len(p.buf) && strings.EqualFold(p.buf[p.pos:p.pos+len(kw)], kw)) { + p.pos = save + return false + } + if p.pos+len(kw) < len(p.buf) && !isSeparator(p.buf[p.pos+len(kw)]) { + p.pos = save + return false + } + p.pos += len(kw) + } + return true +} + +func (p *Parser) tryPunctuation(tok string) bool { + p.skipSpaces() + if !(p.pos+len(tok) <= len(p.buf) && p.buf[p.pos:p.pos+len(tok)] == tok) { + return false + } + p.pos += len(tok) + return true +} + +func (p *Parser) tryName() (string, bool) { + p.skipSpaces() + start, cur := p.pos, p.pos + if !(cur < len(p.buf) && isNameStart(p.buf[cur])) { + return "", false + } + cur++ + for cur < len(p.buf) && isNameContinue(p.buf[cur]) { + cur++ + } + p.pos = cur + return p.buf[start:cur], true +} + +func (p *Parser) parseValue(out *Cell) error { + p.skipSpaces() + if p.pos >= len(p.buf) { + return errors.New("expect value") + } + ch := p.buf[p.pos] + if ch == '"' || ch == '\'' { + return p.parseString(out) + } else if isDigit(ch) || ch == '-' || ch == '+' { + return p.parseInt(out) + } else { + return errors.New("expect value") + } +} + +func (p *Parser) parseString(out *Cell) error { + quote := p.buf[p.pos] + cur := p.pos + 1 + for cur < len(p.buf) { + ch := p.buf[cur] + if ch == '\\' { + cur++ + if cur < len(p.buf) && (p.buf[cur] == '"' || p.buf[cur] == '\'') { + out.Str = append(out.Str, p.buf[cur]) + cur++ + } else { + return errors.New("bad escape") + } + } else if ch == quote { + out.Type = TypeStr + p.pos = cur + 1 + return nil + } else { + out.Str = append(out.Str, p.buf[cur]) + cur++ + } + } + return errors.New("string is not terminated") +} + +func (p *Parser) parseInt(out *Cell) (err error) { + start, cur := p.pos, p.pos + if p.buf[cur] == '-' || p.buf[cur] == '+' { + cur++ + } + for cur < len(p.buf) && isDigit(p.buf[cur]) { + cur++ + } + + if out.I64, err = strconv.ParseInt(p.buf[start:cur], 10, 64); err != nil { + return err + } + out.Type = TypeI64 + p.pos = cur + return nil +} + +func (p *Parser) parseEqual(out *NamedCell) error { + var ok bool + out.column, ok = p.tryName() + if !ok { + return errors.New("expect column") + } + if !p.tryPunctuation("=") { + return errors.New("expect =") + } + return p.parseValue(&out.value) +} + +func (p *Parser) parseSelect(out *StmtSelect) error { + for !p.tryKeyword("FROM") { + if len(out.cols) > 0 && !p.tryPunctuation(",") { + return errors.New("expect comma") + } + if name, ok := p.tryName(); ok { + out.cols = append(out.cols, name) + } else { + return errors.New("expect column") + } + } + if len(out.cols) == 0 { + return errors.New("expect column list") + } + var ok bool + if out.table, ok = p.tryName(); !ok { + return errors.New("expect table name") + } + return p.parseWhere(&out.keys) +} + +func (p *Parser) parseWhere(out *[]NamedCell) error { + if !p.tryKeyword("WHERE") { + return errors.New("expect keyword") + } + for !p.tryPunctuation(";") { + expr := NamedCell{} + if len(*out) > 0 && !p.tryKeyword("AND") { + return errors.New("expect AND") + } + if err := p.parseEqual(&expr); err != nil { + return err + } + *out = append(*out, expr) + } + if len(*out) == 0 { + return errors.New("expect where clause") + } + return nil +} + +func (p *Parser) parseCommaList(item func() error) error { + if !p.tryPunctuation("(") { + return errors.New("expect (") + } + comma := false + for !p.tryPunctuation(")") { + if comma && !p.tryPunctuation(",") { + return errors.New("expect ,") + } + comma = true + if err := item(); err != nil { + return err + } + } + return nil +} + +func (p *Parser) parseNameItem(out *[]string) error { + name, ok := p.tryName() + if !ok { + return errors.New("expect name") + } + *out = append(*out, name) + return nil +} + +func (p *Parser) parseCreateTableItem(out *StmtCreatTable) error { + if p.tryKeyword("PRIMARY", "KEY") { + return p.parseCommaList(func() error { return p.parseNameItem(&out.pkey) }) + } + + var ok bool + col := Column{} + if col.Name, ok = p.tryName(); !ok { + return errors.New("expect name") + } + kind, ok := p.tryName() + if !ok { + return errors.New("expect name") + } + switch kind { + case "int64": + col.Type = TypeI64 + case "string": + col.Type = TypeStr + default: + return errors.New("unknown column type") + } + out.cols = append(out.cols, col) + return nil +} + +func (p *Parser) parseCreateTable(out *StmtCreatTable) error { + var ok bool + if out.table, ok = p.tryName(); !ok { + return errors.New("expect table name") + } + err := p.parseCommaList(func() error { return p.parseCreateTableItem(out) }) + if err != nil { + return err + } + if !p.tryPunctuation(";") { + return errors.New("expect ;") + } + return nil +} + +func (p *Parser) parseValueItem(out *[]Cell) error { + cell := Cell{} + if err := p.parseValue(&cell); err != nil { + return err + } + *out = append(*out, cell) + return nil +} + +func (p *Parser) parseInsert(out *StmtInsert) error { + var ok bool + if out.table, ok = p.tryName(); !ok { + return errors.New("expect table name") + } + if !p.tryKeyword("VALUES") { + return errors.New("expect VALUES") + } + err := p.parseCommaList(func() error { return p.parseValueItem(&out.value) }) + if err != nil { + return err + } + if !p.tryPunctuation(";") { + return errors.New("expect ;") + } + return nil +} + +func (p *Parser) parseUpdate(out *StmtUpdate) error { + var ok bool + if out.table, ok = p.tryName(); !ok { + return errors.New("expect table name") + } + if !p.tryKeyword("SET") { + return errors.New("expect SET") + } + for !p.tryKeyword("WHERE") { + expr := NamedCell{} + if len(out.value) > 0 && !p.tryKeyword(",") { + return errors.New("expect ,") + } + if err := p.parseEqual(&expr); err != nil { + return err + } + out.value = append(out.value, expr) + } + if len(out.value) == 0 { + return errors.New("expect assignment list") + } + p.pos -= len("WHERE") + return p.parseWhere(&out.keys) +} + +func (p *Parser) parseDelete(out *StmtDelete) error { + var ok bool + if out.table, ok = p.tryName(); !ok { + return errors.New("expect table name") + } + return p.parseWhere(&out.keys) +} + +func (p *Parser) parseStmt() (out interface{}, err error) { + if p.tryKeyword("SELECT") { + stmt := &StmtSelect{} + err = p.parseSelect(stmt) + out = stmt + } else if p.tryKeyword("CREATE", "TABLE") { + stmt := &StmtCreatTable{} + err = p.parseCreateTable(stmt) + out = stmt + } else if p.tryKeyword("INSERT", "INTO") { + stmt := &StmtInsert{} + err = p.parseInsert(stmt) + out = stmt + } else if p.tryKeyword("UPDATE") { + stmt := &StmtUpdate{} + err = p.parseUpdate(stmt) + out = stmt + } else if p.tryKeyword("DELETE", "FROM") { + stmt := &StmtDelete{} + err = p.parseDelete(stmt) + out = stmt + } else { + err = errors.New("unknown statement") + } + if err != nil { + return nil, err + } + return out, nil +} + +func (p *Parser) isEnd() bool { + p.skipSpaces() + return p.pos >= len(p.buf) +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/sql_parser_test.go b/0402/sql_parser_test.go new file mode 100644 index 0000000..1fa627a --- /dev/null +++ b/0402/sql_parser_test.go @@ -0,0 +1,110 @@ +package db0402 + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestParseName(t *testing.T) { + p := NewParser(" a b0 _0_ 123 ") + name, ok := p.tryName() + assert.True(t, ok && name == "a") + name, ok = p.tryName() + assert.True(t, ok && name == "b0") + name, ok = p.tryName() + assert.True(t, ok && name == "_0_") + _, ok = p.tryName() + assert.False(t, ok) +} + +func TestParseKeyword(t *testing.T) { + p := NewParser(" select HELLO ") + assert.False(t, p.tryKeyword("sel")) + assert.True(t, p.tryKeyword("SELECT")) + assert.True(t, p.tryKeyword("hello") && p.isEnd()) + + p = NewParser(" select HELLO ") + assert.False(t, p.tryKeyword("select", "hi")) + assert.True(t, p.tryKeyword("select", "hello") && p.isEnd()) +} + +func testParseValue(t *testing.T, s string, ref Cell) { + p := NewParser(s) + out := Cell{} + err := p.parseValue(&out) + assert.Nil(t, err) + assert.True(t, p.isEnd()) + assert.Equal(t, ref, out) +} + +func TestParseValue(t *testing.T) { + testParseValue(t, " -123 ", Cell{Type: TypeI64, I64: -123}) + testParseValue(t, ` 'abc\'\"d' `, Cell{Type: TypeStr, Str: []byte("abc'\"d")}) + testParseValue(t, ` "abc\'\"d" `, Cell{Type: TypeStr, Str: []byte("abc'\"d")}) +} + +func testParseStmt(t *testing.T, s string, ref interface{}) { + p := NewParser(s) + out, err := p.parseStmt() + assert.Nil(t, err) + assert.True(t, p.isEnd()) + assert.Equal(t, ref, out) +} + +func TestParseStmt(t *testing.T) { + var stmt interface{} + s := "select a from t where c=1;" + stmt = &StmtSelect{ + table: "t", + cols: []string{"a"}, + keys: []NamedCell{{column: "c", value: Cell{Type: TypeI64, I64: 1}}}, + } + testParseStmt(t, s, stmt) + + s = "select a,b_02 from T where c=1 and d='e';" + stmt = &StmtSelect{ + table: "T", + cols: []string{"a", "b_02"}, + keys: []NamedCell{ + {column: "c", value: Cell{Type: TypeI64, I64: 1}}, + {column: "d", value: Cell{Type: TypeStr, Str: []byte("e")}}, + }, + } + testParseStmt(t, s, stmt) + + s = "select a, b_02 from T where c = 1 and d = 'e' ; " + testParseStmt(t, s, stmt) + + s = "create table t (a string, b int64, primary key (b));" + stmt = &StmtCreatTable{ + table: "t", + cols: []Column{{"a", TypeStr}, {"b", TypeI64}}, + pkey: []string{"b"}, + } + testParseStmt(t, s, stmt) + + s = "insert into t values (1, 'hi');" + stmt = &StmtInsert{ + table: "t", + value: []Cell{{Type: TypeI64, I64: 1}, {Type: TypeStr, Str: []byte("hi")}}, + } + testParseStmt(t, s, stmt) + + s = "update t set a = 1, b = 2 where c = 3 and d = 4;" + stmt = &StmtUpdate{ + table: "t", + value: []NamedCell{{"a", Cell{Type: TypeI64, I64: 1}}, {"b", Cell{Type: TypeI64, I64: 2}}}, + keys: []NamedCell{{"c", Cell{Type: TypeI64, I64: 3}}, {"d", Cell{Type: TypeI64, I64: 4}}}, + } + testParseStmt(t, s, stmt) + + s = "delete from t where c = 3 and d = 4;" + stmt = &StmtDelete{ + table: "t", + keys: []NamedCell{{"c", Cell{Type: TypeI64, I64: 3}}, {"d", Cell{Type: TypeI64, I64: 4}}}, + } + testParseStmt(t, s, stmt) +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/table.go b/0402/table.go new file mode 100644 index 0000000..8e5420c --- /dev/null +++ b/0402/table.go @@ -0,0 +1,267 @@ +package db0402 + +import ( + "encoding/json" + "errors" + "slices" +) + +type DB struct { + KV KV + tables map[string]Schema +} + +func (db *DB) Open() error { + db.tables = map[string]Schema{} + return db.KV.Open() +} + +func (db *DB) Close() error { return db.KV.Close() } + +func (db *DB) Select(schema *Schema, row Row) (ok bool, err error) { + key := row.EncodeKey(schema) + val, ok, err := db.KV.Get(key) + if err != nil || !ok { + return ok, err + } + if err = row.DecodeVal(schema, val); err != nil { + return false, err + } + return true, nil +} + +func (db *DB) Insert(schema *Schema, row Row) (updated bool, err error) { + key := row.EncodeKey(schema) + val := row.EncodeVal(schema) + return db.KV.SetEx(key, val, ModeInsert) +} + +func (db *DB) Upsert(schema *Schema, row Row) (updated bool, err error) { + key := row.EncodeKey(schema) + val := row.EncodeVal(schema) + return db.KV.SetEx(key, val, ModeUpsert) +} + +func (db *DB) Update(schema *Schema, row Row) (updated bool, err error) { + key := row.EncodeKey(schema) + val := row.EncodeVal(schema) + return db.KV.SetEx(key, val, ModeUpdate) +} + +func (db *DB) Delete(schema *Schema, row Row) (deleted bool, err error) { + key := row.EncodeKey(schema) + return db.KV.Del(key) +} + +type SQLResult struct { + Updated int + Header []string + Values []Row +} + +func (db *DB) ExecStmt(stmt interface{}) (r SQLResult, err error) { + switch ptr := stmt.(type) { + case *StmtCreatTable: + err = db.execCreateTable(ptr) + case *StmtSelect: + r.Header = ptr.cols + r.Values, err = db.execSelect(ptr) + case *StmtInsert: + r.Updated, err = db.execInsert(ptr) + case *StmtUpdate: + r.Updated, err = db.execUpdate(ptr) + case *StmtDelete: + r.Updated, err = db.execDelete(ptr) + default: + panic("unreachable") + } + return +} + +func (db *DB) execCreateTable(stmt *StmtCreatTable) (err error) { + if _, err := db.GetSchema(stmt.table); err == nil { + return errors.New("duplicate table name") + } + + schema := Schema{ + Table: stmt.table, + Cols: stmt.cols, + } + if schema.PKey, err = lookupColumns(stmt.cols, stmt.pkey); err != nil { + return err + } + + val, err := json.Marshal(schema) + check(err == nil) + if _, err = db.KV.Set([]byte("@schema_"+stmt.table), val); err != nil { + return err + } + + db.tables[schema.Table] = schema + return nil +} + +func (db *DB) GetSchema(table string) (Schema, error) { + schema, ok := db.tables[table] + if !ok { + val, ok, err := db.KV.Get([]byte("@schema_" + table)) + if err == nil && ok { + err = json.Unmarshal(val, &schema) + } + if err != nil { + return Schema{}, err + } + if !ok { + return Schema{}, errors.New("table is not found") + } + db.tables[table] = schema + } + return schema, nil +} + +func lookupColumns(cols []Column, names []string) (indices []int, err error) { + for _, name := range names { + idx := slices.IndexFunc(cols, func(col Column) bool { + return col.Name == name + }) + if idx < 0 { + return nil, errors.New("column is not found") + } + indices = append(indices, idx) + } + return +} + +func makePKey(schema *Schema, pkey []NamedCell) (Row, error) { + if len(schema.PKey) != len(pkey) { + return nil, errors.New("not primary key") + } + row := schema.NewRow() + for _, idx1 := range schema.PKey { + col := schema.Cols[idx1] + idx2 := slices.IndexFunc(pkey, func(expr NamedCell) bool { + return expr.column == col.Name && expr.value.Type == col.Type + }) + if idx2 < 0 { + return nil, errors.New("not primary key") + } + row[idx1] = pkey[idx2].value + } + return row, nil +} + +func subsetRow(row Row, indices []int) (out Row) { + for _, idx := range indices { + out = append(out, row[idx]) + } + return +} + +func (db *DB) execSelect(stmt *StmtSelect) ([]Row, error) { + schema, err := db.GetSchema(stmt.table) + if err != nil { + return nil, err + } + indices, err := lookupColumns(schema.Cols, stmt.cols) + if err != nil { + return nil, err + } + + row, err := makePKey(&schema, stmt.keys) + if err != nil { + return nil, err + } + if ok, err := db.Select(&schema, row); err != nil || !ok { + return nil, err + } + + row = subsetRow(row, indices) + return []Row{row}, nil +} + +func (db *DB) execInsert(stmt *StmtInsert) (count int, err error) { + schema, err := db.GetSchema(stmt.table) + if err != nil { + return 0, err + } + if len(schema.Cols) != len(stmt.value) { + return 0, errors.New("schema mismatch") + } + for i := range schema.Cols { + if schema.Cols[i].Type != stmt.value[i].Type { + return 0, errors.New("schema mismatch") + } + } + + updated, err := db.Insert(&schema, stmt.value) + if err != nil { + return 0, err + } + if updated { + count++ + } + return count, nil +} + +func fillNonPKey(schema *Schema, updates []NamedCell, out Row) error { + for _, expr := range updates { + idx := slices.IndexFunc(schema.Cols, func(col Column) bool { + return col.Name == expr.column && col.Type == expr.value.Type + }) + if idx < 0 || slices.Contains(schema.PKey, idx) { + return errors.New("cannot update column") + } + out[idx] = expr.value + } + return nil +} + +func (db *DB) execUpdate(stmt *StmtUpdate) (count int, err error) { + schema, err := db.GetSchema(stmt.table) + if err != nil { + return 0, err + } + + row, err := makePKey(&schema, stmt.keys) + if err != nil { + return 0, err + } + if ok, err := db.Select(&schema, row); err != nil || !ok { + return 0, err + } + + if err = fillNonPKey(&schema, stmt.value, row); err != nil { + return 0, err + } + updated, err := db.Update(&schema, row) + if err != nil { + return 0, err + } + if updated { + count++ + } + return count, nil +} + +func (db *DB) execDelete(stmt *StmtDelete) (count int, err error) { + schema, err := db.GetSchema(stmt.table) + if err != nil { + return 0, err + } + + row, err := makePKey(&schema, stmt.keys) + if err != nil { + return 0, err + } + + updated, err := db.Delete(&schema, row) + if err != nil { + return 0, err + } + if updated { + count++ + } + return count, nil +} + +// QzBQWVJJOUhU https://trialofcode.org/ diff --git a/0402/table_test.go b/0402/table_test.go new file mode 100644 index 0000000..c885caf --- /dev/null +++ b/0402/table_test.go @@ -0,0 +1,126 @@ +package db0402 + +import ( + "os" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestTableByPKey(t *testing.T) { + db := DB{} + db.KV.log.FileName = ".test_db" + defer os.Remove(db.KV.log.FileName) + + os.Remove(db.KV.log.FileName) + err := db.Open() + assert.Nil(t, err) + defer db.Close() + + schema := &Schema{ + Table: "link", + Cols: []Column{ + {Name: "time", Type: TypeI64}, + {Name: "src", Type: TypeStr}, + {Name: "dst", Type: TypeStr}, + }, + PKey: []int{1, 2}, // (src, dst) + } + + row := Row{ + Cell{Type: TypeI64, I64: 123}, + Cell{Type: TypeStr, Str: []byte("a")}, + Cell{Type: TypeStr, Str: []byte("b")}, + } + ok, err := db.Select(schema, row) + assert.True(t, !ok && err == nil) + + updated, err := db.Insert(schema, row) + assert.True(t, updated && err == nil) + + out := Row{ + Cell{}, + Cell{Type: TypeStr, Str: []byte("a")}, + Cell{Type: TypeStr, Str: []byte("b")}, + } + ok, err = db.Select(schema, out) + assert.True(t, ok && err == nil) + assert.Equal(t, row, out) + + row[0].I64 = 456 + updated, err = db.Update(schema, row) + assert.True(t, updated && err == nil) + + ok, err = db.Select(schema, out) + assert.True(t, ok && err == nil) + assert.Equal(t, row, out) + + deleted, err := db.Delete(schema, row) + assert.True(t, deleted && err == nil) + + ok, err = db.Select(schema, row) + assert.True(t, !ok && err == nil) +} + +func parseStmt(t *testing.T, s string) interface{} { + p := NewParser(s) + stmt, err := p.parseStmt() + require.Nil(t, err) + return stmt +} + +func TestSQLByPKey(t *testing.T) { + db := DB{} + db.KV.log.FileName = ".test_db" + defer os.Remove(db.KV.log.FileName) + + os.Remove(db.KV.log.FileName) + err := db.Open() + assert.Nil(t, err) + defer db.Close() + + s := "create table link (time int64, src string, dst string, primary key (src, dst));" + _, err = db.ExecStmt(parseStmt(t, s)) + require.Nil(t, err) + + s = "insert into link values (123, 'bob', 'alice');" + r, err := db.ExecStmt(parseStmt(t, s)) + require.Nil(t, err) + require.Equal(t, 1, r.Updated) + + s = "select time from link where dst = 'alice' and src = 'bob';" + r, err = db.ExecStmt(parseStmt(t, s)) + require.Nil(t, err) + require.Equal(t, []Row{{Cell{Type: TypeI64, I64: 123}}}, r.Values) + + s = "update link set time = 456 where dst = 'alice' and src = 'bob';" + r, err = db.ExecStmt(parseStmt(t, s)) + require.Nil(t, err) + require.Equal(t, 1, r.Updated) + + s = "select time from link where dst = 'alice' and src = 'bob';" + r, err = db.ExecStmt(parseStmt(t, s)) + require.Nil(t, err) + require.Equal(t, []Row{{Cell{Type: TypeI64, I64: 456}}}, r.Values) + + // reopen + err = db.Close() + require.Nil(t, err) + db = DB{} + db.KV.log.FileName = ".test_db" + err = db.Open() + require.Nil(t, err) + + s = "delete from link where src = 'bob' and dst = 'alice';" + r, err = db.ExecStmt(parseStmt(t, s)) + require.Nil(t, err) + require.Equal(t, 1, r.Updated) + + s = "select time from link where dst = 'alice' and src = 'bob';" + r, err = db.ExecStmt(parseStmt(t, s)) + require.Nil(t, err) + require.Equal(t, 0, len(r.Values)) +} + +// QzBQWVJJOUhU https://trialofcode.org/