diff --git a/pkg/types/synthesis.go b/pkg/types/synthesis.go new file mode 100644 index 0000000..53a481b --- /dev/null +++ b/pkg/types/synthesis.go @@ -0,0 +1,59 @@ +package types + +import "time" + +// SynthesisInput contains the input for the synthesis workflow. +type SynthesisInput struct { + Project string `json:"project"` + Source string `json:"source"` + Text string `json:"text"` + Kind string `json:"kind"` // L1, L2, reference + Tags []string `json:"tags,omitempty"` +} + +// SynthesisResult contains the output of the synthesis workflow. +type SynthesisResult struct { + ChunkID string `json:"chunk_id"` + EntitiesExtracted int `json:"entities_extracted"` + FactsExtracted int `json:"facts_extracted"` + Contradictions int `json:"contradictions"` + ReviewQueued int `json:"review_queued"` + Entities []ExtractedEntity `json:"entities"` + Facts []ExtractedFact `json:"facts"` + Duration time.Duration `json:"duration"` +} + +// ExtractedEntity represents an entity found during synthesis. +type ExtractedEntity struct { + Name string `json:"name"` + EntityType string `json:"entity_type"` + Confidence float64 `json:"confidence"` +} + +// ExtractedFact represents a fact extracted during synthesis. +type ExtractedFact struct { + Subject string `json:"subject"` + Predicate string `json:"predicate"` + Object string `json:"object"` + Confidence float64 `json:"confidence"` +} + +// ContradictionResult represents a contradiction detection result. +type ContradictionResult struct { + FactA ExtractedFact `json:"fact_a"` + FactB ExtractedFact `json:"fact_b"` + Severity string `json:"severity"` // low, medium, high + AutoResolved bool `json:"auto_resolved"` + QueuedReview bool `json:"queued_review"` +} + +// PersistInput groups all synthesis results for persistence. +type PersistInput struct { + ChunkID string `json:"chunk_id"` + Project string `json:"project"` + Source string `json:"source"` + Kind string `json:"kind"` + Entities []ExtractedEntity `json:"entities"` + Facts []ExtractedFact `json:"facts"` + Contradictions []ContradictionResult `json:"contradictions"` +} diff --git a/workflow/synthesis.go b/workflow/synthesis.go new file mode 100644 index 0000000..f4e9254 --- /dev/null +++ b/workflow/synthesis.go @@ -0,0 +1,125 @@ +package workflow + +import ( + "fmt" + "time" + + "go.temporal.io/sdk/temporal" + "go.temporal.io/sdk/workflow" + "github.com/rockliang/poimen/workflows/pkg/types" +) + +// Re-export shared types from pkg/types for backward compatibility +type SynthesisInput = types.SynthesisInput +type SynthesisResult = types.SynthesisResult +type ExtractedEntity = types.ExtractedEntity +type ExtractedFact = types.ExtractedFact +type ContradictionResult = types.ContradictionResult +type PersistInput = types.PersistInput + +var synthesisActivityOptions = workflow.ActivityOptions{ + StartToCloseTimeout: 60 * time.Second, + RetryPolicy: &temporal.RetryPolicy{ + InitialInterval: time.Second, + BackoffCoefficient: 2.0, + MaximumInterval: 30 * time.Second, + MaximumAttempts: 3, + }, +} + +// SynthesisWorkflow orchestrates the 4-stage memory synthesis pipeline. +// +// Stage 1: Chunk + embed text +// Stage 2: Extract entities (LLM + reflection) +// Stage 3: Extract facts (pattern + LLM) +// Stage 4: Detect contradictions (pre-filter + LLM) +// +// Each stage is an activity with independent retry policy. +func SynthesisWorkflow(ctx workflow.Context, input SynthesisInput) (*SynthesisResult, error) { + logger := workflow.GetLogger(ctx) + startTime := workflow.Now(ctx) + + logger.Info("synthesis started", + "project", input.Project, + "source", input.Source, + "kind", input.Kind, + ) + + actCtx := workflow.WithActivityOptions(ctx, synthesisActivityOptions) + + // Stage 1: Chunk + Embed + var chunkID string + err := workflow.ExecuteActivity(actCtx, "ChunkAndEmbedActivity", input).Get(ctx, &chunkID) + if err != nil { + return nil, fmt.Errorf("stage 1 chunk+embed: %w", err) + } + logger.Info("stage 1 complete", "chunk_id", chunkID) + + // Stage 2: Entity Extraction + var entities []ExtractedEntity + err = workflow.ExecuteActivity(actCtx, "ExtractEntitiesActivity", chunkID, input.Text).Get(ctx, &entities) + if err != nil { + return nil, fmt.Errorf("stage 2 entity extraction: %w", err) + } + logger.Info("stage 2 complete", "entities", len(entities)) + + // Stage 3: Fact Extraction + var facts []ExtractedFact + err = workflow.ExecuteActivity(actCtx, "ExtractFactsActivity", chunkID, input.Text, entities).Get(ctx, &facts) + if err != nil { + return nil, fmt.Errorf("stage 3 fact extraction: %w", err) + } + logger.Info("stage 3 complete", "facts", len(facts)) + + // Stage 4: Contradiction Detection + var contradictions []ContradictionResult + err = workflow.ExecuteActivity(actCtx, "DetectContradictionsActivity", input.Project, facts).Get(ctx, &contradictions) + if err != nil { + return nil, fmt.Errorf("stage 4 contradiction detection: %w", err) + } + + reviewQueued := 0 + for _, c := range contradictions { + if c.QueuedReview { + reviewQueued++ + } + } + logger.Info("stage 4 complete", "contradictions", len(contradictions), "review_queued", reviewQueued) + + // Stage 5: Persist results + persistInput := PersistInput{ + ChunkID: chunkID, + Project: input.Project, + Source: input.Source, + Kind: input.Kind, + Entities: entities, + Facts: facts, + Contradictions: contradictions, + } + err = workflow.ExecuteActivity(actCtx, "PersistSynthesisActivity", persistInput).Get(ctx, nil) + if err != nil { + return nil, fmt.Errorf("stage 5 persist: %w", err) + } + + duration := workflow.Now(ctx).Sub(startTime) + result := &SynthesisResult{ + ChunkID: chunkID, + EntitiesExtracted: len(entities), + FactsExtracted: len(facts), + Contradictions: len(contradictions), + ReviewQueued: reviewQueued, + Entities: entities, + Facts: facts, + Duration: duration, + } + + logger.Info("synthesis complete", + "chunk_id", chunkID, + "entities", len(entities), + "facts", len(facts), + "contradictions", len(contradictions), + "duration", duration, + ) + + return result, nil +} diff --git a/workflow/synthesis_test.go b/workflow/synthesis_test.go new file mode 100644 index 0000000..f7d4d8a --- /dev/null +++ b/workflow/synthesis_test.go @@ -0,0 +1,179 @@ +package workflow + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "go.temporal.io/sdk/testsuite" +) + +// Stub activity functions for test registration +func ChunkAndEmbedActivity(_ context.Context, _ SynthesisInput) (string, error) { return "", nil } +func ExtractEntitiesActivity(_ context.Context, _ string, _ string) ([]ExtractedEntity, error) { return nil, nil } +func ExtractFactsActivity(_ context.Context, _ string, _ string, _ []ExtractedEntity) ([]ExtractedFact, error) { return nil, nil } +func DetectContradictionsActivity(_ context.Context, _ string, _ []ExtractedFact) ([]ContradictionResult, error) { return nil, nil } +func PersistSynthesisActivity(_ context.Context, _ PersistInput) error { return nil } + +func registerSynthesisActivities(env *testsuite.TestWorkflowEnvironment) { + env.RegisterActivity(ChunkAndEmbedActivity) + env.RegisterActivity(ExtractEntitiesActivity) + env.RegisterActivity(ExtractFactsActivity) + env.RegisterActivity(DetectContradictionsActivity) + env.RegisterActivity(PersistSynthesisActivity) +} + +func TestSynthesisWorkflow_Success(t *testing.T) { + ts := &testsuite.WorkflowTestSuite{} + env := ts.NewTestWorkflowEnvironment() + + input := SynthesisInput{ + Project: "poimen", + Source: "transcript://test-123", + Text: "Kubernetes uses port 8080 for the API server", + Kind: "L1", + } + + registerSynthesisActivities(env) + + // Stage 1: Chunk + Embed + env.OnActivity(ChunkAndEmbedActivity, mock.Anything, input).Return("chunk-abc123", nil) + + // Stage 2: Entity Extraction + entities := []ExtractedEntity{ + {Name: "Kubernetes", EntityType: "tool", Confidence: 0.95}, + {Name: "API server", EntityType: "component", Confidence: 0.90}, + } + env.OnActivity(ExtractEntitiesActivity, mock.Anything, "chunk-abc123", input.Text).Return(entities, nil) + + // Stage 3: Fact Extraction + facts := []ExtractedFact{ + {Subject: "Kubernetes", Predicate: "uses_port", Object: "8080", Confidence: 0.85}, + } + env.OnActivity(ExtractFactsActivity, mock.Anything, "chunk-abc123", input.Text, entities).Return(facts, nil) + + // Stage 4: Contradiction Detection + contradictions := []ContradictionResult{} + env.OnActivity(DetectContradictionsActivity, mock.Anything, "poimen", facts).Return(contradictions, nil) + + // Stage 5: Persist + env.OnActivity(PersistSynthesisActivity, mock.Anything, mock.Anything).Return(nil) + + env.ExecuteWorkflow(SynthesisWorkflow, input) + + assert.True(t, env.IsWorkflowCompleted()) + assert.NoError(t, env.GetWorkflowError()) + + var result SynthesisResult + assert.NoError(t, env.GetWorkflowResult(&result)) + assert.Equal(t, "chunk-abc123", result.ChunkID) + assert.Equal(t, 2, result.EntitiesExtracted) + assert.Equal(t, 1, result.FactsExtracted) + assert.Equal(t, 0, result.Contradictions) + assert.Equal(t, 0, result.ReviewQueued) +} + +func TestSynthesisWorkflow_WithContradictions(t *testing.T) { + ts := &testsuite.WorkflowTestSuite{} + env := ts.NewTestWorkflowEnvironment() + registerSynthesisActivities(env) + + input := SynthesisInput{ + Project: "poimen", + Source: "transcript://test-456", + Text: "Port 8080 is used by nginx", + Kind: "L1", + } + + env.OnActivity(ChunkAndEmbedActivity, mock.Anything, input).Return("chunk-def456", nil) + + entities := []ExtractedEntity{ + {Name: "nginx", EntityType: "tool", Confidence: 0.92}, + } + env.OnActivity(ExtractEntitiesActivity, mock.Anything, "chunk-def456", input.Text).Return(entities, nil) + + facts := []ExtractedFact{ + {Subject: "nginx", Predicate: "uses_port", Object: "8080", Confidence: 0.88}, + } + env.OnActivity(ExtractFactsActivity, mock.Anything, "chunk-def456", input.Text, entities).Return(facts, nil) + + contradictions := []ContradictionResult{ + { + FactA: ExtractedFact{Subject: "Kubernetes", Predicate: "uses_port", Object: "8080"}, + FactB: ExtractedFact{Subject: "nginx", Predicate: "uses_port", Object: "8080"}, + Severity: "medium", + AutoResolved: false, + QueuedReview: true, + }, + } + env.OnActivity(DetectContradictionsActivity, mock.Anything, "poimen", facts).Return(contradictions, nil) + env.OnActivity(PersistSynthesisActivity, mock.Anything, mock.Anything).Return(nil) + + env.ExecuteWorkflow(SynthesisWorkflow, input) + + assert.True(t, env.IsWorkflowCompleted()) + assert.NoError(t, env.GetWorkflowError()) + + var result SynthesisResult + assert.NoError(t, env.GetWorkflowResult(&result)) + assert.Equal(t, 1, result.Contradictions) + assert.Equal(t, 1, result.ReviewQueued) +} + +func TestSynthesisWorkflow_EntityExtractionFails(t *testing.T) { + ts := &testsuite.WorkflowTestSuite{} + env := ts.NewTestWorkflowEnvironment() + registerSynthesisActivities(env) + + input := SynthesisInput{ + Project: "poimen", + Source: "transcript://test-789", + Text: "Some text", + Kind: "L1", + } + + env.OnActivity(ChunkAndEmbedActivity, mock.Anything, input).Return("chunk-xyz", nil) + env.OnActivity(ExtractEntitiesActivity, mock.Anything, "chunk-xyz", input.Text). + Return(nil, assert.AnError) + + env.ExecuteWorkflow(SynthesisWorkflow, input) + + assert.True(t, env.IsWorkflowCompleted()) + assert.Error(t, env.GetWorkflowError()) + assert.Contains(t, env.GetWorkflowError().Error(), "stage 2 entity extraction") +} + +func TestSynthesisWorkflow_ChunkFails(t *testing.T) { + ts := &testsuite.WorkflowTestSuite{} + env := ts.NewTestWorkflowEnvironment() + registerSynthesisActivities(env) + + input := SynthesisInput{ + Project: "poimen", + Source: "transcript://test-fail", + Text: "Bad text", + Kind: "L1", + } + + env.OnActivity(ChunkAndEmbedActivity, mock.Anything, input).Return("", assert.AnError) + + env.ExecuteWorkflow(SynthesisWorkflow, input) + + assert.True(t, env.IsWorkflowCompleted()) + assert.Error(t, env.GetWorkflowError()) + assert.Contains(t, env.GetWorkflowError().Error(), "stage 1 chunk+embed") +} + +func TestSynthesisInput_Fields(t *testing.T) { + input := SynthesisInput{ + Project: "test", + Source: "source://1", + Text: "hello", + Kind: "L2", + Tags: []string{"tag1", "tag2"}, + } + assert.Equal(t, "test", input.Project) + assert.Equal(t, "L2", input.Kind) + assert.Len(t, input.Tags, 2) +}