diff --git a/pkg/db/v1beta1/mysql/mysql.go b/pkg/db/v1beta1/mysql/mysql.go index 4c25a188bed..a010db69e50 100644 --- a/pkg/db/v1beta1/mysql/mysql.go +++ b/pkg/db/v1beta1/mysql/mysql.go @@ -65,6 +65,9 @@ func NewDBInterface(connectTimeout time.Duration) (common.KatibDBInterface, erro } func (d *dbConn) RegisterObservationLog(trialName string, observationLog *v1beta1.ObservationLog) error { + if observationLog == nil { + return nil + } sqlQuery := "INSERT INTO observation_logs (trial_name, time, metric_name, value) VALUES " values := []interface{}{} @@ -81,6 +84,10 @@ func (d *dbConn) RegisterObservationLog(trialName string, observationLog *v1beta sqlQuery += "(?, ?, ?, ?)," values = append(values, trialName, sqlTimeStr, mlog.Metric.Name, mlog.Metric.Value) } + if len(values) == 0 { + // No valid metric logs to insert, skip Prepare/Exec. + return nil + } sqlQuery = sqlQuery[0 : len(sqlQuery)-1] // Prepare the statement diff --git a/pkg/db/v1beta1/mysql/mysql_test.go b/pkg/db/v1beta1/mysql/mysql_test.go index f22df6cb896..1bbd6c0e4a6 100644 --- a/pkg/db/v1beta1/mysql/mysql_test.go +++ b/pkg/db/v1beta1/mysql/mysql_test.go @@ -88,6 +88,43 @@ func TestRegisterObservationLog(t *testing.T) { } +func TestRegisterObservationLogNoValidEntries(t *testing.T) { + obsLog := &api_pb.ObservationLog{ + MetricLogs: []*api_pb.MetricLog{ + { + TimeStamp: "", + Metric: &api_pb.Metric{ + Name: "f1_score", + Value: "88.95", + }, + }, + }, + } + + err := dbInterface.RegisterObservationLog("test1_trial1", obsLog) + if err != nil { + t.Errorf("RegisterObservationLog failed: %v", err) + } +} + +func TestRegisterObservationLogEmptyMetricLogs(t *testing.T) { + obsLog := &api_pb.ObservationLog{ + MetricLogs: []*api_pb.MetricLog{}, + } + + err := dbInterface.RegisterObservationLog("test1_trial1", obsLog) + if err != nil { + t.Errorf("RegisterObservationLog failed: %v", err) + } +} + +func TestRegisterObservationLogNilObservationLog(t *testing.T) { + err := dbInterface.RegisterObservationLog("test1_trial1", nil) + if err != nil { + t.Errorf("RegisterObservationLog failed: %v", err) + } +} + func TestGetObservationLog(t *testing.T) { mock.ExpectQuery("SELECT").WillReturnRows( sqlmock.NewRows([]string{"time", "metric_name", "value"}).AddRow( diff --git a/pkg/db/v1beta1/postgres/postgres.go b/pkg/db/v1beta1/postgres/postgres.go index 808af5ad93c..508d04c230a 100644 --- a/pkg/db/v1beta1/postgres/postgres.go +++ b/pkg/db/v1beta1/postgres/postgres.go @@ -67,6 +67,9 @@ func NewDBInterface(connectTimeout time.Duration) (common.KatibDBInterface, erro } func (d *dbConn) RegisterObservationLog(trialName string, observationLog *v1beta1.ObservationLog) error { + if observationLog == nil { + return nil + } statement := "INSERT INTO observation_logs (trial_name, time, metric_name, value) VALUES " values := []interface{}{} @@ -88,6 +91,11 @@ func (d *dbConn) RegisterObservationLog(trialName string, observationLog *v1beta index_of_qparam += 4 } + if len(values) == 0 { + // No valid metric logs to insert, skip Prepare/Exec. + return nil + } + statement = statement[:len(statement)-1] // Prepare the statement diff --git a/pkg/db/v1beta1/postgres/postgres_test.go b/pkg/db/v1beta1/postgres/postgres_test.go index f478107aab2..a2083ee3beb 100644 --- a/pkg/db/v1beta1/postgres/postgres_test.go +++ b/pkg/db/v1beta1/postgres/postgres_test.go @@ -90,6 +90,43 @@ func TestRegisterObservationLog(t *testing.T) { } +func TestRegisterObservationLogNoValidEntries(t *testing.T) { + obsLog := &api_pb.ObservationLog{ + MetricLogs: []*api_pb.MetricLog{ + { + TimeStamp: "", + Metric: &api_pb.Metric{ + Name: "f1_score", + Value: "88.95", + }, + }, + }, + } + + err := dbInterface.RegisterObservationLog("test1_trial1", obsLog) + if err != nil { + t.Errorf("RegisterObservationLog failed: %v", err) + } +} + +func TestRegisterObservationLogEmptyMetricLogs(t *testing.T) { + obsLog := &api_pb.ObservationLog{ + MetricLogs: []*api_pb.MetricLog{}, + } + + err := dbInterface.RegisterObservationLog("test1_trial1", obsLog) + if err != nil { + t.Errorf("RegisterObservationLog failed: %v", err) + } +} + +func TestRegisterObservationLogNilObservationLog(t *testing.T) { + err := dbInterface.RegisterObservationLog("test1_trial1", nil) + if err != nil { + t.Errorf("RegisterObservationLog failed: %v", err) + } +} + func TestGetObservationLog(t *testing.T) { mock.ExpectQuery("SELECT").WillReturnRows( sqlmock.NewRows([]string{"time", "metric_name", "value"}).AddRow(