From 34b9a569ae4dc9007939d76a432514ad943ba228 Mon Sep 17 00:00:00 2001 From: douxu Date: Tue, 21 Jul 2026 16:13:23 +0800 Subject: [PATCH] feat(data-object)!: enhance measurement mode update workflow - support numeric manual and automatic measurement modes - allow omitted fields and optional manual data values - start or stop dataRT writes when switching measurement modes - atomically replace manual measurement values in Redis - record bounded operations with Unix millisecond timestamps - propagate client tokens through request contexts - return errors instead of panicking when Redis initialization fails - optimize measurement queries and add regression tests --- constants/context.go | 6 + constants/data-object.go | 7 + database/query_component_measurement.go | 61 +++-- database/query_component_measurement_test.go | 42 +++ diagram/context.go | 20 ++ diagram/context_test.go | 38 +++ diagram/redis_hash.go | 9 +- diagram/redis_set.go | 9 +- diagram/redis_zset.go | 29 ++- handler/component_attribute_query.go | 16 +- handler/component_attribute_update.go | 9 +- handler/data_object_attribute_update.go | 177 ++++++++++--- handler/data_object_attribute_update_test.go | 253 +++++++++++++++++-- handler/diagram_node_link.go | 25 +- handler/measurement_load.go | 7 +- handler/mesurement_link.go | 14 +- middleware/token.go | 12 +- middleware/token_test.go | 28 ++ 18 files changed, 654 insertions(+), 108 deletions(-) create mode 100644 diagram/context.go create mode 100644 diagram/context_test.go create mode 100644 middleware/token_test.go diff --git a/constants/context.go b/constants/context.go index dcac3c3..cf3ff99 100644 --- a/constants/context.go +++ b/constants/context.go @@ -1,7 +1,13 @@ // Package constants define constant variable package constants +// ClientTokenContextName is the Gin key used for the configured client token. +const ClientTokenContextName = "client_token" + type contextKey string // MeasurementUUIDKey define measurement uuid key into context const MeasurementUUIDKey contextKey = "measurement_uuid" + +// CtxKeyClientToken is the typed standard-library context key for client token propagation. +const CtxKeyClientToken contextKey = ClientTokenContextName diff --git a/constants/data-object.go b/constants/data-object.go index 0c3212c..f57158e 100644 --- a/constants/data-object.go +++ b/constants/data-object.go @@ -10,3 +10,10 @@ const ( // DataObjectTypeMeasurement represents a component measurement. DataObjectTypeMeasurement DataObjectType = "measurement" ) + +const ( + // MeasurementModeManual indicates that manual value entry is enabled. + MeasurementModeManual int16 = 0 + // MeasurementModeAutomatic indicates that the measurement runs automatically. + MeasurementModeAutomatic int16 = 1 +) diff --git a/database/query_component_measurement.go b/database/query_component_measurement.go index 06562c4..c4a874f 100644 --- a/database/query_component_measurement.go +++ b/database/query_component_measurement.go @@ -9,6 +9,7 @@ import ( "time" "modelRT/common" + "modelRT/constants" "modelRT/orm" "modelRT/sql" @@ -17,21 +18,38 @@ import ( "gorm.io/gorm/clause" ) -const measurementOperationsLimit = 500 +const ( + measurementOperationsLimit = 500 + measurementOperationAppendSQL = "(array_append(operations, ?::jsonb))[GREATEST(cardinality(operations) - ? + 2, 1):]" +) -// QueryMeasurementByID return the result of query circuit diagram component measurement info by id from postgresDB -func QueryMeasurementByID(ctx context.Context, tx *gorm.DB, id int64) (orm.Measurement, error) { +// QueryMeasurementByID returns a measurement by primary key without acquiring +// a row lock. Call QueryMeasurementByIDForUpdate for write workflows. +func QueryMeasurementByID(ctx context.Context, db *gorm.DB, id int64) (orm.Measurement, error) { var measurement orm.Measurement - // ctx超时判断 cancelCtx, cancel := context.WithTimeout(ctx, 5*time.Second) defer cancel() - result := tx.WithContext(cancelCtx). + result := db.WithContext(cancelCtx). Where(sql.MeasurementIDWhere, id). - Clauses(clause.Locking{Strength: "UPDATE"}). - First(&measurement) + Take(&measurement) if result.Error != nil { - return orm.Measurement{}, result.Error + return orm.Measurement{}, fmt.Errorf("query measurement %d: %w", id, result.Error) + } + return measurement, nil +} + +// QueryMeasurementByIDForUpdate locks a measurement row and loads only the +// fields required by the data-object update workflow. +func QueryMeasurementByIDForUpdate(ctx context.Context, tx *gorm.DB, id int64) (orm.Measurement, error) { + var measurement orm.Measurement + result := tx.WithContext(ctx). + Select("id", "mode", "data_source"). + Where(sql.MeasurementIDWhere, id). + Clauses(clause.Locking{Strength: "UPDATE"}). + Take(&measurement) + if result.Error != nil { + return orm.Measurement{}, fmt.Errorf("lock measurement %d: %w", id, result.Error) } return measurement, nil } @@ -48,9 +66,9 @@ func QueryMeasurementByToken(ctx context.Context, tx *gorm.DB, token string) (or // UpdateMeasurementMode stores the data-object mode representation in the // measurement row: false is manual mode (0), true is automatic mode (1). func UpdateMeasurementMode(ctx context.Context, db *gorm.DB, measurementID int64, automatic bool) error { - mode := int16(0) + mode := constants.MeasurementModeManual if automatic { - mode = 1 + mode = constants.MeasurementModeAutomatic } result := db.WithContext(ctx). @@ -68,14 +86,13 @@ func UpdateMeasurementMode(ctx context.Context, db *gorm.DB, measurementID int64 // UpdateMeasurementModeWithOperation changes mode and appends its audit entry // atomically. The operations array retains only its newest 500 entries. -func UpdateMeasurementModeWithOperation(ctx context.Context, db *gorm.DB, measurementID int64, automatic bool, timestamp time.Time) error { - mode := int16(0) - if automatic { - mode = 1 +func UpdateMeasurementModeWithOperation(ctx context.Context, db *gorm.DB, measurementID int64, mode int16, timestamp time.Time) error { + if mode != constants.MeasurementModeManual && mode != constants.MeasurementModeAutomatic { + return fmt.Errorf("measurement mode must be 0 or 1, got %d", mode) } operation := orm.JSONMap{ "command": mode, - "timestamp": timestamp, + "timestamp": timestamp.UnixMilli(), } return updateMeasurementWithOperation(ctx, db, measurementID, map[string]any{"mode": mode}, operation) } @@ -86,7 +103,7 @@ func AppendMeasurementValueOperation(ctx context.Context, db *gorm.DB, measureme operation := orm.JSONMap{ "transaction": transaction, "value": value, - "timestamp": timestamp, + "timestamp": timestamp.UnixMilli(), } return updateMeasurementWithOperation(ctx, db, measurementID, nil, operation) } @@ -97,13 +114,11 @@ func updateMeasurementWithOperation(ctx context.Context, db *gorm.DB, measuremen return fmt.Errorf("encode measurement %d operation: %w", measurementID, err) } - operationExpression := gorm.Expr(`( - CASE - WHEN cardinality(operations) >= ? - THEN operations[(cardinality(operations) - ? + 2):cardinality(operations)] - ELSE operations - END - ) || ARRAY[?::jsonb]`, measurementOperationsLimit, measurementOperationsLimit, string(encodedOperation)) + operationExpression := gorm.Expr( + measurementOperationAppendSQL, + string(encodedOperation), + measurementOperationsLimit, + ) if updates == nil { updates = make(map[string]any, 1) } diff --git a/database/query_component_measurement_test.go b/database/query_component_measurement_test.go index 4c603e1..5e9d920 100644 --- a/database/query_component_measurement_test.go +++ b/database/query_component_measurement_test.go @@ -82,6 +82,48 @@ func TestBuildMeasurementTokenValidationQuery(t *testing.T) { } } +func TestQueryMeasurementByIDDoesNotLockRead(t *testing.T) { + sqlDB, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = sqlDB.Close() }) + db, err := gorm.Open(postgres.New(postgres.Config{Conn: sqlDB}), &gorm.Config{SkipDefaultTransaction: true}) + require.NoError(t, err) + + mock.ExpectQuery(`SELECT \* FROM "measurement" WHERE id = \$1 LIMIT \$2`). + WithArgs(int64(10), 1). + WillReturnRows(sqlmock.NewRows([]string{"id", "mode"}).AddRow(int64(10), int16(1))) + + measurement, err := QueryMeasurementByID(context.Background(), db, 10) + require.NoError(t, err) + assert.Equal(t, int64(10), measurement.ID) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestQueryMeasurementByIDForUpdateSelectsOnlyRequiredFields(t *testing.T) { + sqlDB, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = sqlDB.Close() }) + db, err := gorm.Open(postgres.New(postgres.Config{Conn: sqlDB}), &gorm.Config{SkipDefaultTransaction: true}) + require.NoError(t, err) + + mock.ExpectQuery(`SELECT "id","mode","data_source" FROM "measurement" WHERE id = \$1 LIMIT \$2 FOR UPDATE`). + WithArgs(int64(10), 1). + WillReturnRows(sqlmock.NewRows([]string{"id", "mode", "data_source"}). + AddRow(int64(10), int16(1), `{"type":1}`)) + + measurement, err := QueryMeasurementByIDForUpdate(context.Background(), db, 10) + require.NoError(t, err) + assert.Equal(t, int64(10), measurement.ID) + assert.Equal(t, int16(1), measurement.Mode) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestMeasurementOperationAppendSQLIsSingleLine(t *testing.T) { + assert.NotContains(t, measurementOperationAppendSQL, "\n") + assert.NotContains(t, measurementOperationAppendSQL, "\r") + assert.NotContains(t, measurementOperationAppendSQL, "\t") +} + func TestUpdateMeasurementMode(t *testing.T) { tests := []struct { name string diff --git a/diagram/context.go b/diagram/context.go new file mode 100644 index 0000000..ce6047b --- /dev/null +++ b/diagram/context.go @@ -0,0 +1,20 @@ +package diagram + +import ( + "context" + "fmt" + + "modelRT/common" + "modelRT/constants" +) + +func clientTokenFromContext(ctx context.Context) (string, error) { + if ctx == nil { + return "", common.ErrGetClientToken + } + token, ok := ctx.Value(constants.CtxKeyClientToken).(string) + if !ok || token == "" { + return "", fmt.Errorf("%w: missing or invalid context value", common.ErrGetClientToken) + } + return token, nil +} diff --git a/diagram/context_test.go b/diagram/context_test.go new file mode 100644 index 0000000..33ec1b9 --- /dev/null +++ b/diagram/context_test.go @@ -0,0 +1,38 @@ +package diagram + +import ( + "context" + "testing" + + "modelRT/common" + "modelRT/constants" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestClientTokenFromContext(t *testing.T) { + ctx := context.WithValue(context.Background(), constants.CtxKeyClientToken, "test-token") + token, err := clientTokenFromContext(ctx) + require.NoError(t, err) + assert.Equal(t, "test-token", token) +} + +func TestClientTokenFromContextReturnsErrorWhenMissing(t *testing.T) { + for _, ctx := range []context.Context{nil, context.Background()} { + _, err := clientTokenFromContext(ctx) + require.Error(t, err) + assert.ErrorIs(t, err, common.ErrGetClientToken) + } +} + +func TestRedisConstructorsReturnErrorInsteadOfPanickingWithoutToken(t *testing.T) { + ctx := context.Background() + + _, err := NewRedisZSet(ctx, "zset", 0, false) + assert.ErrorIs(t, err, common.ErrGetClientToken) + _, err = NewRedisSet(ctx, "set", 0, false) + assert.ErrorIs(t, err, common.ErrGetClientToken) + _, err = NewRedisHash(ctx, "hash", 0, false) + assert.ErrorIs(t, err, common.ErrGetClientToken) +} diff --git a/diagram/redis_hash.go b/diagram/redis_hash.go index 2382828..edc39c5 100644 --- a/diagram/redis_hash.go +++ b/diagram/redis_hash.go @@ -18,14 +18,17 @@ type RedisHash struct { } // NewRedisHash define func of new redis hash instance -func NewRedisHash(ctx context.Context, hashKey string, lockLeaseTime uint64, needRefresh bool) *RedisHash { - token := ctx.Value("client_token").(string) +func NewRedisHash(ctx context.Context, hashKey string, lockLeaseTime uint64, needRefresh bool) (*RedisHash, error) { + token, err := clientTokenFromContext(ctx) + if err != nil { + return nil, err + } return &RedisHash{ ctx: ctx, hashKey: hashKey, rwLocker: locker.InitRWLocker(hashKey, token, lockLeaseTime, needRefresh), storageClient: GetRedisClientInstance(), - } + }, nil } // SetRedisHashByMap define func of set redis hash by map struct diff --git a/diagram/redis_set.go b/diagram/redis_set.go index bfb9f6c..61d7064 100644 --- a/diagram/redis_set.go +++ b/diagram/redis_set.go @@ -21,15 +21,18 @@ type RedisSet struct { } // NewRedisSet define func of new redis set instance -func NewRedisSet(ctx context.Context, setKey string, lockLeaseTime uint64, needRefresh bool) *RedisSet { - token := ctx.Value("client_token").(string) +func NewRedisSet(ctx context.Context, setKey string, lockLeaseTime uint64, needRefresh bool) (*RedisSet, error) { + token, err := clientTokenFromContext(ctx) + if err != nil { + return nil, err + } return &RedisSet{ ctx: ctx, key: setKey, rwLocker: locker.InitRWLocker(setKey, token, lockLeaseTime, needRefresh), storageClient: GetRedisClientInstance(), logger: logger.GetLoggerInstance(), - } + }, nil } // SADD define func of add redis set by members diff --git a/diagram/redis_zset.go b/diagram/redis_zset.go index 6884448..d8f5ee1 100644 --- a/diagram/redis_zset.go +++ b/diagram/redis_zset.go @@ -18,13 +18,16 @@ type RedisZSet struct { } // NewRedisZSet define func of new redis zset instance -func NewRedisZSet(ctx context.Context, key string, lockLeaseTime uint64, needRefresh bool) *RedisZSet { - token := ctx.Value("client_token").(string) +func NewRedisZSet(ctx context.Context, key string, lockLeaseTime uint64, needRefresh bool) (*RedisZSet, error) { + token, err := clientTokenFromContext(ctx) + if err != nil { + return nil, err + } return &RedisZSet{ ctx: ctx, rwLocker: locker.InitRWLocker(key, token, lockLeaseTime, needRefresh), storageClient: GetRedisClientInstance(), - } + }, nil } // ZADD define func of add redis zset by members @@ -44,6 +47,26 @@ func (rs *RedisZSet) ZADD(setKey string, score float64, member any) error { return nil } +// ZREPLACE atomically removes all existing members and adds one new member. +func (rs *RedisZSet) ZREPLACE(setKey string, score float64, member any) error { + if err := rs.rwLocker.WLock(rs.ctx); err != nil { + logger.Error(rs.ctx, "lock wLock by setKey failed", "set_key", setKey, "error", err) + return err + } + defer rs.rwLocker.UnWLock(rs.ctx) + + _, err := rs.storageClient.TxPipelined(rs.ctx, func(pipe redis.Pipeliner) error { + pipe.Del(rs.ctx, setKey) + pipe.ZAdd(rs.ctx, setKey, redis.Z{Score: score, Member: member}) + return nil + }) + if err != nil { + logger.Error(rs.ctx, "replace zset member failed", "set_key", setKey, "member", member, "error", err) + return err + } + return nil +} + // ZRANGE define func of returns the specified range of elements in the sorted set stored by key func (rs *RedisZSet) ZRANGE(setKey string, start, stop int64) ([]string, error) { var results []string diff --git a/handler/component_attribute_query.go b/handler/component_attribute_query.go index 133a622..aa08801 100644 --- a/handler/component_attribute_query.go +++ b/handler/component_attribute_query.go @@ -54,7 +54,15 @@ func ComponentAttributeQueryHandler(c *gin.Context) { dbQueryMap := make(map[string][]cacheQueryItem) var secondaryQueryCount int for hSetKey, items := range cacheQueryMap { - hset := diagram.NewRedisHash(ctx, hSetKey, 5000, false) + hset, err := diagram.NewRedisHash(ctx, hSetKey, 5000, false) + if err != nil { + logger.Warn(ctx, "create redis hash failed", "key", hSetKey, "error", err) + for _, item := range items { + dbQueryMap[item.attributeCompTag] = append(dbQueryMap[item.attributeCompTag], item) + secondaryQueryCount++ + } + continue + } cacheData, err := hset.HGetAll() if err != nil { logger.Warn(ctx, "redis hgetall failed", "key", hSetKey, "err", err) @@ -185,7 +193,11 @@ func fillRemainingErrors(results map[string]queryResult, tokens []string, err *e } func backfillRedis(ctx context.Context, hSetKey string, items []cacheQueryItem) { - hset := diagram.NewRedisHash(ctx, hSetKey, 5000, false) + hset, err := diagram.NewRedisHash(ctx, hSetKey, 5000, false) + if err != nil { + logger.Error(ctx, "create redis hash for async backfill failed", "hash_key", hSetKey, "error", err) + return + } fields := make(map[string]any, len(items)) for _, item := range items { if item.attributeVal != "" { diff --git a/handler/component_attribute_update.go b/handler/component_attribute_update.go index 8a10f67..28b4836 100644 --- a/handler/component_attribute_update.go +++ b/handler/component_attribute_update.go @@ -140,7 +140,14 @@ func ComponentAttributeUpdateHandler(c *gin.Context) { } for key, items := range redisUpdateMap { - hset := diagram.NewRedisHash(ctx, key, 5000, false) + hset, err := diagram.NewRedisHash(ctx, key, 5000, false) + if err != nil { + logger.Error(ctx, "create redis hash failed", "hash_key", key, "error", err) + for _, item := range items { + updateResults[item.token] = errcode.ErrCacheSyncWarn.WithCause(err) + } + continue + } fields := make(map[string]any, len(items)) for _, item := range items { diff --git a/handler/data_object_attribute_update.go b/handler/data_object_attribute_update.go index 9be0719..828b3eb 100644 --- a/handler/data_object_attribute_update.go +++ b/handler/data_object_attribute_update.go @@ -28,6 +28,7 @@ type dataObjectAttributeUpdateRequest struct { Token string `json:"token"` Field string `json:"field"` Value json.RawMessage `json:"value"` + Data json.RawMessage `json:"data,omitempty"` } // DataObjectAttributeUpdateHandler updates the writable field of one data object. @@ -53,6 +54,12 @@ func DataObjectAttributeUpdateHandler(c *gin.Context) { renderRespFailure(c, constants.RespCodeServerError, "begin postgres transaction failed", nil) return } + transactionCompleted := false + defer func() { + if !transactionCompleted { + _ = tx.Rollback().Error + } + }() message := "data-object attribute update success" var measurementResult measurementUpdateResult @@ -64,7 +71,12 @@ func DataObjectAttributeUpdateHandler(c *gin.Context) { } err = queryErr case constants.DataObjectTypeMeasurement: - measurementResult, err = updateMeasurementDataObject(ctx, tx, request.Token, field, value, writeMeasurementManualValue, writeMeasurementManualValueToDataRT) + measurementResult, err = updateMeasurementDataObject(ctx, tx, request.Token, field, value, request.Data, measurementUpdateDependencies{ + writeManualValueFunc: writeMeasurementManualValue, + updateDataRTFunc: callRealTimeDataWriteStopInterface, + startDataRTFunc: callRealTimeDataWriteStartInterface, + replaceRedisValueFunc: replaceMeasurementRedisValue, + }) message = measurementResult.message default: err = fmt.Errorf("unsupported data object type %q", dataObjectType) @@ -91,6 +103,7 @@ func DataObjectAttributeUpdateHandler(c *gin.Context) { renderRespFailure(c, constants.RespCodeServerError, "transaction commit failed", nil) return } + transactionCompleted = true renderRespSuccess(c, constants.RespCodeSuccess, message, map[string]any{ "token": request.Token, @@ -109,7 +122,7 @@ func validateDataObjectAttributeUpdate(request dataObjectAttributeUpdateRequest) field := strings.ToLower(strings.TrimSpace(request.Field)) if field == "" { - return "", "", nil, fmt.Errorf("field is required") + field = "value" } dataObjectType, err := model.ClassifyDataObjectToken(request.Token) @@ -131,8 +144,8 @@ func validateDataObjectAttributeUpdate(request dataObjectAttributeUpdateRequest) return dataObjectType, field, value, err case constants.DataObjectTypeMeasurement: parts := strings.Split(request.Token, ".") - if len(parts) == 2 || parts[len(parts)-2] != "bay" { - return "", "", nil, fmt.Errorf("measurement updates require token6=bay") + if len(parts) != 2 && parts[len(parts)-2] != "bay" { + return "", "", nil, fmt.Errorf("measurement updates require token4.token7 or token6=bay") } switch field { case "value": @@ -151,7 +164,7 @@ func validateDataObjectAttributeUpdate(request dataObjectAttributeUpdateRequest) func isWritableParameterAttributeGroup(group string) bool { switch group { - case "rated", "setup", "model", "stable", "craft", "integrity", "behavior": + case "rated", "setup", "model", "stable", "craft", "integrity", "behavior", "base_extend": return true default: return false @@ -195,29 +208,29 @@ func parseMeasurementUpdateValue(raw json.RawMessage) (float64, error) { return number, nil } -func parseMeasurementUpdateMode(raw json.RawMessage) (bool, error) { - var mode bool - if err := json.Unmarshal(raw, &mode); err == nil { - return mode, nil +func parseMeasurementUpdateMode(raw json.RawMessage) (int16, error) { + var mode int16 + if err := json.Unmarshal(raw, &mode); err != nil { + return 0, fmt.Errorf("measurement mode must be 0 (manual) or 1 (automatic)") } - - var text string - if err := json.Unmarshal(raw, &text); err != nil { - return false, fmt.Errorf("measurement mode must be true or false") - } - switch strings.ToLower(text) { - case "true": - return true, nil - case "false": - return false, nil - default: - return false, fmt.Errorf("measurement mode must be true or false") + if mode != constants.MeasurementModeManual && mode != constants.MeasurementModeAutomatic { + return 0, fmt.Errorf("measurement mode must be 0 (manual) or 1 (automatic)") } + return mode, nil } type measurementManualValueWriter func(context.Context, *orm.Measurement, float64) error -type measurementDataRTWriter func(context.Context, orm.JSONMap, float64) error +type measurementDataRTUpdater func(context.Context, orm.JSONMap, *float64) error + +type measurementRedisValueReplacer func(context.Context, *orm.Measurement, float64) error + +type measurementUpdateDependencies struct { + writeManualValueFunc measurementManualValueWriter + updateDataRTFunc measurementDataRTUpdater + startDataRTFunc measurementDataRTUpdater + replaceRedisValueFunc measurementRedisValueReplacer +} type measurementUpdateResult struct { message string @@ -231,35 +244,74 @@ func updateMeasurementDataObject( tx *gorm.DB, token, field string, value any, - writeManualValue measurementManualValueWriter, - writeDataRT measurementDataRTWriter, + modeData json.RawMessage, + dependencies measurementUpdateDependencies, ) (measurementUpdateResult, error) { measurement, _, err := database.QueryMeasurementByDataObjectToken(ctx, tx, token) if err != nil { return measurementUpdateResult{}, err } - lockedMeasurement, err := database.QueryMeasurementByID(ctx, tx, measurement.ID) + lockedMeasurement, err := database.QueryMeasurementByIDForUpdate(ctx, tx, measurement.ID) if err != nil { return measurementUpdateResult{}, fmt.Errorf("lock measurement %d for update: %w", measurement.ID, err) } switch field { case "mode": - mode, ok := value.(bool) + mode, ok := value.(int16) if !ok { return measurementUpdateResult{}, fmt.Errorf("measurement mode has invalid type %T", value) } - currentMode := lockedMeasurement.Mode != 0 - if currentMode == mode { + currentMode, err := measurementModeIsAutomatic(lockedMeasurement.Mode) + if err != nil { + return measurementUpdateResult{}, err + } + targetAutomatic := mode == constants.MeasurementModeAutomatic + if currentMode == targetAutomatic { return measurementUpdateResult{message: fmt.Sprintf("measurement is already in %s mode", measurementModeName(mode))}, nil } + var manualValue *float64 + if currentMode && mode == constants.MeasurementModeManual { + manualValue, err = parseOptionalMeasurementModeData(modeData) + if err != nil { + return measurementUpdateResult{}, err + } + } if err := database.UpdateMeasurementModeWithOperation(ctx, tx, lockedMeasurement.ID, mode, time.Now().UTC()); err != nil { return measurementUpdateResult{}, err } + if currentMode && mode == constants.MeasurementModeManual { + if dependencies.updateDataRTFunc == nil { + return measurementUpdateResult{}, fmt.Errorf("measurement dataRT updater is nil") + } + if err := dependencies.updateDataRTFunc(ctx, lockedMeasurement.DataSource, nil); err != nil { + return measurementUpdateResult{}, fmt.Errorf("stop automatic measurement write to dataRT: %w", err) + } + if manualValue != nil { + if dependencies.replaceRedisValueFunc == nil { + return measurementUpdateResult{}, fmt.Errorf("measurement redis value replacer is nil") + } + if err := dependencies.replaceRedisValueFunc(ctx, &lockedMeasurement, *manualValue); err != nil { + return measurementUpdateResult{}, fmt.Errorf("replace measurement redis value: %w", err) + } + } + } + if !currentMode && mode == constants.MeasurementModeAutomatic { + if dependencies.startDataRTFunc == nil { + return measurementUpdateResult{}, fmt.Errorf("measurement dataRT starter is nil") + } + if err := dependencies.startDataRTFunc(ctx, lockedMeasurement.DataSource, nil); err != nil { + return measurementUpdateResult{}, fmt.Errorf("start automatic measurement write to dataRT: %w", err) + } + } return measurementUpdateResult{message: fmt.Sprintf("measurement mode changed to %s", measurementModeName(mode))}, nil case "value": - if lockedMeasurement.Mode != 0 { + currentMode, err := measurementModeIsAutomatic(lockedMeasurement.Mode) + if err != nil { + return measurementUpdateResult{}, err + } + if currentMode { return measurementUpdateResult{}, fmt.Errorf("measurement value is read-only while mode is automatic") } manualValue, ok := value.(float64) @@ -271,16 +323,16 @@ func updateMeasurementDataObject( value: manualValue, recordFailure: true, } - if writeManualValue == nil { + if dependencies.writeManualValueFunc == nil { return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(fmt.Errorf("measurement manual value writer is nil")) } - if err := writeManualValue(ctx, &lockedMeasurement, manualValue); err != nil { + if err := dependencies.writeManualValueFunc(ctx, &lockedMeasurement, manualValue); err != nil { return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(err) } - if writeDataRT == nil { - return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(fmt.Errorf("measurement dataRT writer is nil")) + if dependencies.updateDataRTFunc == nil { + return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(fmt.Errorf("measurement dataRT updater is nil")) } - if err := writeDataRT(ctx, lockedMeasurement.DataSource, manualValue); err != nil { + if err := dependencies.updateDataRTFunc(ctx, lockedMeasurement.DataSource, &manualValue); err != nil { return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(err) } if err := database.AppendMeasurementValueOperation(ctx, tx, lockedMeasurement.ID, 0, manualValue, time.Now().UTC()); err != nil { @@ -292,13 +344,36 @@ func updateMeasurementDataObject( } } -func measurementModeName(automatic bool) string { - if automatic { +func parseOptionalMeasurementModeData(raw json.RawMessage) (*float64, error) { + trimmed := bytes.TrimSpace(raw) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return nil, nil + } + value, err := parseMeasurementUpdateValue(trimmed) + if err != nil { + return nil, fmt.Errorf("invalid measurement mode data: %w", err) + } + return &value, nil +} + +func measurementModeName(mode int16) string { + if mode == constants.MeasurementModeAutomatic { return "automatic" } return "manual" } +func measurementModeIsAutomatic(mode int16) (bool, error) { + switch mode { + case constants.MeasurementModeManual: + return false, nil + case constants.MeasurementModeAutomatic: + return true, nil + default: + return false, fmt.Errorf("measurement has invalid mode %d", mode) + } +} + func isInvalidDataObjectUpdateError(err error) bool { return errors.Is(err, common.ErrInvalidParameterToken) || errors.Is(err, common.ErrParameterTokenNotFound) || @@ -313,14 +388,38 @@ func writeMeasurementManualValue(ctx context.Context, measurement *orm.Measureme if err != nil { return fmt.Errorf("generate measurement redis key: %w", err) } - zset := diagram.NewRedisZSet(ctx, key, 0, false) + zset, err := diagram.NewRedisZSet(ctx, key, 0, false) + if err != nil { + return fmt.Errorf("create measurement redis zset: %w", err) + } if err := zset.ZADD(key, value, strconv.FormatInt(time.Now().UnixNano(), 10)); err != nil { return fmt.Errorf("write manual measurement value to redis: %w", err) } return nil } -func writeMeasurementManualValueToDataRT(_ context.Context, _ orm.JSONMap, _ float64) error { - // TODO: call the dataRT HTTP API with the measurement data_source and manual value. +func callRealTimeDataWriteStopInterface(_ context.Context, _ orm.JSONMap, _ *float64) error { + // TODO: call the dataRT HTTP API. A nil value stops automatic writes; + // a non-nil value writes the supplied manual measurement value. + return nil +} + +func callRealTimeDataWriteStartInterface(_ context.Context, _ orm.JSONMap, _ *float64) error { + // TODO: call the dataRT HTTP API to start automatic measurement writes. + return nil +} + +func replaceMeasurementRedisValue(ctx context.Context, measurement *orm.Measurement, value float64) error { + key, err := model.GenerateMeasureIdentifier(measurement.DataSource) + if err != nil { + return fmt.Errorf("generate measurement redis key: %w", err) + } + zset, err := diagram.NewRedisZSet(ctx, key, 0, false) + if err != nil { + return fmt.Errorf("create measurement redis zset: %w", err) + } + if err := zset.ZREPLACE(key, value, strconv.FormatInt(time.Now().UnixNano(), 10)); err != nil { + return fmt.Errorf("replace manual measurement value in redis: %w", err) + } return nil } diff --git a/handler/data_object_attribute_update_test.go b/handler/data_object_attribute_update_test.go index 8e56135..03bb246 100644 --- a/handler/data_object_attribute_update_test.go +++ b/handler/data_object_attribute_update_test.go @@ -69,9 +69,10 @@ func TestValidateDataObjectAttributeUpdateMeasurementFields(t *testing.T) { }{ {name: "numeric value", field: "value", value: `15.2`, expected: float64(15.2)}, {name: "numeric string value", field: "value", value: `"15.2"`, expected: float64(15.2)}, - {name: "boolean automatic mode", field: "mode", value: `true`, expected: true}, - {name: "string manual mode", field: "mode", value: `"false"`, expected: false}, - {name: "invalid mode", field: "mode", value: `"manual"`, wantError: true}, + {name: "automatic mode", field: "mode", value: `1`, expected: constants.MeasurementModeAutomatic}, + {name: "manual mode", field: "mode", value: `0`, expected: constants.MeasurementModeManual}, + {name: "boolean mode is rejected", field: "mode", value: `true`, wantError: true}, + {name: "out of range mode", field: "mode", value: `2`, wantError: true}, {name: "unsupported field", field: "name", value: `"measurement"`, wantError: true}, } @@ -94,14 +95,15 @@ func TestValidateDataObjectAttributeUpdateMeasurementFields(t *testing.T) { } } -func TestValidateDataObjectAttributeUpdateRejectsShortMeasurementToken(t *testing.T) { - _, _, _, err := validateDataObjectAttributeUpdate(dataObjectAttributeUpdateRequest{ +func TestValidateDataObjectAttributeUpdateAcceptsToken4Token7Measurement(t *testing.T) { + dataObjectType, field, value, err := validateDataObjectAttributeUpdate(dataObjectAttributeUpdateRequest{ Token: "nspath.measurement", - Field: "value", - Value: json.RawMessage(`1`), + Value: json.RawMessage(`15.2`), }) - require.Error(t, err) - assert.Contains(t, err.Error(), "token6=bay") + require.NoError(t, err) + assert.Equal(t, constants.DataObjectTypeMeasurement, dataObjectType) + assert.Equal(t, "value", field) + assert.Equal(t, float64(15.2), value) } func TestValidateDataObjectAttributeUpdateRequiredFields(t *testing.T) { @@ -110,7 +112,6 @@ func TestValidateDataObjectAttributeUpdateRequiredFields(t *testing.T) { request dataObjectAttributeUpdateRequest }{ {name: "missing token", request: dataObjectAttributeUpdateRequest{Field: "value", Value: json.RawMessage(`1`)}}, - {name: "missing field", request: dataObjectAttributeUpdateRequest{Token: "nspath.measurement", Value: json.RawMessage(`1`)}}, {name: "missing value", request: dataObjectAttributeUpdateRequest{Token: "nspath.measurement", Field: "value"}}, {name: "null value", request: dataObjectAttributeUpdateRequest{Token: "nspath.measurement", Field: "value", Value: json.RawMessage(`null`)}}, } @@ -123,6 +124,59 @@ func TestValidateDataObjectAttributeUpdateRequiredFields(t *testing.T) { } } +func TestValidateDataObjectAttributeUpdateDefaultsEmptyFieldToValue(t *testing.T) { + tests := []struct { + name string + request dataObjectAttributeUpdateRequest + wantType constants.DataObjectType + wantValue any + }{ + { + name: "parameter omitted field", + request: dataObjectAttributeUpdateRequest{ + Token: "nspath.component.rated.attribute", + Value: json.RawMessage(`"15.2"`), + }, + wantType: constants.DataObjectTypeParameter, + wantValue: "15.2", + }, + { + name: "measurement whitespace field", + request: dataObjectAttributeUpdateRequest{ + Token: "nspath.component.bay.measurement", + Field: " ", + Value: json.RawMessage(`15.2`), + }, + wantType: constants.DataObjectTypeMeasurement, + wantValue: float64(15.2), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dataObjectType, field, value, err := validateDataObjectAttributeUpdate(tt.request) + require.NoError(t, err) + assert.Equal(t, tt.wantType, dataObjectType) + assert.Equal(t, "value", field) + assert.Equal(t, tt.wantValue, value) + }) + } +} + +func TestMeasurementModeIsAutomatic(t *testing.T) { + automatic, err := measurementModeIsAutomatic(constants.MeasurementModeAutomatic) + require.NoError(t, err) + assert.True(t, automatic) + + automatic, err = measurementModeIsAutomatic(constants.MeasurementModeManual) + require.NoError(t, err) + assert.False(t, automatic) + + _, err = measurementModeIsAutomatic(-1) + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid mode") +} + func TestUpdateMeasurementDataObjectLocksRowAndUpdatesMode(t *testing.T) { db, mock, closeDB := newDataObjectUpdateTestDB(t) defer closeDB() @@ -131,18 +185,51 @@ func TestUpdateMeasurementDataObjectLocksRowAndUpdatesMode(t *testing.T) { tx := db.Begin() require.NoError(t, tx.Error) expectMeasurementResolution(mock, 0) - mock.ExpectExec(`UPDATE "measurement" SET .*"mode"=\$1.*"operations"=.*WHERE id = \$5`). - WithArgs(int16(1), 500, 500, sqlmock.AnyArg(), int64(10)). + mock.ExpectExec(`UPDATE "measurement" SET .*"mode"=\$1.*"operations"=.*WHERE id = \$4`). + WithArgs(constants.MeasurementModeAutomatic, sqlmock.AnyArg(), 500, int64(10)). WillReturnResult(sqlmock.NewResult(0, 1)) mock.ExpectRollback() - result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "mode", true, nil, nil) + startCalled := false + result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "mode", constants.MeasurementModeAutomatic, nil, measurementUpdateDependencies{ + startDataRTFunc: func(_ context.Context, dataSource orm.JSONMap, value *float64) error { + startCalled = true + assert.Equal(t, float64(1), dataSource["type"]) + assert.Nil(t, value) + return nil + }, + }) require.NoError(t, err) + assert.True(t, startCalled) assert.Contains(t, result.message, "automatic") require.NoError(t, tx.Rollback().Error) require.NoError(t, mock.ExpectationsWereMet()) } +func TestUpdateMeasurementModeToAutomaticReturnsErrorWhenDataRTStartFails(t *testing.T) { + db, mock, closeDB := newDataObjectUpdateTestDB(t) + defer closeDB() + + mock.ExpectBegin() + tx := db.Begin() + require.NoError(t, tx.Error) + expectMeasurementResolution(mock, constants.MeasurementModeManual) + mock.ExpectExec(`UPDATE "measurement" SET .*"mode"=\$1.*"operations"=.*WHERE id = \$4`). + WithArgs(constants.MeasurementModeAutomatic, sqlmock.AnyArg(), 500, int64(10)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + _, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "mode", constants.MeasurementModeAutomatic, nil, measurementUpdateDependencies{ + startDataRTFunc: func(context.Context, orm.JSONMap, *float64) error { + return fmt.Errorf("dataRT unavailable") + }, + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "start automatic measurement write") + require.NoError(t, tx.Rollback().Error) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestUpdateMeasurementDataObjectReturnsMessageWhenModeIsUnchanged(t *testing.T) { db, mock, closeDB := newDataObjectUpdateTestDB(t) defer closeDB() @@ -153,13 +240,129 @@ func TestUpdateMeasurementDataObjectReturnsMessageWhenModeIsUnchanged(t *testing expectMeasurementResolution(mock, 1) mock.ExpectRollback() - result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "mode", true, nil, nil) + result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "mode", constants.MeasurementModeAutomatic, nil, measurementUpdateDependencies{}) require.NoError(t, err) assert.Contains(t, result.message, "already") require.NoError(t, tx.Rollback().Error) require.NoError(t, mock.ExpectationsWereMet()) } +func TestUpdateMeasurementModeToManualWithoutDataOnlyStopsDataRT(t *testing.T) { + db, mock, closeDB := newDataObjectUpdateTestDB(t) + defer closeDB() + + mock.ExpectBegin() + tx := db.Begin() + require.NoError(t, tx.Error) + expectMeasurementResolution(mock, constants.MeasurementModeAutomatic) + mock.ExpectExec(`UPDATE "measurement" SET .*"mode"=\$1.*"operations"=.*WHERE id = \$4`). + WithArgs(constants.MeasurementModeManual, sqlmock.AnyArg(), 500, int64(10)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + stopCalled := false + replaceCalled := false + result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "mode", constants.MeasurementModeManual, nil, measurementUpdateDependencies{ + updateDataRTFunc: func(_ context.Context, dataSource orm.JSONMap, value *float64) error { + stopCalled = true + assert.Equal(t, float64(1), dataSource["type"]) + assert.Nil(t, value) + return nil + }, + replaceRedisValueFunc: func(context.Context, *orm.Measurement, float64) error { + replaceCalled = true + return nil + }, + }) + require.NoError(t, err) + assert.True(t, stopCalled) + assert.False(t, replaceCalled) + assert.Contains(t, result.message, "manual") + require.NoError(t, tx.Rollback().Error) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateMeasurementModeToManualReplacesRedisValueWhenDataProvided(t *testing.T) { + db, mock, closeDB := newDataObjectUpdateTestDB(t) + defer closeDB() + + mock.ExpectBegin() + tx := db.Begin() + require.NoError(t, tx.Error) + expectMeasurementResolution(mock, constants.MeasurementModeAutomatic) + mock.ExpectExec(`UPDATE "measurement" SET .*"mode"=\$1.*"operations"=.*WHERE id = \$4`). + WithArgs(constants.MeasurementModeManual, sqlmock.AnyArg(), 500, int64(10)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + callOrder := make([]string, 0, 2) + result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "mode", constants.MeasurementModeManual, json.RawMessage(`0`), measurementUpdateDependencies{ + updateDataRTFunc: func(_ context.Context, _ orm.JSONMap, value *float64) error { + callOrder = append(callOrder, "stop-dataRT") + assert.Nil(t, value) + return nil + }, + replaceRedisValueFunc: func(_ context.Context, measurement *orm.Measurement, value float64) error { + callOrder = append(callOrder, "replace-redis") + assert.Equal(t, int64(10), measurement.ID) + assert.Equal(t, float64(0), value) + return nil + }, + }) + require.NoError(t, err) + assert.Equal(t, []string{"stop-dataRT", "replace-redis"}, callOrder) + assert.Contains(t, result.message, "manual") + require.NoError(t, tx.Rollback().Error) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateMeasurementModeToManualDoesNotTouchRedisWhenDataRTStopFails(t *testing.T) { + db, mock, closeDB := newDataObjectUpdateTestDB(t) + defer closeDB() + + mock.ExpectBegin() + tx := db.Begin() + require.NoError(t, tx.Error) + expectMeasurementResolution(mock, constants.MeasurementModeAutomatic) + mock.ExpectExec(`UPDATE "measurement" SET .*"mode"=\$1.*"operations"=.*WHERE id = \$4`). + WithArgs(constants.MeasurementModeManual, sqlmock.AnyArg(), 500, int64(10)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + replaceCalled := false + _, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "mode", constants.MeasurementModeManual, json.RawMessage(`15.2`), measurementUpdateDependencies{ + updateDataRTFunc: func(_ context.Context, _ orm.JSONMap, value *float64) error { + assert.Nil(t, value) + return fmt.Errorf("dataRT unavailable") + }, + replaceRedisValueFunc: func(context.Context, *orm.Measurement, float64) error { + replaceCalled = true + return nil + }, + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "stop automatic measurement write") + assert.False(t, replaceCalled) + require.NoError(t, tx.Rollback().Error) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestParseOptionalMeasurementModeData(t *testing.T) { + for _, raw := range []json.RawMessage{nil, json.RawMessage(`null`)} { + value, err := parseOptionalMeasurementModeData(raw) + require.NoError(t, err) + assert.Nil(t, value) + } + + value, err := parseOptionalMeasurementModeData(json.RawMessage(`"15.2"`)) + require.NoError(t, err) + require.NotNil(t, value) + assert.Equal(t, 15.2, *value) + + _, err = parseOptionalMeasurementModeData(json.RawMessage(`"invalid"`)) + require.Error(t, err) +} + func TestUpdateMeasurementDataObjectRejectsValueInAutomaticMode(t *testing.T) { db, mock, closeDB := newDataObjectUpdateTestDB(t) defer closeDB() @@ -170,7 +373,7 @@ func TestUpdateMeasurementDataObjectRejectsValueInAutomaticMode(t *testing.T) { expectMeasurementResolution(mock, 1) mock.ExpectRollback() - _, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "value", float64(15.2), nil, nil) + _, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "value", float64(15.2), nil, measurementUpdateDependencies{}) require.Error(t, err) assert.Contains(t, err.Error(), "read-only") require.NoError(t, tx.Rollback().Error) @@ -185,8 +388,8 @@ func TestUpdateMeasurementDataObjectWritesValueInManualMode(t *testing.T) { tx := db.Begin() require.NoError(t, tx.Error) expectMeasurementResolution(mock, 0) - mock.ExpectExec(`UPDATE "measurement" SET "operations"=.*WHERE id = \$4`). - WithArgs(500, 500, sqlmock.AnyArg(), int64(10)). + mock.ExpectExec(`UPDATE "measurement" SET "operations"=.*WHERE id = \$3`). + WithArgs(sqlmock.AnyArg(), 500, int64(10)). WillReturnResult(sqlmock.NewResult(0, 1)) mock.ExpectRollback() @@ -198,13 +401,17 @@ func TestUpdateMeasurementDataObjectWritesValueInManualMode(t *testing.T) { return nil } dataRTCalled := false - dataRTWriter := func(_ context.Context, dataSource orm.JSONMap, value float64) error { + dataRTWriter := func(_ context.Context, dataSource orm.JSONMap, value *float64) error { dataRTCalled = true - assert.Equal(t, float64(15.2), value) + require.NotNil(t, value) + assert.Equal(t, float64(15.2), *value) assert.Equal(t, float64(1), dataSource["type"]) return nil } - result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "value", float64(15.2), writer, dataRTWriter) + result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "value", float64(15.2), nil, measurementUpdateDependencies{ + writeManualValueFunc: writer, + updateDataRTFunc: dataRTWriter, + }) require.NoError(t, err) assert.True(t, called) assert.True(t, dataRTCalled) @@ -226,7 +433,9 @@ func TestUpdateMeasurementDataObjectReturnsFailureResultAndAppError(t *testing.T writeErr := fmt.Errorf("write value failed") writer := func(context.Context, *orm.Measurement, float64) error { return writeErr } - result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "value", float64(15.2), writer, nil) + result, err := updateMeasurementDataObject(context.Background(), tx, "nspath.measurement", "value", float64(15.2), nil, measurementUpdateDependencies{ + writeManualValueFunc: writer, + }) require.Error(t, err) assert.ErrorIs(t, err, errcode.ErrMeasurementValueUpdateFailed) assert.ErrorIs(t, err, writeErr) @@ -257,7 +466,7 @@ func expectMeasurementResolution(mock sqlmock.Sqlmock, mode int16) { WithArgs(componentUUID). WillReturnRows(sqlmock.NewRows([]string{"global_uuid", "nspath", "tag"}). AddRow(componentUUID, "nspath", "component")) - mock.ExpectQuery(`(?s)SELECT \* FROM "measurement" WHERE id = \$1.*LIMIT \$2 FOR UPDATE`). + mock.ExpectQuery(`SELECT "id","mode","data_source" FROM "measurement" WHERE id = \$1 LIMIT \$2 FOR UPDATE`). WithArgs(int64(10), 1). WillReturnRows(sqlmock.NewRows([]string{ "id", "tag", "mode", "data_source", "component_uuid", diff --git a/handler/diagram_node_link.go b/handler/diagram_node_link.go index ffd09a3..b10b0cd 100644 --- a/handler/diagram_node_link.go +++ b/handler/diagram_node_link.go @@ -85,7 +85,12 @@ func DiagramNodeLinkHandler(c *gin.Context) { return } - prevLinkSet, currLinkSet := generateLinkSet(ctx, nodeLevel, prevNodeInfo) + prevLinkSet, currLinkSet, err := generateLinkSet(ctx, nodeLevel, prevNodeInfo) + if err != nil { + logger.Error(ctx, "create diagram link redis sets failed", "node_id", nodeID, "level", nodeLevel, "error", err) + c.JSON(http.StatusOK, network.FailureResponse{Code: http.StatusInternalServerError, Msg: err.Error()}) + return + } err = processLinkSetData(ctx, action, nodeLevel, prevLinkSet, currLinkSet, prevNodeInfo, currNodeInfo) if err != nil { c.JSON(http.StatusOK, network.FailureResponse{ @@ -113,21 +118,27 @@ func DiagramNodeLinkHandler(c *gin.Context) { }) } -func generateLinkSet(ctx context.Context, level int, prevNodeInfo orm.CircuitDiagramNodeInterface) (*diagram.RedisSet, *diagram.RedisSet) { +func generateLinkSet(ctx context.Context, level int, prevNodeInfo orm.CircuitDiagramNodeInterface) (*diagram.RedisSet, *diagram.RedisSet, error) { config, ok := linkSetConfigs[level] // level not supported if !ok { - return nil, nil + return nil, nil, nil } - currLinkSet := diagram.NewRedisSet(ctx, config.CurrKey, 0, false) + currLinkSet, err := diagram.NewRedisSet(ctx, config.CurrKey, 0, false) + if err != nil { + return nil, nil, err + } if config.PrevIsNil { - return nil, currLinkSet + return nil, currLinkSet, nil } prevLinkSetKey := fmt.Sprintf(config.PrevKeyTemplate, prevNodeInfo.GetTagName()) - prevLinkSet := diagram.NewRedisSet(ctx, prevLinkSetKey, 0, false) - return prevLinkSet, currLinkSet + prevLinkSet, err := diagram.NewRedisSet(ctx, prevLinkSetKey, 0, false) + if err != nil { + return nil, nil, err + } + return prevLinkSet, currLinkSet, nil } func processLinkSetData(ctx context.Context, action string, level int, prevLinkSet, currLinkSet *diagram.RedisSet, prevNodeInfo, currNodeInfo orm.CircuitDiagramNodeInterface) error { diff --git a/handler/measurement_load.go b/handler/measurement_load.go index ddae642..2a57f29 100644 --- a/handler/measurement_load.go +++ b/handler/measurement_load.go @@ -39,7 +39,12 @@ func MeasurementGetHandler(c *gin.Context) { return } - zset := diagram.NewRedisZSet(ctx, request.MeasurementToken, 0, false) + zset, err := diagram.NewRedisZSet(ctx, request.MeasurementToken, 0, false) + if err != nil { + logger.Error(ctx, "failed to create measurement redis zset", "measurement_token", request.MeasurementToken, "error", err) + c.JSON(http.StatusOK, network.FailureResponse{Code: http.StatusInternalServerError, Msg: err.Error()}) + return + } points, err := zset.ZRANGE(request.MeasurementToken, 0, -1) if err != nil { logger.Error(ctx, "failed to get measurement data from redis", "measurement_token", request.MeasurementToken, "error", err) diff --git a/handler/mesurement_link.go b/handler/mesurement_link.go index 8737840..ae5b086 100644 --- a/handler/mesurement_link.go +++ b/handler/mesurement_link.go @@ -75,9 +75,19 @@ func MeasurementLinkHandler(c *gin.Context) { return } - allMeasSet := diagram.NewRedisSet(ctx, constants.RedisAllMeasTagSetKey, 0, false) + allMeasSet, err := diagram.NewRedisSet(ctx, constants.RedisAllMeasTagSetKey, 0, false) + if err != nil { + logger.Error(ctx, "create all-measurement redis set failed", "error", err) + c.JSON(http.StatusOK, network.FailureResponse{Code: http.StatusInternalServerError, Msg: err.Error()}) + return + } compMeasLinkKey := fmt.Sprintf(constants.RedisSpecCompTagMeasSetKey, componentInfo.Tag) - compMeasLinkSet := diagram.NewRedisSet(ctx, compMeasLinkKey, 0, false) + compMeasLinkSet, err := diagram.NewRedisSet(ctx, compMeasLinkKey, 0, false) + if err != nil { + logger.Error(ctx, "create component-measurement redis set failed", "set_key", compMeasLinkKey, "error", err) + c.JSON(http.StatusOK, network.FailureResponse{Code: http.StatusInternalServerError, Msg: err.Error()}) + return + } switch action { case constants.SearchLinkAddAction: diff --git a/middleware/token.go b/middleware/token.go index 6759f40..9d1bbf0 100644 --- a/middleware/token.go +++ b/middleware/token.go @@ -1,12 +1,20 @@ // Package middleware define gin framework middlewares package middleware -import "github.com/gin-gonic/gin" +import ( + "context" + + "modelRT/constants" + + "github.com/gin-gonic/gin" +) // SetTokenMiddleware define a middleware for set token in context func SetTokenMiddleware(clientToken string) gin.HandlerFunc { return func(c *gin.Context) { - c.Set("client_token", clientToken) + c.Set(constants.ClientTokenContextName, clientToken) + requestCtx := context.WithValue(c.Request.Context(), constants.CtxKeyClientToken, clientToken) + c.Request = c.Request.WithContext(requestCtx) c.Next() } } diff --git a/middleware/token_test.go b/middleware/token_test.go new file mode 100644 index 0000000..9d2f01c --- /dev/null +++ b/middleware/token_test.go @@ -0,0 +1,28 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "testing" + + "modelRT/constants" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" +) + +func TestSetTokenMiddlewarePropagatesClientToken(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(SetTokenMiddleware("test-token")) + router.GET("/test", func(c *gin.Context) { + assert.Equal(t, "test-token", c.GetString(constants.ClientTokenContextName)) + assert.Equal(t, "test-token", c.Request.Context().Value(constants.CtxKeyClientToken)) + c.Status(http.StatusNoContent) + }) + + request := httptest.NewRequest(http.MethodGet, "/test", nil) + response := httptest.NewRecorder() + router.ServeHTTP(response, request) + assert.Equal(t, http.StatusNoContent, response.Code) +}