Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 9 additions & 46 deletions persistent/internal/driver/postgres/basic_operations.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ import (
// Insert adds one or more objects into the database in a single batch operation.
// Returns an error if the input is empty or the insert fails.
func (d *driver) Insert(ctx context.Context, objects ...model.DBObject) error {
if d.db == nil {
if d.writeDB == nil {
return errors.New(types.ErrorSessionClosed)
}
if len(objects) == 0 {
Expand Down Expand Up @@ -52,7 +52,7 @@ func (d *driver) Insert(ctx context.Context, objects ...model.DBObject) error {
}

tableName := objs[0].TableName()
if err := d.db.WithContext(ctx).Table(tableName).Create(sliceValue.Interface()).Error; err != nil {
if err := d.writeDB.WithContext(ctx).Table(tableName).Create(sliceValue.Interface()).Error; err != nil {
return err
}
}
Expand All @@ -73,7 +73,7 @@ func (d *driver) Delete(ctx context.Context, object model.DBObject, filters ...m
}

// Start building the query with the table name
db := d.db.WithContext(ctx).Table(tableName)
db := d.writeDB.WithContext(ctx).Table(tableName)
// If we have a filter, use our translator function
if len(filters) == 1 {
db, err = d.translateQuery(db, filters[0], object)
Expand Down Expand Up @@ -111,7 +111,7 @@ func (d *driver) Update(ctx context.Context, object model.DBObject, filters ...m
return errors.New(types.ErrorMultipleDBM)
}

tx := d.db.WithContext(ctx).Table(tableName)
tx := d.writeDB.WithContext(ctx).Table(tableName)

// Apply filters
if len(filters) == 1 {
Expand Down Expand Up @@ -149,7 +149,7 @@ for updating a collection of specific records with different values.
*/
func (d *driver) BulkUpdate(ctx context.Context, objects []model.DBObject, filters ...model.DBM) error {
// Basic validation
if d.db == nil {
if d.writeDB == nil {
return errors.New(types.ErrorSessionClosed)
}
if len(objects) == 0 {
Expand All @@ -160,7 +160,7 @@ func (d *driver) BulkUpdate(ctx context.Context, objects []model.DBObject, filte
}

// Start a transaction
tx := d.db.WithContext(ctx).Begin()
tx := d.writeDB.WithContext(ctx).Begin()
if tx.Error != nil {
return tx.Error
}
Expand Down Expand Up @@ -272,7 +272,7 @@ func (d *driver) UpdateAll(ctx context.Context, row model.DBObject, query, updat
}

// Start a transaction
tx := d.db.WithContext(ctx).Begin()
tx := d.writeDB.WithContext(ctx).Begin()
if tx.Error != nil {
return tx.Error
}
Expand All @@ -284,7 +284,7 @@ func (d *driver) UpdateAll(ctx context.Context, row model.DBObject, query, updat
panic(r) // re-throw panic after rollback
}
}()
db := d.db.WithContext(ctx).Table(tableName)
db := d.writeDB.WithContext(ctx).Table(tableName)

// Check if query is empty
hasFilter := false
Expand Down Expand Up @@ -343,7 +343,7 @@ func (d *driver) Upsert(ctx context.Context, row model.DBObject, query, update m
return err
}

tx := d.db.WithContext(ctx).Begin()
tx := d.writeDB.WithContext(ctx).Begin()
if tx.Error != nil {
return tx.Error
}
Expand Down Expand Up @@ -430,43 +430,6 @@ func (d *driver) fetchUpdatedRow(tx *gorm.DB, table string, query model.DBM, row
return db.First(row).Error
}

func ensureID(originalID model.ObjectID, row model.DBObject, query model.DBM) {
if originalID != "" {
row.SetObjectID(originalID)
} else if idVal, ok := query["id"].(string); ok && idVal != "" {
row.SetObjectID(model.ObjectIDHex(idVal))
}
if row.GetObjectID() == "" {
row.SetObjectID(model.NewObjectID())
}
}

func cloneDBObject(row model.DBObject) model.DBObject {
newRow := reflect.New(reflect.TypeOf(row).Elem()).Interface().(model.DBObject)
newRow.SetObjectID(row.GetObjectID())
return newRow
}

func mergeQueryFields(row model.DBObject, query model.DBM) {
for k, v := range query {
if strings.HasPrefix(k, "_") || k == "$or" {
continue
}
setField(row, k, v) // keeps reflection logic isolated
}
}

func (d *driver) ensureID(originalID model.ObjectID, row model.DBObject, query model.DBM) {
if originalID != "" {
row.SetObjectID(originalID)
} else if qid, ok := query["id"]; ok {
if sid, ok2 := qid.(string); ok2 && sid != "" {
row.SetObjectID(model.ObjectIDHex(sid))

}
}
}

// Helper function to set a field in a struct using reflection
func setField(obj interface{}, name string, value interface{}) {
structValue := reflect.ValueOf(obj)
Expand Down
70 changes: 0 additions & 70 deletions persistent/internal/driver/postgres/basic_operations_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -549,73 +549,3 @@ func TestUpsert(t *testing.T) {
assert.Equal(t, 50, resultItem.Value)
})
}

func TestCloneDBObject(t *testing.T) {
original := &TestObject{
Name: "Original",
Value: 42,
CreatedAt: time.Now(),
}
original.SetObjectID(model.NewObjectID())

clone := cloneDBObject(original)

// Ensure it's a different pointer
assert.NotSame(t, original, clone)

// Ensure it has the same ID
assert.Equal(t, original.GetObjectID(), clone.GetObjectID())

// Ensure other fields are zeroed (because cloneDBObject only copies ID)
cloneObj, ok := clone.(*TestObject)
require.True(t, ok)

assert.Equal(t, "", cloneObj.Name)
assert.Equal(t, 0, cloneObj.Value)
assert.WithinDuration(t, time.Time{}, cloneObj.CreatedAt, time.Second)
}

func TestMergeQueryFields(t *testing.T) {
obj := &TestObject{
Name: "Initial",
Value: 10,
CreatedAt: time.Now(),
}

query := model.DBM{
"name": "Updated Name",
"value": 42,
"_limit": 100, // should be ignored
"$or": []model.DBM{}, // should be ignored
"extra_field": "Extra", // will only work if TestObject has this field; otherwise ignored
}

mergeQueryFields(obj, query)

// Check that allowed fields were updated
assert.Equal(t, "Updated Name", obj.Name)
assert.Equal(t, 42, obj.Value)

}

func TestEnsureID(t *testing.T) {
driver, _ := setupTest(t)
defer teardownTest(t, driver)

// Case 1: originalID is provided → should preserve it
obj1 := &TestObject{}
origID := model.NewObjectID()
driver.ensureID(origID, obj1, model.DBM{})
assert.Equal(t, origID, obj1.GetObjectID())

// Case 2: originalID is empty, but query contains "id"
obj2 := &TestObject{}
queryID := model.NewObjectID()
driver.ensureID("", obj2, model.DBM{"id": queryID.Hex()})
assert.Equal(t, queryID, obj2.GetObjectID())

// Case 3: neither originalID nor query["id"] → ID should remain empty
obj3 := &TestObject{}
driver.ensureID("", obj3, model.DBM{})
assert.Equal(t, model.ObjectID(""), obj3.GetObjectID())
}
2 changes: 1 addition & 1 deletion persistent/internal/driver/postgres/driver.go
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@ func NewPostgresDriver(opts *types.ClientOpts) (*driver, error) {

func (d *driver) validateDBAndTable(object model.DBObject) (string, error) {
// Check if the database connection is valid
if d.db == nil {
if d.writeDB == nil || d.readDB == nil {
return "", errors.New(types.ErrorSessionClosed)
}

Expand Down
2 changes: 1 addition & 1 deletion persistent/internal/driver/postgres/driver_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -122,7 +122,7 @@ func TestValidateDBAndTable(t *testing.T) {
// Create a driver with a valid connection
driver, _ := setupTest(t)

// Close the connection to simulate a nil db
// Close the connection to simulate a nil writeDB
driver.Close()

// Create a mock object with a valid table name
Expand Down
32 changes: 17 additions & 15 deletions persistent/internal/driver/postgres/indexes.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,9 @@
"regexp"
"strings"
)

Check notice on line 13 in persistent/internal/driver/postgres/indexes.go

View check run for this annotation

probelabs / Visor: performance

performance Issue

The compilation of the regular expression for `sanitizeIdentifier` has been moved to a global variable. This is a positive micro-optimization that avoids the performance overhead of recompiling the regex on every function call.
Raw output
This change is a good performance practice and should be maintained.
var sanitizerRegex = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`)

type IndexRow struct {
IndexName string
ColumnName string
Expand All @@ -20,13 +22,6 @@
Comment *string // Using pointer for nullable string
}

func sanitizeIdentifier(s string) (string, error) {
if matched := regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`).MatchString(s); !matched {
return "", fmt.Errorf("invalid identifier: %s", s)
}
return pq.QuoteIdentifier(s), nil // use pq or pgx quoting
}

// CreateIndex creates a database index on the specified table for the given fields.
// Returns an error if index creation fails.
func (d *driver) CreateIndex(ctx context.Context, row model.DBObject, index model.Index) error {
Expand Down Expand Up @@ -135,7 +130,7 @@
createSQL = strings.Replace(createSQL, "CREATE INDEX", "CREATE INDEX CONCURRENTLY", 1)
}

if err := d.db.WithContext(ctx).Exec(createSQL).Error; err != nil {
if err := d.writeDB.WithContext(ctx).Exec(createSQL).Error; err != nil {
return fmt.Errorf("failed to create index: %w", err)
}

Expand All @@ -150,7 +145,7 @@
PRIMARY KEY (table_name, index_name)
)
`
if err := d.db.WithContext(ctx).Exec(metadataSQL).Error; err != nil {
if err := d.writeDB.WithContext(ctx).Exec(metadataSQL).Error; err != nil {
return fmt.Errorf("failed to create metadata table: %w", err)
}

Expand All @@ -160,7 +155,7 @@
ON CONFLICT (table_name, index_name) DO UPDATE
SET is_ttl = TRUE, ttl_seconds = EXCLUDED.ttl_seconds
`
if err := d.db.WithContext(ctx).Exec(ttlSQL, tableName, indexName, index.TTL).Error; err != nil {
if err := d.writeDB.WithContext(ctx).Exec(ttlSQL, tableName, indexName, index.TTL).Error; err != nil {
return fmt.Errorf("failed to store TTL metadata: %w", err)
}
}
Expand Down Expand Up @@ -232,7 +227,7 @@
`

var rows []IndexRow
err = d.db.WithContext(ctx).Raw(query, tableName).Scan(&rows).Error
err = d.readDB.WithContext(ctx).Raw(query, tableName).Scan(&rows).Error
if err != nil {
return nil, fmt.Errorf("failed to query indexes: %w", err)
}
Expand Down Expand Up @@ -274,7 +269,7 @@

func (d *driver) tableExists(ctx context.Context, tableName string) (bool, error) {
var exists bool
err := d.db.WithContext(ctx).Raw("SELECT EXISTS (SELECT FROM information_schema.tables WHERE table_name = ?)", tableName).Scan(&exists).Error
err := d.readDB.WithContext(ctx).Raw("SELECT EXISTS (SELECT FROM information_schema.tables WHERE table_name = ?)", tableName).Scan(&exists).Error
if err != nil {
return false, fmt.Errorf("failed to check if table exists: %w", err)
}
Expand Down Expand Up @@ -316,7 +311,7 @@

var indexNames []string
// Execute the query
err = d.db.WithContext(ctx).Raw(query, tableName).Scan(&indexNames).Error
err = d.writeDB.WithContext(ctx).Raw(query, tableName).Scan(&indexNames).Error
if err != nil {
return fmt.Errorf("failed to query indexes: %w", err)
}
Expand All @@ -326,7 +321,7 @@
}

// Start a transaction
tx := d.db.WithContext(ctx).Begin()
tx := d.writeDB.WithContext(ctx).Begin()
if tx.Error != nil {
return fmt.Errorf("failed to begin transaction: %w", err)
}
Expand Down Expand Up @@ -365,10 +360,17 @@
`

var exists bool
err := d.db.WithContext(ctx).Raw(query, tableName, indexName).Scan(&exists).Error
err := d.readDB.WithContext(ctx).Raw(query, tableName, indexName).Scan(&exists).Error
if err != nil {
return false, err
}

return exists, nil
}

func sanitizeIdentifier(s string) (string, error) {
if matched := sanitizerRegex.MatchString(s); !matched {
return "", fmt.Errorf("invalid identifier: %s", s)
}

Check notice on line 374 in persistent/internal/driver/postgres/indexes.go

View check run for this annotation

probelabs / Visor: style

style Issue

The helper function `sanitizeIdentifier` is defined at the end of the file, far from its usage point in `CreateIndex`. While functionally correct, placing helper functions closer to where they are called improves code organization and readability.
Raw output
Consider moving the `sanitizeIdentifier` function definition to be either just before or just after the `CreateIndex` function to improve code flow and discoverability for future maintainers.
return pq.QuoteIdentifier(s), nil // use pq or pgx quoting
}
2 changes: 1 addition & 1 deletion persistent/internal/driver/postgres/indexes_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ func TestCreateIndex(t *testing.T) {

// Helper function to clean up test data
cleanupTestData := func(tableName string) {
err := driver.db.WithContext(ctx).Exec(fmt.Sprintf("DROP TABLE IF EXISTS %s", tableName)).Error
err := driver.writeDB.WithContext(ctx).Exec(fmt.Sprintf("DROP TABLE IF EXISTS %s", tableName)).Error
if err != nil {
t.Logf("Error cleaning up test data: %v", err)
}
Expand Down
Loading
Loading