diff --git a/backend/cmd/server/main.go b/backend/cmd/server/main.go index a04853e5..debbdc29 100644 --- a/backend/cmd/server/main.go +++ b/backend/cmd/server/main.go @@ -1199,16 +1199,20 @@ func main() { // Initialize clean architecture risk module riskRepo := repository.NewGormRiskRepository(database.DB) riskControlMappingRepo := repository.NewGormRiskControlMappingRepository(database.DB) + riskAssetStore := repository.NewGormRiskAssetStore(database.DB) createRiskUseCase := risk.NewCreateRiskUseCase(riskRepo). WithActivation(activationRecorder). - WithOwnership(ownershipService) + WithOwnership(ownershipService). + WithAssets(riskAssetStore) getRiskUseCase := risk.NewGetRiskUseCase(riskRepo). WithMappings(riskControlMappingRepo). WithOwnership(ownershipService) listRisksUseCase := risk.NewListRisksUseCase(riskRepo). WithMappings(riskControlMappingRepo). WithOwnership(ownershipService) - updateRiskUseCase := risk.NewUpdateRiskUseCase(riskRepo).WithOwnership(ownershipService) + updateRiskUseCase := risk.NewUpdateRiskUseCase(riskRepo). + WithOwnership(ownershipService). + WithAssets(riskAssetStore) deleteRiskUseCase := risk.NewDeleteRiskUseCase(riskRepo) // Cyber Risk Quantification: XAF→USD rate configurable via XAF_USD_RATE // (default ≈ 600 FCFA/USD). Reference ALE bands match the board ExposureModel. @@ -1247,7 +1251,20 @@ func main() { // implementation accepted a performedBy and discarded it, so a supervisor // asking "who reassigned these and when" had no answer. auditChainRepo is // the same hash-chained, append-only store the rest of the trail uses. - WithBulkAction(risk.NewBulkActionUseCase(riskRepo, auditChainRepo)) + WithBulkAction(risk.NewBulkActionUseCase(riskRepo, auditChainRepo)). + // #755 — CSV import: every row validated first, then all of them written + // in one transaction through CreateRiskUseCase, or none. The plan cap is + // checked against the whole file, not just the first row. + WithImport(risk.NewImportRisksUseCase(repository.RunRiskTx(database.DB)). + WithAssets(repository.ListImportAssetRefs(database.DB)). + WithActivation(activationRecorder). + WithCapacity(func(ctx context.Context, tenant uuid.UUID) (int, error) { + _, limit, used, _, err := entitlementService.Capacity(ctx, tenant, ent.LimitRisks) + if err != nil || limit == ent.Unlimited || used < 0 { + return -1, err + } + return max(limit-used, 0), nil + })) // Financial Risk Quantification (spec §9): tenant-wide CFO/CISO dashboard // (portfolio FAIR-lite P10/P50/P90, ALE, worst-case, residual, remediation @@ -1318,6 +1335,9 @@ func main() { protected.Post("/risks/bulk", riskUpdate, riskHandler.BulkAction) protected.Post("/risks", riskCreate, capRisks, riskHandler.CreateRisk) + // #755 — CSV import. capRisks refuses a tenant already at its limit; the use + // case then refuses a file that would carry it past. + protected.Post("/risks/import", riskCreate, capRisks, riskHandler.ImportRisks) protected.Patch("/risks/:id", riskUpdate, riskHandler.UpdateRisk) protected.Post("/risks/:id/review", riskUpdate, riskHandler.MarkReviewed) protected.Post("/risks/:id/transfer-owner", riskUpdate, ownershipTransferHandler.TransferRiskOwner) diff --git a/backend/internal/api/http/handlers/risks.go b/backend/internal/api/http/handlers/risks.go index 2ffecbd1..27d1674d 100644 --- a/backend/internal/api/http/handlers/risks.go +++ b/backend/internal/api/http/handlers/risks.go @@ -517,8 +517,6 @@ func (h *RiskHandler) ImportRisks(c *fiber.Ctx) error { return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": "file is required"}) } - format := risk.ImportFormat(c.FormValue("format", "json")) - // Read file content openFile, err := file.Open() if err != nil { @@ -533,7 +531,7 @@ func (h *RiskHandler) ImportRisks(c *fiber.Ctx) error { } // Execute import - result, err := h.importUC.Execute(c.Context(), tenantID, buffer, format, userID) + result, err := h.importUC.Execute(c.Context(), tenantID, risk.ImportRisksInput{CSV: buffer, ImportedBy: userID}) if err != nil { h.logger.Error().Err(err).Msg("failed to import risks") return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to import risks"}) diff --git a/backend/internal/application/risk/create_risk.go b/backend/internal/application/risk/create_risk.go index 7ebc4dc0..96fea256 100644 --- a/backend/internal/application/risk/create_risk.go +++ b/backend/internal/application/risk/create_risk.go @@ -11,6 +11,7 @@ import ( "github.com/google/uuid" "github.com/opendefender/openrisk/internal/domain" + pkgscoring "github.com/opendefender/openrisk/pkg/scoring" ) // CreateRiskInput represents the input for creating a risk. @@ -32,6 +33,11 @@ type CreateRiskInput struct { // may be unclassified, and forcing a pick at creation only teaches people to // choose the first entry. CategoryID *uuid.UUID + // AssetIDs links the risk to the tenant's assets at creation. Their + // criticality is a term of the score, so they are linked before the score + // is computed and in the same write. Ids that are not the tenant's assets + // are dropped, as the handler always did. + AssetIDs []uuid.UUID Source string // parsed into domain.RiskSource in Execute() ExternalID string CreatedBy uuid.UUID // the authenticated user creating the risk @@ -62,11 +68,20 @@ type CreateRiskUseCase struct { riskRepo domain.RiskRepository activation ActivationRecorder ownership OwnershipManager + assets RiskAssetStore + engine pkgscoring.Engine } // NewCreateRiskUseCase creates a new CreateRiskUseCase. func NewCreateRiskUseCase(riskRepo domain.RiskRepository) *CreateRiskUseCase { - return &CreateRiskUseCase{riskRepo: riskRepo} + return &CreateRiskUseCase{riskRepo: riskRepo, engine: pkgscoring.NewEngine()} +} + +// WithAssets attaches the asset store that resolves and links input.AssetIDs. +// Without it, AssetIDs is ignored and the risk is scored with no asset. +func (uc *CreateRiskUseCase) WithAssets(s RiskAssetStore) *CreateRiskUseCase { + uc.assets = s + return uc } // WithActivation attaches the optional activation recorder. Nil-safe. @@ -156,20 +171,34 @@ func (uc *CreateRiskUseCase) Execute(ctx context.Context, orgID uuid.UUID, input risk.AssignedTo = risk.AssigneeID } - // 3. Compute score (Claude.md formula: P × I, score engine can override later) - risk.Score = risk.Impact * risk.Probability - // Band the score synchronously so the create response is self-consistent - // (score and criticality agree) instead of returning the default 'low' until - // the async ScoreWorker runs ~2s later. The worker refines it once asset - // criticality is folded in; both move together (audit-2026 #246). - risk.Criticality = domain.CriticalityFromScore(risk.Score) + // 3. Resolve the linked assets: their criticality is a term of the score. + var linked []*domain.Asset + if uc.assets != nil && len(input.AssetIDs) > 0 { + linked, err = uc.assets.FindByIDs(ctx, orgID, input.AssetIDs) + if err != nil { + return nil, domain.NewInternalError(fmt.Sprintf("failed to resolve assets: %v", err)) + } + risk.Assets = linked + } + + // 4. Score through the Score Engine, asset criticality included, so the + // stored score is the formula's from the first write — not P × I waiting + // for a Redis event that may never come (#792). + if err := applyScore(uc.engine, risk); err != nil { + return nil, err + } - // 4. Persist - if err := uc.riskRepo.Create(ctx, risk); err != nil { + // 5. Persist the risk and its asset links together. + if len(linked) > 0 { + err = uc.assets.SaveWithAssets(ctx, risk, linked, true) + } else { + err = uc.riskRepo.Create(ctx, risk) + } + if err != nil { return nil, domain.NewInternalError(fmt.Sprintf("failed to create risk: %v", err)) } - // 5. Note the activation milestone. Every creation records an event; only the + // 6. Note the activation milestone. Every creation records an event; only the // FIRST one ticks the checklist (the read model takes MIN(occurred_at)), so // no counting or de-duplication is needed here. if uc.activation != nil { @@ -179,7 +208,7 @@ func (uc *CreateRiskUseCase) Execute(ctx context.Context, orgID uuid.UUID, input }) } - // 6. Announce assignments made at creation (never to the creator themselves). + // 7. Announce assignments made at creation (never to the creator themselves). if uc.ownership != nil && len(ownershipChanges) > 0 { uc.ownership.Notify(ctx, orgID, ownershipChanges, domain.OwnershipSubject{ ResourceType: "risk", diff --git a/backend/internal/application/risk/get_score_breakdown.go b/backend/internal/application/risk/get_score_breakdown.go index feccced4..5b78d3cd 100644 --- a/backend/internal/application/risk/get_score_breakdown.go +++ b/backend/internal/application/risk/get_score_breakdown.go @@ -44,18 +44,10 @@ func (uc *GetScoreBreakdownUseCase) Execute(ctx context.Context, tenantID uuid.U return nil, domain.NewNotFoundError("risk", riskID) } - // 2. Calculate asset criticality — average domain.AssetCriticality.ScoreFactor() - // across every linked asset (not just the first one), consistent with how - // GormRiskRepository.GetRisksByAssetID/RiskHandler now derive it. Defaults - // to MEDIUM's factor (1.5) if no asset is linked. - assetCriticality := domain.CriticalityMedium.ScoreFactor() - if len(risk.Assets) > 0 { - var sum float64 - for _, a := range risk.Assets { - sum += a.Criticality.ScoreFactor() - } - assetCriticality = sum / float64(len(risk.Assets)) - } + // 2. Asset criticality — the one derivation every score writer uses + // (domain.RiskAssetCriticality): average factor of the linked assets, + // neutral when none is linked (#792). + assetCriticality := domain.RiskAssetCriticality(domain.AssetCriticalities(risk.Assets)) // 3. Use Score Engine to compute breakdown // IMPORTANT: All score calculations go through the Score Engine diff --git a/backend/internal/application/risk/get_score_breakdown_test.go b/backend/internal/application/risk/get_score_breakdown_test.go index dcb785c1..47a7e00e 100644 --- a/backend/internal/application/risk/get_score_breakdown_test.go +++ b/backend/internal/application/risk/get_score_breakdown_test.go @@ -29,9 +29,10 @@ func TestGetScoreBreakdown_Success_NoLinkedAssets(t *testing.T) { breakdown, err := uc.Execute(context.Background(), tenantID, riskID) require.NoError(t, err) - // No linked assets → defaults to MEDIUM's factor (1.5): 0.5 * 8.0 * 1.5 = 6.0 - assert.Equal(t, 6.0, breakdown.Score) - assert.Equal(t, 1.5, breakdown.AssetCriticality) + // No linked assets → the neutral factor every score writer stores (#792): + // 0.5 * 8.0 * 1.0 = 4.0 + assert.Equal(t, 4.0, breakdown.Score) + assert.Equal(t, domain.NoAssetCriticalityFactor, breakdown.AssetCriticality) } func TestGetScoreBreakdown_AveragesAcrossAllLinkedAssets(t *testing.T) { diff --git a/backend/internal/application/risk/get_score_working.go b/backend/internal/application/risk/get_score_working.go index c9052876..66a58ec8 100644 --- a/backend/internal/application/risk/get_score_working.go +++ b/backend/internal/application/risk/get_score_working.go @@ -76,26 +76,19 @@ func (uc *GetScoreWorkingUseCase) Execute(ctx context.Context, tenantID, riskID SourcesVisible: canReadAudit && uc.audit != nil, } - // Asset criticality: the average factor over every linked asset, medium when - // none — the same derivation the Score Engine is fed (GetScoreBreakdown, - // GormRiskRepository.GetRisksByAssetID). - ac := domain.CriticalityMedium.ScoreFactor() - if len(r.Assets) > 0 { - var sum float64 - for _, a := range r.Assets { - f := a.Criticality.ScoreFactor() - sum += f - w.Assets = append(w.Assets, domain.ScoreWorkingAsset{ - ID: a.ID.String(), - Name: a.Name, - Criticality: string(a.Criticality), - Factor: f, - }) - } - ac = sum / float64(len(r.Assets)) - } else { - w.AssetCriticalityDefaulted = true + // Asset criticality: the same derivation every score writer uses + // (domain.RiskAssetCriticality) — the average factor over the linked + // assets, neutral when none is linked (#792). + for _, a := range r.Assets { + w.Assets = append(w.Assets, domain.ScoreWorkingAsset{ + ID: a.ID.String(), + Name: a.Name, + Criticality: string(a.Criticality), + Factor: a.Criticality.ScoreFactor(), + }) } + w.AssetCriticalityDefaulted = len(r.Assets) == 0 + ac := domain.RiskAssetCriticality(domain.AssetCriticalities(r.Assets)) b, err := uc.engine.Breakdown(r.Probability, r.Impact, ac, nil) if err != nil { diff --git a/backend/internal/application/risk/get_score_working_test.go b/backend/internal/application/risk/get_score_working_test.go index 0a02439a..ce7a5809 100644 --- a/backend/internal/application/risk/get_score_working_test.go +++ b/backend/internal/application/risk/get_score_working_test.go @@ -149,14 +149,14 @@ func TestGetScoreWorking_InconsistentStoredScoreIsSaidSo(t *testing.T) { assert.InDelta(t, 6.0, w.Computed, 1e-9) } -func TestGetScoreWorking_NoAssetUsesTheDocumentedDefault(t *testing.T) { +func TestGetScoreWorking_NoAssetUsesTheNeutralFactor(t *testing.T) { tenant, r, trail, _ := scoreWorkingFixture(t) r.Assets = nil w, err := NewGetScoreWorkingUseCase(repoReturning(r), trail, pkgscoring.NewEngine()). Execute(context.Background(), tenant, r.ID, true) require.NoError(t, err) assert.True(t, w.AssetCriticalityDefaulted) - assert.InDelta(t, domain.CriticalityMedium.ScoreFactor(), w.Terms[2].Value, 1e-9) + assert.InDelta(t, domain.NoAssetCriticalityFactor, w.Terms[2].Value, 1e-9) assert.Nil(t, w.Terms[2].Source) } diff --git a/backend/internal/application/risk/import_risks.go b/backend/internal/application/risk/import_risks.go index 9e806762..0ea0c8f3 100644 --- a/backend/internal/application/risk/import_risks.go +++ b/backend/internal/application/risk/import_risks.go @@ -8,184 +8,562 @@ package risk import ( "bytes" "context" - "encoding/json" + "encoding/csv" + "errors" "fmt" + "io" + "math" + "sort" + "strconv" + "strings" + "unicode/utf8" "github.com/google/uuid" "github.com/opendefender/openrisk/internal/domain" ) -// ImportFormat represents the supported import formats -type ImportFormat string +// --------------------------------------------------------------------------- +// CSV import of the risk register (#755). +// +// All-or-nothing, like bulk actions (D-036). Every row is validated before +// anything is written; one bad row means nothing is persisted and the caller +// gets the line and column of EVERY error, not just the first. Valid files are +// written inside ONE transaction, and each row goes through CreateRiskUseCase so +// an imported risk is born exactly like one typed into the form: same +// validation, same lifecycle entry, same owner fallback, same initial score. +// +// The scales are the product's: probability in [0,1], impact in [0,10]. The +// template this page used to hand out was on a 1–5 scale for both, and reading +// such a file as-is would produce plausible-looking but wrong scores. A file +// that looks like that scale is refused with a message saying so, never +// converted behind the user's back. +// +// This replaces a use case whose CSV parser returned zero rows and reported +// success, and that no route ever reached. +// --------------------------------------------------------------------------- const ( - ImportFormatCSV ImportFormat = "csv" - ImportFormatJSON ImportFormat = "json" - ImportFormatXLSX ImportFormat = "xlsx" + // MaxImportRows caps one file. Synchronous import inside one transaction is + // adequate at this size; a larger register is several files. + MaxImportRows = 1000 + // MaxImportBytes caps the upload. 1000 rows of long descriptions fit well + // within it. + MaxImportBytes = 2 << 20 ) -// ImportRiskItem represents a single risk item in import data -type ImportRiskItem struct { - Name string `json:"name"` - Description string `json:"description"` - Probability float64 `json:"probability"` - Impact float64 `json:"impact"` - Status string `json:"status"` - Tags []string `json:"tags"` - Frameworks []string `json:"frameworks"` - AssetID *string `json:"asset_id,omitempty"` - Criticality string `json:"criticality"` - Source string `json:"source"` -} - -// ImportResult represents the outcome of an import operation -type ImportResult struct { - Total int `json:"total"` - Succeeded int `json:"succeeded"` - Failed int `json:"failed"` - Created []uuid.UUID `json:"created"` - Errors []ImportError `json:"errors"` - Duplicates []ImportDuplicateWarning `json:"duplicates"` -} - -// ImportError represents an error during import -type ImportError struct { - Row int `json:"row"` - Reason string `json:"reason"` -} - -// ImportDuplicateWarning represents a duplicate detected during import -type ImportDuplicateWarning struct { - Row int `json:"row"` - Name string `json:"name"` - ExistingID uuid.UUID `json:"existing_id"` - Action string `json:"action"` // "skipped" or "overwritten" -} - -// ImportRisksUseCase handles importing risks from file -// ABSOLUTE: Import must be idempotent (same file imported twice should not create duplicates) +// Import column names, lower-case. "name" is accepted as an alias of "title" +// because the register export and older files use it. +const ( + importColTitle = "title" + importColDescription = "description" + importColProbability = "probability" + importColImpact = "impact" + importColTags = "tags" + importColFrameworks = "frameworks" + importColAssets = "assets" +) + +var importColumnAliases = map[string]string{ + "title": importColTitle, + "name": importColTitle, + "description": importColDescription, + "probability": importColProbability, + "impact": importColImpact, + "tags": importColTags, + "frameworks": importColFrameworks, + "framework": importColFrameworks, + "assets": importColAssets, + "asset": importColAssets, +} + +// ImportColumns is the accepted header, in template order. +var ImportColumns = []string{ + importColTitle, importColDescription, importColProbability, + importColImpact, importColTags, importColFrameworks, importColAssets, +} + +// RiskTxRunner runs fn inside one database transaction and hands it a risk +// repository bound to that transaction. An error from fn rolls everything back. +// The asset store is bound to the same transaction, so a risk and its asset +// links are written or rolled back together. +type RiskTxRunner func(ctx context.Context, fn func(repo domain.RiskRepository, assets RiskAssetStore) error) error + +// ImportCapacity reports how many more risks the tenant's plan allows. +// A negative value means unlimited. +type ImportCapacity func(ctx context.Context, tenantID uuid.UUID) (remaining int, err error) + +// ImportAssetRef is one of the tenant's assets as the "assets" column can name +// it: by name (case-insensitive) or by id. +type ImportAssetRef struct { + ID uuid.UUID + Name string +} + +// ImportAssetLister lists the tenant's live assets, tenant-scoped, so the +// "assets" column can be resolved before anything is written. +type ImportAssetLister func(ctx context.Context, tenantID uuid.UUID) ([]ImportAssetRef, error) + +// ImportRisksInput is one uploaded file. +type ImportRisksInput struct { + CSV []byte + ImportedBy uuid.UUID +} + +// ImportRowError locates one problem in the file. Line is the 1-based line in +// the file as a spreadsheet shows it (the header is line 1); 0 means the +// problem concerns the file as a whole. Column is the header name, empty when +// the problem is not about one cell. +// +// Code and Params are the contract: the page renders them in the reader's +// language. Message is the English rendering, for API callers and as the +// page's fallback when it does not know a code. +type ImportRowError struct { + Line int `json:"line"` + Column string `json:"column,omitempty"` + Code string `json:"code"` + Params map[string]string `json:"params,omitempty"` + Message string `json:"message"` +} + +// ImportRisksResult is what the caller is told. Under all-or-nothing, either +// Created is the number of data rows and Errors is empty, or Created is 0 and +// Rejected counts the rows that had at least one error. +type ImportRisksResult struct { + Created int `json:"created"` + Rejected int `json:"rejected"` + RiskIDs []uuid.UUID `json:"risk_ids"` + Errors []ImportRowError `json:"errors"` + // Risks are the created entities, for the caller's post-commit work + // (score events). Not serialised. + Risks []*domain.Risk `json:"-"` +} + +// ImportRejectedError carries a refused file's result. It matches +// domain.ErrValidation so generic handling still answers 4xx. +type ImportRejectedError struct { + Result *ImportRisksResult +} + +func (e *ImportRejectedError) Error() string { + return fmt.Sprintf("import rejected: %d error(s), nothing was imported", len(e.Result.Errors)) +} + +func (e *ImportRejectedError) Unwrap() error { return domain.ErrValidation } + +// ImportOverCapacityError means the file is valid but would take the tenant +// past its plan's risk limit. Nothing is written. +type ImportOverCapacityError struct { + Requested int + Remaining int +} + +func (e *ImportOverCapacityError) Error() string { + return fmt.Sprintf("import of %d risks exceeds the plan: %d more allowed", e.Requested, e.Remaining) +} + +func (e *ImportOverCapacityError) Unwrap() error { return domain.ErrForbidden } + +// ImportRisksUseCase imports a CSV file into the tenant's register. type ImportRisksUseCase struct { - riskRepo domain.RiskRepository + inTx RiskTxRunner + capacity ImportCapacity + activation ActivationRecorder + listAssets ImportAssetLister } -// NewImportRisksUseCase creates a new ImportRisksUseCase -func NewImportRisksUseCase(riskRepo domain.RiskRepository) *ImportRisksUseCase { - return &ImportRisksUseCase{riskRepo: riskRepo} +// NewImportRisksUseCase builds the use case over a transaction runner. +func NewImportRisksUseCase(inTx RiskTxRunner) *ImportRisksUseCase { + return &ImportRisksUseCase{inTx: inTx} } -// Execute imports risks from file data -// Format can be CSV, JSON, or XLSX -// Returns import result with details of successes and failures -func (uc *ImportRisksUseCase) Execute( - ctx context.Context, - tenantID uuid.UUID, - fileContent []byte, - format ImportFormat, - importedBy uuid.UUID, -) (*ImportResult, error) { - result := &ImportResult{ - Created: []uuid.UUID{}, - Errors: []ImportError{}, - Duplicates: []ImportDuplicateWarning{}, +// WithCapacity attaches the plan-limit check. Nil-safe: without it there is no +// limit beyond MaxImportRows. +func (uc *ImportRisksUseCase) WithCapacity(c ImportCapacity) *ImportRisksUseCase { + uc.capacity = c + return uc +} + +// WithAssets attaches the lister that resolves the "assets" column. Without +// it, a file that fills that column is refused rather than imported unlinked. +func (uc *ImportRisksUseCase) WithAssets(l ImportAssetLister) *ImportRisksUseCase { + uc.listAssets = l + return uc +} + +// WithActivation attaches the activation recorder, fed after commit only. +func (uc *ImportRisksUseCase) WithActivation(rec ActivationRecorder) *ImportRisksUseCase { + uc.activation = rec + return uc +} + +// importRow is one parsed, validated data row. +type importRow struct { + line int + input CreateRiskInput + assetRefs []string + // invalid rows are kept only so their assets are checked too, and every + // error in the file is reported in one pass. + invalid bool +} + +// Execute validates the whole file, then creates every row in one transaction. +func (uc *ImportRisksUseCase) Execute(ctx context.Context, tenantID uuid.UUID, input ImportRisksInput) (*ImportRisksResult, error) { + if tenantID == uuid.Nil { + return nil, domain.NewUnauthorizedError("tenant is required") + } + if len(input.CSV) > MaxImportBytes { + return nil, domain.NewValidationError(fmt.Sprintf("file is larger than %d MB", MaxImportBytes>>20)) } - // 1. Parse file based on format - var items []ImportRiskItem - var err error + parsed, errs := parseImportCSV(input.CSV, input.ImportedBy) + assetErrs, err := uc.resolveAssets(ctx, tenantID, parsed) + if err != nil { + return nil, domain.NewInternalError(fmt.Sprintf("failed to list assets: %v", err)) + } + errs = append(errs, assetErrs...) + if len(errs) > 0 { + return nil, &ImportRejectedError{Result: rejected(errs)} + } + rows := parsed - switch format { - case ImportFormatJSON: - items, err = uc.parseJSON(fileContent) - case ImportFormatCSV: - items, err = uc.parseCSV(fileContent) - case ImportFormatXLSX: - items, err = uc.parseXLSX(fileContent) - default: - return nil, domain.NewValidationError(fmt.Sprintf("unsupported import format: %s", format)) + if uc.capacity != nil { + remaining, err := uc.capacity(ctx, tenantID) + // A counting error fails open, as the per-create middleware does. + if err == nil && remaining >= 0 && len(rows) > remaining { + return nil, &ImportOverCapacityError{Requested: len(rows), Remaining: remaining} + } } + created := make([]*domain.Risk, 0, len(rows)) + err = uc.inTx(ctx, func(repo domain.RiskRepository, assets RiskAssetStore) error { + create := NewCreateRiskUseCase(repo).WithAssets(assets) + for _, row := range rows { + r, err := create.Execute(ctx, tenantID, row.input) + if err != nil { + var appErr *domain.AppError + if errors.As(err, &appErr) && errors.Is(err, domain.ErrValidation) { + return &ImportRejectedError{Result: rejected([]ImportRowError{{Line: row.line, Code: "rejected_by_rules", Message: appErr.Message}})} + } + return err + } + created = append(created, r) + } + return nil + }) if err != nil { - return nil, domain.NewInternalError(fmt.Sprintf("failed to parse import file: %v", err)) + var rej *ImportRejectedError + if errors.As(err, &rej) { + return nil, rej + } + return nil, domain.NewInternalError(fmt.Sprintf("import transaction failed: %v", err)) } - result.Total = len(items) + result := &ImportRisksResult{ + Created: len(created), + RiskIDs: make([]uuid.UUID, 0, len(created)), + Errors: []ImportRowError{}, + Risks: created, + } + for _, r := range created { + result.RiskIDs = append(result.RiskIDs, r.ID) + } - // 2. Import each risk - for row, item := range items { - // Validate item - if item.Name == "" { - result.Errors = append(result.Errors, ImportError{ - Row: row + 1, - Reason: "name is required", - }) - result.Failed++ - continue + if uc.activation != nil && len(created) > 0 { + uc.activation.Record(ctx, tenantID, string(domain.ActivationRiskCreated), map[string]interface{}{ + "risk_id": created[0].ID.String(), + "source": string(domain.SourceImport), + "count": len(created), + }) + } + return result, nil +} + +// rejected builds the result of a refused file: nothing created, every +// offending line counted once. +func rejected(errs []ImportRowError) *ImportRisksResult { + lines := map[int]bool{} + for _, e := range errs { + if e.Line > 1 { + lines[e.Line] = true } + } + sort.SliceStable(errs, func(i, j int) bool { return errs[i].Line < errs[j].Line }) + return &ImportRisksResult{Created: 0, Rejected: len(lines), RiskIDs: []uuid.UUID{}, Errors: errs} +} + +// resolveAssets turns each row's "assets" cell into the tenant's asset ids. +// A name must match exactly one live asset of the tenant (case-insensitive); +// an id must be one of the tenant's assets. Anything else is a row error, never +// a silently unlinked risk: the link is a term of the score. +func (uc *ImportRisksUseCase) resolveAssets(ctx context.Context, tenantID uuid.UUID, rows []importRow) ([]ImportRowError, error) { + needed := false + for _, r := range rows { + if len(r.assetRefs) > 0 { + needed = true + break + } + } + if !needed { + return nil, nil + } + if uc.listAssets == nil { + return []ImportRowError{{Line: 0, Column: importColAssets, Code: "assets_unavailable", + Message: "linking assets is not available on this server; remove the assets column"}}, nil + } + refs, err := uc.listAssets(ctx, tenantID) + if err != nil { + return nil, err + } + byID := map[uuid.UUID]bool{} + byName := map[string][]uuid.UUID{} + for _, a := range refs { + byID[a.ID] = true + key := strings.ToLower(strings.TrimSpace(a.Name)) + byName[key] = append(byName[key], a.ID) + } - // Create risk domain entity - newRisk := &domain.Risk{ - ID: uuid.New(), - TenantID: tenantID, - OrganizationID: tenantID, - Name: item.Name, - Title: item.Name, - Description: item.Description, - Probability: item.Probability, - Impact: item.Impact, - Status: domain.RiskOpen, - Tags: item.Tags, - Frameworks: item.Frameworks, - CreatedBy: importedBy, - Source: domain.SourceImport, - } - - // Parse asset ID if provided - if item.AssetID != nil { - if assetID, err := uuid.Parse(*item.AssetID); err == nil { - newRisk.AssetID = &assetID + var errs []ImportRowError + for i := range rows { + seen := map[uuid.UUID]bool{} + for _, ref := range rows[i].assetRefs { + var id uuid.UUID + if parsed, perr := uuid.Parse(ref); perr == nil && byID[parsed] { + id = parsed + } else { + switch ids := byName[strings.ToLower(ref)]; len(ids) { + case 1: + id = ids[0] + case 0: + errs = append(errs, ImportRowError{Line: rows[i].line, Column: importColAssets, Code: "unknown_asset", + Params: map[string]string{"value": ref}, + Message: fmt.Sprintf("no asset named %q in the inventory", ref)}) + continue + default: + errs = append(errs, ImportRowError{Line: rows[i].line, Column: importColAssets, Code: "ambiguous_asset", + Params: map[string]string{"value": ref, "count": strconv.Itoa(len(ids))}, + Message: fmt.Sprintf("%d assets are named %q; use the asset's id instead", len(ids), ref)}) + continue + } + } + if !seen[id] { + seen[id] = true + rows[i].input.AssetIDs = append(rows[i].input.AssetIDs, id) } } + } + return errs, nil +} + +// parseImportCSV reads and validates the file. With errors it may still +// return the rows it could read, so their assets are checked in the same pass; +// such a file is never written. +func parseImportCSV(data []byte, importedBy uuid.UUID) ([]importRow, []ImportRowError) { + data = bytes.TrimPrefix(data, []byte("\xef\xbb\xbf")) // Excel's UTF-8 BOM + if len(bytes.TrimSpace(data)) == 0 { + return nil, []ImportRowError{{Line: 0, Code: "file_empty", Message: "the file is empty"}} + } + if !utf8.Valid(data) { + return nil, []ImportRowError{{Line: 0, Code: "not_utf8", Message: "the file is not UTF-8 text; save it as \"CSV UTF-8\""}} + } + + // French-locale spreadsheets export CSV with ";" and a decimal comma. + firstLine, _, _ := bytes.Cut(data, []byte("\n")) + semicolon := bytes.Count(firstLine, []byte(";")) > bytes.Count(firstLine, []byte(",")) + + reader := csv.NewReader(bytes.NewReader(data)) + if semicolon { + reader.Comma = ';' + } + reader.FieldsPerRecord = -1 // checked per row so the error carries a line + reader.TrimLeadingSpace = true - // Create in repository - if err := uc.riskRepo.Create(ctx, newRisk); err != nil { - result.Errors = append(result.Errors, ImportError{ - Row: row + 1, - Reason: fmt.Sprintf("failed to create risk: %v", err), - }) - result.Failed++ + header, err := reader.Read() + if err != nil { + return nil, []ImportRowError{{Line: 1, Code: "header_unreadable", Params: map[string]string{"detail": err.Error()}, Message: fmt.Sprintf("the header cannot be read: %v", err)}} + } + + var errs []ImportRowError + cols := map[string]int{} + names := make([]string, len(header)) + for i, h := range header { + raw := strings.TrimSpace(h) + canon, ok := importColumnAliases[strings.ToLower(raw)] + if !ok { + errs = append(errs, ImportRowError{Line: 1, Column: raw, Code: "unknown_column", Params: map[string]string{"column": raw, "accepted": strings.Join(ImportColumns, ", ")}, Message: fmt.Sprintf( + "unknown column %q; accepted columns are %s", raw, strings.Join(ImportColumns, ", "))}) continue } + if _, dup := cols[canon]; dup { + errs = append(errs, ImportRowError{Line: 1, Column: raw, Code: "duplicate_column", Params: map[string]string{"column": canon}, Message: fmt.Sprintf("column %q appears twice", canon)}) + continue + } + cols[canon] = i + names[i] = canon + } + for _, required := range []string{importColTitle, importColProbability, importColImpact} { + if _, ok := cols[required]; !ok { + errs = append(errs, ImportRowError{Line: 1, Column: required, Code: "missing_column", Params: map[string]string{"column": required}, Message: fmt.Sprintf("required column %q is missing", required)}) + } + } + if len(errs) > 0 { + return nil, errs + } - result.Created = append(result.Created, newRisk.ID) - result.Succeeded++ + cell := func(rec []string, col string) string { + i, ok := cols[col] + if !ok || i >= len(rec) { + return "" + } + return strings.TrimSpace(rec[i]) + } + number := func(s string) (float64, error) { + if semicolon { + s = strings.Replace(s, ",", ".", 1) + } + v, err := strconv.ParseFloat(s, 64) + if err == nil && (math.IsNaN(v) || math.IsInf(v, 0)) { + err = errors.New("not a finite number") + } + return v, err } - return result, nil -} + var rows []importRow + // legacy stays true while every row reads as the old 1–5 template: whole + // numbers between 1 and 5 for both probability and impact. + legacy := true + dataRows := 0 + for { + rec, err := reader.Read() + if err == io.EOF { + break + } + if err != nil { + line := 0 + var pe *csv.ParseError + if errors.As(err, &pe) { + line = pe.StartLine + } + errs = append(errs, ImportRowError{Line: line, Code: "line_unreadable", Params: map[string]string{"detail": err.Error()}, Message: fmt.Sprintf("the line cannot be read: %v", err)}) + legacy = false + // A broken quote can swallow the rest of the file; stop here rather + // than report a cascade of phantom errors. + break + } + line, _ := reader.FieldPos(0) + if isBlank(rec) { + continue + } + dataRows++ + if dataRows > MaxImportRows { + return nil, []ImportRowError{{Line: 0, Code: "too_many_rows", Params: map[string]string{"max": strconv.Itoa(MaxImportRows)}, Message: fmt.Sprintf("the file has more than %d rows; split it into several files", MaxImportRows)}} + } + if len(rec) > len(header) { + errs = append(errs, ImportRowError{Line: line, Code: "cell_count", Params: map[string]string{"cells": strconv.Itoa(len(rec)), "header": strconv.Itoa(len(header))}, Message: fmt.Sprintf("the line has %d cells but the header has %d", len(rec), len(header))}) + continue + } + + rowOK := true + title := cell(rec, importColTitle) + switch { + case title == "": + errs = append(errs, ImportRowError{Line: line, Column: importColTitle, Code: "required", Message: "title is required"}) + rowOK = false + case utf8.RuneCountInString(title) > 255: + errs = append(errs, ImportRowError{Line: line, Column: importColTitle, Code: "too_long", Params: map[string]string{"max": "255"}, Message: "title must be 255 characters or less"}) + rowOK = false + } + + prob, perr := number(cell(rec, importColProbability)) + switch { + case cell(rec, importColProbability) == "": + errs = append(errs, ImportRowError{Line: line, Column: importColProbability, Code: "required", Message: "probability is required"}) + rowOK = false + case perr != nil: + errs = append(errs, ImportRowError{Line: line, Column: importColProbability, Code: "not_a_number", Params: map[string]string{"value": cell(rec, importColProbability)}, Message: fmt.Sprintf("%q is not a number", cell(rec, importColProbability))}) + rowOK = false + case prob < 0 || prob > 1: + errs = append(errs, ImportRowError{Line: line, Column: importColProbability, Code: "out_of_range", Params: map[string]string{"min": "0", "max": "1", "value": cell(rec, importColProbability)}, Message: fmt.Sprintf("probability must be between 0 and 1 (got %s)", cell(rec, importColProbability))}) + rowOK = false + } + + imp, ierr := number(cell(rec, importColImpact)) + switch { + case cell(rec, importColImpact) == "": + errs = append(errs, ImportRowError{Line: line, Column: importColImpact, Code: "required", Message: "impact is required"}) + rowOK = false + case ierr != nil: + errs = append(errs, ImportRowError{Line: line, Column: importColImpact, Code: "not_a_number", Params: map[string]string{"value": cell(rec, importColImpact)}, Message: fmt.Sprintf("%q is not a number", cell(rec, importColImpact))}) + rowOK = false + case imp < 0 || imp > 10: + errs = append(errs, ImportRowError{Line: line, Column: importColImpact, Code: "out_of_range", Params: map[string]string{"min": "0", "max": "10", "value": cell(rec, importColImpact)}, Message: fmt.Sprintf("impact must be between 0 and 10 (got %s)", cell(rec, importColImpact))}) + rowOK = false + } + + if perr != nil || ierr != nil || !onLegacyScale(prob) || !onLegacyScale(imp) { + legacy = false + } + rows = append(rows, importRow{line: line, invalid: !rowOK, assetRefs: splitList(cell(rec, importColAssets)), input: CreateRiskInput{ + Title: title, + Description: cell(rec, importColDescription), + Probability: prob, + Impact: imp, + Tags: splitList(cell(rec, importColTags)), + Frameworks: splitList(cell(rec, importColFrameworks)), + Source: string(domain.SourceImport), + CreatedBy: importedBy, + }}) + } -// parseJSON parses JSON format import data -func (uc *ImportRisksUseCase) parseJSON(data []byte) ([]ImportRiskItem, error) { - var items []ImportRiskItem - decoder := json.NewDecoder(bytes.NewReader(data)) - if err := decoder.Decode(&items); err != nil { - return nil, fmt.Errorf("failed to decode JSON: %w", err) + if dataRows == 0 && len(errs) == 0 { + return nil, []ImportRowError{{Line: 0, Code: "no_rows", Message: "the file has a header but no risk rows"}} } - return items, nil + if legacy && dataRows > 0 { + // Reported alone: the per-row range errors it also triggers would only + // repeat the same cause once per line. + return nil, []ImportRowError{{Line: 0, Column: importColProbability, Code: "legacy_scale", Message: "this file uses the old 1–5 scale for probability and impact; " + + "OpenRisk expects probability between 0 and 1 and impact between 0 and 10. " + + "Download the current template and convert the values (for example probability 3/5 → 0.6, impact 4/5 → 8)"}} + } + if len(errs) > 0 { + return rows, errs + } + return rows, nil +} + +func onLegacyScale(v float64) bool { + return v >= 1 && v <= 5 && v == math.Trunc(v) } -// parseCSV parses CSV format import data -// Format: name,description,probability,impact,status,tags,frameworks -func (uc *ImportRisksUseCase) parseCSV(data []byte) ([]ImportRiskItem, error) { - // CSV parsing would be implemented here using a CSV library - // For now, return empty to show the structure - // In production, use encoding/csv or a dedicated library - return []ImportRiskItem{}, nil +func isBlank(rec []string) bool { + for _, v := range rec { + if strings.TrimSpace(v) != "" { + return false + } + } + return true } -// parseXLSX parses XLSX format import data -func (uc *ImportRisksUseCase) parseXLSX(data []byte) ([]ImportRiskItem, error) { - // XLSX parsing would be implemented here using a library like excelize - // For now, return empty to show the structure - // In production, use github.com/xuri/excelize or similar - return []ImportRiskItem{}, nil +// splitList reads a multi-value cell. "|" or ";" separate values; a "," does +// too when the cell holds neither, which is how most people type a list. +func splitList(s string) []string { + if s == "" { + return nil + } + sep := "," + switch { + case strings.Contains(s, "|"): + sep = "|" + case strings.Contains(s, ";"): + sep = ";" + } + var out []string + seen := map[string]bool{} + for _, part := range strings.Split(s, sep) { + p := strings.TrimSpace(part) + if p != "" && !seen[p] { + seen[p] = true + out = append(out, p) + } + } + return out } diff --git a/backend/internal/application/risk/import_risks_test.go b/backend/internal/application/risk/import_risks_test.go new file mode 100644 index 00000000..6fafc3de --- /dev/null +++ b/backend/internal/application/risk/import_risks_test.go @@ -0,0 +1,346 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). + +package risk + +import ( + "context" + "errors" + "fmt" + "strings" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/opendefender/openrisk/internal/domain" +) + +// fakeTx stages creates and only "commits" them when fn succeeds, which is the +// property the import depends on. +type fakeTx struct { + committed []*domain.Risk + failOn int // 1-based create call that errors; 0 never + assets []*domain.Asset +} + +func (f *fakeTx) run(ctx context.Context, fn func(repo domain.RiskRepository, assets RiskAssetStore) error) error { + var staged []*domain.Risk + calls := 0 + write := func(r *domain.Risk) error { + calls++ + if calls == f.failOn { + return errors.New("disk full") + } + staged = append(staged, r) + return nil + } + repo := &MockRiskRepository{createFunc: func(_ context.Context, r *domain.Risk) error { return write(r) }} + if err := fn(repo, &fakeImportAssetStore{tx: f, write: write}); err != nil { + return err + } + f.committed = append(f.committed, staged...) + return nil +} + +// list is the ImportAssetLister over the fake's assets, tenant-scoped. +func (f *fakeTx) list(_ context.Context, tenantID uuid.UUID) ([]ImportAssetRef, error) { + var out []ImportAssetRef + for _, a := range f.assets { + if a.TenantID == tenantID { + out = append(out, ImportAssetRef{ID: a.ID, Name: a.Name}) + } + } + return out, nil +} + +type fakeImportAssetStore struct { + tx *fakeTx + write func(*domain.Risk) error +} + +func (s *fakeImportAssetStore) FindByIDs(_ context.Context, tenantID uuid.UUID, ids []uuid.UUID) ([]*domain.Asset, error) { + var out []*domain.Asset + for _, a := range s.tx.assets { + for _, id := range ids { + if a.ID == id && a.TenantID == tenantID { + out = append(out, a) + } + } + } + return out, nil +} + +func (s *fakeImportAssetStore) SaveWithAssets(_ context.Context, r *domain.Risk, assets []*domain.Asset, _ bool) error { + r.Assets = assets + return s.write(r) +} + +func importCSV(t *testing.T, tx *fakeTx, csv string) (*ImportRisksResult, error) { + t.Helper() + return NewImportRisksUseCase(tx.run).Execute(context.Background(), uuid.New(), ImportRisksInput{ + CSV: []byte(csv), ImportedBy: uuid.New(), + }) +} + +func rejectedErrors(t *testing.T, err error) []ImportRowError { + t.Helper() + var rej *ImportRejectedError + require.ErrorAs(t, err, &rej) + require.ErrorIs(t, err, domain.ErrValidation) + assert.Equal(t, 0, rej.Result.Created) + return rej.Result.Errors +} + +func TestImportRisks_Success(t *testing.T) { + tx := &fakeTx{} + res, err := importCSV(t, tx, "title,description,probability,impact,tags,frameworks\n"+ + "Phishing,Credential theft,0.6,8,\"email;people\",ISO27001\n"+ + "Ransomware,,0.3,10,,\n") + require.NoError(t, err) + + assert.Equal(t, 2, res.Created) + assert.Equal(t, 0, res.Rejected) + assert.Empty(t, res.Errors) + require.Len(t, tx.committed, 2) + + r := tx.committed[0] + assert.Equal(t, "Phishing", r.Title) + assert.Equal(t, domain.SourceImport, r.Source) + assert.Equal(t, []string{"email", "people"}, []string(r.Tags)) + assert.Equal(t, []string{"ISO27001"}, []string(r.Frameworks)) + // Scored exactly as CreateRiskUseCase scores a hand-made risk. + assert.InDelta(t, 4.8, r.Score, 1e-9) + assert.Equal(t, domain.CriticalityFromScore(4.8), r.Criticality) + assert.Equal(t, []uuid.UUID{tx.committed[0].ID, tx.committed[1].ID}, res.RiskIDs) +} + +func TestImportRisks_NotFound(t *testing.T) { + // The import's "not found" case is a file with nothing in it to import. + for name, csv := range map[string]string{ + "empty": "", + "header only": "title,probability,impact\n", + "blank rows": "title,probability,impact\n,,\n\n", + } { + t.Run(name, func(t *testing.T) { + tx := &fakeTx{} + _, err := importCSV(t, tx, csv) + errs := rejectedErrors(t, err) + require.Len(t, errs, 1) + assert.Equal(t, 0, errs[0].Line) + assert.Empty(t, tx.committed) + }) + } +} + +func TestImportRisks_Unauthorized(t *testing.T) { + tx := &fakeTx{} + _, err := NewImportRisksUseCase(tx.run).Execute(context.Background(), uuid.Nil, ImportRisksInput{ + CSV: []byte("title,probability,impact\nX,0.5,5\n"), + }) + require.ErrorIs(t, err, domain.ErrUnauthorized) + assert.Empty(t, tx.committed) +} + +func TestImportRisks_Validation_OneBadRowImportsNothing(t *testing.T) { + tx := &fakeTx{} + _, err := importCSV(t, tx, "title,probability,impact\n"+ + "Good,0.5,5\n"+ + ",1.5,abc\n"+ + "Also good,0.1,2\n"+ + "Too big,0.2,11\n") + errs := rejectedErrors(t, err) + assert.Empty(t, tx.committed, "a file with an invalid row must persist nothing") + + got := map[string]bool{} + for _, e := range errs { + got[fmt.Sprintf("%d:%s", e.Line, e.Column)] = true + } + assert.Equal(t, map[string]bool{ + "3:title": true, "3:probability": true, "3:impact": true, "5:impact": true, + }, got, "every error is reported with its line and column") + + var rej *ImportRejectedError + require.ErrorAs(t, err, &rej) + assert.Equal(t, 2, rej.Result.Rejected) +} + +func TestImportRisks_WriteFailureRollsBackEverything(t *testing.T) { + tx := &fakeTx{failOn: 2} + _, err := importCSV(t, tx, "title,probability,impact\nA,0.5,5\nB,0.5,5\nC,0.5,5\n") + require.Error(t, err) + assert.Empty(t, tx.committed) +} + +func TestImportRisks_LegacyScaleIsRefusedNotConverted(t *testing.T) { + for name, csv := range map[string]string{ + "old template": "Title,Description,Probability,Impact\n\"Web API\",\"x\",3,4\n\"DB\",\"y\",4,5\n", + // Every probability is 1: valid on the 0–1 scale, but this is the old + // file's floor, and reading it as "certain" would be silent. + "all ones": "title,probability,impact\nA,1,2\nB,1,5\n", + } { + t.Run(name, func(t *testing.T) { + tx := &fakeTx{} + _, err := importCSV(t, tx, csv) + errs := rejectedErrors(t, err) + require.Len(t, errs, 1) + assert.Contains(t, errs[0].Message, "1–5 scale") + assert.Empty(t, tx.committed) + }) + } +} + +func TestImportRisks_FrenchSpreadsheetExport(t *testing.T) { + tx := &fakeTx{} + res, err := importCSV(t, tx, "\xef\xbb\xbftitre;probability;impact\n") + errs := rejectedErrors(t, err) + assert.Equal(t, "titre", errs[0].Column, "unknown columns are named, not ignored") + assert.Nil(t, res) + + res, err = importCSV(t, tx, "\xef\xbb\xbfTitle;Probability;Impact;Tags\r\nFuite;0,4;7,5;a, b\r\n") + require.NoError(t, err) + require.Equal(t, 1, res.Created) + assert.InDelta(t, 0.4, tx.committed[0].Probability, 1e-9) + assert.InDelta(t, 7.5, tx.committed[0].Impact, 1e-9) + assert.Equal(t, []string{"a", "b"}, []string(tx.committed[0].Tags)) +} + +func TestImportRisks_HeaderErrors(t *testing.T) { + tx := &fakeTx{} + _, err := importCSV(t, tx, "title,probability,status\nA,0.5,Open\n") + errs := rejectedErrors(t, err) + cols := []string{} + for _, e := range errs { + assert.Equal(t, 1, e.Line) + cols = append(cols, e.Column) + } + assert.ElementsMatch(t, []string{"status", "impact"}, cols) +} + +func TestImportRisks_TooManyRows(t *testing.T) { + var b strings.Builder + b.WriteString("title,probability,impact\n") + for i := 0; i <= MaxImportRows; i++ { + fmt.Fprintf(&b, "R%d,0.5,5\n", i) + } + tx := &fakeTx{} + _, err := importCSV(t, tx, b.String()) + errs := rejectedErrors(t, err) + assert.Contains(t, errs[0].Message, "more than") +} + +func TestImportRisks_OverCapacityWritesNothing(t *testing.T) { + tx := &fakeTx{} + uc := NewImportRisksUseCase(tx.run).WithCapacity(func(context.Context, uuid.UUID) (int, error) { return 1, nil }) + _, err := uc.Execute(context.Background(), uuid.New(), ImportRisksInput{ + CSV: []byte("title,probability,impact\nA,0.5,5\nB,0.5,5\n"), + }) + var over *ImportOverCapacityError + require.ErrorAs(t, err, &over) + assert.Equal(t, 2, over.Requested) + assert.Equal(t, 1, over.Remaining) + assert.Empty(t, tx.committed) + + // Unlimited (-1) lets it through. + uc = NewImportRisksUseCase(tx.run).WithCapacity(func(context.Context, uuid.UUID) (int, error) { return -1, nil }) + res, err := uc.Execute(context.Background(), uuid.New(), ImportRisksInput{ + CSV: []byte("title,probability,impact\nA,0.5,5\nB,0.5,5\n"), + }) + require.NoError(t, err) + assert.Equal(t, 2, res.Created) +} + +// The page translates errors from Code and Params, so every error carries a +// code and the codes below are a contract with importRisksSchema.ts. +func TestImportRisks_ErrorsCarryCodesForTranslation(t *testing.T) { + tx := &fakeTx{} + _, err := importCSV(t, tx, "title,probability,impact\n"+ + ",1.5,abc\n"+ + "Too big,0.2,11\n") + errs := rejectedErrors(t, err) + + got := map[string]ImportRowError{} + for _, e := range errs { + require.NotEmpty(t, e.Code, "error without a code: %+v", e) + got[fmt.Sprintf("%d:%s", e.Line, e.Column)] = e + } + assert.Equal(t, "required", got["2:title"].Code) + assert.Equal(t, "out_of_range", got["2:probability"].Code) + assert.Equal(t, map[string]string{"min": "0", "max": "1", "value": "1.5"}, got["2:probability"].Params) + assert.Equal(t, "not_a_number", got["2:impact"].Code) + assert.Equal(t, "abc", got["2:impact"].Params["value"]) + assert.Equal(t, "out_of_range", got["3:impact"].Code) + assert.Equal(t, "10", got["3:impact"].Params["max"]) + + _, err = importCSV(t, tx, "title,probability,impact\nA,3,4\n") + assert.Equal(t, "legacy_scale", rejectedErrors(t, err)[0].Code) + _, err = importCSV(t, tx, "titre,probability,impact\n") + assert.Equal(t, "unknown_column", rejectedErrors(t, err)[0].Code) + assert.Empty(t, tx.committed) +} + +// #755 + #792: the "assets" column links each risk to the tenant's assets by +// name or id, and the stored score is the engine's with their criticality. +func TestImportRisks_AssetsColumnLinksAndScores(t *testing.T) { + tenant := uuid.New() + critical := &domain.Asset{ID: uuid.New(), TenantID: tenant, Name: "Core banking DB", Criticality: domain.CriticalityCritical} + low := &domain.Asset{ID: uuid.New(), TenantID: tenant, Name: "Kiosk", Criticality: domain.CriticalityLow} + tx := &fakeTx{assets: []*domain.Asset{critical, low}} + + res, err := NewImportRisksUseCase(tx.run).WithAssets(tx.list).Execute(context.Background(), tenant, ImportRisksInput{ + CSV: []byte("title,probability,impact,assets\n" + + "Fraud,0.5,6,core banking db\n" + + "Theft,0.5,6," + low.ID.String() + "\n" + + "Unlinked,0.5,6,\n"), + ImportedBy: uuid.New(), + }) + require.NoError(t, err) + require.Equal(t, 3, res.Created) + require.Len(t, tx.committed, 3) + + fraud, theft, unlinked := tx.committed[0], tx.committed[1], tx.committed[2] + require.Len(t, fraud.Assets, 1) + assert.Equal(t, critical.ID, fraud.Assets[0].ID, "names match case-insensitively") + require.Len(t, theft.Assets, 1) + assert.Equal(t, low.ID, theft.Assets[0].ID, "ids are accepted") + assert.Empty(t, unlinked.Assets) + assert.Greater(t, fraud.Score, unlinked.Score, "a critical asset raises the score") + assert.Less(t, theft.Score, unlinked.Score, "a low-criticality asset lowers it") +} + +func TestImportRisks_UnknownOrAmbiguousAssetImportsNothing(t *testing.T) { + tenant := uuid.New() + other := uuid.New() + tx := &fakeTx{assets: []*domain.Asset{ + {ID: uuid.New(), TenantID: tenant, Name: "Web"}, + {ID: uuid.New(), TenantID: tenant, Name: "web"}, + {ID: uuid.New(), TenantID: other, Name: "Payroll"}, + }} + _, err := NewImportRisksUseCase(tx.run).WithAssets(tx.list).Execute(context.Background(), tenant, ImportRisksInput{ + CSV: []byte("title,probability,impact,assets\nA,0.5,6,Web\nB,0.5,6,Payroll\nC,2,6,Ghost\n"), + ImportedBy: uuid.New(), + }) + errs := rejectedErrors(t, err) + assert.Empty(t, tx.committed) + + codes := map[string]bool{} + for _, e := range errs { + codes[fmt.Sprintf("%d:%s:%s", e.Line, e.Column, e.Code)] = true + } + assert.Equal(t, map[string]bool{ + "2:assets:ambiguous_asset": true, + "3:assets:unknown_asset": true, // another tenant's asset is not visible + "4:probability:out_of_range": true, + "4:assets:unknown_asset": true, // checked even on a row with other errors + }, codes) +} + +func TestImportRisks_AssetsColumnWithoutResolverIsRefused(t *testing.T) { + tx := &fakeTx{} + _, err := importCSV(t, tx, "title,probability,impact,assets\nA,0.5,6,Web\n") + assert.Equal(t, "assets_unavailable", rejectedErrors(t, err)[0].Code) + assert.Empty(t, tx.committed) +} diff --git a/backend/internal/application/risk/score_risk.go b/backend/internal/application/risk/score_risk.go new file mode 100644 index 00000000..a724d228 --- /dev/null +++ b/backend/internal/application/risk/score_risk.go @@ -0,0 +1,43 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). + +package risk + +import ( + "context" + "fmt" + + "github.com/google/uuid" + "github.com/opendefender/openrisk/internal/domain" + pkgscoring "github.com/opendefender/openrisk/pkg/scoring" +) + +// RiskAssetStore resolves the assets a risk is linked to and writes the link +// together with the risk, so the score stored on the row is always the one its +// linked assets produce. +type RiskAssetStore interface { + // FindByIDs returns the tenant's assets among ids. Ids of another tenant's + // assets, or of no asset at all, are absent from the result. + FindByIDs(ctx context.Context, tenantID uuid.UUID, ids []uuid.UUID) ([]*domain.Asset, error) + + // SaveWithAssets creates (create=true) or saves the risk and replaces its + // asset links with assets, in one transaction. + SaveWithAssets(ctx context.Context, risk *domain.Risk, assets []*domain.Asset, create bool) error +} + +// applyScore stores on the risk the score the Score Engine computes for its +// current terms — probability, impact and the criticality of risk.Assets — and +// the band that goes with it. It is the only way create and update write a +// score, so a fresh row matches pkg/scoring before the worker ever runs (#792). +func applyScore(engine pkgscoring.Engine, r *domain.Risk) error { + ac := domain.RiskAssetCriticality(domain.AssetCriticalities(r.Assets)) + b, err := engine.Breakdown(r.Probability, r.Impact, ac, nil) + if err != nil { + return domain.NewInternalError(fmt.Sprintf("score engine rejected the risk terms: %v", err)) + } + r.Score = b.Score + r.Criticality = domain.CriticalityLevel(b.Criticality) + return nil +} diff --git a/backend/internal/application/risk/score_risk_test.go b/backend/internal/application/risk/score_risk_test.go new file mode 100644 index 00000000..ba6e594a --- /dev/null +++ b/backend/internal/application/risk/score_risk_test.go @@ -0,0 +1,221 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). + +package risk + +import ( + "context" + "errors" + "testing" + + "github.com/google/uuid" + "github.com/opendefender/openrisk/internal/domain" + pkgscoring "github.com/opendefender/openrisk/pkg/scoring" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// fakeAssetStore is an in-memory RiskAssetStore holding assets of several +// tenants, recording what the use case asked it to write. +type fakeAssetStore struct { + assets []*domain.Asset + saved *domain.Risk + linked []*domain.Asset + created bool + calls int +} + +func (f *fakeAssetStore) FindByIDs(_ context.Context, tenantID uuid.UUID, ids []uuid.UUID) ([]*domain.Asset, error) { + out := []*domain.Asset{} + for _, id := range ids { + for _, a := range f.assets { + if a.ID == id && a.TenantID == tenantID { + out = append(out, a) + } + } + } + return out, nil +} + +func (f *fakeAssetStore) SaveWithAssets(_ context.Context, r *domain.Risk, assets []*domain.Asset, create bool) error { + f.calls++ + f.saved, f.linked, f.created = r, assets, create + return nil +} + +func (f *fakeAssetStore) add(tenant uuid.UUID, crit domain.AssetCriticality) uuid.UUID { + a := &domain.Asset{ID: uuid.New(), TenantID: tenant, Name: string(crit), Criticality: crit} + f.assets = append(f.assets, a) + return a.ID +} + +// assertConsistent runs GET /risks/:id/score-working's use case on the risk as +// stored and requires it to agree with the stored score — the #792 contract. +func assertConsistent(t *testing.T, tenant uuid.UUID, stored *domain.Risk) { + t.Helper() + w, err := NewGetScoreWorkingUseCase(repoReturning(stored), nil, pkgscoring.NewEngine()). + Execute(context.Background(), tenant, stored.ID, false) + require.NoError(t, err) + assert.True(t, w.Consistent, "stored %.3f, working computes %.3f", w.Stored, w.Computed) +} + +func TestCreateRisk_ScoresWithAssetCriticality(t *testing.T) { + tenant := uuid.New() + cases := []struct { + name string + crits []domain.AssetCriticality + want float64 + }{ + // P=0.5, I=2 throughout: the issue's live case is the "one critical" row. + {"zero assets", nil, 1.0}, + {"one critical asset", []domain.AssetCriticality{domain.CriticalityCritical}, 3.0}, + {"several assets", []domain.AssetCriticality{domain.CriticalityLow, domain.CriticalityHigh, domain.CriticalityCritical}, 2.0}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + store := &fakeAssetStore{} + var ids []uuid.UUID + for _, c := range tc.crits { + ids = append(ids, store.add(tenant, c)) + } + var repoCreated *domain.Risk + repo := &MockRiskRepository{createFunc: func(_ context.Context, r *domain.Risk) error { + repoCreated = r + return nil + }} + + got, err := NewCreateRiskUseCase(repo).WithAssets(store).Execute(context.Background(), tenant, CreateRiskInput{ + Title: "t", Probability: 0.5, Impact: 2, AssetIDs: ids, + }) + require.NoError(t, err) + + assert.InDelta(t, tc.want, got.Score, 1e-9) + assert.Equal(t, domain.CriticalityFromScore(tc.want), got.Criticality) + if len(ids) == 0 { + assert.Same(t, got, repoCreated, "no asset: plain create") + assert.Zero(t, store.calls) + } else { + assert.Nil(t, repoCreated, "links and row are written together") + require.Equal(t, 1, store.calls) + assert.True(t, store.created) + assert.Len(t, store.linked, len(ids)) + } + assertConsistent(t, tenant, got) + }) + } +} + +func TestCreateRisk_ForeignAssetIsNotLinkedNorScored(t *testing.T) { + tenant, other := uuid.New(), uuid.New() + store := &fakeAssetStore{} + foreign := store.add(other, domain.CriticalityCritical) + + got, err := NewCreateRiskUseCase(&MockRiskRepository{}).WithAssets(store).Execute(context.Background(), tenant, CreateRiskInput{ + Title: "t", Probability: 0.5, Impact: 2, AssetIDs: []uuid.UUID{foreign}, + }) + require.NoError(t, err) + assert.Empty(t, got.Assets) + assert.InDelta(t, 1.0, got.Score, 1e-9, "another tenant's asset must not weigh on the score") + assert.Zero(t, store.calls) +} + +func TestUpdateRisk_ScoresWithAssetCriticality(t *testing.T) { + tenant := uuid.New() + cases := []struct { + name string + existing []domain.AssetCriticality // links already on the row + relink []domain.AssetCriticality // nil = the update sends no asset_ids + want float64 + }{ + {"zero assets", nil, nil, 1.0}, + {"one asset already linked, edit P only", []domain.AssetCriticality{domain.CriticalityCritical}, nil, 3.0}, + {"several assets sent", nil, []domain.AssetCriticality{domain.CriticalityLow, domain.CriticalityHigh, domain.CriticalityCritical}, 2.0}, + {"links replaced", []domain.AssetCriticality{domain.CriticalityCritical}, []domain.AssetCriticality{domain.CriticalityLow}, 0.5}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + store := &fakeAssetStore{} + existing := &domain.Risk{ID: uuid.New(), TenantID: tenant, Title: "t", Probability: 0.9, Impact: 2, Score: 1.8} + for _, c := range tc.existing { + existing.Assets = append(existing.Assets, &domain.Asset{ID: uuid.New(), TenantID: tenant, Criticality: c}) + } + var relinkIDs []uuid.UUID + if tc.relink != nil { + relinkIDs = []uuid.UUID{} + for _, c := range tc.relink { + relinkIDs = append(relinkIDs, store.add(tenant, c)) + } + } + repoUpdated := false + repo := &MockRiskRepository{ + getByIDFunc: func(_ context.Context, id, tid uuid.UUID) (*domain.Risk, error) { + if id == existing.ID && tid == tenant { + return existing, nil + } + return nil, nil + }, + updateFunc: func(context.Context, *domain.Risk) error { repoUpdated = true; return nil }, + } + p := 0.5 + + got, err := NewUpdateRiskUseCase(repo).WithAssets(store).Execute(context.Background(), tenant, existing.ID, UpdateRiskInput{ + Probability: &p, AssetIDs: relinkIDs, + }) + require.NoError(t, err) + + assert.InDelta(t, tc.want, got.Score, 1e-9) + assert.Equal(t, domain.CriticalityFromScore(tc.want), got.Criticality) + if tc.relink == nil { + assert.True(t, repoUpdated) + assert.Zero(t, store.calls) + } else { + assert.False(t, repoUpdated) + require.Equal(t, 1, store.calls) + assert.False(t, store.created) + assert.Len(t, store.linked, len(tc.relink)) + } + assertConsistent(t, tenant, got) + }) + } +} + +func TestUpdateRisk_WithAssets_NotFound(t *testing.T) { + store := &fakeAssetStore{} + repo := &MockRiskRepository{getByIDFunc: func(context.Context, uuid.UUID, uuid.UUID) (*domain.Risk, error) { return nil, nil }} + + _, err := NewUpdateRiskUseCase(repo).WithAssets(store).Execute(context.Background(), uuid.New(), uuid.New(), UpdateRiskInput{ + AssetIDs: []uuid.UUID{store.add(uuid.New(), domain.CriticalityHigh)}, + }) + assert.True(t, errors.Is(err, domain.ErrNotFound), "got %v", err) + assert.Zero(t, store.calls) +} + +// Another tenant's risk reads as not found through the tenant-scoped GetByID, +// and nothing — links included — is written. +func TestUpdateRisk_WithAssets_Unauthorized(t *testing.T) { + owner, intruder := uuid.New(), uuid.New() + store := &fakeAssetStore{} + riskID := uuid.New() + repo := &MockRiskRepository{getByIDFunc: func(_ context.Context, id, tid uuid.UUID) (*domain.Risk, error) { + if tid == owner { + return &domain.Risk{ID: id, TenantID: owner, Title: "t"}, nil + } + return nil, nil + }} + + _, err := NewUpdateRiskUseCase(repo).WithAssets(store).Execute(context.Background(), intruder, riskID, UpdateRiskInput{ + AssetIDs: []uuid.UUID{store.add(intruder, domain.CriticalityCritical)}, + }) + assert.True(t, errors.Is(err, domain.ErrNotFound), "got %v", err) + assert.Zero(t, store.calls) +} + +func TestRiskAssetCriticality(t *testing.T) { + assert.Equal(t, domain.NoAssetCriticalityFactor, domain.RiskAssetCriticality(nil)) + assert.InDelta(t, 3.0, domain.RiskAssetCriticality([]domain.AssetCriticality{domain.CriticalityCritical}), 1e-9) + assert.InDelta(t, 2.0, domain.RiskAssetCriticality([]domain.AssetCriticality{ + domain.CriticalityLow, domain.CriticalityHigh, domain.CriticalityCritical, + }), 1e-9) +} diff --git a/backend/internal/application/risk/update_risk.go b/backend/internal/application/risk/update_risk.go index 35e676ef..4dfbcba7 100644 --- a/backend/internal/application/risk/update_risk.go +++ b/backend/internal/application/risk/update_risk.go @@ -12,6 +12,7 @@ import ( "github.com/google/uuid" "github.com/opendefender/openrisk/internal/domain" + pkgscoring "github.com/opendefender/openrisk/pkg/scoring" ) // UpdateRiskInput represents the input for updating a risk. @@ -33,6 +34,9 @@ type UpdateRiskInput struct { // CategoryID is tri-state via NullableUUID for the same reason ownership is: // omitting it must not clear it. Category domain.NullableUUID + // AssetIDs, when non-nil, replaces the risk's asset links. nil leaves them + // alone. Ids that are not the tenant's assets are dropped. + AssetIDs []uuid.UUID // CRQ monetary inputs (XAF). Pointers so a partial update can set or clear them. SLEXAF *float64 ARO *float64 @@ -58,10 +62,19 @@ type UpdateRiskInput struct { type UpdateRiskUseCase struct { riskRepo domain.RiskRepository ownership OwnershipManager + assets RiskAssetStore + engine pkgscoring.Engine } func NewUpdateRiskUseCase(riskRepo domain.RiskRepository) *UpdateRiskUseCase { - return &UpdateRiskUseCase{riskRepo: riskRepo} + return &UpdateRiskUseCase{riskRepo: riskRepo, engine: pkgscoring.NewEngine()} +} + +// WithAssets attaches the asset store that resolves and links input.AssetIDs. +// Without it, AssetIDs is ignored and the risk keeps its current links. +func (uc *UpdateRiskUseCase) WithAssets(s RiskAssetStore) *UpdateRiskUseCase { + uc.assets = s + return uc } // WithOwnership attaches the optional ownership manager (membership validation @@ -195,14 +208,35 @@ func (uc *UpdateRiskUseCase) Execute(ctx context.Context, orgID uuid.UUID, riskI risk.AssignedTo = risk.AssigneeID } - // 3. Recompute score + band it synchronously so the update response is - // self-consistent (score and criticality agree) rather than showing a stale - // band until the async ScoreWorker runs (audit-2026 #246). - risk.Score = risk.Impact * risk.Probability - risk.Criticality = domain.CriticalityFromScore(risk.Score) + // 2c. Replace the asset links when the caller sent them; otherwise the + // links GetByID preloaded stay, and still feed the score below. + relink := uc.assets != nil && input.AssetIDs != nil + var linked []*domain.Asset + if relink { + linked = []*domain.Asset{} + if len(input.AssetIDs) > 0 { + linked, err = uc.assets.FindByIDs(ctx, orgID, input.AssetIDs) + if err != nil { + return nil, domain.NewInternalError(fmt.Sprintf("failed to resolve assets: %v", err)) + } + } + risk.Assets = linked + } + + // 3. Recompute the score through the Score Engine with the risk's current + // terms, asset criticality included, so the update response and the row + // both carry the formula's score before the async worker runs (#792). + if err := applyScore(uc.engine, risk); err != nil { + return nil, err + } - // 4. Persist - if err := uc.riskRepo.Update(ctx, risk); err != nil { + // 4. Persist, with the new asset links in the same transaction. + if relink { + err = uc.assets.SaveWithAssets(ctx, risk, linked, false) + } else { + err = uc.riskRepo.Update(ctx, risk) + } + if err != nil { return nil, domain.NewInternalError(fmt.Sprintf("failed to update risk: %v", err)) } diff --git a/backend/internal/domain/asset.go b/backend/internal/domain/asset.go index b171c649..4aa18253 100644 --- a/backend/internal/domain/asset.go +++ b/backend/internal/domain/asset.go @@ -42,6 +42,38 @@ func (c AssetCriticality) ScoreFactor() float64 { } } +// NoAssetCriticalityFactor is the asset-criticality term of a risk with no +// linked asset: neutral, so the score is P × I. It is what the Score Engine +// worker, the handlers and the demo seed have always stored for such a risk. +const NoAssetCriticalityFactor = 1.0 + +// RiskAssetCriticality is THE asset-criticality term of the frozen formula for +// a risk: the average ScoreFactor of its linked assets, or +// NoAssetCriticalityFactor when none is linked. Every writer and every reader +// of a risk score derives the term here, so a stored score and its displayed +// working can never disagree on it (#792). +func RiskAssetCriticality(crits []AssetCriticality) float64 { + if len(crits) == 0 { + return NoAssetCriticalityFactor + } + var sum float64 + for _, c := range crits { + sum += c.ScoreFactor() + } + return sum / float64(len(crits)) +} + +// AssetCriticalities lists the criticality of each asset, for RiskAssetCriticality. +func AssetCriticalities(assets []*Asset) []AssetCriticality { + out := make([]AssetCriticality, 0, len(assets)) + for _, a := range assets { + if a != nil { + out = append(out, a.Criticality) + } + } + return out +} + type Asset struct { ID uuid.UUID `gorm:"type:uuid;default:gen_random_uuid();primaryKey" json:"id"` TenantID uuid.UUID `gorm:"type:uuid;index" json:"tenant_id"` diff --git a/backend/internal/domain/score_working.go b/backend/internal/domain/score_working.go index 566a6f56..3a5a59eb 100644 --- a/backend/internal/domain/score_working.go +++ b/backend/internal/domain/score_working.go @@ -75,8 +75,8 @@ type ScoreWorking struct { Formula string `json:"formula"` Terms []ScoreWorkingTerm `json:"terms"` // Assets are averaged into asset_criticality. Empty with - // AssetCriticalityDefaulted=true means no asset is linked and the engine's - // documented default (medium) applied. + // AssetCriticalityDefaulted=true means no asset is linked and the neutral + // factor NoAssetCriticalityFactor (1.0) applied. Assets []ScoreWorkingAsset `json:"assets"` AssetCriticalityDefaulted bool `json:"asset_criticality_defaulted"` diff --git a/backend/internal/handler/risk_handler.go b/backend/internal/handler/risk_handler.go index 7fac0683..da0cbe30 100644 --- a/backend/internal/handler/risk_handler.go +++ b/backend/internal/handler/risk_handler.go @@ -7,7 +7,7 @@ package handler import ( "context" - "log" + "fmt" "strconv" "strings" @@ -37,6 +37,7 @@ type RiskHandler struct { crq *crq.Quantifier // Cyber Risk Quantification (XAF + USD) presenters *risk.FinancialPresenterFactory // optional: tenant currency + FX bulkActionUC *risk.BulkActionUseCase // #581, attached via WithBulkAction + importRisksUC *risk.ImportRisksUseCase // #755, attached via WithImport } // WithBulkAction attaches the bulk-action use case behind POST /risks/bulk. @@ -334,6 +335,10 @@ func (h *RiskHandler) CreateRisk(c *fiber.Ctx) error { if err != nil { return c.Status(400).JSON(fiber.Map{"error": "validation_failed", "details": err.Error()}) } + assetIDs, err := parseAssetIDs(input.AssetIDs) + if err != nil { + return c.Status(400).JSON(fiber.Map{"error": "validation_failed", "details": err.Error()}) + } ucInput := risk.CreateRiskInput{ Title: input.Title, @@ -344,6 +349,7 @@ func (h *RiskHandler) CreateRisk(c *fiber.Ctx) error { Frameworks: input.Frameworks, Ownership: input.OwnershipPatch, CategoryID: categoryID, + AssetIDs: assetIDs, CreatedBy: createdBy, SLEXAF: input.SLEXAF, ARO: input.ARO, @@ -362,41 +368,24 @@ func (h *RiskHandler) CreateRisk(c *fiber.Ctx) error { return c.Status(400).JSON(fiber.Map{"error": err.Error()}) } - // Link Assets (fallback until AssetRepo introduced) - var linkedAssets []*domain.Asset - if len(input.AssetIDs) > 0 { - // Tenant-scoped unconditionally: this filter used to apply only when the - // middleware context was present, so its absence meant no filter at all. - // The request context is passed so the audit trail attributes the link to - // the caller instead of recording an unattributed write (#486). - query := database.DB.WithContext(stdCtx).Where("organization_id = ?", orgID) - if err := query.Where("id IN ?", input.AssetIDs).Find(&linkedAssets).Error; err == nil { - domainRisk.Assets = linkedAssets - // Save relationships (no direct score compute — publish Redis event instead) - if err := database.DB.WithContext(stdCtx).Model(&domainRisk).Association("Assets").Replace(linkedAssets); err != nil { - log.Printf("Warning: failed to update asset associations for risk %s: %v", domainRisk.ID, err) - } - } - } - // RULE #12: Score Engine is NEVER called directly from handler. - // Always publish Redis event → ScoreWorker listens and recalculates async, - // using the real criticality of whichever assets were just linked instead - // of a hardcoded placeholder. + // The use case already stored the engine's score with the linked assets + // (#792); the event still drives the worker's audit entry and the + // risk.score_updated fan-out, with the same asset term. if h.redisClient != nil { event := events.RiskUpdatedEvent{ RiskID: domainRisk.ID.String(), TenantID: orgID.String(), Probability: float64(domainRisk.Probability), Impact: float64(domainRisk.Impact), - AssetCriticality: averageAssetCriticalityFactor(linkedAssets), + AssetCriticality: domain.RiskAssetCriticality(domain.AssetCriticalities(domainRisk.Assets)), TriggeredBy: createdBy.String(), } _ = h.redisClient.Publish(c.Context(), events.RiskUpdated, event) } var out domain.Risk - if err := database.DB.Preload("Mitigations").Preload("Mitigations.SubActions").Preload("Assets").First(&out, "id = ?", domainRisk.ID).Error; err != nil { + if err := database.DB.Preload("Mitigations").Preload("Mitigations.SubActions").Preload("Assets").First(&out, "id = ? AND tenant_id = ?", domainRisk.ID, orgID).Error; err != nil { h.quantify(domainRisk) return c.Status(201).JSON(domainRisk) } @@ -585,6 +574,13 @@ func (h *RiskHandler) UpdateRisk(c *fiber.Ctx) error { actorID = mwCtx.UserID } + var assetIDs []uuid.UUID + if len(input.AssetIDs) > 0 { + if assetIDs, err = parseAssetIDs(input.AssetIDs); err != nil { + return c.Status(400).JSON(fiber.Map{"error": "validation_failed", "details": err.Error()}) + } + } + ucInput := risk.UpdateRiskInput{ Title: &input.Title, Description: &input.Description, @@ -594,6 +590,7 @@ func (h *RiskHandler) UpdateRisk(c *fiber.Ctx) error { Frameworks: input.Frameworks, Ownership: input.OwnershipPatch, Category: input.Category, + AssetIDs: assetIDs, Actor: actorID, Locale: c.Query("locale", "fr"), SLEXAF: input.SLEXAF, @@ -633,35 +630,13 @@ func (h *RiskHandler) UpdateRisk(c *fiber.Ctx) error { return writeAppError(c, err) } - if len(input.AssetIDs) > 0 { - var linkedAssets []*domain.Asset - query := database.DB - if mwCtx != nil { - query = query.Where("organization_id = ?", mwCtx.OrganizationID) - } - if err := query.Where("id IN ?", input.AssetIDs).Find(&linkedAssets).Error; err == nil { - domainRisk.Assets = linkedAssets - // No direct score compute here (RULE #12) — save the association, - // then publish a Redis event below so the ScoreWorker recalculates - // via the real Score Engine, same as CreateRisk. - if err := database.DB.Model(&domainRisk).Association("Assets").Replace(linkedAssets); err != nil { - log.Printf("Warning: failed to update asset associations for risk %s: %v", domainRisk.ID, err) - } - } - } - var out domain.Risk - hasOut := database.DB.Preload("Mitigations").Preload("Mitigations.SubActions").Preload("Assets").First(&out, "id = ?", riskID).Error == nil + hasOut := database.DB.Preload("Mitigations").Preload("Mitigations.SubActions").Preload("Assets").First(&out, "id = ? AND tenant_id = ?", riskID, orgID).Error == nil // RULE #12: Score Engine is NEVER called directly from handler. // Always publish Redis event → ScoreWorker listens and recalculates async. - // Uses the risk's currently linked assets — freshly replaced above if this - // update touched asset_ids, or its pre-existing ones otherwise — so an - // Impact/Probability-only edit still gets a criticality-adjusted score. - assetsForScoring := domainRisk.Assets - if hasOut { - assetsForScoring = out.Assets - } + // Uses the risk's linked assets as the use case scored them — replaced if + // this update sent asset_ids, its existing ones otherwise (#792). if h.redisClient != nil { userID := uuid.Nil if mwCtx != nil { @@ -672,7 +647,7 @@ func (h *RiskHandler) UpdateRisk(c *fiber.Ctx) error { TenantID: orgID.String(), Probability: float64(domainRisk.Probability), Impact: float64(domainRisk.Impact), - AssetCriticality: averageAssetCriticalityFactor(assetsForScoring), + AssetCriticality: domain.RiskAssetCriticality(domain.AssetCriticalities(domainRisk.Assets)), TriggeredBy: userID.String(), } _ = h.redisClient.Publish(c.Context(), events.RiskUpdated, event) @@ -686,18 +661,18 @@ func (h *RiskHandler) UpdateRisk(c *fiber.Ctx) error { return c.JSON(out) } -// averageAssetCriticalityFactor averages domain.AssetCriticality.ScoreFactor() -// across a risk's linked assets, for the Redis event consumed by ScoreWorker. -// Defaults to 1.0 (neutral) when a risk has no linked assets yet. -func averageAssetCriticalityFactor(assets []*domain.Asset) float64 { - if len(assets) == 0 { - return 1.0 - } - var sum float64 - for _, a := range assets { - sum += a.Criticality.ScoreFactor() +// parseAssetIDs parses the asset_ids of a create or update body. A malformed +// id is a 400: it used to fail the lookup silently and drop every link. +func parseAssetIDs(raw []string) ([]uuid.UUID, error) { + ids := make([]uuid.UUID, 0, len(raw)) + for _, s := range raw { + id, err := uuid.Parse(s) + if err != nil { + return nil, fmt.Errorf("asset_ids: %q is not a valid id", s) + } + ids = append(ids, id) } - return sum / float64(len(assets)) + return ids, nil } // DeleteRisk godoc @@ -726,7 +701,6 @@ func (h *RiskHandler) DeleteRisk(c *fiber.Ctx) error { return c.SendStatus(204) } - // --------------------------------------------------------------------------- // Bulk actions — POST /api/v1/risks/bulk (#581) // --------------------------------------------------------------------------- diff --git a/backend/internal/handler/risk_import_handler.go b/backend/internal/handler/risk_import_handler.go new file mode 100644 index 00000000..cd6421ca --- /dev/null +++ b/backend/internal/handler/risk_import_handler.go @@ -0,0 +1,111 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). + +package handler + +import ( + "errors" + "io" + "path/filepath" + "strings" + + "github.com/gofiber/fiber/v2" + + "github.com/opendefender/openrisk/internal/application/risk" + "github.com/opendefender/openrisk/internal/domain" + "github.com/opendefender/openrisk/pkg/events" +) + +// WithImport attaches the CSV import use case behind POST /risks/import. +func (h *RiskHandler) WithImport(uc *risk.ImportRisksUseCase) *RiskHandler { + h.importRisksUC = uc + return h +} + +// ImportRisks POST /risks/import — multipart field "file", a CSV. +// +// 200 {created, rejected, risk_ids, errors: []} when every row was written. +// 422 {error, message, created: 0, rejected, errors: [{line, column, message}]} +// when any row is invalid: nothing was written. +// 402 limit_reached when the file would take the tenant past its plan. +// +// Tenant and author come from the signed session only, never the request. +func (h *RiskHandler) ImportRisks(c *fiber.Ctx) error { + if h.importRisksUC == nil { + return c.Status(fiber.StatusNotImplemented).JSON(fiber.Map{"error": "risk import is not enabled"}) + } + + file, err := c.FormFile("file") + if err != nil { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": "validation_failed", "message": "Attach a CSV file in the \"file\" field.", + }) + } + if !strings.EqualFold(filepath.Ext(file.Filename), ".csv") { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": "validation_failed", "message": "Only CSV files can be imported.", + }) + } + if file.Size > risk.MaxImportBytes { + return c.Status(fiber.StatusRequestEntityTooLarge).JSON(fiber.Map{ + "error": "validation_failed", "message": "The file is larger than 2 MB. Split it into several files.", + }) + } + f, err := file.Open() + if err != nil { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": "validation_failed", "message": "The file cannot be read."}) + } + defer f.Close() + data, err := io.ReadAll(io.LimitReader(f, risk.MaxImportBytes+1)) + if err != nil { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": "validation_failed", "message": "The file cannot be read."}) + } + + tenant, actor := tenantID(c), userID(c) + result, err := h.importRisksUC.Execute(auditCtx(c), tenant, risk.ImportRisksInput{CSV: data, ImportedBy: actor}) + if err != nil { + var rej *risk.ImportRejectedError + if errors.As(err, &rej) { + return c.Status(fiber.StatusUnprocessableEntity).JSON(fiber.Map{ + "error": "validation_failed", + "message": "Nothing was imported. Fix the lines below and import the file again.", + "created": rej.Result.Created, + "rejected": rej.Result.Rejected, + "risk_ids": rej.Result.RiskIDs, + "errors": rej.Result.Errors, + }) + } + var over *risk.ImportOverCapacityError + if errors.As(err, &over) { + return c.Status(fiber.StatusPaymentRequired).JSON(fiber.Map{ + "code": "limit_reached", + "limit_key": "risks", + "requested": over.Requested, + "remaining": over.Remaining, + "message": "This file would take you past your plan's risk limit. Nothing was imported.", + "upgrade_url": "/settings?tab=billing", + }) + } + return writeAppError(c, err) + } + + // Same contract as CreateRisk: the stored score is already the engine's, + // linked assets included; the event lets the Score Engine fold in the + // signals it computes asynchronously. + if h.redisClient != nil { + for _, r := range result.Risks { + _ = h.redisClient.Publish(c.Context(), events.RiskUpdated, events.RiskUpdatedEvent{ + RiskID: r.ID.String(), + TenantID: tenant.String(), + Probability: r.Probability, + Impact: r.Impact, + AssetCriticality: domain.RiskAssetCriticality(domain.AssetCriticalities(r.Assets)), + TriggeredBy: actor.String(), + }) + } + } + + return c.JSON(result) +} diff --git a/backend/internal/handler/risk_import_test.go b/backend/internal/handler/risk_import_test.go new file mode 100644 index 00000000..9eba99ee --- /dev/null +++ b/backend/internal/handler/risk_import_test.go @@ -0,0 +1,228 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). + +package handler + +import ( + "bytes" + "encoding/json" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gofiber/fiber/v2" + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + applicationrisk "github.com/opendefender/openrisk/internal/application/risk" + "github.com/opendefender/openrisk/internal/domain" + "github.com/opendefender/openrisk/internal/infrastructure/database" + "github.com/opendefender/openrisk/internal/infrastructure/repository" + "github.com/opendefender/openrisk/internal/middleware" + "github.com/opendefender/openrisk/internal/testsupport/sqliteschema" + "github.com/opendefender/openrisk/pkg/crq" +) + +// importApp mounts POST /risks/import exactly as main.go does — behind +// risks:create — over a real sqlite database, so the transaction is real. +type importApp struct { + app *fiber.App + db *gorm.DB + tenant *uuid.UUID + perms *[]string +} + +func newImportApp(t *testing.T) *importApp { + t.Helper() + + dsn := "file:risk_import_" + uuid.New().String() + "?mode=memory&cache=private" + db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{}) + require.NoError(t, err) + require.NoError(t, db.AutoMigrate(&UserT{}, &MitigationT{}, &RiskHistoryT{})) + createRisksTable(t, db) + // The import resolves assets tenant-scoped and links them, so the assets + // table needs its tenant and criticality columns, and the join table. + require.NoError(t, db.Exec(`CREATE TABLE assets (id TEXT PRIMARY KEY, tenant_id TEXT NOT NULL, name TEXT NOT NULL, + criticality TEXT NOT NULL DEFAULT 'MEDIUM', created_at DATETIME, updated_at DATETIME, deleted_at DATETIME)`).Error) + require.NoError(t, sqliteschema.Reconcile(db, "assets", &domain.Asset{})) + require.NoError(t, db.Exec(`CREATE TABLE risk_assets (risk_id TEXT NOT NULL, asset_id TEXT NOT NULL)`).Error) + + orig := database.DB + database.DB = db + t.Cleanup(func() { database.DB = orig }) + + h := &importApp{db: db, tenant: new(uuid.UUID), perms: &[]string{"risks:create"}} + app := fiber.New() + app.Use(func(c *fiber.Ctx) error { + middleware.SetContext(c, &middleware.RequestContext{UserID: uuid.New(), OrganizationID: *h.tenant}) + c.Locals("permissions", *h.perms) + return c.Next() + }) + + riskRepo := repository.NewGormRiskRepository(db) + handler := NewRiskHandler( + applicationrisk.NewCreateRiskUseCase(riskRepo), + applicationrisk.NewGetRiskUseCase(riskRepo), + applicationrisk.NewListRisksUseCase(riskRepo), + applicationrisk.NewUpdateRiskUseCase(riskRepo), + applicationrisk.NewDeleteRiskUseCase(riskRepo), + applicationrisk.NewMarkRiskReviewedUseCase(riskRepo), + applicationrisk.NewTransitionRiskStateUseCase(riskRepo), + nil, + crq.NewQuantifier(0, crq.Reference{}), + ).WithImport(applicationrisk.NewImportRisksUseCase(repository.RunRiskTx(db)). + WithAssets(repository.ListImportAssetRefs(db))) + + app.Post("/api/v1/risks/import", middleware.RequirePermission("risks:create"), handler.ImportRisks) + h.app = app + return h +} + +func (h *importApp) upload(t *testing.T, filename, content string) (int, map[string]any) { + t.Helper() + var body bytes.Buffer + w := multipart.NewWriter(&body) + part, err := w.CreateFormFile("file", filename) + require.NoError(t, err) + _, _ = part.Write([]byte(content)) + require.NoError(t, w.Close()) + + req := httptest.NewRequest(http.MethodPost, "/api/v1/risks/import", &body) + req.Header.Set("Content-Type", w.FormDataContentType()) + resp, err := h.app.Test(req, -1) + require.NoError(t, err) + defer resp.Body.Close() + raw, _ := io.ReadAll(resp.Body) + var decoded map[string]any + _ = json.Unmarshal(raw, &decoded) + return resp.StatusCode, decoded +} + +func (h *importApp) risksOf(t *testing.T, tenant uuid.UUID) []domain.Risk { + t.Helper() + var out []domain.Risk + require.NoError(t, h.db.Where("tenant_id = ?", tenant).Find(&out).Error) + return out +} + +const validImport = "title,description,probability,impact,tags\n" + + "Phishing,Credential theft,0.6,8,email;people\n" + + "Ransomware,,0.3,10,\n" + + "Supplier outage,,0.2,5,\n" + +func TestRiskImportHTTP_Success(t *testing.T) { + h := newImportApp(t) + tenant := uuid.New() + *h.tenant = tenant + + status, body := h.upload(t, "register.csv", validImport) + require.Equal(t, fiber.StatusOK, status, "%v", body) + require.EqualValues(t, 3, body["created"]) + require.EqualValues(t, 0, body["rejected"]) + require.Len(t, body["risk_ids"], 3) + require.Empty(t, body["errors"]) + + rows := h.risksOf(t, tenant) + require.Len(t, rows, 3, "every CSV row is persisted") + for _, r := range rows { + require.Equal(t, domain.SourceImport, r.Source) + require.InDelta(t, r.Probability*r.Impact, r.Score, 1e-9, "each row is scored on create") + require.Equal(t, domain.CriticalityFromScore(r.Score), r.Criticality) + } +} + +func TestRiskImportHTTP_Validation_NothingPersisted(t *testing.T) { + h := newImportApp(t) + tenant := uuid.New() + *h.tenant = tenant + + status, body := h.upload(t, "register.csv", validImport+"Broken,,2,5,\n") + require.Equal(t, fiber.StatusUnprocessableEntity, status, "%v", body) + require.EqualValues(t, 0, body["created"]) + require.EqualValues(t, 1, body["rejected"]) + errs, _ := body["errors"].([]any) + require.Len(t, errs, 1) + first, _ := errs[0].(map[string]any) + require.EqualValues(t, 5, first["line"]) + require.Equal(t, "probability", first["column"]) + + require.Empty(t, h.risksOf(t, tenant), "one invalid row means zero rows persisted") +} + +func TestRiskImportHTTP_Unauthorized(t *testing.T) { + h := newImportApp(t) + tenant := uuid.New() + *h.tenant = tenant + *h.perms = []string{"risks:read"} + + status, _ := h.upload(t, "register.csv", validImport) + require.Equal(t, fiber.StatusForbidden, status) + require.Empty(t, h.risksOf(t, tenant)) +} + +func TestRiskImportHTTP_NotFound_NoFile(t *testing.T) { + h := newImportApp(t) + *h.tenant = uuid.New() + + req := httptest.NewRequest(http.MethodPost, "/api/v1/risks/import", nil) + resp, err := h.app.Test(req, -1) + require.NoError(t, err) + require.Equal(t, fiber.StatusBadRequest, resp.StatusCode) + + status, _ := h.upload(t, "register.xlsx", validImport) + require.Equal(t, fiber.StatusBadRequest, status, "only CSV is accepted") +} + +func TestRiskImportHTTP_RowsLandOnlyInCallersTenant(t *testing.T) { + h := newImportApp(t) + tenantA, tenantB := uuid.New(), uuid.New() + + *h.tenant = tenantB + status, body := h.upload(t, "b.csv", "title,probability,impact\nB own,0.5,5\n") + require.Equal(t, fiber.StatusOK, status, "%v", body) + + *h.tenant = tenantA + status, body = h.upload(t, "a.csv", validImport) + require.Equal(t, fiber.StatusOK, status, "%v", body) + + require.Len(t, h.risksOf(t, tenantA), 3) + bRows := h.risksOf(t, tenantB) + require.Len(t, bRows, 1, "tenant A's import must not land in tenant B") + require.Equal(t, "B own", bRows[0].Title) +} + +// The "assets" column resolves names inside the caller's tenant only, links +// through the same transaction, and a name from another tenant refuses the file. +func TestRiskImportHTTP_AssetsResolveOnlyInCallersTenant(t *testing.T) { + h := newImportApp(t) + tenantA, tenantB := uuid.New(), uuid.New() + ownAsset, foreignAsset := uuid.New(), uuid.New() + require.NoError(t, h.db.Exec(`INSERT INTO assets (id, tenant_id, name, criticality) VALUES (?, ?, 'Core DB', 'CRITICAL'), (?, ?, 'Payroll', 'HIGH')`, + ownAsset, tenantA, foreignAsset, tenantB).Error) + + *h.tenant = tenantA + status, body := h.upload(t, "x.csv", "title,probability,impact,assets\nLeak,0.5,6,Payroll\n") + require.Equal(t, fiber.StatusUnprocessableEntity, status, "%v", body) + require.Empty(t, h.risksOf(t, tenantA)) + + status, body = h.upload(t, "ok.csv", "title,probability,impact,assets\nLeak,0.5,6,core db\nOther,0.5,6,\n") + require.Equal(t, fiber.StatusOK, status, "%v", body) + + var links []struct{ RiskID, AssetID string } + require.NoError(t, h.db.Raw(`SELECT risk_id, asset_id FROM risk_assets`).Scan(&links).Error) + require.Len(t, links, 1) + assert.Equal(t, ownAsset.String(), links[0].AssetID) + + scores := map[string]float64{} + for _, r := range h.risksOf(t, tenantA) { + scores[r.Title] = r.Score + } + assert.Greater(t, scores["Leak"], scores["Other"], "the critical asset is in the stored score") +} diff --git a/backend/internal/infrastructure/repository/gorm_risk_asset_store.go b/backend/internal/infrastructure/repository/gorm_risk_asset_store.go new file mode 100644 index 00000000..6b79a0fe --- /dev/null +++ b/backend/internal/infrastructure/repository/gorm_risk_asset_store.go @@ -0,0 +1,79 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). + +package repository + +import ( + "context" + "fmt" + + "github.com/google/uuid" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/opendefender/openrisk/internal/domain" +) + +// GormRiskAssetStore implements risk.RiskAssetStore: it resolves a tenant's +// assets and writes a risk with its asset links in one transaction (#792). +type GormRiskAssetStore struct { + db *gorm.DB +} + +func NewGormRiskAssetStore(db *gorm.DB) *GormRiskAssetStore { + return &GormRiskAssetStore{db: db} +} + +// FindByIDs returns the tenant's assets among ids; foreign or unknown ids are +// simply absent. +func (s *GormRiskAssetStore) FindByIDs(ctx context.Context, tenantID uuid.UUID, ids []uuid.UUID) ([]*domain.Asset, error) { + if tenantID == uuid.Nil { + return nil, fmt.Errorf("tenant_id is required") + } + assets := []*domain.Asset{} + if len(ids) == 0 { + return assets, nil + } + if err := s.db.WithContext(ctx). + Where("tenant_id = ? AND id IN ?", tenantID, ids). + Find(&assets).Error; err != nil { + return nil, fmt.Errorf("failed to find assets: %w", err) + } + return assets, nil +} + +// SaveWithAssets writes the risk and replaces its links in risk_assets. +// +// risk_assets has no tenant_id: it is gated through its two parents. The risk +// is the caller's (created here, or loaded tenant-scoped by the use case) and +// every asset must carry the risk's tenant, checked before anything is written. +func (s *GormRiskAssetStore) SaveWithAssets(ctx context.Context, risk *domain.Risk, assets []*domain.Asset, create bool) error { + if risk.TenantID == uuid.Nil { + return fmt.Errorf("tenant_id is required") + } + for _, a := range assets { + if a == nil || a.TenantID != risk.TenantID { + return fmt.Errorf("asset does not belong to the risk's tenant") + } + } + return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + // The links are written by Replace below, not as a side effect of the + // row write, so a removed asset is unlinked rather than left behind. + write := tx.Omit(clause.Associations) + var err error + if create { + err = write.Create(risk).Error + } else { + err = write.Save(risk).Error + } + if err != nil { + return err + } + if err := tx.Model(risk).Association("Assets").Replace(assets); err != nil { + return fmt.Errorf("failed to link assets: %w", err) + } + return nil + }) +} diff --git a/backend/internal/infrastructure/repository/gorm_risk_asset_store_pg_test.go b/backend/internal/infrastructure/repository/gorm_risk_asset_store_pg_test.go new file mode 100644 index 00000000..e03c47ca --- /dev/null +++ b/backend/internal/infrastructure/repository/gorm_risk_asset_store_pg_test.go @@ -0,0 +1,106 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). + +package repository + +import ( + "context" + "errors" + "os" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/postgres" + "gorm.io/gorm" + "gorm.io/gorm/logger" + + "github.com/opendefender/openrisk/internal/domain" +) + +// #792 — the risk row carries a score computed from its asset links, so the +// row and the links must be written together. Runs against a migrated +// database; everything happens in an outer transaction that is rolled back. +func TestGormRiskAssetStore_Postgres(t *testing.T) { + dsn := os.Getenv("DATABASE_URL") + if dsn == "" { + t.Skip("DATABASE_URL not set") + } + db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) + require.NoError(t, err) + + rollback := errors.New("rollback") + err = db.Transaction(func(tx *gorm.DB) error { + ctx := context.Background() + store := NewGormRiskAssetStore(tx) + tenant, other := uuid.New(), uuid.New() + + asset := func(tid uuid.UUID, c domain.AssetCriticality) *domain.Asset { + a := &domain.Asset{ID: uuid.New(), TenantID: tid, Name: "a-" + uuid.NewString()[:8], Type: "server", Criticality: c} + require.NoError(t, tx.Create(a).Error) + return a + } + links := func(riskID uuid.UUID) []uuid.UUID { + var ids []uuid.UUID + // risk_assets has no tenant_id; riskID is a risk this test created under tenant. + require.NoError(t, tx.Table("risk_assets").Where("risk_id = ?", riskID).Order("asset_id").Pluck("asset_id", &ids).Error) + return ids + } + riskCount := func(id uuid.UUID) int64 { + var n int64 + require.NoError(t, tx.Model(&domain.Risk{}).Where("id = ? AND tenant_id = ?", id, tenant).Count(&n).Error) + return n + } + + crit, low, high := asset(tenant, domain.CriticalityCritical), asset(tenant, domain.CriticalityLow), asset(tenant, domain.CriticalityHigh) + foreign := asset(other, domain.CriticalityCritical) + + // FindByIDs is tenant-scoped: the other tenant's asset is absent. + found, err := store.FindByIDs(ctx, tenant, []uuid.UUID{crit.ID, foreign.ID}) + require.NoError(t, err) + require.Len(t, found, 1) + assert.Equal(t, crit.ID, found[0].ID) + + // Create writes the row and the links. + r := &domain.Risk{ID: uuid.New(), TenantID: tenant, Title: "pg-792", Probability: 0.5, Impact: 2, Score: 3} + r.SetState(domain.StateDraft) + require.NoError(t, store.SaveWithAssets(ctx, r, []*domain.Asset{crit}, true)) + assert.Equal(t, int64(1), riskCount(r.ID)) + assert.Equal(t, []uuid.UUID{crit.ID}, links(r.ID)) + var stored float64 + require.NoError(t, tx.Table("risks").Where("id = ? AND tenant_id = ?", r.ID, tenant).Pluck("score", &stored).Error) + assert.InDelta(t, 3.0, stored, 1e-9) + + // Save replaces: the critical asset is unlinked, low and high are linked. + r.Score = 0.5 * 2 * 1.5 + require.NoError(t, store.SaveWithAssets(ctx, r, []*domain.Asset{low, high}, false)) + got := links(r.ID) + assert.ElementsMatch(t, []uuid.UUID{low.ID, high.ID}, got) + + // Another tenant's asset is refused before anything is written. + r2 := &domain.Risk{ID: uuid.New(), TenantID: tenant, Title: "pg-792-foreign", Probability: 0.5, Impact: 2} + r2.SetState(domain.StateDraft) + require.Error(t, store.SaveWithAssets(ctx, r2, []*domain.Asset{foreign}, true)) + assert.Zero(t, riskCount(r2.ID)) + + // A link write Postgres refuses rolls the risk row back with it. + for _, sql := range []string{ + `CREATE FUNCTION pg_temp.refuse_link() RETURNS trigger LANGUAGE plpgsql AS + $$ BEGIN RAISE EXCEPTION 'link refused'; END $$`, + `CREATE TRIGGER refuse_link_792 BEFORE INSERT ON risk_assets + FOR EACH ROW EXECUTE FUNCTION pg_temp.refuse_link()`, + } { + require.NoError(t, tx.Exec(sql).Error, sql) + } + r3 := &domain.Risk{ID: uuid.New(), TenantID: tenant, Title: "pg-792-atomic", Probability: 0.5, Impact: 2, Score: 3} + r3.SetState(domain.StateDraft) + require.Error(t, store.SaveWithAssets(ctx, r3, []*domain.Asset{crit}, true)) + assert.Zero(t, riskCount(r3.ID), "the risk row must not survive a failed link") + + return rollback + }) + require.ErrorIs(t, err, rollback) +} diff --git a/backend/internal/infrastructure/repository/gorm_risk_repository.go b/backend/internal/infrastructure/repository/gorm_risk_repository.go index 8902c26a..4901e510 100644 --- a/backend/internal/infrastructure/repository/gorm_risk_repository.go +++ b/backend/internal/infrastructure/repository/gorm_risk_repository.go @@ -22,6 +22,7 @@ import ( "gorm.io/gorm/clause" "github.com/opendefender/openrisk/internal/application/dashboard" + riskapp "github.com/opendefender/openrisk/internal/application/risk" "github.com/opendefender/openrisk/internal/domain" ) @@ -34,6 +35,35 @@ type GormRiskRepository struct { db *gorm.DB } +// RunRiskTx runs fn inside one transaction with a risk repository bound to it, +// for use cases that must write several risks or none (the CSV import, #755). +// The asset store shares the transaction, so asset links roll back with the +// risks. It satisfies application/risk.RiskTxRunner. +func RunRiskTx(db *gorm.DB) func(ctx context.Context, fn func(repo domain.RiskRepository, assets riskapp.RiskAssetStore) error) error { + return func(ctx context.Context, fn func(repo domain.RiskRepository, assets riskapp.RiskAssetStore) error) error { + return db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + return fn(NewGormRiskRepository(tx), NewGormRiskAssetStore(tx)) + }) + } +} + +// ListImportAssetRefs lists a tenant's live assets by id and name, for the +// "assets" column of the CSV import (#755). It satisfies +// application/risk.ImportAssetLister. +func ListImportAssetRefs(db *gorm.DB) func(ctx context.Context, tenantID uuid.UUID) ([]riskapp.ImportAssetRef, error) { + return func(ctx context.Context, tenantID uuid.UUID) ([]riskapp.ImportAssetRef, error) { + if tenantID == uuid.Nil { + return nil, fmt.Errorf("tenant_id is required") + } + var refs []riskapp.ImportAssetRef + err := db.WithContext(ctx).Model(&domain.Asset{}). + Select("id", "name"). + Where("tenant_id = ?", tenantID). + Scan(&refs).Error + return refs, err + } +} + // NewGormRiskRepository creates a new GORM-backed risk repository. func NewGormRiskRepository(db *gorm.DB) *GormRiskRepository { return &GormRiskRepository{db: db} @@ -595,19 +625,14 @@ func (r *GormRiskRepository) GetRisksByAssetID(ctx context.Context, assetID uuid return nil, fmt.Errorf("failed to load linked asset criticalities: %w", err) } - factorSums := make(map[uuid.UUID]float64, len(riskIDs)) - factorCounts := make(map[uuid.UUID]int, len(riskIDs)) + crits := make(map[uuid.UUID][]domain.AssetCriticality, len(riskIDs)) for _, link := range links { - factorSums[link.RiskID] += link.Criticality.ScoreFactor() - factorCounts[link.RiskID]++ + crits[link.RiskID] = append(crits[link.RiskID], link.Criticality) } risks := make([]domain.RiskForScoring, 0, len(riskRows)) for _, row := range riskRows { - factor := 1.0 - if count := factorCounts[row.ID]; count > 0 { - factor = factorSums[row.ID] / float64(count) - } + factor := domain.RiskAssetCriticality(crits[row.ID]) risks = append(risks, domain.RiskForScoring{ ID: row.ID, TenantID: row.TenantID, diff --git a/docs/DECISIONS.md b/docs/DECISIONS.md index 7914a25a..5a9f35f6 100644 --- a/docs/DECISIONS.md +++ b/docs/DECISIONS.md @@ -5,34 +5,34 @@ recommends, and surfaces these in the daily brief. Run `/decide` to clear them. ## Open -### D-061 — 3D tilt on the overall score card (#751 phase 5) · raised 2026-09-29 -**Question** — #751 phase 5 lists "`3d-tilt` reserved for high-value visual elements (e.g. -overall score card)". Do we ship it? - -**Recommendation — B, drop it.** art-director and ux-designer reached the same answer -independently: -- The tilt skews the score arc, which is the data itself, on the one figure a CISO reads at - a glance and an auditor screenshots. -- A card that follows the pointer signals "the whole card is clickable". Only the button - inside `ScoreGauge` is. -- It works only with a mouse, which is a second rendering path with no keyboard or touch - equivalent, and pointer-tracked parallax is a vestibular trigger. -- It would touch `ScoreGauge`, which #824 (phase 3) owns. - -**Options** -- **A — ship a restrained tilt.** At most 2° per axis and an 800px perspective, only on - `(hover:hover) and (pointer:fine)` under `motion-safe`, with no glare and no moving - shadow. It goes on the `DashboardPage.tsx` wrapper, never inside `ScoreGauge`. About half - a day, plus a live pass. -- **B — drop it (recommended).** The item is closed as "evaluated, rejected" on #751. - -**Cost of delay** — none. Phase 5 ships without it, and A can be added later without -rework. - -**Raised by** — #751 phase 5 (art-director + ux-designer specs, 2026-09-29). - ## Resolved +### D-063 — 3D tilt on the overall score card: ship it · decided 2026-10-01 +**Decided (owner)** — Option A: ship a restrained tilt now. The owner's answer was +"livre maintenant la phase 5", read as "ship the tilt", so the agents' recommendation (B, +drop it) is overruled. +**Rationale (owner)** — The tilt was part of the phase 5 plan, and the restraints in option +A answer the objections raised: at most 2° per axis, 800px perspective, only with +`(hover:hover) and (pointer:fine)` under `motion-safe`, no glare and no moving shadow. +It sits on the `DashboardPage.tsx` wrapper and never inside `ScoreGauge`. +**Consequence** — #751 is closed (phase 5 merged in #847), so the work is tracked in +**#856** (`status:ready`, milestone `ds-v1`) with those restraints as acceptance criteria. +**Numbering** — This entry was raised as "D-061" on 2026-09-29, a number already used by +the SSO/MFA decision below. It was renumbered D-063 so every number is unique. +**Unblocked** — #856. + +### D-062 — Asset-criticality term of a risk with no linked asset: 1.0 · decided 2026-10-01 +**Decided (owner)** — Option A: 1.0, neutral. A risk with no linked asset scores P × I. +**Rationale (owner)** — Matches the recommendation. It is what the Score Engine worker, +the handlers and the demo seed have always stored, so no stored score moves and no +backfill is needed. Only the breakdown and score-working views, which displayed 1.5, +change. +**Consequence** — `domain.NoAssetCriticalityFactor = 1.0` and +`domain.RiskAssetCriticality` are the single derivation for every score writer and +reader (PR #855). Moving to 1.5 later means changing that constant and writing a backfill +of stored scores. +**Unblocked** — PR #855 / #792 can merge with no open question. + ### D-061 — SSO accounts turn MFA off with an authenticator code · decided 2026-09-30 **Decided (owner)** — An account with no local password (identity-provider sign-in) confirms turning MFA off with a **current TOTP code** from its authenticator app. Backup diff --git a/frontend/src/components/shared/AutoDetectedBadge.tsx b/frontend/src/components/shared/AutoDetectedBadge.tsx deleted file mode 100644 index 78b07c61..00000000 --- a/frontend/src/components/shared/AutoDetectedBadge.tsx +++ /dev/null @@ -1,77 +0,0 @@ -// Copyright (c) 2026 OpenDefender Contributors -// SPDX-License-Identifier: AGPL-3.0-only -// This program is free software: you can redistribute it and/or modify it under -// the terms of the GNU Affero General Public License v3.0 (see LICENSE). - -import { useFormat } from '../../hooks/useI18n'; -import { motion } from '../../shared/motion'; -import { Zap } from 'lucide-react'; -import { clsx, type ClassValue } from 'clsx'; -import { twMerge } from 'tailwind-merge'; - -function cn(...inputs: ClassValue[]) { - return twMerge(clsx(inputs)); -} - -interface AutoDetectedBadgeProps { - detectedAt?: string; - scanId?: string; - size?: 'sm' | 'md'; - className?: string; -} - -export const AutoDetectedBadge = ({ - detectedAt, - scanId, - size = 'md', - className, -}: AutoDetectedBadgeProps) => { - const fmt = useFormat(); - const sizeClasses = { - sm: 'px-2 py-1 text-xs gap-1', - md: 'px-3 py-1.5 text-sm gap-1.5', - }; - - const iconSizes = { - sm: 12, - md: 14, - }; - - const formatTime = (isoString?: string) => { - if (!isoString) return ''; - const date = new Date(isoString); - return fmt.dateTime(date, { - day: '2-digit', - month: '2-digit', - hour: '2-digit', - minute: '2-digit', - }); - }; - - const tooltipText = detectedAt - ? `Détecté automatiquement par le scanner le ${formatTime(detectedAt)}${scanId ? ` (scan #${scanId})` : ''}` - : 'Auto-détecté par le scanner'; - - return ( - - - Auto - - {/* Tooltip */} -
- {tooltipText} -
-
- - ); -}; diff --git a/frontend/src/components/shared/ProgressBar.tsx b/frontend/src/components/shared/ProgressBar.tsx deleted file mode 100644 index 44ee33f2..00000000 --- a/frontend/src/components/shared/ProgressBar.tsx +++ /dev/null @@ -1,70 +0,0 @@ -// Copyright (c) 2026 OpenDefender Contributors -// SPDX-License-Identifier: AGPL-3.0-only -// This program is free software: you can redistribute it and/or modify it under -// the terms of the GNU Affero General Public License v3.0 (see LICENSE). - -import { motion } from '../../shared/motion'; -import { clsx, type ClassValue } from 'clsx'; -import { twMerge } from 'tailwind-merge'; - -function cn(...inputs: ClassValue[]) { - return twMerge(clsx(inputs)); -} - -interface ProgressBarProps { - value: number; - max: number; - label?: string; - showPercentage?: boolean; - size?: 'sm' | 'md' | 'lg'; - variant?: 'default' | 'success' | 'warning' | 'danger'; - animated?: boolean; - className?: string; -} - -export const ProgressBar = ({ - value, - max, - label, - showPercentage = true, - size = 'md', - variant = 'default', - animated = true, - className, -}: ProgressBarProps) => { - const percentage = Math.min(Math.round((value / max) * 100), 100); - - const sizeClasses = { - sm: 'h-1.5', - md: 'h-2.5', - lg: 'h-3', - }; - - const variantClasses = { - default: 'bg-accent', - success: 'bg-success', - warning: 'bg-warning', - danger: 'bg-danger', - }; - - return ( -
- {(label || showPercentage) && ( -
- {label && {label}} - {showPercentage && ( - {percentage}% - )} -
- )} -
- -
-
- ); -}; diff --git a/frontend/src/components/shared/RiskBadge.tsx b/frontend/src/components/shared/RiskBadge.tsx deleted file mode 100644 index 4a52fd38..00000000 --- a/frontend/src/components/shared/RiskBadge.tsx +++ /dev/null @@ -1,102 +0,0 @@ -// Copyright (c) 2026 OpenDefender Contributors -// SPDX-License-Identifier: AGPL-3.0-only -// This program is free software: you can redistribute it and/or modify it under -// the terms of the GNU Affero General Public License v3.0 (see LICENSE). - -import { motion } from '../../shared/motion'; -import { AlertCircle, AlertTriangle, Info, AlertOctagon } from 'lucide-react'; -import { clsx, type ClassValue } from 'clsx'; -import { twMerge } from 'tailwind-merge'; - -function cn(...inputs: ClassValue[]) { - return twMerge(clsx(inputs)); -} - -type RiskLevel = 'CRITICAL' | 'HIGH' | 'MEDIUM' | 'LOW'; - -interface RiskBadgeProps { - // Accept any string: the backend sends risk.level as lowercase ("medium"), - // while other call sites pass the uppercase RiskLevel union. We normalize - // below instead of trusting the caller to have the exact casing. - level: RiskLevel | string; - animated?: boolean; - size?: 'sm' | 'md' | 'lg'; - className?: string; -} - -const RISK_CONFIGS = { - CRITICAL: { - bg: 'bg-danger/20', - border: 'border-danger/50', - text: 'text-danger-text', - icon: AlertOctagon, - label: 'Critique', - }, - HIGH: { - bg: 'bg-warning/20', - border: 'border-warning/50', - text: 'text-warning-text', - icon: AlertTriangle, - label: 'Élevé', - }, - MEDIUM: { - bg: 'bg-warning/20', - border: 'border-warning/50', - text: 'text-warning-text', - icon: AlertCircle, - label: 'Moyen', - }, - LOW: { - bg: 'bg-success/20', - border: 'border-success/50', - text: 'text-success-text', - icon: Info, - label: 'Bas', - }, -} as const; - -// getRiskConfig normalizes any casing and always returns a valid config -// (defaults to LOW) — an unknown level must never crash the badge, which -// previously white-screened the whole Risk drawer when the backend sent a -// lowercase level like "medium". -const getRiskConfig = (level: RiskLevel | string) => { - const key = String(level ?? '').toUpperCase() as keyof typeof RISK_CONFIGS; - return RISK_CONFIGS[key] ?? RISK_CONFIGS.LOW; -}; - -export const RiskBadge = ({ level, animated = true, size = 'md', className }: RiskBadgeProps) => { - const config = getRiskConfig(level); - const Icon = config.icon; - - const sizeClasses = { - sm: 'px-2 py-1 text-xs', - md: 'px-3 py-1.5 text-sm', - lg: 'px-4 py-2 text-base', - }; - - const iconSizes = { - sm: 14, - md: 16, - lg: 20, - }; - - return ( - - - {config.label} - - ); -}; diff --git a/frontend/src/components/shared/SkeletonTable.tsx b/frontend/src/components/shared/SkeletonTable.tsx deleted file mode 100644 index ab65a07d..00000000 --- a/frontend/src/components/shared/SkeletonTable.tsx +++ /dev/null @@ -1,65 +0,0 @@ -// Copyright (c) 2026 OpenDefender Contributors -// SPDX-License-Identifier: AGPL-3.0-only -// This program is free software: you can redistribute it and/or modify it under -// the terms of the GNU Affero General Public License v3.0 (see LICENSE). - -import { motion, type Variants } from '../../shared/motion'; -import { clsx, type ClassValue } from 'clsx'; -import { twMerge } from 'tailwind-merge'; - -function cn(...inputs: ClassValue[]) { - return twMerge(clsx(inputs)); -} - -interface SkeletonTableProps { - rows?: number; - columns?: number; - className?: string; -} - -export const SkeletonTable = ({ rows = 5, columns = 6, className }: SkeletonTableProps) => { - const pulseVariants: Variants = { - animate: { - opacity: [0.5, 1, 0.5], - transition: { - duration: 2, - repeat: Infinity, - ease: 'easeInOut', - }, - }, - }; - - return ( -
- {/* Header */} -
- {Array.from({ length: columns }).map((_, i) => ( - - ))} -
- - {/* Rows */} - {Array.from({ length: rows }).map((_, rowIdx) => ( -
- {Array.from({ length: columns }).map((_, colIdx) => ( - - ))} -
- ))} -
- ); -}; diff --git a/frontend/src/components/shared/index.ts b/frontend/src/components/shared/index.ts deleted file mode 100644 index 1fd3a5b6..00000000 --- a/frontend/src/components/shared/index.ts +++ /dev/null @@ -1,17 +0,0 @@ -// Copyright (c) 2026 OpenDefender Contributors -// SPDX-License-Identifier: AGPL-3.0-only -// This program is free software: you can redistribute it and/or modify it under -// the terms of the GNU Affero General Public License v3.0 (see LICENSE). - -// Design System Components -export { RiskBadge } from './RiskBadge'; -export type {} from './RiskBadge'; - -export { SkeletonTable } from './SkeletonTable'; -export type {} from './SkeletonTable'; - -export { ProgressBar } from './ProgressBar'; -export type {} from './ProgressBar'; - -export { AutoDetectedBadge } from './AutoDetectedBadge'; -export type {} from './AutoDetectedBadge'; diff --git a/frontend/src/features/risks/__tests__/importRisks.test.tsx b/frontend/src/features/risks/__tests__/importRisks.test.tsx new file mode 100644 index 00000000..f49450be --- /dev/null +++ b/frontend/src/features/risks/__tests__/importRisks.test.tsx @@ -0,0 +1,164 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// +// #755 — the import page reports exactly what the server answered. A refused +// file lists every error by line and column and never shows a success; a +// success toast appears only when the server created at least one risk. + +import { describe, expect, it, vi, beforeEach } from 'vitest'; +import { render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { MemoryRouter } from 'react-router'; +import { AxiosError, AxiosHeaders, type AxiosResponse } from 'axios'; + +const post = vi.fn(); +const toastSuccess = vi.fn(); +const fetchRisks = vi.fn(); + +vi.mock('../../../lib/api', () => ({ + api: { post: (...a: unknown[]) => post(...a), defaults: { baseURL: '' } }, +})); +vi.mock('../../../hooks/useToast', () => ({ + useToast: () => ({ success: toastSuccess, error: vi.fn(), promise: vi.fn() }), +})); +vi.mock('../../../hooks/useRiskStore', () => ({ + useRiskStore: () => ({ fetchRisks }), +})); + +import { ImportRisksPage } from '../../../pages/ImportRisks'; +import { importErrorMessage, importFileSchema } from '../importRisksSchema'; +import { catalogs, translate } from '../../../i18n'; +import { useUIStore } from '../../../store/uiStore'; + +// The catalogue translator the page gets from useI18n, bound to one locale. +const tIn = (locale: 'fr' | 'en') => (key: string, params?: Record) => + translate(catalogs, locale, key, { params }); + +function httpError(status: number, data: unknown): AxiosError { + const response = { + status, + data, + statusText: '', + headers: {}, + config: { headers: new AxiosHeaders() }, + } as AxiosResponse; + return new AxiosError('request failed', String(status), undefined, undefined, response); +} + +async function chooseAndImport( + name = 'register.csv', + content = 'title,probability,impact\nA,0.5,5\n', +) { + const user = userEvent.setup(); + render( + + + , + ); + await user.upload( + screen.getByTestId('import-file-input'), + new File([content], name, { type: 'text/csv' }), + ); + await user.click(screen.getByRole('button', { name: /^import$/i })); +} + +describe('ImportRisksPage', () => { + beforeEach(() => { + useUIStore.getState().setLang('en'); + post.mockReset(); + toastSuccess.mockReset(); + fetchRisks.mockReset(); + }); + + it('shows the count the server created and refreshes the register', async () => { + post.mockResolvedValue({ + data: { created: 3, rejected: 0, risk_ids: ['a', 'b', 'c'], errors: [] }, + }); + await chooseAndImport(); + + expect(await screen.findByText('3 risk(s) imported.')).toBeInTheDocument(); + expect(toastSuccess).toHaveBeenCalledTimes(1); + expect(fetchRisks).toHaveBeenCalledTimes(1); + }); + + it('lists every error by line and column and shows no success when the file is refused', async () => { + post.mockRejectedValue( + httpError(422, { + error: 'validation_failed', + created: 0, + rejected: 2, + risk_ids: [], + errors: [ + { + line: 3, + column: 'probability', + message: 'probability must be between 0 and 1 (got 3)', + }, + { line: 5, column: 'title', message: 'title is required' }, + ], + }), + ); + await chooseAndImport(); + + const panel = await screen.findByTestId('import-outcome-error'); + expect(panel).toHaveTextContent('Nothing was imported: 2 error(s) on 2 row(s).'); + expect(panel).toHaveTextContent('probability must be between 0 and 1 (got 3)'); + expect(screen.getByRole('cell', { name: '5' })).toBeInTheDocument(); + expect(screen.getByRole('cell', { name: 'title' })).toBeInTheDocument(); + expect(toastSuccess).not.toHaveBeenCalled(); + expect(fetchRisks).not.toHaveBeenCalled(); + }); + + it('explains a plan limit instead of a generic failure', async () => { + post.mockRejectedValue(httpError(402, { code: 'limit_reached', requested: 40, remaining: 5 })); + await chooseAndImport(); + + const panel = await screen.findByTestId('import-outcome-error'); + expect(panel).toHaveTextContent('The file has 40 risk(s); your plan allows 5 more.'); + expect(toastSuccess).not.toHaveBeenCalled(); + }); + + it('refuses a non-CSV file before sending anything', async () => { + const user = userEvent.setup({ applyAccept: false }); + render( + + + , + ); + await user.upload(screen.getByTestId('import-file-input'), new File(['x'], 'register.xlsx')); + + expect(await screen.findByRole('alert')).toHaveTextContent('Only CSV files are accepted.'); + expect(screen.queryByRole('button', { name: /^import$/i })).not.toBeInTheDocument(); + await waitFor(() => expect(post).not.toHaveBeenCalled()); + }); +}); + +describe('importFileSchema', () => { + it('accepts a CSV and refuses empty or oversized files', () => { + const schema = importFileSchema(tIn('en')); + expect(schema.safeParse(new File(['a'], 'r.CSV')).success).toBe(true); + expect(schema.safeParse(new File([], 'r.csv')).success).toBe(false); + expect(schema.safeParse(new File([new Uint8Array(2 * 1024 * 1024 + 1)], 'r.csv')).success).toBe( + false, + ); + }); + + it('renders server error codes in the reader’s language and falls back to the server text', () => { + const fr = tIn('fr'); + expect( + importErrorMessage( + { + line: 3, + column: 'probability', + code: 'out_of_range', + params: { min: '0', max: '1', value: '3' }, + message: 'probability must be between 0 and 1 (got 3)', + }, + fr, + ), + ).toBe('Doit être entre 0 et 1 (valeur : 3).'); + expect( + importErrorMessage({ line: 2, code: 'some_future_code', message: 'server text' }, fr), + ).toBe('server text'); + }); +}); diff --git a/frontend/src/features/risks/components/ScoreWorking.tsx b/frontend/src/features/risks/components/ScoreWorking.tsx index ac935d7e..8ecb50f7 100644 --- a/frontend/src/features/risks/components/ScoreWorking.tsx +++ b/frontend/src/features/risks/components/ScoreWorking.tsx @@ -179,8 +179,8 @@ export function ScoreWorking({ riskId, storedScore }: { riskId: string; storedSc return (
{tr( - 'Aucun actif lié : le moteur applique sa valeur par défaut (criticité moyenne).', - 'No linked asset: the engine applies its default (medium criticality).', + 'Aucun actif lié : le facteur est neutre (1,0), le score vaut P × I.', + 'No linked asset: the factor is neutral (1.0), so the score is P × I.', )}
); diff --git a/frontend/src/features/risks/importRisksSchema.ts b/frontend/src/features/risks/importRisksSchema.ts new file mode 100644 index 00000000..80a47ccc --- /dev/null +++ b/frontend/src/features/risks/importRisksSchema.ts @@ -0,0 +1,27 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). +// +// The risk import (#755). The contract shared with the asset import lives in +// shared/csvImport; this file keeps the risk template and re-exports the rest +// for the risk page and its tests. + +export { + MAX_IMPORT_BYTES, + importErrorMessage, + importFileSchema, + importLimitSchema, + importRejectedSchema, + importRowErrorSchema, + importSuccessSchema, + type ImportRowError, +} from '../../shared/csvImport/csvImportSchema'; + +/** The current template: the product's scales, P in [0,1] and I in [0,10]. */ +export const IMPORT_TEMPLATE = [ + 'title,description,probability,impact,tags,frameworks,assets', + '"Phishing campaign against finance staff","Credential theft leading to fraudulent transfers",0.6,8,"email;people","ISO27001"', + '"Ransomware on file servers","Encryption of shared drives, no tested restore",0.3,10,"backup","ISO27001;NIST CSF"', + '"Cloud provider outage","Loss of the hosted CRM for more than 24 hours",0.2,5,"supplier",', +].join('\n'); diff --git a/frontend/src/locales/en.json b/frontend/src/locales/en.json index b576dcb9..67c6a15d 100644 --- a/frontend/src/locales/en.json +++ b/frontend/src/locales/en.json @@ -93,19 +93,13 @@ "bulkChangeStatus": "Change Status", "bulkAssignTo": "Assign To", "bulkAddTags": "Add Tags", - "import": "Import Risks", "export": "Export Risks", "importFile": "Import File", "exportFormat": "Export Format", "csv": "CSV", "json": "JSON", "xlsx": "Excel", - "dragDropHint": "Drag and drop your file here or click to browse", - "importPreview": "Import Preview", "importResults": "Import Results", - "successCount": "{count} risk(s) imported successfully", - "errorCount": "{count} error(s) during import", - "templateDownload": "Download CSV Template", "noRisks": "No Risks Found", "noRisksDescription": "Start by creating your first risk to manage it here.", "createFirstRisk": "Create My First Risk", @@ -157,9 +151,7 @@ "failedToCreateRisk": "Failed to create risk", "failedToUpdateRisk": "Failed to update risk", "failedToDeleteRisk": "Failed to delete risk", - "failedToImportRisks": "Failed to import risks", "failedToExportRisks": "Failed to export risks", - "invalidFile": "Invalid file format", "maxFileSizeExceeded": "File size exceeds maximum limit", "validationError": "Validation error", "serverError": "Server error", @@ -182,8 +174,6 @@ "mitigationAddedSuccess": "Mitigation plan added successfully", "mitigationUpdatedSuccess": "Mitigation plan updated successfully", "mitigationDeletedSuccess": "Mitigation plan deleted successfully", - "importStarted": "Import in progress...", - "importCompleted": "Import completed", "exportCompleted": "Export completed" }, "filters": { @@ -922,5 +912,65 @@ "switchError": "The switch failed. You are still in {name}.", "onlyOne": "You belong to this organization only. An invitation from another organization will add it here.", "settings": "Organization settings" + }, + "csvImport": { + "expectedFormat": "Expected format", + "downloadTemplate": "Download the template", + "dropHere": "Drop a CSV file here or click to choose one", + "limits": "CSV, up to 2 MB and 1000 rows", + "removeFile": "Remove the file", + "import": "Import", + "invalidFile": "Invalid file.", + "onlyCsv": "Only CSV files are accepted.", + "fileEmpty": "The file is empty.", + "tooLarge": "The file is larger than 2 MB. Split it into several files.", + "unexpectedResponse": "Unexpected response from the server.", + "serverFailed": "The server could not process the file. Try again in a moment.", + "nothingImported": "Nothing was imported.", + "rejectedRows": "Nothing was imported: {errors} error(s) on {rows} row(s).", + "fileRefused": "Nothing was imported: the file was refused.", + "fixRows": "Fix these rows in your file, then import it again.", + "line": "Line", + "column": "Column", + "problem": "Problem", + "file": "File", + "seePlans": "See plans", + "errors": { + "file_empty": "The file is empty.", + "not_utf8": "The file is not UTF-8 text: save it as \"CSV UTF-8\".", + "header_unreadable": "The header cannot be read.", + "unknown_column": "Unknown column \"{column}\". Accepted columns: {accepted}.", + "duplicate_column": "Column \"{column}\" appears twice.", + "missing_column": "Required column \"{column}\" is missing.", + "line_unreadable": "This line cannot be read (unclosed quote?).", + "too_many_rows": "The file has more than {max} rows: split it into several files.", + "cell_count": "The line has {cells} cells but the header has {header}.", + "required": "A value is required.", + "too_long": "At most {max} characters.", + "not_a_number": "\"{value}\" is not a number.", + "out_of_range": "Must be between {min} and {max} (got {value}).", + "no_rows": "The file has a header but no rows.", + "unknown_asset": "No asset named \"{value}\" in the inventory.", + "ambiguous_asset": "{count} assets are named \"{value}\": use its id instead.", + "assets_unavailable": "Linking assets is not available on this server: remove the assets column.", + "legacy_scale": "This file uses the old 1–5 scale. OpenRisk expects probability between 0 and 1 and impact between 0 and 10. Download the current template and convert the values (probability 3/5 → 0.6, impact 4/5 → 8)." + }, + "risks": { + "back": "Risk register", + "title": "Import risks", + "intro": "A CSV file, one row per risk. Every row is checked before anything is imported: if a single row is invalid, nothing is imported and every error is listed.", + "colTitle": "Required, at most 255 characters.", + "colProbability": "Required, between 0 and 1 (e.g. 0.6).", + "colImpact": "Required, between 0 and 10 (e.g. 8).", + "colOptional": "Optional. Separate several values with \";\".", + "colAssets": "Optional. Names or ids of assets in your inventory, separated by \";\". Their criticality is part of the score; a name that is unknown or shared by several assets refuses the row.", + "note": "Excel exports using \";\" and a decimal comma are accepted. Files on the old 1–5 scale are refused: convert them (probability 3/5 → 0.6, impact 4/5 → 8).", + "open": "Open the register", + "created": "{count} risk(s) imported", + "emptyFile": "The file contained no risks.", + "forbidden": "You are not allowed to create risks.", + "limitTitle": "Nothing was imported: this file exceeds your plan’s risk limit.", + "limitDetail": "The file has {requested} risk(s); your plan allows {remaining} more." + } } } diff --git a/frontend/src/locales/fr.json b/frontend/src/locales/fr.json index 25351ae7..bc206e7e 100644 --- a/frontend/src/locales/fr.json +++ b/frontend/src/locales/fr.json @@ -93,19 +93,13 @@ "bulkChangeStatus": "Changer le statut", "bulkAssignTo": "Assigner à", "bulkAddTags": "Ajouter des étiquettes", - "import": "Importer des risques", "export": "Exporter les risques", "importFile": "Importer un fichier", "exportFormat": "Format d'export", "csv": "CSV", "json": "JSON", "xlsx": "Excel", - "dragDropHint": "Glissez-déposez votre fichier ici ou cliquez pour parcourir", - "importPreview": "Aperçu de l'import", "importResults": "Résultats de l'import", - "successCount": "{count} risque(s) importé(s) avec succès", - "errorCount": "{count} erreur(s) lors de l'import", - "templateDownload": "Télécharger le modèle CSV", "noRisks": "Aucun risque trouvé", "noRisksDescription": "Commencez par créer votre premier risque pour le gérer ici.", "createFirstRisk": "Créer mon premier risque", @@ -157,9 +151,7 @@ "failedToCreateRisk": "Échec de la création du risque", "failedToUpdateRisk": "Échec de la mise à jour du risque", "failedToDeleteRisk": "Échec de la suppression du risque", - "failedToImportRisks": "Échec de l'import des risques", "failedToExportRisks": "Échec de l'export des risques", - "invalidFile": "Format de fichier invalide", "maxFileSizeExceeded": "La taille du fichier dépasse la limite", "validationError": "Erreur de validation", "serverError": "Erreur serveur", @@ -182,8 +174,6 @@ "mitigationAddedSuccess": "Plan d'atténuation ajouté avec succès", "mitigationUpdatedSuccess": "Plan d'atténuation mis à jour avec succès", "mitigationDeletedSuccess": "Plan d'atténuation supprimé avec succès", - "importStarted": "Import en cours...", - "importCompleted": "Import terminé", "exportCompleted": "Export terminé" }, "filters": { @@ -922,5 +912,65 @@ "switchError": "Le changement d'organisation a échoué. Vous êtes toujours dans {name}.", "onlyOne": "Vous n'appartenez qu'à cette organisation. Une invitation d'une autre organisation l'ajoutera ici.", "settings": "Paramètres de l'organisation" + }, + "csvImport": { + "expectedFormat": "Format attendu", + "downloadTemplate": "Télécharger le modèle", + "dropHere": "Glissez un fichier CSV ici ou cliquez pour le choisir", + "limits": "CSV, 2 Mo et 1000 lignes au plus", + "removeFile": "Retirer le fichier", + "import": "Importer", + "invalidFile": "Fichier invalide.", + "onlyCsv": "Seuls les fichiers CSV sont acceptés.", + "fileEmpty": "Le fichier est vide.", + "tooLarge": "Le fichier dépasse 2 Mo. Découpez-le en plusieurs fichiers.", + "unexpectedResponse": "Réponse inattendue du serveur.", + "serverFailed": "Le serveur n’a pas pu traiter le fichier. Réessayez dans un instant.", + "nothingImported": "Rien n’a été importé.", + "rejectedRows": "Rien n’a été importé : {errors} erreur(s) sur {rows} ligne(s).", + "fileRefused": "Rien n’a été importé : le fichier a été refusé.", + "fixRows": "Corrigez ces lignes dans votre fichier puis importez-le à nouveau.", + "line": "Ligne", + "column": "Colonne", + "problem": "Problème", + "file": "Fichier", + "seePlans": "Voir les plans", + "errors": { + "file_empty": "Le fichier est vide.", + "not_utf8": "Le fichier n’est pas en UTF-8 : enregistrez-le au format « CSV UTF-8 ».", + "header_unreadable": "L’en-tête est illisible.", + "unknown_column": "Colonne inconnue « {column} ». Colonnes acceptées : {accepted}.", + "duplicate_column": "La colonne « {column} » apparaît deux fois.", + "missing_column": "La colonne obligatoire « {column} » est absente.", + "line_unreadable": "Cette ligne est illisible (guillemet non fermé ?).", + "too_many_rows": "Le fichier dépasse {max} lignes : découpez-le en plusieurs fichiers.", + "cell_count": "La ligne a {cells} cellules, l’en-tête en a {header}.", + "required": "Valeur obligatoire.", + "too_long": "{max} caractères au plus.", + "not_a_number": "« {value} » n’est pas un nombre.", + "out_of_range": "Doit être entre {min} et {max} (valeur : {value}).", + "no_rows": "Le fichier a un en-tête mais aucune ligne.", + "unknown_asset": "Aucun actif « {value} » dans l’inventaire.", + "ambiguous_asset": "{count} actifs s’appellent « {value} » : indiquez son identifiant.", + "assets_unavailable": "La liaison aux actifs n’est pas disponible sur ce serveur : retirez la colonne assets.", + "legacy_scale": "Ce fichier utilise l’ancienne échelle 1–5. OpenRisk attend une probabilité entre 0 et 1 et un impact entre 0 et 10. Téléchargez le modèle actuel et convertissez les valeurs (probabilité 3/5 → 0,6 ; impact 4/5 → 8)." + }, + "risks": { + "back": "Registre des risques", + "title": "Importer des risques", + "intro": "Un fichier CSV, une ligne par risque. Toutes les lignes sont vérifiées avant l’import : si une seule est invalide, rien n’est importé et chaque erreur vous est indiquée.", + "colTitle": "Obligatoire, 255 caractères au plus.", + "colProbability": "Obligatoire, entre 0 et 1 (ex. 0.6).", + "colImpact": "Obligatoire, entre 0 et 10 (ex. 8).", + "colOptional": "Facultatifs. Plusieurs valeurs séparées par « ; ».", + "colAssets": "Facultatif. Noms ou identifiants d’actifs de votre inventaire, séparés par « ; ». Leur criticité entre dans le score ; un nom inconnu ou porté par plusieurs actifs fait refuser la ligne.", + "note": "Les exports Excel en « ; » avec virgule décimale sont acceptés. Les fichiers sur l’ancienne échelle 1–5 sont refusés : convertissez-les (probabilité 3/5 → 0,6 ; impact 4/5 → 8).", + "open": "Voir le registre", + "created": "{count} risque(s) importé(s)", + "emptyFile": "Le fichier ne contenait aucun risque.", + "forbidden": "Vous n’avez pas le droit de créer des risques.", + "limitTitle": "Rien n’a été importé : ce fichier dépasse la limite de risques de votre plan.", + "limitDetail": "Le fichier contient {requested} risque(s) ; votre plan en permet encore {remaining}." + } } } diff --git a/frontend/src/pages/ImportRisks.tsx b/frontend/src/pages/ImportRisks.tsx index 6c022d62..4b2384ff 100644 --- a/frontend/src/pages/ImportRisks.tsx +++ b/frontend/src/pages/ImportRisks.tsx @@ -3,387 +3,44 @@ // This program is free software: you can redistribute it and/or modify it under // the terms of the GNU Affero General Public License v3.0 (see LICENSE). -import { useState, useRef, useCallback } from 'react'; -import { Link } from 'react-router'; -import { motion, AnimatePresence } from '../shared/motion'; -import { - Upload, - AlertCircle, - CheckCircle2, - FileJson, - FileText, - FileSpreadsheet, - Download, - ArrowLeft, - X, -} from 'lucide-react'; -import { useToast } from '../hooks/useToast'; -import { useI18n, interpolate } from '../hooks/useI18n'; -import { Button } from '../shared/ds'; -import { api } from '../lib/api'; -import { useRiskStore } from '../hooks/useRiskStore'; -import { SkeletonTable } from '../components/shared'; -import { clsx, type ClassValue } from 'clsx'; -import { twMerge } from 'tailwind-merge'; - -function cn(...inputs: ClassValue[]) { - return twMerge(clsx(inputs)); -} - -interface ImportResult { - success: number; - failed: number; - errors: Array<{ row: number; message: string }>; -} +// CSV import of the risk register (#755). The page itself is the shared +// CsvImportPage; this file says what is particular to risks. -type DragState = 'idle' | 'dragging' | 'processing'; -type FileFormat = 'csv' | 'json' | 'xlsx'; +import { useRiskStore } from '../hooks/useRiskStore'; +import { useI18n } from '../hooks/useI18n'; +import { CsvImportPage, type CsvImportConfig } from '../shared/csvImport/CsvImportPage'; +import { IMPORT_TEMPLATE } from '../features/risks/importRisksSchema'; export const ImportRisksPage = () => { const { t } = useI18n(); - const { success, error, promise } = useToast(); const { fetchRisks } = useRiskStore(); - const [dragState, setDragState] = useState('idle'); - const [selectedFile, setSelectedFile] = useState(null); - const [preview, setPreview] = useState([]); - const [importResult, setImportResult] = useState(null); - const [isImporting, setIsImporting] = useState(false); - const [mappedColumns, setMappedColumns] = useState>({}); - const fileInputRef = useRef(null); - - // Get file icon based on format - const getFileIcon = (format: FileFormat) => { - switch (format) { - case 'json': - return ; - case 'xlsx': - return ; - default: - return ; - } - }; - - // Parse file and show preview - const handleFileSelect = useCallback( - async (file: File) => { - const format = file.name.split('.').pop()?.toLowerCase() as FileFormat | undefined; - - if (!['csv', 'json', 'xlsx'].includes(format || '')) { - error(t('errors.invalidFile')); - return; - } - - setSelectedFile(file); - setDragState('processing'); - - try { - let data: any[] = []; - - if (format === 'json') { - const text = await file.text(); - data = JSON.parse(text); - } else if (format === 'csv') { - // Simple CSV parser (production would use a library) - const text = await file.text(); - const lines = text.split('\n'); - const headers = lines[0].split(',').map((h) => h.trim()); - data = lines.slice(1).map((line) => { - const values = line.split(','); - return headers.reduce( - (acc, header, i) => { - acc[header] = values[i]?.trim() || ''; - return acc; - }, - {} as Record, - ); - }); - } else if (format === 'xlsx') { - error(t('common.loading')); // Placeholder - need excelize library - return; - } - - // Show first 10 rows as preview - setPreview(data.slice(0, 10)); - setDragState('idle'); - success(t('messages.importStarted')); - } catch (err) { - error(interpolate(t('errors.failedToImportRisks'), {})); - setDragState('idle'); - } - }, - [t, error, success], - ); - - // Handle drag and drop - const handleDragOver = (e: React.DragEvent) => { - e.preventDefault(); - setDragState('dragging'); - }; - - const handleDragLeave = () => { - setDragState('idle'); - }; - - const handleDrop = (e: React.DragEvent) => { - e.preventDefault(); - const files = e.dataTransfer.files; - if (files.length > 0) { - handleFileSelect(files[0]); - } - setDragState('idle'); - }; - - // Handle file input change - const handleFileInputChange = (e: React.ChangeEvent) => { - if (e.target.files?.length) { - handleFileSelect(e.target.files[0]); - } + const config: CsvImportConfig = { + endpoint: '/risks/import', + back: { to: '/risks', label: t('csvImport.risks.back') }, + title: t('csvImport.risks.title'), + intro: t('csvImport.risks.intro'), + columns: [ + { name: 'title', help: t('csvImport.risks.colTitle') }, + { name: 'probability', help: t('csvImport.risks.colProbability') }, + { name: 'impact', help: t('csvImport.risks.colImpact') }, + { name: 'description, tags, frameworks', help: t('csvImport.risks.colOptional') }, + { name: 'assets', help: t('csvImport.risks.colAssets') }, + ], + note: t('csvImport.risks.note'), + template: IMPORT_TEMPLATE, + templateFilename: 'openrisk-risks-template.csv', + open: { to: '/risks', label: t('csvImport.risks.open') }, + created: (count) => t('csvImport.risks.created', { count }), + emptyFile: t('csvImport.risks.emptyFile'), + forbidden: t('csvImport.risks.forbidden'), + limitTitle: t('csvImport.risks.limitTitle'), + limitDetail: (requested, remaining) => + t('csvImport.risks.limitDetail', { requested, remaining }), + onCreated: () => void fetchRisks(), }; - // Submit import - const handleImport = async () => { - if (!selectedFile) return; - - setIsImporting(true); - - try { - const formData = new FormData(); - formData.append('file', selectedFile); - - const importRequest = api.post('/risks/import', formData, { - headers: { 'Content-Type': 'multipart/form-data' }, - }); - - promise(importRequest, { - loading: t('messages.importStarted'), - success: t('messages.importCompleted'), - error: t('errors.failedToImportRisks'), - }); - - const response = await importRequest; - const result: ImportResult = response.data; - setImportResult(result); - - // Refresh risks list - await fetchRisks(); - } catch (err) { - console.error('Import failed:', err); - } finally { - setIsImporting(false); - } - }; - - // Download template - const handleDownloadTemplate = () => { - const template = `Title,Description,Probability,Impact,Status,Framework,Tags,Assets -"Web API Vulnerability","Unvalidated API endpoints",3,4,Open,OWASP,"API,Security","API-001" -"Database Compromise","SQL injection risk",4,5,Open,NIST,"Database","DB-001"`; - - const blob = new Blob([template], { type: 'text/csv' }); - const url = window.URL.createObjectURL(blob); - const a = document.createElement('a'); - a.href = url; - a.download = 'risks-template.csv'; - a.click(); - window.URL.revokeObjectURL(url); - }; - - return ( -
- {/* Header */} -
- - {t('risks.title')} - -

