Sfoglia il codice sorgente

optimize primary key access

Danilo Fragoso 1 settimana fa
parent
commit
f9238c2049
5 ha cambiato i file con 319 aggiunte e 29 eliminazioni
  1. 69 12
      pkg/executor/executor.go
  2. 51 8
      pkg/storage/schema.go
  3. 52 0
      pkg/storage/schema_test.go
  4. 96 9
      pkg/storage/table.go
  5. 51 0
      pkg/storage/table_test.go

+ 69 - 12
pkg/executor/executor.go

@@ -401,20 +401,32 @@ func (e *Executor) executeSelect(stmt *parser.SelectStmt) (*Result, error) {
 		// Check if we can use an index
 		colName, colValue, isEquality := e.extractIndexableCondition(stmt.Where)
 		if isEquality {
-			// Look for an index on this column
-			indexes, _ := e.schema.ListTableIndexes(tableName)
-			for _, idx := range indexes {
-				if len(idx.Columns) == 1 && strings.EqualFold(idx.Columns[0].Name, colName) {
-					// Use this index
-					rows, err = e.table.SelectByIndex(tableName, idx.Name, colValue)
-					if err == nil {
-						usedIndex = true
-						// Normalize rows from index
-						for i := range rows {
-							normalizeRowBySchema(rows[i], schema)
+			if strings.EqualFold(schema.PrimaryKey, colName) {
+				row, getErr := e.table.GetByPK(tableName, fmt.Sprintf("%v", colValue))
+				if getErr == nil {
+					normalizeRowBySchema(row, schema)
+					rows = []storage.Row{row}
+					usedIndex = true
+				} else if getErr == storage.ErrKeyNotFound {
+					rows = []storage.Row{}
+					usedIndex = true
+				} else {
+					return nil, getErr
+				}
+			} else {
+				// Look for an index on this column
+				indexes, _ := e.schema.ListTableIndexes(tableName)
+				for _, idx := range indexes {
+					if len(idx.Columns) == 1 && strings.EqualFold(idx.Columns[0].Name, colName) {
+						rows, err = e.table.SelectByIndex(tableName, idx.Name, colValue)
+						if err == nil {
+							usedIndex = true
+							for i := range rows {
+								normalizeRowBySchema(rows[i], schema)
+							}
 						}
+						break
 					}
-					break
 				}
 			}
 		}
@@ -2324,6 +2336,32 @@ func (e *Executor) executeUpdate(stmt *parser.UpdateStmt) (*Result, error) {
 		}
 		return updates, nil
 	}
+	if stmt.Where != nil {
+		column, value, equality := e.extractIndexableCondition(stmt.Where)
+		updatesPrimaryKey := false
+		for _, assignment := range stmt.Set {
+			if strings.EqualFold(assignment.Column, schema.PrimaryKey) {
+				updatesPrimaryKey = true
+				break
+			}
+		}
+		if equality && strings.EqualFold(column, schema.PrimaryKey) && !updatesPrimaryKey {
+			oldRow, updated, err := e.table.UpdateByPK(tableName, fmt.Sprintf("%v", value), updateFn)
+			if err != nil {
+				return nil, err
+			}
+			count := 0
+			if updated {
+				count = 1
+				if e.inTransaction {
+					e.txLog = append(e.txLog, txLogEntry{operation: "UPDATE", table: tableName, key: fmt.Sprintf("%v", oldRow[schema.PrimaryKey]), oldData: oldRow})
+				}
+			}
+			result := NewResult("UPDATE")
+			result.SetRowCount(count)
+			return result, nil
+		}
+	}
 
 	var oldRows []storage.Row
 	if e.inTransaction {
@@ -2364,6 +2402,25 @@ func (e *Executor) executeDelete(stmt *parser.DeleteStmt) (*Result, error) {
 			return toBool(val)
 		}
 	}
+	if stmt.Where != nil {
+		column, value, equality := e.extractIndexableCondition(stmt.Where)
+		if equality && strings.EqualFold(column, schema.PrimaryKey) {
+			oldRow, deleted, err := e.table.DeleteByPK(tableName, fmt.Sprintf("%v", value))
+			if err != nil {
+				return nil, err
+			}
+			count := 0
+			if deleted {
+				count = 1
+				if e.inTransaction {
+					e.txLog = append(e.txLog, txLogEntry{operation: "DELETE", table: tableName, key: fmt.Sprintf("%v", oldRow[schema.PrimaryKey]), oldData: oldRow})
+				}
+			}
+			result := NewResult("DELETE")
+			result.SetRowCount(count)
+			return result, nil
+		}
+	}
 
 	var oldRows []storage.Row
 	if e.inTransaction {

+ 51 - 8
pkg/storage/schema.go

@@ -50,6 +50,9 @@ type SchemaManager struct {
 	pool             *KVPool
 	database         string
 	cache            map[string]*Schema
+	indexCache       map[string]*Index
+	indexListCache   []string
+	indexListCached  bool
 	rowIDInitialized map[string]bool
 	version          uint64
 	mu               sync.RWMutex
@@ -77,6 +80,7 @@ func NewSchemaManager(pool *KVPool, database string) *SchemaManager {
 		pool:             pool,
 		database:         database,
 		cache:            make(map[string]*Schema),
+		indexCache:       make(map[string]*Index),
 		rowIDInitialized: make(map[string]bool),
 		tableLocks:       make(map[string]*sync.RWMutex),
 	}
@@ -303,6 +307,15 @@ func cloneSchema(schema *Schema) *Schema {
 	return &cloned
 }
 
+func cloneIndex(index *Index) *Index {
+	if index == nil {
+		return nil
+	}
+	cloned := *index
+	cloned.Columns = append([]IndexColumn(nil), index.Columns...)
+	return &cloned
+}
+
 // TableExists checks if a table exists.
 func (m *SchemaManager) TableExists(name string) bool {
 	_, err := m.GetSchema(name)
@@ -694,6 +707,7 @@ func (m *SchemaManager) CreateIndex(index *Index) error {
 	if err := m.addToIndexList(index.Name); err != nil {
 		return err
 	}
+	m.indexCache[strings.ToLower(index.Name)] = cloneIndex(index)
 	m.bumpVersionLocked()
 	return nil
 }
@@ -714,6 +728,7 @@ func (m *SchemaManager) DropIndex(name string) error {
 	if err := m.removeFromIndexList(name); err != nil {
 		return err
 	}
+	delete(m.indexCache, strings.ToLower(name))
 	m.bumpVersionLocked()
 	return nil
 }
@@ -722,6 +737,9 @@ func (m *SchemaManager) DropIndex(name string) error {
 func (m *SchemaManager) IndexExists(name string) bool {
 	m.mu.RLock()
 	defer m.mu.RUnlock()
+	if _, ok := m.indexCache[strings.ToLower(name)]; ok {
+		return true
+	}
 
 	key := m.indexKey(name)
 	err := m.pool.WithClient(func(c *KVClient) error {
@@ -733,8 +751,12 @@ func (m *SchemaManager) IndexExists(name string) bool {
 
 // GetIndex retrieves an index by name.
 func (m *SchemaManager) GetIndex(name string) (*Index, error) {
-	m.mu.RLock()
-	defer m.mu.RUnlock()
+	m.mu.Lock()
+	defer m.mu.Unlock()
+	cacheKey := strings.ToLower(name)
+	if index, ok := m.indexCache[cacheKey]; ok {
+		return cloneIndex(index), nil
+	}
 
 	key := m.indexKey(name)
 	var data string
@@ -752,13 +774,17 @@ func (m *SchemaManager) GetIndex(name string) (*Index, error) {
 		return nil, fmt.Errorf("failed to parse index: %w", err)
 	}
 
-	return &index, nil
+	m.indexCache[cacheKey] = &index
+	return cloneIndex(&index), nil
 }
 
 // ListIndexes returns all index names.
 func (m *SchemaManager) ListIndexes() ([]string, error) {
-	m.mu.RLock()
-	defer m.mu.RUnlock()
+	m.mu.Lock()
+	defer m.mu.Unlock()
+	if m.indexListCached {
+		return append([]string(nil), m.indexListCache...), nil
+	}
 
 	key := m.indexListKey()
 	var data string
@@ -767,15 +793,22 @@ func (m *SchemaManager) ListIndexes() ([]string, error) {
 		data, err = c.Read(key)
 		return err
 	})
-	if err != nil {
+	if err == ErrKeyNotFound {
+		m.indexListCache = nil
+		m.indexListCached = true
 		return []string{}, nil
 	}
+	if err != nil {
+		return nil, err
+	}
 
 	var indexes []string
 	if err := json.Unmarshal([]byte(data), &indexes); err != nil {
 		return []string{}, nil
 	}
 
+	m.indexListCache = append([]string(nil), indexes...)
+	m.indexListCached = true
 	return indexes, nil
 }
 
@@ -818,9 +851,14 @@ func (m *SchemaManager) addToIndexList(name string) error {
 	indexes = append(indexes, name)
 	newData, _ := json.Marshal(indexes)
 
-	return m.pool.WithClient(func(c *KVClient) error {
+	err = m.pool.WithClient(func(c *KVClient) error {
 		return c.Write(key, string(newData))
 	})
+	if err == nil {
+		m.indexListCache = append([]string(nil), indexes...)
+		m.indexListCached = true
+	}
+	return err
 }
 
 // removeFromIndexList removes an index name from the list.
@@ -847,9 +885,14 @@ func (m *SchemaManager) removeFromIndexList(name string) error {
 	}
 
 	newData, _ := json.Marshal(newIndexes)
-	return m.pool.WithClient(func(c *KVClient) error {
+	err = m.pool.WithClient(func(c *KVClient) error {
 		return c.Write(key, string(newData))
 	})
+	if err == nil {
+		m.indexListCache = append([]string(nil), newIndexes...)
+		m.indexListCached = true
+	}
+	return err
 }
 
 // AddColumn adds a new column to a table.

+ 52 - 0
pkg/storage/schema_test.go

@@ -29,6 +29,8 @@ type testKVServer struct {
 	scanNexts    int
 	scanCloses   int
 	keyOnlyOpens int
+	gets         int
+	multiGets    int
 	closers      []net.Conn
 }
 
@@ -127,6 +129,12 @@ func (s *testKVServer) keyOnlyOpenCount() int {
 	return s.keyOnlyOpens
 }
 
+func (s *testKVServer) readStats() (gets, multiGets int) {
+	s.mu.Lock()
+	defer s.mu.Unlock()
+	return s.gets, s.multiGets
+}
+
 func (s *testKVServer) handle(conn net.Conn) {
 	defer conn.Close()
 
@@ -158,6 +166,7 @@ func (s *testKVServer) execute(opcode uint16, payload []byte, scans map[uint64]*
 			return errorBody("InvalidPayload")
 		}
 		s.mu.Lock()
+		s.gets++
 		value, found := s.data[string(key)]
 		lsn := s.lsns[string(key)]
 		s.mu.Unlock()
@@ -227,6 +236,7 @@ func (s *testKVServer) execute(opcode uint16, payload []byte, scans map[uint64]*
 		putU16(body[0:2], statusOK)
 		putU32(body[2:6], uint32(len(keys)))
 		s.mu.Lock()
+		s.multiGets++
 		for _, key := range keys {
 			value, found := s.data[string(key)]
 			if !found {
@@ -652,3 +662,45 @@ func TestIndexIsDerivedFromRowsAfterRestart(t *testing.T) {
 		t.Fatalf("expected no durable index entry writes, got %d", got)
 	}
 }
+
+func TestListTableIndexesCachesMetadata(t *testing.T) {
+	kv := newTestKVServer(t)
+	defer kv.close()
+	pool := newTestKVPool(kv, 2, 5*time.Second)
+	defer pool.Close()
+
+	schemas := NewSchemaManager(pool, "testdb")
+	if err := schemas.CreateTable(&Schema{
+		Name: "items",
+		Columns: []Column{
+			{Name: "id", Type: "INTEGER", PrimaryKey: true},
+			{Name: "kind", Type: "TEXT"},
+		},
+	}); err != nil {
+		t.Fatal(err)
+	}
+	if err := schemas.CreateIndex(&Index{
+		Name: "idx_items_kind", Table: "items", Columns: []IndexColumn{{Name: "kind"}},
+	}); err != nil {
+		t.Fatal(err)
+	}
+
+	restarted := NewSchemaManager(pool, "testdb")
+	getsBefore, _ := kv.readStats()
+	if indexes, err := restarted.ListTableIndexes("items"); err != nil || len(indexes) != 1 {
+		t.Fatalf("first list: indexes=%v err=%v", indexes, err)
+	}
+	getsAfterFirst, _ := kv.readStats()
+	if getsAfterFirst <= getsBefore {
+		t.Fatal("first index metadata lookup did not read durable metadata")
+	}
+	for i := 0; i < 10; i++ {
+		if indexes, err := restarted.ListTableIndexes("items"); err != nil || len(indexes) != 1 {
+			t.Fatalf("cached list %d: indexes=%v err=%v", i, indexes, err)
+		}
+	}
+	getsAfterCached, _ := kv.readStats()
+	if getsAfterCached != getsAfterFirst {
+		t.Fatalf("cached index metadata issued %d extra reads", getsAfterCached-getsAfterFirst)
+	}
+}

+ 96 - 9
pkg/storage/table.go

@@ -826,6 +826,57 @@ func (m *TableManager) UpdateFunc(table string, updateFn func(Row) (Row, error),
 	return count, nil
 }
 
+// UpdateByPK updates one row without scanning the table.
+func (m *TableManager) UpdateByPK(table, pk string, updateFn func(Row) (Row, error)) (Row, bool, error) {
+	tl := m.tableLock(table)
+	tl.Lock()
+	defer tl.Unlock()
+
+	schema, err := m.schema.GetSchema(table)
+	if err != nil {
+		return nil, false, err
+	}
+	row, err := m.getByPKUnlocked(table, pk)
+	if err == ErrKeyNotFound {
+		return nil, false, nil
+	}
+	if err != nil {
+		return nil, false, err
+	}
+
+	oldRow := cloneRow(row)
+	m.updateIndexesForRow(table, oldRow, false)
+	updates, err := updateFn(row)
+	if err != nil {
+		m.updateIndexesForRow(table, oldRow, true)
+		return nil, false, err
+	}
+	for name, value := range updates {
+		for _, column := range schema.Columns {
+			if strings.EqualFold(name, column.Name) {
+				row[column.Name] = value
+				break
+			}
+		}
+	}
+
+	data, err := encodeRow(row)
+	if err != nil {
+		m.updateIndexesForRow(table, oldRow, true)
+		return nil, false, err
+	}
+	err = m.pool.WithClient(func(client *KVClient) error {
+		_, err := client.Put([]byte(m.dataKey(table, pk)), data)
+		return err
+	})
+	if err != nil {
+		m.updateIndexesForRow(table, oldRow, true)
+		return nil, false, err
+	}
+	m.updateIndexesForRow(table, row, true)
+	return oldRow, true, nil
+}
+
 // Delete deletes rows matching the filter.
 func (m *TableManager) Delete(table string, filter func(Row) bool) (int, error) {
 	// Get all rows
@@ -867,6 +918,37 @@ func (m *TableManager) Delete(table string, filter func(Row) bool) (int, error)
 	return count, nil
 }
 
+// DeleteByPK deletes one row without scanning the table.
+func (m *TableManager) DeleteByPK(table, pk string) (Row, bool, error) {
+	tl := m.tableLock(table)
+	tl.Lock()
+	defer tl.Unlock()
+
+	schema, err := m.schema.GetSchema(table)
+	if err != nil {
+		return nil, false, err
+	}
+	row, err := m.getByPKUnlocked(table, pk)
+	if err == ErrKeyNotFound {
+		return nil, false, nil
+	}
+	if err != nil {
+		return nil, false, err
+	}
+
+	m.updateIndexesForRow(table, row, false)
+	err = m.pool.WithClient(func(client *KVClient) error {
+		_, err := client.Del([]byte(m.dataKey(table, pk)))
+		return err
+	})
+	if err != nil {
+		m.updateIndexesForRow(table, row, true)
+		return nil, false, err
+	}
+	m.incrCount(table, schema.CreatedAt, -1)
+	return row, true, nil
+}
+
 // GetByPK retrieves a row by primary key.
 func (m *TableManager) GetByPK(table string, pk string) (Row, error) {
 	tl := m.tableLock(table)
@@ -876,7 +958,10 @@ func (m *TableManager) GetByPK(table string, pk string) (Row, error) {
 	if !m.schema.TableExists(table) {
 		return nil, fmt.Errorf("table not found: %s", table)
 	}
+	return m.getByPKUnlocked(table, pk)
+}
 
+func (m *TableManager) getByPKUnlocked(table, pk string) (Row, error) {
 	key := m.dataKey(table, pk)
 	var value []byte
 
@@ -889,9 +974,6 @@ func (m *TableManager) GetByPK(table string, pk string) (Row, error) {
 		return nil
 	})
 	if err != nil {
-		if err == ErrKeyNotFound {
-			return nil, fmt.Errorf("row not found: %s", pk)
-		}
 		return nil, err
 	}
 
@@ -1274,14 +1356,19 @@ func (m *TableManager) SelectByIndex(table, indexName string, colValue interface
 
 	rows := make([]Row, 0, len(primaryKeys))
 	err = m.pool.WithClient(func(client *KVClient) error {
-		for _, primaryKey := range primaryKeys {
-			result, err := client.Get([]byte(m.dataKey(table, primaryKey)))
-			if err == ErrKeyNotFound {
+		keys := make([][]byte, len(primaryKeys))
+		for i, primaryKey := range primaryKeys {
+			keys[i] = []byte(m.dataKey(table, primaryKey))
+		}
+		results, err := client.MultiGet(keys)
+		if err != nil {
+			return err
+		}
+		for i, result := range results {
+			if !result.Found {
+				primaryKey := primaryKeys[i]
 				return fmt.Errorf("index %s references missing primary key %s", indexName, primaryKey)
 			}
-			if err != nil {
-				return err
-			}
 			row, err := decodeRow(result.Value)
 			if err != nil {
 				return err

+ 51 - 0
pkg/storage/table_test.go

@@ -257,13 +257,64 @@ func TestSelectByIndexUsesPointReadsAfterBuild(t *testing.T) {
 		t.Fatalf("initial indexed select: len=%d err=%v", len(rows), err)
 	}
 	opensBefore, _, _ := kv.scanStats()
+	_, multiGetsBefore := kv.readStats()
 	if rows, err := tables.SelectByIndex("items", "idx_kind", "k1"); err != nil || len(rows) != 10 {
 		t.Fatalf("cached indexed select: len=%d err=%v", len(rows), err)
 	}
 	opensAfter, _, _ := kv.scanStats()
+	_, multiGetsAfter := kv.readStats()
 	if opensAfter != opensBefore {
 		t.Fatalf("indexed select opened %d table scans after index build", opensAfter-opensBefore)
 	}
+	if multiGetsAfter-multiGetsBefore != 1 {
+		t.Fatalf("indexed select issued %d multi-get requests, want 1", multiGetsAfter-multiGetsBefore)
+	}
+}
+
+func TestPrimaryKeyMutationsDoNotScan(t *testing.T) {
+	kv := newTestKVServer(t)
+	defer kv.close()
+	pool := newTestKVPool(kv, 4, 5*time.Second)
+	defer pool.Close()
+	schemas := NewSchemaManager(pool, "testdb")
+	tables := NewTableManager(pool, schemas, "testdb")
+	if err := schemas.CreateTable(&Schema{
+		Name: "items",
+		Columns: []Column{
+			{Name: "id", Type: "TEXT", PrimaryKey: true},
+			{Name: "value", Type: "INTEGER"},
+		},
+	}); err != nil {
+		t.Fatal(err)
+	}
+	for i := 0; i < 20; i++ {
+		if err := tables.Insert("items", Row{"id": fmt.Sprintf("item-%d", i), "value": int64(i)}); err != nil {
+			t.Fatal(err)
+		}
+	}
+
+	opensBefore, _, _ := kv.scanStats()
+	oldRow, updated, err := tables.UpdateByPK("items", "item-10", func(Row) (Row, error) {
+		return Row{"value": int64(99)}, nil
+	})
+	if err != nil || !updated || oldRow["value"] != int64(10) {
+		t.Fatalf("point update: updated=%v old=%v err=%v", updated, oldRow, err)
+	}
+	deletedRow, deleted, err := tables.DeleteByPK("items", "item-11")
+	if err != nil || !deleted || deletedRow["value"] != int64(11) {
+		t.Fatalf("point delete: deleted=%v old=%v err=%v", deleted, deletedRow, err)
+	}
+	opensAfter, _, _ := kv.scanStats()
+	if opensAfter != opensBefore {
+		t.Fatalf("primary-key mutations opened %d scans", opensAfter-opensBefore)
+	}
+	row, err := tables.GetByPK("items", "item-10")
+	if err != nil || row["value"] != int64(99) {
+		t.Fatalf("updated row=%v err=%v", row, err)
+	}
+	if _, err := tables.GetByPK("items", "item-11"); err != ErrKeyNotFound {
+		t.Fatalf("deleted row error=%v, want ErrKeyNotFound", err)
+	}
 }
 
 func TestCountFastResetsAfterDirectDropAndRecreate(t *testing.T) {