diff --git a/common/data_object_errors.go b/common/data_object_errors.go new file mode 100644 index 0000000..afbb7e0 --- /dev/null +++ b/common/data_object_errors.go @@ -0,0 +1,23 @@ +// Package common define common error variables +package common + +import "errors" + +var ( + // ErrUnsupportedParameterField indicates that a requested parameter field is not supported. + ErrUnsupportedParameterField = errors.New("unsupported parameter field") + // ErrInvalidParameterToken indicates that a token cannot represent a parameter. + ErrInvalidParameterToken = errors.New("invalid parameter token") + // ErrParameterTokenNotFound indicates that no parameter matches the token hierarchy. + ErrParameterTokenNotFound = errors.New("parameter token not found") + // ErrAmbiguousParameterToken indicates that a token matches more than one parameter. + ErrAmbiguousParameterToken = errors.New("ambiguous parameter token") + // ErrUnsupportedMeasurementField define error of unsupport measurement field + ErrUnsupportedMeasurementField = errors.New("unsupported measurement field") + // ErrInvalidMeasurementToken indicates that a token cannot represent a measurement. + ErrInvalidMeasurementToken = errors.New("invalid measurement token") + // ErrMeasurementTokenNotFound indicates that no measurement matches the token hierarchy. + ErrMeasurementTokenNotFound = errors.New("measurement token not found") + // ErrAmbiguousMeasurementToken indicates that a token matches more than one measurement. + ErrAmbiguousMeasurementToken = errors.New("ambiguous measurement token") +) diff --git a/common/errcode/bussiness_error.go b/common/errcode/bussiness_error.go index 6ab5999..f84cbe1 100644 --- a/common/errcode/bussiness_error.go +++ b/common/errcode/bussiness_error.go @@ -38,6 +38,9 @@ var ( // ErrCommitTxFailed indicates that the PostgreSQL transaction could not be committed successfully. ErrCommitTxFailed = newError(50005, "postgres database transaction commit failed") + // ErrMeasurementValueUpdateFailed indicates that a manual measurement value transaction failed. + ErrMeasurementValueUpdateFailed = newError(50006, "measurement manual value update failed") + // ErrCachedQueryFailed define variable to indicates an error occurred while attempting to fetch data from the Redis cache. ErrCachedQueryFailed = newError(60001, "query redis cached data failed") diff --git a/common/errcode/error.go b/common/errcode/error.go index 1d12276..ec0941e 100644 --- a/common/errcode/error.go +++ b/common/errcode/error.go @@ -66,8 +66,8 @@ func Wrap(msg string, err error) *AppError { return appErr } -// UnWrap define func return the error wrapped in structure -func (e *AppError) UnWrap() error { +// Unwrap returns the underlying cause for errors.Is and errors.As traversal. +func (e *AppError) Unwrap() error { return e.cause } diff --git a/constants/context.go b/constants/context.go index dcac3c3..cf3ff99 100644 --- a/constants/context.go +++ b/constants/context.go @@ -1,7 +1,13 @@ // Package constants define constant variable package constants +// ClientTokenContextName is the Gin key used for the configured client token. +const ClientTokenContextName = "client_token" + type contextKey string // MeasurementUUIDKey define measurement uuid key into context const MeasurementUUIDKey contextKey = "measurement_uuid" + +// CtxKeyClientToken is the typed standard-library context key for client token propagation. +const CtxKeyClientToken contextKey = ClientTokenContextName diff --git a/constants/data-object.go b/constants/data-object.go new file mode 100644 index 0000000..f57158e --- /dev/null +++ b/constants/data-object.go @@ -0,0 +1,19 @@ +// Package constants define constant variable +package constants + +// DataObjectType identifies the kind of object represented by a data object token. +type DataObjectType string + +const ( + // DataObjectTypeParameter represents a component parameter. + DataObjectTypeParameter DataObjectType = "parameter" + // DataObjectTypeMeasurement represents a component measurement. + DataObjectTypeMeasurement DataObjectType = "measurement" +) + +const ( + // MeasurementModeManual indicates that manual value entry is enabled. + MeasurementModeManual int16 = 0 + // MeasurementModeAutomatic indicates that the measurement runs automatically. + MeasurementModeAutomatic int16 = 1 +) diff --git a/constants/parameter_table.go b/constants/parameter_table.go new file mode 100644 index 0000000..bcbb12c --- /dev/null +++ b/constants/parameter_table.go @@ -0,0 +1,26 @@ +// Package constants define constant variable +package constants + +import "strings" + +var supportedParameterTableSuffixes = [...]string{ + "base_extend", + "rated", + "setup", + "model", + "stable", + "craft", + "integrity", + "behavior", +} + +// IsSupportedParameterTableName reports whether a dynamic parameter table has +// one of the supported attribute-group suffixes. +func IsSupportedParameterTableName(tableName string) bool { + for _, suffix := range supportedParameterTableSuffixes { + if strings.HasSuffix(tableName, "_"+suffix) { + return true + } + } + return false +} diff --git a/constants/parameter_table_test.go b/constants/parameter_table_test.go new file mode 100644 index 0000000..7277978 --- /dev/null +++ b/constants/parameter_table_test.go @@ -0,0 +1,25 @@ +package constants + +import "testing" + +func TestIsSupportedParameterTableName(t *testing.T) { + tests := []struct { + name string + tableName string + want bool + }{ + {name: "bay table is excluded", tableName: "ct_ct_demo_bay", want: false}, + {name: "model table is included", tableName: "cable_cable_demo_model", want: true}, + {name: "base extend table is included", tableName: "cable_cable_demo_base_extend", want: true}, + {name: "suffix must start at separator", tableName: "cable_cable_demomodel", want: false}, + {name: "empty table name", tableName: "", want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := IsSupportedParameterTableName(tt.tableName); got != tt.want { + t.Fatalf("IsSupportedParameterTableName(%q) = %v, want %v", tt.tableName, got, tt.want) + } + }) + } +} diff --git a/database/query_bay_columns.go b/database/query_bay_columns.go new file mode 100644 index 0000000..5a5e8dc --- /dev/null +++ b/database/query_bay_columns.go @@ -0,0 +1,28 @@ +package database + +import ( + "context" + "strings" + + "modelRT/orm" + + "gorm.io/gorm" +) + +// QueryBayDevColumnNames returns the bay table columns exposed as token7 +// candidates under token6=bay. +func QueryBayDevColumnNames(ctx context.Context, db *gorm.DB) ([]string, error) { + columnTypes, err := db.WithContext(ctx).Migrator().ColumnTypes((&orm.Bay{}).TableName()) + if err != nil { + return nil, err + } + + columnNames := make([]string, 0, len(columnTypes)) + for _, columnType := range columnTypes { + columnName := columnType.Name() + if strings.HasPrefix(columnName, "dev_") { + columnNames = append(columnNames, columnName) + } + } + return columnNames, nil +} diff --git a/database/query_component_measurement.go b/database/query_component_measurement.go index ffdba8a..74a91ac 100644 --- a/database/query_component_measurement.go +++ b/database/query_component_measurement.go @@ -3,24 +3,270 @@ package database import ( "context" + "encoding/json" "fmt" + "strings" + "time" + "modelRT/common" + "modelRT/constants" "modelRT/orm" + "modelRT/sql" "golang.org/x/sync/errgroup" "gorm.io/gorm" + "gorm.io/gorm/clause" ) -type ZoneWithParent struct { - orm.Zone - GridTag string `gorm:"column:grid_tag"` +const ( + measurementOperationsLimit = 500 + measurementOperationAppendSQL = "(array_append(operations, ?::jsonb))[GREATEST(cardinality(operations) - ? + 2, 1):]" +) + +// 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 + cancelCtx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + result := db.WithContext(cancelCtx). + Where(sql.MeasurementIDWhere, id). + Take(&measurement) + + if result.Error != nil { + return orm.Measurement{}, fmt.Errorf("query measurement %d: %w", id, result.Error) + } + return measurement, nil } -type StationWithParent struct { - orm.Zone - ZoneTag string `gorm:"column:zone_tag"` +// 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 } +// QueryMeasurementByToken define function query circuit diagram component measurement info by token from postgresDB +func QueryMeasurementByToken(ctx context.Context, tx *gorm.DB, token string) (orm.Measurement, error) { + measurement, _, err := QueryMeasurementByDataObjectToken(ctx, tx, token) + if err != nil { + return orm.Measurement{}, err + } + return *measurement, nil +} + +// 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 := constants.MeasurementModeManual + if automatic { + mode = constants.MeasurementModeAutomatic + } + + result := db.WithContext(ctx). + Model(&orm.Measurement{}). + Where("id = ?", measurementID). + Update("mode", mode) + if result.Error != nil { + return fmt.Errorf("update measurement %d mode: %w", measurementID, result.Error) + } + if result.RowsAffected == 0 { + return fmt.Errorf("update measurement %d mode affected no rows", measurementID) + } + return nil +} + +// 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, 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.UnixMilli(), + } + return updateMeasurementWithOperation(ctx, db, measurementID, map[string]any{"mode": mode}, operation) +} + +// AppendMeasurementValueOperation appends the audit result of a manual-value +// transaction without changing other measurement columns. +func AppendMeasurementValueOperation(ctx context.Context, db *gorm.DB, measurementID int64, transaction int, value float64, timestamp time.Time) error { + operation := orm.JSONMap{ + "transaction": transaction, + "value": value, + "timestamp": timestamp.UnixMilli(), + } + return updateMeasurementWithOperation(ctx, db, measurementID, nil, operation) +} + +func updateMeasurementWithOperation(ctx context.Context, db *gorm.DB, measurementID int64, updates map[string]any, operation orm.JSONMap) error { + encodedOperation, err := json.Marshal(operation) + if err != nil { + return fmt.Errorf("encode measurement %d operation: %w", measurementID, err) + } + + operationExpression := gorm.Expr( + measurementOperationAppendSQL, + string(encodedOperation), + measurementOperationsLimit, + ) + if updates == nil { + updates = make(map[string]any, 1) + } + updates["operations"] = operationExpression + + result := db.WithContext(ctx). + Model(&orm.Measurement{}). + Where("id = ?", measurementID). + Updates(updates) + if result.Error != nil { + return fmt.Errorf("update measurement %d operation: %w", measurementID, result.Error) + } + if result.RowsAffected == 0 { + return fmt.Errorf("update measurement %d operation affected no rows", measurementID) + } + return nil +} + +// ValidateMeasurementToken checks whether token uniquely identifies an existing +// measurement through the measurement, component, bay, station, zone, and grid +// relationships. Supported formats are token1-token7, token4-token7, and +// token4.token7. +func ValidateMeasurementToken(ctx context.Context, db *gorm.DB, token string) error { + query, args, err := buildMeasurementTokenValidationQuery(token) + if err != nil { + return err + } + + var count int64 + if err := db.WithContext(ctx).Raw(query, args...).Scan(&count).Error; err != nil { + return fmt.Errorf("query measurement token %q: %w", token, err) + } + + switch { + case count == 0: + return fmt.Errorf("%w: %q", common.ErrMeasurementTokenNotFound, token) + case count > 1: + return fmt.Errorf("%w: %q matched %d records", common.ErrAmbiguousMeasurementToken, token, count) + default: + return nil + } +} + +// QueryMeasurementByDataObjectToken validates token and returns the existing +// measurement and its owning component for attribute response construction. +func QueryMeasurementByDataObjectToken(ctx context.Context, db *gorm.DB, token string) (*orm.Measurement, *orm.Component, error) { + validationQuery, args, err := buildMeasurementTokenValidationQuery(token) + if err != nil { + return nil, nil, err + } + + query := buildMeasurementRowsQuery(validationQuery) + var rows []orm.Measurement + if err := db.WithContext(ctx).Raw(query, args...).Scan(&rows).Error; err != nil { + return nil, nil, fmt.Errorf("query measurement token %q: %w", token, err) + } + + switch len(rows) { + case 0: + return nil, nil, fmt.Errorf("%w: %q", common.ErrMeasurementTokenNotFound, token) + case 1: + // Continue by loading the owning component. + default: + return nil, nil, fmt.Errorf("%w: %q matched more than one record", common.ErrAmbiguousMeasurementToken, token) + } + + var component orm.Component + result := db.WithContext(ctx). + Raw(compactMeasurementSQL(sql.MeasurementComponentByUUID), rows[0].ComponentUUID). + Scan(&component) + if result.Error != nil { + return nil, nil, fmt.Errorf("query component for measurement token %q: %w", token, result.Error) + } + if result.RowsAffected == 0 { + return nil, nil, fmt.Errorf("%w: component for %q", common.ErrMeasurementTokenNotFound, token) + } + + return &rows[0], &component, nil +} + +func buildMeasurementRowsQuery(validationQuery string) string { + measurementQuery := strings.Replace( + validationQuery, + sql.MeasurementCountSelect, + sql.MeasurementRowsSelect, + 1, + ) + return compactMeasurementSQL(strings.Join([]string{measurementQuery, sql.MeasurementLimitTwo}, "\n")) +} + +func compactMeasurementSQL(statement string) string { + return strings.Join(strings.Fields(statement), " ") +} + +func buildMeasurementTokenValidationQuery(token string) (string, []any, error) { + parts := strings.Split(token, ".") + for _, part := range parts { + if part == "" { + return "", nil, fmt.Errorf("%w %q: token segment cannot be empty", common.ErrInvalidMeasurementToken, token) + } + } + + switch len(parts) { + case 7: + if parts[5] != "bay" { + return "", nil, fmt.Errorf("%w %q: token6 must be bay", common.ErrInvalidMeasurementToken, token) + } + query := compactMeasurementSQL(strings.Join([]string{ + sql.MeasurementTokenValidationQueryBase, + sql.MeasurementSevenPartTokenWhere, + }, "\n")) + return query, []any{parts[0], parts[1], parts[2], parts[3], parts[4], parts[6]}, nil + case 4: + if parts[2] != "bay" { + return "", nil, fmt.Errorf("%w %q: token6 must be bay", common.ErrInvalidMeasurementToken, token) + } + query := compactMeasurementSQL(strings.Join([]string{ + sql.MeasurementTokenValidationQueryBase, + sql.MeasurementFourPartTokenWhere, + }, "\n")) + return query, []any{parts[0], parts[1], parts[3]}, nil + case 2: + query := compactMeasurementSQL(strings.Join([]string{ + sql.MeasurementTokenValidationQueryBase, + sql.MeasurementTwoPartTokenWhere, + }, "\n")) + return query, []any{parts[0], parts[1]}, nil + default: + return "", nil, fmt.Errorf("%w %q: expected 2, 4, or 7 segments, got %d", common.ErrInvalidMeasurementToken, token, len(parts)) + } +} + +// GetAllMeasurements define func to query all measurement info from postgresDB +func GetAllMeasurements(ctx context.Context, tx *gorm.DB) ([]orm.Measurement, error) { + var measurements []orm.Measurement + + // ctx超时判断 + cancelCtx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + result := tx.WithContext(cancelCtx).Clauses(clause.Locking{Strength: "UPDATE"}).Find(&measurements) + if result.Error != nil { + return nil, result.Error + } + return measurements, nil +} + +// GetFullMeasurementSet queries all hierarchy tags required to build +// measurement recommendations. func GetFullMeasurementSet(ctx context.Context, db *gorm.DB) (*orm.MeasurementSet, error) { mSet := &orm.MeasurementSet{ GridToZoneTags: make(map[string][]string), @@ -33,10 +279,35 @@ func GetFullMeasurementSet(ctx context.Context, db *gorm.DB) (*orm.MeasurementSe g, gctx := errgroup.WithContext(ctx) db = db.WithContext(gctx) + var bayLinkedCompTags []string + var bayDevColumnNames []string + + g.Go(func() error { + var linkedComponents []struct { + CompTag string `gorm:"column:comp_tag"` + } + if err := db.Raw(compactMeasurementSQL(sql.MeasurementBayLinkedComponentTags)).Scan(&linkedComponents).Error; err != nil { + return fmt.Errorf("query bay-linked components: %w", err) + } + bayLinkedCompTags = make([]string, 0, len(linkedComponents)) + for _, component := range linkedComponents { + bayLinkedCompTags = append(bayLinkedCompTags, component.CompTag) + } + return nil + }) + + g.Go(func() error { + var err error + bayDevColumnNames, err = QueryBayDevColumnNames(gctx, db) + if err != nil { + return fmt.Errorf("query bay dev columns: %w", err) + } + return nil + }) g.Go(func() error { var grids []orm.Grid - if err := db.Table("grid").Select("tagname").Scan(&grids).Error; err != nil { + if err := db.Raw(compactMeasurementSQL(sql.MeasurementGridTags)).Scan(&grids).Error; err != nil { return fmt.Errorf("query grids: %w", err) } for _, grid := range grids { @@ -52,16 +323,13 @@ func GetFullMeasurementSet(ctx context.Context, db *gorm.DB) (*orm.MeasurementSe orm.Zone GridTag string `gorm:"column:grid_tag"` } - if err := db.Table("zone"). - Select("zone.*, grid.tagname as grid_tag"). - Joins("left join grid on zone.grid_id = grid.id"). - Scan(&zones).Error; err != nil { + if err := db.Raw(compactMeasurementSQL(sql.MeasurementZoneHierarchy)).Scan(&zones).Error; err != nil { return fmt.Errorf("query zones: %w", err) } - for _, z := range zones { - mSet.AllZoneTags = append(mSet.AllZoneTags, z.TAGNAME) - if z.GridTag != "" { - mSet.GridToZoneTags[z.GridTag] = append(mSet.GridToZoneTags[z.GridTag], z.TAGNAME) + for _, zone := range zones { + mSet.AllZoneTags = append(mSet.AllZoneTags, zone.TAGNAME) + if zone.GridTag != "" { + mSet.GridToZoneTags[zone.GridTag] = append(mSet.GridToZoneTags[zone.GridTag], zone.TAGNAME) } } return nil @@ -72,40 +340,40 @@ func GetFullMeasurementSet(ctx context.Context, db *gorm.DB) (*orm.MeasurementSe orm.Station ZoneTag string `gorm:"column:zone_tag"` } - if err := db.Table("station"). - Select("station.*, zone.tagname as zone_tag"). - Joins("left join zone on station.zone_id = zone.id"). - Scan(&stations).Error; err != nil { + if err := db.Raw(compactMeasurementSQL(sql.MeasurementStationHierarchy)).Scan(&stations).Error; err != nil { return fmt.Errorf("query stations: %w", err) } - for _, s := range stations { - mSet.AllStationTags = append(mSet.AllStationTags, s.TAGNAME) - if s.ZoneTag != "" { - mSet.ZoneToStationTags[s.ZoneTag] = append(mSet.ZoneToStationTags[s.ZoneTag], s.TAGNAME) + for _, station := range stations { + mSet.AllStationTags = append(mSet.AllStationTags, station.TAGNAME) + if station.ZoneTag != "" { + mSet.ZoneToStationTags[station.ZoneTag] = append(mSet.ZoneToStationTags[station.ZoneTag], station.TAGNAME) } } return nil }) g.Go(func() error { - var comps []struct { + var components []struct { orm.Component StationTag string `gorm:"column:station_tag"` } - if err := db.Table("component"). - Select("component.*, station.tagname as station_tag"). - Joins("left join station on component.station_id = station.id"). - Scan(&comps).Error; err != nil { + if err := db.Raw(compactMeasurementSQL(sql.MeasurementComponentHierarchy)).Scan(&components).Error; err != nil { return fmt.Errorf("query components: %w", err) } - for _, c := range comps { - mSet.AllCompNSPaths = append(mSet.AllCompNSPaths, c.NSPath) - mSet.AllCompTags = append(mSet.AllCompTags, c.Tag) - if c.StationTag != "" { - mSet.StationToCompNSPaths[c.StationTag] = append(mSet.StationToCompNSPaths[c.StationTag], c.NSPath) + for _, component := range components { + mSet.AllCompNSPaths = append(mSet.AllCompNSPaths, component.NSPath) + mSet.AllCompTags = append(mSet.AllCompTags, component.Tag) + if component.StationTag != "" { + mSet.StationToCompNSPaths[component.StationTag] = append( + mSet.StationToCompNSPaths[component.StationTag], + component.NSPath, + ) } - if c.NSPath != "" { - mSet.CompNSPathToCompTags[c.NSPath] = append(mSet.CompNSPathToCompTags[c.NSPath], c.Tag) + if component.NSPath != "" { + mSet.CompNSPathToCompTags[component.NSPath] = append( + mSet.CompNSPathToCompTags[component.NSPath], + component.Tag, + ) } } return nil @@ -118,20 +386,19 @@ func GetFullMeasurementSet(ctx context.Context, db *gorm.DB) (*orm.MeasurementSe CompNSPath string `gorm:"column:comp_nspath"` BayTag string `gorm:"column:bay_tag"` } - if err := db.Table("measurement"). - Select("measurement.*, component.tag as comp_tag, component.nspath as comp_nspath, bay.tag as bay_tag"). - Joins("left join component on measurement.component_uuid = component.global_uuid"). - Joins("left join bay on measurement.bay_uuid = bay.bay_uuid"). - Scan(&measurements).Error; err != nil { + if err := db.Raw(compactMeasurementSQL(sql.MeasurementTagHierarchy)).Scan(&measurements).Error; err != nil { return fmt.Errorf("query measurements: %w", err) } - for _, m := range measurements { - mSet.AllMeasTags = append(mSet.AllMeasTags, m.Tag) - if m.CompTag != "" { - mSet.CompTagToMeasTags[m.CompTag] = append(mSet.CompTagToMeasTags[m.CompTag], m.Tag) + for _, measurement := range measurements { + mSet.AllMeasTags = append(mSet.AllMeasTags, measurement.Tag) + if measurement.CompTag != "" { + mSet.CompTagToMeasTags[measurement.CompTag] = append( + mSet.CompTagToMeasTags[measurement.CompTag], + measurement.Tag, + ) } - if m.CompNSPath != "" && m.CompNSPath == m.BayTag { - mSet.CompNSPathToMeasTags[m.CompNSPath] = append(mSet.CompNSPathToMeasTags[m.CompNSPath], m.Tag) + if measurement.CompNSPath != "" && measurement.CompNSPath == measurement.BayTag { + mSet.CompNSPathToMeasTags[measurement.CompNSPath] = append(mSet.CompNSPathToMeasTags[measurement.CompNSPath], measurement.Tag) } } return nil @@ -141,6 +408,22 @@ func GetFullMeasurementSet(ctx context.Context, db *gorm.DB) (*orm.MeasurementSe return nil, err } + appendBayDevCandidates(mSet, bayLinkedCompTags, bayDevColumnNames) + mSet.AllConfigTags = append(mSet.AllConfigTags, "bay") return mSet, nil } + +func appendBayDevCandidates(mSet *orm.MeasurementSet, bayLinkedCompTags, bayDevColumnNames []string) { + if mSet == nil || len(bayLinkedCompTags) == 0 || len(bayDevColumnNames) == 0 { + return + } + + mSet.AllMeasTags = append(mSet.AllMeasTags, bayDevColumnNames...) + for _, compTag := range bayLinkedCompTags { + mSet.CompTagToMeasTags[compTag] = append( + mSet.CompTagToMeasTags[compTag], + bayDevColumnNames..., + ) + } +} diff --git a/database/query_component_measurement_bay_test.go b/database/query_component_measurement_bay_test.go new file mode 100644 index 0000000..db0d0a7 --- /dev/null +++ b/database/query_component_measurement_bay_test.go @@ -0,0 +1,52 @@ +package database + +import ( + "testing" + + "modelRT/orm" + + "github.com/stretchr/testify/require" +) + +func TestAppendBayDevCandidatesRequiresBayLinkedComponent(t *testing.T) { + measurementSet := &orm.MeasurementSet{ + AllMeasTags: []string{"current"}, + CompTagToMeasTags: map[string][]string{ + "linked-component": {"current"}, + "unlinked-component": {"voltage"}, + }, + } + + appendBayDevCandidates( + measurementSet, + []string{"linked-component"}, + []string{"dev_instruct", "dev_dyn_sense", "dev_fault_record"}, + ) + + require.Equal(t, + []string{"current", "dev_instruct", "dev_dyn_sense", "dev_fault_record"}, + measurementSet.AllMeasTags, + ) + require.Equal(t, + []string{"current", "dev_instruct", "dev_dyn_sense", "dev_fault_record"}, + measurementSet.CompTagToMeasTags["linked-component"], + ) + require.Equal(t, + []string{"voltage"}, + measurementSet.CompTagToMeasTags["unlinked-component"], + ) +} + +func TestAppendBayDevCandidatesWithoutBayLinkDoesNothing(t *testing.T) { + measurementSet := &orm.MeasurementSet{ + AllMeasTags: []string{"current"}, + CompTagToMeasTags: map[string][]string{ + "component": {"current"}, + }, + } + + appendBayDevCandidates(measurementSet, nil, []string{"dev_instruct"}) + + require.Equal(t, []string{"current"}, measurementSet.AllMeasTags) + require.Equal(t, []string{"current"}, measurementSet.CompTagToMeasTags["component"]) +} diff --git a/database/query_component_measurement_test.go b/database/query_component_measurement_test.go new file mode 100644 index 0000000..5e9d920 --- /dev/null +++ b/database/query_component_measurement_test.go @@ -0,0 +1,271 @@ +package database + +import ( + "context" + "errors" + "fmt" + "regexp" + "strings" + "testing" + + "modelRT/common" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/postgres" + "gorm.io/gorm" +) + +func TestBuildMeasurementTokenValidationQuery(t *testing.T) { + tests := []struct { + name string + token string + wantArgs []any + wantWhere string + wantErr bool + }{ + { + name: "seven-part token", + token: "grid.zone.station.nspath.component.bay.measurement", + wantArgs: []any{"grid", "zone", "station", "nspath", "component", "measurement"}, + wantWhere: "WHERE g.tagname = ?", + }, + { + name: "four-part token", + token: "nspath.component.bay.measurement", + wantArgs: []any{"nspath", "component", "measurement"}, + wantWhere: "WHERE c.nspath = ?", + }, + { + name: "two-part token", + token: "nspath.measurement", + wantArgs: []any{"nspath", "measurement"}, + wantWhere: "WHERE c.nspath = ?", + }, + { + name: "non-bay group", + token: "nspath.component.rated.attribute", + wantErr: true, + }, + { + name: "empty segment", + token: "nspath..bay.measurement", + wantErr: true, + }, + { + name: "invalid segment count", + token: "grid.zone.station", + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + query, args, err := buildMeasurementTokenValidationQuery(tt.token) + if tt.wantErr { + require.Error(t, err) + assert.ErrorIs(t, err, common.ErrInvalidMeasurementToken) + return + } + + require.NoError(t, err) + assert.Contains(t, query, "INNER JOIN component AS c ON c.global_uuid = m.component_uuid") + assert.Contains(t, query, "INNER JOIN bay AS b ON b.bay_uuid = m.bay_uuid") + assert.Contains(t, query, tt.wantWhere) + assert.NotContains(t, query, "grid_idWHERE") + assert.Regexp(t, `grid_id\s+WHERE`, query) + assert.NotContains(t, query, "\n") + assert.NotContains(t, query, "\t") + assert.Equal(t, tt.wantArgs, args) + }) + } +} + +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 + automatic bool + wantMode int16 + }{ + {name: "manual", automatic: false, wantMode: 0}, + {name: "automatic", automatic: true, wantMode: 1}, + } + + for _, tt := range tests { + t.Run(tt.name, func(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.ExpectExec(regexp.QuoteMeta(`UPDATE "measurement" SET "mode"=$1 WHERE id = $2`)). + WithArgs(tt.wantMode, int64(10)). + WillReturnResult(sqlmock.NewResult(0, 1)) + + err = UpdateMeasurementMode(context.Background(), db, 10, tt.automatic) + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +func TestValidateMeasurementToken(t *testing.T) { + tests := []struct { + name string + count int64 + queryErr error + wantErr error + }{ + {name: "exists", count: 1}, + {name: "not found", count: 0, wantErr: common.ErrMeasurementTokenNotFound}, + {name: "ambiguous", count: 2, wantErr: common.ErrAmbiguousMeasurementToken}, + {name: "query failure", queryErr: errors.New("database unavailable")}, + } + + const token = "nspath.measurement" + query, _, err := buildMeasurementTokenValidationQuery(token) + require.NoError(t, err) + expectedQuery := query + for i := 1; strings.Contains(expectedQuery, "?"); i++ { + expectedQuery = strings.Replace(expectedQuery, "?", fmt.Sprintf("$%d", i), 1) + } + + for _, tt := range tests { + t.Run(tt.name, func(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{}) + require.NoError(t, err) + + expectation := mock.ExpectQuery(regexp.QuoteMeta(expectedQuery)). + WithArgs("nspath", "measurement") + if tt.queryErr != nil { + expectation.WillReturnError(tt.queryErr) + } else { + expectation.WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(tt.count)) + } + + err = ValidateMeasurementToken(context.Background(), db, token) + if tt.wantErr != nil { + assert.ErrorIs(t, err, tt.wantErr) + } else if tt.queryErr != nil { + require.Error(t, err) + assert.Contains(t, err.Error(), tt.queryErr.Error()) + } else { + require.NoError(t, err) + } + require.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +func TestBuildMeasurementRowsQuerySeparatesLimitClause(t *testing.T) { + validationQuery, _, err := buildMeasurementTokenValidationQuery("nspath.measurement") + require.NoError(t, err) + + query := buildMeasurementRowsQuery(validationQuery) + assert.NotContains(t, query, "?LIMIT") + assert.Regexp(t, `m\.tag = \?\s+LIMIT 2$`, query) + assert.NotContains(t, query, "\n") + assert.NotContains(t, query, "\t") +} + +func TestQueryMeasurementByDataObjectToken(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{}) + require.NoError(t, err) + + const componentUUID = "70c190f2-8a60-42a9-b143-ec5f87e0aa6b" + mock.ExpectQuery(`(?s)SELECT m\.\*.*WHERE c\.nspath = \$1.*AND m\.tag = \$2.*LIMIT 2`). + WithArgs("nspath", "measurement"). + WillReturnRows(sqlmock.NewRows([]string{ + "id", + "tag", + "name", + "mode", + "size", + "data_source", + "event_plan", + "binding", + "component_uuid", + }).AddRow( + int64(10), + "measurement", + "A phase current", + int16(1), + 10, + `{"type":1,"io_address":{"channel":"tm1"}}`, + `{"enabled":true}`, + `{"ct":{"ratio":1}}`, + componentUUID, + )) + mock.ExpectQuery(`(?s)SELECT global_uuid, nspath, tag, grid, zone, station.*WHERE global_uuid = \$1.*LIMIT 1`). + WithArgs(componentUUID). + WillReturnRows(sqlmock.NewRows([]string{ + "global_uuid", + "nspath", + "tag", + "grid", + "zone", + "station", + }).AddRow(componentUUID, "nspath", "component", "grid", "zone", "station")) + + measurement, component, err := QueryMeasurementByDataObjectToken(context.Background(), db, "nspath.measurement") + require.NoError(t, err) + assert.Equal(t, int64(10), measurement.ID) + assert.Equal(t, int16(1), measurement.Mode) + assert.Equal(t, float64(1), measurement.DataSource["type"]) + assert.Equal(t, "grid", component.GridName) + assert.Equal(t, "component", component.Tag) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/database/query_component_parameter.go b/database/query_component_parameter.go new file mode 100644 index 0000000..6a965f6 --- /dev/null +++ b/database/query_component_parameter.go @@ -0,0 +1,248 @@ +package database + +import ( + "context" + "fmt" + "regexp" + "strings" + + "modelRT/common" + "modelRT/constants" + "modelRT/model" + "modelRT/orm" + modelsql "modelRT/sql" + + "gorm.io/gorm" +) + +var parameterTableNamePattern = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`) + +// ParameterDataObject contains the resolved metadata needed to query a +// parameter attribute after its token has passed hierarchy validation. +type ParameterDataObject struct { + Component orm.Component + Project orm.ProjectManager + TableName string + AttributeGroup string + AttributeName string + AttributeType string +} + +// QueryParameterByDataObjectToken validates a four-part or seven-part +// parameter token. The component group resolves directly to the component +// table; other groups resolve through project_manager and a dynamic table. +func QueryParameterByDataObjectToken(ctx context.Context, db *gorm.DB, token string) (*ParameterDataObject, error) { + componentQuery, componentArgs, parts, err := buildParameterComponentQuery(token) + if err != nil { + return nil, err + } + + var components []orm.Component + if err := db.WithContext(ctx).Raw(componentQuery, componentArgs...).Scan(&components).Error; err != nil { + return nil, fmt.Errorf("query component for parameter token %q: %w", token, err) + } + + switch len(components) { + case 0: + return nil, fmt.Errorf("%w: component hierarchy for %q", common.ErrParameterTokenNotFound, token) + case 1: + // Continue by resolving the component model and attribute group. + default: + return nil, fmt.Errorf("%w: component hierarchy for %q matched more than one record", common.ErrAmbiguousParameterToken, token) + } + + attributeGroup := parts[len(parts)-2] + attributeName := parts[len(parts)-1] + component := components[0] + if attributeGroup == "component" { + attributeType, err := queryParameterAttributeType(ctx, db, "component", attributeName, token) + if err != nil { + return nil, err + } + return &ParameterDataObject{ + Component: component, + TableName: "component", + AttributeGroup: attributeGroup, + AttributeName: attributeName, + AttributeType: attributeType, + }, nil + } + + var projects []orm.ProjectManager + if err := db.WithContext(ctx). + Where("tag = ? AND group_name = ?", component.ModelName, attributeGroup). + Limit(2). + Find(&projects).Error; err != nil { + return nil, fmt.Errorf("query project mapping for parameter token %q: %w", token, err) + } + + switch len(projects) { + case 0: + return nil, fmt.Errorf("%w: model %q does not define attribute group %q", common.ErrParameterTokenNotFound, component.ModelName, attributeGroup) + case 1: + // Continue by validating the dynamic table and attribute. + default: + return nil, fmt.Errorf("%w: model %q and attribute group %q matched more than one project", common.ErrAmbiguousParameterToken, component.ModelName, attributeGroup) + } + + project := projects[0] + if !validParameterTableName(project.Name) { + return nil, fmt.Errorf("project mapping for parameter token %q contains invalid table name %q", token, project.Name) + } + attributeType, err := queryParameterAttributeType(ctx, db, project.Name, attributeName, token) + if err != nil { + return nil, err + } + + var recordCount int64 + if err := db.WithContext(ctx). + Table(project.Name). + Where("global_uuid = ? AND attribute_group = ?", component.GlobalUUID, attributeGroup). + Count(&recordCount).Error; err != nil { + return nil, fmt.Errorf("query dynamic record for parameter token %q: %w", token, err) + } + + switch { + case recordCount == 0: + return nil, fmt.Errorf("%w: component %q has no %q parameter record", common.ErrParameterTokenNotFound, component.Tag, attributeGroup) + case recordCount > 1: + return nil, fmt.Errorf("%w: component %q has %d %q parameter records", common.ErrAmbiguousParameterToken, component.Tag, recordCount, attributeGroup) + } + + return &ParameterDataObject{ + Component: component, + Project: project, + TableName: project.Name, + AttributeGroup: attributeGroup, + AttributeName: attributeName, + AttributeType: attributeType, + }, nil +} + +// QueryParameterDataObjectValue returns token7 from the component row or from +// a dynamic parameter row identified during token validation. +func QueryParameterDataObjectValue(ctx context.Context, db *gorm.DB, parameter *ParameterDataObject) (any, error) { + if parameter == nil { + return nil, fmt.Errorf("parameter data object is nil") + } + + var record map[string]any + query := db.WithContext(ctx).Table(parameter.TableName) + if parameter.AttributeGroup == "component" { + query = query.Where("tag = ?", parameter.Component.Tag) + } else { + query = query.Where("global_uuid = ? AND attribute_group = ?", parameter.Component.GlobalUUID, parameter.AttributeGroup) + } + result := query.Take(&record) + if result.Error != nil { + return nil, fmt.Errorf("query parameter value from table %q: %w", parameter.TableName, result.Error) + } + + value, ok := record[parameter.AttributeName] + if !ok { + return nil, fmt.Errorf("parameter column %q is missing from table %q result", parameter.AttributeName, parameter.TableName) + } + return value, nil +} + +// UpdateParameterDataObjectValue writes token7 to the dynamic parameter row +// resolved from a data-object token. Component-table attributes are not +// supported by the data-object update API. +func UpdateParameterDataObjectValue(ctx context.Context, db *gorm.DB, parameter *ParameterDataObject, value any) error { + if parameter == nil { + return fmt.Errorf("parameter data object is nil") + } + if parameter.AttributeGroup == "component" { + return fmt.Errorf("component data-object updates are not supported") + } + if !validParameterTableName(parameter.TableName) { + return fmt.Errorf("invalid parameter table name %q", parameter.TableName) + } + + result := db.WithContext(ctx). + Table(parameter.TableName). + Where("global_uuid = ? AND attribute_group = ?", parameter.Component.GlobalUUID, parameter.AttributeGroup). + Update(parameter.AttributeName, value) + if result.Error != nil { + return fmt.Errorf("update parameter %q in table %q: %w", parameter.AttributeName, parameter.TableName, result.Error) + } + if result.RowsAffected == 0 { + return fmt.Errorf("update parameter %q in table %q affected no rows", parameter.AttributeName, parameter.TableName) + } + return nil +} + +func validParameterTableName(tableName string) bool { + return parameterTableNamePattern.MatchString(tableName) && constants.IsSupportedParameterTableName(tableName) +} + +// QueryParameterAttributeDescription returns the display name registered for +// token7 in basic.attribute. +func QueryParameterAttributeDescription(ctx context.Context, db *gorm.DB, attributeName string) (string, error) { + var rows []struct { + Description string `gorm:"column:attribute_name"` + } + if err := db.WithContext(ctx). + Raw(modelsql.ParameterAttributeDescription, attributeName). + Scan(&rows).Error; err != nil { + return "", fmt.Errorf("query parameter description for attribute %q: %w", attributeName, err) + } + + switch len(rows) { + case 0: + return "", fmt.Errorf("parameter description not found for attribute %q", attributeName) + case 1: + return rows[0].Description, nil + default: + return "", fmt.Errorf("ambiguous parameter description for attribute %q", attributeName) + } +} + +func queryParameterAttributeType(ctx context.Context, db *gorm.DB, tableName, attributeName, token string) (string, error) { + var attributeType string + result := db.WithContext(ctx). + Raw(modelsql.ParameterAttributeColumnType, tableName, attributeName). + Scan(&attributeType) + if result.Error != nil { + return "", fmt.Errorf("validate attribute column for parameter token %q: %w", token, result.Error) + } + if result.RowsAffected == 0 || attributeType == "" { + return "", fmt.Errorf("%w: column %q does not exist in parameter table %q", common.ErrParameterTokenNotFound, attributeName, tableName) + } + return strings.ToUpper(attributeType), nil +} + +func buildParameterComponentQuery(token string) (string, []any, []string, error) { + dataObjectType, err := model.ClassifyDataObjectToken(token) + if err != nil { + return "", nil, nil, fmt.Errorf("%w %q: %v", common.ErrInvalidParameterToken, token, err) + } + if dataObjectType != constants.DataObjectTypeParameter { + return "", nil, nil, fmt.Errorf("%w %q: token does not identify a parameter", common.ErrInvalidParameterToken, token) + } + + parts := strings.Split(token, ".") + var where string + var args []any + switch len(parts) { + case 7: + where = modelsql.ParameterSevenPartTokenWhere + args = []any{parts[0], parts[1], parts[2], parts[3], parts[4]} + case 4: + where = modelsql.ParameterFourPartTokenWhere + args = []any{parts[0], parts[1]} + default: + return "", nil, nil, fmt.Errorf("%w %q: expected 4 or 7 segments", common.ErrInvalidParameterToken, token) + } + + query := compactParameterSQL(strings.Join([]string{ + modelsql.ParameterComponentQueryBase, + where, + modelsql.ParameterLimitTwo, + }, "\n")) + return query, args, parts, nil +} + +func compactParameterSQL(statement string) string { + return strings.Join(strings.Fields(statement), " ") +} diff --git a/database/query_component_parameter_test.go b/database/query_component_parameter_test.go new file mode 100644 index 0000000..c3a6418 --- /dev/null +++ b/database/query_component_parameter_test.go @@ -0,0 +1,286 @@ +package database + +import ( + "context" + "regexp" + "testing" + + "modelRT/common" + "modelRT/orm" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/gofrs/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/postgres" + "gorm.io/gorm" +) + +func TestBuildParameterComponentQuery(t *testing.T) { + tests := []struct { + name string + token string + wantArgs []any + wantWhere string + wantErr bool + }{ + { + name: "seven-part parameter", + token: "grid.zone.station.nspath.component.stable.attribute", + wantArgs: []any{"grid", "zone", "station", "nspath", "component"}, + wantWhere: "WHERE g.tagname = ?", + }, + { + name: "four-part local parameter", + token: "nspath.component.rated.attribute", + wantArgs: []any{"nspath", "component"}, + wantWhere: "s.is_local = TRUE", + }, + { + name: "measurement token", + token: "nspath.component.bay.measurement", + wantErr: true, + }, + { + name: "empty segment", + token: "nspath..stable.attribute", + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + query, args, _, err := buildParameterComponentQuery(tt.token) + if tt.wantErr { + require.Error(t, err) + assert.ErrorIs(t, err, common.ErrInvalidParameterToken) + return + } + + require.NoError(t, err) + assert.Contains(t, query, "INNER JOIN station AS s ON s.id = c.station_id") + assert.Contains(t, query, tt.wantWhere) + assert.Equal(t, tt.wantArgs, args) + assert.NotContains(t, query, "\n") + assert.NotContains(t, query, "\t") + }) + } +} + +func TestQueryParameterByDataObjectToken(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{}) + require.NoError(t, err) + + const ( + token = "grid.zone.station.nspath.component.stable.rated_voltage" + componentUUID = "70c190f2-8a60-42a9-b143-ec5f87e0aa6b" + ) + + mock.ExpectQuery(`(?s)SELECT c\.\*.*WHERE g\.tagname = \$1.*AND c\.tag = \$5.*LIMIT 2`). + WithArgs("grid", "zone", "station", "nspath", "component"). + WillReturnRows(sqlmock.NewRows([]string{ + "global_uuid", "nspath", "tag", "model_name", "station_id", + }).AddRow(componentUUID, "nspath", "component", "bus_1", int64(10))) + + mock.ExpectQuery(`SELECT \* FROM "project_manager" WHERE tag = \$1 AND group_name = \$2 LIMIT \$3`). + WithArgs("bus_1", "stable", 2). + WillReturnRows(sqlmock.NewRows([]string{ + "id", "name", "tag", "meta_model", "group_name", "link_type", "check_state", "ispublic", + }).AddRow( + int32(1), + "bus_bus_1_stable", + "bus_1", + "bus", + "stable", + int32(0), + `{"checkState":[{"name":"rated_voltage","checked":1}]}`, + false, + )) + + mock.ExpectQuery(`(?s)SELECT pg_catalog\.format_type.*pg_catalog\.pg_attribute.*c\.relname = \$1.*a\.attname = \$2.*LIMIT 1`). + WithArgs("bus_bus_1_stable", "rated_voltage"). + WillReturnRows(sqlmock.NewRows([]string{"format_type"}).AddRow("double precision")) + + mock.ExpectQuery(regexp.QuoteMeta(`SELECT count(*) FROM "bus_bus_1_stable" WHERE global_uuid = $1 AND attribute_group = $2`)). + WithArgs(componentUUID, "stable"). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(int64(1))) + + parameter, err := QueryParameterByDataObjectToken(context.Background(), db, token) + require.NoError(t, err) + assert.Equal(t, "component", parameter.Component.Tag) + assert.Equal(t, "bus_bus_1_stable", parameter.Project.Name) + assert.Equal(t, "bus_bus_1_stable", parameter.TableName) + assert.Equal(t, "stable", parameter.AttributeGroup) + assert.Equal(t, "rated_voltage", parameter.AttributeName) + assert.Equal(t, "DOUBLE PRECISION", parameter.AttributeType) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestQueryParameterDataObjectValue(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{}) + require.NoError(t, err) + + const componentUUID = "70c190f2-8a60-42a9-b143-ec5f87e0aa6b" + parsedUUID, err := uuid.FromString(componentUUID) + require.NoError(t, err) + parameter := &ParameterDataObject{ + Component: orm.Component{GlobalUUID: parsedUUID}, + Project: orm.ProjectManager{ + Name: "bus_bus_1_stable", + }, + TableName: "bus_bus_1_stable", + AttributeGroup: "stable", + AttributeName: "rated_voltage", + } + + mock.ExpectQuery(regexp.QuoteMeta(`SELECT * FROM "bus_bus_1_stable" WHERE global_uuid = $1 AND attribute_group = $2 LIMIT $3`)). + WithArgs(componentUUID, "stable", 1). + WillReturnRows(sqlmock.NewRows([]string{ + "global_uuid", "attribute_group", "rated_voltage", + }).AddRow(componentUUID, "stable", float64(220))) + + value, err := QueryParameterDataObjectValue(context.Background(), db, parameter) + require.NoError(t, err) + assert.Equal(t, float64(220), value) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateParameterDataObjectValue(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) + + const componentUUID = "70c190f2-8a60-42a9-b143-ec5f87e0aa6b" + parsedUUID, err := uuid.FromString(componentUUID) + require.NoError(t, err) + parameter := &ParameterDataObject{ + Component: orm.Component{GlobalUUID: parsedUUID}, + TableName: "bus_bus_1_rated", + AttributeGroup: "rated", + AttributeName: "unom_kv", + } + + mock.ExpectExec(regexp.QuoteMeta(`UPDATE "bus_bus_1_rated" SET "unom_kv"=$1 WHERE global_uuid = $2 AND attribute_group = $3`)). + WithArgs("15.2", componentUUID, "rated"). + WillReturnResult(sqlmock.NewResult(0, 1)) + + err = UpdateParameterDataObjectValue(context.Background(), db, parameter, "15.2") + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateParameterDataObjectValueRejectsComponent(t *testing.T) { + err := UpdateParameterDataObjectValue(context.Background(), &gorm.DB{}, &ParameterDataObject{ + TableName: "component", + AttributeGroup: "component", + AttributeName: "global_uuid", + }, "uuid") + require.Error(t, err) + assert.Contains(t, err.Error(), "not supported") +} + +func TestQueryComponentParameterByDataObjectToken(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{}) + require.NoError(t, err) + + const ( + token = "grid.zone.station.nspath.component.component.description" + componentUUID = "70c190f2-8a60-42a9-b143-ec5f87e0aa6b" + ) + mock.ExpectQuery(`(?s)SELECT c\.\*.*WHERE g\.tagname = \$1.*AND c\.tag = \$5.*LIMIT 2`). + WithArgs("grid", "zone", "station", "nspath", "component"). + WillReturnRows(sqlmock.NewRows([]string{ + "global_uuid", "nspath", "tag", "model_name", "grid", "zone", "station", "station_id", + }).AddRow(componentUUID, "nspath", "component", "bus_1", "grid", "zone", "station", int64(10))) + mock.ExpectQuery(`(?s)SELECT pg_catalog\.format_type.*c\.relname = \$1.*a\.attname = \$2.*LIMIT 1`). + WithArgs("component", "description"). + WillReturnRows(sqlmock.NewRows([]string{"format_type"}).AddRow("character varying(512)")) + + parameter, err := QueryParameterByDataObjectToken(context.Background(), db, token) + require.NoError(t, err) + assert.Equal(t, "component", parameter.TableName) + assert.Equal(t, "component", parameter.AttributeGroup) + assert.Equal(t, "description", parameter.AttributeName) + assert.Equal(t, "CHARACTER VARYING(512)", parameter.AttributeType) + assert.Empty(t, parameter.Project.Name) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestQueryComponentParameterValue(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{}) + require.NoError(t, err) + + parameter := &ParameterDataObject{ + Component: orm.Component{Tag: "component"}, + TableName: "component", + AttributeGroup: "component", + AttributeName: "description", + } + mock.ExpectQuery(regexp.QuoteMeta(`SELECT * FROM "component" WHERE tag = $1 LIMIT $2`)). + WithArgs("component", 1). + WillReturnRows(sqlmock.NewRows([]string{"tag", "description"}).AddRow("component", "测试组件")) + + value, err := QueryParameterDataObjectValue(context.Background(), db, parameter) + require.NoError(t, err) + assert.Equal(t, "测试组件", value) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestQueryParameterAttributeDescription(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{}) + require.NoError(t, err) + + mock.ExpectQuery(`(?s)SELECT attribute_name.*FROM basic\.attribute.*WHERE attribute = \$1.*LIMIT 2`). + WithArgs("rated_voltage"). + WillReturnRows(sqlmock.NewRows([]string{"attribute_name"}).AddRow("额定电压")) + + description, err := QueryParameterAttributeDescription(context.Background(), db, "rated_voltage") + require.NoError(t, err) + assert.Equal(t, "额定电压", description) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestQueryParameterByDataObjectTokenComponentNotFound(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{}) + require.NoError(t, err) + + mock.ExpectQuery(`(?s)SELECT c\.\*.*WHERE c\.nspath = \$1.*AND c\.tag = \$2.*s\.is_local = TRUE.*LIMIT 2`). + WithArgs("nspath", "component"). + WillReturnRows(sqlmock.NewRows([]string{"global_uuid"})) + + _, err = QueryParameterByDataObjectToken( + context.Background(), + db, + "nspath.component.stable.rated_voltage", + ) + require.Error(t, err) + assert.ErrorIs(t, err, common.ErrParameterTokenNotFound) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/database/query_measurement.go b/database/query_measurement.go deleted file mode 100644 index aa51e89..0000000 --- a/database/query_measurement.go +++ /dev/null @@ -1,62 +0,0 @@ -// Package database define database operation functions -package database - -import ( - "context" - "time" - - "modelRT/orm" - - "gorm.io/gorm" - "gorm.io/gorm/clause" -) - -// 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) { - var measurement orm.Measurement - // ctx超时判断 - cancelCtx, cancel := context.WithTimeout(ctx, 5*time.Second) - defer cancel() - result := tx.WithContext(cancelCtx). - Where("id = ?", id). - Clauses(clause.Locking{Strength: "UPDATE"}). - First(&measurement) - - if result.Error != nil { - return orm.Measurement{}, result.Error - } - return measurement, nil -} - -// QueryMeasurementByToken define function query circuit diagram component measurement info by token from postgresDB -func QueryMeasurementByToken(ctx context.Context, tx *gorm.DB, token string) (orm.Measurement, error) { - // TODO parse token to avoid SQL injection - - var component orm.Measurement - // ctx超时判断 - cancelCtx, cancel := context.WithTimeout(ctx, 5*time.Second) - defer cancel() - result := tx.WithContext(cancelCtx). - Where(" = ?", token). - Clauses(clause.Locking{Strength: "UPDATE"}). - First(&component) - - if result.Error != nil { - return orm.Measurement{}, result.Error - } - return component, nil -} - -// GetAllMeasurements define func to query all measurement info from postgresDB -func GetAllMeasurements(ctx context.Context, tx *gorm.DB) ([]orm.Measurement, error) { - var measurements []orm.Measurement - - // ctx超时判断 - cancelCtx, cancel := context.WithTimeout(ctx, 5*time.Second) - defer cancel() - result := tx.WithContext(cancelCtx).Clauses(clause.Locking{Strength: "UPDATE"}).Find(&measurements) - if result.Error != nil { - return nil, result.Error - } - return measurements, nil -} diff --git a/deploy/redis-test-data/real-time-compute/compute_data_injection.go b/deploy/redis-test-data/real-time-compute/compute_data_injection.go index 2f68ef7..8707336 100644 --- a/deploy/redis-test-data/real-time-compute/compute_data_injection.go +++ b/deploy/redis-test-data/real-time-compute/compute_data_injection.go @@ -88,7 +88,7 @@ func generateNormalData(baseValue, normalBase float64) []float64 { func main() { rootCtx := context.Background() - pgURI := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s", "192.168.1.101", 5432, "postgres", "coslight", "demo") + pgURI := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s", "localhost", 5432, "postgres", "coslight", "develop_env") postgresDBClient, err := gorm.Open(postgres.Open(pgURI)) if err != nil { @@ -164,7 +164,6 @@ func main() { } datas = generateMixedData(highMin, lowMin, highBase, lowBase, baseValue, normalBase) - // log.Printf("key:%s\n datas:%v\n", key, datas) allHigh := true for i := highStart; i < highEnd; i++ { diff --git a/deploy/redis-test-data/util/rand.go b/deploy/redis-test-data/util/rand.go index 7df9c92..2e31302 100644 --- a/deploy/redis-test-data/util/rand.go +++ b/deploy/redis-test-data/util/rand.go @@ -3,6 +3,7 @@ package util import ( "fmt" + "strings" "modelRT/orm" ) @@ -61,7 +62,7 @@ func ProcessMeasurements(measurements []orm.Measurement) map[string]CalculationR device, _ := ioAddress["device"].(string) channel, _ := ioAddress["channel"].(string) - result := fmt.Sprintf("%s:%s:phasor:%s", station, device, channel) + result := strings.ToLower(fmt.Sprintf("%s:%s:phasor:%s", station, device, channel)) if measurement.EventPlan == nil { continue } diff --git a/diagram/context.go b/diagram/context.go new file mode 100644 index 0000000..ce6047b --- /dev/null +++ b/diagram/context.go @@ -0,0 +1,20 @@ +package diagram + +import ( + "context" + "fmt" + + "modelRT/common" + "modelRT/constants" +) + +func clientTokenFromContext(ctx context.Context) (string, error) { + if ctx == nil { + return "", common.ErrGetClientToken + } + token, ok := ctx.Value(constants.CtxKeyClientToken).(string) + if !ok || token == "" { + return "", fmt.Errorf("%w: missing or invalid context value", common.ErrGetClientToken) + } + return token, nil +} diff --git a/diagram/context_test.go b/diagram/context_test.go new file mode 100644 index 0000000..33ec1b9 --- /dev/null +++ b/diagram/context_test.go @@ -0,0 +1,38 @@ +package diagram + +import ( + "context" + "testing" + + "modelRT/common" + "modelRT/constants" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestClientTokenFromContext(t *testing.T) { + ctx := context.WithValue(context.Background(), constants.CtxKeyClientToken, "test-token") + token, err := clientTokenFromContext(ctx) + require.NoError(t, err) + assert.Equal(t, "test-token", token) +} + +func TestClientTokenFromContextReturnsErrorWhenMissing(t *testing.T) { + for _, ctx := range []context.Context{nil, context.Background()} { + _, err := clientTokenFromContext(ctx) + require.Error(t, err) + assert.ErrorIs(t, err, common.ErrGetClientToken) + } +} + +func TestRedisConstructorsReturnErrorInsteadOfPanickingWithoutToken(t *testing.T) { + ctx := context.Background() + + _, err := NewRedisZSet(ctx, "zset", 0, false) + assert.ErrorIs(t, err, common.ErrGetClientToken) + _, err = NewRedisSet(ctx, "set", 0, false) + assert.ErrorIs(t, err, common.ErrGetClientToken) + _, err = NewRedisHash(ctx, "hash", 0, false) + assert.ErrorIs(t, err, common.ErrGetClientToken) +} diff --git a/diagram/redis_client.go b/diagram/redis_client.go index 4b673a8..247c98c 100644 --- a/diagram/redis_client.go +++ b/diagram/redis_client.go @@ -3,6 +3,8 @@ package diagram import ( "context" + "fmt" + "strconv" "github.com/redis/go-redis/v9" ) @@ -12,6 +14,46 @@ type RedisClient struct { Client *redis.Client } +// QueryLatestMeasurementValue returns the score whose member contains the +// 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) { + if rc.Client == nil { + return 0, fmt.Errorf("redis client is not initialized") + } + + members, err := rc.Client.ZRangeWithScores(ctx, key, 0, -1).Result() + if err != nil { + return 0, err + } + return latestMeasurementValue(members, key) +} + +func latestMeasurementValue(members []redis.Z, key string) (float64, error) { + if len(members) == 0 { + return 0, fmt.Errorf("real-time measurement value not found for key %q", key) + } + + var latestTimestamp int64 + var latestValue float64 + found := false + 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 + } + } + if !found { + return 0, fmt.Errorf("real-time measurement timestamps are invalid for key %q", key) + } + return latestValue, nil +} + // NewRedisClient define func of new redis client instance func NewRedisClient() *RedisClient { return &RedisClient{ diff --git a/diagram/redis_client_test.go b/diagram/redis_client_test.go new file mode 100644 index 0000000..7416ba1 --- /dev/null +++ b/diagram/redis_client_test.go @@ -0,0 +1,28 @@ +package diagram + +import ( + "testing" + + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestLatestMeasurementValueUsesMemberTimestamp(t *testing.T) { + value, err := latestMeasurementValue([]redis.Z{ + {Member: "100", Score: 999}, + {Member: "300", Score: 12}, + {Member: "200", Score: 500}, + }, "measurement-key") + + require.NoError(t, err) + assert.Equal(t, float64(12), value) +} + +func TestLatestMeasurementValueRejectsMissingOrInvalidTimestamps(t *testing.T) { + _, err := latestMeasurementValue(nil, "measurement-key") + require.Error(t, err) + + _, err = latestMeasurementValue([]redis.Z{{Member: "invalid", Score: 1}}, "measurement-key") + require.Error(t, err) +} diff --git a/diagram/redis_hash.go b/diagram/redis_hash.go index 2382828..edc39c5 100644 --- a/diagram/redis_hash.go +++ b/diagram/redis_hash.go @@ -18,14 +18,17 @@ type RedisHash struct { } // NewRedisHash define func of new redis hash instance -func NewRedisHash(ctx context.Context, hashKey string, lockLeaseTime uint64, needRefresh bool) *RedisHash { - token := ctx.Value("client_token").(string) +func NewRedisHash(ctx context.Context, hashKey string, lockLeaseTime uint64, needRefresh bool) (*RedisHash, error) { + token, err := clientTokenFromContext(ctx) + if err != nil { + return nil, err + } return &RedisHash{ ctx: ctx, hashKey: hashKey, rwLocker: locker.InitRWLocker(hashKey, token, lockLeaseTime, needRefresh), storageClient: GetRedisClientInstance(), - } + }, nil } // SetRedisHashByMap define func of set redis hash by map struct diff --git a/diagram/redis_set.go b/diagram/redis_set.go index bfb9f6c..61d7064 100644 --- a/diagram/redis_set.go +++ b/diagram/redis_set.go @@ -21,15 +21,18 @@ type RedisSet struct { } // NewRedisSet define func of new redis set instance -func NewRedisSet(ctx context.Context, setKey string, lockLeaseTime uint64, needRefresh bool) *RedisSet { - token := ctx.Value("client_token").(string) +func NewRedisSet(ctx context.Context, setKey string, lockLeaseTime uint64, needRefresh bool) (*RedisSet, error) { + token, err := clientTokenFromContext(ctx) + if err != nil { + return nil, err + } return &RedisSet{ ctx: ctx, key: setKey, rwLocker: locker.InitRWLocker(setKey, token, lockLeaseTime, needRefresh), storageClient: GetRedisClientInstance(), logger: logger.GetLoggerInstance(), - } + }, nil } // SADD define func of add redis set by members diff --git a/diagram/redis_zset.go b/diagram/redis_zset.go index 6884448..d8f5ee1 100644 --- a/diagram/redis_zset.go +++ b/diagram/redis_zset.go @@ -18,13 +18,16 @@ type RedisZSet struct { } // NewRedisZSet define func of new redis zset instance -func NewRedisZSet(ctx context.Context, key string, lockLeaseTime uint64, needRefresh bool) *RedisZSet { - token := ctx.Value("client_token").(string) +func NewRedisZSet(ctx context.Context, key string, lockLeaseTime uint64, needRefresh bool) (*RedisZSet, error) { + token, err := clientTokenFromContext(ctx) + if err != nil { + return nil, err + } return &RedisZSet{ ctx: ctx, rwLocker: locker.InitRWLocker(key, token, lockLeaseTime, needRefresh), storageClient: GetRedisClientInstance(), - } + }, nil } // ZADD define func of add redis zset by members @@ -44,6 +47,26 @@ func (rs *RedisZSet) ZADD(setKey string, score float64, member any) error { return nil } +// ZREPLACE atomically removes all existing members and adds one new member. +func (rs *RedisZSet) ZREPLACE(setKey string, score float64, member any) error { + if err := rs.rwLocker.WLock(rs.ctx); err != nil { + logger.Error(rs.ctx, "lock wLock by setKey failed", "set_key", setKey, "error", err) + return err + } + defer rs.rwLocker.UnWLock(rs.ctx) + + _, err := rs.storageClient.TxPipelined(rs.ctx, func(pipe redis.Pipeliner) error { + pipe.Del(rs.ctx, setKey) + pipe.ZAdd(rs.ctx, setKey, redis.Z{Score: score, Member: member}) + return nil + }) + if err != nil { + logger.Error(rs.ctx, "replace zset member failed", "set_key", setKey, "member", member, "error", err) + return err + } + return nil +} + // ZRANGE define func of returns the specified range of elements in the sorted set stored by key func (rs *RedisZSet) ZRANGE(setKey string, start, stop int64) ([]string, error) { var results []string diff --git a/docs/docs.go b/docs/docs.go index a3451d5..51d1bc2 100644 --- a/docs/docs.go +++ b/docs/docs.go @@ -23,6 +23,57 @@ const docTemplate = `{ "host": "{{.Host}}", "basePath": "{{.BasePath}}", "paths": { + "/data-object/recommend": { + "get": { + "description": "根据用户输入的字符串,从 Redis 中查询可能的测量点或结构路径,并提供推荐列表。", + "consumes": [ + "application/json" + ], + "produces": [ + "application/json" + ], + "tags": [ + "DataObject Recommend" + ], + "summary": "测量点推荐(搜索框自动补全)", + "parameters": [ + { + "type": "string", + "example": "\"grid1\"", + "description": "推荐关键词,例如 'grid1' 或 'grid1.'", + "name": "input", + "in": "query", + "required": true + } + ], + "responses": { + "200": { + "description": "返回推荐列表成功", + "schema": { + "allOf": [ + { + "$ref": "#/definitions/network.SuccessResponse" + }, + { + "type": "object", + "properties": { + "payload": { + "$ref": "#/definitions/network.DataObjectRecommendPayload" + } + } + } + ] + } + }, + "400": { + "description": "返回推荐列表失败", + "schema": { + "$ref": "#/definitions/network.FailureResponse" + } + } + } + } + }, "/data/realtime": { "get": { "description": "根据用户输入的组件token,从 dataRT 服务中持续获取测点实时数据", @@ -87,57 +138,6 @@ const docTemplate = `{ } } }, - "/measurement/recommend": { - "get": { - "description": "根据用户输入的字符串,从 Redis 中查询可能的测量点或结构路径,并提供推荐列表。", - "consumes": [ - "application/json" - ], - "produces": [ - "application/json" - ], - "tags": [ - "Measurement Recommend" - ], - "summary": "测量点推荐(搜索框自动补全)", - "parameters": [ - { - "type": "string", - "example": "\"grid1\"", - "description": "推荐关键词,例如 'grid1' 或 'grid1.'", - "name": "input", - "in": "query", - "required": true - } - ], - "responses": { - "200": { - "description": "返回推荐列表成功", - "schema": { - "allOf": [ - { - "$ref": "#/definitions/network.SuccessResponse" - }, - { - "type": "object", - "properties": { - "payload": { - "$ref": "#/definitions/network.MeasurementRecommendPayload" - } - } - } - ] - } - }, - "400": { - "description": "返回推荐列表失败", - "schema": { - "$ref": "#/definitions/network.FailureResponse" - } - } - } - } - }, "/model/diagram_load/{page_id}": { "get": { "description": "load circuit diagram info by page id", @@ -487,23 +487,7 @@ const docTemplate = `{ } } }, - "network.FailureResponse": { - "type": "object", - "properties": { - "code": { - "type": "integer", - "example": 3000 - }, - "msg": { - "type": "string", - "example": "process completed with partial failures" - }, - "payload": { - "type": "object" - } - } - }, - "network.MeasurementRecommendPayload": { + "network.DataObjectRecommendPayload": { "type": "object", "properties": { "input": { @@ -527,6 +511,22 @@ const docTemplate = `{ } } }, + "network.FailureResponse": { + "type": "object", + "properties": { + "code": { + "type": "integer", + "example": 3000 + }, + "msg": { + "type": "string", + "example": "process completed with partial failures" + }, + "payload": { + "type": "object" + } + } + }, "network.RealTimeDataPayload": { "type": "object", "properties": { diff --git a/docs/swagger.json b/docs/swagger.json index 14b2253..16b355a 100644 --- a/docs/swagger.json +++ b/docs/swagger.json @@ -17,6 +17,57 @@ "host": "localhost:8080", "basePath": "/api/v1", "paths": { + "/data-object/recommend": { + "get": { + "description": "根据用户输入的字符串,从 Redis 中查询可能的测量点或结构路径,并提供推荐列表。", + "consumes": [ + "application/json" + ], + "produces": [ + "application/json" + ], + "tags": [ + "DataObject Recommend" + ], + "summary": "测量点推荐(搜索框自动补全)", + "parameters": [ + { + "type": "string", + "example": "\"grid1\"", + "description": "推荐关键词,例如 'grid1' 或 'grid1.'", + "name": "input", + "in": "query", + "required": true + } + ], + "responses": { + "200": { + "description": "返回推荐列表成功", + "schema": { + "allOf": [ + { + "$ref": "#/definitions/network.SuccessResponse" + }, + { + "type": "object", + "properties": { + "payload": { + "$ref": "#/definitions/network.DataObjectRecommendPayload" + } + } + } + ] + } + }, + "400": { + "description": "返回推荐列表失败", + "schema": { + "$ref": "#/definitions/network.FailureResponse" + } + } + } + } + }, "/data/realtime": { "get": { "description": "根据用户输入的组件token,从 dataRT 服务中持续获取测点实时数据", @@ -81,57 +132,6 @@ } } }, - "/measurement/recommend": { - "get": { - "description": "根据用户输入的字符串,从 Redis 中查询可能的测量点或结构路径,并提供推荐列表。", - "consumes": [ - "application/json" - ], - "produces": [ - "application/json" - ], - "tags": [ - "Measurement Recommend" - ], - "summary": "测量点推荐(搜索框自动补全)", - "parameters": [ - { - "type": "string", - "example": "\"grid1\"", - "description": "推荐关键词,例如 'grid1' 或 'grid1.'", - "name": "input", - "in": "query", - "required": true - } - ], - "responses": { - "200": { - "description": "返回推荐列表成功", - "schema": { - "allOf": [ - { - "$ref": "#/definitions/network.SuccessResponse" - }, - { - "type": "object", - "properties": { - "payload": { - "$ref": "#/definitions/network.MeasurementRecommendPayload" - } - } - } - ] - } - }, - "400": { - "description": "返回推荐列表失败", - "schema": { - "$ref": "#/definitions/network.FailureResponse" - } - } - } - } - }, "/model/diagram_load/{page_id}": { "get": { "description": "load circuit diagram info by page id", @@ -481,23 +481,7 @@ } } }, - "network.FailureResponse": { - "type": "object", - "properties": { - "code": { - "type": "integer", - "example": 3000 - }, - "msg": { - "type": "string", - "example": "process completed with partial failures" - }, - "payload": { - "type": "object" - } - } - }, - "network.MeasurementRecommendPayload": { + "network.DataObjectRecommendPayload": { "type": "object", "properties": { "input": { @@ -521,6 +505,22 @@ } } }, + "network.FailureResponse": { + "type": "object", + "properties": { + "code": { + "type": "integer", + "example": 3000 + }, + "msg": { + "type": "string", + "example": "process completed with partial failures" + }, + "payload": { + "type": "object" + } + } + }, "network.RealTimeDataPayload": { "type": "object", "properties": { diff --git a/docs/swagger.yaml b/docs/swagger.yaml index 92cbdc8..4f0d2e2 100644 --- a/docs/swagger.yaml +++ b/docs/swagger.yaml @@ -59,18 +59,7 @@ definitions: example: 3 type: integer type: object - network.FailureResponse: - properties: - code: - example: 3000 - type: integer - msg: - example: process completed with partial failures - type: string - payload: - type: object - type: object - network.MeasurementRecommendPayload: + network.DataObjectRecommendPayload: properties: input: example: transformfeeder1_220. @@ -87,6 +76,17 @@ definitions: type: string type: array type: object + network.FailureResponse: + properties: + code: + example: 3000 + type: integer + msg: + example: process completed with partial failures + type: string + payload: + type: object + type: object network.RealTimeDataPayload: properties: sub_pos: @@ -169,6 +169,37 @@ info: title: ModelRT 实时模型服务 API 文档 version: "1.0" paths: + /data-object/recommend: + get: + consumes: + - application/json + description: 根据用户输入的字符串,从 Redis 中查询可能的测量点或结构路径,并提供推荐列表。 + parameters: + - description: 推荐关键词,例如 'grid1' 或 'grid1.' + example: '"grid1"' + in: query + name: input + required: true + type: string + produces: + - application/json + responses: + "200": + description: 返回推荐列表成功 + schema: + allOf: + - $ref: '#/definitions/network.SuccessResponse' + - properties: + payload: + $ref: '#/definitions/network.DataObjectRecommendPayload' + type: object + "400": + description: 返回推荐列表失败 + schema: + $ref: '#/definitions/network.FailureResponse' + summary: 测量点推荐(搜索框自动补全) + tags: + - DataObject Recommend /data/realtime: get: consumes: @@ -209,37 +240,6 @@ paths: summary: 获取实时测点数据 tags: - RealTime Component - /measurement/recommend: - get: - consumes: - - application/json - description: 根据用户输入的字符串,从 Redis 中查询可能的测量点或结构路径,并提供推荐列表。 - parameters: - - description: 推荐关键词,例如 'grid1' 或 'grid1.' - example: '"grid1"' - in: query - name: input - required: true - type: string - produces: - - application/json - responses: - "200": - description: 返回推荐列表成功 - schema: - allOf: - - $ref: '#/definitions/network.SuccessResponse' - - properties: - payload: - $ref: '#/definitions/network.MeasurementRecommendPayload' - type: object - "400": - description: 返回推荐列表失败 - schema: - $ref: '#/definitions/network.FailureResponse' - summary: 测量点推荐(搜索框自动补全) - tags: - - Measurement Recommend /model/diagram_load/{page_id}: get: consumes: diff --git a/go.mod b/go.mod index 441ab18..d32ddda 100644 --- a/go.mod +++ b/go.mod @@ -11,6 +11,7 @@ require ( github.com/gofrs/uuid v4.4.0+incompatible github.com/gomodule/redigo v1.8.9 github.com/gorilla/websocket v1.5.3 + github.com/jackc/pgx/v5 v5.5.5 github.com/json-iterator/go v1.1.12 github.com/natefinch/lumberjack v2.0.0+incompatible github.com/panjf2000/ants/v2 v2.10.0 @@ -62,7 +63,6 @@ require ( github.com/hashicorp/hcl v1.0.0 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a // indirect - github.com/jackc/pgx/v5 v5.5.5 // indirect github.com/jackc/puddle/v2 v2.2.1 // indirect github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect diff --git a/handler/async_task_create_handler.go b/handler/async_task_create_handler.go index 6da67c5..9afa1d9 100644 --- a/handler/async_task_create_handler.go +++ b/handler/async_task_create_handler.go @@ -154,8 +154,7 @@ func validateBatchImportParams(params map[string]any) bool { func validateTestTaskParams(params map[string]any) bool { // Test task has optional parameters, all are valid // sleep_duration defaults to 60 seconds if not provided - // TODO Add more validation logic for test task parameters if needed - fmt.Println(params) + fmt.Println("Test task parameters:", params) return true } diff --git a/handler/component_attribute_query.go b/handler/component_attribute_query.go index bfa24d8..aa08801 100644 --- a/handler/component_attribute_query.go +++ b/handler/component_attribute_query.go @@ -8,8 +8,6 @@ import ( "slices" "strings" - "github.com/gofrs/uuid" - "modelRT/common/errcode" "modelRT/constants" "modelRT/database" @@ -18,6 +16,7 @@ import ( "modelRT/orm" "github.com/gin-gonic/gin" + "github.com/gofrs/uuid" ) // ComponentAttributeQueryHandler define circuit diagram component attribute value query process API @@ -55,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) @@ -186,7 +193,11 @@ func fillRemainingErrors(results map[string]queryResult, tokens []string, err *e } func backfillRedis(ctx context.Context, hSetKey string, items []cacheQueryItem) { - hset := diagram.NewRedisHash(ctx, hSetKey, 5000, false) + hset, err := diagram.NewRedisHash(ctx, hSetKey, 5000, false) + if err != nil { + logger.Error(ctx, "create redis hash for async backfill failed", "hash_key", hSetKey, "error", err) + return + } fields := make(map[string]any, len(items)) for _, item := range items { if item.attributeVal != "" { diff --git a/handler/component_attribute_update.go b/handler/component_attribute_update.go index 8a10f67..28b4836 100644 --- a/handler/component_attribute_update.go +++ b/handler/component_attribute_update.go @@ -140,7 +140,14 @@ func ComponentAttributeUpdateHandler(c *gin.Context) { } for key, items := range redisUpdateMap { - hset := diagram.NewRedisHash(ctx, key, 5000, false) + hset, err := diagram.NewRedisHash(ctx, key, 5000, false) + if err != nil { + logger.Error(ctx, "create redis hash failed", "hash_key", key, "error", err) + for _, item := range items { + updateResults[item.token] = errcode.ErrCacheSyncWarn.WithCause(err) + } + continue + } fields := make(map[string]any, len(items)) for _, item := range items { diff --git a/handler/data_object_attribute_query.go b/handler/data_object_attribute_query.go new file mode 100644 index 0000000..d4592e9 --- /dev/null +++ b/handler/data_object_attribute_query.go @@ -0,0 +1,337 @@ +// Package handler provides HTTP handlers for various endpoints. +package handler + +import ( + "context" + "errors" + "fmt" + "strings" + + "modelRT/common" + "modelRT/common/errcode" + "modelRT/constants" + "modelRT/database" + "modelRT/diagram" + "modelRT/logger" + "modelRT/model" + "modelRT/orm" + + "github.com/gin-gonic/gin" +) + +// 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 { + logger.Error(ctx, "query token from query parameters failed", "error", err, "url", c.Request.RequestURI) + renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) + return + } + + dataObjectType, err := model.ClassifyDataObjectToken(token) + if err != nil { + logger.Error(ctx, "classify data object token failed", "error", err, "url", c.Request.RequestURI) + renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) + return + } + + if err := validateDataObjectField(dataObjectType, field); err != nil { + logger.Warn(ctx, "validate data object field failed", "token", token, "field", field, "error", err) + renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) + 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) + 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 + } + } + + 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) + } +} + +func parseDataObjectAttributeQuery(c *gin.Context) (string, string, error) { + token := c.Query("token") + if token == "" { + return "", "", fmt.Errorf("token is missing from query parameters") + } + + field := strings.ToLower(c.Query("field")) + if field == "" { + field = "value" + } + return token, field, nil +} + +type measurementValueLoader func(context.Context, orm.JSONMap) (any, error) + +type parameterValueLoader func(context.Context, *database.ParameterDataObject) (any, error) + +type parameterDescriptionLoader func(context.Context, string) (string, error) + +var measurementDataObjectFields = map[string]struct{}{ + "value": {}, + "mode": {}, + "meta": {}, + "type": {}, + "name": {}, + "description": {}, + "id": {}, + "size": {}, + "data_source": {}, + "event_plan": {}, + "binding": {}, +} + +var parameterDataObjectFields = map[string]struct{}{ + "value": {}, + "name": {}, + "meta": {}, + "type": {}, + "description": {}, + "id": {}, +} + +type dataObjectAttributeQueryResult struct { + Token string `json:"token"` + Field string `json:"field"` + Code int `json:"code"` + Msg string `json:"msg"` + Value any `json:"value"` +} + +func validateDataObjectField(dataObjectType constants.DataObjectType, field string) error { + field = strings.ToLower(field) + switch dataObjectType { + case constants.DataObjectTypeMeasurement: + if _, ok := measurementDataObjectFields[field]; ok { + return nil + } + return fmt.Errorf("%w: %s", common.ErrUnsupportedMeasurementField, field) + case constants.DataObjectTypeParameter: + if _, ok := parameterDataObjectFields[field]; ok { + return nil + } + return fmt.Errorf("%w: %s", common.ErrUnsupportedParameterField, field) + default: + return fmt.Errorf("invalid data object type %q", dataObjectType) + } +} + +func buildParameterAttributeValue( + ctx context.Context, + field string, + parameter *database.ParameterDataObject, + loadValue parameterValueLoader, + loadDescription parameterDescriptionLoader, +) (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") + } + 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 { + return nil, fmt.Errorf("measurement value loader is nil") + } + return loadValue(ctx, measurement.DataSource) + case "mode": + return measurement.Mode, nil + case "meta": + return "MEASUREMENT", nil + case "type": + return model.MeasurementTypeFromDataSource(measurement.DataSource) + 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 + 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 + default: + return nil, fmt.Errorf("%w: %s", common.ErrUnsupportedMeasurementField, field) + } +} + +func queryMeasurementRealtimeValue(ctx context.Context, dataSource orm.JSONMap) (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) + if err != nil { + return nil, fmt.Errorf("query real-time measurement value by key %q: %w", queryKey, err) + } + return value, nil +} diff --git a/handler/data_object_attribute_query_test.go b/handler/data_object_attribute_query_test.go new file mode 100644 index 0000000..5ff9efc --- /dev/null +++ b/handler/data_object_attribute_query_test.go @@ -0,0 +1,276 @@ +package handler + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "modelRT/common" + "modelRT/constants" + "modelRT/database" + "modelRT/orm" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestParseDataObjectAttributeQuery(t *testing.T) { + const token = "nspath.component.bay.measurement" + + tests := []struct { + name string + target string + wantToken string + wantField string + wantErr string + }{ + { + name: "reads token and field from query parameters", + target: "/data-object/attribute?token=" + token + "&field=NAME", + wantToken: token, + wantField: "name", + }, + { + name: "defaults missing field to value", + target: "/data-object/attribute?token=" + token, + wantToken: token, + wantField: "value", + }, + { + name: "defaults empty field to value", + target: "/data-object/attribute?token=" + token + "&field=", + wantToken: token, + wantField: "value", + }, + { + name: "rejects missing token", + target: "/data-object/attribute?field=value", + wantErr: "token is missing from query parameters", + }, + { + name: "rejects empty token", + target: "/data-object/attribute?token=&field=value", + wantErr: "token is missing from query parameters", + }, + { + name: "does not read legacy path parameters", + target: "/data-object/attribute/" + token + "/value", + wantErr: "token is missing from query parameters", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Request = httptest.NewRequest(http.MethodGet, tt.target, nil) + + token, field, err := parseDataObjectAttributeQuery(ctx) + if tt.wantErr != "" { + require.EqualError(t, err, tt.wantErr) + assert.Empty(t, token) + assert.Empty(t, field) + return + } + require.NoError(t, err) + assert.Equal(t, tt.wantToken, token) + assert.Equal(t, tt.wantField, field) + }) + } +} + +func TestValidateDataObjectField(t *testing.T) { + tests := []struct { + name string + token string + dataObjectType constants.DataObjectType + field string + wantErr error + }{ + {name: "bay value", token: "nspath.component.bay.measurement", dataObjectType: constants.DataObjectTypeMeasurement, field: "value"}, + {name: "bay name", token: "nspath.component.bay.measurement", dataObjectType: constants.DataObjectTypeMeasurement, field: "name"}, + {name: "bay binding", token: "nspath.component.bay.measurement", dataObjectType: constants.DataObjectTypeMeasurement, field: "binding"}, + {name: "bay field is case insensitive", token: "nspath.component.bay.measurement", dataObjectType: constants.DataObjectTypeMeasurement, field: "DATA_SOURCE"}, + {name: "bay unsupported", token: "nspath.component.bay.measurement", dataObjectType: constants.DataObjectTypeMeasurement, field: "unknown", wantErr: common.ErrUnsupportedMeasurementField}, + {name: "parameter value", token: "nspath.component.rated.attribute", dataObjectType: constants.DataObjectTypeParameter, field: "value"}, + {name: "parameter name", token: "nspath.component.rated.attribute", dataObjectType: constants.DataObjectTypeParameter, field: "name"}, + {name: "parameter rejects mode", token: "nspath.component.rated.attribute", dataObjectType: constants.DataObjectTypeParameter, field: "mode", wantErr: common.ErrUnsupportedParameterField}, + {name: "component name", token: "nspath.component.component.name", dataObjectType: constants.DataObjectTypeParameter, field: "name"}, + {name: "component rejects mode", token: "nspath.component.component.name", dataObjectType: constants.DataObjectTypeParameter, field: "mode", wantErr: common.ErrUnsupportedParameterField}, + {name: "parameter rejects size", token: "nspath.component.rated.attribute", dataObjectType: constants.DataObjectTypeParameter, field: "size", wantErr: common.ErrUnsupportedParameterField}, + {name: "invalid type", token: "token", dataObjectType: constants.DataObjectType("unknown"), field: "value", wantErr: assert.AnError}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := validateDataObjectField(tt.dataObjectType, tt.field) + if tt.wantErr == nil { + require.NoError(t, err) + return + } + require.Error(t, err) + if tt.wantErr != assert.AnError { + assert.ErrorIs(t, err, tt.wantErr) + } + }) + } +} + +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 + } + + 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( + context.Background(), + "unknown", + &database.ParameterDataObject{}, + nil, + nil, + ) + require.Error(t, err) + assert.ErrorIs(t, err, common.ErrUnsupportedParameterField) +} + +func TestBuildMeasurementAttributeValue(t *testing.T) { + dataSource := orm.JSONMap{ + "type": float64(1), + "io_address": map[string]any{ + "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相电流", + Mode: 1, + Size: 10, + DataSource: dataSource, + EventPlan: eventPlan, + Binding: binding, + } + component := &orm.Component{ + GridName: "grid000", + ZoneName: "zone000", + StationName: "station000", + NSPath: "110kV_TV", + Tag: "cable_22", + } + + 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) + }) + } +} + +func TestBuildMeasurementAttributeValueRejectsUnsupportedField(t *testing.T) { + _, err := buildMeasurementAttributeValue( + context.Background(), + "unknown", + &orm.Measurement{}, + &orm.Component{}, + nil, + ) + require.Error(t, err) + assert.ErrorIs(t, err, common.ErrUnsupportedMeasurementField) +} + +func TestBuildMeasurementAttributeValueMode(t *testing.T) { + component := &orm.Component{} + tests := []struct { + name string + mode int16 + expected int16 + }{ + {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}, + } + + 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, + ) + require.NoError(t, err) + assert.Equal(t, tt.expected, actual) + }) + } +} diff --git a/handler/data_object_attribute_update.go b/handler/data_object_attribute_update.go new file mode 100644 index 0000000..828b3eb --- /dev/null +++ b/handler/data_object_attribute_update.go @@ -0,0 +1,425 @@ +// Package handler provides HTTP handlers for various endpoints. +package handler + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "strconv" + "strings" + "time" + + "modelRT/common" + "modelRT/common/errcode" + "modelRT/constants" + "modelRT/database" + "modelRT/diagram" + "modelRT/logger" + "modelRT/model" + "modelRT/orm" + + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +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. +func DataObjectAttributeUpdateHandler(c *gin.Context) { + ctx := c.Request.Context() + var request dataObjectAttributeUpdateRequest + if err := c.ShouldBindJSON(&request); err != nil { + logger.Error(ctx, "unmarshal data-object update request failed", "error", err) + renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) + return + } + + dataObjectType, field, value, err := validateDataObjectAttributeUpdate(request) + if err != nil { + logger.Warn(ctx, "validate data-object update request failed", "token", request.Token, "field", request.Field, "error", err) + renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) + return + } + + tx := database.GetPostgresDBClient().WithContext(ctx).Begin() + if tx.Error != nil { + logger.Error(ctx, "begin data-object update transaction failed", "error", tx.Error) + 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 + switch dataObjectType { + case constants.DataObjectTypeParameter: + parameter, queryErr := database.QueryParameterByDataObjectToken(ctx, tx, request.Token) + if queryErr == nil { + queryErr = database.UpdateParameterDataObjectValue(ctx, tx, parameter, 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, + }) + message = measurementResult.message + default: + err = fmt.Errorf("unsupported data object type %q", dataObjectType) + } + + if err != nil { + _ = tx.Rollback().Error + if measurementResult.recordFailure { + 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.Warn(ctx, "update data-object attribute failed", "token", request.Token, "field", field, "error", err) + if isInvalidDataObjectUpdateError(err) { + renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) + return + } + renderRespFailure(c, constants.RespCodeFailed, err.Error(), nil) + return + } + + if err := tx.Commit().Error; err != nil { + 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 + } + transactionCompleted = true + + renderRespSuccess(c, constants.RespCodeSuccess, message, map[string]any{ + "token": request.Token, + "field": field, + "value": value, + }) +} + +func validateDataObjectAttributeUpdate(request dataObjectAttributeUpdateRequest) (constants.DataObjectType, string, any, error) { + if request.Token == "" { + return "", "", nil, fmt.Errorf("token is required") + } + if len(bytes.TrimSpace(request.Value)) == 0 || bytes.Equal(bytes.TrimSpace(request.Value), []byte("null")) { + return "", "", nil, fmt.Errorf("value is required") + } + + field := strings.ToLower(strings.TrimSpace(request.Field)) + if field == "" { + field = "value" + } + + dataObjectType, err := model.ClassifyDataObjectToken(request.Token) + if err != nil { + return "", "", nil, err + } + + switch dataObjectType { + case constants.DataObjectTypeParameter: + parts := strings.Split(request.Token, ".") + attributeGroup := parts[len(parts)-2] + if !isWritableParameterAttributeGroup(attributeGroup) { + return "", "", nil, fmt.Errorf("parameter updates do not support token6=%s", attributeGroup) + } + if field != "value" { + return "", "", nil, fmt.Errorf("parameter data objects only support updating field value") + } + value, err := decodeDataObjectUpdateValue(request.Value) + 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 token4.token7 or token6=bay") + } + switch field { + case "value": + value, err := parseMeasurementUpdateValue(request.Value) + return dataObjectType, field, value, err + case "mode": + value, err := parseMeasurementUpdateMode(request.Value) + return dataObjectType, field, value, err + default: + return "", "", nil, fmt.Errorf("measurement data objects only support updating fields value and mode") + } + default: + return "", "", nil, fmt.Errorf("unsupported data object type %q", dataObjectType) + } +} + +func isWritableParameterAttributeGroup(group string) bool { + switch group { + case "rated", "setup", "model", "stable", "craft", "integrity", "behavior", "base_extend": + return true + default: + return false + } +} + +func decodeDataObjectUpdateValue(raw json.RawMessage) (any, error) { + decoder := json.NewDecoder(bytes.NewReader(raw)) + decoder.UseNumber() + var value any + if err := decoder.Decode(&value); err != nil { + return nil, fmt.Errorf("decode update value: %w", err) + } + if number, ok := value.(json.Number); ok { + if integer, err := number.Int64(); err == nil { + return integer, nil + } + decimal, err := number.Float64() + if err != nil { + return nil, fmt.Errorf("invalid numeric update value %q: %w", number, err) + } + return decimal, nil + } + return value, nil +} + +func parseMeasurementUpdateValue(raw json.RawMessage) (float64, error) { + var number float64 + if err := json.Unmarshal(raw, &number); err == nil { + return number, nil + } + + var text string + if err := json.Unmarshal(raw, &text); err != nil { + return 0, fmt.Errorf("measurement value must be a number or numeric string") + } + number, err := strconv.ParseFloat(text, 64) + if err != nil { + return 0, fmt.Errorf("measurement value %q is not numeric: %w", text, err) + } + return number, 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)") + } + 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 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 + measurementID int64 + value float64 + recordFailure bool +} + +func updateMeasurementDataObject( + ctx context.Context, + tx *gorm.DB, + token, field string, + value any, + modeData json.RawMessage, + dependencies measurementUpdateDependencies, +) (measurementUpdateResult, error) { + measurement, _, err := database.QueryMeasurementByDataObjectToken(ctx, tx, token) + if err != nil { + return measurementUpdateResult{}, err + } + + 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.(int16) + if !ok { + return measurementUpdateResult{}, fmt.Errorf("measurement mode has invalid type %T", value) + } + 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": + 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) + if !ok { + return measurementUpdateResult{}, fmt.Errorf("measurement value has invalid type %T", value) + } + failureResult := measurementUpdateResult{ + measurementID: lockedMeasurement.ID, + value: manualValue, + recordFailure: true, + } + if dependencies.writeManualValueFunc == nil { + return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(fmt.Errorf("measurement manual value writer is nil")) + } + if err := dependencies.writeManualValueFunc(ctx, &lockedMeasurement, manualValue); err != nil { + return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(err) + } + if dependencies.updateDataRTFunc == nil { + return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(fmt.Errorf("measurement dataRT updater is 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 { + return failureResult, errcode.ErrMeasurementValueUpdateFailed.WithCause(err) + } + return measurementUpdateResult{message: "measurement manual value updated"}, nil + default: + return measurementUpdateResult{}, fmt.Errorf("unsupported measurement update field %q", field) + } +} + +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) || + errors.Is(err, common.ErrAmbiguousParameterToken) || + errors.Is(err, common.ErrInvalidMeasurementToken) || + errors.Is(err, common.ErrMeasurementTokenNotFound) || + 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. + return nil +} + +func callRealTimeDataWriteStartInterface(_ context.Context, _ orm.JSONMap, _ *float64) error { + // TODO: call the dataRT HTTP API to start automatic measurement writes. + return nil +} + +func replaceMeasurementRedisValue(ctx context.Context, measurement *orm.Measurement, value float64) error { + key, err := model.GenerateMeasureIdentifier(measurement.DataSource) + if err != nil { + return fmt.Errorf("generate measurement redis key: %w", err) + } + zset, err := diagram.NewRedisZSet(ctx, key, 0, false) + if err != nil { + return fmt.Errorf("create measurement redis zset: %w", err) + } + if err := zset.ZREPLACE(key, value, strconv.FormatInt(time.Now().UnixNano(), 10)); err != nil { + return fmt.Errorf("replace manual measurement value in redis: %w", err) + } + return nil +} diff --git a/handler/data_object_attribute_update_test.go b/handler/data_object_attribute_update_test.go new file mode 100644 index 0000000..03bb246 --- /dev/null +++ b/handler/data_object_attribute_update_test.go @@ -0,0 +1,474 @@ +package handler + +import ( + "context" + "encoding/json" + "fmt" + "testing" + + "modelRT/common/errcode" + "modelRT/constants" + "modelRT/orm" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/postgres" + "gorm.io/gorm" +) + +func TestValidateDataObjectAttributeUpdateParameterGroups(t *testing.T) { + groups := []string{ + "rated", + "setup", + "model", + "stable", + "craft", + "integrity", + "behavior", + } + + for _, group := range groups { + t.Run(group, func(t *testing.T) { + request := dataObjectAttributeUpdateRequest{ + Token: fmt.Sprintf("nspath.component.%s.attribute", group), + Field: "VALUE", + Value: json.RawMessage(`"15.2"`), + } + + dataObjectType, field, value, err := validateDataObjectAttributeUpdate(request) + require.NoError(t, err) + assert.Equal(t, constants.DataObjectTypeParameter, dataObjectType) + assert.Equal(t, "value", field) + assert.Equal(t, "15.2", value) + }) + } +} + +func TestValidateDataObjectAttributeUpdateRejectsUnsupportedParameterGroups(t *testing.T) { + for _, group := range []string{"component", "base_extend"} { + t.Run(group, func(t *testing.T) { + _, _, _, err := validateDataObjectAttributeUpdate(dataObjectAttributeUpdateRequest{ + Token: fmt.Sprintf("nspath.component.%s.attribute", group), + Field: "value", + Value: json.RawMessage(`"uuid"`), + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "do not support token6="+group) + }) + } +} + +func TestValidateDataObjectAttributeUpdateMeasurementFields(t *testing.T) { + tests := []struct { + name string + field string + value string + expected any + wantError bool + }{ + {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: "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}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dataObjectType, field, value, err := validateDataObjectAttributeUpdate(dataObjectAttributeUpdateRequest{ + Token: "nspath.component.bay.measurement", + Field: tt.field, + Value: json.RawMessage(tt.value), + }) + if tt.wantError { + require.Error(t, err) + return + } + require.NoError(t, err) + assert.Equal(t, constants.DataObjectTypeMeasurement, dataObjectType) + assert.Equal(t, tt.field, field) + assert.Equal(t, tt.expected, value) + }) + } +} + +func TestValidateDataObjectAttributeUpdateAcceptsToken4Token7Measurement(t *testing.T) { + dataObjectType, field, value, err := validateDataObjectAttributeUpdate(dataObjectAttributeUpdateRequest{ + Token: "nspath.measurement", + Value: json.RawMessage(`15.2`), + }) + 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) { + tests := []struct { + name string + request dataObjectAttributeUpdateRequest + }{ + {name: "missing token", request: dataObjectAttributeUpdateRequest{Field: "value", 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`)}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, _, _, err := validateDataObjectAttributeUpdate(tt.request) + require.Error(t, err) + }) + } +} + +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() + + mock.ExpectBegin() + tx := db.Begin() + require.NoError(t, tx.Error) + expectMeasurementResolution(mock, 0) + 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() + + 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() + + mock.ExpectBegin() + tx := db.Begin() + require.NoError(t, tx.Error) + expectMeasurementResolution(mock, 1) + mock.ExpectRollback() + + 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() + + mock.ExpectBegin() + tx := db.Begin() + require.NoError(t, tx.Error) + expectMeasurementResolution(mock, 1) + mock.ExpectRollback() + + _, 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) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateMeasurementDataObjectWritesValueInManualMode(t *testing.T) { + db, mock, closeDB := newDataObjectUpdateTestDB(t) + defer closeDB() + + mock.ExpectBegin() + tx := db.Begin() + require.NoError(t, tx.Error) + expectMeasurementResolution(mock, 0) + mock.ExpectExec(`UPDATE "measurement" SET "operations"=.*WHERE id = \$3`). + WithArgs(sqlmock.AnyArg(), 500, int64(10)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + called := false + writer := func(_ context.Context, measurement *orm.Measurement, value float64) error { + called = true + assert.Equal(t, int64(10), measurement.ID) + assert.Equal(t, float64(15.2), value) + return nil + } + dataRTCalled := false + dataRTWriter := func(_ context.Context, dataSource orm.JSONMap, value *float64) error { + dataRTCalled = true + 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), nil, measurementUpdateDependencies{ + writeManualValueFunc: writer, + updateDataRTFunc: dataRTWriter, + }) + require.NoError(t, err) + assert.True(t, called) + assert.True(t, dataRTCalled) + assert.Contains(t, result.message, "updated") + assert.False(t, result.recordFailure) + require.NoError(t, tx.Rollback().Error) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateMeasurementDataObjectReturnsFailureResultAndAppError(t *testing.T) { + db, mock, closeDB := newDataObjectUpdateTestDB(t) + defer closeDB() + + mock.ExpectBegin() + tx := db.Begin() + require.NoError(t, tx.Error) + expectMeasurementResolution(mock, 0) + mock.ExpectRollback() + + 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), nil, measurementUpdateDependencies{ + writeManualValueFunc: writer, + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrMeasurementValueUpdateFailed) + assert.ErrorIs(t, err, writeErr) + assert.True(t, result.recordFailure) + 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()) +} + +func newDataObjectUpdateTestDB(t *testing.T) (*gorm.DB, sqlmock.Sqlmock, func()) { + t.Helper() + sqlDB, mock, err := sqlmock.New() + require.NoError(t, err) + db, err := gorm.Open(postgres.New(postgres.Config{Conn: sqlDB}), &gorm.Config{SkipDefaultTransaction: true}) + require.NoError(t, err) + return db, mock, func() { _ = sqlDB.Close() } +} + +func expectMeasurementResolution(mock sqlmock.Sqlmock, mode int16) { + const componentUUID = "70c190f2-8a60-42a9-b143-ec5f87e0aa6b" + mock.ExpectQuery(`(?s)SELECT m\.\*.*WHERE c\.nspath = \$1.*AND m\.tag = \$2.*LIMIT 2`). + WithArgs("nspath", "measurement"). + WillReturnRows(sqlmock.NewRows([]string{ + "id", "tag", "mode", "data_source", "component_uuid", + }).AddRow(int64(10), "measurement", mode, `{"type":1,"io_address":{"station":"station","device":"device","channel":"tm1"}}`, componentUUID)) + mock.ExpectQuery(`(?s)SELECT global_uuid, nspath, tag, grid, zone, station.*WHERE global_uuid = \$1.*LIMIT 1`). + WithArgs(componentUUID). + WillReturnRows(sqlmock.NewRows([]string{"global_uuid", "nspath", "tag"}). + AddRow(componentUUID, "nspath", "component")) + 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", + }).AddRow(int64(10), "measurement", mode, `{"type":1,"io_address":{"station":"station","device":"device","channel":"tm1"}}`, componentUUID)) +} diff --git a/handler/measurement_recommend.go b/handler/data_object_recommend.go similarity index 84% rename from handler/measurement_recommend.go rename to handler/data_object_recommend.go index 40ce7bd..75554d8 100644 --- a/handler/measurement_recommend.go +++ b/handler/data_object_recommend.go @@ -13,14 +13,14 @@ import ( "github.com/gin-gonic/gin" ) -// MeasurementRecommendHandler define measurement recommend API +// DataObjectRecommendHandler define data-object recommend API // @Summary 测量点推荐(搜索框自动补全) // @Description 根据用户输入的字符串,从 Redis 中查询可能的测量点或结构路径,并提供推荐列表。 -// @Tags Measurement Recommend +// @Tags DataObject Recommend // @Accept json // @Produce json // @Param input query string true "推荐关键词,例如 'grid1' 或 'grid1.'" Example("grid1") -// @Success 200 {object} network.SuccessResponse{payload=network.MeasurementRecommendPayload} "返回推荐列表成功" +// @Success 200 {object} network.SuccessResponse{payload=network.DataObjectRecommendPayload} "返回推荐列表成功" // // @Example 200 { // "code": 200, @@ -43,25 +43,25 @@ import ( // "msg": "failed to get recommend data from redis", // } // -// @Router /measurement/recommend [get] -func MeasurementRecommendHandler(c *gin.Context) { +// @Router /data-object/recommend [get] +func DataObjectRecommendHandler(c *gin.Context) { ctx := c.Request.Context() - var request network.MeasurementRecommendRequest + var request network.DataObjectRecommendRequest if err := c.ShouldBindQuery(&request); err != nil { - logger.Error(ctx, "failed to bind measurement recommend request", "error", err) + logger.Error(ctx, "failed to bind data object recommend request", "error", err) renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), nil) return } - if err := validateMeasurementRecommendInput(request.Input); err != nil { - logger.Warn(ctx, "invalid measurement recommend input", "input", request.Input, "error", err) + if err := validateDataObjectRecommendInput(request.Input); err != nil { + logger.Warn(ctx, "invalid data object recommend input", "input", request.Input, "error", err) renderRespFailure(c, constants.RespCodeInvalidParams, err.Error(), map[string]any{ "input": request.Input, }) return } recommendResults := model.RedisSearchRecommend(ctx, request.Input) - payload := network.MeasurementRecommendPayload{ + payload := network.DataObjectRecommendPayload{ Input: request.Input, RecommendedList: make([]string, 0), } @@ -117,7 +117,7 @@ func orderedRecommendResults(recommendResults map[string]model.SearchResult) []m return results } -func validateMeasurementRecommendInput(input string) error { +func validateDataObjectRecommendInput(input string) error { if strings.Contains(input, "..") { return errors.New("input contains continuous dots") } diff --git a/handler/diagram_node_link.go b/handler/diagram_node_link.go index ffd09a3..b10b0cd 100644 --- a/handler/diagram_node_link.go +++ b/handler/diagram_node_link.go @@ -85,7 +85,12 @@ func DiagramNodeLinkHandler(c *gin.Context) { return } - prevLinkSet, currLinkSet := generateLinkSet(ctx, nodeLevel, prevNodeInfo) + prevLinkSet, currLinkSet, err := generateLinkSet(ctx, nodeLevel, prevNodeInfo) + if err != nil { + logger.Error(ctx, "create diagram link redis sets failed", "node_id", nodeID, "level", nodeLevel, "error", err) + c.JSON(http.StatusOK, network.FailureResponse{Code: http.StatusInternalServerError, Msg: err.Error()}) + return + } err = processLinkSetData(ctx, action, nodeLevel, prevLinkSet, currLinkSet, prevNodeInfo, currNodeInfo) if err != nil { c.JSON(http.StatusOK, network.FailureResponse{ @@ -113,21 +118,27 @@ func DiagramNodeLinkHandler(c *gin.Context) { }) } -func generateLinkSet(ctx context.Context, level int, prevNodeInfo orm.CircuitDiagramNodeInterface) (*diagram.RedisSet, *diagram.RedisSet) { +func generateLinkSet(ctx context.Context, level int, prevNodeInfo orm.CircuitDiagramNodeInterface) (*diagram.RedisSet, *diagram.RedisSet, error) { config, ok := linkSetConfigs[level] // level not supported if !ok { - return nil, nil + return nil, nil, nil } - currLinkSet := diagram.NewRedisSet(ctx, config.CurrKey, 0, false) + currLinkSet, err := diagram.NewRedisSet(ctx, config.CurrKey, 0, false) + if err != nil { + return nil, nil, err + } if config.PrevIsNil { - return nil, currLinkSet + return nil, currLinkSet, nil } prevLinkSetKey := fmt.Sprintf(config.PrevKeyTemplate, prevNodeInfo.GetTagName()) - prevLinkSet := diagram.NewRedisSet(ctx, prevLinkSetKey, 0, false) - return prevLinkSet, currLinkSet + prevLinkSet, err := diagram.NewRedisSet(ctx, prevLinkSetKey, 0, false) + if err != nil { + return nil, nil, err + } + return prevLinkSet, currLinkSet, nil } func processLinkSetData(ctx context.Context, action string, level int, prevLinkSet, currLinkSet *diagram.RedisSet, prevNodeInfo, currNodeInfo orm.CircuitDiagramNodeInterface) error { diff --git a/handler/measurement_load.go b/handler/measurement_load.go index ddae642..2a57f29 100644 --- a/handler/measurement_load.go +++ b/handler/measurement_load.go @@ -39,7 +39,12 @@ func MeasurementGetHandler(c *gin.Context) { return } - zset := diagram.NewRedisZSet(ctx, request.MeasurementToken, 0, false) + zset, err := diagram.NewRedisZSet(ctx, request.MeasurementToken, 0, false) + if err != nil { + logger.Error(ctx, "failed to create measurement redis zset", "measurement_token", request.MeasurementToken, "error", err) + c.JSON(http.StatusOK, network.FailureResponse{Code: http.StatusInternalServerError, Msg: err.Error()}) + return + } points, err := zset.ZRANGE(request.MeasurementToken, 0, -1) if err != nil { logger.Error(ctx, "failed to get measurement data from redis", "measurement_token", request.MeasurementToken, "error", err) diff --git a/handler/measurement_recommend_test.go b/handler/measurement_recommend_test.go index 1b29aeb..f997bb1 100644 --- a/handler/measurement_recommend_test.go +++ b/handler/measurement_recommend_test.go @@ -32,7 +32,7 @@ func TestValidateMeasurementRecommendInput(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - err := validateMeasurementRecommendInput(tt.input) + err := validateDataObjectRecommendInput(tt.input) if tt.valid && err != nil { t.Fatalf("expected valid input, got error %v", err) } diff --git a/handler/mesurement_link.go b/handler/mesurement_link.go index 8737840..ae5b086 100644 --- a/handler/mesurement_link.go +++ b/handler/mesurement_link.go @@ -75,9 +75,19 @@ func MeasurementLinkHandler(c *gin.Context) { return } - allMeasSet := diagram.NewRedisSet(ctx, constants.RedisAllMeasTagSetKey, 0, false) + allMeasSet, err := diagram.NewRedisSet(ctx, constants.RedisAllMeasTagSetKey, 0, false) + if err != nil { + logger.Error(ctx, "create all-measurement redis set failed", "error", err) + c.JSON(http.StatusOK, network.FailureResponse{Code: http.StatusInternalServerError, Msg: err.Error()}) + return + } compMeasLinkKey := fmt.Sprintf(constants.RedisSpecCompTagMeasSetKey, componentInfo.Tag) - compMeasLinkSet := diagram.NewRedisSet(ctx, compMeasLinkKey, 0, false) + compMeasLinkSet, err := diagram.NewRedisSet(ctx, compMeasLinkKey, 0, false) + if err != nil { + logger.Error(ctx, "create component-measurement redis set failed", "set_key", compMeasLinkKey, "error", err) + c.JSON(http.StatusOK, network.FailureResponse{Code: http.StatusInternalServerError, Msg: err.Error()}) + return + } switch action { case constants.SearchLinkAddAction: diff --git a/logger/caller_test.go b/logger/caller_test.go new file mode 100644 index 0000000..035d41f --- /dev/null +++ b/logger/caller_test.go @@ -0,0 +1,50 @@ +package logger_test + +import ( + "context" + "io" + "os" + "strings" + "testing" + "time" + + "modelRT/config" + "modelRT/constants" + "modelRT/logger" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestCallerPointsToBusinessCode(t *testing.T) { + reader, writer, err := os.Pipe() + require.NoError(t, err) + originalStdout := os.Stdout + os.Stdout = writer + t.Cleanup(func() { + os.Stdout = originalStdout + _ = reader.Close() + _ = writer.Close() + }) + + logger.InitLoggerInstance(config.LoggerConfig{ + Mode: constants.DevelopmentLogMode, + Level: "info", + }) + + logger.Info(context.Background(), "facade caller test") + logger.NewGormLogger().Trace(context.Background(), time.Now(), func() (string, int64) { + return "SELECT 1", 1 + }, nil) + require.NoError(t, writer.Close()) + + outputBytes, err := io.ReadAll(reader) + require.NoError(t, err) + output := string(outputBytes) + + assert.NotContains(t, output, "logger/facede.go") + assert.NotContains(t, output, `"func":"modelRT/logger.Info"`) + assert.NotContains(t, output, "gorm.io/gorm") + assert.Contains(t, output, "caller_test.go") + assert.True(t, strings.Count(output, "modelRT/logger_test.TestCallerPointsToBusinessCode") >= 2) +} diff --git a/logger/facede.go b/logger/facede.go index f2ea9c4..0e27628 100644 --- a/logger/facede.go +++ b/logger/facede.go @@ -18,8 +18,6 @@ type facade struct { _logger *zap.Logger } -const facadeCallerSkip = 2 - // Debug define facade func of debug level log func Debug(ctx context.Context, msg string, kv ...any) { logFacade().log(ctx, zapcore.DebugLevel, msg, kv...) @@ -41,16 +39,17 @@ func Error(ctx context.Context, msg string, kv ...any) { } func (f *facade) log(ctx context.Context, lvl zapcore.Level, msg string, kv ...any) { - f.logSkip(ctx, lvl, 1, msg, kv...) + f.logSkip(ctx, lvl, 0, msg, kv...) } func (f *facade) logSkip(ctx context.Context, lvl zapcore.Level, extraSkip int, msg string, kv ...any) { - fields := makeLogFieldsSkip(ctx, extraSkip, kv...) - logger := f._logger - if extraSkip > 0 { - logger = logger.WithOptions(zap.AddCallerSkip(extraSkip)) + caller := resolveLoggerCaller(extraSkip) + fields := makeLogFieldsWithCaller(ctx, caller, kv...) + ce := f._logger.Check(lvl, msg) + if ce == nil { + return } - ce := logger.Check(lvl, msg) + setCheckedEntryCaller(ce, caller) ce.Write(fields...) } @@ -72,7 +71,7 @@ func InfoSkip(ctx context.Context, extraSkip int, msg string, kv ...any) { func logFacade() *facade { fOnce.Do(func() { f = &facade{ - _logger: GetLoggerInstance().WithOptions(zap.AddCallerSkip(facadeCallerSkip)), + _logger: GetLoggerInstance(), } }) return f diff --git a/logger/gorm_logger.go b/logger/gorm_logger.go index 2387b5f..59d7979 100644 --- a/logger/gorm_logger.go +++ b/logger/gorm_logger.go @@ -50,12 +50,12 @@ func (l *GormLogger) Trace(ctx context.Context, begin time.Time, fc func() (sql // get gorm exec sql and rows affected sql, rows := fc() if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { - ErrorSkip(ctx, 1, "SQL ERROR", "sql", sql, "rows", rows, "dur(ms)", duration) + ErrorSkip(ctx, 0, "SQL ERROR", "sql", sql, "rows", rows, "dur(ms)", duration) return } if duration > l.SlowThreshold.Milliseconds() { - WarnSkip(ctx, 1, "SQL SLOW", "sql", sql, "rows", rows, "dur(ms)", duration) + WarnSkip(ctx, 0, "SQL SLOW", "sql", sql, "rows", rows, "dur(ms)", duration) } else { - InfoSkip(ctx, 1, "SQL INFO", "sql", sql, "rows", rows, "dur(ms)", duration) + InfoSkip(ctx, 0, "SQL INFO", "sql", sql, "rows", rows, "dur(ms)", duration) } } diff --git a/logger/logger.go b/logger/logger.go index f7fa247..f40700d 100644 --- a/logger/logger.go +++ b/logger/logger.go @@ -5,6 +5,7 @@ import ( "context" "path" "runtime" + "strings" "go.opentelemetry.io/otel/trace" "go.uber.org/zap" @@ -41,8 +42,13 @@ func (l *logger) Error(msg string, kv ...any) { } func (l *logger) log(lvl zapcore.Level, msg string, kv ...any) { - fields := makeLogFields(l.ctx, kv...) + caller := resolveLoggerCaller(0) + fields := makeLogFieldsWithCaller(l.ctx, caller, kv...) ce := l._logger.Check(lvl, msg) + if ce == nil { + return + } + setCheckedEntryCaller(ce, caller) ce.Write(fields...) } @@ -51,6 +57,10 @@ func makeLogFields(ctx context.Context, kv ...any) []zap.Field { } func makeLogFieldsSkip(ctx context.Context, extraSkip int, kv ...any) []zap.Field { + return makeLogFieldsWithCaller(ctx, resolveLoggerCaller(extraSkip), kv...) +} + +func makeLogFieldsWithCaller(ctx context.Context, caller loggerCaller, kv ...any) []zap.Field { if len(kv)%2 != 0 { kv = append(kv, "unknown") } @@ -60,8 +70,7 @@ func makeLogFieldsSkip(ctx context.Context, extraSkip int, kv ...any) []zap.Fiel spanID := spanCtx.SpanID().String() kv = append(kv, "traceID", traceID, "spanID", spanID) - funcName, file, line := getLoggerCallerInfoSkip(extraSkip) - kv = append(kv, "func", funcName, "file", file, "line", line) + kv = append(kv, "func", caller.funcName, "file", caller.shortFile, "line", caller.line) fields := make([]zap.Field, 0, len(kv)/2) for i := 0; i < len(kv); i += 2 { key := kv[i].(string) @@ -95,13 +104,59 @@ func getLoggerCallerInfo() (funcName, file string, line int) { // getLoggerCallerInfoSkip returns caller info with additional skip frames beyond the standard depth. func getLoggerCallerInfoSkip(extraSkip int) (funcName, file string, line int) { - pc, file, line, ok := runtime.Caller(4 + extraSkip) - if !ok { + caller := resolveLoggerCaller(extraSkip) + return caller.funcName, caller.shortFile, caller.line +} + +type loggerCaller struct { + pc uintptr + funcName string + fullFile string + shortFile string + line int +} + +func resolveLoggerCaller(extraSkip int) loggerCaller { + pcs := make([]uintptr, 32) + count := runtime.Callers(2, pcs) + frames := runtime.CallersFrames(pcs[:count]) + + for { + frame, more := frames.Next() + if !isLoggerInfrastructureFrame(frame.Function) { + if extraSkip > 0 { + extraSkip-- + } else { + return loggerCaller{ + pc: frame.PC, + funcName: frame.Function, + fullFile: frame.File, + shortFile: path.Base(frame.File), + line: frame.Line, + } + } + } + if !more { + return loggerCaller{} + } + } +} + +func isLoggerInfrastructureFrame(function string) bool { + return strings.HasPrefix(function, "modelRT/logger.") || + strings.HasPrefix(function, "gorm.io/gorm") +} + +func setCheckedEntryCaller(entry *zapcore.CheckedEntry, caller loggerCaller) { + if caller.pc == 0 { return } - file = path.Base(file) - funcName = runtime.FuncForPC(pc).Name() - return + entry.Entry.Caller = zapcore.EntryCaller{ + Defined: true, + PC: caller.pc, + File: caller.fullFile, + Line: caller.line, + } } // New returns a logger bound to ctx. Trace fields (traceID, spanID) are extracted diff --git a/middleware/token.go b/middleware/token.go index 6759f40..9d1bbf0 100644 --- a/middleware/token.go +++ b/middleware/token.go @@ -1,12 +1,20 @@ // Package middleware define gin framework middlewares package middleware -import "github.com/gin-gonic/gin" +import ( + "context" + + "modelRT/constants" + + "github.com/gin-gonic/gin" +) // SetTokenMiddleware define a middleware for set token in context func SetTokenMiddleware(clientToken string) gin.HandlerFunc { return func(c *gin.Context) { - c.Set("client_token", clientToken) + c.Set(constants.ClientTokenContextName, clientToken) + requestCtx := context.WithValue(c.Request.Context(), constants.CtxKeyClientToken, clientToken) + c.Request = c.Request.WithContext(requestCtx) c.Next() } } diff --git a/middleware/token_test.go b/middleware/token_test.go new file mode 100644 index 0000000..9d2f01c --- /dev/null +++ b/middleware/token_test.go @@ -0,0 +1,28 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "testing" + + "modelRT/constants" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" +) + +func TestSetTokenMiddlewarePropagatesClientToken(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(SetTokenMiddleware("test-token")) + router.GET("/test", func(c *gin.Context) { + assert.Equal(t, "test-token", c.GetString(constants.ClientTokenContextName)) + assert.Equal(t, "test-token", c.Request.Context().Value(constants.CtxKeyClientToken)) + c.Status(http.StatusNoContent) + }) + + request := httptest.NewRequest(http.MethodGet, "/test", nil) + response := httptest.NewRecorder() + router.ServeHTTP(response, request) + assert.Equal(t, http.StatusNoContent, response.Code) +} diff --git a/model/attribute_group_recommend_model.go b/model/attribute_group_recommend_model.go index 22253df..7a7632f 100644 --- a/model/attribute_group_recommend_model.go +++ b/model/attribute_group_recommend_model.go @@ -33,13 +33,19 @@ func TraverseAttributeGroupTables(ctx context.Context, db *gorm.DB, compTagToFul var tableNames []string excludedTables := []string{"component", ""} + var projectTableNames []string result := db.Model(&orm.ProjectManager{}). Where("name NOT IN ?", excludedTables). - Pluck("name", &tableNames) + Pluck("name", &projectTableNames) if result.Error != nil && result.Error != gorm.ErrRecordNotFound { logger.Error(ctx, "query name column data from postgres table failed", "err", result.Error) return result.Error } + for _, tableName := range projectTableNames { + if constants.IsSupportedParameterTableName(tableName) { + tableNames = append(tableNames, tableName) + } + } if len(tableNames) == 0 { logger.Info(ctx, "query from postgres successed, but no records found") diff --git a/model/data_object_token.go b/model/data_object_token.go new file mode 100644 index 0000000..32b2f04 --- /dev/null +++ b/model/data_object_token.go @@ -0,0 +1,53 @@ +// Package model defines data models and domain rules for model runtime service. +package model + +import ( + "fmt" + "slices" + "strings" + + "modelRT/constants" +) + +var parameterAttributeGroups = map[string]struct{}{ + "component": {}, + "base_extend": {}, + "rated": {}, + "setup": {}, + "model": {}, + "stable": {}, + "craft": {}, + "integrity": {}, + "behavior": {}, +} + +// ClassifyDataObjectToken determines whether token identifies a parameter or a +// measurement. Seven-part and four-part tokens are classified by token6, while +// two-part tokens are treated as measurements at the current stage. +func ClassifyDataObjectToken(token string) (constants.DataObjectType, error) { + parts := strings.Split(token, ".") + if slices.Contains(parts, "") { + return "", fmt.Errorf("invalid data object token %q: token segment cannot be empty", token) + } + + switch len(parts) { + case 2: + return constants.DataObjectTypeMeasurement, nil + case 4, 7: + token6Index := 2 + if len(parts) == 7 { + token6Index = 5 + } + + token6 := parts[token6Index] + if _, ok := parameterAttributeGroups[token6]; ok { + return constants.DataObjectTypeParameter, nil + } + if token6 == "bay" { + return constants.DataObjectTypeMeasurement, nil + } + return "", fmt.Errorf("invalid data object token %q: unsupported token6 %q", token, token6) + default: + return "", fmt.Errorf("invalid data object token %q: expected 2, 4, or 7 segments, got %d", token, len(parts)) + } +} diff --git a/model/data_object_token_test.go b/model/data_object_token_test.go new file mode 100644 index 0000000..d26a802 --- /dev/null +++ b/model/data_object_token_test.go @@ -0,0 +1,109 @@ +package model + +import ( + "testing" + + "modelRT/constants" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestClassifyDataObjectToken(t *testing.T) { + tests := []struct { + name string + token string + expected constants.DataObjectType + wantErr string + }{ + { + name: "seven-part measurement", + token: "grid.zone.station.nspath.component.bay.measurement", + expected: constants.DataObjectTypeMeasurement, + }, + { + name: "four-part measurement", + token: "nspath.component.bay.measurement", + expected: constants.DataObjectTypeMeasurement, + }, + { + name: "two-part measurement", + token: "nspath.measurement", + expected: constants.DataObjectTypeMeasurement, + }, + { + name: "seven-part parameter", + token: "grid.zone.station.nspath.component.rated.voltage", + expected: constants.DataObjectTypeParameter, + }, + { + name: "four-part parameter", + token: "nspath.component.base_extend.description", + expected: constants.DataObjectTypeParameter, + }, + { + name: "component group", + token: "nspath.component.component.name", + expected: constants.DataObjectTypeParameter, + }, + { + name: "unknown group", + token: "nspath.component.unknown.name", + wantErr: "unsupported token6", + }, + { + name: "invalid segment count", + token: "grid.zone.station", + wantErr: "expected 2, 4, or 7 segments", + }, + { + name: "empty segment", + token: "nspath..bay.measurement", + wantErr: "token segment cannot be empty", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + actual, err := ClassifyDataObjectToken(tt.token) + if tt.wantErr != "" { + require.Error(t, err) + assert.Contains(t, err.Error(), tt.wantErr) + assert.Empty(t, actual) + return + } + + require.NoError(t, err) + assert.Equal(t, tt.expected, actual) + }) + } +} + +func TestClassifyDataObjectTokenParameterGroups(t *testing.T) { + groups := []string{ + "component", + "base_extend", + "rated", + "setup", + "model", + "stable", + "craft", + "integrity", + "behavior", + } + + for _, group := range groups { + t.Run(group, func(t *testing.T) { + tokens := []string{ + "nspath.component." + group + ".attribute", + "grid.zone.station.nspath.component." + group + ".attribute", + } + + for _, token := range tokens { + actual, err := ClassifyDataObjectToken(token) + require.NoError(t, err) + assert.Equal(t, constants.DataObjectTypeParameter, actual) + } + }) + } +} diff --git a/model/measurement_attribute.go b/model/measurement_attribute.go new file mode 100644 index 0000000..51c0cc0 --- /dev/null +++ b/model/measurement_attribute.go @@ -0,0 +1,78 @@ +package model + +import ( + "fmt" + "strings" + + "modelRT/constants" + "modelRT/orm" +) + +var allowedMeasurementTypes = map[string]struct{}{ + "TM": {}, + "TS": {}, + "TC": {}, + "TA": {}, + "SP": {}, +} + +// MeasurementTypeFromDataSource returns the two-character measurement type +// encoded in a CL3611 channel. Only TM, TS, TC, TA, and SP are valid. +func MeasurementTypeFromDataSource(dataSource orm.JSONMap) (string, error) { + dataSourceType, err := integerJSONValue(dataSource["type"]) + if err != nil { + return "", fmt.Errorf("invalid measurement data_source type: %w", err) + } + if dataSourceType != constants.DataSourceTypeCL3611 { + return "", fmt.Errorf("measurement type requires data_source type %d, got %d", constants.DataSourceTypeCL3611, dataSourceType) + } + + ioAddress, ok := dataSource["io_address"].(map[string]any) + if !ok { + if value, jsonMapOK := dataSource["io_address"].(orm.JSONMap); jsonMapOK { + ioAddress = map[string]any(value) + } else { + return "", fmt.Errorf("measurement data_source io_address is not an object") + } + } + + channel, ok := ioAddress["channel"].(string) + if !ok || len(channel) < 2 { + return "", fmt.Errorf("measurement data_source channel must contain at least two characters") + } + + measurementType := strings.ToUpper(channel[:2]) + if _, ok := allowedMeasurementTypes[measurementType]; !ok { + return "", fmt.Errorf("unsupported measurement type %q", measurementType) + } + return measurementType, nil +} + +func integerJSONValue(value any) (int, error) { + switch typed := value.(type) { + case int: + return typed, nil + case int8: + return int(typed), nil + case int16: + return int(typed), nil + case int32: + return int(typed), nil + case int64: + return int(typed), nil + case float32: + converted := int(typed) + if typed != float32(converted) { + return 0, fmt.Errorf("expected integer, got %v", typed) + } + return converted, nil + case float64: + converted := int(typed) + if typed != float64(converted) { + return 0, fmt.Errorf("expected integer, got %v", typed) + } + return converted, nil + default: + return 0, fmt.Errorf("expected integer, got %T", value) + } +} diff --git a/model/measurement_attribute_test.go b/model/measurement_attribute_test.go new file mode 100644 index 0000000..8ca8e38 --- /dev/null +++ b/model/measurement_attribute_test.go @@ -0,0 +1,40 @@ +package model + +import ( + "strings" + "testing" + + "modelRT/orm" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestMeasurementTypeFromDataSource(t *testing.T) { + for _, measurementType := range []string{"TM", "TS", "TC", "TA", "SP"} { + t.Run(measurementType, func(t *testing.T) { + actual, err := MeasurementTypeFromDataSource(orm.JSONMap{ + "type": float64(1), + "io_address": map[string]any{ + "channel": strings.ToLower(measurementType) + "1_test", + }, + }) + require.NoError(t, err) + assert.Equal(t, measurementType, actual) + }) + } +} + +func TestMeasurementTypeFromDataSourceRejectsInvalidValues(t *testing.T) { + tests := []orm.JSONMap{ + {"type": float64(2), "io_address": map[string]any{"channel": "tm1"}}, + {"type": float64(1), "io_address": map[string]any{"channel": "xx1"}}, + {"type": float64(1), "io_address": map[string]any{"channel": "t"}}, + {"type": "1", "io_address": map[string]any{"channel": "tm1"}}, + } + + for _, dataSource := range tests { + _, err := MeasurementTypeFromDataSource(dataSource) + require.Error(t, err) + } +} diff --git a/model/measurement_identifier_test.go b/model/measurement_identifier_test.go new file mode 100644 index 0000000..aa8925b --- /dev/null +++ b/model/measurement_identifier_test.go @@ -0,0 +1,35 @@ +package model + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGenerateMeasureIdentifierSupportsJSONNumbers(t *testing.T) { + identifier, err := GenerateMeasureIdentifier(map[string]any{ + "type": float64(2), + "io_address": map[string]any{ + "station": "Station000", + "packet": float64(10), + "offset": float64(35), + }, + }) + + require.NoError(t, err) + assert.Equal(t, "station000:104:10:35", identifier) +} + +func TestGenerateMeasureIdentifierRejectsFractionalJSONNumbers(t *testing.T) { + _, err := GenerateMeasureIdentifier(map[string]any{ + "type": float64(2), + "io_address": map[string]any{ + "station": "station000", + "packet": float64(10.5), + "offset": float64(35), + }, + }) + + require.Error(t, err) +} diff --git a/model/measurement_protol_model.go b/model/measurement_protol_model.go index 228f51f..cfc4a1f 100644 --- a/model/measurement_protol_model.go +++ b/model/measurement_protol_model.go @@ -245,24 +245,18 @@ func GenerateMeasureIdentifier(source map[string]any) (string, error) { if !ok { return "", fmt.Errorf("Power104: missing packet field") } - var packet int - switch v := packetVal.(type) { - case int: - packet = v - default: - return "", fmt.Errorf("Power104:invalid packet format") + packet, err := integerJSONValue(packetVal) + if err != nil { + return "", fmt.Errorf("Power104:invalid packet format: %w", err) } offsetVal, ok := ioAddress["offset"] if !ok { return "", fmt.Errorf("Power104:missing offset field") } - var offset int - switch v := offsetVal.(type) { - case int: - offset = v - default: - return "", fmt.Errorf("Power104:invalid offset format") + offset, err := integerJSONValue(offsetVal) + if err != nil { + return "", fmt.Errorf("Power104:invalid offset format: %w", err) } return concatP104WithPlus(station, packet, offset), nil default: diff --git a/network/data_object_request.go b/network/data_object_request.go new file mode 100644 index 0000000..479fd1e --- /dev/null +++ b/network/data_object_request.go @@ -0,0 +1,7 @@ +// Package network define struct of network operation +package network + +// DataObjectRecommendRequest defines the request payload for an data object recommend +type DataObjectRecommendRequest struct { + Input string `form:"input,omitempty" example:"grid1"` +} diff --git a/network/measurement_request.go b/network/measurement_request.go index c04704a..c73c15d 100644 --- a/network/measurement_request.go +++ b/network/measurement_request.go @@ -6,8 +6,3 @@ type MeasurementGetRequest struct { MeasurementID int64 `json:"measurement_id" example:"1001"` MeasurementToken string `json:"token" example:"some-token"` } - -// MeasurementRecommendRequest defines the request payload for an measurement recommend -type MeasurementRecommendRequest struct { - Input string `form:"input,omitempty" example:"grid1"` -} diff --git a/network/response.go b/network/response.go index 10e7589..811bbb2 100644 --- a/network/response.go +++ b/network/response.go @@ -22,13 +22,6 @@ type WSResponse struct { Payload any `json:"payload,omitempty" swaggertype:"object"` } -// MeasurementRecommendPayload define struct of represents the data payload for the successful recommendation response. -type MeasurementRecommendPayload struct { - Input string `json:"input" example:"transformfeeder1_220."` - Offset int `json:"offset" example:"21"` - RecommendedList []string `json:"recommended_list" example:"[\"I_A_rms\", \"I_B_rms\",\"I_C_rms\"]"` -} - // TargetResult define struct of target item in real time data subscription response payload type TargetResult struct { ID string `json:"id" example:"grid1.zone1.station1.ns1.tag1.transformfeeder1_220.I_A_rms"` @@ -41,3 +34,10 @@ type RealTimeSubPayload struct { ClientID string `json:"client_id" example:"5d72f2d9-e33a-4f1b-9c76-88a44b9a953e" description:"用于标识不同client的监控请求ID"` TargetResults []TargetResult `json:"targets"` } + +// DataObjectRecommendPayload define struct of represents the data payload for the successful recommendation response +type DataObjectRecommendPayload struct { + Input string `json:"input" example:"transformfeeder1_220."` + Offset int `json:"offset" example:"21"` + RecommendedList []string `json:"recommended_list" example:"[\"I_A_rms\", \"I_B_rms\",\"I_C_rms\"]"` +} diff --git a/orm/async_task.go b/orm/async_task.go index 37709bf..d2900f5 100644 --- a/orm/async_task.go +++ b/orm/async_task.go @@ -117,7 +117,7 @@ func (a *AsyncTask) IsFailed() bool { return a.Status == AsyncTaskStatusFailed } -// IsValidTaskType checks if the task type is valid +// IsValidAsyncTaskType checks if the task type is valid func IsValidAsyncTaskType(taskType string) bool { switch AsyncTaskType(taskType) { case AsyncTaskTypeTopologyAnalysis, AsyncTaskTypePerformanceAnalysis, diff --git a/orm/busbar_section.go b/orm/busbar_section.go index d953d2f..b177385 100644 --- a/orm/busbar_section.go +++ b/orm/busbar_section.go @@ -91,10 +91,7 @@ func NewBusbarSection(name string) (*BusbarSection, error) { } func (b *BusbarSection) BusNameLenCheck() bool { - if len([]rune(b.BusbarName)) > 20 { - return false - } - return true + return len([]rune(b.BusbarName)) <= 20 } func (b *BusbarSection) BusVoltageCheck() bool { @@ -105,8 +102,5 @@ func (b *BusbarSection) BusVoltageCheck() bool { } func (b *BusbarSection) BusDescLenCheck() bool { - if len([]rune(b.BusbarDesc)) > 100 { - return false - } - return true + return len([]rune(b.BusbarDesc)) <= 100 } diff --git a/orm/circuit_diagram_measurement.go b/orm/circuit_diagram_measurement.go index e281559..c52b237 100644 --- a/orm/circuit_diagram_measurement.go +++ b/orm/circuit_diagram_measurement.go @@ -9,18 +9,20 @@ import ( // Measurement structure define abstracted info set of electrical measurement type Measurement struct { - ID int64 `gorm:"column:id;primaryKey;autoIncrement"` - Tag string `gorm:"column:tag;size:64;not null;default:''"` - Name string `gorm:"column:name;size:64;not null;default:''"` - Type int16 `gorm:"column:type;not null;default:-1"` - Size int `gorm:"column:size;not null;default:-1"` - DataSource JSONMap `gorm:"column:data_source;type:jsonb;not null;default:'{}'"` - EventPlan JSONMap `gorm:"column:event_plan;type:jsonb;not null;default:'{}'"` - Binding JSONMap `gorm:"column:binding;type:jsonb;not null;default:'{\"ct\":{\"ratio\":1.0,\"polarity\":1,\"index\":0},\"pt\":{\"ratio\":1.0,\"polarity\":1,\"index\":0}}'"` - BayUUID uuid.UUID `gorm:"column:bay_uuid;type:uuid;not null"` - ComponentUUID uuid.UUID `gorm:"column:component_uuid;type:uuid;not null"` - Op int `gorm:"column:op;not null;default:-1"` - TS time.Time `gorm:"column:ts;type:timestamptz;not null;default:CURRENT_TIMESTAMP"` + ID int64 `gorm:"column:id;primaryKey;autoIncrement"` + Tag string `gorm:"column:tag;size:64;not null;default:'';uniqueIndex"` + Name string `gorm:"column:name;size:64;not null;default:''"` + Type int16 `gorm:"column:type;not null;default:-1"` + Size int `gorm:"column:size;not null;default:-1"` + Mode int16 `gorm:"column:mode;not null;default:1"` + Operations JSONMapArray `gorm:"column:operations;type:jsonb[];not null;default:'{}'"` + DataSource JSONMap `gorm:"column:data_source;type:jsonb;not null;default:'{}'"` + EventPlan JSONMap `gorm:"column:event_plan;type:jsonb;not null;default:'{}'"` + Binding JSONMap `gorm:"column:binding;type:jsonb;not null;default:'{\"ct\":{\"ratio\":1.0,\"polarity\":1,\"index\":0},\"pt\":{\"ratio\":1.0,\"polarity\":1,\"index\":0}}'"` + BayUUID uuid.UUID `gorm:"column:bay_uuid;type:uuid;not null"` + ComponentUUID uuid.UUID `gorm:"column:component_uuid;type:uuid;not null"` + Op int `gorm:"column:op;not null;default:-1"` + TS time.Time `gorm:"column:ts;type:timestamptz;not null;default:CURRENT_TIMESTAMP"` } // TableName func respresent return table name of Measurement diff --git a/orm/jsonb_serializer.go b/orm/jsonb_serializer.go index d43f7d2..c4898e4 100644 --- a/orm/jsonb_serializer.go +++ b/orm/jsonb_serializer.go @@ -5,6 +5,9 @@ import ( "database/sql/driver" "encoding/json" "errors" + "fmt" + + "github.com/jackc/pgx/v5/pgtype" ) // JSONMap define struct of implements the sql.Scanner and driver.Valuer interfaces for handling JSONB fields @@ -36,3 +39,48 @@ func (j *JSONMap) Scan(value any) error { } return json.Unmarshal(source, j) } + +// JSONMapArray represents a PostgreSQL jsonb[] column. +type JSONMapArray []JSONMap + +// Value encodes the slice as a PostgreSQL jsonb array. +func (j JSONMapArray) Value() (driver.Value, error) { + items := make(pgtype.FlatArray[map[string]any], len(j)) + for index, item := range j { + items[index] = map[string]any(item) + } + encoded, err := pgtype.NewMap().Encode(pgtype.JSONBArrayOID, pgtype.TextFormatCode, items, nil) + if err != nil { + return nil, fmt.Errorf("encode JSONMapArray: %w", err) + } + return string(encoded), nil +} + +// Scan decodes a PostgreSQL jsonb array. +func (j *JSONMapArray) Scan(value any) error { + if value == nil { + *j = nil + return nil + } + + var source []byte + switch typedValue := value.(type) { + case []byte: + source = typedValue + case string: + source = []byte(typedValue) + default: + return fmt.Errorf("unsupported data type %T for JSONMapArray Scan", value) + } + + var items pgtype.FlatArray[map[string]any] + if err := pgtype.NewMap().Scan(pgtype.JSONBArrayOID, pgtype.TextFormatCode, source, &items); err != nil { + return fmt.Errorf("decode JSONMapArray: %w", err) + } + result := make(JSONMapArray, len(items)) + for index, item := range items { + result[index] = JSONMap(item) + } + *j = result + return nil +} diff --git a/orm/jsonb_serializer_test.go b/orm/jsonb_serializer_test.go new file mode 100644 index 0000000..4478dec --- /dev/null +++ b/orm/jsonb_serializer_test.go @@ -0,0 +1,21 @@ +package orm + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestJSONMapArrayValueAndScan(t *testing.T) { + original := JSONMapArray{ + {"command": float64(0), "timestamp": "2026-07-20T00:00:00Z"}, + {"transaction": float64(1), "value": 15.2}, + } + + encoded, err := original.Value() + require.NoError(t, err) + + var decoded JSONMapArray + require.NoError(t, decoded.Scan(encoded)) + require.Equal(t, original, decoded) +} diff --git a/router/data_object.go b/router/data_object.go new file mode 100644 index 0000000..f83af20 --- /dev/null +++ b/router/data_object.go @@ -0,0 +1,17 @@ +// Package router provides router config +package router + +import ( + "modelRT/handler" + + "github.com/gin-gonic/gin" +) + +// registerDataObjectRoutes define func of register data object routes +func registerDataObjectRoutes(rg *gin.RouterGroup, middlewares ...gin.HandlerFunc) { + g := rg.Group("/data-object/") + g.Use(middlewares...) + g.GET("attribute", handler.DataObjectAttributeQueryHandler) + g.PATCH("attribute", handler.DataObjectAttributeUpdateHandler) + g.GET("recommend", handler.DataObjectRecommendHandler) +} diff --git a/router/measurement.go b/router/measurement.go index 3d94078..9ef3307 100644 --- a/router/measurement.go +++ b/router/measurement.go @@ -12,5 +12,4 @@ func registerMeasurementRoutes(rg *gin.RouterGroup, middlewares ...gin.HandlerFu g := rg.Group("/measurement/") g.Use(middlewares...) g.GET("load", handler.MeasurementGetHandler) - g.GET("recommend", handler.MeasurementRecommendHandler) } diff --git a/router/router.go b/router/router.go index f242cb9..b6ca534 100644 --- a/router/router.go +++ b/router/router.go @@ -27,5 +27,6 @@ func RegisterRoutes(engine *gin.Engine, clientToken string) { registerDataRoutes(routeGroup) registerMonitorRoutes(routeGroup) registerComponentRoutes(routeGroup, middleware.SetTokenMiddleware(clientToken)) + registerDataObjectRoutes(routeGroup, middleware.SetTokenMiddleware(clientToken)) registerAsyncTaskRoutes(routeGroup, middleware.SetTokenMiddleware(clientToken)) } diff --git a/sql/data_object_measurement.go b/sql/data_object_measurement.go new file mode 100644 index 0000000..bd734f6 --- /dev/null +++ b/sql/data_object_measurement.go @@ -0,0 +1,100 @@ +// Package sql defines reusable database SQL statements. +package sql + +const ( + // MeasurementCountSelect selects the number of measurements matching a + // token hierarchy and is used during token existence validation. + MeasurementCountSelect = "SELECT COUNT(*)" + + // MeasurementRowsSelect selects complete measurement rows after a token has + // been parsed into its hierarchy conditions. + MeasurementRowsSelect = "SELECT m.*" + + // MeasurementLimitTwo limits token resolution to two rows so callers can + // distinguish a unique match from an ambiguous token without loading all matches. + MeasurementLimitTwo = "LIMIT 2" + + // MeasurementIDWhere is the GORM condition used to query a measurement by + // its database primary key. + MeasurementIDWhere = "id = ?" + + // MeasurementTokenValidationQueryBase contains the common hierarchy joins + // used by all supported measurement token formats. + MeasurementTokenValidationQueryBase = MeasurementCountSelect + ` + FROM measurement AS m + INNER JOIN component AS c ON c.global_uuid = m.component_uuid + INNER JOIN bay AS b ON b.bay_uuid = m.bay_uuid + INNER JOIN station AS s ON s.id = c.station_id + INNER JOIN zone AS z ON z.id = s.zone_id + INNER JOIN grid AS g ON g.id = z.grid_id` + + // MeasurementSevenPartTokenWhere matches a complete token in the form + // token1.token2.token3.token4.token5.token6.token7. Token6 is validated as + // "bay" before this condition is used. + MeasurementSevenPartTokenWhere = `WHERE g.tagname = ? + AND z.tagname = ? AND s.tagname = ? + AND c.nspath = ? AND c.tag = ? + AND m.tag = ?` + + // MeasurementFourPartTokenWhere matches a local measurement token in the + // form token4.token5.token6.token7. Token6 is validated as "bay" before use. + MeasurementFourPartTokenWhere = `WHERE c.nspath = ? AND c.tag = ? + AND m.tag = ?` + + // MeasurementTwoPartTokenWhere matches the short measurement token format + // token4.token7 using component namespace path and measurement tag. + MeasurementTwoPartTokenWhere = `WHERE c.nspath = ? AND m.tag = ?` + + // MeasurementComponentByUUID returns the component hierarchy fields needed + // to construct a measurement's canonical name and seven-part ID. + MeasurementComponentByUUID = `SELECT global_uuid, nspath, + tag, grid, zone, station FROM component + WHERE global_uuid = ? LIMIT 1` + + // MeasurementGridTags returns every grid tag used to construct the first + // level of the measurement recommendation hierarchy. + MeasurementGridTags = `SELECT tagname FROM grid` + + // MeasurementZoneHierarchy returns zones together with their parent grid + // tags for building the grid-to-zone recommendation mapping. + MeasurementZoneHierarchy = `SELECT zone.*, + grid.tagname AS grid_tag FROM zone + LEFT JOIN grid ON zone.grid_id = grid.id` + + // MeasurementStationHierarchy returns stations together with their parent + // zone tags for building the zone-to-station recommendation mapping. + MeasurementStationHierarchy = `SELECT station.*, zone.tagname AS zone_tag + FROM station + LEFT JOIN zone ON station.zone_id = zone.id` + + // MeasurementComponentHierarchy returns components together with their + // parent station tags for building station, namespace, and component mappings. + MeasurementComponentHierarchy = `SELECT component.*, + station.tagname AS station_tag + FROM component LEFT JOIN station + ON component.station_id = station.id` + + // MeasurementTagHierarchy returns measurements together with their owning + // component tags for building the component-to-measurement mapping. + MeasurementTagHierarchy = ` + SELECT measurement.*, + component.tag AS comp_tag, + component.nspath AS comp_nspath, + bay.tag AS bay_tag + FROM measurement + LEFT JOIN component + ON measurement.component_uuid = component.global_uuid + LEFT JOIN bay + ON measurement.bay_uuid = bay.bay_uuid` + + // MeasurementBayLinkedComponentTags returns components that have at least + // one measurement whose bay_uuid resolves to an existing bay record. + MeasurementBayLinkedComponentTags = ` + SELECT DISTINCT component.tag AS comp_tag + FROM component + INNER JOIN measurement + ON component.global_uuid = measurement.component_uuid + INNER JOIN bay + ON measurement.bay_uuid = bay.bay_uuid + WHERE component.tag <> ''` +) diff --git a/sql/data_object_parameter.go b/sql/data_object_parameter.go new file mode 100644 index 0000000..f3e2ba2 --- /dev/null +++ b/sql/data_object_parameter.go @@ -0,0 +1,50 @@ +// Package sql defines reusable database SQL statements. +package sql + +const ( + // ParameterComponentQueryBase selects the component owning a parameter and + // joins its complete hierarchy so seven-part tokens can be validated as one path. + ParameterComponentQueryBase = `SELECT c.* + FROM component AS c + INNER JOIN station AS s ON s.id = c.station_id + INNER JOIN zone AS z ON z.id = s.zone_id + INNER JOIN grid AS g ON g.id = z.grid_id` + + // ParameterSevenPartTokenWhere matches the complete parameter token prefix + // token1.token2.token3.token4.token5. + ParameterSevenPartTokenWhere = `WHERE g.tagname = ? + AND z.tagname = ? + AND s.tagname = ? + AND c.nspath = ? + AND c.tag = ?` + + // ParameterFourPartTokenWhere matches token4.token5 and restricts the short + // token form to components belonging to a local station. + ParameterFourPartTokenWhere = `WHERE c.nspath = ? + AND c.tag = ? + AND s.is_local = TRUE` + + // ParameterLimitTwo allows callers to distinguish a unique component from + // an ambiguous token without loading every matching row. + ParameterLimitTwo = `LIMIT 2` + + // ParameterAttributeColumnType checks that token7 is an actual column of + // the dynamic parameter table and returns its PostgreSQL display type. + ParameterAttributeColumnType = `SELECT pg_catalog.format_type(a.atttypid, a.atttypmod) + FROM pg_catalog.pg_attribute AS a + INNER JOIN pg_catalog.pg_class AS c ON c.oid = a.attrelid + INNER JOIN pg_catalog.pg_namespace AS n ON n.oid = c.relnamespace + WHERE n.nspname = 'public' + AND c.relname = ? + AND a.attname = ? + AND a.attnum > 0 + AND NOT a.attisdropped + LIMIT 1` + + // ParameterAttributeDescription returns the display name of token7 from the + // basic attribute metadata table. Two rows are enough to detect ambiguity. + ParameterAttributeDescription = `SELECT attribute_name + FROM basic.attribute + WHERE attribute = ? + LIMIT 2` +) diff --git a/sql/topologic.go b/sql/topologic.go index f440e27..9c38eae 100644 --- a/sql/topologic.go +++ b/sql/topologic.go @@ -2,7 +2,7 @@ package sql // RecursiveSQL define topologic table recursive query statement -var RecursiveSQL = `WITH RECURSIVE recursive_tree as ( +const RecursiveSQL = `WITH RECURSIVE recursive_tree as ( SELECT uuid_from,uuid_to,flag FROM "topologic" WHERE uuid_from = ? @@ -16,7 +16,7 @@ var RecursiveSQL = `WITH RECURSIVE recursive_tree as ( // RecursiveTopologicByStartSQL returns every directed edge reachable from the // supplied start component. It tracks the visited node path inside PostgreSQL // so cycles in topologic data cannot recurse forever. -var RecursiveTopologicByStartSQL = `WITH RECURSIVE recursive_tree as ( +const RecursiveTopologicByStartSQL = `WITH RECURSIVE recursive_tree as ( SELECT uuid_from, uuid_to, flag, ARRAY[uuid_from, uuid_to] AS path FROM "topologic" WHERE uuid_from = ?