diff --git a/database/query_measurement_initialization.go b/database/query_measurement_initialization.go index 393c128..6a6e5f0 100644 --- a/database/query_measurement_initialization.go +++ b/database/query_measurement_initialization.go @@ -59,6 +59,13 @@ func validateMeasurementInitializationRecords(records []model.MeasurementInitial if _, err := model.MeasurementTypeString(record.MeasurementType); err != nil { return fmt.Errorf("measurement %q: %w", record.MeasurementTag, err) } + if record.MeasurementSize <= 0 { + return fmt.Errorf( + "measurement %q window size must be greater than 0, got %d", + record.MeasurementTag, + record.MeasurementSize, + ) + } if record.MeasurementDataSource == nil { return fmt.Errorf("measurement %q has null data_source", record.MeasurementTag) } diff --git a/database/query_measurement_initialization_test.go b/database/query_measurement_initialization_test.go index 132bd8e..c3553df 100644 --- a/database/query_measurement_initialization_test.go +++ b/database/query_measurement_initialization_test.go @@ -90,6 +90,12 @@ func TestValidateMeasurementInitializationRecords(t *testing.T) { err = validateMeasurementInitializationRecords([]model.MeasurementInitializationRecord{invalidType}) require.Error(t, err) assert.Contains(t, err.Error(), "unsupported measurement type -1") + + invalidSize := record + invalidSize.MeasurementSize = 0 + err = validateMeasurementInitializationRecords([]model.MeasurementInitializationRecord{invalidSize}) + require.Error(t, err) + assert.Contains(t, err.Error(), "window size must be greater than 0") } func measurementInitializationRows() *sqlmock.Rows { @@ -124,6 +130,7 @@ func validMeasurementInitializationRecord() model.MeasurementInitializationRecor MeasurementTag: "measurement", MeasurementType: 0, MeasurementMode: 1, + MeasurementSize: 1, MeasurementDataSource: map[string]any{}, MeasurementEventPlan: map[string]any{}, MeasurementBinding: map[string]any{}, diff --git a/diagram/redis_client.go b/diagram/redis_client.go index 247c98c..ddfe34b 100644 --- a/diagram/redis_client.go +++ b/diagram/redis_client.go @@ -4,6 +4,7 @@ package diagram import ( "context" "fmt" + "sort" "strconv" "github.com/redis/go-redis/v9" @@ -18,40 +19,73 @@ type RedisClient struct { // greatest numeric timestamp. Measurement ZSets currently store timestamp in // member and measurement value in score. func (rc *RedisClient) QueryLatestMeasurementValue(ctx context.Context, key string) (float64, error) { + values, err := rc.QueryLatestMeasurementValues(ctx, key, 1) + if err != nil { + return 0, err + } + return values[0], nil +} + +// QueryLatestMeasurementValues returns up to size scores ordered by their +// numeric member timestamps from newest to oldest. +func (rc *RedisClient) QueryLatestMeasurementValues(ctx context.Context, key string, size int) ([]float64, error) { if rc.Client == nil { - return 0, fmt.Errorf("redis client is not initialized") + return nil, fmt.Errorf("redis client is not initialized") + } + if size <= 0 { + return nil, fmt.Errorf("measurement window size must be greater than 0, got %d", size) } members, err := rc.Client.ZRangeWithScores(ctx, key, 0, -1).Result() if err != nil { - return 0, err + return nil, err } - return latestMeasurementValue(members, key) + return latestMeasurementValues(members, key, size) } func latestMeasurementValue(members []redis.Z, key string) (float64, error) { + values, err := latestMeasurementValues(members, key, 1) + if err != nil { + return 0, err + } + return values[0], nil +} + +func latestMeasurementValues(members []redis.Z, key string, size int) ([]float64, error) { + if size <= 0 { + return nil, fmt.Errorf("measurement window size must be greater than 0, got %d", size) + } if len(members) == 0 { - return 0, fmt.Errorf("real-time measurement value not found for key %q", key) + return nil, fmt.Errorf("real-time measurement value not found for key %q", key) } - var latestTimestamp int64 - var latestValue float64 - found := false + type timestampedValue struct { + timestamp int64 + value float64 + } + values := make([]timestampedValue, 0, len(members)) for _, member := range members { timestamp, err := strconv.ParseInt(fmt.Sprint(member.Member), 10, 64) if err != nil { continue } - if !found || timestamp > latestTimestamp { - latestTimestamp = timestamp - latestValue = member.Score - found = true - } + values = append(values, timestampedValue{timestamp: timestamp, value: member.Score}) } - if !found { - return 0, fmt.Errorf("real-time measurement timestamps are invalid for key %q", key) + if len(values) == 0 { + return nil, fmt.Errorf("real-time measurement timestamps are invalid for key %q", key) } - return latestValue, nil + + sort.Slice(values, func(i, j int) bool { + return values[i].timestamp > values[j].timestamp + }) + if size > len(values) { + size = len(values) + } + result := make([]float64, size) + for index := range size { + result[index] = values[index].value + } + return result, nil } // NewRedisClient define func of new redis client instance diff --git a/diagram/redis_client_test.go b/diagram/redis_client_test.go index 7416ba1..14960db 100644 --- a/diagram/redis_client_test.go +++ b/diagram/redis_client_test.go @@ -26,3 +26,32 @@ func TestLatestMeasurementValueRejectsMissingOrInvalidTimestamps(t *testing.T) { _, err = latestMeasurementValue([]redis.Z{{Member: "invalid", Score: 1}}, "measurement-key") require.Error(t, err) } + +func TestLatestMeasurementValuesReturnsNewestWindow(t *testing.T) { + values, err := latestMeasurementValues([]redis.Z{ + {Member: "100", Score: 10}, + {Member: "400", Score: 40}, + {Member: "invalid", Score: 999}, + {Member: "200", Score: 20}, + {Member: "300", Score: 30}, + }, "measurement-key", 3) + + require.NoError(t, err) + assert.Equal(t, []float64{40, 30, 20}, values) +} + +func TestLatestMeasurementValuesReturnsAvailableWindow(t *testing.T) { + values, err := latestMeasurementValues([]redis.Z{ + {Member: "100", Score: 10}, + {Member: "200", Score: 20}, + }, "measurement-key", 5) + + require.NoError(t, err) + assert.Equal(t, []float64{20, 10}, values) +} + +func TestLatestMeasurementValuesRejectsInvalidSize(t *testing.T) { + _, err := latestMeasurementValues([]redis.Z{{Member: "100", Score: 10}}, "measurement-key", 0) + require.Error(t, err) + assert.Contains(t, err.Error(), "window size must be greater than 0") +} diff --git a/handler/data_object_attribute_query.go b/handler/data_object_attribute_query.go index 77c2faf..84b0e4f 100644 --- a/handler/data_object_attribute_query.go +++ b/handler/data_object_attribute_query.go @@ -3,26 +3,27 @@ package handler import ( "context" + "encoding/json" "errors" "fmt" + "strconv" "strings" "modelRT/common" "modelRT/common/errcode" "modelRT/constants" - "modelRT/database" "modelRT/diagram" "modelRT/logger" "modelRT/model" "modelRT/orm" "github.com/gin-gonic/gin" + "github.com/redis/go-redis/v9" ) // DataObjectAttributeQueryHandler define data object attribute value query process API func DataObjectAttributeQueryHandler(c *gin.Context) { ctx := c.Request.Context() - pgClient := database.GetPostgresDBClient() token, field, err := parseDataObjectAttributeQuery(c) if err != nil { @@ -44,112 +45,37 @@ func DataObjectAttributeQueryHandler(c *gin.Context) { return } - var parameter *database.ParameterDataObject - var measurement *orm.Measurement - var measurementComponent *orm.Component - switch dataObjectType { - case constants.DataObjectTypeParameter: - // 参量支持两种形式token4.token5.token6.token7与token1.token2.token3.token4.token5.token6.token7 - parameter, err = database.QueryParameterByDataObjectToken(ctx, pgClient, token) - if err != nil { - if errors.Is(err, common.ErrInvalidParameterToken) || - errors.Is(err, common.ErrParameterTokenNotFound) || - errors.Is(err, common.ErrAmbiguousParameterToken) { - logger.Warn(ctx, "validate parameter token failed", "token", token, "error", err) - renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) - return - } - - logger.Error(ctx, "query parameter token from postgres failed", "token", token, "error", err) - renderRespFailure(c, constants.RespCodeServerError, "validate parameter token failed", nil) + value, err := queryDataObjectAttributeValue( + ctx, + dataObjectType, + token, + field, + loadDataObjectHashField, + loadMeasurementValueMetadata, + queryMeasurementRealtimeValue, + ) + if err != nil { + if isDataObjectTokenNotFound(err) { + logger.Warn(ctx, "query data-object token from redis failed", "token", token, "field", field, "error", err) + renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) return } - case constants.DataObjectTypeMeasurement: - // 量测支持token1.token2.token3.token4.token5.token6.token7、token4.token5.token6.token7、token4.token7 - measurement, measurementComponent, err = database.QueryMeasurementByDataObjectToken(ctx, pgClient, token) - if err != nil { - if errors.Is(err, common.ErrInvalidMeasurementToken) || - errors.Is(err, common.ErrMeasurementTokenNotFound) || - errors.Is(err, common.ErrAmbiguousMeasurementToken) { - logger.Warn(ctx, "validate measurement token failed", "token", token, "error", err) - renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) - return - } - logger.Error(ctx, "query measurement token from postgres failed", "token", token, "error", err) - renderRespFailure(c, constants.RespCodeServerError, "validate measurement token failed", nil) - return - } + logger.Error(ctx, "query data-object attribute from redis failed", "token", token, "field", field, "error", err) + renderRespFailure(c, constants.RespCodeServerError, dataObjectAttributeFailureMessage(dataObjectType), nil) + return } - switch dataObjectType { - case constants.DataObjectTypeParameter: - value, err := buildParameterAttributeValue( - ctx, - field, - parameter, - func(ctx context.Context, parameter *database.ParameterDataObject) (any, error) { - return database.QueryParameterDataObjectValue(ctx, pgClient, parameter) - }, - func(ctx context.Context, attributeName string) (string, error) { - return database.QueryParameterAttributeDescription(ctx, pgClient, attributeName) - }, - ) - if err != nil { - if errors.Is(err, common.ErrUnsupportedParameterField) { - logger.Warn(ctx, "query unsupported parameter field", "token", token, "field", field, "error", err) - renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) - return - } - - logger.Error(ctx, "build parameter attribute value failed", "token", token, "field", field, "error", err) - renderRespFailure(c, constants.RespCodeServerError, "query parameter attribute failed", nil) - return - } - - result := dataObjectAttributeQueryResult{ - Token: token, - Field: field, - Code: errcode.ErrProcessSuccess.Code(), - Msg: errcode.ErrProcessSuccess.Msg(), - Value: value, - } - renderRespSuccess(c, constants.RespCodeSuccess, "query parameter attribute success", map[string]any{ - "attributes": []dataObjectAttributeQueryResult{result}, - }) - case constants.DataObjectTypeMeasurement: - value, err := buildMeasurementAttributeValue( - ctx, - field, - measurement, - measurementComponent, - queryMeasurementRealtimeValue, - ) - if err != nil { - if errors.Is(err, common.ErrUnsupportedMeasurementField) { - logger.Warn(ctx, "query unsupported measurement field", "token", token, "field", field, "error", err) - renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) - return - } - - logger.Error(ctx, "build measurement attribute value failed", "token", token, "field", field, "error", err) - renderRespFailure(c, constants.RespCodeServerError, "query measurement attribute failed", nil) - return - } - - result := dataObjectAttributeQueryResult{ - Token: token, - Field: field, - Code: errcode.ErrProcessSuccess.Code(), - Msg: errcode.ErrProcessSuccess.Msg(), - Value: value, - } - renderRespSuccess(c, constants.RespCodeSuccess, "query measurement attribute success", map[string]any{ - "attributes": []dataObjectAttributeQueryResult{result}, - }) - default: - renderRespFailure(c, constants.RespCodeInvalidParams, "invalid data object type", nil) + result := dataObjectAttributeQueryResult{ + Token: token, + Field: field, + Code: errcode.ErrProcessSuccess.Code(), + Msg: errcode.ErrProcessSuccess.Msg(), + Value: value, } + renderRespSuccess(c, constants.RespCodeSuccess, dataObjectAttributeSuccessMessage(dataObjectType), map[string]any{ + "attributes": []dataObjectAttributeQueryResult{result}, + }) } func parseDataObjectAttributeQuery(c *gin.Context) (string, string, error) { @@ -165,11 +91,11 @@ func parseDataObjectAttributeQuery(c *gin.Context) (string, string, error) { return token, field, nil } -type measurementValueLoader func(context.Context, orm.JSONMap) (any, error) +type dataObjectHashFieldLoader func(context.Context, constants.DataObjectType, string, string) (string, error) -type parameterValueLoader func(context.Context, *database.ParameterDataObject) (any, error) +type measurementValueMetadataLoader func(context.Context, string) (orm.JSONMap, int, error) -type parameterDescriptionLoader func(context.Context, string) (string, error) +type measurementValueLoader func(context.Context, orm.JSONMap, int) (any, error) var measurementDataObjectFields = map[string]struct{}{ "value": {}, @@ -220,116 +146,242 @@ func validateDataObjectField(dataObjectType constants.DataObjectType, field stri } } -func buildParameterAttributeValue( +func queryDataObjectAttributeValue( ctx context.Context, + dataObjectType constants.DataObjectType, + token string, field string, - parameter *database.ParameterDataObject, - loadValue parameterValueLoader, - loadDescription parameterDescriptionLoader, + loadHashField dataObjectHashFieldLoader, + loadMeasurementMetadata measurementValueMetadataLoader, + loadMeasurementValue measurementValueLoader, ) (any, error) { - if parameter == nil { - return nil, fmt.Errorf("parameter data object is nil") - } - - component := parameter.Component - switch field { - case "value": - if loadValue == nil { - return nil, fmt.Errorf("parameter value loader is nil") + if dataObjectType == constants.DataObjectTypeMeasurement && field == "value" { + if loadMeasurementMetadata == nil { + return nil, fmt.Errorf("measurement value metadata loader is nil") } - return loadValue(ctx, parameter) - case "meta": - return "PARAM", nil - case "type": - return parameter.AttributeType, nil - case "name": - return strings.Join([]string{ - component.NSPath, - component.Tag, - parameter.AttributeGroup, - parameter.AttributeName, - }, "."), nil - case "description": - if loadDescription == nil { - return nil, fmt.Errorf("parameter description loader is nil") - } - return loadDescription(ctx, parameter.AttributeName) - case "id": - return strings.Join([]string{ - component.GridName, - component.ZoneName, - component.StationName, - component.NSPath, - component.Tag, - parameter.AttributeGroup, - parameter.AttributeName, - }, "."), nil - default: - return nil, fmt.Errorf("%w: %s", common.ErrUnsupportedParameterField, field) - } -} - -func buildMeasurementAttributeValue( - ctx context.Context, - field string, - measurement *orm.Measurement, - component *orm.Component, - loadValue measurementValueLoader, -) (any, error) { - if measurement == nil { - return nil, fmt.Errorf("measurement is nil") - } - if component == nil { - return nil, fmt.Errorf("measurement component is nil") - } - - switch field { - case "value": - if loadValue == nil { + if loadMeasurementValue == nil { return nil, fmt.Errorf("measurement value loader is nil") } - return loadValue(ctx, measurement.DataSource) + dataSource, size, err := loadMeasurementMetadata(ctx, token) + if err != nil { + return nil, err + } + return loadMeasurementValue(ctx, dataSource, size) + } + + if loadHashField == nil { + return nil, fmt.Errorf("data-object hash field loader is nil") + } + rawValue, err := loadHashField(ctx, dataObjectType, token, field) + if err != nil { + return nil, err + } + if dataObjectType == constants.DataObjectTypeParameter && field == "value" { + attributeType, err := loadHashField(ctx, dataObjectType, token, "type") + if err != nil { + return nil, err + } + return decodeParameterHashValue(rawValue, attributeType) + } + return decodeDataObjectHashField(dataObjectType, field, rawValue) +} + +func loadDataObjectHashField( + ctx context.Context, + dataObjectType constants.DataObjectType, + token string, + field string, +) (string, error) { + rdb := diagram.GetRedisClientInstance() + if rdb == nil { + return "", fmt.Errorf("redis client is not initialized") + } + value, err := rdb.HGet(ctx, token, field).Result() + if errors.Is(err, redis.Nil) { + exists, existsErr := rdb.Exists(ctx, token).Result() + if existsErr != nil { + return "", fmt.Errorf("check redis data-object hash %q: %w", token, existsErr) + } + if exists > 0 { + return "", fmt.Errorf("redis data-object hash %q does not contain field %q", token, field) + } + switch dataObjectType { + case constants.DataObjectTypeParameter: + return "", fmt.Errorf("%w: %q", common.ErrParameterTokenNotFound, token) + case constants.DataObjectTypeMeasurement: + return "", fmt.Errorf("%w: %q", common.ErrMeasurementTokenNotFound, token) + default: + return "", fmt.Errorf("invalid data object type %q", dataObjectType) + } + } + if err != nil { + return "", fmt.Errorf("query redis hash %q field %q: %w", token, field, err) + } + return value, nil +} + +func loadMeasurementValueMetadata(ctx context.Context, token string) (orm.JSONMap, int, error) { + rdb := diagram.GetRedisClientInstance() + if rdb == nil { + return nil, 0, fmt.Errorf("redis client is not initialized") + } + + values, err := rdb.HMGet(ctx, token, "data_source", "size").Result() + if err != nil { + return nil, 0, fmt.Errorf("query redis hash %q measurement value metadata: %w", token, err) + } + if len(values) != 2 { + return nil, 0, fmt.Errorf("redis hash %q returned %d measurement metadata fields", token, len(values)) + } + if values[0] == nil || values[1] == nil { + exists, existsErr := rdb.Exists(ctx, token).Result() + if existsErr != nil { + return nil, 0, fmt.Errorf("check redis data-object hash %q: %w", token, existsErr) + } + if exists == 0 { + return nil, 0, fmt.Errorf("%w: %q", common.ErrMeasurementTokenNotFound, token) + } + + missingFields := make([]string, 0, 2) + if values[0] == nil { + missingFields = append(missingFields, "data_source") + } + if values[1] == nil { + missingFields = append(missingFields, "size") + } + return nil, 0, fmt.Errorf( + "redis measurement hash %q does not contain field(s) %s", + token, + strings.Join(missingFields, ", "), + ) + } + + rawDataSource, ok := values[0].(string) + if !ok { + return nil, 0, fmt.Errorf("redis measurement hash %q data_source has type %T", token, values[0]) + } + var dataSource orm.JSONMap + if err := json.Unmarshal([]byte(rawDataSource), &dataSource); err != nil { + return nil, 0, fmt.Errorf("decode measurement data_source from redis hash %q: %w", token, err) + } + + rawSize, ok := values[1].(string) + if !ok { + return nil, 0, fmt.Errorf("redis measurement hash %q size has type %T", token, values[1]) + } + size, err := strconv.Atoi(rawSize) + if err != nil { + return nil, 0, fmt.Errorf("decode measurement size %q: %w", rawSize, err) + } + if size <= 0 { + return nil, 0, fmt.Errorf("measurement window size must be greater than 0, got %d", size) + } + return dataSource, size, nil +} + +func decodeDataObjectHashField( + dataObjectType constants.DataObjectType, + field string, + rawValue string, +) (any, error) { + if dataObjectType != constants.DataObjectTypeMeasurement { + return rawValue, nil + } + + switch field { case "mode": - return measurement.Mode, nil - case "meta": - return "MEASUREMENT", nil - case "type": - return model.MeasurementTypeString(measurement.Type) - case "name": - // The resolved measurement and component prove that token4.token7 exists. - return component.NSPath + "." + measurement.Tag, nil - case "description": - return measurement.Name, nil - case "id": - return strings.Join([]string{ - component.GridName, - component.ZoneName, - component.StationName, - component.NSPath, - component.Tag, - "bay", - measurement.Tag, - }, "."), nil + value, err := strconv.ParseInt(rawValue, 10, 16) + if err != nil { + return nil, fmt.Errorf("decode measurement mode %q: %w", rawValue, err) + } + return int16(value), nil case "size": - return measurement.Size, nil - case "data_source": - return measurement.DataSource, nil - case "event_plan": - return measurement.EventPlan, nil - case "binding": - return measurement.Binding, nil + value, err := strconv.Atoi(rawValue) + if err != nil { + return nil, fmt.Errorf("decode measurement size %q: %w", rawValue, err) + } + return value, nil + case "data_source", "event_plan", "binding": + return decodeRedisJSON(rawValue) default: - return nil, fmt.Errorf("%w: %s", common.ErrUnsupportedMeasurementField, field) + return rawValue, nil } } -func queryMeasurementRealtimeValue(ctx context.Context, dataSource orm.JSONMap) (any, error) { +func decodeParameterHashValue(rawValue, attributeType string) (any, error) { + normalizedType := strings.ToUpper(strings.TrimSpace(attributeType)) + if rawValue == "null" { + return nil, nil + } + + switch { + case normalizedType == "BOOLEAN": + value, err := strconv.ParseBool(rawValue) + if err != nil { + return nil, fmt.Errorf("decode parameter boolean value %q: %w", rawValue, err) + } + return value, nil + case normalizedType == "SMALLINT", + normalizedType == "INTEGER", + normalizedType == "BIGINT": + value, err := strconv.ParseInt(rawValue, 10, 64) + if err != nil { + return nil, fmt.Errorf("decode parameter integer value %q: %w", rawValue, err) + } + return value, nil + case normalizedType == "REAL", + normalizedType == "DOUBLE PRECISION", + strings.HasPrefix(normalizedType, "NUMERIC"), + strings.HasPrefix(normalizedType, "DECIMAL"): + if _, err := strconv.ParseFloat(rawValue, 64); err != nil { + return nil, fmt.Errorf("decode parameter numeric value %q: %w", rawValue, err) + } + return json.Number(rawValue), nil + case normalizedType == "JSON", + normalizedType == "JSONB", + strings.HasSuffix(normalizedType, "[]"): + return decodeRedisJSON(rawValue) + default: + return rawValue, nil + } +} + +func decodeRedisJSON(rawValue string) (any, error) { + var value any + decoder := json.NewDecoder(strings.NewReader(rawValue)) + decoder.UseNumber() + if err := decoder.Decode(&value); err != nil { + return nil, fmt.Errorf("decode redis JSON value: %w", err) + } + return value, nil +} + +func isDataObjectTokenNotFound(err error) bool { + return errors.Is(err, common.ErrParameterTokenNotFound) || + errors.Is(err, common.ErrMeasurementTokenNotFound) +} + +func dataObjectAttributeSuccessMessage(dataObjectType constants.DataObjectType) string { + if dataObjectType == constants.DataObjectTypeParameter { + return "query parameter attribute success" + } + return "query measurement attribute success" +} + +func dataObjectAttributeFailureMessage(dataObjectType constants.DataObjectType) string { + if dataObjectType == constants.DataObjectTypeParameter { + return "query parameter attribute failed" + } + return "query measurement attribute failed" +} + +func queryMeasurementRealtimeValue(ctx context.Context, dataSource orm.JSONMap, size int) (any, error) { queryKey, err := model.GenerateMeasureIdentifier(dataSource) if err != nil { return nil, fmt.Errorf("generate measurement redis key: %w", err) } - value, err := diagram.NewRedisClient().QueryLatestMeasurementValue(ctx, queryKey) + value, err := diagram.NewRedisClient().QueryLatestMeasurementValues(ctx, queryKey, size) if err != nil { return nil, fmt.Errorf("query real-time measurement value by key %q: %w", queryKey, err) } diff --git a/handler/data_object_attribute_query_test.go b/handler/data_object_attribute_query_test.go index f887def..efe7d74 100644 --- a/handler/data_object_attribute_query_test.go +++ b/handler/data_object_attribute_query_test.go @@ -2,13 +2,14 @@ package handler import ( "context" + "encoding/json" + "fmt" "net/http" "net/http/httptest" "testing" "modelRT/common" "modelRT/constants" - "modelRT/database" "modelRT/orm" "github.com/gin-gonic/gin" @@ -117,161 +118,156 @@ func TestValidateDataObjectField(t *testing.T) { } } -func TestBuildParameterAttributeValue(t *testing.T) { - parameter := &database.ParameterDataObject{ - Component: orm.Component{ - GridName: "grid000", - ZoneName: "zone000", - StationName: "station000", - NSPath: "110kV_TV", - Tag: "cable_22", - }, - AttributeGroup: "rated", - AttributeName: "rated_voltage", - AttributeType: "DOUBLE PRECISION", - } - loader := func(_ context.Context, actual *database.ParameterDataObject) (any, error) { - assert.Same(t, parameter, actual) - return float64(220), nil - } - descriptionLoader := func(_ context.Context, attributeName string) (string, error) { - assert.Equal(t, "rated_voltage", attributeName) - return "额定电压", nil +func TestQueryParameterAttributeValueFromRedisHash(t *testing.T) { + fields := map[string]string{ + "value": "220.50", + "type": "DOUBLE PRECISION", + "name": "110kV_TV.cable_22.rated.rated_voltage", + "description": "额定电压", } + loader := hashFieldLoaderForTest(fields) - tests := []struct { - field string - expected any - }{ - {field: "value", expected: float64(220)}, - {field: "meta", expected: "PARAM"}, - {field: "type", expected: "DOUBLE PRECISION"}, - {field: "name", expected: "110kV_TV.cable_22.rated.rated_voltage"}, - {field: "description", expected: "额定电压"}, - {field: "id", expected: "grid000.zone000.station000.110kV_TV.cable_22.rated.rated_voltage"}, - } - - for _, tt := range tests { - t.Run(tt.field, func(t *testing.T) { - actual, err := buildParameterAttributeValue( - context.Background(), - tt.field, - parameter, - loader, - descriptionLoader, - ) - require.NoError(t, err) - assert.Equal(t, tt.expected, actual) - }) - } -} - -func TestBuildParameterAttributeValueRejectsUnsupportedField(t *testing.T) { - _, err := buildParameterAttributeValue( + value, err := queryDataObjectAttributeValue( context.Background(), - "unknown", - &database.ParameterDataObject{}, + constants.DataObjectTypeParameter, + "parameter-token", + "value", + loader, nil, nil, ) - require.Error(t, err) - assert.ErrorIs(t, err, common.ErrUnsupportedParameterField) + require.NoError(t, err) + assert.Equal(t, json.Number("220.50"), value) + + description, err := queryDataObjectAttributeValue( + context.Background(), + constants.DataObjectTypeParameter, + "parameter-token", + "description", + loader, + nil, + nil, + ) + require.NoError(t, err) + assert.Equal(t, "额定电压", description) } -func TestBuildMeasurementAttributeValue(t *testing.T) { +func TestQueryMeasurementAttributeValueFromRedisHash(t *testing.T) { + fields := map[string]string{ + "mode": "1", + "size": "10", + "name": "110kV_TV.IA_rms", + "data_source": `{"type":1,"io_address":{"channel":"tm1p"}}`, + "event_plan": `{"enabled":true}`, + } + loader := hashFieldLoaderForTest(fields) + + mode, err := queryDataObjectAttributeValue( + context.Background(), + constants.DataObjectTypeMeasurement, + "measurement-token", + "mode", + loader, + nil, + nil, + ) + require.NoError(t, err) + assert.Equal(t, int16(1), mode) + + eventPlan, err := queryDataObjectAttributeValue( + context.Background(), + constants.DataObjectTypeMeasurement, + "measurement-token", + "event_plan", + loader, + nil, + nil, + ) + require.NoError(t, err) + assert.Equal(t, map[string]any{"enabled": true}, eventPlan) +} + +func TestQueryMeasurementRealtimeValueUsesDataSourceFromRedisHash(t *testing.T) { dataSource := orm.JSONMap{ "type": float64(1), "io_address": map[string]any{ + "station": "001", "channel": "tm1p", }, } - eventPlan := orm.JSONMap{"enabled": true} - binding := orm.JSONMap{"ct": map[string]any{"ratio": float64(2)}} - measurement := &orm.Measurement{ - Tag: "IA_rms", - Name: "A相电流", - Type: 0, - Mode: 1, - Size: 10, - DataSource: dataSource, - EventPlan: eventPlan, - Binding: binding, + metadataLoader := func(_ context.Context, token string) (orm.JSONMap, int, error) { + assert.Equal(t, "measurement-token", token) + return dataSource, 2, nil } - component := &orm.Component{ - GridName: "grid000", - ZoneName: "zone000", - StationName: "station000", - NSPath: "110kV_TV", - Tag: "cable_22", + valueLoader := func(_ context.Context, dataSource orm.JSONMap, size int) (any, error) { + assert.Equal(t, float64(1), dataSource["type"]) + assert.Equal(t, "001", dataSource["io_address"].(map[string]any)["station"]) + assert.Equal(t, 2, size) + return []float64{220, 219.5}, nil } - loader := func(_ context.Context, source orm.JSONMap) (any, error) { - assert.Equal(t, dataSource, source) - return float64(220), nil - } - - tests := []struct { - field string - expected any - }{ - {field: "value", expected: float64(220)}, - {field: "mode", expected: int16(1)}, - {field: "meta", expected: "MEASUREMENT"}, - {field: "type", expected: "TM"}, - {field: "name", expected: "110kV_TV.IA_rms"}, - {field: "description", expected: "A相电流"}, - {field: "id", expected: "grid000.zone000.station000.110kV_TV.cable_22.bay.IA_rms"}, - {field: "size", expected: 10}, - {field: "data_source", expected: dataSource}, - {field: "event_plan", expected: eventPlan}, - {field: "binding", expected: binding}, - } - - for _, tt := range tests { - t.Run(tt.field, func(t *testing.T) { - actual, err := buildMeasurementAttributeValue(context.Background(), tt.field, measurement, component, loader) - require.NoError(t, err) - assert.Equal(t, tt.expected, actual) - }) - } + value, err := queryDataObjectAttributeValue( + context.Background(), + constants.DataObjectTypeMeasurement, + "measurement-token", + "value", + nil, + metadataLoader, + valueLoader, + ) + require.NoError(t, err) + assert.Equal(t, []float64{220, 219.5}, value) } -func TestBuildMeasurementAttributeValueRejectsUnsupportedField(t *testing.T) { - _, err := buildMeasurementAttributeValue( +func TestQueryDataObjectAttributeValuePropagatesRedisTokenNotFound(t *testing.T) { + loader := func(context.Context, constants.DataObjectType, string, string) (string, error) { + return "", fmt.Errorf("%w: token", common.ErrParameterTokenNotFound) + } + + _, err := queryDataObjectAttributeValue( context.Background(), - "unknown", - &orm.Measurement{}, - &orm.Component{}, + constants.DataObjectTypeParameter, + "missing-token", + "name", + loader, + nil, nil, ) require.Error(t, err) - assert.ErrorIs(t, err, common.ErrUnsupportedMeasurementField) + assert.ErrorIs(t, err, common.ErrParameterTokenNotFound) + assert.True(t, isDataObjectTokenNotFound(err)) } -func TestBuildMeasurementAttributeValueMode(t *testing.T) { - component := &orm.Component{} +func TestDecodeParameterHashValue(t *testing.T) { tests := []struct { - name string - mode int16 - expected int16 + name string + rawValue string + attributeType string + expected any }{ - {name: "collected value", mode: 1, expected: 1}, - {name: "manually assigned value", mode: 0, expected: 0}, - {name: "other positive mode", mode: 2, expected: 2}, - {name: "negative mode", mode: -1, expected: -1}, + {name: "boolean", rawValue: "true", attributeType: "BOOLEAN", expected: true}, + {name: "integer", rawValue: "42", attributeType: "INTEGER", expected: int64(42)}, + {name: "numeric", rawValue: "1234567890.123456789", attributeType: "NUMERIC(30,9)", expected: json.Number("1234567890.123456789")}, + {name: "jsonb", rawValue: `{"key":"value"}`, attributeType: "JSONB", expected: map[string]any{"key": "value"}}, + {name: "string", rawValue: "cable", attributeType: "CHARACTER VARYING(64)", expected: "cable"}, + {name: "null", rawValue: "null", attributeType: "INTEGER", expected: nil}, } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - actual, err := buildMeasurementAttributeValue( - context.Background(), - "mode", - &orm.Measurement{Mode: tt.mode}, - component, - nil, - ) + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + actual, err := decodeParameterHashValue(test.rawValue, test.attributeType) require.NoError(t, err) - assert.Equal(t, tt.expected, actual) + assert.Equal(t, test.expected, actual) }) } } + +func hashFieldLoaderForTest(fields map[string]string) dataObjectHashFieldLoader { + return func(_ context.Context, _ constants.DataObjectType, _ string, field string) (string, error) { + value, exists := fields[field] + if !exists { + return "", fmt.Errorf("field %q not found", field) + } + return value, nil + } +} diff --git a/handler/data_object_attribute_update.go b/handler/data_object_attribute_update.go index 828b3eb..feb1038 100644 --- a/handler/data_object_attribute_update.go +++ b/handler/data_object_attribute_update.go @@ -61,6 +61,7 @@ func DataObjectAttributeUpdateHandler(c *gin.Context) { } }() + redisChanges := NewRedisChangeSet(diagram.GetRedisClientInstance()) message := "data-object attribute update success" var measurementResult measurementUpdateResult switch dataObjectType { @@ -69,14 +70,24 @@ func DataObjectAttributeUpdateHandler(c *gin.Context) { if queryErr == nil { queryErr = database.UpdateParameterDataObjectValue(ctx, tx, parameter, value) } + if queryErr == nil { + queryErr = redisChanges.AddDataObjectHashChange(ctx, dataObjectType, request.Token, field, value) + } err = queryErr case constants.DataObjectTypeMeasurement: measurementResult, err = updateMeasurementDataObject(ctx, tx, request.Token, field, value, request.Data, measurementUpdateDependencies{ - writeManualValueFunc: writeMeasurementManualValue, - updateDataRTFunc: callRealTimeDataWriteStopInterface, - startDataRTFunc: callRealTimeDataWriteStartInterface, - replaceRedisValueFunc: replaceMeasurementRedisValue, + writeManualValueFunc: func(ctx context.Context, measurement *orm.Measurement, value float64) error { + return redisChanges.AddMeasurementValueChange(ctx, measurement, value, false) + }, + updateDataRTFunc: callRealTimeDataWriteStopInterface, + startDataRTFunc: callRealTimeDataWriteStartInterface, + replaceRedisValueFunc: func(ctx context.Context, measurement *orm.Measurement, value float64) error { + return redisChanges.AddMeasurementValueChange(ctx, measurement, value, true) + }, }) + if err == nil && measurementResult.modeChanged { + err = redisChanges.AddDataObjectHashChange(ctx, dataObjectType, request.Token, "mode", measurementResult.mode) + } message = measurementResult.message default: err = fmt.Errorf("unsupported data object type %q", dataObjectType) @@ -84,7 +95,7 @@ func DataObjectAttributeUpdateHandler(c *gin.Context) { if err != nil { _ = tx.Rollback().Error - if measurementResult.recordFailure { + if measurementResult.recordFailureOnError { if logErr := database.AppendMeasurementValueOperation(ctx, database.GetPostgresDBClient(), measurementResult.measurementID, 1, measurementResult.value, time.Now().UTC()); logErr != nil { logger.Error(ctx, "append failed measurement value operation failed", "measurement_id", measurementResult.measurementID, "error", logErr) } @@ -98,7 +109,29 @@ func DataObjectAttributeUpdateHandler(c *gin.Context) { return } + if err := redisChanges.Apply(ctx); err != nil { + _ = tx.Rollback().Error + if measurementResult.recordFailureOnError { + if logErr := database.AppendMeasurementValueOperation(ctx, database.GetPostgresDBClient(), measurementResult.measurementID, 1, measurementResult.value, time.Now().UTC()); logErr != nil { + logger.Error(ctx, "append failed measurement value operation failed", "measurement_id", measurementResult.measurementID, "error", logErr) + } + } + logger.Error(ctx, "apply redis data-object changes failed", "token", request.Token, "field", field, "error", err) + renderRespFailure(c, constants.RespCodeFailed, err.Error(), nil) + return + } + if err := tx.Commit().Error; err != nil { + revertCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), redisChangeRestoreTimeout) + defer cancel() + if redisErr := redisChanges.Revert(revertCtx); redisErr != nil { + logger.Error(ctx, "revert redis data-object changes failed", "token", request.Token, "field", field, "error", redisErr) + } + if measurementResult.recordFailureOnError { + if logErr := database.AppendMeasurementValueOperation(ctx, database.GetPostgresDBClient(), measurementResult.measurementID, 1, measurementResult.value, time.Now().UTC()); logErr != nil { + logger.Error(ctx, "append failed measurement value operation failed", "measurement_id", measurementResult.measurementID, "error", logErr) + } + } logger.Error(ctx, "commit data-object update transaction failed", "token", request.Token, "field", field, "error", err) renderRespFailure(c, constants.RespCodeServerError, "transaction commit failed", nil) return @@ -233,10 +266,12 @@ type measurementUpdateDependencies struct { } type measurementUpdateResult struct { - message string - measurementID int64 - value float64 - recordFailure bool + message string + measurementID int64 + value float64 + recordFailureOnError bool + mode int16 + modeChanged bool } func updateMeasurementDataObject( @@ -305,7 +340,11 @@ func updateMeasurementDataObject( 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 + return measurementUpdateResult{ + message: fmt.Sprintf("measurement mode changed to %s", measurementModeName(mode)), + mode: mode, + modeChanged: true, + }, nil case "value": currentMode, err := measurementModeIsAutomatic(lockedMeasurement.Mode) if err != nil { @@ -319,9 +358,9 @@ func updateMeasurementDataObject( return measurementUpdateResult{}, fmt.Errorf("measurement value has invalid type %T", value) } failureResult := measurementUpdateResult{ - measurementID: lockedMeasurement.ID, - value: manualValue, - recordFailure: true, + measurementID: lockedMeasurement.ID, + value: manualValue, + recordFailureOnError: true, } if dependencies.writeManualValueFunc == nil { return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(fmt.Errorf("measurement manual value writer is nil")) @@ -338,7 +377,8 @@ func updateMeasurementDataObject( if err := database.AppendMeasurementValueOperation(ctx, tx, lockedMeasurement.ID, 0, manualValue, time.Now().UTC()); err != nil { return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(err) } - return measurementUpdateResult{message: "measurement manual value updated"}, nil + failureResult.message = "measurement manual value updated" + return failureResult, nil default: return measurementUpdateResult{}, fmt.Errorf("unsupported measurement update field %q", field) } @@ -383,21 +423,6 @@ func isInvalidDataObjectUpdateError(err error) bool { errors.Is(err, common.ErrAmbiguousMeasurementToken) } -func writeMeasurementManualValue(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.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 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. @@ -408,18 +433,3 @@ func callRealTimeDataWriteStartInterface(_ context.Context, _ orm.JSONMap, _ *fl // 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 03bb246..9f1e2b9 100644 --- a/handler/data_object_attribute_update_test.go +++ b/handler/data_object_attribute_update_test.go @@ -26,6 +26,7 @@ func TestValidateDataObjectAttributeUpdateParameterGroups(t *testing.T) { "craft", "integrity", "behavior", + "base_extend", } for _, group := range groups { @@ -46,7 +47,7 @@ func TestValidateDataObjectAttributeUpdateParameterGroups(t *testing.T) { } func TestValidateDataObjectAttributeUpdateRejectsUnsupportedParameterGroups(t *testing.T) { - for _, group := range []string{"component", "base_extend"} { + for _, group := range []string{"component"} { t.Run(group, func(t *testing.T) { _, _, _, err := validateDataObjectAttributeUpdate(dataObjectAttributeUpdateRequest{ Token: fmt.Sprintf("nspath.component.%s.attribute", group), @@ -416,7 +417,9 @@ func TestUpdateMeasurementDataObjectWritesValueInManualMode(t *testing.T) { assert.True(t, called) assert.True(t, dataRTCalled) assert.Contains(t, result.message, "updated") - assert.False(t, result.recordFailure) + assert.True(t, result.recordFailureOnError) + assert.Equal(t, int64(10), result.measurementID) + assert.Equal(t, float64(15.2), result.value) require.NoError(t, tx.Rollback().Error) require.NoError(t, mock.ExpectationsWereMet()) } @@ -439,7 +442,7 @@ func TestUpdateMeasurementDataObjectReturnsFailureResultAndAppError(t *testing.T require.Error(t, err) assert.ErrorIs(t, err, errcode.ErrMeasurementValueUpdateFailed) assert.ErrorIs(t, err, writeErr) - assert.True(t, result.recordFailure) + assert.True(t, result.recordFailureOnError) assert.Equal(t, int64(10), result.measurementID) assert.Equal(t, float64(15.2), result.value) require.NoError(t, tx.Rollback().Error) diff --git a/handler/data_object_redis_change.go b/handler/data_object_redis_change.go new file mode 100644 index 0000000..f8d1f3b --- /dev/null +++ b/handler/data_object_redis_change.go @@ -0,0 +1,429 @@ +package handler + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "sort" + "strconv" + "strings" + "time" + + "modelRT/constants" + "modelRT/model" + "modelRT/orm" + + "github.com/redis/go-redis/v9" +) + +const redisChangeRestoreTimeout = 5 * time.Second + +type redisHashChange struct { + key string + field string + oldValue string + newValue string +} + +type redisZSetChange struct { + key string + oldValues []redis.Z + newValues []redis.Z +} + +// RedisChangeSet keeps the Redis changes belonging to one PostgreSQL +// transaction. Changes are prepared first and applied together immediately +// before the PostgreSQL transaction is committed. +type RedisChangeSet struct { + client *redis.Client + hashChanges []redisHashChange + zsetChanges []redisZSetChange + applied bool +} + +func NewRedisChangeSet(client *redis.Client) *RedisChangeSet { + return &RedisChangeSet{client: client} +} + +func (changes *RedisChangeSet) AddDataObjectHashChange( + ctx context.Context, + dataObjectType constants.DataObjectType, + token, field string, + value any, +) error { + if changes == nil || changes.client == nil { + return fmt.Errorf("redis client is not initialized") + } + + metadata, err := changes.client.HMGet(ctx, token, "id", "name").Result() + if err != nil { + return fmt.Errorf("query redis data-object aliases for %q: %w", token, err) + } + if len(metadata) != 2 || metadata[0] == nil || metadata[1] == nil { + return fmt.Errorf("redis data-object hash %q does not contain id and name", token) + } + + id, ok := metadata[0].(string) + if !ok { + return fmt.Errorf("redis data-object hash %q id has type %T", token, metadata[0]) + } + name, ok := metadata[1].(string) + if !ok { + return fmt.Errorf("redis data-object hash %q name has type %T", token, metadata[1]) + } + keys, err := dataObjectRedisAliasKeys(dataObjectType, id, name) + if err != nil { + return err + } + if !containsString(keys, token) { + return fmt.Errorf("redis data-object hash %q metadata points to different aliases", token) + } + + newValue, err := redisChangeString(value) + if err != nil { + return fmt.Errorf("encode redis data-object value: %w", err) + } + for _, key := range keys { + keyType, err := changes.client.Type(ctx, key).Result() + if err != nil { + return fmt.Errorf("query redis key type for %q: %w", key, err) + } + if keyType == "none" && dataObjectType == constants.DataObjectTypeParameter && key == name && key != token { + // A parameter short alias is only initialized for a local station. + continue + } + if keyType != "hash" { + return fmt.Errorf("redis data-object key %q has type %q, expected hash", key, keyType) + } + oldValue, err := changes.client.HGet(ctx, key, field).Result() + if errors.Is(err, redis.Nil) { + return fmt.Errorf("redis data-object hash %q does not contain field %q", key, field) + } + if err != nil { + return fmt.Errorf("query redis hash %q field %q: %w", key, field, err) + } + changes.hashChanges = append(changes.hashChanges, redisHashChange{ + key: key, + field: field, + oldValue: oldValue, + newValue: newValue, + }) + } + return nil +} + +func (changes *RedisChangeSet) AddMeasurementValueChange( + ctx context.Context, + measurement *orm.Measurement, + value float64, + replace bool, +) error { + if changes == nil || changes.client == nil { + return fmt.Errorf("redis client is not initialized") + } + if measurement == nil { + return fmt.Errorf("measurement is nil") + } + key, err := model.GenerateMeasureIdentifier(measurement.DataSource) + if err != nil { + return fmt.Errorf("generate measurement redis key: %w", err) + } + keyType, err := changes.client.Type(ctx, key).Result() + if err != nil { + return fmt.Errorf("query measurement redis key type for %q: %w", key, err) + } + if keyType != "none" && keyType != "zset" { + return fmt.Errorf("measurement redis key %q has type %q, expected zset", key, keyType) + } + oldValues, err := changes.client.ZRangeWithScores(ctx, key, 0, -1).Result() + if err != nil { + return fmt.Errorf("query measurement redis values for %q: %w", key, err) + } + + newMember := strconv.FormatInt(time.Now().UnixNano(), 10) + newValues := []redis.Z{{Score: value, Member: newMember}} + if !replace { + newValues = mergeRedisZValues(oldValues, newValues...) + } + changes.zsetChanges = append(changes.zsetChanges, redisZSetChange{ + key: key, + oldValues: normalizeRedisZValues(oldValues), + newValues: normalizeRedisZValues(newValues), + }) + return nil +} + +func (changes *RedisChangeSet) Apply(ctx context.Context) error { + if changes == nil || changes.client == nil { + return fmt.Errorf("redis client is not initialized") + } + if len(changes.hashChanges) == 0 && len(changes.zsetChanges) == 0 { + return nil + } + + keys := changes.keys() + err := changes.client.Watch(ctx, func(tx *redis.Tx) error { + if err := changes.verify(ctx, tx, false); err != nil { + return err + } + _, err := tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + for _, change := range changes.hashChanges { + pipe.HSet(ctx, change.key, change.field, change.newValue) + } + for _, change := range changes.zsetChanges { + pipe.Del(ctx, change.key) + if len(change.newValues) > 0 { + pipe.ZAdd(ctx, change.key, change.newValues...) + } + } + return nil + }) + return err + }, keys...) + if err != nil { + // A connection error can leave EXEC's outcome unknown. Restore only + // when Redis still contains either the prepared or the applied state. + restoreCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), redisChangeRestoreTimeout) + defer cancel() + if restoreErr := changes.restoreAfterApplyFailure(restoreCtx); restoreErr != nil { + return fmt.Errorf("apply redis changes: %w; restore redis changes: %v", err, restoreErr) + } + return fmt.Errorf("apply redis changes: %w", err) + } + changes.applied = true + return nil +} + +// Revert restores Redis after a PostgreSQL commit failure. It uses WATCH and +// only restores values that still match this change set. +func (changes *RedisChangeSet) Revert(ctx context.Context) error { + if changes == nil || !changes.applied { + return nil + } + return changes.restoreOldValues(ctx, true) +} + +func (changes *RedisChangeSet) restoreOldValues(ctx context.Context, compareNew bool) error { + keys := changes.keys() + return changes.client.Watch(ctx, func(tx *redis.Tx) error { + if compareNew { + if err := changes.verify(ctx, tx, true); err != nil { + return err + } + } + _, err := tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + for _, change := range changes.hashChanges { + pipe.HSet(ctx, change.key, change.field, change.oldValue) + } + for _, change := range changes.zsetChanges { + pipe.Del(ctx, change.key) + if len(change.oldValues) > 0 { + pipe.ZAdd(ctx, change.key, change.oldValues...) + } + } + return nil + }) + return err + }, keys...) +} + +func (changes *RedisChangeSet) restoreAfterApplyFailure(ctx context.Context) error { + keys := changes.keys() + return changes.client.Watch(ctx, func(tx *redis.Tx) error { + for _, change := range changes.hashChanges { + actual, err := tx.HGet(ctx, change.key, change.field).Result() + if err != nil { + return fmt.Errorf("verify redis hash %q field %q after apply failure: %w", change.key, change.field, err) + } + if actual != change.oldValue && actual != change.newValue { + return fmt.Errorf("redis hash %q field %q changed concurrently", change.key, change.field) + } + } + for _, change := range changes.zsetChanges { + actual, err := tx.ZRangeWithScores(ctx, change.key, 0, -1).Result() + if err != nil { + return fmt.Errorf("verify redis zset %q after apply failure: %w", change.key, err) + } + if !equalRedisZValues(actual, change.oldValues) && !equalRedisZValues(actual, change.newValues) { + return fmt.Errorf("redis zset %q changed concurrently", change.key) + } + } + _, err := tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + for _, change := range changes.hashChanges { + pipe.HSet(ctx, change.key, change.field, change.oldValue) + } + for _, change := range changes.zsetChanges { + pipe.Del(ctx, change.key) + if len(change.oldValues) > 0 { + pipe.ZAdd(ctx, change.key, change.oldValues...) + } + } + return nil + }) + return err + }, keys...) +} + +func (changes *RedisChangeSet) verify(ctx context.Context, tx *redis.Tx, expectNew bool) error { + for _, change := range changes.hashChanges { + expected := change.oldValue + if expectNew { + expected = change.newValue + } + actual, err := tx.HGet(ctx, change.key, change.field).Result() + if err != nil { + return fmt.Errorf("verify redis hash %q field %q: %w", change.key, change.field, err) + } + if actual != expected { + return fmt.Errorf("redis hash %q field %q changed concurrently", change.key, change.field) + } + } + for _, change := range changes.zsetChanges { + expected := change.oldValues + if expectNew { + expected = change.newValues + } + actual, err := tx.ZRangeWithScores(ctx, change.key, 0, -1).Result() + if err != nil { + return fmt.Errorf("verify redis zset %q: %w", change.key, err) + } + if !equalRedisZValues(actual, expected) { + return fmt.Errorf("redis zset %q changed concurrently", change.key) + } + } + return nil +} + +func (changes *RedisChangeSet) keys() []string { + seen := make(map[string]struct{}, len(changes.hashChanges)+len(changes.zsetChanges)) + keys := make([]string, 0, len(seen)) + for _, change := range changes.hashChanges { + if _, ok := seen[change.key]; !ok { + seen[change.key] = struct{}{} + keys = append(keys, change.key) + } + } + for _, change := range changes.zsetChanges { + if _, ok := seen[change.key]; !ok { + seen[change.key] = struct{}{} + keys = append(keys, change.key) + } + } + sort.Strings(keys) + return keys +} + +func dataObjectRedisAliasKeys(dataObjectType constants.DataObjectType, id, name string) ([]string, error) { + switch dataObjectType { + case constants.DataObjectTypeParameter: + if len(strings.Split(id, ".")) != 7 || len(strings.Split(name, ".")) != 4 { + return nil, fmt.Errorf("invalid parameter redis aliases id=%q name=%q", id, name) + } + return uniqueStrings(id, name), nil + case constants.DataObjectTypeMeasurement: + parts := strings.Split(id, ".") + if len(parts) != 7 || len(strings.Split(name, ".")) != 2 { + return nil, fmt.Errorf("invalid measurement redis aliases id=%q name=%q", id, name) + } + return uniqueStrings(id, strings.Join(parts[3:], "."), name), nil + default: + return nil, fmt.Errorf("unsupported data object type %q", dataObjectType) + } +} + +func redisChangeString(value any) (string, error) { + switch typedValue := value.(type) { + case string: + return typedValue, nil + case []byte: + return string(typedValue), nil + case nil: + return "null", nil + case bool: + return strconv.FormatBool(typedValue), nil + case int: + return strconv.Itoa(typedValue), nil + case int16: + return strconv.FormatInt(int64(typedValue), 10), nil + case int64: + return strconv.FormatInt(typedValue, 10), nil + case float64: + return strconv.FormatFloat(typedValue, 'f', -1, 64), nil + default: + encoded, err := json.Marshal(typedValue) + if err != nil { + return "", err + } + return string(encoded), nil + } +} + +func cloneRedisZValues(values []redis.Z) []redis.Z { + cloned := make([]redis.Z, len(values)) + copy(cloned, values) + return cloned +} + +func mergeRedisZValues(current []redis.Z, additions ...redis.Z) []redis.Z { + valuesByMember := make(map[string]redis.Z, len(current)+len(additions)) + for _, value := range current { + valuesByMember[fmt.Sprint(value.Member)] = value + } + for _, value := range additions { + valuesByMember[fmt.Sprint(value.Member)] = value + } + values := make([]redis.Z, 0, len(valuesByMember)) + for _, value := range valuesByMember { + values = append(values, value) + } + return normalizeRedisZValues(values) +} + +func normalizeRedisZValues(values []redis.Z) []redis.Z { + normalized := cloneRedisZValues(values) + sort.Slice(normalized, func(i, j int) bool { + if normalized[i].Score != normalized[j].Score { + return normalized[i].Score < normalized[j].Score + } + return fmt.Sprint(normalized[i].Member) < fmt.Sprint(normalized[j].Member) + }) + return normalized +} + +func equalRedisZValues(left, right []redis.Z) bool { + left = normalizeRedisZValues(left) + right = normalizeRedisZValues(right) + if len(left) != len(right) { + return false + } + for index := range left { + if left[index].Score != right[index].Score || + fmt.Sprint(left[index].Member) != fmt.Sprint(right[index].Member) { + return false + } + } + return true +} + +func uniqueStrings(values ...string) []string { + seen := make(map[string]struct{}, len(values)) + result := make([]string, 0, len(values)) + for _, value := range values { + if _, ok := seen[value]; ok { + continue + } + seen[value] = struct{}{} + result = append(result, value) + } + return result +} + +func containsString(values []string, target string) bool { + for _, value := range values { + if value == target { + return true + } + } + return false +} diff --git a/handler/data_object_redis_change_test.go b/handler/data_object_redis_change_test.go new file mode 100644 index 0000000..3961e2e --- /dev/null +++ b/handler/data_object_redis_change_test.go @@ -0,0 +1,79 @@ +package handler + +import ( + "testing" + + "modelRT/constants" + + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDataObjectRedisAliasKeys(t *testing.T) { + parameterKeys, err := dataObjectRedisAliasKeys( + constants.DataObjectTypeParameter, + "grid.zone.station.nspath.component.rated.voltage", + "nspath.component.rated.voltage", + ) + require.NoError(t, err) + assert.Equal(t, []string{ + "grid.zone.station.nspath.component.rated.voltage", + "nspath.component.rated.voltage", + }, parameterKeys) + + measurementKeys, err := dataObjectRedisAliasKeys( + constants.DataObjectTypeMeasurement, + "grid.zone.station.nspath.component.bay.current", + "nspath.current", + ) + require.NoError(t, err) + assert.Equal(t, []string{ + "grid.zone.station.nspath.component.bay.current", + "nspath.component.bay.current", + "nspath.current", + }, measurementKeys) +} + +func TestDataObjectRedisAliasKeysRejectsInvalidMetadata(t *testing.T) { + _, err := dataObjectRedisAliasKeys(constants.DataObjectTypeParameter, "short.id", "short.name") + require.Error(t, err) + + _, err = dataObjectRedisAliasKeys(constants.DataObjectTypeMeasurement, "short.id", "short.name") + require.Error(t, err) +} + +func TestRedisChangeStringUsesCacheRepresentations(t *testing.T) { + tests := []struct { + name string + value any + want string + }{ + {name: "string", value: "15.2", want: "15.2"}, + {name: "integer", value: int64(15), want: "15"}, + {name: "decimal", value: 15.2, want: "15.2"}, + {name: "boolean", value: true, want: "true"}, + {name: "object", value: map[string]any{"enabled": true}, want: `{"enabled":true}`}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got, err := redisChangeString(test.value) + require.NoError(t, err) + assert.Equal(t, test.want, got) + }) + } +} + +func TestMergeRedisZValuesReplacesDuplicateMemberAndSorts(t *testing.T) { + values := mergeRedisZValues( + []redis.Z{ + {Score: 20, Member: "2"}, + {Score: 10, Member: "1"}, + }, + redis.Z{Score: 30, Member: "1"}, + ) + assert.Equal(t, []redis.Z{ + {Score: 20, Member: "2"}, + {Score: 30, Member: "1"}, + }, values) +} diff --git a/model/measurement_data_object_init.go b/model/measurement_data_object_init.go index 04a987c..e0823bb 100644 --- a/model/measurement_data_object_init.go +++ b/model/measurement_data_object_init.go @@ -73,6 +73,13 @@ func buildMeasurementDataObjectHashes(records []MeasurementInitializationRecord) record.MeasurementMode, ) } + if record.MeasurementSize <= 0 { + return nil, fmt.Errorf( + "measurement %q window size must be greater than 0, got %d", + record.MeasurementTag, + record.MeasurementSize, + ) + } if record.MeasurementDataSource == nil || record.MeasurementEventPlan == nil || record.MeasurementBinding == nil { diff --git a/model/measurement_data_object_init_test.go b/model/measurement_data_object_init_test.go index 549ef74..e236c30 100644 --- a/model/measurement_data_object_init_test.go +++ b/model/measurement_data_object_init_test.go @@ -65,6 +65,15 @@ func TestBuildMeasurementDataObjectHashesRejectsUnsupportedMeasurementType(t *te assert.Contains(t, err.Error(), "unsupported measurement type 5") } +func TestBuildMeasurementDataObjectHashesRejectsInvalidWindowSize(t *testing.T) { + record := measurementInitializationRecordForTest() + record.MeasurementSize = 0 + + _, err := buildMeasurementDataObjectHashes([]MeasurementInitializationRecord{record}) + require.Error(t, err) + assert.Contains(t, err.Error(), "window size must be greater than 0") +} + func measurementInitializationRecordForTest() MeasurementInitializationRecord { return MeasurementInitializationRecord{ GridTag: "grid000",