{t('risks.import')}

-

{t('risks.dragDropHint')}

-
- - {/* Main content area */} - {!importResult ? ( -
- {/* Drag & Drop Zone */} - fileInputRef.current?.click()} - > - {dragState === 'processing' ? ( -
-
- -
-

{t('common.loading')}

-
- ) : ( - <> -
- - - -
-

- {t('risks.dragDropHint')} -

-

CSV, JSON, XLSX

- - )} - - -
- - {/* Template Download */} -
- -
- - {/* Preview */} - {preview.length > 0 && ( - -
-

- {t('risks.importPreview')} -

- -
- - {/* Preview Table */} -
- - - - {Object.keys(preview[0] || {}) - .slice(0, 6) - .map((key) => ( - - ))} - - - - {preview.slice(0, 5).map((row, i) => ( - - {Object.values(row) - .slice(0, 6) - .map((val, j) => ( - - ))} - - ))} - -
- {key} -
- {String(val)} -
-
- - {/* Import Button */} -
- -
-
- )} -
- ) : ( - /* Results */ - - {importResult.success > 0 && ( -
- -
-

{t('messages.importCompleted')}

-

- {interpolate(t('risks.successCount'), { count: importResult.success })} -

-
-
- )} - - {importResult.errors.length > 0 && ( -
- -
-

- {interpolate(t('risks.errorCount'), { count: importResult.errors.length })} -

-
- {importResult.errors.slice(0, 5).map((err, i) => ( -

- Row {err.row}: {err.message} -

- ))} - {importResult.errors.length > 5 && ( -

- ...and {importResult.errors.length - 5} more -

- )} -
-
-
- )} - - {/* Done Button */} -
- -
-
- )} -
- ); + return ; }; export default ImportRisksPage; diff --git a/frontend/src/shared/csvImport/CsvImportPage.tsx b/frontend/src/shared/csvImport/CsvImportPage.tsx new file mode 100644 index 00000000..e4f11fcb --- /dev/null +++ b/frontend/src/shared/csvImport/CsvImportPage.tsx @@ -0,0 +1,381 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). + +// The CSV import page (#755), written once so every import gets the same +// contract. The server validates every row and writes all of them or none. This page reports exactly what the server answered: how many records +// were created, or, when the file was refused, the line, column and reason of +// every error. It never says "imported" unless the server created at least one. + +import { useRef, useState } from 'react'; +import { Link } from 'react-router'; +import axios from 'axios'; +import { AlertCircle, ArrowLeft, CheckCircle2, Download, FileText, Upload, X } from 'lucide-react'; + +import { Button, cn } from '../ds'; +import { api } from '../../lib/api'; +import { useToast } from '../../hooks/useToast'; +import { useI18n } from '../../hooks/useI18n'; +import { + importErrorMessage, + importFileSchema, + importLimitSchema, + importRejectedSchema, + importSuccessSchema, + type ImportRowError, + type T, +} from './csvImportSchema'; + +/** What differs from one import to the next. Strings are already + * in the reader's language: the page builds its config with its own t. */ +export interface CsvImportConfig { + endpoint: string; + back: { to: string; label: string }; + title: string; + intro: string; + columns: ReadonlyArray<{ name: string; help: string }>; + note?: string; + template: string; + templateFilename: string; + /** Where the created records can be seen. */ + open: { to: string; label: string }; + created: (n: number) => string; + emptyFile: string; + forbidden: string; + limitTitle: string; + limitDetail: (requested: number, remaining: number) => string; + /** Called after a successful import, e.g. to refresh a store. */ + onCreated?: () => void; +} + +type Outcome = + | { kind: 'created'; created: number } + | { kind: 'rejected'; rejected: number; errors: ImportRowError[] } + | { kind: 'limit'; requested?: number; remaining?: number } + | { kind: 'failed'; message: string }; + +export const CsvImportPage = ({ config }: { config: CsvImportConfig }) => { + const { t } = useI18n(); + const { success } = useToast(); + + const inputRef = useRef(null); + const [file, setFile] = useState(null); + const [fileError, setFileError] = useState(null); + const [dragging, setDragging] = useState(false); + const [submitting, setSubmitting] = useState(false); + const [outcome, setOutcome] = useState(null); + + const pick = (f: File | undefined) => { + setOutcome(null); + if (!f) return; + const parsed = importFileSchema(t).safeParse(f); + if (!parsed.success) { + setFile(null); + setFileError(parsed.error.issues[0]?.message ?? t('csvImport.invalidFile')); + return; + } + setFileError(null); + setFile(f); + }; + + const reset = () => { + setFile(null); + setFileError(null); + setOutcome(null); + if (inputRef.current) inputRef.current.value = ''; + }; + + const submit = async () => { + if (!file) return; + setSubmitting(true); + setOutcome(null); + try { + const fd = new FormData(); + fd.append('file', file); + const res = await api.post(config.endpoint, fd, { + headers: { 'Content-Type': 'multipart/form-data' }, + }); + const parsed = importSuccessSchema.safeParse(res.data); + if (!parsed.success) { + setOutcome({ + kind: 'failed', + message: t('csvImport.unexpectedResponse'), + }); + return; + } + // The server never answers 200 with zero rows, but if it ever did this + // page must not dress it up as a success. + if (parsed.data.created === 0) { + setOutcome({ kind: 'failed', message: config.emptyFile }); + return; + } + setOutcome({ kind: 'created', created: parsed.data.created }); + success(config.created(parsed.data.created)); + config.onCreated?.(); + } catch (err) { + setOutcome(toOutcome(err, t, config.forbidden)); + } finally { + setSubmitting(false); + } + }; + + const downloadTemplate = () => { + const blob = new Blob([config.template + '\n'], { type: 'text/csv;charset=utf-8' }); + const url = URL.createObjectURL(blob); + const a = document.createElement('a'); + a.href = url; + a.download = config.templateFilename; + a.click(); + URL.revokeObjectURL(url); + }; + + return ( +
+
+ + {config.back.label} + +

