diff --git a/internal/framework/contracts/contracts.go b/internal/framework/contracts/contracts.go new file mode 100644 index 0000000..53a4eee --- /dev/null +++ b/internal/framework/contracts/contracts.go @@ -0,0 +1,97 @@ +package contracts + +import ( + "context" + "encoding/json" + + "gitea.maximumdirect.net/eric/notarius/internal/core/artifacts" + "gitea.maximumdirect.net/eric/notarius/internal/core/source" +) + +type LLMMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type StructuredCompletionRequest struct { + StageName string `json:"stage_name"` + Messages []LLMMessage `json:"messages"` + Model string `json:"model,omitempty"` + ResponseSchemaName string `json:"response_schema_name,omitempty"` + ResponseSchema json.RawMessage `json:"response_schema,omitempty"` +} + +type StructuredCompletionResponse struct { + Content json.RawMessage `json:"content"` + Provider string `json:"provider,omitempty"` + Model string `json:"model,omitempty"` + PromptTokens int `json:"prompt_tokens,omitempty"` + CompletionTokens int `json:"completion_tokens,omitempty"` + TotalTokens int `json:"total_tokens,omitempty"` +} + +type StructuredLLMClient interface { + CompleteStructured(ctx context.Context, req StructuredCompletionRequest, out any) (StructuredCompletionResponse, error) +} + +type ParseRequest struct { + SourceID string `json:"source_id,omitempty"` + Path string `json:"path,omitempty"` + Raw []byte `json:"-"` + Metadata map[string]any `json:"metadata,omitempty"` +} + +type InputAdapter interface { + Key() string + Parse(ctx context.Context, req ParseRequest) (*source.SourceDocument, error) +} + +type ExtractionRequest struct { + Source *source.SourceDocument `json:"-"` + LLMClient StructuredLLMClient `json:"-"` + Metadata map[string]any `json:"metadata,omitempty"` +} + +type ExtractionResult struct { + Candidates []artifacts.Candidate `json:"candidates,omitempty"` + Warnings []Warning `json:"warnings,omitempty"` +} + +type Extractor interface { + Key() string + ArtifactType() string + SchemaVersion() string + Validators() []Validator + Extract(ctx context.Context, req ExtractionRequest) (ExtractionResult, error) +} + +type ValidationRequest struct { + Source *source.SourceDocument `json:"-"` + Candidates []artifacts.Candidate `json:"candidates"` + Metadata map[string]any `json:"metadata,omitempty"` +} + +type ValidationDecision struct { + CandidateIndex int `json:"candidate_index"` + Approved bool `json:"approved"` + ReasonCode string `json:"reason_code"` + Message string `json:"message"` + DiagnosticArtifactPath string `json:"diagnostic_artifact_path,omitempty"` +} + +type ValidationResult struct { + ValidatorName string `json:"validator_name"` + Decisions []ValidationDecision `json:"decisions"` + Warnings []Warning `json:"warnings,omitempty"` +} + +type Validator interface { + Name() string + Validate(ctx context.Context, req ValidationRequest) (ValidationResult, error) +} + +type Warning struct { + Scope string `json:"scope,omitempty"` + ReasonCode string `json:"reason_code"` + Message string `json:"message"` +} diff --git a/internal/framework/contracts/contracts_test.go b/internal/framework/contracts/contracts_test.go new file mode 100644 index 0000000..642dc80 --- /dev/null +++ b/internal/framework/contracts/contracts_test.go @@ -0,0 +1,165 @@ +package contracts + +import ( + "context" + "encoding/json" + "testing" + + "gitea.maximumdirect.net/eric/notarius/internal/core/artifacts" + "gitea.maximumdirect.net/eric/notarius/internal/core/source" +) + +var _ InputAdapter = fakeAdapter{} +var _ Extractor = fakeExtractor{} +var _ Validator = fakeValidator{} +var _ StructuredLLMClient = fakeLLMClient{} + +func TestFakeExtractorReturnsCandidateAndValidator(t *testing.T) { + validator := fakeValidator{name: "generic-validator"} + extractor := fakeExtractor{ + key: "generic-extractor", + artifactType: "generic-artifact", + schemaVersion: "v1", + validators: []Validator{validator}, + } + doc := &source.SourceDocument{ + ID: "source-1", + Kind: "document", + Format: "text/plain", + Digest: "sha256:abc123", + Units: []source.SourceUnit{ + {ID: "u1", Kind: "section", Text: "Source text."}, + }, + } + + result, err := extractor.Extract(context.Background(), ExtractionRequest{Source: doc}) + if err != nil { + t.Fatalf("Extract() error = %v, want nil", err) + } + + if extractor.Key() != "generic-extractor" { + t.Fatalf("Key() = %q, want generic-extractor", extractor.Key()) + } + if extractor.ArtifactType() != "generic-artifact" { + t.Fatalf("ArtifactType() = %q, want generic-artifact", extractor.ArtifactType()) + } + if extractor.SchemaVersion() != "v1" { + t.Fatalf("SchemaVersion() = %q, want v1", extractor.SchemaVersion()) + } + if len(extractor.Validators()) != 1 { + t.Fatalf("len(Validators()) = %d, want 1", len(extractor.Validators())) + } + if extractor.Validators()[0].Name() != "generic-validator" { + t.Fatalf("Validators()[0].Name() = %q, want generic-validator", extractor.Validators()[0].Name()) + } + if len(result.Candidates) != 1 { + t.Fatalf("len(Candidates) = %d, want 1", len(result.Candidates)) + } + + candidate := result.Candidates[0] + if candidate.Index != 0 { + t.Fatalf("Candidate.Index = %d, want 0", candidate.Index) + } + if candidate.ExtractorKey != extractor.Key() { + t.Fatalf("Candidate.ExtractorKey = %q, want %q", candidate.ExtractorKey, extractor.Key()) + } + if candidate.ArtifactType != extractor.ArtifactType() { + t.Fatalf("Candidate.ArtifactType = %q, want %q", candidate.ArtifactType, extractor.ArtifactType()) + } + if candidate.SchemaVersion != extractor.SchemaVersion() { + t.Fatalf("Candidate.SchemaVersion = %q, want %q", candidate.SchemaVersion, extractor.SchemaVersion()) + } + if string(candidate.Payload) != `{"value":"example"}` { + t.Fatalf("Candidate.Payload = %s, want example payload", candidate.Payload) + } +} + +type fakeAdapter struct { + key string + doc *source.SourceDocument +} + +func (adapter fakeAdapter) Key() string { + return adapter.key +} + +func (adapter fakeAdapter) Parse(ctx context.Context, req ParseRequest) (*source.SourceDocument, error) { + return adapter.doc, nil +} + +type fakeExtractor struct { + key string + artifactType string + schemaVersion string + validators []Validator +} + +func (extractor fakeExtractor) Key() string { + return extractor.key +} + +func (extractor fakeExtractor) ArtifactType() string { + return extractor.artifactType +} + +func (extractor fakeExtractor) SchemaVersion() string { + return extractor.schemaVersion +} + +func (extractor fakeExtractor) Validators() []Validator { + return extractor.validators +} + +func (extractor fakeExtractor) Extract(ctx context.Context, req ExtractionRequest) (ExtractionResult, error) { + return ExtractionResult{ + Candidates: []artifacts.Candidate{ + { + Index: 0, + ExtractorKey: extractor.key, + ArtifactType: extractor.artifactType, + SchemaVersion: extractor.schemaVersion, + Payload: json.RawMessage(`{"value":"example"}`), + SourceRefs: []source.SourceRef{ + { + SourceID: req.Source.ID, + StartUnitID: req.Source.Units[0].ID, + EndUnitID: req.Source.Units[0].ID, + }, + }, + }, + }, + }, nil +} + +type fakeValidator struct { + name string +} + +func (validator fakeValidator) Name() string { + return validator.name +} + +func (validator fakeValidator) Validate(ctx context.Context, req ValidationRequest) (ValidationResult, error) { + decisions := make([]ValidationDecision, 0, len(req.Candidates)) + for _, candidate := range req.Candidates { + decisions = append(decisions, ValidationDecision{ + CandidateIndex: candidate.Index, + Approved: true, + ReasonCode: "accepted", + Message: "candidate accepted", + }) + } + + return ValidationResult{ + ValidatorName: validator.name, + Decisions: decisions, + }, nil +} + +type fakeLLMClient struct{} + +func (client fakeLLMClient) CompleteStructured(ctx context.Context, req StructuredCompletionRequest, out any) (StructuredCompletionResponse, error) { + return StructuredCompletionResponse{ + Content: json.RawMessage(`{"value":"example"}`), + }, nil +}