From f5bc32769f1df2a712e637b36df8a678b0b00c66 Mon Sep 17 00:00:00 2001 From: WalMaa Date: Tue, 21 Apr 2026 21:40:59 +0300 Subject: [PATCH 01/14] rule jsonb attributes and unique rule name check --- .../migrations/20260225122322_init_schema.sql | 10 +-- skemr-api/db/queries/database_entities.sql | 2 +- skemr-api/db/queries/rules.sql | 10 ++- skemr-api/db/sqlc/database_entities.sql.go | 10 +-- skemr-api/db/sqlc/models.go | 6 +- skemr-api/db/sqlc/querier.go | 1 + skemr-api/db/sqlc/rules.sql.go | 61 ++++++++++++---- .../internal/controller/rule_controller.go | 14 ++-- skemr-api/internal/dbreflect/schema_sync.go | 70 +++++++------------ skemr-api/internal/dto/common.go | 2 + skemr-api/internal/errormsg/errors.go | 1 + skemr-api/internal/mapper/mapper_util.go | 9 +++ skemr-api/internal/mapper/rule_mapper.go | 24 +++++-- skemr-api/internal/service/rule_service.go | 23 ++++++ skemr-common/models/rule.go | 5 ++ sqlc.yml | 6 ++ 16 files changed, 174 insertions(+), 80 deletions(-) diff --git a/skemr-api/db/migrations/20260225122322_init_schema.sql b/skemr-api/db/migrations/20260225122322_init_schema.sql index e5e6362..852f6f8 100644 --- a/skemr-api/db/migrations/20260225122322_init_schema.sql +++ b/skemr-api/db/migrations/20260225122322_init_schema.sql @@ -120,7 +120,7 @@ CREATE TABLE tables CREATE TABLE database_entities ( id uuid PRIMARY KEY DEFAULT gen_random_uuid(), - fingerprint text, -- this is used to track the same entity across syncs even if it is renamed. + fingerprint text NOT NULL, -- this is used to track the same entity across syncs even if it is renamed. project_id uuid NOT NULL REFERENCES projects (id) ON DELETE CASCADE, database_id uuid NOT NULL REFERENCES databases (id) ON DELETE CASCADE, status database_entity_status NOT NULL DEFAULT 'active', @@ -131,7 +131,7 @@ CREATE TABLE database_entities -- generic identity at this node name text NOT NULL, -- e.g. "public", "users", "email", "my_view" - attributes jsonb, -- Store any additional metadata about the entity here + attributes jsonb NOT NULL DEFAULT '{}'::jsonb, -- Store any additional metadata about the entity here created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), @@ -143,9 +143,11 @@ CREATE TABLE database_entities CREATE TABLE rules ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), - name TEXT NOT NULL, -- User defined for rule + name TEXT NOT NULL, -- Defined by user type rule_type NOT NULL, + attributes jsonb NOT NULL DEFAULT '{}'::jsonb, -- Metadata about the rule, removal_date for deprecated types for example database_entity_id uuid NOT NULL REFERENCES database_entities (id) ON DELETE CASCADE, database_id uuid NOT NULL REFERENCES databases (id) ON DELETE CASCADE, - created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + CONSTRAINT unique_rule_name_per_database UNIQUE (name, database_id) ); \ No newline at end of file diff --git a/skemr-api/db/queries/database_entities.sql b/skemr-api/db/queries/database_entities.sql index ccfd177..e2a1f42 100644 --- a/skemr-api/db/queries/database_entities.sql +++ b/skemr-api/db/queries/database_entities.sql @@ -60,7 +60,7 @@ LIMIT 1; -- name: CreateDatabaseEntity :one INSERT INTO database_entities (project_id, database_id, entity_type, parent_id, name, attributes, fingerprint) -VALUES (@project_id, @database_id, @entity_type, @parent_id, @name, @attributes, @fingerprint) +VALUES (@project_id, @database_id, @entity_type, @parent_id, @name, COALESCE(@attributes, '{}'::jsonb), @fingerprint) RETURNING *; -- name: UpdateDatabaseEntityName :one diff --git a/skemr-api/db/queries/rules.sql b/skemr-api/db/queries/rules.sql index bb940f0..fee0231 100644 --- a/skemr-api/db/queries/rules.sql +++ b/skemr-api/db/queries/rules.sql @@ -4,6 +4,12 @@ FROM rules WHERE database_id = @database_id AND id = @rule_id LIMIT 1; +-- name: GetRuleByDatabaseAndName :one +SELECT * +FROM rules +WHERE database_id = @database_id AND name = @name +LIMIT 1; + -- name: GetRuleWithEntity :one SELECT sqlc.embed(r), @@ -15,8 +21,8 @@ LIMIT 1; -- name: CreateRule :one INSERT INTO rules - (name, type, database_entity_id, database_id) -VALUES (@name, @type, @database_entity_id, @database_id) + (name, type, database_entity_id, database_id, attributes) +VALUES (@name, @type, @database_entity_id, @database_id, COALESCE(@attributes, '{}'::jsonb)) RETURNING *; -- name: UpdateRule :exec diff --git a/skemr-api/db/sqlc/database_entities.sql.go b/skemr-api/db/sqlc/database_entities.sql.go index bc6c7b7..3049476 100644 --- a/skemr-api/db/sqlc/database_entities.sql.go +++ b/skemr-api/db/sqlc/database_entities.sql.go @@ -15,7 +15,7 @@ import ( const createDatabaseEntity = `-- name: CreateDatabaseEntity :one INSERT INTO database_entities (project_id, database_id, entity_type, parent_id, name, attributes, fingerprint) -VALUES ($1, $2, $3, $4, $5, $6, $7) +VALUES ($1, $2, $3, $4, $5, COALESCE($6, '{}'::jsonb), $7) RETURNING id, fingerprint, project_id, database_id, status, deleted_at, first_seen_at, entity_type, parent_id, name, attributes, created_at ` @@ -25,8 +25,8 @@ type CreateDatabaseEntityParams struct { EntityType DatabaseEntityType `json:"entity_type"` ParentID *uuid.UUID `json:"parent_id"` Name string `json:"name"` - Attributes []byte `json:"attributes"` - Fingerprint pgtype.Text `json:"fingerprint"` + Attributes interface{} `json:"attributes"` + Fingerprint string `json:"fingerprint"` } func (q *Queries) CreateDatabaseEntity(ctx context.Context, arg CreateDatabaseEntityParams) (DatabaseEntity, error) { @@ -305,8 +305,8 @@ LIMIT 1 ` type GetDatabaseEntityByFingerprintParams struct { - DatabaseID uuid.UUID `json:"database_id"` - Fingerprint pgtype.Text `json:"fingerprint"` + DatabaseID uuid.UUID `json:"database_id"` + Fingerprint string `json:"fingerprint"` } func (q *Queries) GetDatabaseEntityByFingerprint(ctx context.Context, arg GetDatabaseEntityByFingerprintParams) (DatabaseEntity, error) { diff --git a/skemr-api/db/sqlc/models.go b/skemr-api/db/sqlc/models.go index 49784f4..c978834 100644 --- a/skemr-api/db/sqlc/models.go +++ b/skemr-api/db/sqlc/models.go @@ -6,6 +6,7 @@ package sqlc import ( "database/sql/driver" + "encoding/json" "fmt" "github.com/google/uuid" @@ -293,7 +294,7 @@ type Database struct { type DatabaseEntity struct { ID uuid.UUID `json:"id"` - Fingerprint pgtype.Text `json:"fingerprint"` + Fingerprint string `json:"fingerprint"` ProjectID uuid.UUID `json:"project_id"` DatabaseID uuid.UUID `json:"database_id"` Status DatabaseEntityStatus `json:"status"` @@ -302,7 +303,7 @@ type DatabaseEntity struct { EntityType DatabaseEntityType `json:"entity_type"` ParentID *uuid.UUID `json:"parent_id"` Name string `json:"name"` - Attributes []byte `json:"attributes"` + Attributes json.RawMessage `json:"attributes"` CreatedAt pgtype.Timestamptz `json:"created_at"` } @@ -347,6 +348,7 @@ type Rule struct { ID uuid.UUID `json:"id"` Name string `json:"name"` Type RuleType `json:"type"` + Attributes json.RawMessage `json:"attributes"` DatabaseEntityID uuid.UUID `json:"database_entity_id"` DatabaseID uuid.UUID `json:"database_id"` CreatedAt pgtype.Timestamptz `json:"created_at"` diff --git a/skemr-api/db/sqlc/querier.go b/skemr-api/db/sqlc/querier.go index 70f45ef..c8a05dc 100644 --- a/skemr-api/db/sqlc/querier.go +++ b/skemr-api/db/sqlc/querier.go @@ -39,6 +39,7 @@ type Querier interface { GetProjectSecretKeyByID(ctx context.Context, arg GetProjectSecretKeyByIDParams) (GetProjectSecretKeyByIDRow, error) GetProjects(ctx context.Context) ([]Project, error) GetRule(ctx context.Context, arg GetRuleParams) (Rule, error) + GetRuleByDatabaseAndName(ctx context.Context, arg GetRuleByDatabaseAndNameParams) (Rule, error) GetRuleWithEntity(ctx context.Context, arg GetRuleWithEntityParams) (GetRuleWithEntityRow, error) GetRulesWithEntities(ctx context.Context, databaseID uuid.UUID) ([]GetRulesWithEntitiesRow, error) ListDatabasesByProject(ctx context.Context, projectID uuid.UUID) ([]Database, error) diff --git a/skemr-api/db/sqlc/rules.sql.go b/skemr-api/db/sqlc/rules.sql.go index 6b6ece9..7ea51da 100644 --- a/skemr-api/db/sqlc/rules.sql.go +++ b/skemr-api/db/sqlc/rules.sql.go @@ -13,16 +13,17 @@ import ( const createRule = `-- name: CreateRule :one INSERT INTO rules - (name, type, database_entity_id, database_id) -VALUES ($1, $2, $3, $4) -RETURNING id, name, type, database_entity_id, database_id, created_at + (name, type, database_entity_id, database_id, attributes) +VALUES ($1, $2, $3, $4, COALESCE($5, '{}'::jsonb)) +RETURNING id, name, type, attributes, database_entity_id, database_id, created_at ` type CreateRuleParams struct { - Name string `json:"name"` - Type RuleType `json:"type"` - DatabaseEntityID uuid.UUID `json:"database_entity_id"` - DatabaseID uuid.UUID `json:"database_id"` + Name string `json:"name"` + Type RuleType `json:"type"` + DatabaseEntityID uuid.UUID `json:"database_entity_id"` + DatabaseID uuid.UUID `json:"database_id"` + Attributes interface{} `json:"attributes"` } func (q *Queries) CreateRule(ctx context.Context, arg CreateRuleParams) (Rule, error) { @@ -31,12 +32,14 @@ func (q *Queries) CreateRule(ctx context.Context, arg CreateRuleParams) (Rule, e arg.Type, arg.DatabaseEntityID, arg.DatabaseID, + arg.Attributes, ) var i Rule err := row.Scan( &i.ID, &i.Name, &i.Type, + &i.Attributes, &i.DatabaseEntityID, &i.DatabaseID, &i.CreatedAt, @@ -61,7 +64,7 @@ func (q *Queries) DeleteRule(ctx context.Context, arg DeleteRuleParams) error { } const getRule = `-- name: GetRule :one -SELECT id, name, type, database_entity_id, database_id, created_at +SELECT id, name, type, attributes, database_entity_id, database_id, created_at FROM rules WHERE database_id = $1 AND id = $2 LIMIT 1 @@ -79,6 +82,34 @@ func (q *Queries) GetRule(ctx context.Context, arg GetRuleParams) (Rule, error) &i.ID, &i.Name, &i.Type, + &i.Attributes, + &i.DatabaseEntityID, + &i.DatabaseID, + &i.CreatedAt, + ) + return i, err +} + +const getRuleByDatabaseAndName = `-- name: GetRuleByDatabaseAndName :one +SELECT id, name, type, attributes, database_entity_id, database_id, created_at +FROM rules +WHERE database_id = $1 AND name = $2 +LIMIT 1 +` + +type GetRuleByDatabaseAndNameParams struct { + DatabaseID uuid.UUID `json:"database_id"` + Name string `json:"name"` +} + +func (q *Queries) GetRuleByDatabaseAndName(ctx context.Context, arg GetRuleByDatabaseAndNameParams) (Rule, error) { + row := q.db.QueryRow(ctx, getRuleByDatabaseAndName, arg.DatabaseID, arg.Name) + var i Rule + err := row.Scan( + &i.ID, + &i.Name, + &i.Type, + &i.Attributes, &i.DatabaseEntityID, &i.DatabaseID, &i.CreatedAt, @@ -88,7 +119,7 @@ func (q *Queries) GetRule(ctx context.Context, arg GetRuleParams) (Rule, error) const getRuleWithEntity = `-- name: GetRuleWithEntity :one SELECT - r.id, r.name, r.type, r.database_entity_id, r.database_id, r.created_at, + r.id, r.name, r.type, r.attributes, r.database_entity_id, r.database_id, r.created_at, de.id, de.fingerprint, de.project_id, de.database_id, de.status, de.deleted_at, de.first_seen_at, de.entity_type, de.parent_id, de.name, de.attributes, de.created_at FROM rules r JOIN database_entities de ON r.database_entity_id = de.id @@ -113,6 +144,7 @@ func (q *Queries) GetRuleWithEntity(ctx context.Context, arg GetRuleWithEntityPa &i.Rule.ID, &i.Rule.Name, &i.Rule.Type, + &i.Rule.Attributes, &i.Rule.DatabaseEntityID, &i.Rule.DatabaseID, &i.Rule.CreatedAt, @@ -134,7 +166,7 @@ func (q *Queries) GetRuleWithEntity(ctx context.Context, arg GetRuleWithEntityPa const getRulesWithEntities = `-- name: GetRulesWithEntities :many SELECT - r.id, r.name, r.type, r.database_entity_id, r.database_id, r.created_at, + r.id, r.name, r.type, r.attributes, r.database_entity_id, r.database_id, r.created_at, de.id, de.fingerprint, de.project_id, de.database_id, de.status, de.deleted_at, de.first_seen_at, de.entity_type, de.parent_id, de.name, de.attributes, de.created_at FROM rules r JOIN database_entities de ON r.database_entity_id = de.id @@ -159,6 +191,7 @@ func (q *Queries) GetRulesWithEntities(ctx context.Context, databaseID uuid.UUID &i.Rule.ID, &i.Rule.Name, &i.Rule.Type, + &i.Rule.Attributes, &i.Rule.DatabaseEntityID, &i.Rule.DatabaseID, &i.Rule.CreatedAt, @@ -186,7 +219,7 @@ func (q *Queries) GetRulesWithEntities(ctx context.Context, databaseID uuid.UUID } const listRulesByCriteria = `-- name: ListRulesByCriteria :many -SELECT id, name, type, database_entity_id, database_id, created_at +SELECT id, name, type, attributes, database_entity_id, database_id, created_at FROM rules WHERE database_id = $1 AND (database_entity_id = $2 OR $2 IS NULL) @@ -210,6 +243,7 @@ func (q *Queries) ListRulesByCriteria(ctx context.Context, arg ListRulesByCriter &i.ID, &i.Name, &i.Type, + &i.Attributes, &i.DatabaseEntityID, &i.DatabaseID, &i.CreatedAt, @@ -225,7 +259,7 @@ func (q *Queries) ListRulesByCriteria(ctx context.Context, arg ListRulesByCriter } const listRulesByDatabaseId = `-- name: ListRulesByDatabaseId :many -SELECT id, name, type, database_entity_id, database_id, created_at +SELECT id, name, type, attributes, database_entity_id, database_id, created_at FROM rules WHERE database_id = $1 ` @@ -243,6 +277,7 @@ func (q *Queries) ListRulesByDatabaseId(ctx context.Context, databaseID uuid.UUI &i.ID, &i.Name, &i.Type, + &i.Attributes, &i.DatabaseEntityID, &i.DatabaseID, &i.CreatedAt, @@ -262,7 +297,7 @@ UPDATE rules SET name = $2, type = $3 WHERE id = $1 -RETURNING id, name, type, database_entity_id, database_id, created_at +RETURNING id, name, type, attributes, database_entity_id, database_id, created_at ` type UpdateRuleParams struct { diff --git a/skemr-api/internal/controller/rule_controller.go b/skemr-api/internal/controller/rule_controller.go index 78510e5..8dc98bb 100644 --- a/skemr-api/internal/controller/rule_controller.go +++ b/skemr-api/internal/controller/rule_controller.go @@ -1,13 +1,14 @@ package controller import ( - "encoding/json" + "fmt" "net/http" "github.com/go-chi/chi/v5" "github.com/go-chi/render" "github.com/google/uuid" "github.com/walmaa/skemr-api/internal/dto" + "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/service" ) @@ -84,17 +85,20 @@ func (h *RuleController) createRule(w http.ResponseWriter, r *http.Request) { } projectID, ok := r.Context().Value("projectId").(uuid.UUID) if !ok { - http.Error(w, "projectId not found in context", http.StatusBadRequest) + err = fmt.Errorf("projectId not found in context") + errormsg.WriteErrorResponse(w, r, err) return } var body dto.RuleCreationDto - if err := json.NewDecoder(r.Body).Decode(&body); err != nil { - http.Error(w, err.Error(), http.StatusBadRequest) + + if err := render.Decode(r, &body); err != nil { + errormsg.WriteErrorResponse(w, r, err) return } rule, err := h.Service.CreateRule(r.Context(), projectID, databaseId, body) if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) + + errormsg.WriteErrorResponse(w, r, err) return } diff --git a/skemr-api/internal/dbreflect/schema_sync.go b/skemr-api/internal/dbreflect/schema_sync.go index 12de518..c45d5f4 100644 --- a/skemr-api/internal/dbreflect/schema_sync.go +++ b/skemr-api/internal/dbreflect/schema_sync.go @@ -188,11 +188,8 @@ func (s *SchemaSyncService) updateSchema(c context.Context, schemaRef SchemaRef, fingerprint := GenerateSchemaFingerprint(schemaRef, database.ID) schema, err = s.db.GetDatabaseEntityByFingerprint(c, sqlc.GetDatabaseEntityByFingerprintParams{ - DatabaseID: database.ID, - Fingerprint: pgtype.Text{ - String: fingerprint, - Valid: true, - }, + DatabaseID: database.ID, + Fingerprint: fingerprint, }) if err != nil && !errors.Is(err, pgx.ErrNoRows) { @@ -219,14 +216,11 @@ func (s *SchemaSyncService) updateSchema(c context.Context, schemaRef SchemaRef, // If that schema does not exist yet, save it args := sqlc.CreateDatabaseEntityParams{ - ProjectID: database.ProjectID, - EntityType: sqlc.DatabaseEntityTypeSchema, - DatabaseID: database.ID, - Name: schemaRef.Name, - Fingerprint: pgtype.Text{ - String: fingerprint, - Valid: true, - }, + ProjectID: database.ProjectID, + EntityType: sqlc.DatabaseEntityTypeSchema, + DatabaseID: database.ID, + Name: schemaRef.Name, + Fingerprint: fingerprint, } schema, err := s.db.CreateDatabaseEntity(c, args) if err != nil { @@ -277,11 +271,8 @@ func (s *SchemaSyncService) UpdateTable(c context.Context, tableRef TableRef, da fingerprint := GenerateTableFingerprint(tableRef) table, err = s.db.GetDatabaseEntityByFingerprint(c, sqlc.GetDatabaseEntityByFingerprintParams{ - DatabaseID: database.ID, - Fingerprint: pgtype.Text{ - String: fingerprint, - Valid: true, - }, + DatabaseID: database.ID, + Fingerprint: fingerprint, }) if err != nil && !errors.Is(err, pgx.ErrNoRows) { @@ -309,19 +300,16 @@ func (s *SchemaSyncService) UpdateTable(c context.Context, tableRef TableRef, da } args := sqlc.CreateDatabaseEntityParams{ - ProjectID: database.ProjectID, - EntityType: sqlc.DatabaseEntityTypeTable, - ParentID: &schemaId, - DatabaseID: database.ID, - Name: tableRef.Name, - Fingerprint: pgtype.Text{ - String: GenerateTableFingerprint(tableRef), - Valid: true, - }, + ProjectID: database.ProjectID, + EntityType: sqlc.DatabaseEntityTypeTable, + ParentID: &schemaId, + DatabaseID: database.ID, + Name: tableRef.Name, + Fingerprint: fingerprint, } table, err := s.db.CreateDatabaseEntity(c, args) if err != nil { - slog.Error("error creating schema", "error", err) + slog.Error("error creating table", "error", err) return sqlc.DatabaseEntity{}, err } slog.Info("Table created", "schema", table.Name) @@ -366,11 +354,8 @@ func (s *SchemaSyncService) SyncColumn(c context.Context, columnRef ColumnRef, d fingerprint := GenerateColumnFingerprint(columnRef, tableId) column, err = s.db.GetDatabaseEntityByFingerprint(c, sqlc.GetDatabaseEntityByFingerprintParams{ - DatabaseID: database.ID, - Fingerprint: pgtype.Text{ - String: fingerprint, - Valid: true, - }, + DatabaseID: database.ID, + Fingerprint: fingerprint, }) if err != nil && !errors.Is(err, pgx.ErrNoRows) { @@ -397,20 +382,17 @@ func (s *SchemaSyncService) SyncColumn(c context.Context, columnRef ColumnRef, d // If that column does not exist yet, save it args := sqlc.CreateDatabaseEntityParams{ - ProjectID: database.ProjectID, - EntityType: sqlc.DatabaseEntityTypeColumn, - ParentID: &tableId, - DatabaseID: database.ID, - Name: columnRef.Name, - Attributes: attributesJson, - Fingerprint: pgtype.Text{ - String: GenerateColumnFingerprint(columnRef, tableId), - Valid: true, - }, + ProjectID: database.ProjectID, + EntityType: sqlc.DatabaseEntityTypeColumn, + ParentID: &tableId, + DatabaseID: database.ID, + Name: columnRef.Name, + Attributes: attributesJson, + Fingerprint: fingerprint, } column, err := s.db.CreateDatabaseEntity(c, args) if err != nil { - slog.Error("error creating schema", "error", err) + slog.Error("error creating column", "error", err) return sqlc.DatabaseEntity{}, err } slog.Info("Column created", "name", column.Name) diff --git a/skemr-api/internal/dto/common.go b/skemr-api/internal/dto/common.go index e3af023..cf90b97 100644 --- a/skemr-api/internal/dto/common.go +++ b/skemr-api/internal/dto/common.go @@ -2,6 +2,7 @@ package dto import ( "github.com/google/uuid" + "github.com/walmaa/skemr-common/models" ) type ProjectCreationDto struct { @@ -39,6 +40,7 @@ const ( type RuleCreationDto struct { Name string RuleType RuleType + Attributes models.RuleAttributes `json:"attributes" validate:"omitempty,json"` DataBaseEntityId uuid.UUID } diff --git a/skemr-api/internal/errormsg/errors.go b/skemr-api/internal/errormsg/errors.go index cd227e1..caac1ab 100644 --- a/skemr-api/internal/errormsg/errors.go +++ b/skemr-api/internal/errormsg/errors.go @@ -34,4 +34,5 @@ var ( ErrProjectNotFound = "project not found" ErrInvalidIdFormat = "invalid id format" ErrExpiryTimeInPast = "expiry time is in the past" + ErrRuleWithSameName = "rule with the same name already exists" ) diff --git a/skemr-api/internal/mapper/mapper_util.go b/skemr-api/internal/mapper/mapper_util.go index cca9c97..507c1e5 100644 --- a/skemr-api/internal/mapper/mapper_util.go +++ b/skemr-api/internal/mapper/mapper_util.go @@ -10,6 +10,15 @@ import ( "github.com/walmaa/skemr-api/internal/dto" ) +func ToBytes(v interface{}) []byte { + b, err := json.Marshal(v) + if err != nil { + slog.Error("Unable to marshal JSON", err) + return nil + } + return b +} + func Text(v *string) pgtype.Text { if v == nil { return pgtype.Text{ diff --git a/skemr-api/internal/mapper/rule_mapper.go b/skemr-api/internal/mapper/rule_mapper.go index ba8ec5f..6b83064 100644 --- a/skemr-api/internal/mapper/rule_mapper.go +++ b/skemr-api/internal/mapper/rule_mapper.go @@ -1,6 +1,9 @@ package mapper import ( + "encoding/json" + "log/slog" + "github.com/google/uuid" "github.com/walmaa/skemr-api/db/sqlc" "github.com/walmaa/skemr-api/internal/dto" @@ -9,17 +12,29 @@ import ( func ToDomainRule(e sqlc.Rule) models.Rule { return models.Rule{ - ID: e.ID, - Name: e.Name, - RuleType: models.RuleType(e.Type), - CreatedAt: Time(&e.CreatedAt), + ID: e.ID, + Name: e.Name, + RuleType: models.RuleType(e.Type), + Attributes: ToRuleAttributes(e.Attributes), + CreatedAt: Time(&e.CreatedAt), + } +} + +func ToRuleAttributes(attributes []byte) models.RuleAttributes { + var ruleAttributes models.RuleAttributes + err := json.Unmarshal(attributes, &ruleAttributes) + if err != nil { + slog.Error("Unable to unmarshal rule attributes", "error", err) + panic(err) } + return ruleAttributes } func ToDomainRuleWithEntity(e sqlc.GetRuleWithEntityRow) models.Rule { return models.Rule{ ID: e.Rule.ID, Name: e.Rule.Name, + Attributes: ToRuleAttributes(e.Rule.Attributes), RuleType: models.RuleType(e.Rule.Type), DataBaseEntity: ToDomainDatabaseEntity(e.DatabaseEntity), CreatedAt: Time(&e.Rule.CreatedAt), @@ -47,6 +62,7 @@ func ToSqlcCreateRule(databaseId uuid.UUID, dto dto.RuleCreationDto) sqlc.Create Name: dto.Name, Type: sqlc.RuleType(dto.RuleType), DatabaseID: databaseId, + Attributes: ToBytes(dto.Attributes), DatabaseEntityID: dto.DataBaseEntityId, } } diff --git a/skemr-api/internal/service/rule_service.go b/skemr-api/internal/service/rule_service.go index cd15f73..676dda0 100644 --- a/skemr-api/internal/service/rule_service.go +++ b/skemr-api/internal/service/rule_service.go @@ -2,11 +2,15 @@ package service import ( "context" + "errors" "log/slog" + "net/http" "github.com/google/uuid" + "github.com/jackc/pgx/v5" "github.com/walmaa/skemr-api/db/sqlc" "github.com/walmaa/skemr-api/internal/dto" + "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/mapper" "github.com/walmaa/skemr-common/models" ) @@ -64,6 +68,25 @@ func (r *RuleService) CreateRule(c context.Context, projectID uuid.UUID, databas return models.Rule{}, err } + // Check if a rule with the same name already exists + exists, err := r.db.GetRuleByDatabaseAndName(c, sqlc.GetRuleByDatabaseAndNameParams{ + DatabaseID: databaseId, + Name: dto.Name, + }) + + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + slog.Error("Error checking for existing rule", "name", dto.Name, "err", err) + return models.Rule{}, err + } + + if exists.ID != uuid.Nil { + slog.Warn("Rule with the same name already exists", "name", dto.Name) + return models.Rule{}, &models.ErrorResponse{ + Message: errormsg.ErrRuleWithSameName, + Status: http.StatusConflict, + } + } + rule, err := r.db.CreateRule(c, mapper.ToSqlcCreateRule(databaseId, dto)) if err != nil { slog.Error("Unable to create a Rule", err) diff --git a/skemr-common/models/rule.go b/skemr-common/models/rule.go index 3063c53..5f9ea3b 100644 --- a/skemr-common/models/rule.go +++ b/skemr-common/models/rule.go @@ -6,9 +6,14 @@ import ( "github.com/google/uuid" ) +type RuleAttributes struct { + RemovalDate *string `json:"removalDate" validate:"omitempty,date_format=2006-01-02T15:04:05Z07:00"` +} + type Rule struct { ID uuid.UUID `json:"id"` Name string `json:"name"` + Attributes RuleAttributes `json:"attributes"` RuleType RuleType `json:"ruleType"` DataBaseEntity DatabaseEntity `json:"databaseEntity"` CreatedAt time.Time `json:"createdAt"` diff --git a/sqlc.yml b/sqlc.yml index 195dcae..492c35c 100644 --- a/sqlc.yml +++ b/sqlc.yml @@ -22,3 +22,9 @@ sql: import: "github.com/google/uuid" type: "UUID" pointer: true + - db_type: "jsonb" + go_type: + import: "encoding/json" + type: "RawMessage" + pointer: false + From 30c364ccd2a23dfef007dc0c10596895025f8640 Mon Sep 17 00:00:00 2001 From: WalMaa Date: Sun, 3 May 2026 10:55:10 +0300 Subject: [PATCH 02/14] added rule attribs --- skemr-api/internal/controller/rule_controller.go | 3 ++- skemr-common/models/rule.go | 2 +- skemr-frontend/src/types/types.ts | 5 +++++ 3 files changed, 8 insertions(+), 2 deletions(-) diff --git a/skemr-api/internal/controller/rule_controller.go b/skemr-api/internal/controller/rule_controller.go index 8dc98bb..7c90062 100644 --- a/skemr-api/internal/controller/rule_controller.go +++ b/skemr-api/internal/controller/rule_controller.go @@ -95,9 +95,10 @@ func (h *RuleController) createRule(w http.ResponseWriter, r *http.Request) { errormsg.WriteErrorResponse(w, r, err) return } + rule, err := h.Service.CreateRule(r.Context(), projectID, databaseId, body) - if err != nil { + if err != nil { errormsg.WriteErrorResponse(w, r, err) return } diff --git a/skemr-common/models/rule.go b/skemr-common/models/rule.go index 5f9ea3b..19353df 100644 --- a/skemr-common/models/rule.go +++ b/skemr-common/models/rule.go @@ -7,7 +7,7 @@ import ( ) type RuleAttributes struct { - RemovalDate *string `json:"removalDate" validate:"omitempty,date_format=2006-01-02T15:04:05Z07:00"` + DeprecatedRemovalDate *string `json:"deprecatedRemovalDate" validate:"omitempty,date_format=2006-01-02T15:04:05Z07:00"` } type Rule struct { diff --git a/skemr-frontend/src/types/types.ts b/skemr-frontend/src/types/types.ts index d0e4ba6..0b76298 100644 --- a/skemr-frontend/src/types/types.ts +++ b/skemr-frontend/src/types/types.ts @@ -56,6 +56,11 @@ export interface Rule { ruleType: DatabaseRuleType; createdAt: string; databaseEntity: DatabaseEntity; + attributes: RuleAttributes; +} + +export interface RuleAttributes { + deprecatedRemovalDate: string | null; } export type RuleCreationDto = { From 28f7ef2c6c4062c873e51781fb971686c571d6c2 Mon Sep 17 00:00:00 2001 From: WalMaa Date: Tue, 5 May 2026 13:18:30 +0300 Subject: [PATCH 03/14] Merge branch 'main' into dev --- .github/workflows/build-and-test.yml | 12 +- skemr-api/config/config.go | 62 +++++++--- .../project_access_tokens_controller.go | 6 +- skemr-api/internal/mapper/mapper_util.go | 2 +- .../internal/service/access_token_service.go | 24 ++-- .../internal/service/database_service.go | 10 +- skemr-api/internal/service/project_service.go | 2 +- skemr-api/internal/service/rule_service.go | 18 +-- skemr-api/test/mocks/querier_mock.go | 66 ++++++++++ .../(project)/projects/$projectId/route.tsx | 2 +- .../projects/$projectId/settings/index.tsx | 117 ++++++++++++++---- 11 files changed, 246 insertions(+), 75 deletions(-) diff --git a/.github/workflows/build-and-test.yml b/.github/workflows/build-and-test.yml index d0e6681..ca76e1f 100644 --- a/.github/workflows/build-and-test.yml +++ b/.github/workflows/build-and-test.yml @@ -17,9 +17,10 @@ jobs: - name: Set up Go uses: actions/setup-go@v6 with: - go-version: '1.25' + go-version-file: go.work + cache: 'true' - - uses: sqlc-dev/setup-sqlc@v4 + - uses: sqlc-dev/setup-sqlc@v5 with: sqlc-version: '1.25.0' @@ -32,8 +33,11 @@ jobs: - name: Create mocks run: mockery - - name: Build - run: go build ./... + - name: Build control plane + run: go build ./skemr-api/... + + - name: Build CLI + run: go build ./skemr-cli/... - name: Test run: go test -race -shuffle=on ./skemr-cli/... ./skemr-api/... diff --git a/skemr-api/config/config.go b/skemr-api/config/config.go index 3097147..7c4e9c0 100644 --- a/skemr-api/config/config.go +++ b/skemr-api/config/config.go @@ -8,6 +8,12 @@ import ( "github.com/spf13/viper" ) +const ( + defaultAppPort = 8080 + defaultDatabasePort = 5432 + defaultRedisPort = 6379 +) + type Config struct { App struct { Env string @@ -39,23 +45,23 @@ func LoadConfig() (*Config, error) { if env == "dev" { if err := godotenv.Load(".env"); err != nil { // Log the error but continue, as environment variables might still be set - slog.Error("Warning: Could not load .env file: %v", err) + slog.Error("Warning: Could not load .env file", "err", err) } } // Set defaults first viper.SetDefault("app.env", env) - viper.SetDefault("app.port", 8080) + viper.SetDefault("app.port", defaultAppPort) // Database defaults viper.SetDefault("database.host", "localhost") - viper.SetDefault("database.port", 5432) + viper.SetDefault("database.port", defaultDatabasePort) viper.SetDefault("database.user", "postgres") viper.SetDefault("database.password", "pass") viper.SetDefault("database.name", "postgres") viper.SetDefault("database.sslmode", "disable") // Redis defaults viper.SetDefault("redis.host", "localhost") - viper.SetDefault("redis.port", 6379) + viper.SetDefault("redis.port", defaultRedisPort) viper.SetDefault("redis.password", "") viper.SetDefault("redis.db", 0) @@ -63,20 +69,44 @@ func LoadConfig() (*Config, error) { viper.AutomaticEnv() // Bind environment variables to viper keys - viper.BindEnv("app.env", "APP_ENV") - viper.BindEnv("app.port", "APP_PORT") + if err := viper.BindEnv("app.env", "APP_ENV"); err != nil { + return nil, err + } + if err := viper.BindEnv("app.port", "APP_PORT"); err != nil { + return nil, err + } // Database env vars - viper.BindEnv("database.host", "DB_HOST") - viper.BindEnv("database.port", "DB_PORT") - viper.BindEnv("database.user", "DB_USER") - viper.BindEnv("database.password", "DB_PASSWORD") - viper.BindEnv("database.name", "DB_NAME") - viper.BindEnv("database.sslmode", "DB_SSLMODE") + if err := viper.BindEnv("database.host", "DB_HOST"); err != nil { + return nil, err + } + if err := viper.BindEnv("database.port", "DB_PORT"); err != nil { + return nil, err + } + if err := viper.BindEnv("database.user", "DB_USER"); err != nil { + return nil, err + } + if err := viper.BindEnv("database.password", "DB_PASSWORD"); err != nil { + return nil, err + } + if err := viper.BindEnv("database.name", "DB_NAME"); err != nil { + return nil, err + } + if err := viper.BindEnv("database.sslmode", "DB_SSLMODE"); err != nil { + return nil, err + } // Redis env vars - viper.BindEnv("redis.host", "REDIS_HOST") - viper.BindEnv("redis.port", "REDIS_PORT") - viper.BindEnv("redis.password", "REDIS_PASSWORD") - viper.BindEnv("redis.db", "REDIS_DB") + if err := viper.BindEnv("redis.host", "REDIS_HOST"); err != nil { + return nil, err + } + if err := viper.BindEnv("redis.port", "REDIS_PORT"); err != nil { + return nil, err + } + if err := viper.BindEnv("redis.password", "REDIS_PASSWORD"); err != nil { + return nil, err + } + if err := viper.BindEnv("redis.db", "REDIS_DB"); err != nil { + return nil, err + } var cfg Config diff --git a/skemr-api/internal/controller/project_access_tokens_controller.go b/skemr-api/internal/controller/project_access_tokens_controller.go index c4cb66d..f86af12 100644 --- a/skemr-api/internal/controller/project_access_tokens_controller.go +++ b/skemr-api/internal/controller/project_access_tokens_controller.go @@ -77,7 +77,7 @@ func (h *ProjectSecretsController) getSecrets(w http.ResponseWriter, r *http.Req tokens, err := h.Service.GetTokens(c, projectId) if err != nil { - slog.Error("Error getting tokens", err) + slog.Error("Error getting tokens", "err", err) return } @@ -85,7 +85,7 @@ func (h *ProjectSecretsController) getSecrets(w http.ResponseWriter, r *http.Req render.JSON(w, r, tokens) } -func (h *ProjectSecretsController) updateSecret(w http.ResponseWriter, r *http.Request) { +func (h *ProjectSecretsController) updateSecret(_ http.ResponseWriter, _ *http.Request) { // TODO: Implement update logic } @@ -101,7 +101,7 @@ func (h *ProjectSecretsController) deleteSecret(w http.ResponseWriter, r *http.R err = h.Service.DeleteToken(c, projectId, secretId) if err != nil { - slog.Error("Error deleting token", err) + slog.Error("Error deleting token", "err", err) http.Error(w, "Error deleting token", http.StatusInternalServerError) return } diff --git a/skemr-api/internal/mapper/mapper_util.go b/skemr-api/internal/mapper/mapper_util.go index 507c1e5..07e4dab 100644 --- a/skemr-api/internal/mapper/mapper_util.go +++ b/skemr-api/internal/mapper/mapper_util.go @@ -40,7 +40,7 @@ func ToMap(b []byte) map[string]interface{} { var m map[string]interface{} err := json.Unmarshal(b, &m) if err != nil { - slog.Error("Unable to unmarshal JSON", err) + slog.Error("Unable to unmarshal JSON", "err", err) return nil } return m diff --git a/skemr-api/internal/service/access_token_service.go b/skemr-api/internal/service/access_token_service.go index cb97a64..f2953ec 100644 --- a/skemr-api/internal/service/access_token_service.go +++ b/skemr-api/internal/service/access_token_service.go @@ -3,7 +3,6 @@ package service import ( "context" "errors" - "fmt" "log/slog" "net/http" "strings" @@ -37,14 +36,14 @@ func (s *AccessTokenService) CreateToken(c context.Context, projectId uuid.UUID, project, err := CheckProjectExists(c, s.db, projectId) if err != nil { - slog.Error("Unable to get project", err) + slog.Error("Unable to get project", "err", err) return "", err } tokenToShow, prefix, secret, err := tokens.GenerateToken(prefixLength, secretLength) if err != nil { - slog.Error("Unable to generate token", err) + slog.Error("Unable to generate token", "err", err) return "", err } @@ -52,7 +51,7 @@ func (s *AccessTokenService) CreateToken(c context.Context, projectId uuid.UUID, verifier, err := tokens.HashSecret(secret, tokens.DefaultParams) if err != nil { - slog.Error("Unable to hash token", err) + slog.Error("Unable to hash token", "err", err) return "", err } expires := pgtype.Timestamptz{ @@ -62,12 +61,11 @@ func (s *AccessTokenService) CreateToken(c context.Context, projectId uuid.UUID, if dto.ExpiresAt != "" { expiry, err := time.Parse(time.RFC3339, dto.ExpiresAt) if err != nil { - slog.Error("Unable to parse expiry time", err) + slog.Error("Unable to parse expiry time", "err", err) return "", err } if expiry.Before(time.Now()) { slog.Error("Expiry time is in the past") - err = fmt.Errorf("expiry time is in the past") return "", &models.ErrorResponse{ Message: errormsg.ErrExpiryTimeInPast, Status: http.StatusBadRequest, @@ -89,7 +87,7 @@ func (s *AccessTokenService) CreateToken(c context.Context, projectId uuid.UUID, ExpiresAt: expires, }) if err != nil { - slog.Error("Error saving a project access token", err) + slog.Error("Error saving a project access token", "err", err) return "", err } @@ -101,13 +99,13 @@ func (s *AccessTokenService) GetTokens(c context.Context, projectId uuid.UUID) ( project, err := CheckProjectExists(c, s.db, projectId) if err != nil { - slog.Error("Unable to get project", err) + slog.Error("Unable to get project", "err", err) return nil, err } accessTokens, err := s.db.GetProjectAccessTokens(c, project.ID) if err != nil { - slog.Error("Unable to get tokens", err) + slog.Error("Unable to get tokens", "err", err) return nil, err } @@ -121,7 +119,7 @@ func (s *AccessTokenService) DeleteToken(c context.Context, projectId uuid.UUID, project, err := CheckProjectExists(c, s.db, projectId) if err != nil { - slog.Error("Unable to get project") + slog.Error("Unable to get project", "err", err) return err } @@ -130,7 +128,7 @@ func (s *AccessTokenService) DeleteToken(c context.Context, projectId uuid.UUID, SecretID: secretId, }) if err != nil { - slog.Error("Unable to delete project access token", err) + slog.Error("Unable to delete project access token", "err", err) return err } @@ -172,14 +170,14 @@ func (s *AccessTokenService) ValidateToken(c context.Context, projectId uuid.UUI } if err != nil { - slog.Error("Unable to get token hash from database", err) + slog.Error("Unable to get token hash from database", "err", err) return false, err } ok, err := tokens.VerifySecret(secretPart, hash) if err != nil { - slog.Error("Error verifying token", err) + slog.Error("Error verifying token", "err", err) return false, err } diff --git a/skemr-api/internal/service/database_service.go b/skemr-api/internal/service/database_service.go index 9a27058..04c51ae 100644 --- a/skemr-api/internal/service/database_service.go +++ b/skemr-api/internal/service/database_service.go @@ -76,7 +76,7 @@ func (r *DatabaseService) CreateDatabase(c context.Context, projectId uuid.UUID, database, err := r.db.CreateDatabase(c, mapper.ToCreateDatabaseParams(projectId, dto)) if err != nil { - slog.Error("Error creating database", err) + slog.Error("Error creating database", "err", err) return models.Database{}, err } @@ -94,7 +94,7 @@ func (r *DatabaseService) createDatabaseSyncTask(databaseId uuid.UUID) { } _, err = r.taskClient.Enqueue(task) if err != nil { - slog.Error("Error in task", err) + slog.Error("Error in task", "err", err) } } @@ -105,14 +105,14 @@ func (r *DatabaseService) EnqueueManualDatabaseSync(c context.Context, projectId project, err := CheckProjectExists(c, r.db, projectId) if err != nil { - slog.Error("Error fetching project") + slog.Error("Error fetching project", "err", err) return err } database, err := CheckDatabaseExists(c, r.db, project.ID, databaseId) if err != nil { - slog.Error("Error getting database") + slog.Error("Error getting database", "err", err) return err } @@ -174,7 +174,7 @@ func (r *DatabaseService) UpdateDatabase(c context.Context, projectId uuid.UUID, database, err := r.db.UpdateDatabase(c, mapper.ToUpdateDatabaseParams(databaseId, dto)) if err != nil { - slog.Error("Error updating database", err) + slog.Error("Error updating database", "err", err) return models.Database{}, err } diff --git a/skemr-api/internal/service/project_service.go b/skemr-api/internal/service/project_service.go index 645a710..845ac0a 100644 --- a/skemr-api/internal/service/project_service.go +++ b/skemr-api/internal/service/project_service.go @@ -67,7 +67,7 @@ func (r *ProjectService) GetProject(c context.Context, projectId uuid.UUID) (mod project, err := r.db.GetProject(c, projectId) if err != nil { - slog.Error("Error getting project", err) + slog.Error("Error getting project", "err", err) return models.Project{}, err } diff --git a/skemr-api/internal/service/rule_service.go b/skemr-api/internal/service/rule_service.go index 676dda0..342b33c 100644 --- a/skemr-api/internal/service/rule_service.go +++ b/skemr-api/internal/service/rule_service.go @@ -35,7 +35,7 @@ func (r *RuleService) GetRule(c context.Context, projectID uuid.UUID, databaseID database, err := CheckDatabaseExists(c, r.db, project.ID, databaseID) if err != nil { - slog.Error("Error fetching database", err) + slog.Error("Error fetching database", "err", err) return models.Rule{}, err } @@ -45,7 +45,7 @@ func (r *RuleService) GetRule(c context.Context, projectID uuid.UUID, databaseID }) if err != nil { - slog.Error("Unable to fetch rule", "error", err) + slog.Error("Unable to fetch rule", "err", err) return models.Rule{}, err } @@ -64,7 +64,7 @@ func (r *RuleService) CreateRule(c context.Context, projectID uuid.UUID, databas _, err = CheckDatabaseExists(c, r.db, project.ID, databaseId) if err != nil { - slog.Error("Error fetching database", err) + slog.Error("Error fetching database", "err", err) return models.Rule{}, err } @@ -89,7 +89,7 @@ func (r *RuleService) CreateRule(c context.Context, projectID uuid.UUID, databas rule, err := r.db.CreateRule(c, mapper.ToSqlcCreateRule(databaseId, dto)) if err != nil { - slog.Error("Unable to create a Rule", err) + slog.Error("Unable to create a Rule", "err", err) return models.Rule{}, err } @@ -108,11 +108,15 @@ func (r *RuleService) ListRulesByDatabase(c context.Context, projectID uuid.UUID database, err := CheckDatabaseExists(c, r.db, project.ID, databaseID) if err != nil { - slog.Error("Error fetching database", err) + slog.Error("Error fetching database", "err", err) return []models.Rule{}, err } rules, err := r.db.GetRulesWithEntities(c, database.ID) + if err != nil { + slog.Error("Unable to get rules", "err", err) + return []models.Rule{}, err + } return mapper.ToDomainRulesWithEntity(rules), nil } @@ -123,14 +127,14 @@ func (r *RuleService) DeleteRule(c context.Context, projectID uuid.UUID, databas project, err := CheckProjectExists(c, r.db, projectID) if err != nil { - slog.Error("Error fetching project", err) + slog.Error("Error fetching project", "err", err) return err } database, err := CheckDatabaseExists(c, r.db, project.ID, databaseID) if err != nil { - slog.Error("Error fetching database", err) + slog.Error("Error fetching database", "err", err) return err } diff --git a/skemr-api/test/mocks/querier_mock.go b/skemr-api/test/mocks/querier_mock.go index 36a58e3..121d233 100644 --- a/skemr-api/test/mocks/querier_mock.go +++ b/skemr-api/test/mocks/querier_mock.go @@ -1397,6 +1397,72 @@ func (_c *MockQuerier_GetDatabaseEntityByProjectIdAndId_Call) RunAndReturn(run f return _c } +// GetHashByPrefixAndProjectID provides a mock function for the type MockQuerier +func (_mock *MockQuerier) GetHashByPrefixAndProjectID(ctx context.Context, arg sqlc.GetHashByPrefixAndProjectIDParams) (string, error) { + ret := _mock.Called(ctx, arg) + + if len(ret) == 0 { + panic("no return value specified for GetHashByPrefixAndProjectID") + } + + var r0 string + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, sqlc.GetHashByPrefixAndProjectIDParams) (string, error)); ok { + return returnFunc(ctx, arg) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, sqlc.GetHashByPrefixAndProjectIDParams) string); ok { + r0 = returnFunc(ctx, arg) + } else { + r0 = ret.Get(0).(string) + } + if returnFunc, ok := ret.Get(1).(func(context.Context, sqlc.GetHashByPrefixAndProjectIDParams) error); ok { + r1 = returnFunc(ctx, arg) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockQuerier_GetHashByPrefixAndProjectID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'GetHashByPrefixAndProjectID' +type MockQuerier_GetHashByPrefixAndProjectID_Call struct { + *mock.Call +} + +// GetHashByPrefixAndProjectID is a helper method to define mock.On call +// - ctx context.Context +// - arg sqlc.GetHashByPrefixAndProjectIDParams +func (_e *MockQuerier_Expecter) GetHashByPrefixAndProjectID(ctx interface{}, arg interface{}) *MockQuerier_GetHashByPrefixAndProjectID_Call { + return &MockQuerier_GetHashByPrefixAndProjectID_Call{Call: _e.mock.On("GetHashByPrefixAndProjectID", ctx, arg)} +} + +func (_c *MockQuerier_GetHashByPrefixAndProjectID_Call) Run(run func(ctx context.Context, arg sqlc.GetHashByPrefixAndProjectIDParams)) *MockQuerier_GetHashByPrefixAndProjectID_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 sqlc.GetHashByPrefixAndProjectIDParams + if args[1] != nil { + arg1 = args[1].(sqlc.GetHashByPrefixAndProjectIDParams) + } + run( + arg0, + arg1, + ) + }) + return _c +} + +func (_c *MockQuerier_GetHashByPrefixAndProjectID_Call) Return(s string, err error) *MockQuerier_GetHashByPrefixAndProjectID_Call { + _c.Call.Return(s, err) + return _c +} + +func (_c *MockQuerier_GetHashByPrefixAndProjectID_Call) RunAndReturn(run func(ctx context.Context, arg sqlc.GetHashByPrefixAndProjectIDParams) (string, error)) *MockQuerier_GetHashByPrefixAndProjectID_Call { + _c.Call.Return(run) + return _c +} + // GetProject provides a mock function for the type MockQuerier func (_mock *MockQuerier) GetProject(ctx context.Context, id uuid.UUID) (sqlc.Project, error) { ret := _mock.Called(ctx, id) diff --git a/skemr-frontend/src/routes/(project)/projects/$projectId/route.tsx b/skemr-frontend/src/routes/(project)/projects/$projectId/route.tsx index ac6f51e..1e146f4 100644 --- a/skemr-frontend/src/routes/(project)/projects/$projectId/route.tsx +++ b/skemr-frontend/src/routes/(project)/projects/$projectId/route.tsx @@ -26,7 +26,7 @@ function RouteComponent() { -
+
diff --git a/skemr-frontend/src/routes/(project)/projects/$projectId/settings/index.tsx b/skemr-frontend/src/routes/(project)/projects/$projectId/settings/index.tsx index 5d86f46..3c9dbe1 100644 --- a/skemr-frontend/src/routes/(project)/projects/$projectId/settings/index.tsx +++ b/skemr-frontend/src/routes/(project)/projects/$projectId/settings/index.tsx @@ -11,7 +11,6 @@ import { import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; import { Label } from "@/components/ui/label"; -import { Switch } from "@/components/ui/switch"; import { AlertDialog, AlertDialogAction, @@ -25,7 +24,8 @@ import { } from "@/components/ui/alert-dialog"; import { Alert } from "@/components/ui/alert"; import { useTheme } from "@/components/theme-provider"; -import { MoonIcon, SunIcon, WarningIcon } from "@phosphor-icons/react"; +import { cn } from "@/lib/utils"; +import { CheckIcon, WarningIcon } from "@phosphor-icons/react"; import { toast } from "sonner"; export const Route = createFileRoute( @@ -49,7 +49,6 @@ function RouteComponent() { () => projectName.length > 0 && confirmationText === projectName, [confirmationText, projectName], ); - const isDarkMode = theme === "dark"; const handleDeleteProject = async () => { if (!projectId || !canDelete) { @@ -66,7 +65,7 @@ function RouteComponent() { }; return ( -
+

Settings

@@ -80,26 +79,25 @@ function RouteComponent() { Customize how the interface looks. -

-
- -

- Enable dark mode across the app. This preference is saved on - this device. -

-
-
- - - setTheme(checked ? "dark" : "light") - } - aria-label="Toggle dark mode" - /> - -
+
+ + +
@@ -196,3 +194,74 @@ function RouteComponent() {
); } + +type ThemeValue = "light" | "dark" | "system"; + +type ThemeOptionProps = { + label: string; + value: ThemeValue; + selected: boolean; + onSelect: (theme: ThemeValue) => void; +}; + +function ThemeOption({ label, value, selected, onSelect }: ThemeOptionProps) { + return ( + + ); +} + +function ThemePreview({ value }: { value: ThemeValue }) { + if (value === "system") { + return ( +
+ + +
+ ); + } + + return ; +} + +function PreviewPane({ mode }: { mode: "light" | "dark" }) { + const dark = mode === "dark"; + + return ( +
+
+ Aa +
+ ); +} From 005155501dde76dc68fb692b04cf335c88136c48 Mon Sep 17 00:00:00 2001 From: WalMaa Date: Tue, 5 May 2026 15:29:38 +0300 Subject: [PATCH 04/14] modularized parser tests and created parsing for namespaces --- skemr-cli/parser/pg_parser.go | 40 ++++++- skemr-cli/parser/pg_parser_column_test.go | 52 +++++++++ skemr-cli/parser/pg_parser_database_test.go | 40 +++++++ skemr-cli/parser/pg_parser_namespace_test.go | 40 +++++++ skemr-cli/parser/pg_parser_table_test.go | 40 +++++++ skemr-cli/parser/pg_parser_test.go | 111 ------------------- skemr-cli/rulengn/rule_engine_test.go | 7 +- skemr-common/models/database_entity.go | 10 +- 8 files changed, 219 insertions(+), 121 deletions(-) create mode 100644 skemr-cli/parser/pg_parser_column_test.go create mode 100644 skemr-cli/parser/pg_parser_database_test.go create mode 100644 skemr-cli/parser/pg_parser_namespace_test.go create mode 100644 skemr-cli/parser/pg_parser_table_test.go diff --git a/skemr-cli/parser/pg_parser.go b/skemr-cli/parser/pg_parser.go index 66a19e7..df457ce 100644 --- a/skemr-cli/parser/pg_parser.go +++ b/skemr-cli/parser/pg_parser.go @@ -17,6 +17,12 @@ type StatementAction struct { type SqlAction string const ( + + // Namespace level actions + SqlActionCreateNamespace SqlAction = "CREATE SCHEMA" + SqlActionRenameNamespace SqlAction = "RENAME SCHEMA" + SqlActionDropNamespace SqlAction = "DROP SCHEMA" + // Database level actions SqlActionCreateDatabase SqlAction = "CREATE DATABASE" SqlActionRenameDatabase SqlAction = "RENAME DATABASE" @@ -72,6 +78,10 @@ func parseStatement(stmt *pgquery.RawStmt, original string) (StatementAction, er statementAction, err = parseCreateDatabaseStmt(createDbStmt) } + if createSchemaStmt := node.GetCreateSchemaStmt(); createSchemaStmt != nil { + statementAction, err = parseCreateSchema(createSchemaStmt) + } + if err != nil { slog.Error("Error parsing statement", "error", err, "statement", stmt.String()) return StatementAction{}, err @@ -169,9 +179,20 @@ func parseCreateDatabaseStmt(createDbStmt *pgquery.CreatedbStmt) (StatementActio }, nil } +func parseCreateSchema(createSchemaStmt *pgquery.CreateSchemaStmt) (StatementAction, error) { + schemaName := createSchemaStmt.Schemaname + action := SqlActionCreateNamespace + + return StatementAction{ + Target: schemaName, + Action: action, + Relation: "", + }, nil +} + func parseRenameStmt(renameStmt *pgquery.RenameStmt) (StatementAction, error) { relName := "" - target := renameStmt.Subname + target := "" action := SqlActionUndefined switch renameStmt.GetRenameType() { @@ -180,15 +201,19 @@ func parseRenameStmt(renameStmt *pgquery.RenameStmt) (StatementAction, error) { action = SqlActionRenameTable target = renameStmt.Relation.Relname - // If renaming a database case pgquery.ObjectType_OBJECT_DATABASE: action = SqlActionRenameDatabase + target = renameStmt.Subname // If renaming a column case pgquery.ObjectType_OBJECT_COLUMN: action = SqlActionRenameColumn + target = renameStmt.Subname relName = renameStmt.Relation.Relname - + // if renaming a namespace + case pgquery.ObjectType_OBJECT_SCHEMA: + action = SqlActionRenameNamespace + target = renameStmt.Subname } return StatementAction{ @@ -212,7 +237,7 @@ func parseDropDatabase(node *pgquery.Node) (StatementAction, error) { func parseDrop(dropStmt *pgquery.DropStmt) (StatementAction, error) { relName := "" target := "" - action := SqlActionDropTable + action := SqlActionUndefined // If we are dropping a table if dropStmt.RemoveType == pgquery.ObjectType_OBJECT_TABLE { @@ -222,6 +247,13 @@ func parseDrop(dropStmt *pgquery.DropStmt) (StatementAction, error) { action = SqlActionDropTable } + // If we are dropping a namespace + if dropStmt.RemoveType == pgquery.ObjectType_OBJECT_SCHEMA { + schemaName := dropStmt.GetObjects()[0].GetString_().GetSval() + target = schemaName + action = SqlActionDropNamespace + } + return StatementAction{ Target: target, Action: action, diff --git a/skemr-cli/parser/pg_parser_column_test.go b/skemr-cli/parser/pg_parser_column_test.go new file mode 100644 index 0000000..4a13932 --- /dev/null +++ b/skemr-cli/parser/pg_parser_column_test.go @@ -0,0 +1,52 @@ +package parser + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestParseSqlDropColumn(t *testing.T) { + sql := "ALTER TABLE rules DROP COLUMN name" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "name", statementAction[0].Target, "Expected target 'name'") + assert.Equal(t, SqlActionDropColumn, statementAction[0].Action, "Expected action 'DROP COLUMN'") + assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") +} + +func TestParseSqlRenameColumn(t *testing.T) { + sql := "ALTER TABLE rules RENAME COLUMN name TO new_name" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "name", statementAction[0].Target, "Expected target 'name'") + assert.Equal(t, SqlActionRenameColumn, statementAction[0].Action, "Expected action 'RENAME COLUMN'") + assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") +} + +func TestParseDropQualifiedColumn(t *testing.T) { + sql := "ALTER TABLE public.rules DROP COLUMN name" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "name", statementAction[0].Target, "Expected target 'name'") + assert.Equal(t, SqlActionDropColumn, statementAction[0].Action, "Expected action 'DROP COLUMN'") + assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") +} + +func TestParseSqlModifyColumnDataType(t *testing.T) { + sql := "ALTER TABLE rules ALTER COLUMN name TYPE VARCHAR(255)" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "name", statementAction[0].Target, "Expected target 'name'") + assert.Equal(t, SqlActionModifyDataType, statementAction[0].Action, "Expected action 'MODIFY DATA TYPE'") + assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") + +} diff --git a/skemr-cli/parser/pg_parser_database_test.go b/skemr-cli/parser/pg_parser_database_test.go new file mode 100644 index 0000000..b132c80 --- /dev/null +++ b/skemr-cli/parser/pg_parser_database_test.go @@ -0,0 +1,40 @@ +package parser + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestParseSqlDropDataBase(t *testing.T) { + sql := "DROP DATABASE postgres" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "postgres", statementAction[0].Target, "Expected target 'postgres'") + assert.Equal(t, SqlActionDropDatabase, statementAction[0].Action, "Expected action 'SqlActionDropDatabase'") + assert.Equal(t, "", statementAction[0].Relation) +} + +func TestParseSqlCreateDataBase(t *testing.T) { + sql := "CREATE DATABASE skemr_db" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "skemr_db", statementAction[0].Target, "Expected target 'skemr_db'") + assert.Equal(t, SqlActionCreateDatabase, statementAction[0].Action, "Expected action 'CREATE DATABASE'") + assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for CREATE DATABASE") +} + +func TestParseSqlRenameDataBase(t *testing.T) { + sql := "ALTER DATABASE skemr_db RENAME TO skemr_database" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "skemr_db", statementAction[0].Target, "Expected target 'skemr_db'") + assert.Equal(t, SqlActionRenameDatabase, statementAction[0].Action, "Expected action 'RENAME DATABASE'") + assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for RENAME DATABASE") +} diff --git a/skemr-cli/parser/pg_parser_namespace_test.go b/skemr-cli/parser/pg_parser_namespace_test.go new file mode 100644 index 0000000..d77e501 --- /dev/null +++ b/skemr-cli/parser/pg_parser_namespace_test.go @@ -0,0 +1,40 @@ +package parser + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestParseSqlDropSchema(t *testing.T) { + sql := "DROP SCHEMA public" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "public", statementAction[0].Target, "Expected target 'public'") + assert.Equal(t, SqlActionDropNamespace, statementAction[0].Action, "Expected action 'DROP SCHEMA'") + assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for DROP SCHEMA") +} + +func TestParseSqlCreateSchema(t *testing.T) { + sql := "CREATE SCHEMA public" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "public", statementAction[0].Target, "Expected target 'public'") + assert.Equal(t, SqlActionCreateNamespace, statementAction[0].Action, "Expected action 'CREATE SCHEMA'") + assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for CREATE SCHEMA") +} + +func TestParseSqlRenameSchema(t *testing.T) { + sql := "ALTER SCHEMA public RENAME TO new_public" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "public", statementAction[0].Target, "Expected target 'public'") + assert.Equal(t, SqlActionRenameNamespace, statementAction[0].Action, "Expected action 'RENAME SCHEMA'") + assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for RENAME SCHEMA") +} diff --git a/skemr-cli/parser/pg_parser_table_test.go b/skemr-cli/parser/pg_parser_table_test.go new file mode 100644 index 0000000..61ae577 --- /dev/null +++ b/skemr-cli/parser/pg_parser_table_test.go @@ -0,0 +1,40 @@ +package parser + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestParseSqlDropTable(t *testing.T) { + sql := "DROP TABLE rules" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "rules", statementAction[0].Target, "Expected empty target for DROP TABLE") + assert.Equal(t, SqlActionDropTable, statementAction[0].Action, "Expected action 'DROP TABLE'") + assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") +} + +func TestParseSqlDropTableCascade(t *testing.T) { + sql := "DROP TABLE rules CASCADE" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "rules", statementAction[0].Target, "Expected target 'rules'") + assert.Equal(t, SqlActionDropTable, statementAction[0].Action, "Expected action 'DROP TABLE'") + assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") +} + +func TestParseSqlRenameTable(t *testing.T) { + sql := "ALTER TABLE rules RENAME TO new_rules" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "rules", statementAction[0].Target, "Expected target 'rules'") + assert.Equal(t, SqlActionRenameTable, statementAction[0].Action, "Expected action 'RENAME TABLE'") + assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for RENAME TABLE") +} diff --git a/skemr-cli/parser/pg_parser_test.go b/skemr-cli/parser/pg_parser_test.go index d1d1bbc..fc10dce 100644 --- a/skemr-cli/parser/pg_parser_test.go +++ b/skemr-cli/parser/pg_parser_test.go @@ -7,40 +7,6 @@ import ( "github.com/stretchr/testify/assert" ) -func TestParseSqlDropColumn(t *testing.T) { - sql := "ALTER TABLE rules DROP COLUMN name" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "name", statementAction[0].Target, "Expected target 'name'") - assert.Equal(t, SqlActionDropColumn, statementAction[0].Action, "Expected action 'DROP COLUMN'") - assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") - -} - -func TestParseSqlDropTable(t *testing.T) { - sql := "DROP TABLE rules" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "rules", statementAction[0].Target, "Expected empty target for DROP TABLE") - assert.Equal(t, SqlActionDropTable, statementAction[0].Action, "Expected action 'DROP TABLE'") - assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") -} - -func TestParseSqlDropTableCascade(t *testing.T) { - sql := "DROP TABLE rules CASCADE" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "rules", statementAction[0].Target, "Expected target 'rules'") - assert.Equal(t, SqlActionDropTable, statementAction[0].Action, "Expected action 'DROP TABLE'") - assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") -} - func TestParseCreateIndex(t *testing.T) { sql := "CREATE INDEX idx_name ON rules (name)" statementAction, err := ParseSql(sql) @@ -52,83 +18,6 @@ func TestParseCreateIndex(t *testing.T) { assert.Equal(t, "", statementAction[0].Relation, "Expected relation ''") } -func TestParseSqlRenameColumn(t *testing.T) { - sql := "ALTER TABLE rules RENAME COLUMN name TO new_name" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "name", statementAction[0].Target, "Expected target 'new_name'") - assert.Equal(t, SqlActionRenameColumn, statementAction[0].Action, "Expected action 'RENAME COLUMN'") - assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") -} - -func TestParseDropQualifiedColumn(t *testing.T) { - sql := "ALTER TABLE public.rules DROP COLUMN name" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "name", statementAction[0].Target, "Expected target 'name'") - assert.Equal(t, SqlActionDropColumn, statementAction[0].Action, "Expected action 'DROP COLUMN'") - assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") -} - -func TestParseSqlModifyDataType(t *testing.T) { - sql := "ALTER TABLE rules ALTER COLUMN name TYPE VARCHAR(255)" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - if statementAction[0].Target != "name" || statementAction[0].Action != SqlActionModifyDataType || statementAction[0].Relation != "rules" { - t.Fatalf("Expected target 'name', action 'MODIFY DATA TYPE', relation 'rules', got target '%s', action '%s', relation '%s'", statementAction[0].Target, statementAction[0].Action, statementAction[0].Relation) - } -} - -func TestParseSqlRenameTable(t *testing.T) { - sql := "ALTER TABLE rules RENAME TO new_rules" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "rules", statementAction[0].Target, "Expected target 'rules'") - assert.Equal(t, SqlActionRenameTable, statementAction[0].Action, "Expected action 'RENAME TABLE'") - assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for RENAME TABLE") -} - -func TestParseSqlDropDataBase(t *testing.T) { - sql := "DROP DATABASE postgres" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "postgres", statementAction[0].Target, "Expected target 'postgres'") - assert.Equal(t, SqlActionDropDatabase, statementAction[0].Action, "Expected action 'SqlActionDropDatabase'") - assert.Equal(t, "", statementAction[0].Relation) -} - -func TestParseSqlCreateDataBase(t *testing.T) { - sql := "CREATE DATABASE skemr_db" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "skemr_db", statementAction[0].Target, "Expected target 'skemr_db'") - assert.Equal(t, SqlActionCreateDatabase, statementAction[0].Action, "Expected action 'CREATE DATABASE'") - assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for CREATE DATABASE") -} - -func TestParseSqlRenameDataBase(t *testing.T) { - sql := "ALTER DATABASE skemr_db RENAME TO skemr_database" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "skemr_db", statementAction[0].Target, "Expected target 'skemr_db'") - assert.Equal(t, SqlActionRenameDatabase, statementAction[0].Action, "Expected action 'RENAME DATABASE'") - assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for RENAME DATABASE") -} - func TestParseSqlUndefined(t *testing.T) { sql := "This is not a valid SQL statement" statementAction, err := ParseSql(sql) diff --git a/skemr-cli/rulengn/rule_engine_test.go b/skemr-cli/rulengn/rule_engine_test.go index a8993bc..5dcaf13 100644 --- a/skemr-cli/rulengn/rule_engine_test.go +++ b/skemr-cli/rulengn/rule_engine_test.go @@ -31,6 +31,11 @@ var entities = []models.DatabaseEntity{ Name: "age", Type: models.DatabaseEntityTypeColumn, }, + { + ID: uuid.New(), + Name: "public", + Type: models.DatabaseEntityTypeNamespace, + }, } func TestProcessStatementsReturnsResult(t *testing.T) { @@ -261,7 +266,7 @@ func TestAdvisoryRuleTrigger(t *testing.T) { } // If tables A and B have identical column names, and there is a rule that locks on column name on table A, -// Then dropping the column on table B should not trigger the rule, but dropping the column on table A should trigger the rule. +// Then dropping the column on table B should not trigger the rule. However, dropping the column on table A should trigger the rule. func TestIdenticalColumnNameRule(t *testing.T) { tableAId := uuid.New() tableBId := uuid.New() diff --git a/skemr-common/models/database_entity.go b/skemr-common/models/database_entity.go index 89181f0..4e75b34 100644 --- a/skemr-common/models/database_entity.go +++ b/skemr-common/models/database_entity.go @@ -16,17 +16,17 @@ const ( ) const ( - DatabaseEntityTypeDatabase DatabaseEntityType = "database" - DatabaseEntityTypeSchema DatabaseEntityType = "schema" - DatabaseEntityTypeTable DatabaseEntityType = "table" - DatabaseEntityTypeColumn DatabaseEntityType = "column" + DatabaseEntityTypeDatabase DatabaseEntityType = "database" + DatabaseEntityTypeNamespace DatabaseEntityType = "namespace" + DatabaseEntityTypeTable DatabaseEntityType = "table" + DatabaseEntityTypeColumn DatabaseEntityType = "column" ) type DatabaseEntity struct { ID uuid.UUID `json:"id"` Name string `json:"name"` // Name of the entity "public", "users", "email", "my_view" Type DatabaseEntityType `json:"type"` - ParentId *uuid.UUID `json:"parentId"` // in case of column, references table. table references schema etc. + ParentId *uuid.UUID `json:"parentId"` // in case of column, references table. table references namespace etc. Status DatabaseEntityStatus `json:"status"` CreatedAt time.Time `json:"createdAt"` DeletedAt *time.Time `json:"deletedAt"` From 20c9a63584026aa5485f67f7e96d8ae2e48c6b7e Mon Sep 17 00:00:00 2001 From: WalMaa Date: Sun, 17 May 2026 15:29:00 +0300 Subject: [PATCH 05/14] added test coverage and default namespace fallback for migration parsing --- .../internal/dbreflect/identity_generator.go | 6 +- skemr-api/internal/dbreflect/schema_sync.go | 18 +-- .../internal/dbreflect/schema_sync_test.go | 16 +-- skemr-cli/cmd/validate.go | 13 +- skemr-cli/controlplaneclient/http_client.go | 18 ++- skemr-cli/parser/pg_parser.go | 46 +++++-- skemr-cli/parser/pg_parser_column_test.go | 20 +++ skemr-cli/parser/pg_parser_table_test.go | 52 ++++++++ skemr-cli/parser/pg_parser_test.go | 11 -- skemr-cli/rulengn/rule_engine.go | 39 +++++- skemr-cli/rulengn/rule_engine_test.go | 120 +++++++++++++++++- skemr-cli/test/sql/migration-3.sql | 6 +- 12 files changed, 295 insertions(+), 70 deletions(-) diff --git a/skemr-api/internal/dbreflect/identity_generator.go b/skemr-api/internal/dbreflect/identity_generator.go index 0206d76..4e7c8af 100644 --- a/skemr-api/internal/dbreflect/identity_generator.go +++ b/skemr-api/internal/dbreflect/identity_generator.go @@ -20,9 +20,9 @@ func GenerateTableFingerprint(tableRef TableRef) string { return fmt.Sprintf("table:%s:%s", tableRef.ColumnShape, tableRef.PrimaryKey) } -// GenerateSchemaFingerprint generates a unique identifier for a schema based on its properties. +// GenerateNamespaceFingerprint generates a unique identifier for a schema based on its properties. // The principle is to create a stable identifier that remains consistent across schema renames and db instance changes (backup restores). // The format is schema:{database_id}:{schema_fingerprint} -func GenerateSchemaFingerprint(schemaRef SchemaRef, databaseId uuid.UUID) string { - return fmt.Sprintf("schema:%s:%s", databaseId.String(), schemaRef.Fingerprint) +func GenerateNamespaceFingerprint(schemaRef SchemaRef) string { + return fmt.Sprintf("namespace:%s", schemaRef.Fingerprint) } diff --git a/skemr-api/internal/dbreflect/schema_sync.go b/skemr-api/internal/dbreflect/schema_sync.go index c45d5f4..bbf3359 100644 --- a/skemr-api/internal/dbreflect/schema_sync.go +++ b/skemr-api/internal/dbreflect/schema_sync.go @@ -114,7 +114,7 @@ func (s *SchemaSyncService) SyncSchema(c context.Context, database models.Databa // For each schema, get tables and columns for _, schema := range schemaRefs { - schema, err := s.updateSchema(c, schema, database) + schema, err := s.updateNamespace(c, schema, database) if err != nil { return err } @@ -126,7 +126,7 @@ func (s *SchemaSyncService) SyncSchema(c context.Context, database models.Databa return fmt.Errorf("error getting tables in schema %q: %w", schema.Name, err) } for _, tableRef := range tables { - table, err := s.UpdateTable(c, tableRef, database, schema.ID) + table, err := s.updateTable(c, tableRef, database, schema.ID) if err != nil { return fmt.Errorf("error updating tables: %w", err) @@ -139,7 +139,7 @@ func (s *SchemaSyncService) SyncSchema(c context.Context, database models.Databa return fmt.Errorf("error getting columns in table %q.%q: %w", schema.Name, tableRef.Name, err) } for _, column := range columns { - column, err := s.SyncColumn(c, column, database, table.ID) + column, err := s.updateColumn(c, column, database, table.ID) if err != nil { return fmt.Errorf("Error updating column: %w", err) } @@ -166,10 +166,10 @@ func (s *SchemaSyncService) SyncSchema(c context.Context, database models.Databa return nil } -// updateSchema checks if a schema with the given name exists for the database. +// updateNamespace checks if a namespace (schema in postgres) with the given name exists for the database. // If it does not exist, it creates a new schema entity. // If it does exist, it currently does nothing but can be extended to update schema attributes if needed. -func (s *SchemaSyncService) updateSchema(c context.Context, schemaRef SchemaRef, database models.Database) (sqlc.DatabaseEntity, error) { +func (s *SchemaSyncService) updateNamespace(c context.Context, schemaRef SchemaRef, database models.Database) (sqlc.DatabaseEntity, error) { args := sqlc.GetDatabaseEntityByDatabaseIdAndTypeAndParentAndNameParams{ DatabaseID: database.ID, EntityType: sqlc.DatabaseEntityTypeSchema, @@ -185,7 +185,7 @@ func (s *SchemaSyncService) updateSchema(c context.Context, schemaRef SchemaRef, // If that schema does not exist by name, check by fingerprint to see if it is the same schema with an updated name. - fingerprint := GenerateSchemaFingerprint(schemaRef, database.ID) + fingerprint := GenerateNamespaceFingerprint(schemaRef) schema, err = s.db.GetDatabaseEntityByFingerprint(c, sqlc.GetDatabaseEntityByFingerprintParams{ DatabaseID: database.ID, @@ -251,7 +251,7 @@ func (s *SchemaSyncService) markEntityAsDeleted(c context.Context, entityId uuid return nil } -func (s *SchemaSyncService) UpdateTable(c context.Context, tableRef TableRef, database models.Database, schemaId uuid.UUID) (sqlc.DatabaseEntity, error) { +func (s *SchemaSyncService) updateTable(c context.Context, tableRef TableRef, database models.Database, schemaId uuid.UUID) (sqlc.DatabaseEntity, error) { args := sqlc.GetDatabaseEntityByDatabaseIdAndTypeAndParentAndNameParams{ DatabaseID: database.ID, ParentID: &schemaId, @@ -322,8 +322,8 @@ func (s *SchemaSyncService) UpdateTable(c context.Context, tableRef TableRef, da return table, err } -// SyncColumn syncs a database column to a database entity. -func (s *SchemaSyncService) SyncColumn(c context.Context, columnRef ColumnRef, database models.Database, tableId uuid.UUID) (sqlc.DatabaseEntity, error) { +// updateColumn syncs a database column to a database entity. +func (s *SchemaSyncService) updateColumn(c context.Context, columnRef ColumnRef, database models.Database, tableId uuid.UUID) (sqlc.DatabaseEntity, error) { args := sqlc.GetDatabaseEntityByDatabaseIdAndTypeAndParentAndNameParams{ DatabaseID: database.ID, ParentID: &tableId, diff --git a/skemr-api/internal/dbreflect/schema_sync_test.go b/skemr-api/internal/dbreflect/schema_sync_test.go index 9eb9f43..c9e142c 100644 --- a/skemr-api/internal/dbreflect/schema_sync_test.go +++ b/skemr-api/internal/dbreflect/schema_sync_test.go @@ -6,7 +6,6 @@ import ( "github.com/google/uuid" "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgtype" "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" "github.com/walmaa/skemr-api/db/sqlc" @@ -86,17 +85,14 @@ func TestUpdateSchemaCreatesNew(t *testing.T) { mockDB.On("GetDatabaseEntityByDatabaseIdAndTypeAndParentAndName", mock.Anything, mock.Anything).Return(sqlc.DatabaseEntity{}, pgx.ErrNoRows) mockDB.On("GetDatabaseEntityByFingerprint", mock.Anything, mock.Anything).Return(sqlc.DatabaseEntity{}, pgx.ErrNoRows) mockDB.On("CreateDatabaseEntity", mock.Anything, mock.Anything).Return(sqlc.DatabaseEntity{ - ID: uuid.New(), - Name: schemaName, - DatabaseID: dataBaseId, - Fingerprint: pgtype.Text{ - String: "fingerprint", - Valid: true, - }, + ID: uuid.New(), + Name: schemaName, + DatabaseID: dataBaseId, + Fingerprint: "fingerprint", }, nil) syncService := NewSchemaSyncService(mockDB, func(_ models.Database) DatabaseConnector { return mockConnector }) - schema, err := syncService.updateSchema(c, schemaRef, database) + schema, err := syncService.updateNamespace(c, schemaRef, database) require.NoError(t, err) require.Equal(t, schema.Name, schemaName) @@ -127,7 +123,7 @@ func TestUpdateSchemaUpdatesExisting(t *testing.T) { }, nil) syncService := NewSchemaSyncService(mockDB, func(_ models.Database) DatabaseConnector { return mockConnector }) - schema, err := syncService.updateSchema(c, schemaRef, database) + schema, err := syncService.updateNamespace(c, schemaRef, database) require.NoError(t, err) require.Equal(t, schema.Name, schemaName) diff --git a/skemr-cli/cmd/validate.go b/skemr-cli/cmd/validate.go index 381c036..58578a3 100644 --- a/skemr-cli/cmd/validate.go +++ b/skemr-cli/cmd/validate.go @@ -47,6 +47,10 @@ var validateCmd = &cobra.Command{ slog.Debug("Validate command executed", "project_id", projectId, "database_id", databaseId, "migration_files_dir", migrationFilesDir) ruleEngine := rulengn.NewRuleEngine() + // Process files + filePaths := make([]string, 0) + collectFilePathsFromDir(&filePaths, migrationFilesDir) + // Get rules rules, err := controlplaneclient.GetRules(c, projectId, databaseId, token) @@ -56,17 +60,14 @@ var validateCmd = &cobra.Command{ } slog.Debug("Fetched rules from control plane", "ruleCount", len(rules)) - // Fetch all database entities. This is used to match columns to right parents (tables) + // Fetch all database entities. This is used to match columns to the right parents (tables) entities, err := controlplaneclient.GetDatabaseEntities(c, projectId, databaseId, token) - slog.Debug("Fetched database entities from control plane", "entityCount", len(entities)) if err != nil { + slog.Error("Error fetching database entities from control plane", "err", err) os.Exit(1) } - - // Process files - filePaths := make([]string, 0) - collectFilePathsFromDir(&filePaths, migrationFilesDir) + slog.Debug("Fetched database entities from control plane", "entityCount", len(entities)) // Rule check dtos := make([]rulengn.MigrationFileDto, len(filePaths)) diff --git a/skemr-cli/controlplaneclient/http_client.go b/skemr-cli/controlplaneclient/http_client.go index 5c31bde..0f4922c 100644 --- a/skemr-cli/controlplaneclient/http_client.go +++ b/skemr-cli/controlplaneclient/http_client.go @@ -81,12 +81,17 @@ func GetDatabaseEntity(ctx context.Context, projectId string, databaseId string, return nil, err } + defer func(Body io.ReadCloser) { + err := Body.Close() + if err != nil { + slog.Error("Error closing response body", "error", err) + } + }(resp.Body) + if resp.StatusCode != http.StatusOK { return nil, fmt.Errorf("Error getting database entity, status code: %d", resp.StatusCode) } - defer resp.Body.Close() - var out models.DatabaseEntity if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { slog.Error("Error decoding response body", "error", err) @@ -116,12 +121,17 @@ func GetDatabaseEntities(ctx context.Context, projectId string, databaseId strin return nil, err } + defer func(Body io.ReadCloser) { + err := Body.Close() + if err != nil { + slog.Error("Error closing response body", "error", err) + } + }(resp.Body) + if resp.StatusCode != http.StatusOK { return nil, fmt.Errorf("Error getting database entities, status code: %d", resp.StatusCode) } - defer resp.Body.Close() - var out []models.DatabaseEntity if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { return nil, err diff --git a/skemr-cli/parser/pg_parser.go b/skemr-cli/parser/pg_parser.go index df457ce..d504751 100644 --- a/skemr-cli/parser/pg_parser.go +++ b/skemr-cli/parser/pg_parser.go @@ -8,10 +8,11 @@ import ( ) type StatementAction struct { - Target string // e.g., column name, table name, database name - Action SqlAction // The type of action performed (e.g., CREATE, DROP, ALTER) - Relation string // e.g., table name for column actions - Original string // The original SQL statement for reference + Target string // e.g., column name, table name, database name + Action SqlAction // The type of action performed (e.g., CREATE, DROP, ALTER) + Relation string // e.g., table name for column actions + Namespace string // Namespace for database and table level actions + Original string // The original SQL statement for reference } type SqlAction string @@ -160,11 +161,13 @@ func parseCreateStmt(createStmt *pgquery.CreateStmt) (StatementAction, error) { relName := createStmt.Relation.Relname target := relName action := SqlActionCreateTable + namespace := createStmt.Relation.Schemaname return StatementAction{ - Target: target, - Action: action, - Relation: relName, + Target: target, + Action: action, + Relation: relName, + Namespace: namespace, }, nil } @@ -194,6 +197,7 @@ func parseRenameStmt(renameStmt *pgquery.RenameStmt) (StatementAction, error) { relName := "" target := "" action := SqlActionUndefined + namespace := "" switch renameStmt.GetRenameType() { // If renaming a table @@ -201,6 +205,7 @@ func parseRenameStmt(renameStmt *pgquery.RenameStmt) (StatementAction, error) { action = SqlActionRenameTable target = renameStmt.Relation.Relname + namespace = renameStmt.Relation.Schemaname // If renaming a database case pgquery.ObjectType_OBJECT_DATABASE: action = SqlActionRenameDatabase @@ -217,9 +222,10 @@ func parseRenameStmt(renameStmt *pgquery.RenameStmt) (StatementAction, error) { } return StatementAction{ - Target: target, - Action: action, - Relation: relName, + Target: target, + Action: action, + Relation: relName, + Namespace: namespace, }, nil } @@ -238,10 +244,21 @@ func parseDrop(dropStmt *pgquery.DropStmt) (StatementAction, error) { relName := "" target := "" action := SqlActionUndefined + namespace := "" // If we are dropping a table if dropStmt.RemoveType == pgquery.ObjectType_OBJECT_TABLE { - tableName := dropStmt.GetObjects()[0].GetList().Items[0].GetString_().GetSval() + // if qualified name (namespace.table), the table name is the second item in the list + statementItems := dropStmt.GetObjects()[0].GetList().Items + tableName := "" + + if len(statementItems) > 1 { + namespace = statementItems[0].GetString_().GetSval() + tableName = statementItems[1].GetString_().GetSval() + } else { + tableName = statementItems[0].GetString_().GetSval() + } + relName = tableName target = tableName action = SqlActionDropTable @@ -255,9 +272,10 @@ func parseDrop(dropStmt *pgquery.DropStmt) (StatementAction, error) { } return StatementAction{ - Target: target, - Action: action, - Relation: relName, + Target: target, + Action: action, + Relation: relName, + Namespace: namespace, }, nil } diff --git a/skemr-cli/parser/pg_parser_column_test.go b/skemr-cli/parser/pg_parser_column_test.go index 4a13932..750ba1c 100644 --- a/skemr-cli/parser/pg_parser_column_test.go +++ b/skemr-cli/parser/pg_parser_column_test.go @@ -50,3 +50,23 @@ func TestParseSqlModifyColumnDataType(t *testing.T) { assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") } + +func TestParseSqlAddColumn(t *testing.T) { + sql := "ALTER TABLE rules ADD COLUMN description TEXT" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + assert.Equal(t, "description", statementAction[0].Target, "Expected target 'description'") + assert.Equal(t, SqlActionAddColumn, statementAction[0].Action, "Expected action 'ADD COLUMN'") + assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") +} + +func TestParseSqlAddColumnWithTimestamp(t *testing.T) { + sql := "ALTER TABLE orders ADD COLUMN updated_at TIMESTAMP;" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + assert.Equal(t, "updated_at", statementAction[0].Target, "Expected target 'updated_at'") + assert.Equal(t, SqlActionAddColumn, statementAction[0].Action, "Expected action 'ADD COLUMN'") + assert.Equal(t, "orders", statementAction[0].Relation, "Expected relation 'orders'") +} diff --git a/skemr-cli/parser/pg_parser_table_test.go b/skemr-cli/parser/pg_parser_table_test.go index 61ae577..79bbdc3 100644 --- a/skemr-cli/parser/pg_parser_table_test.go +++ b/skemr-cli/parser/pg_parser_table_test.go @@ -17,6 +17,25 @@ func TestParseSqlDropTable(t *testing.T) { assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") } +func BenchmarkParseSqlDropTable(b *testing.B) { + for b.Loop() { + sql := "DROP TABLE rules" + _, err := ParseSql(sql) + assert.Nil(b, err) + } +} + +func TestParseSqlCreateTable(t *testing.T) { + sql := "CREATE TABLE rules (id SERIAL PRIMARY KEY, name VARCHAR(255))" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "rules", statementAction[0].Target, "Expected target 'rules'") + assert.Equal(t, SqlActionCreateTable, statementAction[0].Action, "Expected action 'CREATE TABLE'") + assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") +} + func TestParseSqlDropTableCascade(t *testing.T) { sql := "DROP TABLE rules CASCADE" statementAction, err := ParseSql(sql) @@ -38,3 +57,36 @@ func TestParseSqlRenameTable(t *testing.T) { assert.Equal(t, SqlActionRenameTable, statementAction[0].Action, "Expected action 'RENAME TABLE'") assert.Equal(t, "", statementAction[0].Relation, "Expected empty relation for RENAME TABLE") } + +func TestParseDropQualifiedTable(t *testing.T) { + sql := "DROP TABLE other.rules" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "rules", statementAction[0].Target, "Expected target 'rules'") + assert.Equal(t, SqlActionDropTable, statementAction[0].Action, "Expected action 'DROP TABLE'") + assert.Equal(t, "other", statementAction[0].Namespace, "Expected namespace 'other'") +} + +func TestParseSqlRenameQualifiedTable(t *testing.T) { + sql := "ALTER TABLE other.rules RENAME TO new_rules" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "rules", statementAction[0].Target, "Expected target 'rules'") + assert.Equal(t, SqlActionRenameTable, statementAction[0].Action, "Expected action 'RENAME TABLE'") + assert.Equal(t, "other", statementAction[0].Namespace, "Expected namespace 'other'") +} + +func TestParseSqlCreateQualifiedTable(t *testing.T) { + sql := "CREATE TABLE other.rules (id SERIAL PRIMARY KEY, name VARCHAR(255))" + statementAction, err := ParseSql(sql) + + assert.Nil(t, err) + + assert.Equal(t, "rules", statementAction[0].Target, "Expected target 'rules'") + assert.Equal(t, SqlActionCreateTable, statementAction[0].Action, "Expected action 'CREATE TABLE'") + assert.Equal(t, "other", statementAction[0].Namespace, "Expected namespace 'other'") +} diff --git a/skemr-cli/parser/pg_parser_test.go b/skemr-cli/parser/pg_parser_test.go index fc10dce..c1a6809 100644 --- a/skemr-cli/parser/pg_parser_test.go +++ b/skemr-cli/parser/pg_parser_test.go @@ -26,17 +26,6 @@ func TestParseSqlUndefined(t *testing.T) { assert.Nil(t, statementAction, "Expected statementAction to be nil for invalid SQL") } -func TestParseSqlAddColumn(t *testing.T) { - sql := "ALTER TABLE rules ADD COLUMN description TEXT" - statementAction, err := ParseSql(sql) - - assert.Nil(t, err) - - assert.Equal(t, "description", statementAction[0].Target, "Expected target 'description'") - assert.Equal(t, SqlActionAddColumn, statementAction[0].Action, "Expected action 'ADD COLUMN'") - assert.Equal(t, "rules", statementAction[0].Relation, "Expected relation 'rules'") -} - func TestParseSqlInsertRow(t *testing.T) { sql := "INSERT INTO rules (name, scope) VALUES ('rule1', 'table')" statementAction, err := ParseSql(sql) diff --git a/skemr-cli/rulengn/rule_engine.go b/skemr-cli/rulengn/rule_engine.go index b8622ca..00ce929 100644 --- a/skemr-cli/rulengn/rule_engine.go +++ b/skemr-cli/rulengn/rule_engine.go @@ -39,7 +39,7 @@ func (r *RuleEngine) ProcessMigrationFiles(c context.Context, statements []Migra stmt := statement wg.Go(func() { slog.Debug("Processing migration file", "file", stmt.File) - stmtResults, err := r.CheckStatement(stmt, rules, entities) + stmtResults, err := r.checkStatement(stmt, rules, entities) if err != nil { slog.Error("Error checking statement", slog.String("statement", stmt.File), slog.String("error", err.Error())) return @@ -71,8 +71,8 @@ func (r *RuleEngine) ProcessMigrationFiles(c context.Context, statements []Migra return resultsSlice, nil } -// CheckStatement checks if the given SQL statement matches any rules in the database for the specified project. -func (r *RuleEngine) CheckStatement(migrationFileDto MigrationFileDto, rules []models.Rule, entities []models.DatabaseEntity) ([]StatementResult, error) { +// checkStatement checks if the given SQL statement matches any rules in the database for the specified project. +func (r *RuleEngine) checkStatement(migrationFileDto MigrationFileDto, rules []models.Rule, entities []models.DatabaseEntity) ([]StatementResult, error) { slog.Debug("Checking migration file", "file", migrationFileDto.File) file, err := os.ReadFile(migrationFileDto.File) @@ -92,7 +92,7 @@ func (r *RuleEngine) CheckStatement(migrationFileDto MigrationFileDto, rules []m for _, rule := range rules { slog.Debug("Evaluating rule against migration file", slog.String("rule_name", rule.Name), slog.String("migration_file", migrationFileDto.File)) for _, action := range statementActions { - // If the database entity is a column, it is not enough to match the name of the column in the rule with the name of the column in the statement action, + // If the database entity is a column, it is not enough to match the name of the column in the rule with the name of the column in the statement action; // we also need to check if the columns are in the same table. This is because there could be multiple columns with the same name in different tables, // and we don't want to trigger a rule violation if the column in the statement action is not the same as the column in the rule. if rule.DataBaseEntity.Type == models.DatabaseEntityTypeColumn { @@ -107,7 +107,7 @@ func (r *RuleEngine) CheckStatement(migrationFileDto MigrationFileDto, rules []m } parentEntity := entities[i] - slog.Debug("Found parent database entity for rule", slog.String("rule_name", rule.Name), slog.String("parent_entity_name", parentEntity.Name)) + slog.Debug("Found parent database entity (table) for rule", slog.String("rule_name", rule.Name), slog.String("parent_entity_name", parentEntity.Name)) // Check if the parent entity name matches the table name in the statement action if parentEntity.Name != action.Relation { @@ -115,6 +115,35 @@ func (r *RuleEngine) CheckStatement(migrationFileDto MigrationFileDto, rules []m continue } } + + // Same logic for tables, we need to check if the namespace matches. + if rule.DataBaseEntity.Type == models.DatabaseEntityTypeTable { + // Get the parent database entity (namespace) for the table in the rule + i := slices.IndexFunc(entities, func(entity models.DatabaseEntity) bool { + return entity.ID == *rule.DataBaseEntity.ParentId + }) + + if i == -1 { + slog.Warn("Parent database entity not found for rule", slog.String("rule_name", rule.Name), slog.String("parent_id", rule.DataBaseEntity.ParentId.String())) + continue + } + + parentEntity := entities[i] + slog.Debug("Found parent database (namespace) entity for rule", slog.String("rule_name", rule.Name), slog.String("parent_entity_name", parentEntity.Name)) + + // Check if the parent entity name matches the namespace in the statement action + namespace := action.Namespace + if namespace == "" { + slog.Debug("Action does not have a defined namespace, using default namespace of 'public'", slog.String("rule_name", rule.Name), slog.String("action_namespace", action.Namespace)) + namespace = "public" + } + + if parentEntity.Name != namespace { + slog.Debug("Parent entity name does not match action namespace, skipping rule evaluation for this action", slog.String("rule_name", rule.Name), slog.String("parent_entity_name", parentEntity.Name), slog.String("action_namespace", action.Namespace)) + continue + } + } + if rule.DataBaseEntity.Name == action.Target { slog.Debug("Rule target matches migrationFileDto target", slog.String("rule_database_entity", rule.DataBaseEntity.Name), slog.String("statement_target", action.Target)) switch rule.RuleType { diff --git a/skemr-cli/rulengn/rule_engine_test.go b/skemr-cli/rulengn/rule_engine_test.go index 5dcaf13..86c0353 100644 --- a/skemr-cli/rulengn/rule_engine_test.go +++ b/skemr-cli/rulengn/rule_engine_test.go @@ -3,6 +3,7 @@ package rulengn import ( "os" "path/filepath" + "strconv" "testing" "github.com/google/uuid" @@ -91,7 +92,7 @@ func TestDeprecatedRuleTrigger(t *testing.T) { }, } - result, err := ruleEngine.CheckStatement(MigrationFileDto{File: migrationFile}, rules, entities) + result, err := ruleEngine.checkStatement(MigrationFileDto{File: migrationFile}, rules, entities) assert.NoError(t, err) // Assert the result @@ -124,7 +125,7 @@ func TestWarnRuleTrigger(t *testing.T) { }, } - result, err := ruleEngine.CheckStatement(MigrationFileDto{File: migrationFile}, rules, entities) + result, err := ruleEngine.checkStatement(MigrationFileDto{File: migrationFile}, rules, entities) assert.NoError(t, err) // Assert the result @@ -157,7 +158,7 @@ func TestLockedRuleViolation(t *testing.T) { }, } - result, err := ruleEngine.CheckStatement(MigrationFileDto{File: migrationFile}, rules, entities) + result, err := ruleEngine.checkStatement(MigrationFileDto{File: migrationFile}, rules, entities) assert.NoError(t, err) // Assert the result @@ -190,7 +191,7 @@ func TestLockedTableRuleViolationOnColumnAdd(t *testing.T) { }, } - result, err := ruleEngine.CheckStatement(MigrationFileDto{File: migrationFile}, rules, entities) + result, err := ruleEngine.checkStatement(MigrationFileDto{File: migrationFile}, rules, entities) assert.NoError(t, err) // Assert the result @@ -222,7 +223,7 @@ func TestLockedColumnRuleViolationWithQualifiedTable(t *testing.T) { }, } - result, err := ruleEngine.CheckStatement(MigrationFileDto{File: migrationFile}, rules, entities) + result, err := ruleEngine.checkStatement(MigrationFileDto{File: migrationFile}, rules, entities) assert.NoError(t, err) assert.Equal(t, 1, len(result)) @@ -255,7 +256,7 @@ func TestAdvisoryRuleTrigger(t *testing.T) { }, } - result, err := ruleEngine.CheckStatement(MigrationFileDto{File: migrationFile}, rules, entities) + result, err := ruleEngine.checkStatement(MigrationFileDto{File: migrationFile}, rules, entities) assert.NoError(t, err) // Assert the result @@ -317,7 +318,7 @@ func TestIdenticalColumnNameRule(t *testing.T) { }, } - result, err := ruleEngine.CheckStatement(MigrationFileDto{File: migrationFile}, rules, entities) + result, err := ruleEngine.checkStatement(MigrationFileDto{File: migrationFile}, rules, entities) assert.NoError(t, err) // Assert the result @@ -326,3 +327,108 @@ func TestIdenticalColumnNameRule(t *testing.T) { assert.Equal(t, "Locked Age Column on Users Table", result[0].Rule.Name) assert.Equal(t, migrationFile, result[0].File) } + +// If namespaces A and B have identical table names, and there is a locked rule on table A +// Then dropping the table on namespace B should not trigger the rule. However, dropping the table on namespace A should trigger the rule. +func TestIdenticalTableNameWithNamespacesRule(t *testing.T) { + namespaceAId := uuid.New() + namespaceBId := uuid.New() + tableA := models.DatabaseEntity{ + ID: uuid.New(), + Name: "users", + Type: models.DatabaseEntityTypeTable, + ParentId: &namespaceAId, + } + tableB := models.DatabaseEntity{ + ID: uuid.New(), + Name: "users", + Type: models.DatabaseEntityTypeTable, + ParentId: &namespaceBId, + } + + entities := []models.DatabaseEntity{ + { + ID: namespaceAId, + Name: "public", + Type: models.DatabaseEntityTypeNamespace, + }, + tableA, + { + ID: namespaceBId, + Name: "other", + Type: models.DatabaseEntityTypeNamespace, + }, + tableB, + } + ruleEngine := NewRuleEngine() + + // Create a temporary migration file + tmpFile, err := os.CreateTemp(t.TempDir(), "migration-*.sql") + assert.NoError(t, err) + defer func() { _ = tmpFile.Close() }() + content := "DROP TABLE other.users;\nDROP TABLE public.users;" + _, err = tmpFile.WriteString(content) + assert.NoError(t, err) + migrationFile := tmpFile.Name() + + rules := []models.Rule{ + { + ID: uuid.New(), + Name: "Locked Users Table in Public Namespace", + RuleType: models.RuleTypeLocked, + DataBaseEntity: tableA, + }, + } + + result, err := ruleEngine.checkStatement(MigrationFileDto{File: migrationFile}, rules, entities) + assert.NoError(t, err) + + // Assert the result + assert.Equal(t, 1, len(result)) + assert.Equal(t, models.RuleTypeLocked, result[0].Type) + assert.Equal(t, "Locked Users Table in Public Namespace", result[0].Rule.Name) + assert.Equal(t, migrationFile, result[0].File) +} + +func BenchmarkRuleEngine_checkStatement(b *testing.B) { + // Create 100 rules + rules := make([]models.Rule, 100) + for i := range rules { + rules[i] = models.Rule{ + ID: uuid.New(), + Name: "Rule " + strconv.Itoa(i), + RuleType: models.RuleTypeDeprecated, + DataBaseEntity: models.DatabaseEntity{ + Name: "age", + }, + } + } + + // create 500 entities + entities := make([]models.DatabaseEntity, 500) + for i := range entities { + entities[i] = models.DatabaseEntity{ + ID: uuid.New(), + Name: "entity" + strconv.Itoa(i), + Type: models.DatabaseEntityTypeTable, + } + } + + ruleEngine := NewRuleEngine() + + // Create a temporary migration file containing 1000 statements + tmpFile, err := os.CreateTemp(b.TempDir(), "migration-*.sql") + assert.NoError(b, err) + defer func() { _ = tmpFile.Close() }() + for i := 0; i < 1000; i++ { + _, err = tmpFile.WriteString("ALTER TABLE users DROP COLUMN age;\n") + assert.NoError(b, err) + } + migrationFile := tmpFile.Name() + + b.ResetTimer() + for b.Loop() { + _, err := ruleEngine.checkStatement(MigrationFileDto{File: migrationFile}, rules, entities) + assert.NoError(b, err) + } +} diff --git a/skemr-cli/test/sql/migration-3.sql b/skemr-cli/test/sql/migration-3.sql index 84b0b24..d7aacd5 100644 --- a/skemr-cli/test/sql/migration-3.sql +++ b/skemr-cli/test/sql/migration-3.sql @@ -1 +1,5 @@ -DROP TABLE orders CASCADE; \ No newline at end of file +DROP TABLE orders CASCADE; + +ALTER TABLE orders ADD COLUMN updated_at TIMESTAMP; + +ALTER SCHEMA analytics RENAME TO console; \ No newline at end of file From 7d91b72dc6e4a215682f81642c217341ecc05fb7 Mon Sep 17 00:00:00 2001 From: WalMaa Date: Mon, 18 May 2026 10:09:47 +0300 Subject: [PATCH 06/14] adjusted cli reporter output for non violating rules --- skemr-cli/reporter/reporter.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/skemr-cli/reporter/reporter.go b/skemr-cli/reporter/reporter.go index f32da45..e5a26ed 100644 --- a/skemr-cli/reporter/reporter.go +++ b/skemr-cli/reporter/reporter.go @@ -114,7 +114,7 @@ func printSection(title string, results []rulengn.StatementResult, isError bool) if isError { message = fmt.Sprintf("rule \"%s\" violated", res.Rule.Name) } - fmt.Fprintf(os.Stdout, " - %-30s by statement \"%s\" in file: %s\n", message, res.Statement, res.File) + fmt.Fprintf(os.Stdout, " - rule \"%-30s\" triggered by statement \"%s\" in file: %s\n", message, res.Statement, res.File) } } From c90cc9c1ff9f4066d19091835bb4be52d8e1872d Mon Sep 17 00:00:00 2001 From: WalMaa Date: Thu, 21 May 2026 10:38:09 +0300 Subject: [PATCH 07/14] rule service tests and migration to interface. schema to include better consistency checks with foreign keys --- .mockery.yml | 9 + skemr-api/cmd/server/main.go | 3 +- .../migrations/20260225122322_init_schema.sql | 30 +- skemr-api/db/queries/database_entities.sql | 8 + skemr-api/db/queries/rules.sql | 36 ++- skemr-api/db/sqlc/database_entities.sql.go | 35 +++ skemr-api/db/sqlc/querier.go | 3 +- skemr-api/db/sqlc/rules.sql.go | 48 +-- .../internal/controller/rule_controller.go | 9 + skemr-api/internal/dto/common.go | 13 +- skemr-api/internal/errormsg/errors.go | 13 +- skemr-api/internal/mapper/mapper_util.go | 2 +- skemr-api/internal/mapper/rule_mapper.go | 4 + skemr-api/internal/service/rule_service.go | 80 ++--- .../internal/service/rule_service_test.go | 281 ++++++++++++++++++ skemr-api/internal/service/scope_resolver.go | 78 +++++ skemr-api/internal/validation/validation.go | 2 + skemr-api/test/mocks/querier_mock.go | 162 +++++++++- 18 files changed, 693 insertions(+), 123 deletions(-) create mode 100644 skemr-api/internal/service/rule_service_test.go create mode 100644 skemr-api/internal/service/scope_resolver.go diff --git a/.mockery.yml b/.mockery.yml index a14e51f..5359c2e 100644 --- a/.mockery.yml +++ b/.mockery.yml @@ -7,3 +7,12 @@ packages: Querier: config: filename: querier_mock.go + github.com/walmaa/skemr-api/internal/service: + interfaces: + ScopeResolver: + config: + filename: scope_resolver_mock.go + RuleStore: + config: + filename: rule_store_mock.go + diff --git a/skemr-api/cmd/server/main.go b/skemr-api/cmd/server/main.go index 6bbf211..4511346 100644 --- a/skemr-api/cmd/server/main.go +++ b/skemr-api/cmd/server/main.go @@ -93,11 +93,12 @@ func main() { }) queries := sqlc.New(conn) + scopeResolver := service.NewScopeResolver(queries) projectService := service.NewProjectService(queries) databaseService := service.NewDatabaseService(queries, taskClient) webhookService := service.NewWebhookService(queries) projectSecretsService := service.NewAccessTokenService(queries) - ruleService := service.NewRuleService(queries) + ruleService := service.NewRuleService(queries, scopeResolver) databaseEntityService := service.NewDatabaseEntityService(queries) integrationService := service.NewIntegrationService(ruleService) diff --git a/skemr-api/db/migrations/20260225122322_init_schema.sql b/skemr-api/db/migrations/20260225122322_init_schema.sql index 852f6f8..fed3ea6 100644 --- a/skemr-api/db/migrations/20260225122322_init_schema.sql +++ b/skemr-api/db/migrations/20260225122322_init_schema.sql @@ -96,7 +96,8 @@ CREATE TABLE databases failed_connection_attempts INTEGER NOT NULL DEFAULT 0, created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), - CONSTRAINT unique_database_name_per_project UNIQUE (display_name, project_id) + CONSTRAINT unique_database_name_per_project UNIQUE (display_name, project_id), + CONSTRAINT databases_id_project_id_unique UNIQUE (id, project_id) ); CREATE TABLE migration_statements @@ -121,13 +122,13 @@ CREATE TABLE database_entities ( id uuid PRIMARY KEY DEFAULT gen_random_uuid(), fingerprint text NOT NULL, -- this is used to track the same entity across syncs even if it is renamed. - project_id uuid NOT NULL REFERENCES projects (id) ON DELETE CASCADE, - database_id uuid NOT NULL REFERENCES databases (id) ON DELETE CASCADE, + project_id uuid NOT NULL, + database_id uuid NOT NULL, status database_entity_status NOT NULL DEFAULT 'active', deleted_at TIMESTAMPTZ NULL, -- Set when status is 'deleted' to track when it was deleted first_seen_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), -- Track when we first saw this entity entity_type database_entity_type NOT NULL, - parent_id uuid NULL REFERENCES database_entities (id), + parent_id uuid NULL, -- generic identity at this node name text NOT NULL, -- e.g. "public", "users", "email", "my_view" @@ -135,7 +136,20 @@ CREATE TABLE database_entities created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + CONSTRAINT database_entities_parent_same_database_fkey + FOREIGN KEY (parent_id, database_id) + REFERENCES database_entities (id, database_id) ON DELETE CASCADE, -- Ensure parent_id references an entity in the same database + + CONSTRAINT database_entities_database_project_fkey + FOREIGN KEY (database_id, project_id) + REFERENCES databases (id, project_id) + ON DELETE CASCADE, -- Ensure entities are only in the same project + + CONSTRAINT database_entities_id_database_id_unique + UNIQUE (id, database_id), -- for rule composite fkey + UNIQUE NULLS NOT DISTINCT (database_id, name, entity_type, parent_id) -- Ensure we do not map the same entity twice, use NULLS NOT DISTINCT so parentless are not duplicated + ); @@ -146,8 +160,12 @@ CREATE TABLE rules name TEXT NOT NULL, -- Defined by user type rule_type NOT NULL, attributes jsonb NOT NULL DEFAULT '{}'::jsonb, -- Metadata about the rule, removal_date for deprecated types for example - database_entity_id uuid NOT NULL REFERENCES database_entities (id) ON DELETE CASCADE, + database_entity_id uuid NOT NULL REFERENCES database_entities (id), database_id uuid NOT NULL REFERENCES databases (id) ON DELETE CASCADE, created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), - CONSTRAINT unique_rule_name_per_database UNIQUE (name, database_id) + CONSTRAINT unique_rule_name_per_database + UNIQUE (name, database_id), + CONSTRAINT rules_database_entity_database_fkey + FOREIGN KEY (database_entity_id, database_id) + REFERENCES database_entities (id, database_id) ON DELETE CASCADE --- ensure rules are only applied to entities in the same database ); \ No newline at end of file diff --git a/skemr-api/db/queries/database_entities.sql b/skemr-api/db/queries/database_entities.sql index e2a1f42..d03ceb1 100644 --- a/skemr-api/db/queries/database_entities.sql +++ b/skemr-api/db/queries/database_entities.sql @@ -11,6 +11,14 @@ WHERE id = @id AND project_id = @project_id LIMIT 1; +-- name: GetDatabaseEntityByProjectIdDatabaseIdAndId :one +SELECT * +FROM database_entities +WHERE id = @id + AND project_id = @project_id + AND database_id = @database_id +LIMIT 1; + -- name: GetDatabaseEntitiesByProjectId :many SELECT * FROM database_entities diff --git a/skemr-api/db/queries/rules.sql b/skemr-api/db/queries/rules.sql index fee0231..fd58c7a 100644 --- a/skemr-api/db/queries/rules.sql +++ b/skemr-api/db/queries/rules.sql @@ -1,22 +1,26 @@ -- name: GetRule :one SELECT * -FROM rules -WHERE database_id = @database_id AND id = @rule_id +FROM rules r +WHERE r.database_id = @database_id + AND r.id = @rule_id LIMIT 1; -- name: GetRuleByDatabaseAndName :one SELECT * FROM rules -WHERE database_id = @database_id AND name = @name +WHERE database_id = @database_id + AND name = @name LIMIT 1; -- name: GetRuleWithEntity :one -SELECT - sqlc.embed(r), - sqlc.embed(de) +SELECT sqlc.embed(r), + sqlc.embed(de) FROM rules r -JOIN database_entities de ON r.database_entity_id = de.id -WHERE r.database_id = @database_id AND r.id = @rule_id + JOIN databases d ON r.database_id = d.id + JOIN database_entities de ON r.database_entity_id = de.id +WHERE d.project_id = @project_id + AND r.database_id = @database_id + AND r.id = @rule_id LIMIT 1; -- name: CreateRule :one @@ -35,7 +39,8 @@ RETURNING *; -- name: DeleteRule :exec DELETE FROM rules -WHERE database_id = @database_id AND id = @rule_id; +WHERE database_id = @database_id + AND id = @rule_id; -- name: ListRulesByDatabaseId :many @@ -44,12 +49,13 @@ FROM rules WHERE database_id = @database_id; -- name: GetRulesWithEntities :many -SELECT - sqlc.embed(r), - sqlc.embed(de) -FROM rules r -JOIN database_entities de ON r.database_entity_id = de.id -WHERE r.database_id = @database_id; +SELECT sqlc.embed(rules), + sqlc.embed(database_entities) +FROM rules + JOIN databases ON rules.database_id = databases.id + JOIN database_entities ON rules.database_entity_id = database_entities.id +WHERE rules.database_id = @database_id + AND databases.project_id = @project_id; -- name: ListRulesByCriteria :many diff --git a/skemr-api/db/sqlc/database_entities.sql.go b/skemr-api/db/sqlc/database_entities.sql.go index 3049476..a19228f 100644 --- a/skemr-api/db/sqlc/database_entities.sql.go +++ b/skemr-api/db/sqlc/database_entities.sql.go @@ -362,6 +362,41 @@ func (q *Queries) GetDatabaseEntityByProjectIdAndId(ctx context.Context, arg Get return i, err } +const getDatabaseEntityByProjectIdDatabaseIdAndId = `-- name: GetDatabaseEntityByProjectIdDatabaseIdAndId :one +SELECT id, fingerprint, project_id, database_id, status, deleted_at, first_seen_at, entity_type, parent_id, name, attributes, created_at +FROM database_entities +WHERE id = $1 + AND project_id = $2 + AND database_id = $3 +LIMIT 1 +` + +type GetDatabaseEntityByProjectIdDatabaseIdAndIdParams struct { + ID uuid.UUID `json:"id"` + ProjectID uuid.UUID `json:"project_id"` + DatabaseID uuid.UUID `json:"database_id"` +} + +func (q *Queries) GetDatabaseEntityByProjectIdDatabaseIdAndId(ctx context.Context, arg GetDatabaseEntityByProjectIdDatabaseIdAndIdParams) (DatabaseEntity, error) { + row := q.db.QueryRow(ctx, getDatabaseEntityByProjectIdDatabaseIdAndId, arg.ID, arg.ProjectID, arg.DatabaseID) + var i DatabaseEntity + err := row.Scan( + &i.ID, + &i.Fingerprint, + &i.ProjectID, + &i.DatabaseID, + &i.Status, + &i.DeletedAt, + &i.FirstSeenAt, + &i.EntityType, + &i.ParentID, + &i.Name, + &i.Attributes, + &i.CreatedAt, + ) + return i, err +} + const updateDatabaseEntity = `-- name: UpdateDatabaseEntity :one UPDATE database_entities SET name = COALESCE($1, name), diff --git a/skemr-api/db/sqlc/querier.go b/skemr-api/db/sqlc/querier.go index c8a05dc..5232f56 100644 --- a/skemr-api/db/sqlc/querier.go +++ b/skemr-api/db/sqlc/querier.go @@ -32,6 +32,7 @@ type Querier interface { GetDatabaseEntityByDatabaseIdAndTypeAndParentAndName(ctx context.Context, arg GetDatabaseEntityByDatabaseIdAndTypeAndParentAndNameParams) (DatabaseEntity, error) GetDatabaseEntityByFingerprint(ctx context.Context, arg GetDatabaseEntityByFingerprintParams) (DatabaseEntity, error) GetDatabaseEntityByProjectIdAndId(ctx context.Context, arg GetDatabaseEntityByProjectIdAndIdParams) (DatabaseEntity, error) + GetDatabaseEntityByProjectIdDatabaseIdAndId(ctx context.Context, arg GetDatabaseEntityByProjectIdDatabaseIdAndIdParams) (DatabaseEntity, error) GetHashByPrefixAndProjectID(ctx context.Context, arg GetHashByPrefixAndProjectIDParams) (string, error) GetProject(ctx context.Context, id uuid.UUID) (Project, error) GetProjectAccessTokens(ctx context.Context, projectID uuid.UUID) ([]ProjectAccessToken, error) @@ -41,7 +42,7 @@ type Querier interface { GetRule(ctx context.Context, arg GetRuleParams) (Rule, error) GetRuleByDatabaseAndName(ctx context.Context, arg GetRuleByDatabaseAndNameParams) (Rule, error) GetRuleWithEntity(ctx context.Context, arg GetRuleWithEntityParams) (GetRuleWithEntityRow, error) - GetRulesWithEntities(ctx context.Context, databaseID uuid.UUID) ([]GetRulesWithEntitiesRow, error) + GetRulesWithEntities(ctx context.Context, arg GetRulesWithEntitiesParams) ([]GetRulesWithEntitiesRow, error) ListDatabasesByProject(ctx context.Context, projectID uuid.UUID) ([]Database, error) ListRulesByCriteria(ctx context.Context, arg ListRulesByCriteriaParams) ([]Rule, error) ListRulesByDatabaseId(ctx context.Context, databaseID uuid.UUID) ([]Rule, error) diff --git a/skemr-api/db/sqlc/rules.sql.go b/skemr-api/db/sqlc/rules.sql.go index 7ea51da..fe6b61c 100644 --- a/skemr-api/db/sqlc/rules.sql.go +++ b/skemr-api/db/sqlc/rules.sql.go @@ -50,7 +50,8 @@ func (q *Queries) CreateRule(ctx context.Context, arg CreateRuleParams) (Rule, e const deleteRule = `-- name: DeleteRule :exec DELETE FROM rules -WHERE database_id = $1 AND id = $2 +WHERE database_id = $1 + AND id = $2 ` type DeleteRuleParams struct { @@ -65,8 +66,9 @@ func (q *Queries) DeleteRule(ctx context.Context, arg DeleteRuleParams) error { const getRule = `-- name: GetRule :one SELECT id, name, type, attributes, database_entity_id, database_id, created_at -FROM rules -WHERE database_id = $1 AND id = $2 +FROM rules r +WHERE r.database_id = $1 + AND r.id = $2 LIMIT 1 ` @@ -93,7 +95,8 @@ func (q *Queries) GetRule(ctx context.Context, arg GetRuleParams) (Rule, error) const getRuleByDatabaseAndName = `-- name: GetRuleByDatabaseAndName :one SELECT id, name, type, attributes, database_entity_id, database_id, created_at FROM rules -WHERE database_id = $1 AND name = $2 +WHERE database_id = $1 + AND name = $2 LIMIT 1 ` @@ -118,16 +121,19 @@ func (q *Queries) GetRuleByDatabaseAndName(ctx context.Context, arg GetRuleByDat } const getRuleWithEntity = `-- name: GetRuleWithEntity :one -SELECT - r.id, r.name, r.type, r.attributes, r.database_entity_id, r.database_id, r.created_at, - de.id, de.fingerprint, de.project_id, de.database_id, de.status, de.deleted_at, de.first_seen_at, de.entity_type, de.parent_id, de.name, de.attributes, de.created_at +SELECT r.id, r.name, r.type, r.attributes, r.database_entity_id, r.database_id, r.created_at, + de.id, de.fingerprint, de.project_id, de.database_id, de.status, de.deleted_at, de.first_seen_at, de.entity_type, de.parent_id, de.name, de.attributes, de.created_at FROM rules r -JOIN database_entities de ON r.database_entity_id = de.id -WHERE r.database_id = $1 AND r.id = $2 + JOIN databases d ON r.database_id = d.id + JOIN database_entities de ON r.database_entity_id = de.id +WHERE d.project_id = $1 + AND r.database_id = $2 + AND r.id = $3 LIMIT 1 ` type GetRuleWithEntityParams struct { + ProjectID uuid.UUID `json:"project_id"` DatabaseID uuid.UUID `json:"database_id"` RuleID uuid.UUID `json:"rule_id"` } @@ -138,7 +144,7 @@ type GetRuleWithEntityRow struct { } func (q *Queries) GetRuleWithEntity(ctx context.Context, arg GetRuleWithEntityParams) (GetRuleWithEntityRow, error) { - row := q.db.QueryRow(ctx, getRuleWithEntity, arg.DatabaseID, arg.RuleID) + row := q.db.QueryRow(ctx, getRuleWithEntity, arg.ProjectID, arg.DatabaseID, arg.RuleID) var i GetRuleWithEntityRow err := row.Scan( &i.Rule.ID, @@ -165,21 +171,27 @@ func (q *Queries) GetRuleWithEntity(ctx context.Context, arg GetRuleWithEntityPa } const getRulesWithEntities = `-- name: GetRulesWithEntities :many -SELECT - r.id, r.name, r.type, r.attributes, r.database_entity_id, r.database_id, r.created_at, - de.id, de.fingerprint, de.project_id, de.database_id, de.status, de.deleted_at, de.first_seen_at, de.entity_type, de.parent_id, de.name, de.attributes, de.created_at -FROM rules r -JOIN database_entities de ON r.database_entity_id = de.id -WHERE r.database_id = $1 +SELECT rules.id, rules.name, rules.type, rules.attributes, rules.database_entity_id, rules.database_id, rules.created_at, + database_entities.id, database_entities.fingerprint, database_entities.project_id, database_entities.database_id, database_entities.status, database_entities.deleted_at, database_entities.first_seen_at, database_entities.entity_type, database_entities.parent_id, database_entities.name, database_entities.attributes, database_entities.created_at +FROM rules + JOIN databases ON rules.database_id = databases.id + JOIN database_entities ON rules.database_entity_id = database_entities.id +WHERE rules.database_id = $1 + AND databases.project_id = $2 ` +type GetRulesWithEntitiesParams struct { + DatabaseID uuid.UUID `json:"database_id"` + ProjectID uuid.UUID `json:"project_id"` +} + type GetRulesWithEntitiesRow struct { Rule Rule `json:"rule"` DatabaseEntity DatabaseEntity `json:"database_entity"` } -func (q *Queries) GetRulesWithEntities(ctx context.Context, databaseID uuid.UUID) ([]GetRulesWithEntitiesRow, error) { - rows, err := q.db.Query(ctx, getRulesWithEntities, databaseID) +func (q *Queries) GetRulesWithEntities(ctx context.Context, arg GetRulesWithEntitiesParams) ([]GetRulesWithEntitiesRow, error) { + rows, err := q.db.Query(ctx, getRulesWithEntities, arg.DatabaseID, arg.ProjectID) if err != nil { return nil, err } diff --git a/skemr-api/internal/controller/rule_controller.go b/skemr-api/internal/controller/rule_controller.go index 7c90062..dff02f2 100644 --- a/skemr-api/internal/controller/rule_controller.go +++ b/skemr-api/internal/controller/rule_controller.go @@ -10,6 +10,7 @@ import ( "github.com/walmaa/skemr-api/internal/dto" "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/service" + "github.com/walmaa/skemr-api/internal/validation" ) type RuleController struct { @@ -96,6 +97,14 @@ func (h *RuleController) createRule(w http.ResponseWriter, r *http.Request) { return } + err = validation.Validate.Struct(body) + + if err != nil { + errorResponse := validation.CreateErrorResponse(err) + errormsg.WriteErrorResponse(w, r, &errorResponse) + return + } + rule, err := h.Service.CreateRule(r.Context(), projectID, databaseId, body) if err != nil { diff --git a/skemr-api/internal/dto/common.go b/skemr-api/internal/dto/common.go index cf90b97..5fe360d 100644 --- a/skemr-api/internal/dto/common.go +++ b/skemr-api/internal/dto/common.go @@ -39,20 +39,11 @@ const ( type RuleCreationDto struct { Name string - RuleType RuleType + RuleType models.RuleType `json:"ruleType" validate:"required,oneof=locked deprecated advisory warning"` Attributes models.RuleAttributes `json:"attributes" validate:"omitempty,json"` - DataBaseEntityId uuid.UUID + DataBaseEntityId uuid.UUID `json:"databaseEntityId" validate:"required,uuid4"` } -type RuleType string - -const ( - RuleTypeLocked RuleType = "locked" - RuleTypeWarn RuleType = "warn" - RuleTypeAdvisory RuleType = "advisory" - RuleTypeDeprecated RuleType = "deprecated" -) - type SecretCreationDto struct { Name string `json:"name" validate:"required,min=2,max=100"` ExpiresAt string `json:"expiresAt" validate:"omitempty,datetime=2006-01-02T15:04:05Z07:00"` diff --git a/skemr-api/internal/errormsg/errors.go b/skemr-api/internal/errormsg/errors.go index caac1ab..dbcba36 100644 --- a/skemr-api/internal/errormsg/errors.go +++ b/skemr-api/internal/errormsg/errors.go @@ -29,10 +29,11 @@ func WriteErrorResponse(w http.ResponseWriter, r *http.Request, err error) { } var ( - ErrDatabaseAlreadyExists = "database already exists" - ErrDatabaseNotFound = "database not found" - ErrProjectNotFound = "project not found" - ErrInvalidIdFormat = "invalid id format" - ErrExpiryTimeInPast = "expiry time is in the past" - ErrRuleWithSameName = "rule with the same name already exists" + ErrDatabaseAlreadyExists = "database already exists" + ErrDatabaseNotFound = "database not found" + ErrProjectNotFound = "project not found" + ErrInvalidIdFormat = "invalid id format" + ErrExpiryTimeInPast = "expiry time is in the past" + ErrRuleWithSameName = "rule with the same name already exists" + ErrDatabaseEntityNotFound = "database entity not found" ) diff --git a/skemr-api/internal/mapper/mapper_util.go b/skemr-api/internal/mapper/mapper_util.go index 07e4dab..5d75535 100644 --- a/skemr-api/internal/mapper/mapper_util.go +++ b/skemr-api/internal/mapper/mapper_util.go @@ -13,7 +13,7 @@ import ( func ToBytes(v interface{}) []byte { b, err := json.Marshal(v) if err != nil { - slog.Error("Unable to marshal JSON", err) + slog.Error("Unable to marshal JSON", "err", err) return nil } return b diff --git a/skemr-api/internal/mapper/rule_mapper.go b/skemr-api/internal/mapper/rule_mapper.go index 6b83064..59cf156 100644 --- a/skemr-api/internal/mapper/rule_mapper.go +++ b/skemr-api/internal/mapper/rule_mapper.go @@ -21,6 +21,10 @@ func ToDomainRule(e sqlc.Rule) models.Rule { } func ToRuleAttributes(attributes []byte) models.RuleAttributes { + if attributes == nil { + return models.RuleAttributes{} + } + var ruleAttributes models.RuleAttributes err := json.Unmarshal(attributes, &ruleAttributes) if err != nil { diff --git a/skemr-api/internal/service/rule_service.go b/skemr-api/internal/service/rule_service.go index 342b33c..3046359 100644 --- a/skemr-api/internal/service/rule_service.go +++ b/skemr-api/internal/service/rule_service.go @@ -16,31 +16,28 @@ import ( ) type RuleService struct { - db sqlc.Querier + ruleStore RuleStore + scopeResolver ScopeResolver } -func NewRuleService(q sqlc.Querier) *RuleService { - return &RuleService{db: q} +type RuleStore interface { + GetRuleWithEntity(ctx context.Context, params sqlc.GetRuleWithEntityParams) (sqlc.GetRuleWithEntityRow, error) + GetRuleByDatabaseAndName(ctx context.Context, params sqlc.GetRuleByDatabaseAndNameParams) (sqlc.Rule, error) + CreateRule(ctx context.Context, dto sqlc.CreateRuleParams) (sqlc.Rule, error) + GetRulesWithEntities(ctx context.Context, row sqlc.GetRulesWithEntitiesParams) ([]sqlc.GetRulesWithEntitiesRow, error) + DeleteRule(ctx context.Context, params sqlc.DeleteRuleParams) error } -func (r *RuleService) GetRule(c context.Context, projectID uuid.UUID, databaseID uuid.UUID, ruleID uuid.UUID) (models.Rule, error) { - slog.Info("Fetching rule", "ruleID", ruleID) - - project, err := CheckProjectExists(c, r.db, projectID) - - if err != nil { - return models.Rule{}, err - } - - database, err := CheckDatabaseExists(c, r.db, project.ID, databaseID) +func NewRuleService(ruleStore RuleStore, resolver ScopeResolver) *RuleService { + return &RuleService{ruleStore: ruleStore, scopeResolver: resolver} +} - if err != nil { - slog.Error("Error fetching database", "err", err) - return models.Rule{}, err - } +func (r *RuleService) GetRule(c context.Context, projectID uuid.UUID, databaseID uuid.UUID, ruleID uuid.UUID) (models.Rule, error) { + slog.Info("Fetching rule", "ruleID", ruleID, "databaseID", databaseID, "projectID", projectID) - rule, err := r.db.GetRuleWithEntity(c, sqlc.GetRuleWithEntityParams{ - DatabaseID: database.ID, + rule, err := r.ruleStore.GetRuleWithEntity(c, sqlc.GetRuleWithEntityParams{ + ProjectID: projectID, + DatabaseID: databaseID, RuleID: ruleID, }) @@ -53,23 +50,24 @@ func (r *RuleService) GetRule(c context.Context, projectID uuid.UUID, databaseID } func (r *RuleService) CreateRule(c context.Context, projectID uuid.UUID, databaseId uuid.UUID, dto dto.RuleCreationDto) (models.Rule, error) { - slog.Info("Creating rule") + slog.Info("Creating rule", "name", dto.Name, "databaseID", databaseId, "projectID", projectID) - project, err := CheckProjectExists(c, r.db, projectID) + _, err := r.scopeResolver.RequireDatabase(c, projectID, databaseId) if err != nil { + slog.Error("Error fetching database", "err", err) return models.Rule{}, err } - _, err = CheckDatabaseExists(c, r.db, project.ID, databaseId) + _, err = r.scopeResolver.RequireDatabaseEntity(c, projectID, databaseId, dto.DataBaseEntityId) if err != nil { - slog.Error("Error fetching database", "err", err) + slog.Error("Error fetching database entity", "err", err) return models.Rule{}, err } // Check if a rule with the same name already exists - exists, err := r.db.GetRuleByDatabaseAndName(c, sqlc.GetRuleByDatabaseAndNameParams{ + exists, err := r.ruleStore.GetRuleByDatabaseAndName(c, sqlc.GetRuleByDatabaseAndNameParams{ DatabaseID: databaseId, Name: dto.Name, }) @@ -87,7 +85,7 @@ func (r *RuleService) CreateRule(c context.Context, projectID uuid.UUID, databas } } - rule, err := r.db.CreateRule(c, mapper.ToSqlcCreateRule(databaseId, dto)) + rule, err := r.ruleStore.CreateRule(c, mapper.ToSqlcCreateRule(databaseId, dto)) if err != nil { slog.Error("Unable to create a Rule", "err", err) return models.Rule{}, err @@ -99,20 +97,11 @@ func (r *RuleService) CreateRule(c context.Context, projectID uuid.UUID, databas func (r *RuleService) ListRulesByDatabase(c context.Context, projectID uuid.UUID, databaseID uuid.UUID) ([]models.Rule, error) { slog.Info("Listing rules", "projectID", projectID, "databaseID", databaseID) - project, err := CheckProjectExists(c, r.db, projectID) - - if err != nil { - return []models.Rule{}, err - } - - database, err := CheckDatabaseExists(c, r.db, project.ID, databaseID) - - if err != nil { - slog.Error("Error fetching database", "err", err) - return []models.Rule{}, err - } + rules, err := r.ruleStore.GetRulesWithEntities(c, sqlc.GetRulesWithEntitiesParams{ + DatabaseID: databaseID, + ProjectID: projectID, + }) - rules, err := r.db.GetRulesWithEntities(c, database.ID) if err != nil { slog.Error("Unable to get rules", "err", err) return []models.Rule{}, err @@ -121,25 +110,18 @@ func (r *RuleService) ListRulesByDatabase(c context.Context, projectID uuid.UUID } -func (r *RuleService) DeleteRule(c context.Context, projectID uuid.UUID, databaseID uuid.UUID, ruleID uuid.UUID) error { +func (r *RuleService) DeleteRule(c context.Context, projectID uuid.UUID, databaseId uuid.UUID, ruleID uuid.UUID) error { slog.Info("Deleting rule", "ruleID", ruleID) - project, err := CheckProjectExists(c, r.db, projectID) - - if err != nil { - slog.Error("Error fetching project", "err", err) - return err - } - - database, err := CheckDatabaseExists(c, r.db, project.ID, databaseID) + _, err := r.scopeResolver.RequireDatabase(c, projectID, databaseId) if err != nil { slog.Error("Error fetching database", "err", err) return err } - err = r.db.DeleteRule(c, sqlc.DeleteRuleParams{ - DatabaseID: database.ID, + err = r.ruleStore.DeleteRule(c, sqlc.DeleteRuleParams{ + DatabaseID: databaseId, RuleID: ruleID, }) if err != nil { diff --git a/skemr-api/internal/service/rule_service_test.go b/skemr-api/internal/service/rule_service_test.go new file mode 100644 index 0000000..7ee755a --- /dev/null +++ b/skemr-api/internal/service/rule_service_test.go @@ -0,0 +1,281 @@ +package service + +import ( + "net/http" + "testing" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + "github.com/walmaa/skemr-api/db/sqlc" + "github.com/walmaa/skemr-api/internal/dto" + "github.com/walmaa/skemr-api/internal/errormsg" + "github.com/walmaa/skemr-api/test/mocks" + "github.com/walmaa/skemr-common/models" +) + +func TestCreateLockedRule(t *testing.T) { + projectId := uuid.New() + databaseId := uuid.New() + scopeResolver := mocks.NewMockScopeResolver(t) + ruleStore := mocks.NewMockRuleStore(t) + ruleType := models.RuleTypeLocked + + svc := NewRuleService(ruleStore, scopeResolver) + + input := dto.RuleCreationDto{ + Name: "test", + RuleType: ruleType, + Attributes: models.RuleAttributes{}, + DataBaseEntityId: uuid.UUID{}, + } + + scopeResolver.On("RequireDatabase", mock.Anything, projectId, databaseId).Return(models.Database{}, nil) + + scopeResolver.On("RequireDatabaseEntity", mock.Anything, projectId, databaseId, input.DataBaseEntityId).Return(models.DatabaseEntity{}, nil) + ruleStore.On("GetRuleByDatabaseAndName", mock.Anything, mock.Anything).Return(sqlc.Rule{}, pgx.ErrNoRows) + ruleStore.On("CreateRule", mock.Anything, mock.Anything).Return(sqlc.Rule{ + ID: uuid.New(), + DatabaseID: databaseId, + Name: "test", + Attributes: nil, + Type: sqlc.RuleTypeLocked, + }, nil) + + rule, err := svc.CreateRule(t.Context(), projectId, databaseId, input) + + require.NoError(t, err) + require.Equal(t, ruleType, rule.RuleType) + ruleStore.AssertCalled(t, "CreateRule", mock.Anything, mock.Anything) +} + +func TestCreateDeprecatedRule(t *testing.T) { + projectId := uuid.New() + databaseId := uuid.New() + scopeResolver := mocks.NewMockScopeResolver(t) + ruleStore := mocks.NewMockRuleStore(t) + ruleType := models.RuleTypeDeprecated + + svc := NewRuleService(ruleStore, scopeResolver) + + input := dto.RuleCreationDto{ + Name: "test", + RuleType: ruleType, + Attributes: models.RuleAttributes{}, + DataBaseEntityId: uuid.UUID{}, + } + + scopeResolver.On("RequireDatabase", mock.Anything, projectId, databaseId).Return(models.Database{}, nil) + + scopeResolver.On("RequireDatabaseEntity", mock.Anything, projectId, databaseId, input.DataBaseEntityId).Return(models.DatabaseEntity{}, nil) + ruleStore.On("GetRuleByDatabaseAndName", mock.Anything, mock.Anything).Return(sqlc.Rule{}, pgx.ErrNoRows) + ruleStore.On("CreateRule", mock.Anything, mock.Anything).Return(sqlc.Rule{ + ID: uuid.New(), + DatabaseID: databaseId, + Name: "test", + Attributes: nil, + Type: sqlc.RuleTypeDeprecated, + }, nil) + + rule, err := svc.CreateRule(t.Context(), projectId, databaseId, input) + + require.NoError(t, err) + require.Equal(t, ruleType, rule.RuleType) + ruleStore.AssertCalled(t, "CreateRule", mock.Anything, mock.Anything) +} + +func TestCreateAdvisoryRule(t *testing.T) { + projectId := uuid.New() + databaseId := uuid.New() + scopeResolver := mocks.NewMockScopeResolver(t) + ruleStore := mocks.NewMockRuleStore(t) + ruleType := models.RuleTypeAdvisory + + svc := NewRuleService(ruleStore, scopeResolver) + + input := dto.RuleCreationDto{ + Name: "test", + RuleType: ruleType, + Attributes: models.RuleAttributes{}, + DataBaseEntityId: uuid.UUID{}, + } + + scopeResolver.On("RequireDatabase", mock.Anything, projectId, databaseId).Return(models.Database{}, nil) + + scopeResolver.On("RequireDatabaseEntity", mock.Anything, projectId, databaseId, input.DataBaseEntityId).Return(models.DatabaseEntity{}, nil) + ruleStore.On("GetRuleByDatabaseAndName", mock.Anything, mock.Anything).Return(sqlc.Rule{}, pgx.ErrNoRows) + ruleStore.On("CreateRule", mock.Anything, mock.Anything).Return(sqlc.Rule{ + ID: uuid.New(), + DatabaseID: databaseId, + Name: "test", + Attributes: nil, + Type: sqlc.RuleTypeAdvisory, + }, nil) + + rule, err := svc.CreateRule(t.Context(), projectId, databaseId, input) + + require.NoError(t, err) + require.Equal(t, ruleType, rule.RuleType) + ruleStore.AssertCalled(t, "CreateRule", mock.Anything, mock.Anything) +} + +func TestCreateWarningRule(t *testing.T) { + projectId := uuid.New() + databaseId := uuid.New() + scopeResolver := mocks.NewMockScopeResolver(t) + ruleStore := mocks.NewMockRuleStore(t) + ruleType := models.RuleTypeWarn + + svc := NewRuleService(ruleStore, scopeResolver) + + input := dto.RuleCreationDto{ + Name: "test", + RuleType: ruleType, + Attributes: models.RuleAttributes{}, + DataBaseEntityId: uuid.UUID{}, + } + + scopeResolver.On("RequireDatabase", mock.Anything, projectId, databaseId).Return(models.Database{}, nil) + + scopeResolver.On("RequireDatabaseEntity", mock.Anything, projectId, databaseId, input.DataBaseEntityId).Return(models.DatabaseEntity{}, nil) + ruleStore.On("GetRuleByDatabaseAndName", mock.Anything, mock.Anything).Return(sqlc.Rule{}, pgx.ErrNoRows) + ruleStore.On("CreateRule", mock.Anything, mock.Anything).Return(sqlc.Rule{ + ID: uuid.New(), + DatabaseID: databaseId, + Name: "test", + Attributes: nil, + Type: sqlc.RuleTypeWarn, + }, nil) + + rule, err := svc.CreateRule(t.Context(), projectId, databaseId, input) + + require.NoError(t, err) + require.Equal(t, ruleType, rule.RuleType) + ruleStore.AssertCalled(t, "CreateRule", mock.Anything, mock.Anything) +} + +func TestCreateRuleWithoutDatabaseEntity(t *testing.T) { + projectId := uuid.New() + databaseId := uuid.New() + scopeResolver := mocks.NewMockScopeResolver(t) + ruleStore := mocks.NewMockRuleStore(t) + + svc := NewRuleService(ruleStore, scopeResolver) + + input := dto.RuleCreationDto{ + Name: "test", + RuleType: models.RuleTypeLocked, + Attributes: models.RuleAttributes{}, + DataBaseEntityId: uuid.UUID{}, + } + + scopeResolver.On("RequireDatabase", mock.Anything, projectId, databaseId).Return(models.Database{}, nil) + + scopeResolver.On("RequireDatabaseEntity", mock.Anything, projectId, databaseId, input.DataBaseEntityId).Return(models.DatabaseEntity{}, &models.ErrorResponse{ + Message: errormsg.ErrDatabaseEntityNotFound, + Errors: nil, + Status: http.StatusBadRequest, + }) + + _, err := svc.CreateRule(t.Context(), projectId, databaseId, input) + require.Error(t, err) + require.Equal(t, errormsg.ErrDatabaseEntityNotFound, err.Error()) + + ruleStore.AssertNotCalled(t, "CreateRule", mock.Anything, mock.Anything) +} + +func TestCreateRuleWithNonExistentDatabaseEntity(t *testing.T) { + projectId := uuid.New() + databaseId := uuid.New() + scopeResolver := mocks.NewMockScopeResolver(t) + ruleStore := mocks.NewMockRuleStore(t) + + svc := NewRuleService(ruleStore, scopeResolver) + + input := dto.RuleCreationDto{ + Name: "test", + RuleType: models.RuleTypeLocked, + Attributes: models.RuleAttributes{}, + DataBaseEntityId: uuid.New(), + } + + scopeResolver.On("RequireDatabase", mock.Anything, projectId, databaseId).Return(models.Database{}, nil) + + scopeResolver.On("RequireDatabaseEntity", mock.Anything, projectId, databaseId, input.DataBaseEntityId).Return(models.DatabaseEntity{}, &models.ErrorResponse{ + Message: errormsg.ErrDatabaseEntityNotFound, + Errors: nil, + Status: http.StatusBadRequest, + }) + + _, err := svc.CreateRule(t.Context(), projectId, databaseId, input) + require.Error(t, err) + require.Equal(t, errormsg.ErrDatabaseEntityNotFound, err.Error()) + + ruleStore.AssertNotCalled(t, "CreateRule", mock.Anything, mock.Anything) + +} + +func TestCreateRuleWithNonExistentProjectOrDatabase(t *testing.T) { + projectId := uuid.New() + databaseId := uuid.New() + scopeResolver := mocks.NewMockScopeResolver(t) + ruleStore := mocks.NewMockRuleStore(t) + + svc := NewRuleService(ruleStore, scopeResolver) + + input := dto.RuleCreationDto{ + Name: "test", + RuleType: models.RuleTypeLocked, + Attributes: models.RuleAttributes{}, + DataBaseEntityId: uuid.New(), + } + + scopeResolver.On("RequireDatabase", mock.Anything, projectId, databaseId).Return(models.Database{}, &models.ErrorResponse{ + Message: errormsg.ErrDatabaseNotFound, + Errors: nil, + Status: http.StatusBadRequest, + }) + + _, err := svc.CreateRule(t.Context(), projectId, databaseId, input) + require.Error(t, err) + require.Equal(t, errormsg.ErrDatabaseNotFound, err.Error()) + + ruleStore.AssertNotCalled(t, "CreateRule", mock.Anything, mock.Anything) + ruleStore.AssertNotCalled(t, "RequireDatabaseEntity", mock.Anything, projectId, databaseId, input.DataBaseEntityId) +} + +func TestCreateRuleWithExistingName(t *testing.T) { + projectId := uuid.New() + databaseId := uuid.New() + scopeResolver := mocks.NewMockScopeResolver(t) + ruleStore := mocks.NewMockRuleStore(t) + + svc := NewRuleService(ruleStore, scopeResolver) + + input := dto.RuleCreationDto{ + Name: "test", + RuleType: models.RuleTypeLocked, + Attributes: models.RuleAttributes{}, + DataBaseEntityId: uuid.New(), + } + + scopeResolver.On("RequireDatabase", mock.Anything, projectId, databaseId).Return(models.Database{}, nil) + + scopeResolver.On("RequireDatabaseEntity", mock.Anything, projectId, databaseId, input.DataBaseEntityId).Return(models.DatabaseEntity{}, nil) + + ruleStore.On("GetRuleByDatabaseAndName", mock.Anything, sqlc.GetRuleByDatabaseAndNameParams{ + DatabaseID: databaseId, + Name: "test", + }).Return(sqlc.Rule{ + ID: uuid.New(), + DatabaseID: databaseId, + Name: "test", + }, nil) + + _, err := svc.CreateRule(t.Context(), projectId, databaseId, input) + require.Error(t, err) + require.Equal(t, errormsg.ErrRuleWithSameName, err.Error()) + + ruleStore.AssertNotCalled(t, "CreateRule", mock.Anything, mock.Anything) +} diff --git a/skemr-api/internal/service/scope_resolver.go b/skemr-api/internal/service/scope_resolver.go new file mode 100644 index 0000000..1861f00 --- /dev/null +++ b/skemr-api/internal/service/scope_resolver.go @@ -0,0 +1,78 @@ +package service + +import ( + "context" + "errors" + "log/slog" + "net/http" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "github.com/walmaa/skemr-api/db/sqlc" + "github.com/walmaa/skemr-api/internal/errormsg" + "github.com/walmaa/skemr-api/internal/mapper" + "github.com/walmaa/skemr-common/models" +) + +type ScopeResolver interface { + RequireDatabase(c context.Context, projectId uuid.UUID, databaseId uuid.UUID) (models.Database, error) + RequireDatabaseEntity(c context.Context, projectId uuid.UUID, databaseId uuid.UUID, entityId uuid.UUID) (models.DatabaseEntity, error) +} +type SqlcScopeResolver struct { + db sqlc.Querier +} + +func NewScopeResolver(q sqlc.Querier) *SqlcScopeResolver { + return &SqlcScopeResolver{db: q} +} + +func (s *SqlcScopeResolver) RequireDatabase(c context.Context, projectId uuid.UUID, databaseId uuid.UUID) (models.Database, error) { + slog.Info("Getting database", "databaseId", databaseId, "projectId", projectId) + + database, err := s.db.GetDatabaseByIDAndProjectID(c, sqlc.GetDatabaseByIDAndProjectIDParams{ + ID: databaseId, + ProjectID: projectId, + }) + + if errors.Is(err, pgx.ErrNoRows) { + slog.Info("Database not found", "databaseId", databaseId, "projectId", projectId) + return models.Database{}, &models.ErrorResponse{ + Message: errormsg.ErrDatabaseNotFound, + Errors: nil, + Status: http.StatusBadRequest, + } + } + + if err != nil { + slog.Error("Unable to get database", "databaseId", databaseId, "projectId", projectId, "err", err) + return models.Database{}, err + } + + return mapper.ToDomainDatabase(database), nil +} + +func (s *SqlcScopeResolver) RequireDatabaseEntity(c context.Context, projectId uuid.UUID, databaseId uuid.UUID, entityId uuid.UUID) (models.DatabaseEntity, error) { + slog.Info("Getting database entity", "entityId", entityId, "databaseId", databaseId, "projectId", projectId) + + entity, err := s.db.GetDatabaseEntityByProjectIdDatabaseIdAndId(c, sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams{ + ID: entityId, + ProjectID: projectId, + DatabaseID: databaseId, + }) + + if errors.Is(err, pgx.ErrNoRows) { + slog.Info("Database entity not found", "entityId", entityId, "databaseId", databaseId, "projectId", projectId) + return models.DatabaseEntity{}, &models.ErrorResponse{ + Message: errormsg.ErrDatabaseEntityNotFound, + Errors: nil, + Status: http.StatusBadRequest, + } + } + + if err != nil { + slog.Error("Unable to get database entity", "entityId", entityId, "databaseId", databaseId, "projectId", projectId, "err", err) + return models.DatabaseEntity{}, err + } + + return mapper.ToDomainDatabaseEntity(entity), nil +} diff --git a/skemr-api/internal/validation/validation.go b/skemr-api/internal/validation/validation.go index b9e510e..b0f3fed 100644 --- a/skemr-api/internal/validation/validation.go +++ b/skemr-api/internal/validation/validation.go @@ -1,6 +1,7 @@ package validation import ( + "net/http" "strings" "github.com/go-playground/validator/v10" @@ -18,6 +19,7 @@ func CreateErrorResponse(err error) models.ErrorResponse { errorResponse := models.ErrorResponse{ Message: "Validation failed", Errors: make(map[string]string), + Status: http.StatusBadRequest, } for _, fieldErr := range validationErrors { diff --git a/skemr-api/test/mocks/querier_mock.go b/skemr-api/test/mocks/querier_mock.go index 121d233..cfd91c0 100644 --- a/skemr-api/test/mocks/querier_mock.go +++ b/skemr-api/test/mocks/querier_mock.go @@ -1397,6 +1397,72 @@ func (_c *MockQuerier_GetDatabaseEntityByProjectIdAndId_Call) RunAndReturn(run f return _c } +// GetDatabaseEntityByProjectIdDatabaseIdAndId provides a mock function for the type MockQuerier +func (_mock *MockQuerier) GetDatabaseEntityByProjectIdDatabaseIdAndId(ctx context.Context, arg sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams) (sqlc.DatabaseEntity, error) { + ret := _mock.Called(ctx, arg) + + if len(ret) == 0 { + panic("no return value specified for GetDatabaseEntityByProjectIdDatabaseIdAndId") + } + + var r0 sqlc.DatabaseEntity + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams) (sqlc.DatabaseEntity, error)); ok { + return returnFunc(ctx, arg) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams) sqlc.DatabaseEntity); ok { + r0 = returnFunc(ctx, arg) + } else { + r0 = ret.Get(0).(sqlc.DatabaseEntity) + } + if returnFunc, ok := ret.Get(1).(func(context.Context, sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams) error); ok { + r1 = returnFunc(ctx, arg) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'GetDatabaseEntityByProjectIdDatabaseIdAndId' +type MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call struct { + *mock.Call +} + +// GetDatabaseEntityByProjectIdDatabaseIdAndId is a helper method to define mock.On call +// - ctx context.Context +// - arg sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams +func (_e *MockQuerier_Expecter) GetDatabaseEntityByProjectIdDatabaseIdAndId(ctx interface{}, arg interface{}) *MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call { + return &MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call{Call: _e.mock.On("GetDatabaseEntityByProjectIdDatabaseIdAndId", ctx, arg)} +} + +func (_c *MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call) Run(run func(ctx context.Context, arg sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams)) *MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams + if args[1] != nil { + arg1 = args[1].(sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams) + } + run( + arg0, + arg1, + ) + }) + return _c +} + +func (_c *MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call) Return(databaseEntity sqlc.DatabaseEntity, err error) *MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call { + _c.Call.Return(databaseEntity, err) + return _c +} + +func (_c *MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call) RunAndReturn(run func(ctx context.Context, arg sqlc.GetDatabaseEntityByProjectIdDatabaseIdAndIdParams) (sqlc.DatabaseEntity, error)) *MockQuerier_GetDatabaseEntityByProjectIdDatabaseIdAndId_Call { + _c.Call.Return(run) + return _c +} + // GetHashByPrefixAndProjectID provides a mock function for the type MockQuerier func (_mock *MockQuerier) GetHashByPrefixAndProjectID(ctx context.Context, arg sqlc.GetHashByPrefixAndProjectIDParams) (string, error) { ret := _mock.Called(ctx, arg) @@ -1857,6 +1923,72 @@ func (_c *MockQuerier_GetRule_Call) RunAndReturn(run func(ctx context.Context, a return _c } +// GetRuleByDatabaseAndName provides a mock function for the type MockQuerier +func (_mock *MockQuerier) GetRuleByDatabaseAndName(ctx context.Context, arg sqlc.GetRuleByDatabaseAndNameParams) (sqlc.Rule, error) { + ret := _mock.Called(ctx, arg) + + if len(ret) == 0 { + panic("no return value specified for GetRuleByDatabaseAndName") + } + + var r0 sqlc.Rule + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, sqlc.GetRuleByDatabaseAndNameParams) (sqlc.Rule, error)); ok { + return returnFunc(ctx, arg) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, sqlc.GetRuleByDatabaseAndNameParams) sqlc.Rule); ok { + r0 = returnFunc(ctx, arg) + } else { + r0 = ret.Get(0).(sqlc.Rule) + } + if returnFunc, ok := ret.Get(1).(func(context.Context, sqlc.GetRuleByDatabaseAndNameParams) error); ok { + r1 = returnFunc(ctx, arg) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockQuerier_GetRuleByDatabaseAndName_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'GetRuleByDatabaseAndName' +type MockQuerier_GetRuleByDatabaseAndName_Call struct { + *mock.Call +} + +// GetRuleByDatabaseAndName is a helper method to define mock.On call +// - ctx context.Context +// - arg sqlc.GetRuleByDatabaseAndNameParams +func (_e *MockQuerier_Expecter) GetRuleByDatabaseAndName(ctx interface{}, arg interface{}) *MockQuerier_GetRuleByDatabaseAndName_Call { + return &MockQuerier_GetRuleByDatabaseAndName_Call{Call: _e.mock.On("GetRuleByDatabaseAndName", ctx, arg)} +} + +func (_c *MockQuerier_GetRuleByDatabaseAndName_Call) Run(run func(ctx context.Context, arg sqlc.GetRuleByDatabaseAndNameParams)) *MockQuerier_GetRuleByDatabaseAndName_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 sqlc.GetRuleByDatabaseAndNameParams + if args[1] != nil { + arg1 = args[1].(sqlc.GetRuleByDatabaseAndNameParams) + } + run( + arg0, + arg1, + ) + }) + return _c +} + +func (_c *MockQuerier_GetRuleByDatabaseAndName_Call) Return(rule sqlc.Rule, err error) *MockQuerier_GetRuleByDatabaseAndName_Call { + _c.Call.Return(rule, err) + return _c +} + +func (_c *MockQuerier_GetRuleByDatabaseAndName_Call) RunAndReturn(run func(ctx context.Context, arg sqlc.GetRuleByDatabaseAndNameParams) (sqlc.Rule, error)) *MockQuerier_GetRuleByDatabaseAndName_Call { + _c.Call.Return(run) + return _c +} + // GetRuleWithEntity provides a mock function for the type MockQuerier func (_mock *MockQuerier) GetRuleWithEntity(ctx context.Context, arg sqlc.GetRuleWithEntityParams) (sqlc.GetRuleWithEntityRow, error) { ret := _mock.Called(ctx, arg) @@ -1924,8 +2056,8 @@ func (_c *MockQuerier_GetRuleWithEntity_Call) RunAndReturn(run func(ctx context. } // GetRulesWithEntities provides a mock function for the type MockQuerier -func (_mock *MockQuerier) GetRulesWithEntities(ctx context.Context, databaseID uuid.UUID) ([]sqlc.GetRulesWithEntitiesRow, error) { - ret := _mock.Called(ctx, databaseID) +func (_mock *MockQuerier) GetRulesWithEntities(ctx context.Context, arg sqlc.GetRulesWithEntitiesParams) ([]sqlc.GetRulesWithEntitiesRow, error) { + ret := _mock.Called(ctx, arg) if len(ret) == 0 { panic("no return value specified for GetRulesWithEntities") @@ -1933,18 +2065,18 @@ func (_mock *MockQuerier) GetRulesWithEntities(ctx context.Context, databaseID u var r0 []sqlc.GetRulesWithEntitiesRow var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, uuid.UUID) ([]sqlc.GetRulesWithEntitiesRow, error)); ok { - return returnFunc(ctx, databaseID) + if returnFunc, ok := ret.Get(0).(func(context.Context, sqlc.GetRulesWithEntitiesParams) ([]sqlc.GetRulesWithEntitiesRow, error)); ok { + return returnFunc(ctx, arg) } - if returnFunc, ok := ret.Get(0).(func(context.Context, uuid.UUID) []sqlc.GetRulesWithEntitiesRow); ok { - r0 = returnFunc(ctx, databaseID) + if returnFunc, ok := ret.Get(0).(func(context.Context, sqlc.GetRulesWithEntitiesParams) []sqlc.GetRulesWithEntitiesRow); ok { + r0 = returnFunc(ctx, arg) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]sqlc.GetRulesWithEntitiesRow) } } - if returnFunc, ok := ret.Get(1).(func(context.Context, uuid.UUID) error); ok { - r1 = returnFunc(ctx, databaseID) + if returnFunc, ok := ret.Get(1).(func(context.Context, sqlc.GetRulesWithEntitiesParams) error); ok { + r1 = returnFunc(ctx, arg) } else { r1 = ret.Error(1) } @@ -1958,20 +2090,20 @@ type MockQuerier_GetRulesWithEntities_Call struct { // GetRulesWithEntities is a helper method to define mock.On call // - ctx context.Context -// - databaseID uuid.UUID -func (_e *MockQuerier_Expecter) GetRulesWithEntities(ctx interface{}, databaseID interface{}) *MockQuerier_GetRulesWithEntities_Call { - return &MockQuerier_GetRulesWithEntities_Call{Call: _e.mock.On("GetRulesWithEntities", ctx, databaseID)} +// - arg sqlc.GetRulesWithEntitiesParams +func (_e *MockQuerier_Expecter) GetRulesWithEntities(ctx interface{}, arg interface{}) *MockQuerier_GetRulesWithEntities_Call { + return &MockQuerier_GetRulesWithEntities_Call{Call: _e.mock.On("GetRulesWithEntities", ctx, arg)} } -func (_c *MockQuerier_GetRulesWithEntities_Call) Run(run func(ctx context.Context, databaseID uuid.UUID)) *MockQuerier_GetRulesWithEntities_Call { +func (_c *MockQuerier_GetRulesWithEntities_Call) Run(run func(ctx context.Context, arg sqlc.GetRulesWithEntitiesParams)) *MockQuerier_GetRulesWithEntities_Call { _c.Call.Run(func(args mock.Arguments) { var arg0 context.Context if args[0] != nil { arg0 = args[0].(context.Context) } - var arg1 uuid.UUID + var arg1 sqlc.GetRulesWithEntitiesParams if args[1] != nil { - arg1 = args[1].(uuid.UUID) + arg1 = args[1].(sqlc.GetRulesWithEntitiesParams) } run( arg0, @@ -1986,7 +2118,7 @@ func (_c *MockQuerier_GetRulesWithEntities_Call) Return(getRulesWithEntitiesRows return _c } -func (_c *MockQuerier_GetRulesWithEntities_Call) RunAndReturn(run func(ctx context.Context, databaseID uuid.UUID) ([]sqlc.GetRulesWithEntitiesRow, error)) *MockQuerier_GetRulesWithEntities_Call { +func (_c *MockQuerier_GetRulesWithEntities_Call) RunAndReturn(run func(ctx context.Context, arg sqlc.GetRulesWithEntitiesParams) ([]sqlc.GetRulesWithEntitiesRow, error)) *MockQuerier_GetRulesWithEntities_Call { _c.Call.Return(run) return _c } From ec95c2adf453f748de7333c09ba8599b9eb1f26c Mon Sep 17 00:00:00 2001 From: WalMaa Date: Thu, 21 May 2026 10:41:42 +0300 Subject: [PATCH 08/14] build and test gha on any branch --- .github/workflows/build-and-test.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/build-and-test.yml b/.github/workflows/build-and-test.yml index ca76e1f..b631f67 100644 --- a/.github/workflows/build-and-test.yml +++ b/.github/workflows/build-and-test.yml @@ -2,7 +2,7 @@ name: build-and-test.yml on: push: branches: - - main + - * pull_request: branches: - main From 00ad1ef7f72b822e5d91e9b5959fd51f7f9736d5 Mon Sep 17 00:00:00 2001 From: WalMaa Date: Thu, 21 May 2026 10:43:18 +0300 Subject: [PATCH 09/14] build and test gha on any branch --- .github/workflows/build-and-test.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/build-and-test.yml b/.github/workflows/build-and-test.yml index b631f67..89283a8 100644 --- a/.github/workflows/build-and-test.yml +++ b/.github/workflows/build-and-test.yml @@ -2,7 +2,7 @@ name: build-and-test.yml on: push: branches: - - * + - '**' pull_request: branches: - main From 868b26462b137408e7fc762f10f3a0f0ac30d21b Mon Sep 17 00:00:00 2001 From: WalMaa Date: Sun, 24 May 2026 12:07:12 +0300 Subject: [PATCH 10/14] database change implementation to control plane --- skemr-api/cmd/server/main.go | 2 + .../migrations/20260225122322_init_schema.sql | 20 ++-- skemr-api/db/queries/database_changes.sql | 19 ++++ skemr-api/db/sqlc/models.go | 17 ++-- skemr-api/db/sqlc/querier.go | 3 + .../controller/database_change_controller.go | 74 +++++++++++++++ .../internal/controller/rule_controller.go | 2 +- skemr-api/internal/dbreflect/schema_sync.go | 59 ++++++++---- skemr-api/internal/errormsg/errors.go | 20 ++-- .../internal/mapper/database_change_mapper.go | 24 +++++ skemr-api/internal/routers/router.go | 3 + .../service/database_change_service.go | 94 +++++++++++++++++++ skemr-common/models/database_change.go | 15 +++ 13 files changed, 309 insertions(+), 43 deletions(-) create mode 100644 skemr-api/db/queries/database_changes.sql create mode 100644 skemr-api/internal/controller/database_change_controller.go create mode 100644 skemr-api/internal/mapper/database_change_mapper.go create mode 100644 skemr-api/internal/service/database_change_service.go create mode 100644 skemr-common/models/database_change.go diff --git a/skemr-api/cmd/server/main.go b/skemr-api/cmd/server/main.go index 4511346..d8e1245 100644 --- a/skemr-api/cmd/server/main.go +++ b/skemr-api/cmd/server/main.go @@ -95,6 +95,7 @@ func main() { queries := sqlc.New(conn) scopeResolver := service.NewScopeResolver(queries) projectService := service.NewProjectService(queries) + databaseChangeService := service.NewDatabaseChangeService(queries, scopeResolver) databaseService := service.NewDatabaseService(queries, taskClient) webhookService := service.NewWebhookService(queries) projectSecretsService := service.NewAccessTokenService(queries) @@ -118,6 +119,7 @@ func main() { RuleService: ruleService, DatabaseEntityService: databaseEntityService, IntegrationService: integrationService, + DatabaseChangeService: databaseChangeService, } // Initialize router diff --git a/skemr-api/db/migrations/20260225122322_init_schema.sql b/skemr-api/db/migrations/20260225122322_init_schema.sql index fed3ea6..df936a7 100644 --- a/skemr-api/db/migrations/20260225122322_init_schema.sql +++ b/skemr-api/db/migrations/20260225122322_init_schema.sql @@ -100,17 +100,6 @@ CREATE TABLE databases CONSTRAINT databases_id_project_id_unique UNIQUE (id, project_id) ); -CREATE TABLE migration_statements -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), - raw_statement TEXT NOT NULL, - action migration_statement_action NOT NULL, - status migration_status NOT NULL DEFAULT 'pending', - target TEXT, - relation_name TEXT -); - - CREATE TABLE tables ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -152,6 +141,15 @@ CREATE TABLE database_entities ); +CREATE TABLE database_changes +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + database_id UUID NOT NULL REFERENCES databases (id) ON DELETE CASCADE, + entity_id UUID NOT NULL REFERENCES database_entities (id) ON DELETE CASCADE, + action migration_statement_action NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + -- Rules specify the protection mechanisms for databases, schemas, tables, and columns. CREATE TABLE rules diff --git a/skemr-api/db/queries/database_changes.sql b/skemr-api/db/queries/database_changes.sql new file mode 100644 index 0000000..25dbf98 --- /dev/null +++ b/skemr-api/db/queries/database_changes.sql @@ -0,0 +1,19 @@ +-- name: CreateDatabaseChange :one +INSERT INTO database_changes + (database_id, entity_id, action) +VALUES (@database_id, @entity_id, @action) +RETURNING *; + +-- name: GetDatabaseChangeByDatabaseIdAndId :one +SELECT * +FROM database_changes c +WHERE c.id = @id + AND c.database_id = @database_id +LIMIT 1; + +-- name: GetDatabaseChangesByDatabaseIdAndId :many +SELECT * +FROM database_changes c +WHERE c.database_id = @database_id +ORDER BY c.created_at DESC +LIMIT sqlc.narg('limit')::int OFFSET sqlc.narg('offset')::int; diff --git a/skemr-api/db/sqlc/models.go b/skemr-api/db/sqlc/models.go index c978834..6cdf315 100644 --- a/skemr-api/db/sqlc/models.go +++ b/skemr-api/db/sqlc/models.go @@ -292,6 +292,14 @@ type Database struct { UpdatedAt pgtype.Timestamptz `json:"updated_at"` } +type DatabaseChange struct { + ID uuid.UUID `json:"id"` + DatabaseID uuid.UUID `json:"database_id"` + EntityID uuid.UUID `json:"entity_id"` + Action MigrationStatementAction `json:"action"` + CreatedAt pgtype.Timestamptz `json:"created_at"` +} + type DatabaseEntity struct { ID uuid.UUID `json:"id"` Fingerprint string `json:"fingerprint"` @@ -307,15 +315,6 @@ type DatabaseEntity struct { CreatedAt pgtype.Timestamptz `json:"created_at"` } -type MigrationStatement struct { - ID uuid.UUID `json:"id"` - RawStatement string `json:"raw_statement"` - Action MigrationStatementAction `json:"action"` - Status MigrationStatus `json:"status"` - Target pgtype.Text `json:"target"` - RelationName pgtype.Text `json:"relation_name"` -} - type Project struct { ID uuid.UUID `json:"id"` Name string `json:"name"` diff --git a/skemr-api/db/sqlc/querier.go b/skemr-api/db/sqlc/querier.go index 5232f56..e7c4c06 100644 --- a/skemr-api/db/sqlc/querier.go +++ b/skemr-api/db/sqlc/querier.go @@ -12,6 +12,7 @@ import ( type Querier interface { CreateDatabase(ctx context.Context, arg CreateDatabaseParams) (Database, error) + CreateDatabaseChange(ctx context.Context, arg CreateDatabaseChangeParams) (DatabaseChange, error) CreateDatabaseEntity(ctx context.Context, arg CreateDatabaseEntityParams) (DatabaseEntity, error) CreateProject(ctx context.Context, name string) (Project, error) CreateProjectSecretKey(ctx context.Context, arg CreateProjectSecretKeyParams) (CreateProjectSecretKeyRow, error) @@ -24,6 +25,8 @@ type Querier interface { GetDatabaseByIDAndProjectID(ctx context.Context, arg GetDatabaseByIDAndProjectIDParams) (Database, error) GetDatabaseByIdAndProject(ctx context.Context, arg GetDatabaseByIdAndProjectParams) (Database, error) GetDatabaseByNameAndProject(ctx context.Context, arg GetDatabaseByNameAndProjectParams) (Database, error) + GetDatabaseChangeByDatabaseIdAndId(ctx context.Context, arg GetDatabaseChangeByDatabaseIdAndIdParams) (DatabaseChange, error) + GetDatabaseChangesByDatabaseIdAndId(ctx context.Context, arg GetDatabaseChangesByDatabaseIdAndIdParams) ([]DatabaseChange, error) GetDatabaseEntities(ctx context.Context, arg GetDatabaseEntitiesParams) ([]DatabaseEntity, error) GetDatabaseEntitiesByDatabaseId(ctx context.Context, databaseID uuid.UUID) ([]DatabaseEntity, error) GetDatabaseEntitiesByDatabaseIdAndParentId(ctx context.Context, arg GetDatabaseEntitiesByDatabaseIdAndParentIdParams) ([]DatabaseEntity, error) diff --git a/skemr-api/internal/controller/database_change_controller.go b/skemr-api/internal/controller/database_change_controller.go new file mode 100644 index 0000000..06181ca --- /dev/null +++ b/skemr-api/internal/controller/database_change_controller.go @@ -0,0 +1,74 @@ +package controller + +import ( + "net/http" + + "github.com/go-chi/chi/v5" + "github.com/go-chi/render" + "github.com/google/uuid" + "github.com/walmaa/skemr-api/internal/errormsg" + "github.com/walmaa/skemr-api/internal/service" +) + +type DatabaseChangeController struct { + Service *service.DatabaseChangeService +} + +func NewDatabaseChangeController(s *service.DatabaseChangeService) *DatabaseChangeController { + return &DatabaseChangeController{Service: s} +} + +func (h *DatabaseChangeController) RegisterRoutes(r chi.Router) { + r.Route("/databases/{databaseId}/changes", func(r chi.Router) { + + r.Get("/", h.listDatabaseChanges) + r.Get("/{databaseChangeId}", h.GetDatabaseChange) + }) +} + +func (h *DatabaseChangeController) listDatabaseChanges(w http.ResponseWriter, r *http.Request) { + projectID, ok := r.Context().Value("projectId").(uuid.UUID) + if !ok { + http.Error(w, "projectId not found in context", http.StatusBadRequest) + return + } + databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) + if err != nil { + http.Error(w, "invalid databaseId", http.StatusBadRequest) + return + } + + databaseChanges, err := h.Service.GetDatabaseChanges(r.Context(), projectID, databaseId, 100, 0) + + if err != nil { + errormsg.WriteErrorResponse(w, r, err) + return + } + render.JSON(w, r, databaseChanges) +} + +func (h *DatabaseChangeController) GetDatabaseChange(w http.ResponseWriter, r *http.Request) { + projectID, ok := r.Context().Value("projectID").(uuid.UUID) + if !ok { + http.Error(w, "projectId not found in context", http.StatusBadRequest) + return + } + databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) + if err != nil { + http.Error(w, "invalid databaseId", http.StatusBadRequest) + return + } + + databaseChangeId, err := uuid.Parse(chi.URLParam(r, "databaseChangeId")) + if err != nil { + http.Error(w, "invalid databaseChangeId", http.StatusBadRequest) + } + + databaseChange, err := h.Service.GetDatabaseChange(r.Context(), projectID, databaseId, databaseChangeId) + + if err != nil { + errormsg.WriteErrorResponse(w, r, err) + return + } + render.JSON(w, r, databaseChange) +} diff --git a/skemr-api/internal/controller/rule_controller.go b/skemr-api/internal/controller/rule_controller.go index dff02f2..e7ce54c 100644 --- a/skemr-api/internal/controller/rule_controller.go +++ b/skemr-api/internal/controller/rule_controller.go @@ -31,7 +31,7 @@ func (h *RuleController) RegisterRoutes(r chi.Router) { } func (h *RuleController) GetRule(w http.ResponseWriter, r *http.Request) { - projectID, ok := r.Context().Value("projectID").(uuid.UUID) + projectID, ok := r.Context().Value("projectId").(uuid.UUID) if !ok { http.Error(w, "projectId not found in context", http.StatusBadRequest) return diff --git a/skemr-api/internal/dbreflect/schema_sync.go b/skemr-api/internal/dbreflect/schema_sync.go index bbf3359..0286b4a 100644 --- a/skemr-api/internal/dbreflect/schema_sync.go +++ b/skemr-api/internal/dbreflect/schema_sync.go @@ -155,7 +155,7 @@ func (s *SchemaSyncService) SyncSchema(c context.Context, database models.Databa // Mark any entities that were not found in the new schema as deleted for _, entity := range savedEntities { if !slices.Contains(currentEntityIds, entity.ID) { - err := s.markEntityAsDeleted(c, entity.ID) + err := s.markEntityAsDeleted(c, database.ID, entity.ID) if err != nil { slog.Error("Error marking entity as deleted", "entityId", entity.ID, "error", err) // Do not return error as we want to continue marking other entities as deleted @@ -175,7 +175,7 @@ func (s *SchemaSyncService) updateNamespace(c context.Context, schemaRef SchemaR EntityType: sqlc.DatabaseEntityTypeSchema, Name: schemaRef.Name, } - schema, err := s.db.GetDatabaseEntityByDatabaseIdAndTypeAndParentAndName(c, args) + namespace, err := s.db.GetDatabaseEntityByDatabaseIdAndTypeAndParentAndName(c, args) if err != nil && !errors.Is(err, pgx.ErrNoRows) { slog.Error("error getting schema", "error", err.Error()) return sqlc.DatabaseEntity{}, err @@ -187,7 +187,7 @@ func (s *SchemaSyncService) updateNamespace(c context.Context, schemaRef SchemaR fingerprint := GenerateNamespaceFingerprint(schemaRef) - schema, err = s.db.GetDatabaseEntityByFingerprint(c, sqlc.GetDatabaseEntityByFingerprintParams{ + namespace, err = s.db.GetDatabaseEntityByFingerprint(c, sqlc.GetDatabaseEntityByFingerprintParams{ DatabaseID: database.ID, Fingerprint: fingerprint, }) @@ -199,19 +199,21 @@ func (s *SchemaSyncService) updateNamespace(c context.Context, schemaRef SchemaR // If found by fingerprint, update the name to the new name. if !errors.Is(err, pgx.ErrNoRows) { - slog.Debug("Schema found by fingerprint, updating name", "oldName", schema.Name, "newName", schemaRef.Name) + slog.Debug("Schema found by fingerprint, updating name", "oldName", namespace.Name, "newName", schemaRef.Name) - schema, err = s.db.UpdateDatabaseEntityName(c, sqlc.UpdateDatabaseEntityNameParams{ - ID: schema.ID, + namespace, err = s.db.UpdateDatabaseEntityName(c, sqlc.UpdateDatabaseEntityNameParams{ + ID: namespace.ID, Name: schemaRef.Name, }) if err != nil { slog.Error("error updating schema name", "error", err.Error()) return sqlc.DatabaseEntity{}, err } - slog.Info("Schema renamed", "oldName", schema.Name, "newName", schemaRef.Name) - schema.Name = schemaRef.Name - return schema, nil + + s.markDatabaseChange(c, database.ID, namespace.ID, models.MigrationStatementActionUpdate) + slog.Info("Schema renamed", "oldName", namespace.Name, "newName", schemaRef.Name) + namespace.Name = schemaRef.Name + return namespace, nil } // If that schema does not exist yet, save it @@ -222,35 +224,53 @@ func (s *SchemaSyncService) updateNamespace(c context.Context, schemaRef SchemaR Name: schemaRef.Name, Fingerprint: fingerprint, } - schema, err := s.db.CreateDatabaseEntity(c, args) + namespace, err := s.db.CreateDatabaseEntity(c, args) if err != nil { slog.Error("error creating schema", "error", err) return sqlc.DatabaseEntity{}, err } - slog.Info("schema created", "schema", schema.Name) - return schema, err + + s.markDatabaseChange(c, database.ID, namespace.ID, models.MigrationStatementActionCreate) + slog.Info("namespace created", "namespace", namespace.Name) + return namespace, err } - // TODO: If schema does exist, update if needed. Currently no-op - slog.Info("schema exists", "schema", schema) + // TODO: If namespace does exist, update if needed. Currently no-op + slog.Info("namespace exists", "namespace", namespace) - return schema, err + return namespace, err } // markEntityAsDeleted sets the status of the entity to "deleted". // This is used for entities that were not found in the new schema during sync. // We want to keep these entities in the database to show how the schema has changed. -func (s *SchemaSyncService) markEntityAsDeleted(c context.Context, entityId uuid.UUID) error { +func (s *SchemaSyncService) markEntityAsDeleted(c context.Context, databaseId uuid.UUID, entityId uuid.UUID) error { slog.Debug("Marking entity as deleted", "entityId", entityId) err := s.db.UpdateDatabaseEntityAsDeleted(c, entityId) if err != nil { slog.Error("error marking entity as deleted", "error", err) return err } + + s.markDatabaseChange(c, databaseId, entityId, models.MigrationStatementActionDelete) + return nil } +func (s *SchemaSyncService) markDatabaseChange(c context.Context, databaseId uuid.UUID, entityId uuid.UUID, action models.MigrationStatementAction) { + slog.Debug("Marking database change as deleted", "databaseId", databaseId, "entityId", entityId) + _, err := s.db.CreateDatabaseChange(c, sqlc.CreateDatabaseChangeParams{ + DatabaseID: databaseId, + EntityID: entityId, + Action: sqlc.MigrationStatementAction(action), + }) + + if err != nil { + slog.Error("error creating database change", "error", err) + } +} + func (s *SchemaSyncService) updateTable(c context.Context, tableRef TableRef, database models.Database, schemaId uuid.UUID) (sqlc.DatabaseEntity, error) { args := sqlc.GetDatabaseEntityByDatabaseIdAndTypeAndParentAndNameParams{ DatabaseID: database.ID, @@ -294,6 +314,7 @@ func (s *SchemaSyncService) updateTable(c context.Context, tableRef TableRef, da slog.Error("error updating table name", "error", err.Error()) return sqlc.DatabaseEntity{}, err } + s.markDatabaseChange(c, database.ID, table.ID, models.MigrationStatementActionUpdate) slog.Info("Table renamed", "oldName", table.Name, "newName", tableRef.Name) table.Name = tableRef.Name return table, nil @@ -312,6 +333,8 @@ func (s *SchemaSyncService) updateTable(c context.Context, tableRef TableRef, da slog.Error("error creating table", "error", err) return sqlc.DatabaseEntity{}, err } + + s.markDatabaseChange(c, database.ID, table.ID, models.MigrationStatementActionCreate) slog.Info("Table created", "schema", table.Name) return table, err } @@ -375,6 +398,8 @@ func (s *SchemaSyncService) updateColumn(c context.Context, columnRef ColumnRef, slog.Error("error updating column name", "error", err.Error()) return sqlc.DatabaseEntity{}, err } + + s.markDatabaseChange(c, database.ID, column.ID, models.MigrationStatementActionUpdate) slog.Info("Column renamed", "oldName", column.Name, "newName", columnRef.Name) column.Name = columnRef.Name return column, nil @@ -395,6 +420,8 @@ func (s *SchemaSyncService) updateColumn(c context.Context, columnRef ColumnRef, slog.Error("error creating column", "error", err) return sqlc.DatabaseEntity{}, err } + + s.markDatabaseChange(c, database.ID, column.ID, models.MigrationStatementActionCreate) slog.Info("Column created", "name", column.Name) return column, err } diff --git a/skemr-api/internal/errormsg/errors.go b/skemr-api/internal/errormsg/errors.go index dbcba36..cf1ac30 100644 --- a/skemr-api/internal/errormsg/errors.go +++ b/skemr-api/internal/errormsg/errors.go @@ -29,11 +29,19 @@ func WriteErrorResponse(w http.ResponseWriter, r *http.Request, err error) { } var ( - ErrDatabaseAlreadyExists = "database already exists" - ErrDatabaseNotFound = "database not found" - ErrProjectNotFound = "project not found" - ErrInvalidIdFormat = "invalid id format" - ErrExpiryTimeInPast = "expiry time is in the past" - ErrRuleWithSameName = "rule with the same name already exists" + // DatabaseChange + ErrDatabaseChangeNotFound = "database change not found" + ErrDatabaseChangeFetchFailed = "failed to fetch database changes" + // Database + ErrDatabaseAlreadyExists = "database already exists" + ErrDatabaseNotFound = "database not found" + // Project + ErrProjectNotFound = "project not found" + // Validation + ErrInvalidIdFormat = "invalid id format" + ErrExpiryTimeInPast = "expiry time is in the past" + // Rule + ErrRuleWithSameName = "rule with the same name already exists" + // DatabaseEntity ErrDatabaseEntityNotFound = "database entity not found" ) diff --git a/skemr-api/internal/mapper/database_change_mapper.go b/skemr-api/internal/mapper/database_change_mapper.go new file mode 100644 index 0000000..758fb06 --- /dev/null +++ b/skemr-api/internal/mapper/database_change_mapper.go @@ -0,0 +1,24 @@ +package mapper + +import ( + "github.com/walmaa/skemr-api/db/sqlc" + "github.com/walmaa/skemr-common/models" +) + +func ToDomainDatabaseChange(change sqlc.DatabaseChange) models.DatabaseChange { + return models.DatabaseChange{ + Id: change.ID, + DatabaseId: change.DatabaseID, + EntityId: change.EntityID, + Action: models.MigrationStatementAction(change.Action), + CreatedAt: Time(&change.CreatedAt), + } +} + +func ToDomainDatabaseChanges(changes []sqlc.DatabaseChange) []models.DatabaseChange { + databaseChanges := make([]models.DatabaseChange, len(changes)) + for i, change := range changes { + databaseChanges[i] = ToDomainDatabaseChange(change) + } + return databaseChanges +} diff --git a/skemr-api/internal/routers/router.go b/skemr-api/internal/routers/router.go index a3e26a9..18fea2c 100644 --- a/skemr-api/internal/routers/router.go +++ b/skemr-api/internal/routers/router.go @@ -21,6 +21,7 @@ type Services struct { AccessTokenService *service.AccessTokenService DatabaseEntityService *service.DatabaseEntityService IntegrationService *service.IntegrationService + DatabaseChangeService *service.DatabaseChangeService } func InitRouter(services *Services) http.Handler { @@ -76,10 +77,12 @@ func InitRouter(services *Services) http.Handler { projectSecretsController := controller.NewProjectSecretsController(services.AccessTokenService) ruleController := controller.NewRuleController(services.RuleService) databaseEntityController := controller.NewDatabaseEntityController(services.DatabaseEntityService) + databaseChangeController := controller.NewDatabaseChangeController(services.DatabaseChangeService) databaseController.RegisterRoutes(r) projectSecretsController.RegisterRoutes(r) ruleController.RegisterRoutes(r) databaseEntityController.RegisterRoutes(r) + databaseChangeController.RegisterRoutes(r) r.Get("/", projectController.GetProject) r.Delete("/", projectController.DeleteProject) diff --git a/skemr-api/internal/service/database_change_service.go b/skemr-api/internal/service/database_change_service.go new file mode 100644 index 0000000..c6e6d19 --- /dev/null +++ b/skemr-api/internal/service/database_change_service.go @@ -0,0 +1,94 @@ +package service + +import ( + "context" + "errors" + "log/slog" + "net/http" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgtype" + "github.com/walmaa/skemr-api/db/sqlc" + "github.com/walmaa/skemr-api/internal/errormsg" + "github.com/walmaa/skemr-api/internal/mapper" + "github.com/walmaa/skemr-common/models" +) + +type DatabaseChangeStore interface { + GetDatabaseChangeByDatabaseIdAndId(ctx context.Context, arg sqlc.GetDatabaseChangeByDatabaseIdAndIdParams) (sqlc.DatabaseChange, error) + GetDatabaseChangesByDatabaseIdAndId(c context.Context, params sqlc.GetDatabaseChangesByDatabaseIdAndIdParams) ([]sqlc.DatabaseChange, error) +} + +type DatabaseChangeService struct { + store DatabaseChangeStore + scopeResolver ScopeResolver +} + +func NewDatabaseChangeService(store DatabaseChangeStore, resolver ScopeResolver) *DatabaseChangeService { + return &DatabaseChangeService{store: store, scopeResolver: resolver} +} + +func (s *DatabaseChangeService) GetDatabaseChange(c context.Context, projectId uuid.UUID, databaseId uuid.UUID, id uuid.UUID) (models.DatabaseChange, error) { + slog.Info("Getting database change", "id", id, "projectId", projectId, "databaseId", databaseId) + + _, err := s.scopeResolver.RequireDatabase(c, projectId, databaseId) + if err != nil { + slog.Error("Error fetching database", "err", err) + return models.DatabaseChange{}, err + } + + databaseChange, err := s.store.GetDatabaseChangeByDatabaseIdAndId(c, sqlc.GetDatabaseChangeByDatabaseIdAndIdParams{ + ID: id, + DatabaseID: databaseId, + }) + + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + slog.Error("Error fetching database change", "err", err) + return models.DatabaseChange{}, &models.ErrorResponse{} + } + + if errors.Is(err, pgx.ErrNoRows) { + slog.Info("Database change not found", "id", id, "projectId", projectId, "databaseId", databaseId) + return models.DatabaseChange{}, &models.ErrorResponse{ + Message: errormsg.ErrDatabaseChangeNotFound, + Status: http.StatusNotFound, + Errors: nil, + } + } + + return mapper.ToDomainDatabaseChange(databaseChange), nil +} + +func (s *DatabaseChangeService) GetDatabaseChanges(c context.Context, projectId uuid.UUID, databaseId uuid.UUID, limit int, offset int) ([]models.DatabaseChange, error) { + slog.Info("Getting database changes", "limit", limit, "offset", offset, "projectId", projectId, "databaseId", databaseId) + + _, err := s.scopeResolver.RequireDatabase(c, projectId, databaseId) + if err != nil { + slog.Error("Error fetching database", "err", err) + return nil, err + } + + databaseChanges, err := s.store.GetDatabaseChangesByDatabaseIdAndId(c, sqlc.GetDatabaseChangesByDatabaseIdAndIdParams{ + DatabaseID: databaseId, + Offset: pgtype.Int4{ + Int32: int32(offset), + Valid: true, + }, + Limit: pgtype.Int4{ + Int32: int32(limit), + Valid: true, + }, + }) + + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + slog.Error("Error fetching database changes", "err", err) + return nil, &models.ErrorResponse{ + Errors: nil, + Status: http.StatusInternalServerError, + Message: "Error fetching database changes", + } + } + + return mapper.ToDomainDatabaseChanges(databaseChanges), nil +} diff --git a/skemr-common/models/database_change.go b/skemr-common/models/database_change.go new file mode 100644 index 0000000..dd499f9 --- /dev/null +++ b/skemr-common/models/database_change.go @@ -0,0 +1,15 @@ +package models + +import ( + "time" + + "github.com/google/uuid" +) + +type DatabaseChange struct { + Id uuid.UUID `json:"id"` + DatabaseId uuid.UUID `json:"databaseId"` + EntityId uuid.UUID `json:"entityId"` + Action MigrationStatementAction `json:"action"` + CreatedAt time.Time `json:"createdAt"` +} From 41e3a3024e48c7d172217b5b6a67221348b6c576 Mon Sep 17 00:00:00 2001 From: WalMaa Date: Mon, 25 May 2026 10:01:12 +0300 Subject: [PATCH 11/14] pipeline run persistence model and service --- .../migrations/20260225122322_init_schema.sql | 11 ++ skemr-api/db/queries/pipeline_runs.sql | 18 +++ skemr-api/db/sqlc/models.go | 10 ++ skemr-api/db/sqlc/querier.go | 3 + .../controller/database_change_controller.go | 2 +- skemr-api/internal/dto/common.go | 6 + skemr-api/internal/errormsg/errors.go | 3 + .../internal/mapper/database_change_mapper.go | 9 +- .../internal/mapper/pipeline_run_mapper.go | 25 ++++ .../internal/service/pipeline_run_service.go | 134 ++++++++++++++++++ skemr-common/models/database_change.go | 9 +- skemr-common/models/pipeline_run.go | 16 +++ 12 files changed, 235 insertions(+), 11 deletions(-) create mode 100644 skemr-api/db/queries/pipeline_runs.sql create mode 100644 skemr-api/internal/mapper/pipeline_run_mapper.go create mode 100644 skemr-api/internal/service/pipeline_run_service.go create mode 100644 skemr-common/models/pipeline_run.go diff --git a/skemr-api/db/migrations/20260225122322_init_schema.sql b/skemr-api/db/migrations/20260225122322_init_schema.sql index df936a7..469d5e5 100644 --- a/skemr-api/db/migrations/20260225122322_init_schema.sql +++ b/skemr-api/db/migrations/20260225122322_init_schema.sql @@ -150,6 +150,17 @@ CREATE TABLE database_changes created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); +CREATE TABLE pipeline_runs +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + database_id UUID NOT NULL REFERENCES databases (id) ON DELETE CASCADE, + status migration_status NOT NULL DEFAULT 'pending', + environment TEXT, + started_at TIMESTAMPTZ, + completed_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + -- Rules specify the protection mechanisms for databases, schemas, tables, and columns. CREATE TABLE rules diff --git a/skemr-api/db/queries/pipeline_runs.sql b/skemr-api/db/queries/pipeline_runs.sql new file mode 100644 index 0000000..90f884e --- /dev/null +++ b/skemr-api/db/queries/pipeline_runs.sql @@ -0,0 +1,18 @@ +-- name: GetPipelineRunByDatabaseIdAndId :one +SELECT * +FROM pipeline_runs +WHERE database_id = @database_id + AND id = @id +LIMIT 1; + +-- name: GetPipelineRunsByDatabaseId :many +SELECT * +FROM pipeline_runs +WHERE database_id = @database_id +ORDER BY created_at DESC; + +-- name: CreatePipelineRun :one +INSERT INTO pipeline_runs + (database_id, status, environment, completed_at) +VALUES (@database_id, @status, @environment, @completed_at) +RETURNING *; \ No newline at end of file diff --git a/skemr-api/db/sqlc/models.go b/skemr-api/db/sqlc/models.go index 6cdf315..a218857 100644 --- a/skemr-api/db/sqlc/models.go +++ b/skemr-api/db/sqlc/models.go @@ -315,6 +315,16 @@ type DatabaseEntity struct { CreatedAt pgtype.Timestamptz `json:"created_at"` } +type PipelineRun struct { + ID uuid.UUID `json:"id"` + DatabaseID uuid.UUID `json:"database_id"` + Status MigrationStatus `json:"status"` + Environment pgtype.Text `json:"environment"` + StartedAt pgtype.Timestamptz `json:"started_at"` + CompletedAt pgtype.Timestamptz `json:"completed_at"` + CreatedAt pgtype.Timestamptz `json:"created_at"` +} + type Project struct { ID uuid.UUID `json:"id"` Name string `json:"name"` diff --git a/skemr-api/db/sqlc/querier.go b/skemr-api/db/sqlc/querier.go index e7c4c06..7cb8c9d 100644 --- a/skemr-api/db/sqlc/querier.go +++ b/skemr-api/db/sqlc/querier.go @@ -14,6 +14,7 @@ type Querier interface { CreateDatabase(ctx context.Context, arg CreateDatabaseParams) (Database, error) CreateDatabaseChange(ctx context.Context, arg CreateDatabaseChangeParams) (DatabaseChange, error) CreateDatabaseEntity(ctx context.Context, arg CreateDatabaseEntityParams) (DatabaseEntity, error) + CreatePipelineRun(ctx context.Context, arg CreatePipelineRunParams) (PipelineRun, error) CreateProject(ctx context.Context, name string) (Project, error) CreateProjectSecretKey(ctx context.Context, arg CreateProjectSecretKeyParams) (CreateProjectSecretKeyRow, error) CreateRule(ctx context.Context, arg CreateRuleParams) (Rule, error) @@ -37,6 +38,8 @@ type Querier interface { GetDatabaseEntityByProjectIdAndId(ctx context.Context, arg GetDatabaseEntityByProjectIdAndIdParams) (DatabaseEntity, error) GetDatabaseEntityByProjectIdDatabaseIdAndId(ctx context.Context, arg GetDatabaseEntityByProjectIdDatabaseIdAndIdParams) (DatabaseEntity, error) GetHashByPrefixAndProjectID(ctx context.Context, arg GetHashByPrefixAndProjectIDParams) (string, error) + GetPipelineRunByDatabaseIdAndId(ctx context.Context, arg GetPipelineRunByDatabaseIdAndIdParams) (PipelineRun, error) + GetPipelineRunsByDatabaseId(ctx context.Context, databaseID uuid.UUID) ([]PipelineRun, error) GetProject(ctx context.Context, id uuid.UUID) (Project, error) GetProjectAccessTokens(ctx context.Context, projectID uuid.UUID) ([]ProjectAccessToken, error) GetProjectBySecretPrefix(ctx context.Context, prefix string) (Project, error) diff --git a/skemr-api/internal/controller/database_change_controller.go b/skemr-api/internal/controller/database_change_controller.go index 06181ca..0a374a1 100644 --- a/skemr-api/internal/controller/database_change_controller.go +++ b/skemr-api/internal/controller/database_change_controller.go @@ -48,7 +48,7 @@ func (h *DatabaseChangeController) listDatabaseChanges(w http.ResponseWriter, r } func (h *DatabaseChangeController) GetDatabaseChange(w http.ResponseWriter, r *http.Request) { - projectID, ok := r.Context().Value("projectID").(uuid.UUID) + projectID, ok := r.Context().Value("projectId").(uuid.UUID) if !ok { http.Error(w, "projectId not found in context", http.StatusBadRequest) return diff --git a/skemr-api/internal/dto/common.go b/skemr-api/internal/dto/common.go index 5fe360d..2730c5d 100644 --- a/skemr-api/internal/dto/common.go +++ b/skemr-api/internal/dto/common.go @@ -48,3 +48,9 @@ type SecretCreationDto struct { Name string `json:"name" validate:"required,min=2,max=100"` ExpiresAt string `json:"expiresAt" validate:"omitempty,datetime=2006-01-02T15:04:05Z07:00"` } + +type PipelineRunCreationDto struct { + Status models.MigrationStatus `json:"status" validate:"required,oneof=completed failed"` + Environment string `json:"environment" validate:"required"` + CompletedAt string `json:"completedAt" validate:"required,datetime=2006-01-02T15:04:05Z07:00"` +} diff --git a/skemr-api/internal/errormsg/errors.go b/skemr-api/internal/errormsg/errors.go index cf1ac30..67e0fc8 100644 --- a/skemr-api/internal/errormsg/errors.go +++ b/skemr-api/internal/errormsg/errors.go @@ -44,4 +44,7 @@ var ( ErrRuleWithSameName = "rule with the same name already exists" // DatabaseEntity ErrDatabaseEntityNotFound = "database entity not found" + // PipelineRun + ErrPipelineRunNotFound = "pipeline run not found" + ErrPipelineRunFetchFailed = "failed to fetch pipeline runs" ) diff --git a/skemr-api/internal/mapper/database_change_mapper.go b/skemr-api/internal/mapper/database_change_mapper.go index 758fb06..36ae89c 100644 --- a/skemr-api/internal/mapper/database_change_mapper.go +++ b/skemr-api/internal/mapper/database_change_mapper.go @@ -7,11 +7,10 @@ import ( func ToDomainDatabaseChange(change sqlc.DatabaseChange) models.DatabaseChange { return models.DatabaseChange{ - Id: change.ID, - DatabaseId: change.DatabaseID, - EntityId: change.EntityID, - Action: models.MigrationStatementAction(change.Action), - CreatedAt: Time(&change.CreatedAt), + Id: change.ID, + EntityId: change.EntityID, + Action: models.MigrationStatementAction(change.Action), + CreatedAt: Time(&change.CreatedAt), } } diff --git a/skemr-api/internal/mapper/pipeline_run_mapper.go b/skemr-api/internal/mapper/pipeline_run_mapper.go new file mode 100644 index 0000000..7c69b51 --- /dev/null +++ b/skemr-api/internal/mapper/pipeline_run_mapper.go @@ -0,0 +1,25 @@ +package mapper + +import ( + "github.com/walmaa/skemr-api/db/sqlc" + "github.com/walmaa/skemr-common/models" +) + +func ToDomainPipelineRun(r sqlc.PipelineRun) models.PipelineRun { + return models.PipelineRun{ + ID: r.ID, + Status: models.MigrationStatus(r.Status), + Environment: r.Environment.String, + StartedAt: Time(&r.StartedAt), + CompletedAt: Time(&r.CompletedAt), + CreatedAt: Time(&r.CreatedAt), + } +} + +func ToDomainPipelineRuns(r []sqlc.PipelineRun) []models.PipelineRun { + pipelineRuns := make([]models.PipelineRun, len(r)) + for i, run := range r { + pipelineRuns[i] = ToDomainPipelineRun(run) + } + return pipelineRuns +} diff --git a/skemr-api/internal/service/pipeline_run_service.go b/skemr-api/internal/service/pipeline_run_service.go new file mode 100644 index 0000000..f3fc27a --- /dev/null +++ b/skemr-api/internal/service/pipeline_run_service.go @@ -0,0 +1,134 @@ +package service + +import ( + "context" + "errors" + "log/slog" + "net/http" + "time" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgtype" + "github.com/walmaa/skemr-api/db/sqlc" + "github.com/walmaa/skemr-api/internal/dto" + "github.com/walmaa/skemr-api/internal/errormsg" + "github.com/walmaa/skemr-api/internal/mapper" + "github.com/walmaa/skemr-common/models" +) + +type PipelineRunStore interface { + GetPipelineRunByDatabaseIdAndId(ctx context.Context, arg sqlc.GetPipelineRunByDatabaseIdAndIdParams) (sqlc.PipelineRun, error) + GetPipelineRunsByDatabaseId(ctx context.Context, databaseId uuid.UUID) ([]sqlc.PipelineRun, error) + CreatePipelineRun(ctx context.Context, pipelineRun sqlc.CreatePipelineRunParams) (sqlc.PipelineRun, error) +} + +type PipelineRunService struct { + store PipelineRunStore + scopeResolver ScopeResolver +} + +func NewPipelineRunService(store PipelineRunStore, resolver ScopeResolver) *PipelineRunService { + return &PipelineRunService{store: store, scopeResolver: resolver} +} + +func (s *PipelineRunService) GetPipelineRun(c context.Context, projectId uuid.UUID, databaseId uuid.UUID, id uuid.UUID) (models.PipelineRun, error) { + slog.Info("Getting pipeline run", "id", id, "databaseId", databaseId, "projectId", projectId) + _, err := s.scopeResolver.RequireDatabase(c, projectId, databaseId) + if err != nil { + slog.Error("Error fetching database", "err", err) + return models.PipelineRun{}, err + } + + pipelineRun, err := s.store.GetPipelineRunByDatabaseIdAndId(c, sqlc.GetPipelineRunByDatabaseIdAndIdParams{ + DatabaseID: databaseId, + ID: id, + }) + + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + slog.Error("Error fetching pipeline run", "err", err) + return models.PipelineRun{}, &models.ErrorResponse{ + Errors: nil, + Status: http.StatusInternalServerError, + Message: errormsg.ErrPipelineRunFetchFailed, + } + } + + if errors.Is(err, pgx.ErrNoRows) { + slog.Info("Pipeline run not found", "id", id, "databaseId", databaseId, "projectId", projectId) + return models.PipelineRun{}, &models.ErrorResponse{ + Errors: nil, + Status: http.StatusNotFound, + Message: errormsg.ErrPipelineRunNotFound, + } + } + + return mapper.ToDomainPipelineRun(pipelineRun), nil +} + +func (s *PipelineRunService) GetPipelineRuns(c context.Context, projectId uuid.UUID, databaseId uuid.UUID) ([]models.PipelineRun, error) { + slog.Info("Getting pipeline runs", "databaseId", databaseId, "projectId", projectId) + _, err := s.scopeResolver.RequireDatabase(c, projectId, databaseId) + if err != nil { + slog.Error("Error fetching database", "err", err) + return nil, err + } + + pipelineRuns, err := s.store.GetPipelineRunsByDatabaseId(c, databaseId) + + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + slog.Error("Error fetching pipeline runs", "err", err) + return nil, &models.ErrorResponse{ + Errors: nil, + Status: http.StatusInternalServerError, + Message: errormsg.ErrPipelineRunFetchFailed, + } + } + + if errors.Is(err, pgx.ErrNoRows) { + slog.Info("No pipeline runs found", "databaseId", databaseId, "projectId", projectId) + return nil, &models.ErrorResponse{ + Errors: nil, + Status: http.StatusNotFound, + Message: errormsg.ErrPipelineRunNotFound, + } + } + + return mapper.ToDomainPipelineRuns(pipelineRuns), nil +} + +func (s *PipelineRunService) CreatePipelineRun(c context.Context, projectId uuid.UUID, databaseId uuid.UUID, pipelineRun dto.PipelineRunCreationDto) (sqlc.PipelineRun, error) { + slog.Info("Creating pipeline run", "pipelineRun", pipelineRun, "databaseId", databaseId, "projectId", projectId) + _, err := s.scopeResolver.RequireDatabase(c, projectId, databaseId) + if err != nil { + slog.Error("Error fetching database", "err", err) + return sqlc.PipelineRun{}, err + } + + completedAt, err := time.Parse(time.RFC3339, pipelineRun.CompletedAt) + + if err != nil { + slog.Error("Error parsing completedAt", "err", err) + return sqlc.PipelineRun{}, err + } + + createdPipelineRun, err := s.store.CreatePipelineRun(c, sqlc.CreatePipelineRunParams{ + DatabaseID: databaseId, + Status: sqlc.MigrationStatus(pipelineRun.Status), + Environment: pgtype.Text{ + Valid: true, + String: pipelineRun.Environment, + }, + CompletedAt: pgtype.Timestamptz{ + Valid: true, + Time: completedAt, + }, + }) + + if err != nil { + slog.Error("Error creating pipeline run", "err", err) + return sqlc.PipelineRun{}, err + } + + return createdPipelineRun, nil +} diff --git a/skemr-common/models/database_change.go b/skemr-common/models/database_change.go index dd499f9..4fd442e 100644 --- a/skemr-common/models/database_change.go +++ b/skemr-common/models/database_change.go @@ -7,9 +7,8 @@ import ( ) type DatabaseChange struct { - Id uuid.UUID `json:"id"` - DatabaseId uuid.UUID `json:"databaseId"` - EntityId uuid.UUID `json:"entityId"` - Action MigrationStatementAction `json:"action"` - CreatedAt time.Time `json:"createdAt"` + Id uuid.UUID `json:"id"` + EntityId uuid.UUID `json:"entityId"` + Action MigrationStatementAction `json:"action"` + CreatedAt time.Time `json:"createdAt"` } diff --git a/skemr-common/models/pipeline_run.go b/skemr-common/models/pipeline_run.go new file mode 100644 index 0000000..efd4300 --- /dev/null +++ b/skemr-common/models/pipeline_run.go @@ -0,0 +1,16 @@ +package models + +import ( + "time" + + "github.com/google/uuid" +) + +type PipelineRun struct { + ID uuid.UUID `json:"id"` + Status MigrationStatus `json:"status"` + Environment string `json:"environment"` + StartedAt time.Time `json:"startedAt"` + CompletedAt time.Time `json:"completedAt"` + CreatedAt time.Time `json:"createdAt"` +} From 7ee8e9f4522cc8ae86c80e9bcbd97699542aacc2 Mon Sep 17 00:00:00 2001 From: WalMaa Date: Mon, 25 May 2026 10:38:13 +0300 Subject: [PATCH 12/14] pipeline integration endpoint --- skemr-api/cmd/server/main.go | 18 +++--- .../controller/integration_controller.go | 55 ++++++++++++++++++- skemr-api/internal/errormsg/errors.go | 12 +++- .../middleware/access_token_middleware.go | 3 +- skemr-api/internal/routers/router.go | 1 + .../internal/service/access_token_service.go | 2 +- .../internal/service/integration_service.go | 12 +++- .../internal/service/pipeline_run_service.go | 10 ++-- 8 files changed, 90 insertions(+), 23 deletions(-) diff --git a/skemr-api/cmd/server/main.go b/skemr-api/cmd/server/main.go index d8e1245..f1a9662 100644 --- a/skemr-api/cmd/server/main.go +++ b/skemr-api/cmd/server/main.go @@ -92,6 +92,11 @@ func main() { TLSConfig: nil, }) + if cfg.App.Env == "dev" { + runSchema(conn) + seedTestData(conn) + } + queries := sqlc.New(conn) scopeResolver := service.NewScopeResolver(queries) projectService := service.NewProjectService(queries) @@ -101,14 +106,8 @@ func main() { projectSecretsService := service.NewAccessTokenService(queries) ruleService := service.NewRuleService(queries, scopeResolver) databaseEntityService := service.NewDatabaseEntityService(queries) - integrationService := service.NewIntegrationService(ruleService) - - if cfg.App.Env == "dev" { - runSchema(conn) - seedTestData(conn) - } - - worker.StartTaskWorkers(queries, cfg) + pipelineRunService := service.NewPipelineRunService(queries, scopeResolver) + integrationService := service.NewIntegrationService(ruleService, pipelineRunService) // Initialize services services := &routers.Services{ @@ -120,8 +119,11 @@ func main() { DatabaseEntityService: databaseEntityService, IntegrationService: integrationService, DatabaseChangeService: databaseChangeService, + PipelineRunService: pipelineRunService, } + worker.StartTaskWorkers(queries, cfg) + // Initialize router router := routers.InitRouter(services) diff --git a/skemr-api/internal/controller/integration_controller.go b/skemr-api/internal/controller/integration_controller.go index 98d09e4..2966135 100644 --- a/skemr-api/internal/controller/integration_controller.go +++ b/skemr-api/internal/controller/integration_controller.go @@ -6,21 +6,24 @@ import ( "github.com/go-chi/chi/v5" "github.com/go-chi/render" "github.com/google/uuid" + "github.com/walmaa/skemr-api/internal/dto" "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/service" + "github.com/walmaa/skemr-api/internal/validation" "github.com/walmaa/skemr-common/models" ) type IntegrationController struct { - IntegrationService *service.IntegrationService + integrationService *service.IntegrationService } func NewIntegrationController(s *service.IntegrationService) *IntegrationController { - return &IntegrationController{IntegrationService: s} + return &IntegrationController{integrationService: s} } func (h *IntegrationController) RegisterRoutes(r chi.Router) { r.Get("/ci-cd/rules", h.listRulesByDatabase) + r.Post("/ci-cd/pipeline-runs", h.createPipelineRun) } @@ -47,7 +50,7 @@ func (h *IntegrationController) listRulesByDatabase(w http.ResponseWriter, r *ht return } - rules, err := h.IntegrationService.ListRulesByDatabase(r.Context(), projectId, databaseId) + rules, err := h.integrationService.ListRulesByDatabase(r.Context(), projectId, databaseId) if err != nil { errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ Message: "Error fetching rules", @@ -59,3 +62,49 @@ func (h *IntegrationController) listRulesByDatabase(w http.ResponseWriter, r *ht } render.JSON(w, r, rules) } + +func (h *IntegrationController) createPipelineRun(w http.ResponseWriter, r *http.Request) { + + databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) + if err != nil { + errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ + Message: "Invalid database ID format", + Status: http.StatusBadRequest, + Errors: nil, + }, + ) + return + } + + projectId, err := uuid.Parse(chi.URLParam(r, "projectId")) + if err != nil { + errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ + Message: "Invalid project ID format", + Status: http.StatusBadRequest, + Errors: nil, + }, + ) + return + } + + var req dto.PipelineRunCreationDto + if err := render.DecodeJSON(r.Body, &req); err != nil { + errormsg.WriteInvalidRequestBodyErrorResponse(w, r) + return + } + + err = validation.Validate.Struct(req) + + if err != nil { + errorResponse := validation.CreateErrorResponse(err) + errormsg.WriteErrorResponse(w, r, &errorResponse) + return + } + + rule, err := h.integrationService.CreatePipeLineRun(r.Context(), projectId, databaseId, req) + if err != nil { + errormsg.WriteErrorResponse(w, r, err) + return + } + render.JSON(w, r, rule) +} diff --git a/skemr-api/internal/errormsg/errors.go b/skemr-api/internal/errormsg/errors.go index 67e0fc8..03b8bdc 100644 --- a/skemr-api/internal/errormsg/errors.go +++ b/skemr-api/internal/errormsg/errors.go @@ -28,6 +28,13 @@ func WriteErrorResponse(w http.ResponseWriter, r *http.Request, err error) { ) } +func WriteInvalidRequestBodyErrorResponse(w http.ResponseWriter, r *http.Request) { + WriteErrorResponse(w, r, &models.ErrorResponse{ + Message: "Invalid request body", + Status: http.StatusBadRequest, + }) +} + var ( // DatabaseChange ErrDatabaseChangeNotFound = "database change not found" @@ -45,6 +52,7 @@ var ( // DatabaseEntity ErrDatabaseEntityNotFound = "database entity not found" // PipelineRun - ErrPipelineRunNotFound = "pipeline run not found" - ErrPipelineRunFetchFailed = "failed to fetch pipeline runs" + ErrPipelineRunNotFound = "pipeline run not found" + ErrPipelineRunFetchFailed = "failed to fetch pipeline runs" + ErrPipelineRunCreateFailed = "failed to create pipeline run" ) diff --git a/skemr-api/internal/middleware/access_token_middleware.go b/skemr-api/internal/middleware/access_token_middleware.go index b0f7419..73db030 100644 --- a/skemr-api/internal/middleware/access_token_middleware.go +++ b/skemr-api/internal/middleware/access_token_middleware.go @@ -3,6 +3,7 @@ package middleware import ( "context" "net/http" + "strings" "github.com/go-chi/chi/v5" "github.com/google/uuid" @@ -48,7 +49,7 @@ func AccessTokenMiddleware(service *service.AccessTokenService) func(next http.H } token := tokenHeaderValue[len("Bearer "):] - ok, err := authenticateToken(c, service, projectId, token) + ok, err := authenticateToken(c, service, projectId, strings.TrimSpace(token)) if err != nil { errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ diff --git a/skemr-api/internal/routers/router.go b/skemr-api/internal/routers/router.go index 18fea2c..eef8d08 100644 --- a/skemr-api/internal/routers/router.go +++ b/skemr-api/internal/routers/router.go @@ -22,6 +22,7 @@ type Services struct { DatabaseEntityService *service.DatabaseEntityService IntegrationService *service.IntegrationService DatabaseChangeService *service.DatabaseChangeService + PipelineRunService *service.PipelineRunService } func InitRouter(services *Services) http.Handler { diff --git a/skemr-api/internal/service/access_token_service.go b/skemr-api/internal/service/access_token_service.go index f2953ec..a267cbb 100644 --- a/skemr-api/internal/service/access_token_service.go +++ b/skemr-api/internal/service/access_token_service.go @@ -137,7 +137,7 @@ func (s *AccessTokenService) DeleteToken(c context.Context, projectId uuid.UUID, } func (s *AccessTokenService) ValidateToken(c context.Context, projectId uuid.UUID, token string) (bool, error) { - slog.Info("Validating token") + slog.Info("Validating token", "projectId", projectId) // Extract the prefix to find the token in the database diff --git a/skemr-api/internal/service/integration_service.go b/skemr-api/internal/service/integration_service.go index 6e3fbea..2aad778 100644 --- a/skemr-api/internal/service/integration_service.go +++ b/skemr-api/internal/service/integration_service.go @@ -4,17 +4,23 @@ import ( "context" "github.com/google/uuid" + "github.com/walmaa/skemr-api/internal/dto" "github.com/walmaa/skemr-common/models" ) type IntegrationService struct { - RuleService *RuleService + RuleService *RuleService + PipelineRunService *PipelineRunService } -func NewIntegrationService(ruleService *RuleService) *IntegrationService { - return &IntegrationService{RuleService: ruleService} +func NewIntegrationService(ruleService *RuleService, pipelineRunService *PipelineRunService) *IntegrationService { + return &IntegrationService{RuleService: ruleService, PipelineRunService: pipelineRunService} } func (s *IntegrationService) ListRulesByDatabase(c context.Context, projectID uuid.UUID, databaseID uuid.UUID) ([]models.Rule, error) { return s.RuleService.ListRulesByDatabase(c, projectID, databaseID) } + +func (s *IntegrationService) CreatePipeLineRun(c context.Context, projectID uuid.UUID, databaseID uuid.UUID, dto dto.PipelineRunCreationDto) (models.PipelineRun, error) { + return s.PipelineRunService.CreatePipelineRun(c, projectID, databaseID, dto) +} diff --git a/skemr-api/internal/service/pipeline_run_service.go b/skemr-api/internal/service/pipeline_run_service.go index f3fc27a..771a439 100644 --- a/skemr-api/internal/service/pipeline_run_service.go +++ b/skemr-api/internal/service/pipeline_run_service.go @@ -97,19 +97,19 @@ func (s *PipelineRunService) GetPipelineRuns(c context.Context, projectId uuid.U return mapper.ToDomainPipelineRuns(pipelineRuns), nil } -func (s *PipelineRunService) CreatePipelineRun(c context.Context, projectId uuid.UUID, databaseId uuid.UUID, pipelineRun dto.PipelineRunCreationDto) (sqlc.PipelineRun, error) { +func (s *PipelineRunService) CreatePipelineRun(c context.Context, projectId uuid.UUID, databaseId uuid.UUID, pipelineRun dto.PipelineRunCreationDto) (models.PipelineRun, error) { slog.Info("Creating pipeline run", "pipelineRun", pipelineRun, "databaseId", databaseId, "projectId", projectId) _, err := s.scopeResolver.RequireDatabase(c, projectId, databaseId) if err != nil { slog.Error("Error fetching database", "err", err) - return sqlc.PipelineRun{}, err + return models.PipelineRun{}, err } completedAt, err := time.Parse(time.RFC3339, pipelineRun.CompletedAt) if err != nil { slog.Error("Error parsing completedAt", "err", err) - return sqlc.PipelineRun{}, err + return models.PipelineRun{}, err } createdPipelineRun, err := s.store.CreatePipelineRun(c, sqlc.CreatePipelineRunParams{ @@ -127,8 +127,8 @@ func (s *PipelineRunService) CreatePipelineRun(c context.Context, projectId uuid if err != nil { slog.Error("Error creating pipeline run", "err", err) - return sqlc.PipelineRun{}, err + return models.PipelineRun{}, err } - return createdPipelineRun, nil + return mapper.ToDomainPipelineRun(createdPipelineRun), nil } From 06b3fe68889c5157fa7a74caefc605f1cbe42e3e Mon Sep 17 00:00:00 2001 From: WalMaa Date: Mon, 25 May 2026 10:47:58 +0300 Subject: [PATCH 13/14] pipeline listing endpoints --- .../controller/pipeline_run_controller.go | 75 +++++++++++++++++++ skemr-api/internal/routers/router.go | 6 ++ 2 files changed, 81 insertions(+) create mode 100644 skemr-api/internal/controller/pipeline_run_controller.go diff --git a/skemr-api/internal/controller/pipeline_run_controller.go b/skemr-api/internal/controller/pipeline_run_controller.go new file mode 100644 index 0000000..08978a5 --- /dev/null +++ b/skemr-api/internal/controller/pipeline_run_controller.go @@ -0,0 +1,75 @@ +package controller + +import ( + "net/http" + + "github.com/go-chi/chi/v5" + "github.com/go-chi/render" + "github.com/google/uuid" + "github.com/walmaa/skemr-api/internal/errormsg" + "github.com/walmaa/skemr-api/internal/service" +) + +type PipelineRunController struct { + Service *service.PipelineRunService +} + +func NewPipelineRunController(s *service.PipelineRunService) *PipelineRunController { + return &PipelineRunController{Service: s} +} + +func (h *PipelineRunController) RegisterRoutes(r chi.Router) { + r.Route("/databases/{databaseId}/pipeline-runs", func(r chi.Router) { + r.Get("/", h.listPipelineRuns) + r.Get("/{pipelineRunId}", h.getPipelineRun) + }) +} + +func (h *PipelineRunController) listPipelineRuns(w http.ResponseWriter, r *http.Request) { + projectID, ok := r.Context().Value("projectId").(uuid.UUID) + if !ok { + http.Error(w, "projectId not found in context", http.StatusBadRequest) + return + } + databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) + if err != nil { + http.Error(w, "invalid databaseId", http.StatusBadRequest) + return + } + + pipelineRuns, err := h.Service.GetPipelineRuns(r.Context(), projectID, databaseId) + + if err != nil { + errormsg.WriteErrorResponse(w, r, err) + return + } + render.JSON(w, r, pipelineRuns) +} + +func (h *PipelineRunController) getPipelineRun(w http.ResponseWriter, r *http.Request) { + projectID, ok := r.Context().Value("projectId").(uuid.UUID) + if !ok { + http.Error(w, "projectId not found in context", http.StatusBadRequest) + return + } + databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) + if err != nil { + http.Error(w, "invalid databaseId", http.StatusBadRequest) + return + } + + pipelineRunId, err := uuid.Parse(chi.URLParam(r, "pipelineRunId")) + + if err != nil { + http.Error(w, "invalid pipelineRunId", http.StatusBadRequest) + return + } + + pipelineRun, err := h.Service.GetPipelineRun(r.Context(), projectID, databaseId, pipelineRunId) + if err != nil { + errormsg.WriteErrorResponse(w, r, err) + } + + render.JSON(w, r, pipelineRun) + +} diff --git a/skemr-api/internal/routers/router.go b/skemr-api/internal/routers/router.go index eef8d08..7e0c3dd 100644 --- a/skemr-api/internal/routers/router.go +++ b/skemr-api/internal/routers/router.go @@ -74,16 +74,22 @@ func InitRouter(services *Services) http.Handler { // Project level routes r.Route("/projects/{projectId}", func(r chi.Router) { r.Use(middleware.ProjectIDMiddleware) + + // define controllers databaseController := controller.NewDatabaseController(services.DatabaseService) projectSecretsController := controller.NewProjectSecretsController(services.AccessTokenService) ruleController := controller.NewRuleController(services.RuleService) databaseEntityController := controller.NewDatabaseEntityController(services.DatabaseEntityService) databaseChangeController := controller.NewDatabaseChangeController(services.DatabaseChangeService) + pipelineRunController := controller.NewPipelineRunController(services.PipelineRunService) + + // register routes databaseController.RegisterRoutes(r) projectSecretsController.RegisterRoutes(r) ruleController.RegisterRoutes(r) databaseEntityController.RegisterRoutes(r) databaseChangeController.RegisterRoutes(r) + pipelineRunController.RegisterRoutes(r) r.Get("/", projectController.GetProject) r.Delete("/", projectController.DeleteProject) From cb1396f191e431fc60d2ca9efe4814a40a8cecc6 Mon Sep 17 00:00:00 2001 From: WalMaa Date: Mon, 25 May 2026 11:11:22 +0300 Subject: [PATCH 14/14] streamlined url param parsing with a helper --- .../controller/database_change_controller.go | 29 ++++----- .../controller/database_controller.go | 60 ++++++++++------- .../database_entities_controller.go | 31 +++++---- .../controller/integration_controller.go | 43 +++--------- skemr-api/internal/controller/param_parser.go | 34 ++++++++++ .../controller/pipeline_run_controller.go | 22 +++---- .../project_access_tokens_controller.go | 25 ++++--- .../internal/controller/project_controller.go | 25 +++---- .../internal/controller/rule_controller.go | 65 +++++++++---------- .../middleware/access_token_middleware.go | 22 +------ .../middleware/project_id_middleware.go | 7 +- 11 files changed, 182 insertions(+), 181 deletions(-) create mode 100644 skemr-api/internal/controller/param_parser.go diff --git a/skemr-api/internal/controller/database_change_controller.go b/skemr-api/internal/controller/database_change_controller.go index 0a374a1..39cdd47 100644 --- a/skemr-api/internal/controller/database_change_controller.go +++ b/skemr-api/internal/controller/database_change_controller.go @@ -5,7 +5,6 @@ import ( "github.com/go-chi/chi/v5" "github.com/go-chi/render" - "github.com/google/uuid" "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/service" ) @@ -27,18 +26,17 @@ func (h *DatabaseChangeController) RegisterRoutes(r chi.Router) { } func (h *DatabaseChangeController) listDatabaseChanges(w http.ResponseWriter, r *http.Request) { - projectID, ok := r.Context().Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") if !ok { - http.Error(w, "projectId not found in context", http.StatusBadRequest) return } - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "invalid databaseId", http.StatusBadRequest) + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } - databaseChanges, err := h.Service.GetDatabaseChanges(r.Context(), projectID, databaseId, 100, 0) + databaseChanges, err := h.Service.GetDatabaseChanges(r.Context(), projectId, databaseId, 100, 0) if err != nil { errormsg.WriteErrorResponse(w, r, err) @@ -48,23 +46,22 @@ func (h *DatabaseChangeController) listDatabaseChanges(w http.ResponseWriter, r } func (h *DatabaseChangeController) GetDatabaseChange(w http.ResponseWriter, r *http.Request) { - projectID, ok := r.Context().Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") if !ok { - http.Error(w, "projectId not found in context", http.StatusBadRequest) return } - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "invalid databaseId", http.StatusBadRequest) + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } - databaseChangeId, err := uuid.Parse(chi.URLParam(r, "databaseChangeId")) - if err != nil { - http.Error(w, "invalid databaseChangeId", http.StatusBadRequest) + databaseChangeId, ok := ParseUUIDParam(w, r, "databaseChangeId") + if !ok { + return } - databaseChange, err := h.Service.GetDatabaseChange(r.Context(), projectID, databaseId, databaseChangeId) + databaseChange, err := h.Service.GetDatabaseChange(r.Context(), projectId, databaseId, databaseChangeId) if err != nil { errormsg.WriteErrorResponse(w, r, err) diff --git a/skemr-api/internal/controller/database_controller.go b/skemr-api/internal/controller/database_controller.go index 1a36cc8..bbb57c4 100644 --- a/skemr-api/internal/controller/database_controller.go +++ b/skemr-api/internal/controller/database_controller.go @@ -6,7 +6,6 @@ import ( "github.com/go-chi/chi/v5" "github.com/go-chi/render" - "github.com/google/uuid" "github.com/walmaa/skemr-api/internal/dto" "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/service" @@ -33,12 +32,13 @@ func (h *DatabaseController) RegisterRoutes(r chi.Router) { } func (h *DatabaseController) deleteDatabase(w http.ResponseWriter, r *http.Request) { - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "Invalid ID format", http.StatusBadRequest) + + // TODO: scope this to the project + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } - err = h.Service.DeleteDatabase(r.Context(), databaseId) + err := h.Service.DeleteDatabase(r.Context(), databaseId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return @@ -47,12 +47,12 @@ func (h *DatabaseController) deleteDatabase(w http.ResponseWriter, r *http.Reque } func (h *DatabaseController) listDatabasesByProject(w http.ResponseWriter, r *http.Request) { - id, err := uuid.Parse(chi.URLParam(r, "projectId")) - if err != nil { - http.Error(w, "Invalid project ID format", http.StatusBadRequest) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { return } - databases, err := h.Service.ListDatabasesByProject(r.Context(), id) + + databases, err := h.Service.ListDatabasesByProject(r.Context(), projectId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return @@ -61,14 +61,14 @@ func (h *DatabaseController) listDatabasesByProject(w http.ResponseWriter, r *ht } func (h *DatabaseController) createDatabase(w http.ResponseWriter, r *http.Request) { - projectId, err := uuid.Parse(chi.URLParam(r, "projectId")) - if err != nil { - http.Error(w, "Invalid project ID format", http.StatusBadRequest) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { return } + var body dto.DatabaseCreationDto - err = render.Decode(r, &body) + err := render.Decode(r, &body) if err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return @@ -91,12 +91,16 @@ func (h *DatabaseController) createDatabase(w http.ResponseWriter, r *http.Reque } func (h *DatabaseController) updateDatabase(w http.ResponseWriter, r *http.Request) { - projectId := r.Context().Value("projectId").(uuid.UUID) - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "Invalid database ID format", http.StatusBadRequest) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { + return + } + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } + var body dto.DatabaseUpdateDto if err := json.NewDecoder(r.Body).Decode(&body); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) @@ -111,13 +115,17 @@ func (h *DatabaseController) updateDatabase(w http.ResponseWriter, r *http.Reque } func (h *DatabaseController) syncDatabase(w http.ResponseWriter, r *http.Request) { - projectId := r.Context().Value("projectId").(uuid.UUID) - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "Invalid database ID format", http.StatusBadRequest) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { + return + } + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } - err = h.Service.EnqueueManualDatabaseSync(r.Context(), projectId, databaseId) + + err := h.Service.EnqueueManualDatabaseSync(r.Context(), projectId, databaseId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return @@ -126,11 +134,13 @@ func (h *DatabaseController) syncDatabase(w http.ResponseWriter, r *http.Request } func (h *DatabaseController) getDatabase(w http.ResponseWriter, r *http.Request) { - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "Invalid ID format", http.StatusBadRequest) + + // TODO: scope this to the project + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } + database, err := h.Service.GetDatabase(r.Context(), databaseId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) diff --git a/skemr-api/internal/controller/database_entities_controller.go b/skemr-api/internal/controller/database_entities_controller.go index f6fcc4d..9ef42b6 100644 --- a/skemr-api/internal/controller/database_entities_controller.go +++ b/skemr-api/internal/controller/database_entities_controller.go @@ -30,23 +30,22 @@ type Query struct { } func (h *DatabaseEntityController) GetDatabaseEntity(w http.ResponseWriter, r *http.Request) { - ctx := r.Context() - projectId, ok := ctx.Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") if !ok { - http.Error(w, "projectId not found in context", http.StatusBadRequest) return } - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "Invalid database ID format", http.StatusBadRequest) + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } - entityId, err := uuid.Parse(chi.URLParam(r, "entityId")) - if err != nil { - http.Error(w, "Invalid entity ID format", http.StatusBadRequest) + + entityId, ok := ParseUUIDParam(w, r, "entityId") + if !ok { return } - entity, err := h.Service.GetDatabaseEntityByID(ctx, projectId, databaseId, entityId) + + entity, err := h.Service.GetDatabaseEntityByID(r.Context(), projectId, databaseId, entityId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return @@ -57,13 +56,17 @@ func (h *DatabaseEntityController) GetDatabaseEntity(w http.ResponseWriter, r *h func (h *DatabaseEntityController) GetDatabaseEntities(w http.ResponseWriter, r *http.Request) { ctx := r.Context() - projectId := ctx.Value("projectId").(uuid.UUID) - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "Invalid database ID format", http.StatusBadRequest) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { return } + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { + return + } + entityTypeQuery := r.URL.Query().Get("type") var entityType *models.DatabaseEntityType if entityTypeQuery != "" { diff --git a/skemr-api/internal/controller/integration_controller.go b/skemr-api/internal/controller/integration_controller.go index 2966135..0c51e6d 100644 --- a/skemr-api/internal/controller/integration_controller.go +++ b/skemr-api/internal/controller/integration_controller.go @@ -5,7 +5,6 @@ import ( "github.com/go-chi/chi/v5" "github.com/go-chi/render" - "github.com/google/uuid" "github.com/walmaa/skemr-api/internal/dto" "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/service" @@ -28,25 +27,13 @@ func (h *IntegrationController) RegisterRoutes(r chi.Router) { } func (h *IntegrationController) listRulesByDatabase(w http.ResponseWriter, r *http.Request) { - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ - Message: "Invalid database ID format", - Status: http.StatusBadRequest, - Errors: nil, - }, - ) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { return } - projectId, err := uuid.Parse(chi.URLParam(r, "projectId")) - if err != nil { - errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ - Message: "Invalid project ID format", - Status: http.StatusBadRequest, - Errors: nil, - }, - ) + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } @@ -65,25 +52,13 @@ func (h *IntegrationController) listRulesByDatabase(w http.ResponseWriter, r *ht func (h *IntegrationController) createPipelineRun(w http.ResponseWriter, r *http.Request) { - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ - Message: "Invalid database ID format", - Status: http.StatusBadRequest, - Errors: nil, - }, - ) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { return } - projectId, err := uuid.Parse(chi.URLParam(r, "projectId")) - if err != nil { - errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ - Message: "Invalid project ID format", - Status: http.StatusBadRequest, - Errors: nil, - }, - ) + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } @@ -93,7 +68,7 @@ func (h *IntegrationController) createPipelineRun(w http.ResponseWriter, r *http return } - err = validation.Validate.Struct(req) + err := validation.Validate.Struct(req) if err != nil { errorResponse := validation.CreateErrorResponse(err) diff --git a/skemr-api/internal/controller/param_parser.go b/skemr-api/internal/controller/param_parser.go new file mode 100644 index 0000000..473104e --- /dev/null +++ b/skemr-api/internal/controller/param_parser.go @@ -0,0 +1,34 @@ +package controller + +import ( + "net/http" + + "github.com/go-chi/chi/v5" + "github.com/google/uuid" + "github.com/walmaa/skemr-api/internal/errormsg" + "github.com/walmaa/skemr-common/models" +) + +func ParseUUIDParam(w http.ResponseWriter, r *http.Request, name string) (uuid.UUID, bool) { + raw := chi.URLParam(r, name) + if raw == "" { + errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ + Message: "missing " + name, + Errors: nil, + Status: http.StatusBadRequest, + }) + return uuid.Nil, false + } + + id, err := uuid.Parse(raw) + if err != nil { + errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ + Message: "invalid " + name, + Errors: nil, + Status: http.StatusBadRequest, + }) + return uuid.Nil, false + } + + return id, true +} diff --git a/skemr-api/internal/controller/pipeline_run_controller.go b/skemr-api/internal/controller/pipeline_run_controller.go index 08978a5..6830dfd 100644 --- a/skemr-api/internal/controller/pipeline_run_controller.go +++ b/skemr-api/internal/controller/pipeline_run_controller.go @@ -26,18 +26,17 @@ func (h *PipelineRunController) RegisterRoutes(r chi.Router) { } func (h *PipelineRunController) listPipelineRuns(w http.ResponseWriter, r *http.Request) { - projectID, ok := r.Context().Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") if !ok { - http.Error(w, "projectId not found in context", http.StatusBadRequest) return } - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "invalid databaseId", http.StatusBadRequest) + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } - pipelineRuns, err := h.Service.GetPipelineRuns(r.Context(), projectID, databaseId) + pipelineRuns, err := h.Service.GetPipelineRuns(r.Context(), projectId, databaseId) if err != nil { errormsg.WriteErrorResponse(w, r, err) @@ -47,14 +46,13 @@ func (h *PipelineRunController) listPipelineRuns(w http.ResponseWriter, r *http. } func (h *PipelineRunController) getPipelineRun(w http.ResponseWriter, r *http.Request) { - projectID, ok := r.Context().Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") if !ok { - http.Error(w, "projectId not found in context", http.StatusBadRequest) return } - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "invalid databaseId", http.StatusBadRequest) + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } @@ -65,7 +63,7 @@ func (h *PipelineRunController) getPipelineRun(w http.ResponseWriter, r *http.Re return } - pipelineRun, err := h.Service.GetPipelineRun(r.Context(), projectID, databaseId, pipelineRunId) + pipelineRun, err := h.Service.GetPipelineRun(r.Context(), projectId, databaseId, pipelineRunId) if err != nil { errormsg.WriteErrorResponse(w, r, err) } diff --git a/skemr-api/internal/controller/project_access_tokens_controller.go b/skemr-api/internal/controller/project_access_tokens_controller.go index f86af12..996f267 100644 --- a/skemr-api/internal/controller/project_access_tokens_controller.go +++ b/skemr-api/internal/controller/project_access_tokens_controller.go @@ -33,8 +33,11 @@ func (h *ProjectSecretsController) RegisterRoutes(r chi.Router) { } func (h *ProjectSecretsController) createToken(w http.ResponseWriter, r *http.Request) { - c := r.Context() - projectId := c.Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { + return + } + var body dto.SecretCreationDto if err := render.Decode(r, &body); err != nil { http.Error(w, "Invalid request", http.StatusBadRequest) @@ -48,7 +51,7 @@ func (h *ProjectSecretsController) createToken(w http.ResponseWriter, r *http.Re return } - token, err := h.Service.CreateToken(c, projectId, body) + token, err := h.Service.CreateToken(r.Context(), projectId, body) if err != nil { errormsg.WriteErrorResponse(w, r, err) return @@ -72,10 +75,12 @@ func (h *ProjectSecretsController) getSecret(w http.ResponseWriter, r *http.Requ } func (h *ProjectSecretsController) getSecrets(w http.ResponseWriter, r *http.Request) { - c := r.Context() - projectId := c.Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { + return + } - tokens, err := h.Service.GetTokens(c, projectId) + tokens, err := h.Service.GetTokens(r.Context(), projectId) if err != nil { slog.Error("Error getting tokens", "err", err) return @@ -90,8 +95,10 @@ func (h *ProjectSecretsController) updateSecret(_ http.ResponseWriter, _ *http.R } func (h *ProjectSecretsController) deleteSecret(w http.ResponseWriter, r *http.Request) { - c := r.Context() - projectId := c.Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { + return + } secretId, err := uuid.Parse(chi.URLParam(r, "secretId")) @@ -99,7 +106,7 @@ func (h *ProjectSecretsController) deleteSecret(w http.ResponseWriter, r *http.R http.Error(w, "Invalid Secret ID", http.StatusBadRequest) } - err = h.Service.DeleteToken(c, projectId, secretId) + err = h.Service.DeleteToken(r.Context(), projectId, secretId) if err != nil { slog.Error("Error deleting token", "err", err) http.Error(w, "Error deleting token", http.StatusInternalServerError) diff --git a/skemr-api/internal/controller/project_controller.go b/skemr-api/internal/controller/project_controller.go index 539c38a..3935c94 100644 --- a/skemr-api/internal/controller/project_controller.go +++ b/skemr-api/internal/controller/project_controller.go @@ -4,14 +4,11 @@ import ( "encoding/json" "net/http" - "github.com/go-chi/chi/v5" "github.com/go-chi/render" - "github.com/google/uuid" "github.com/walmaa/skemr-api/internal/dto" "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/service" "github.com/walmaa/skemr-api/internal/validation" - "github.com/walmaa/skemr-common/models" ) type ProjectController struct { @@ -60,15 +57,12 @@ func (h *ProjectController) CreateProject(w http.ResponseWriter, r *http.Request } func (h *ProjectController) GetProject(w http.ResponseWriter, r *http.Request) { - projectID, err := uuid.Parse(chi.URLParam(r, "projectId")) - if err != nil { - errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ - Message: errormsg.ErrInvalidIdFormat, - Status: http.StatusBadRequest, - }) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { return } - project, err := h.Service.GetProject(r.Context(), projectID) + + project, err := h.Service.GetProject(r.Context(), projectId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return @@ -79,15 +73,12 @@ func (h *ProjectController) GetProject(w http.ResponseWriter, r *http.Request) { } func (h *ProjectController) DeleteProject(w http.ResponseWriter, r *http.Request) { - projectID, err := uuid.Parse(chi.URLParam(r, "projectId")) - if err != nil { - errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ - Message: errormsg.ErrInvalidIdFormat, - Status: http.StatusBadRequest, - }) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { return } - err = h.Service.DeleteProject(r.Context(), projectID) + + err := h.Service.DeleteProject(r.Context(), projectId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return diff --git a/skemr-api/internal/controller/rule_controller.go b/skemr-api/internal/controller/rule_controller.go index e7ce54c..a837596 100644 --- a/skemr-api/internal/controller/rule_controller.go +++ b/skemr-api/internal/controller/rule_controller.go @@ -1,12 +1,10 @@ package controller import ( - "fmt" "net/http" "github.com/go-chi/chi/v5" "github.com/go-chi/render" - "github.com/google/uuid" "github.com/walmaa/skemr-api/internal/dto" "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/service" @@ -31,22 +29,22 @@ func (h *RuleController) RegisterRoutes(r chi.Router) { } func (h *RuleController) GetRule(w http.ResponseWriter, r *http.Request) { - projectID, ok := r.Context().Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") if !ok { - http.Error(w, "projectId not found in context", http.StatusBadRequest) return } - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "invalid databaseId", http.StatusBadRequest) + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } - ruleId, err := uuid.Parse(chi.URLParam(r, "ruleId")) - if err != nil { - http.Error(w, "invalid ruleId", http.StatusBadRequest) + + ruleId, ok := ParseUUIDParam(w, r, "ruleId") + if !ok { return } - rule, err := h.Service.GetRule(r.Context(), projectID, databaseId, ruleId) + + rule, err := h.Service.GetRule(r.Context(), projectId, databaseId, ruleId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return @@ -55,22 +53,22 @@ func (h *RuleController) GetRule(w http.ResponseWriter, r *http.Request) { } func (h *RuleController) deleteRule(w http.ResponseWriter, r *http.Request) { - projectID, ok := r.Context().Value("projectId").(uuid.UUID) + projectId, ok := ParseUUIDParam(w, r, "projectId") if !ok { - http.Error(w, "projectId not found in context", http.StatusBadRequest) return } - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "invalid databaseId", http.StatusBadRequest) + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") + if !ok { return } - ruleId, err := uuid.Parse(chi.URLParam(r, "ruleId")) - if err != nil { - http.Error(w, "invalid ruleId", http.StatusBadRequest) + + ruleId, ok := ParseUUIDParam(w, r, "ruleId") + if !ok { return } - err = h.Service.DeleteRule(r.Context(), projectID, databaseId, ruleId) + + err := h.Service.DeleteRule(r.Context(), projectId, databaseId, ruleId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return @@ -79,17 +77,16 @@ func (h *RuleController) deleteRule(w http.ResponseWriter, r *http.Request) { } func (h *RuleController) createRule(w http.ResponseWriter, r *http.Request) { - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "invalid databaseId", http.StatusBadRequest) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { return } - projectID, ok := r.Context().Value("projectId").(uuid.UUID) + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") if !ok { - err = fmt.Errorf("projectId not found in context") - errormsg.WriteErrorResponse(w, r, err) return } + var body dto.RuleCreationDto if err := render.Decode(r, &body); err != nil { @@ -97,7 +94,7 @@ func (h *RuleController) createRule(w http.ResponseWriter, r *http.Request) { return } - err = validation.Validate.Struct(body) + err := validation.Validate.Struct(body) if err != nil { errorResponse := validation.CreateErrorResponse(err) @@ -105,7 +102,7 @@ func (h *RuleController) createRule(w http.ResponseWriter, r *http.Request) { return } - rule, err := h.Service.CreateRule(r.Context(), projectID, databaseId, body) + rule, err := h.Service.CreateRule(r.Context(), projectId, databaseId, body) if err != nil { errormsg.WriteErrorResponse(w, r, err) @@ -117,17 +114,17 @@ func (h *RuleController) createRule(w http.ResponseWriter, r *http.Request) { } func (h *RuleController) ListRules(w http.ResponseWriter, r *http.Request) { - databaseId, err := uuid.Parse(chi.URLParam(r, "databaseId")) - if err != nil { - http.Error(w, "invalid databaseId", http.StatusBadRequest) + projectId, ok := ParseUUIDParam(w, r, "projectId") + if !ok { return } - projectID, ok := r.Context().Value("projectId").(uuid.UUID) + + databaseId, ok := ParseUUIDParam(w, r, "databaseId") if !ok { - http.Error(w, "projectId not found in context", http.StatusBadRequest) return } - rules, err := h.Service.ListRulesByDatabase(r.Context(), projectID, databaseId) + + rules, err := h.Service.ListRulesByDatabase(r.Context(), projectId, databaseId) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return diff --git a/skemr-api/internal/middleware/access_token_middleware.go b/skemr-api/internal/middleware/access_token_middleware.go index 73db030..3ad274d 100644 --- a/skemr-api/internal/middleware/access_token_middleware.go +++ b/skemr-api/internal/middleware/access_token_middleware.go @@ -5,8 +5,8 @@ import ( "net/http" "strings" - "github.com/go-chi/chi/v5" "github.com/google/uuid" + "github.com/walmaa/skemr-api/internal/controller" "github.com/walmaa/skemr-api/internal/errormsg" "github.com/walmaa/skemr-api/internal/service" "github.com/walmaa/skemr-common/models" @@ -16,25 +16,9 @@ func AccessTokenMiddleware(service *service.AccessTokenService) func(next http.H return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { c := r.Context() - projectIdParam := chi.URLParam(r, "projectId") - if projectIdParam == "" { - errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ - Message: "projectIdParam is required", - Errors: nil, - Status: http.StatusBadRequest, - }) - return - } - - projectId, err := uuid.Parse(projectIdParam) - - if err != nil { - errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ - Message: "Invalid projectId format", - Errors: nil, - Status: http.StatusBadRequest, - }) + projectId, ok := controller.ParseUUIDParam(w, r, "projectId") + if !ok { return } diff --git a/skemr-api/internal/middleware/project_id_middleware.go b/skemr-api/internal/middleware/project_id_middleware.go index 40a5938..9b7bc12 100644 --- a/skemr-api/internal/middleware/project_id_middleware.go +++ b/skemr-api/internal/middleware/project_id_middleware.go @@ -6,6 +6,8 @@ import ( "github.com/go-chi/chi/v5" "github.com/google/uuid" + "github.com/walmaa/skemr-api/internal/errormsg" + "github.com/walmaa/skemr-common/models" ) const CtxProjectID = "projectId" @@ -19,7 +21,10 @@ func ProjectIDMiddleware(next http.Handler) http.Handler { } id, err := uuid.Parse(param) if err != nil { - http.Error(w, "invalid projectId", http.StatusBadRequest) + errormsg.WriteErrorResponse(w, r, &models.ErrorResponse{ + Status: http.StatusBadRequest, + Message: "Invalid project id", + }) return } ctx := context.WithValue(r.Context(), CtxProjectID, id)