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
This commit is contained in:
douxu 2026-07-21 16:13:23 +08:00
parent b85c2e129d
commit 34b9a569ae
18 changed files with 654 additions and 108 deletions

View File

@ -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

View File

@ -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
)

View File

@ -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)
}

View File

@ -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

20
diagram/context.go Normal file
View File

@ -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
}

38
diagram/context_test.go Normal file
View File

@ -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)
}

View File

@ -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

View File

@ -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

View File

@ -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

View File

@ -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 != "" {

View File

@ -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 {

View File

@ -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
}

View File

@ -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",

View File

@ -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 {

View File

@ -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)

View File

@ -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:

View File

@ -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()
}
}

28
middleware/token_test.go Normal file
View File

@ -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)
}