{config.title}

+

{config.intro}

+
+ + {/* Format reference */} +
+
+

+ {t('csvImport.expectedFormat')} +

+ +
+
+ {config.columns.map((c) => ( +
+
{c.name}
+
{c.help}
+
+ ))} +
+ {config.note &&

{config.note}

} +
+ + {/* File picker */} +
inputRef.current?.click()} + onKeyDown={(e) => { + if (e.key === 'Enter' || e.key === ' ') { + e.preventDefault(); + inputRef.current?.click(); + } + }} + onDragOver={(e) => { + e.preventDefault(); + setDragging(true); + }} + onDragLeave={() => setDragging(false)} + onDrop={(e) => { + e.preventDefault(); + setDragging(false); + pick(e.dataTransfer.files[0]); + }} + className={cn( + 'rounded-xl border-2 border-dashed p-10 text-center cursor-pointer transition-colors', + 'focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-accent', + dragging + ? 'border-accent-line bg-accent-soft' + : 'border-border-default hover:bg-surface-1', + )} + > + +

{t('csvImport.dropHere')}

+

{t('csvImport.limits')}

+ pick(e.target.files?.[0])} + data-testid="import-file-input" + /> +
+ {fileError && ( + + )} + + {file && ( +
+
+ +
+

{file.name}

+

{formatSize(file.size)}

+
+
+
+ + +
+
+ )} + +
+ {outcome && } +
+
+ ); +}; + +function OutcomePanel({ outcome, config, t }: { outcome: Outcome; config: CsvImportConfig; t: T }) { + if (outcome.kind === 'created') { + return ( +
+ +
+

{config.created(outcome.created)}.

+ + {config.open.label} + +
+
+ ); + } + + const title = + outcome.kind === 'rejected' + ? outcome.rejected > 0 + ? t('csvImport.rejectedRows', { errors: outcome.errors.length, rows: outcome.rejected }) + : // The file as a whole was refused (header, scale, size): no row to count. + t('csvImport.fileRefused') + : outcome.kind === 'limit' + ? config.limitTitle + : t('csvImport.nothingImported'); + + return ( +
+
+ +
+

{title}

+ {outcome.kind === 'limit' && ( +

+ {outcome.requested !== undefined && outcome.remaining !== undefined + ? config.limitDetail(outcome.requested, outcome.remaining) + ' ' + : ''} + + {t('csvImport.seePlans')} + +

+ )} + {outcome.kind === 'failed' && ( +

{outcome.message}

+ )} + {outcome.kind === 'rejected' && ( + <> +

{t('csvImport.fixRows')}

+
+ + + + + + + + + + {outcome.errors.map((e, i) => ( + + + + + + ))} + +
+ {t('csvImport.line')} + + {t('csvImport.column')} + + {t('csvImport.problem')} +
+ {e.line > 0 ? e.line : t('csvImport.file')} + {e.column ?? '—'}{importErrorMessage(e, t)}
+
+ + )} +
+
+
+ ); +} + +function toOutcome(err: unknown, t: T, forbidden: string): Outcome { + if (axios.isAxiosError(err) && err.response) { + const { status, data } = err.response; + if (status === 422) { + const parsed = importRejectedSchema.safeParse(data); + if (parsed.success) { + return { kind: 'rejected', rejected: parsed.data.rejected, errors: parsed.data.errors }; + } + } + if (status === 402) { + const parsed = importLimitSchema.safeParse(data); + if (parsed.success) { + return { + kind: 'limit', + requested: parsed.data.requested, + remaining: parsed.data.remaining, + }; + } + } + if (status === 403) { + return { kind: 'failed', message: forbidden }; + } + const message = + typeof data === 'object' && + data !== null && + 'message' in data && + typeof data.message === 'string' + ? data.message + : null; + if (message) return { kind: 'failed', message }; + } + return { + kind: 'failed', + message: t('csvImport.serverFailed'), + }; +} + +function formatSize(bytes: number): string { + if (bytes < 1024) return `${bytes} B`; + if (bytes < 1024 * 1024) return `${(bytes / 1024).toFixed(1)} KB`; + return `${(bytes / (1024 * 1024)).toFixed(1)} MB`; +} diff --git a/frontend/src/shared/csvImport/csvImportSchema.ts b/frontend/src/shared/csvImport/csvImportSchema.ts new file mode 100644 index 00000000..7a8b8725 --- /dev/null +++ b/frontend/src/shared/csvImport/csvImportSchema.ts @@ -0,0 +1,92 @@ +// Copyright (c) 2026 OpenDefender Contributors +// SPDX-License-Identifier: AGPL-3.0-only +// This program is free software: you can redistribute it and/or modify it under +// the terms of the GNU Affero General Public License v3.0 (see LICENSE). +// +// Contract of the CSV imports, starting with POST /risks/import (#755). The server is the authority on +// every row; the client only refuses what it can know without reading the file, +// and parses every response so the page never shows a number it was not sent. + +import { z } from 'zod'; + +/** Mirrors MaxImportBytes of the server import use cases. */ +export const MAX_IMPORT_BYTES = 2 * 1024 * 1024; + +/** The catalogue translator, as useI18n().t provides it. */ +export type T = (key: string, params?: Record) => string; + +export function importFileSchema(t: T) { + return z + .instanceof(File) + .refine((f) => f.name.toLowerCase().endsWith('.csv'), { message: t('csvImport.onlyCsv') }) + .refine((f) => f.size > 0, { message: t('csvImport.fileEmpty') }) + .refine((f) => f.size <= MAX_IMPORT_BYTES, { message: t('csvImport.tooLarge') }); +} + +export const importRowErrorSchema = z.object({ + line: z.number().int(), + column: z.string().optional(), + /** Stable code; the page renders it in the reader's language. */ + code: z.string().optional(), + params: z.record(z.string(), z.string()).optional(), + /** The server's English rendering, shown when a code is unknown. */ + message: z.string(), +}); +export type ImportRowError = z.infer; + +/** Codes the catalogue knows, under csvImport.errors. */ +const KNOWN_CODES = new Set([ + 'file_empty', + 'not_utf8', + 'header_unreadable', + 'unknown_column', + 'duplicate_column', + 'missing_column', + 'line_unreadable', + 'too_many_rows', + 'cell_count', + 'required', + 'too_long', + 'not_a_number', + 'out_of_range', + 'no_rows', + 'unknown_asset', + 'ambiguous_asset', + 'assets_unavailable', + 'legacy_scale', +]); + +/** + * Renders one server error in the reader's language from its code and params + * (csvImport.errors in the catalogue). The codes mirror ImportRowError on the + * server; an unknown code falls back to the server's English message rather + * than to nothing. + */ +export function importErrorMessage(e: ImportRowError, t: T): string { + if (!e.code || !KNOWN_CODES.has(e.code)) return e.message; + return t(`csvImport.errors.${e.code}`, { column: e.column ?? '', ...e.params }); +} + +/** 200: every row was written. */ +export const importSuccessSchema = z.object({ + created: z.number().int(), + rejected: z.number().int(), + errors: z.array(importRowErrorSchema), +}); +export type ImportSuccess = z.infer; + +/** 422: at least one row was invalid, nothing was written. */ +export const importRejectedSchema = z.object({ + created: z.literal(0), + rejected: z.number().int(), + errors: z.array(importRowErrorSchema).min(1), +}); +export type ImportRejected = z.infer; + +/** 402: the file would take the tenant past its plan's risk limit. */ +export const importLimitSchema = z.object({ + code: z.literal('limit_reached'), + requested: z.number().int().optional(), + remaining: z.number().int().optional(), +}); +export type ImportLimit = z.infer;