Compare commits
40
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
84b4ca120f | ||
|
|
a461e9799a | ||
|
|
9624f0e18d | ||
|
|
2a3b080e29 | ||
|
|
00fc83c081 | ||
|
|
0da90fdd7a | ||
|
|
924aa398b6 | ||
|
|
4a34c8e672 | ||
|
|
ebf95506cd | ||
|
|
66c17e821f | ||
|
|
86ad8e7b5e | ||
|
|
e8984b055c | ||
|
|
0d70da4f31 | ||
|
|
fd2ebce8e1 | ||
|
|
1c16869126 | ||
|
|
a0e64224a7 | ||
|
|
5a465b145c | ||
|
|
687bdb21e0 | ||
|
|
ab79cf33bc | ||
|
|
25a4787022 | ||
|
|
db71919207 | ||
|
|
9b9e99da3e | ||
|
|
fd3d2787c3 | ||
|
|
f69295db6a | ||
|
|
5ef14ad5ec | ||
|
|
978a33377c | ||
|
|
6a87833c7f | ||
|
|
121cad1ad5 | ||
|
|
c7fcd3c6f9 | ||
|
|
ca96769736 | ||
|
|
6360466a28 | ||
|
|
71f3bfae65 | ||
|
|
b14d124049 | ||
|
|
00f1dad5df | ||
|
|
cb94314bcc | ||
|
|
75a01a9444 | ||
|
|
e00762bb0b | ||
|
|
b0313ae818 | ||
|
|
cb8a3fe12a | ||
|
|
00d40e3bbe |
@@ -0,0 +1,22 @@
|
||||
# Workflows Backend Configuration
|
||||
|
||||
# Database (memory-db CNPG in K8s)
|
||||
# Option A: Direct DATABASE_URL
|
||||
DATABASE_URL=postgresql://app:[email protected]:5432/memory?sslmode=disable
|
||||
|
||||
# Option B: Individual env vars (used if DATABASE_URL is empty)
|
||||
DATABASE_HOST=memory-db-rw.poimen.svc.cluster.local
|
||||
DATABASE_PORT=5432
|
||||
DATABASE_NAME=memory
|
||||
DATABASE_USER=app
|
||||
DATABASE_PASSWORD=PASSWORD
|
||||
|
||||
# Temporal
|
||||
TEMPORAL_HOST_PORT=localhost:7233
|
||||
TEMPORAL_NAMESPACE=default
|
||||
|
||||
# API Server
|
||||
API_PORT=8080
|
||||
|
||||
# Logging
|
||||
VERBOSE=false
|
||||
@@ -0,0 +1,9 @@
|
||||
# Ignore markdown docs except agent-prompts and README
|
||||
*.md
|
||||
!README.md
|
||||
!agent-prompts/*.md
|
||||
!agent-prompts/**/*.md
|
||||
# Binaries
|
||||
starter
|
||||
worker
|
||||
poimen
|
||||
@@ -0,0 +1,278 @@
|
||||
================================================================================
|
||||
POIMEN ROUTING WORKFLOW - COMPLETE SPECIFICATION SUMMARY
|
||||
================================================================================
|
||||
|
||||
STATUS: ✅ READY FOR IMPLEMENTATION
|
||||
|
||||
Created: August 31, 2025
|
||||
Total Documentation: 3,856 lines across 5 files
|
||||
Implementation Effort: 60-70 hours (3-4 weeks, 1-2 engineers)
|
||||
|
||||
================================================================================
|
||||
📚 DOCUMENTATION CREATED
|
||||
================================================================================
|
||||
|
||||
1. ROUTING_WORKFLOW_SPEC.md (1,198 lines, 32KB)
|
||||
├─ Architecture overview
|
||||
├─ ActivityKnowledgeBase.json format
|
||||
├─ llm-router Activity (intelligent generator)
|
||||
├─ RoutingWorkflow (generic executor)
|
||||
├─ Go type definitions (copy-paste ready)
|
||||
├─ 7 implementation architecture sections
|
||||
├─ CronWorkflowSpec (scheduled workflows)
|
||||
├─ Execution flow with examples
|
||||
├─ Validation rules
|
||||
└─ Complete reference
|
||||
|
||||
2. IMPLEMENTATION_TASKS.md (1,158 lines, 27KB)
|
||||
├─ Phase 1: Foundation (8-10 hours, 4 tasks)
|
||||
├─ Phase 2: LLM-Router (12-15 hours, 5 tasks)
|
||||
├─ Phase 3: RoutingWorkflow (15-18 hours, 6 tasks)
|
||||
├─ Phase 4: API/CLI (12-15 hours, 4 tasks)
|
||||
├─ Phase 5: Testing (8-12 hours, 4 tasks)
|
||||
├─ Phase 6: Documentation (5-8 hours, 4 tasks)
|
||||
├─ Total: 27 specific, actionable tasks
|
||||
├─ Each with: effort estimate, acceptance criteria, dependencies
|
||||
├─ Timeline: 3-4 weeks
|
||||
├─ Resource allocation: 1-2 engineers
|
||||
├─ Blockers to watch
|
||||
└─ Success criteria per phase
|
||||
|
||||
3. CRON_JOBS_QUICK_REFERENCE.md (252 lines, 5.7KB)
|
||||
├─ Cron syntax examples (daily, hourly, weekly, etc)
|
||||
├─ How llm-router detects scheduled jobs
|
||||
├─ Execution tracking
|
||||
├─ Proposed API endpoints
|
||||
├─ One-time vs Cron comparison table
|
||||
└─ Quick lookup reference
|
||||
|
||||
4. DESIGN_MASTER_REVIEW.md (888 lines, 25KB)
|
||||
├─ Executive summary for stakeholders
|
||||
├─ Problem/solution statement
|
||||
├─ 3 patterns (Sequential, Await-Task-Complete, Retry)
|
||||
├─ 3 entry points (CLI, API, Legacy)
|
||||
├─ Before/after comparison
|
||||
├─ Implementation timeline
|
||||
├─ KMSvc questions (Q1-Q6)
|
||||
├─ Risks & mitigations
|
||||
├─ Success criteria
|
||||
├─ Approval checklist
|
||||
└─ Complete example workflows
|
||||
|
||||
5. README_IMPLEMENTATION.md (360 lines, 9KB)
|
||||
├─ Quick start guide
|
||||
├─ Documentation structure explanation
|
||||
├─ Week-by-week breakdown
|
||||
├─ How to start today
|
||||
├─ Success criteria per phase
|
||||
├─ Effort summary table
|
||||
├─ Key features matrix
|
||||
├─ Tips for success
|
||||
├─ Learning resources
|
||||
└─ Decision maker's checklist
|
||||
|
||||
================================================================================
|
||||
🎯 THE SYSTEM ARCHITECTURE
|
||||
================================================================================
|
||||
|
||||
User Input (one-time or scheduled):
|
||||
"Analyze repo for security and quality"
|
||||
or
|
||||
"Scan all repos daily at 2 AM"
|
||||
|
||||
↓
|
||||
|
||||
[llm-router Activity] - Intelligent Workflow Generator
|
||||
├─ Reads: ActivityKnowledgeBase.json (metadata about activities)
|
||||
├─ Uses LLM to understand intent
|
||||
├─ Selects activities: Clone → Analyze → SecurityScan → Combine → Notify
|
||||
├─ Orders by dependencies
|
||||
├─ Decides timeout for each (from knowledge base)
|
||||
├─ Decides retry policy (from isFlaky flag)
|
||||
├─ Chains parameters (JSONPath: ${Clone.output.path})
|
||||
├─ Detects if scheduled (cron)
|
||||
└─ Generates: WorkflowSpec or CronWorkflowSpec (JSON)
|
||||
|
||||
↓
|
||||
|
||||
[RoutingWorkflow] - Generic Executor
|
||||
├─ Takes JSON spec from llm-router
|
||||
├─ Executes states in order
|
||||
├─ Respects timeout/retry for each activity
|
||||
├─ Handles errors with catch blocks
|
||||
└─ Returns results
|
||||
|
||||
↓
|
||||
|
||||
[Temporal] - Distributed Workflow Engine
|
||||
├─ For one-time: Executes immediately
|
||||
├─ For cron: Schedules and runs on schedule
|
||||
├─ Provides durability (replay guarantee)
|
||||
├─ Tracks execution history
|
||||
└─ Handles retries automatically
|
||||
|
||||
↓
|
||||
|
||||
[Results] - Final Output
|
||||
├─ Execution history
|
||||
├─ Step-by-step results
|
||||
├─ Performance metrics
|
||||
└─ Status updates
|
||||
|
||||
================================================================================
|
||||
✨ KEY FEATURES
|
||||
================================================================================
|
||||
|
||||
✅ One-time workflows (instant execution via API/CLI)
|
||||
✅ Scheduled workflows (cron jobs with full history)
|
||||
✅ Intelligent routing (LLM decides what to run)
|
||||
✅ Smart timeouts (from ActivityKnowledgeBase.json)
|
||||
✅ Smart retries (3x for flaky, 1x for stable)
|
||||
✅ Error handling (catch blocks for graceful failures)
|
||||
✅ Parameter chaining (JSONPath: ${step.output.field})
|
||||
✅ Parallel execution (multiple branches)
|
||||
✅ Temporal durability (automatic replay on failure)
|
||||
✅ HTTP API (for programmatic access)
|
||||
✅ CLI (for command-line access)
|
||||
✅ Execution tracking (full history)
|
||||
✅ Backward compatible (legacy CLI still works)
|
||||
|
||||
================================================================================
|
||||
📊 IMPLEMENTATION BREAKDOWN
|
||||
================================================================================
|
||||
|
||||
PHASE 1: Foundation (8-10 hours)
|
||||
Task 1.1: Go types (2h)
|
||||
Task 1.2: ActivityKnowledgeBase.json (3h)
|
||||
Task 1.3: KB loader (2h)
|
||||
Task 1.4: Validator (3h)
|
||||
→ Deliverable: Core data structures
|
||||
|
||||
PHASE 2: LLM-Router (12-15 hours)
|
||||
Task 2.1: JSONPath resolver (3h)
|
||||
Task 2.2: Activity skeleton (2h)
|
||||
Task 2.3: LLM intent analysis (5h) ⚠️ HIGHEST RISK
|
||||
Task 2.4: Spec builder (4h)
|
||||
Task 2.5: Cron builder (2h)
|
||||
→ Deliverable: Intelligent workflow generation
|
||||
|
||||
PHASE 3: RoutingWorkflow (15-18 hours)
|
||||
Task 3.1: Executor dispatcher (1h)
|
||||
Task 3.2: Task executor (2h)
|
||||
Task 3.3: Pass/Fail executors (1h)
|
||||
Task 3.4: Main workflow engine (4h)
|
||||
Task 3.5: Register in worker (1h)
|
||||
Task 3.6: Helper functions (2h)
|
||||
→ Deliverable: Generic workflow executor
|
||||
|
||||
PHASE 4: API/CLI (12-15 hours)
|
||||
Task 4.1: API handlers (4h)
|
||||
Task 4.2: CLI commands (5h)
|
||||
Task 4.3: Server bootstrap (2h)
|
||||
Task 4.4: Validation (2h)
|
||||
→ Deliverable: HTTP API + CLI
|
||||
|
||||
PHASE 5: Testing (8-12 hours)
|
||||
Task 5.1: Unit tests (3h)
|
||||
Task 5.2: Integration tests (4h)
|
||||
Task 5.3: E2E tests (4h)
|
||||
Task 5.4: Load tests (2h)
|
||||
→ Deliverable: >90% coverage, all scenarios pass
|
||||
|
||||
PHASE 6: Documentation (5-8 hours)
|
||||
Task 6.1: API documentation (2h)
|
||||
Task 6.2: CLI documentation (1h)
|
||||
Task 6.3: Deployment guide (2h)
|
||||
Task 6.4: User guide & examples (2h)
|
||||
→ Deliverable: Complete documentation
|
||||
|
||||
TOTAL: 60-70 hours (3-4 weeks, 1-2 engineers)
|
||||
|
||||
================================================================================
|
||||
🚀 HOW TO START TODAY
|
||||
================================================================================
|
||||
|
||||
Step 1: Review Documentation (1-2 hours)
|
||||
→ Read ROUTING_WORKFLOW_SPEC.md (understand design)
|
||||
→ Read IMPLEMENTATION_TASKS.md (understand tasks)
|
||||
→ Read README_IMPLEMENTATION.md (quick start)
|
||||
|
||||
Step 2: Assign Tasks (30 minutes)
|
||||
→ Engineer 1: Tasks 1.1-1.4, 2.1-2.5, 3.1-3.6
|
||||
→ Engineer 2: Tasks 4.1-4.4, 5.1-5.4, 6.1-6.4
|
||||
|
||||
Step 3: Begin Implementation (immediately)
|
||||
→ Start with Task 1.1: Create Go types (types.go)
|
||||
→ 2 hours to completion
|
||||
→ Then proceed to Task 1.2 (ActivityKnowledgeBase.json)
|
||||
|
||||
Step 4: Daily Sync
|
||||
→ Report progress
|
||||
→ Unblock dependencies
|
||||
→ Adjust timeline if needed
|
||||
|
||||
================================================================================
|
||||
✅ SUCCESS CRITERIA
|
||||
================================================================================
|
||||
|
||||
Phase 1: All types compile, KB loads, validator works
|
||||
Phase 2: llm-router generates valid specs, detects cron
|
||||
Phase 3: RoutingWorkflow executes any spec, handles errors
|
||||
Phase 4: HTTP API + CLI fully functional
|
||||
Phase 5: >90% code coverage, all tests pass
|
||||
Phase 6: Complete documentation, ready to ship
|
||||
|
||||
✅ DONE WHEN:
|
||||
- All code compiles without warnings
|
||||
- All tests pass (unit, integration, E2E, load)
|
||||
- Documentation complete
|
||||
- Can deploy to Kubernetes
|
||||
- Can submit workflows from API/CLI
|
||||
- Can create cron jobs
|
||||
- Performance targets met (<200ms submit, <100ms poll)
|
||||
|
||||
================================================================================
|
||||
📁 FILE LOCATIONS
|
||||
================================================================================
|
||||
|
||||
Core Specification:
|
||||
~/workplace/Poimen/workflows/ROUTING_WORKFLOW_SPEC.md
|
||||
|
||||
Task Breakdown:
|
||||
~/workplace/Poimen/workflows/IMPLEMENTATION_TASKS.md
|
||||
|
||||
Cron Reference:
|
||||
~/workplace/Poimen/workflows/CRON_JOBS_QUICK_REFERENCE.md
|
||||
|
||||
Stakeholder Review:
|
||||
~/workplace/Poimen/workflows/DESIGN_MASTER_REVIEW.md
|
||||
|
||||
Quick Start Guide:
|
||||
~/workplace/Poimen/workflows/README_IMPLEMENTATION.md
|
||||
|
||||
This Summary:
|
||||
~/workplace/Poimen/workflows/COMPLETE_SPECIFICATION_SUMMARY.txt
|
||||
|
||||
================================================================================
|
||||
🎓 RECOMMENDATION
|
||||
================================================================================
|
||||
|
||||
This specification is:
|
||||
✅ Complete - covers all aspects of the system
|
||||
✅ Implementable - all code patterns shown
|
||||
✅ Testable - success criteria clearly defined
|
||||
✅ Maintainable - well-documented
|
||||
✅ Scalable - designed for production use
|
||||
|
||||
NEXT STEPS:
|
||||
1. Get stakeholder approval (use DESIGN_MASTER_REVIEW.md)
|
||||
2. Assign engineers (use IMPLEMENTATION_TASKS.md)
|
||||
3. Start Phase 1, Task 1.1 today
|
||||
4. Daily standup on progress
|
||||
5. Gate each phase before moving to next
|
||||
|
||||
TIMELINE: 3-4 weeks to complete implementation ⏱️
|
||||
|
||||
STATUS: 🟢 READY TO BUILD
|
||||
|
||||
================================================================================
|
||||
+385
@@ -0,0 +1,385 @@
|
||||
# Multi-stage build for Poimen Temporal Worker
|
||||
# Stage 1: Builder - Compile Go binary and set up tools
|
||||
FROM golang:1.25-alpine AS builder
|
||||
|
||||
WORKDIR /build
|
||||
|
||||
# Install system dependencies (ast-grep, git, build essentials)
|
||||
RUN apk add --no-cache \
|
||||
git \
|
||||
curl \
|
||||
wget \
|
||||
ca-certificates \
|
||||
gcc \
|
||||
musl-dev \
|
||||
bash \
|
||||
&& echo "[builder] System dependencies installed"
|
||||
|
||||
# Install ast-grep CLI tool
|
||||
RUN curl -fsSL https://github.com/ast-grep/ast-grep/releases/download/0.24.0/sg-x86_64-unknown-linux-musl.tar.gz \
|
||||
| tar xzf - -C /usr/local/bin \
|
||||
&& chmod +x /usr/local/bin/sg \
|
||||
&& sg --version \
|
||||
&& echo "[builder] ast-grep installed"
|
||||
|
||||
# Install Node.js for pi CLI and browser-use
|
||||
RUN apk add --no-cache nodejs npm \
|
||||
&& echo "[builder] Node.js installed"
|
||||
|
||||
# Install pi CLI globally
|
||||
RUN npm install -g @earendil-works/pi-coding-agent --unsafe-perm \
|
||||
&& pi --version \
|
||||
&& echo "[builder] pi CLI installed"
|
||||
|
||||
# Install browser-use CLI for browser automation
|
||||
RUN npm install -g browser-use --unsafe-perm \
|
||||
&& browser-use --version \
|
||||
&& echo "[builder] browser-use CLI installed"
|
||||
|
||||
# Set up pi home directory and skills
|
||||
RUN mkdir -p ~/.pi/agent/skills ~/.pi/agent/agents \
|
||||
&& echo "[builder] pi directories created"
|
||||
|
||||
# Stage 2: Download pi skills (caveman & andrej karpathy)
|
||||
# Clone caveman skill from pi-agent repo
|
||||
RUN cd /tmp && git clone https://github.com/earendil-works/pi-agent.git pi-repo \
|
||||
&& mkdir -p ~/.pi/agent/skills/caveman \
|
||||
&& cp -r pi-repo/examples/skills/caveman/* ~/.pi/agent/skills/caveman/ 2>/dev/null || true \
|
||||
&& echo "[builder] caveman skill installed"
|
||||
|
||||
# Create andrej karpathy skill manually (reference/training patterns)
|
||||
RUN mkdir -p ~/.pi/agent/skills/andrej-karpathy && cat > ~/.pi/agent/skills/andrej-karpathy/SKILL.md << 'EOF'
|
||||
# Andrej Karpathy LLM & AI Principles Skill
|
||||
|
||||
Build neural networks and LLM systems with proven patterns from Andrej Karpathy.
|
||||
Topics: attention mechanisms, transformer training, inference optimization, edge cases.
|
||||
|
||||
## Key Principles
|
||||
|
||||
### 1. Simplicity First
|
||||
- Start with minimal implementation
|
||||
- Add complexity only when justified
|
||||
- Test each component independently
|
||||
- Use debugging tools effectively
|
||||
|
||||
### 2. Neural Network Architecture
|
||||
- Understand backward pass deeply
|
||||
- Implement from scratch when possible
|
||||
- Use visualization for debugging
|
||||
- Profile before optimizing
|
||||
|
||||
### 3. LLM Training Patterns
|
||||
- Quality data > quantity
|
||||
- Curriculum learning for complex tasks
|
||||
- Loss landscape visualization
|
||||
- Checkpoint strategy matters
|
||||
|
||||
### 4. Inference Optimization
|
||||
- Quantization without quality loss
|
||||
- KV cache management
|
||||
- Batch processing strategies
|
||||
- Latency profiling
|
||||
|
||||
### 5. Failure Analysis
|
||||
- Log intermediate activations
|
||||
- Check gradient flow
|
||||
- Validate data pipeline
|
||||
- Test edge cases explicitly
|
||||
|
||||
## Usage in Poimen
|
||||
|
||||
Apply when:
|
||||
- Designing workflow stages (like training curricula)
|
||||
- Optimizing inference (planner/judge/implementer prompts)
|
||||
- Debugging convergence issues (retry patterns)
|
||||
- Scaling to production (quantization patterns)
|
||||
|
||||
## Resources
|
||||
- github.com/karpathy/minGPT - Minimal GPT implementation
|
||||
- youtube: "Neural Networks: Zero to Hero" series
|
||||
- Papers: Attention Is All You Need, GPT series whitepapers
|
||||
EOF
|
||||
&& echo "[builder] andrej-karpathy skill created"
|
||||
|
||||
# Create browser-use skill for web testing & automation
|
||||
RUN mkdir -p ~/.pi/agent/skills/browser-use && cat > ~/.pi/agent/skills/browser-use/SKILL.md << 'EOF'
|
||||
# browser-use: Browser Automation Skill
|
||||
|
||||
Automate web browser interactions for testing, verification, and UI validation.
|
||||
Topics: headless browser control, visual testing, form automation, screenshot capture.
|
||||
|
||||
## Key Capabilities
|
||||
|
||||
### 1. Browser Control
|
||||
- Launch headless Chrome/Firefox
|
||||
- Navigate to URLs
|
||||
- Wait for elements/navigation
|
||||
- Handle popups/dialogs
|
||||
|
||||
### 2. Interaction Patterns
|
||||
- Click buttons/links
|
||||
- Fill forms (text, dropdown, checkbox)
|
||||
- Drag & drop
|
||||
- Keyboard input
|
||||
|
||||
### 3. Verification & Capture
|
||||
- Screenshot capture
|
||||
- Element inspection
|
||||
- Accessibility checks
|
||||
- Network monitoring
|
||||
|
||||
### 4. Wait Strategies
|
||||
- Wait for element visible
|
||||
- Wait for navigation
|
||||
- Wait for condition (custom JS)
|
||||
- Timeout handling
|
||||
|
||||
### 5. Error Recovery
|
||||
- Retry failed actions
|
||||
- Handle stale elements
|
||||
- Browser crash recovery
|
||||
- Memory leak prevention
|
||||
|
||||
## Usage in Poimen Phases
|
||||
|
||||
### Phase T2 (Implementation)
|
||||
- Test generated UI code in real browser
|
||||
- Verify visual layout matches spec
|
||||
- Validate form inputs work correctly
|
||||
|
||||
### Phase T3 (Verification)
|
||||
- Visual regression testing
|
||||
- Accessibility validation (ARIA, keyboard nav)
|
||||
- Cross-browser verification
|
||||
|
||||
### Phase T6 (Integration)
|
||||
- End-to-end workflow testing
|
||||
- External service integration testing
|
||||
- User journey verification
|
||||
|
||||
### Phase T9 (Release)
|
||||
- Pre-release smoke tests
|
||||
- Deployment verification
|
||||
- Production canary testing
|
||||
|
||||
## Example Workflows
|
||||
|
||||
```bash
|
||||
# Launch browser and take screenshot
|
||||
browser-use screenshot "https://example.com" --file output.png
|
||||
|
||||
# Fill form and submit
|
||||
browser-use interact "https://example.com" \
|
||||
--click "#submit-btn" \
|
||||
--type "#email" "[email protected]" \
|
||||
--type "#password" "secretpass" \
|
||||
--click ".submit"
|
||||
|
||||
# Wait for dynamic content and extract data
|
||||
browser-use extract "https://example.com" \
|
||||
--wait ".dynamic-content" \
|
||||
--selector ".data-row" \
|
||||
--output json
|
||||
|
||||
# Accessibility audit
|
||||
browser-use audit "https://example.com" \
|
||||
--check wcag2a \
|
||||
--report a11y-report.html
|
||||
```
|
||||
|
||||
## Integration with Poimen
|
||||
|
||||
Pre-generated code can be tested:
|
||||
```bash
|
||||
# Generate code (T2)
|
||||
implementer_output = "function handleClick() { ... }"
|
||||
|
||||
# Verify in browser (T3)
|
||||
browser-use interact "http://localhost:3000" \
|
||||
--click ".test-button" \
|
||||
--screenshot result.png
|
||||
|
||||
# Compare with expected
|
||||
verify_visual_match(result.png, expected.png)
|
||||
```
|
||||
|
||||
## Performance Notes
|
||||
|
||||
- Startup: ~2-5s per browser
|
||||
- Action latency: 100-500ms per interaction
|
||||
- Screenshot: 500ms-2s (depends on page size)
|
||||
- Keep browser alive for batch operations (pool management)
|
||||
|
||||
## Error Handling
|
||||
|
||||
- Transient: Network timeout → retry with backoff
|
||||
- Permanent: Element not found → fail and log
|
||||
- Flaky: Wait strategies → increase timeout gradually
|
||||
- Memory: Reuse browser instances → kill after 10 uses
|
||||
|
||||
## Resources
|
||||
- docs.browseruse.com - Official documentation
|
||||
- github.com/browser-use/browser-use - Source code
|
||||
- Chrome DevTools Protocol - Advanced browser control
|
||||
EOF
|
||||
&& echo "[builder] browser-use skill created"
|
||||
|
||||
# Copy Go source code
|
||||
COPY . /build/
|
||||
|
||||
# Download Go dependencies
|
||||
RUN go mod download \
|
||||
&& echo "[builder] Go dependencies downloaded"
|
||||
|
||||
# Build worker binary
|
||||
RUN CGO_ENABLED=1 GOOS=linux go build -o /build/worker ./cmd/worker \
|
||||
&& echo "[builder] Worker binary built"
|
||||
|
||||
# Verify binary
|
||||
RUN file /build/worker && ls -lh /build/worker
|
||||
|
||||
# Stage 3: Runtime - Minimal base image with runtime dependencies
|
||||
FROM alpine:3.20
|
||||
|
||||
LABEL maintainer="Poimen Team"
|
||||
LABEL description="Poimen Temporal Worker with memory service, ast-grep, and browser automation"
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install runtime dependencies (including Chromium for browser-use)
|
||||
RUN apk add --no-cache \
|
||||
ca-certificates \
|
||||
git \
|
||||
bash \
|
||||
curl \
|
||||
jq \
|
||||
chromium \
|
||||
chromium-chromedriver \
|
||||
&& echo "[runtime] Runtime dependencies installed"
|
||||
|
||||
# Install Node.js for pi CLI and browser-use
|
||||
RUN apk add --no-cache nodejs npm \
|
||||
&& echo "[runtime] Node.js installed"
|
||||
|
||||
# Install pi CLI in runtime image
|
||||
RUN npm install -g @earendil-works/pi-coding-agent --unsafe-perm \
|
||||
&& pi --version \
|
||||
&& echo "[runtime] pi CLI installed"
|
||||
|
||||
# Install browser-use CLI in runtime image
|
||||
RUN npm install -g browser-use --unsafe-perm \
|
||||
&& browser-use --version \
|
||||
&& echo "[runtime] browser-use CLI installed"
|
||||
|
||||
# Copy ast-grep binary from builder
|
||||
COPY --from=builder /usr/local/bin/sg /usr/local/bin/sg
|
||||
RUN chmod +x /usr/local/bin/sg && sg --version \
|
||||
&& echo "[runtime] ast-grep copied"
|
||||
|
||||
# Copy pi skills from builder
|
||||
COPY --from=builder /root/.pi /root/.pi
|
||||
RUN ls -la /root/.pi/agent/skills/ \
|
||||
&& echo "[runtime] pi skills configured"
|
||||
|
||||
# Copy worker binary from builder
|
||||
COPY --from=builder /build/worker /app/worker
|
||||
RUN chmod +x /app/worker && file /app/worker \
|
||||
&& echo "[runtime] Worker binary copied"
|
||||
|
||||
# Create app directory structure
|
||||
RUN mkdir -p /app/work /app/logs /app/screenshots \
|
||||
&& chmod 755 /app/work /app/logs /app/screenshots \
|
||||
&& echo "[runtime] App directories created"
|
||||
|
||||
# Health check endpoint
|
||||
EXPOSE 8081
|
||||
|
||||
# Worker task queue listener
|
||||
ENV TEMPORAL_NAMESPACE=poimen-harness \
|
||||
TEMPORAL_HOSTPORT=temporal-frontend.temporal:7233 \
|
||||
MEMORY_SERVICE_URL=http://memory-service.poimen:5000 \
|
||||
MEMORY_SERVICE_TOKEN= \
|
||||
ANTHROPIC_API_KEY= \
|
||||
PI_SKILLS_PATH=/root/.pi/agent/skills \
|
||||
AST_GREP_BIN=/usr/local/bin/sg \
|
||||
BROWSER_USE_BIN=/usr/local/bin/browser-use \
|
||||
CHROMIUM_BIN=/usr/bin/chromium-browser \
|
||||
SCREENSHOTS_DIR=/app/screenshots
|
||||
|
||||
# Entrypoint script with startup diagnostics
|
||||
COPY --chmod=755 << 'EOF' /app/entrypoint.sh
|
||||
#!/bin/bash
|
||||
set -e
|
||||
|
||||
echo "[$(date)] ========== POIMEN WORKER STARTUP =========="
|
||||
echo "[$(date)] Container: $HOSTNAME"
|
||||
echo "[$(date)] Image: $(cat /etc/os-release | grep PRETTY_NAME | cut -d= -f2)"
|
||||
|
||||
# Verify CLI tools
|
||||
echo "[$(date)] ✓ Checking CLI tools..."
|
||||
echo " - Go version: $(go version 2>/dev/null || echo 'N/A')"
|
||||
echo " - ast-grep: $(sg --version 2>&1 | head -1)"
|
||||
echo " - pi: $(pi --version 2>&1 | head -1)"
|
||||
echo " - browser-use: $(browser-use --version 2>&1 | head -1)"
|
||||
echo " - chromium: $(chromium-browser --version 2>&1 || echo 'Not found')"
|
||||
echo " - git: $(git --version)"
|
||||
echo " - node: $(node --version)"
|
||||
echo " - npm: $(npm --version)"
|
||||
|
||||
# Verify pi skills
|
||||
echo "[$(date)] ✓ Checking pi skills..."
|
||||
if [ -d "$PI_SKILLS_PATH" ]; then
|
||||
echo " - Skills path: $PI_SKILLS_PATH"
|
||||
ls -1 "$PI_SKILLS_PATH" | sed 's/^/ ✓ /'
|
||||
else
|
||||
echo " - WARNING: Skills path not found: $PI_SKILLS_PATH"
|
||||
fi
|
||||
|
||||
# Verify browser tools
|
||||
echo "[$(date)] ✓ Checking browser automation tools..."
|
||||
echo " - Chromium binary: $CHROMIUM_BIN"
|
||||
echo " - Screenshots directory: $SCREENSHOTS_DIR"
|
||||
if [ -d "$SCREENSHOTS_DIR" ]; then
|
||||
echo " - Screenshots dir ready ($(du -sh $SCREENSHOTS_DIR 2>/dev/null | cut -f1 || echo '0B'))"
|
||||
fi
|
||||
|
||||
# Check environment variables
|
||||
echo "[$(date)] ✓ Configuration loaded:"
|
||||
echo " - TEMPORAL_NAMESPACE: $TEMPORAL_NAMESPACE"
|
||||
echo " - TEMPORAL_HOSTPORT: $TEMPORAL_HOSTPORT"
|
||||
echo " - MEMORY_SERVICE_URL: ${MEMORY_SERVICE_URL:-(not set)}"
|
||||
echo " - PI_SKILLS_PATH: $PI_SKILLS_PATH"
|
||||
echo " - CHROMIUM_BIN: $CHROMIUM_BIN"
|
||||
|
||||
# Verify memory service connectivity (optional, non-blocking)
|
||||
if [ ! -z "$MEMORY_SERVICE_URL" ]; then
|
||||
echo "[$(date)] ✓ Testing memory service connectivity..."
|
||||
if curl -sf "$MEMORY_SERVICE_URL/health" > /dev/null 2>&1; then
|
||||
echo " - Memory service: HEALTHY"
|
||||
else
|
||||
echo " - Memory service: UNREACHABLE (will retry in worker)"
|
||||
fi
|
||||
fi
|
||||
|
||||
# Test browser automation (optional, non-blocking)
|
||||
echo "[$(date)] ✓ Testing browser automation..."
|
||||
if command -v chromium-browser &> /dev/null && command -v browser-use &> /dev/null; then
|
||||
echo " - Chromium available: YES"
|
||||
echo " - browser-use available: YES"
|
||||
echo " - Browser automation: READY"
|
||||
else
|
||||
echo " - Browser automation: WARNING - missing dependencies"
|
||||
fi
|
||||
|
||||
echo "[$(date)] ========== STARTING WORKER =========="
|
||||
exec /app/worker
|
||||
EOF
|
||||
|
||||
RUN chmod +x /app/entrypoint.sh
|
||||
|
||||
# Run worker with diagnostics
|
||||
ENTRYPOINT ["/app/entrypoint.sh"]
|
||||
|
||||
# Health check
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD curl -f http://localhost:8081/health || exit 1
|
||||
@@ -1,337 +0,0 @@
|
||||
# Multi-Agent Dev Orchestrator (Temporal + Go) — Implementation Plan
|
||||
|
||||
See the full design doc at `/Users/rockliang/.claude/plans/considered-u-are-a-curried-reef.md`.
|
||||
|
||||
## Quick Summary
|
||||
|
||||
Build a Temporal-based orchestrator that drives multi-agent software dev work on arbitrary target repos. Three roles (Planner/reasoning, Judge/reasoning, Implementer/cheaper model) collaborate on hierarchical tasks (`T0` milestone split into `T0.1`-`T0.9` subtasks). Orchestrator owns config (system prompt, skills, activity timeouts/retries), live-updatable via signals. All code runs on shared FS where target repo sits; git concurrency handled via worktrees + advisory lock. System testable with mocked activities + real end-to-end against `temporal.riotpiao.com`.
|
||||
|
||||
## Implementation Track
|
||||
|
||||
**Milestone T0**: Nine subtasks, each with its own verification gate = completion criterion.
|
||||
|
||||
| Task | Scope | Status |
|
||||
|---|---|---|
|
||||
| [T0.1](#t01-repo-scaffold) | Repo scaffold: `go.mod`, `statemachine/`, `action/`, `cmd/`, `prompts/`, `internal/`, `tests/` | [ ] |
|
||||
| [T0.2](#t02-shared-types) | Shared types: `ModelSpec`, `PromptSpec`, `OrchestratorConfig`, `ActivityTuning`, `PiRetryPolicy` | [ ] |
|
||||
| [T0.3](#t03-git-and-locking) | Git & locking: `CloneRepoActivity`, worktrees, squash-merge, `orchestrator.lock` | [ ] |
|
||||
| [T0.4](#t04-pi-and-error-classification) | `PrepareSkillsActivity`, `classifyPiErr` (4xx/5xx/504), stream timeout learning | [ ] |
|
||||
| [T0.5](#t05-llm-agents-and-prompts) | Planner/Judge/Implementer activities, LLM client, prompt templates + live customization | [ ] |
|
||||
| [T0.6](#t06-taskunit-workflow) | TaskUnit workflow: retry loops, timeout escalation, lessons injection | [ ] |
|
||||
| [T0.7](#t07-orchestrator-workflow) | Orchestrator workflow: config state, signals, fan-out/fan-in, `continue-as-new`, 504 learning | [ ] |
|
||||
| [T0.8](#t08-worker-and-starter) | Worker & starter CLIs, env/config loading, Temporal registration | [ ] |
|
||||
| [T0.9](#t09-end-to-end) | Full e2e against real cluster + scratch repo: all 7 verification items | [ ] |
|
||||
|
||||
---
|
||||
|
||||
### T0.1: Repo Scaffold
|
||||
|
||||
Create directory structure, `go.mod`, empty stubs.
|
||||
|
||||
**Verification:** `go build ./...` succeeds; layout matches plan.
|
||||
|
||||
<details>
|
||||
<summary>Details</summary>
|
||||
|
||||
```
|
||||
/go.mod
|
||||
/cmd/worker/main.go
|
||||
/cmd/starter/main.go
|
||||
/statemachine/types.go types.go signals.go orchestrator.go taskunit.go
|
||||
/action/planner.go implementer.go judge.go git.go skills.go integration_test.go lessons.go llm/client.go
|
||||
/prompts/registry.go planner/default.tmpl judge/default.tmpl implementer/default.tmpl
|
||||
/internal/config/config.go lock/flock.go
|
||||
/tests/taskunit_workflow_test.go orchestrator_workflow_test.go
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
---
|
||||
|
||||
### T0.2: Shared Types
|
||||
|
||||
Implement `statemachine/types.go` with all config/input/output structs. Document defaults.
|
||||
|
||||
**Verification:** Unit test asserts all defaults (5m / 2s / 30s / 2.0 / 30s stream / 2m stream-max).
|
||||
|
||||
<details>
|
||||
<summary>Details</summary>
|
||||
|
||||
- `ModelSpec`: ModelID, Thinking, Effort
|
||||
- `PromptSpec`: TemplateRef, RawTemplate, Variables, Model, LessonsRef
|
||||
- `OrchestratorInput`, `OrchestratorOutput`
|
||||
- `TaskUnitInput`, `TaskUnitOutput`
|
||||
- `ActivityTuning`: ImplementerBaseTimeout, ImplementerMaxRetries, JudgeTimeout, PiRetry
|
||||
- `PiRetryPolicy`: ScheduleToCloseTimeout (5m), InitialInterval (2s), MaximumInterval (30s), BackoffCoefficient (2.0), StreamTimeout (30s), StreamTimeoutMax (2m)
|
||||
- `OrchestratorConfig`: SystemPrompt, Skills, RolePrompts, Tuning
|
||||
|
||||
</details>
|
||||
|
||||
---
|
||||
|
||||
### T0.3: Git & Locking
|
||||
|
||||
Implement `action/git.go` + `internal/lock/flock.go`.
|
||||
|
||||
**Verification:** Test against local scratch repo: clone-if-empty vs fetch-if-exists, worktree lifecycle (add/commit/remove), squash-merge produces exactly one commit on main.
|
||||
|
||||
<details>
|
||||
<summary>Details</summary>
|
||||
|
||||
**Activities:**
|
||||
- `CloneRepoActivity(ctx, {RemoteURL, TargetRepoPath}) error` — idempotent `git clone` or `git fetch`
|
||||
- `GitWorktreeAddActivity(ctx, {RepoPath, TaskID}) (string, error)` — returns worktree path, guarded by lock
|
||||
- `GitCommitActivity(ctx, {WorktreePath, Message}) error` — commits in worktree (no lock needed)
|
||||
- `GitPushActivity(ctx, {RepoPath}) error` — guarded by lock
|
||||
- `GitSquashMergeActivity(ctx, {RepoPath, Branches, Message}) error` — guarded by lock
|
||||
|
||||
**Lock helper (`internal/lock/`):**
|
||||
- `Lock(path string) error`, `Unlock(path string) error` using `golang.org/x/sys/unix.Flock` or `fcntl` equivalent
|
||||
|
||||
**Squash-merge sequence:**
|
||||
```
|
||||
fetch origin main
|
||||
checkout main && pull --ff-only origin main
|
||||
for b in branches:
|
||||
merge --squash $b
|
||||
commit -m "T0: squash merge subtasks..."
|
||||
push origin main
|
||||
for b in branches:
|
||||
worktree remove worktrees/$id --force
|
||||
branch -D $b
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
---
|
||||
|
||||
### T0.4: Pi & Error Classification
|
||||
|
||||
Implement `action/skills.go` with `PrepareSkillsActivity` and `classifyPiErr`.
|
||||
|
||||
**Verification:** Unit test all three error buckets (4xx/5xx/504) against a mocked `pi` HTTP client.
|
||||
|
||||
<details>
|
||||
<summary>Details</summary>
|
||||
|
||||
**`PrepareSkillsActivity`:**
|
||||
- Input: `{Skills []SkillRef, StreamTimeout time.Duration}`
|
||||
- For each skill, `pi clone-or-fetch <skill-url>` (idempotent)
|
||||
- Each skill guarded by its own lock
|
||||
|
||||
**Error classification:**
|
||||
```go
|
||||
func classifyPiErr(err error) error {
|
||||
// 4xx -> NonRetryableApplicationError "PiClientError"
|
||||
// 504 -> ApplicationError "PiStreamTimeout"
|
||||
// others -> retryable
|
||||
}
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
---
|
||||
|
||||
### T0.5: LLM Agents & Prompts
|
||||
|
||||
Implement LLM activities + prompt templates.
|
||||
|
||||
**Verification:** Unit test renders a `PromptSpec` (system prompt + template override + raw template) and calls mock Anthropic client.
|
||||
|
||||
<details>
|
||||
<summary>Details</summary>
|
||||
|
||||
**Activities:**
|
||||
- `PlanningActivity(ctx, {OrchestratorConfig, BoardState}) (TaskDispatch, error)` — reads board/INDEX.md, calls Planner
|
||||
- `ImplementerActivity(ctx, {PromptSpec, WortkreeePath, Lessons}) (ImplementOutput, error)` — tool-call agent loop
|
||||
- `JudgeActivity(ctx, {PromptSpec, Diff, IntegrationTestResult}) (Verdict, Critique, error)` — reviews correctness
|
||||
- `RunIntegrationTestActivity(ctx, {WortkreeePath, TestCmd}) (pass/fail, logs, error)` — shells out
|
||||
|
||||
**Prompt templates:**
|
||||
- `planner/default.tmpl`: expects `{{.SystemPrompt}}`, `{{.TaskBoard}}`, etc.
|
||||
- `judge/default.tmpl`: expects `{{.SystemPrompt}}`, `{{.Diff}}`, `{{.TestResult}}`
|
||||
- `implementer/default.tmpl`: expects `{{.SystemPrompt}}`, `{{.Task}}`, `{{.Lessons}}`
|
||||
|
||||
**`prompts/registry.go`:**
|
||||
- `go:embed prompts/*.tmpl`
|
||||
- `Render(templateRef string, variables map[string]any) (string, error)`
|
||||
|
||||
**`action/llm/client.go`:**
|
||||
- Thin Anthropic client wrapper
|
||||
- Read `ANTHROPIC_API_KEY` from env
|
||||
- Call `messages.Create` with model/thinking/effort from `ModelSpec`
|
||||
|
||||
</details>
|
||||
|
||||
---
|
||||
|
||||
### T0.6: TaskUnit Workflow
|
||||
|
||||
Implement `statemachine/taskunit.go` with retry loops & timeout escalation.
|
||||
|
||||
**Verification (Testsuite):**
|
||||
- Pass-first-try
|
||||
- Fail-then-pass-after-lesson-injection
|
||||
- Retries-exhausted
|
||||
- Timeout-escalation (both judges)
|
||||
|
||||
<details>
|
||||
<summary>Details</summary>
|
||||
|
||||
**Flow:**
|
||||
1. `GitWorktreeAddActivity` → get isolated working tree
|
||||
2. Retry loop:
|
||||
- Track `timeoutAttempt`, `judgeAttempt` separately
|
||||
- `ImplementerActivity` with timeout = `BaseTimeout * timeoutAttempt`
|
||||
- If timeout, increment `timeoutAttempt` and retry (duration grows)
|
||||
- If success, call `RunIntegrationTestActivity`
|
||||
- Call `JudgeActivity`
|
||||
- If judge pass, commit in worktree and return
|
||||
- If judge fail, append to lessons, increment `judgeAttempt`, retry (lessons injected next time)
|
||||
- If retries exhausted, return fail verdict to orchestrator
|
||||
|
||||
**Key detail:** HeartbeatTimeout = (BaseTimeout * timeoutAttempt) / 4, scales with escalation.
|
||||
|
||||
</details>
|
||||
|
||||
---
|
||||
|
||||
### T0.7: Orchestrator Workflow
|
||||
|
||||
Implement `statemachine/orchestrator.go` with config state, signals, fan-out/fan-in, `continue-as-new`, 504 learning.
|
||||
|
||||
**Verification (Testsuite):**
|
||||
- Fan-out/fan-in correctness
|
||||
- Squash-merge triggers on submilestone complete
|
||||
- `continue-as-new` at cycle cap, carries `OrchestratorConfig` forward
|
||||
- `update-*` signals mutate config without touching in-flight TaskUnit
|
||||
- `PiStreamTimeout` doubles `StreamTimeout` and persists it
|
||||
|
||||
<details>
|
||||
<summary>Details</summary>
|
||||
|
||||
**Per-cycle logic:**
|
||||
1. If `config.Skills` changed, call `PrepareSkillsActivity` once (wraps it for 504 learning)
|
||||
2. Call `PlanningActivity` → get dispatch decision
|
||||
3. Fan out: `workflow.ExecuteChildWorkflow(TaskUnitWorkflow, ...)` for each dispatched T0.x
|
||||
4. Await all via `workflow.Selector`
|
||||
5. Call `PlanningActivity` again to update board + commit + push
|
||||
6. If submilestone complete, call `GitSquashMergeActivity`
|
||||
7. Increment cycle count
|
||||
8. If cycle count >= cap, `workflow.NewContinueAsNewError(ctx, ..., nextInput)`
|
||||
|
||||
**Signal handlers:**
|
||||
- `pause`, `resume`: gate the cycle loop
|
||||
- `abort-task(taskID)`: forward via `SignalExternalWorkflow` to TaskUnit
|
||||
- `inject-lesson`: append to lessons store
|
||||
- `update-system-prompt`, `update-skills`, `update-role-prompt`, `update-tuning`: mutate `config.*`
|
||||
|
||||
**504 learning wrapper (pseudo-code):**
|
||||
```go
|
||||
for {
|
||||
r := config.Tuning.PiRetry
|
||||
err := ExecuteActivity(..., PrepareSkillsActivity, Input{...StreamTimeout: r.StreamTimeout})
|
||||
if isPiStreamTimeout(err) && r.StreamTimeout < r.StreamTimeoutMax {
|
||||
config.Tuning.PiRetry.StreamTimeout *= 2
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
---
|
||||
|
||||
### T0.8: Worker & Starter CLIs
|
||||
|
||||
Implement `cmd/worker/main.go` and `cmd/starter/main.go`.
|
||||
|
||||
**Verification:**
|
||||
- `go run ./cmd/worker` connects to `temporal.riotpiao.com:7233` without error
|
||||
- `go run ./cmd/starter --dry-run` starts a workflow that appears in Temporal Web UI
|
||||
|
||||
<details>
|
||||
<summary>Details</summary>
|
||||
|
||||
**`cmd/worker/main.go`:**
|
||||
```go
|
||||
config := loadConfig() // reads env: TEMPORAL_NAMESPACE, TEMPORAL_TLS_CERT, TEMPORAL_TLS_KEY, ANTHROPIC_API_KEY
|
||||
c, err := client.Dial(client.Options{HostPort: "temporal.riotpiao.com:7233", ...TLS...})
|
||||
w, err := worker.New(c, "default", worker.Options{})
|
||||
// register both workflows
|
||||
w.RegisterWorkflow(statemachine.OrchestratorWorkflow)
|
||||
w.RegisterWorkflow(statemachine.TaskUnitWorkflow)
|
||||
// register all activities
|
||||
w.RegisterActivity(action.CloneRepoActivity)
|
||||
w.RegisterActivity(action.GitWorktreeAddActivity)
|
||||
// ... etc
|
||||
w.Run()
|
||||
```
|
||||
|
||||
**`cmd/starter/main.go`:**
|
||||
```go
|
||||
flag.String("repo", "", "target repo path")
|
||||
flag.String("remote", "", "remote URL")
|
||||
flag.String("milestone", "T0", "milestone ID")
|
||||
flag.Bool("dry-run", false, "disable git push/merge")
|
||||
flag.String("planner-model", "claude-opus-5", "planner model ID")
|
||||
// ... judge, implementer models
|
||||
// build OrchestratorInput, call client.ExecuteWorkflow
|
||||
```
|
||||
|
||||
**`internal/config/config.go`:**
|
||||
- Load Temporal settings from env
|
||||
- Load ANTHROPIC_API_KEY from env
|
||||
- Return filled config struct
|
||||
|
||||
</details>
|
||||
|
||||
---
|
||||
|
||||
### T0.9: End-to-End Test
|
||||
|
||||
Run against real `temporal.riotpiao.com` + disposable scratch repo.
|
||||
|
||||
**Verification (all 7 items in the plan):**
|
||||
1. Clone bootstrap: fresh clone when repo path empty
|
||||
2. Full cycle: dispatch subtasks, judge pass/fail, commit, squash-merge
|
||||
3. Live signal updates: change prompt/skills mid-run, next dispatch sees them
|
||||
4. 5xx retry-then-succeed + always-503 exhausts at 5m mark
|
||||
5. 504 stream-timeout learning: doubles and is actually used, capped at max
|
||||
6. `continue-as-new` history bounded
|
||||
7. Squash-merge result: main has one squashed commit per submilestone
|
||||
|
||||
<details>
|
||||
<summary>Details</summary>
|
||||
|
||||
**Fixture repo structure:**
|
||||
```
|
||||
tasks/
|
||||
INDEX.md (guidelines)
|
||||
board.json (task list, T0.1-T0.3 with trivial definitions)
|
||||
```
|
||||
|
||||
**Example subtask:** "Create file `output.txt` with content 'hello world'"
|
||||
|
||||
**Run sequence:**
|
||||
1. `go run ./cmd/starter --repo /tmp/fixture --remote [email protected]:scratch/workflow-test.git --dry-run`
|
||||
2. Monitor Temporal Web UI for workflow progress
|
||||
3. Midway, send signals: `temporal workflow signal --workflow-id orch-... --name update-role-prompt ...`
|
||||
4. Confirm next dispatch uses new prompt (assert marker in output file)
|
||||
5. Confirm board updated, lessons file exists (if any failure happened)
|
||||
6. Remove `--dry-run`, repeat against real remote
|
||||
7. Assert final state: real commits on remote, squash-merge on main
|
||||
|
||||
</details>
|
||||
|
||||
---
|
||||
|
||||
## Next Steps
|
||||
|
||||
1. Approve this scaffold (PLAN.md + tasks/INDEX.md + board)
|
||||
2. Start T0.1 → checkout branch `task/T0.1` → scaffold repo structure
|
||||
3. Each task: implement, test locally, verify against criterion
|
||||
4. Mark on board: [x] when verification passes
|
||||
5. T0.9: final e2e run
|
||||
6. Squash all T0.* branches into main, push
|
||||
@@ -1,269 +0,0 @@
|
||||
# Temporal Integration for Poimen Workflows
|
||||
|
||||
## Overview
|
||||
|
||||
This project uses **Temporal** for distributed workflow orchestration. Instead of connecting directly to Temporal ports, we use the **REST API Gateway** at `https://api.riotpiao.com/workflow`.
|
||||
|
||||
**Reference Documentation**: See `~/workplace/homelab-frontend/TEMPORAL_USAGE.md` for full API details.
|
||||
|
||||
---
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Configuration
|
||||
|
||||
The Temporal connection is configured via environment variables:
|
||||
|
||||
```bash
|
||||
TEMPORAL_NAMESPACE=poimen-harness # Default namespace
|
||||
TEMPORAL_HOSTPORT=api.riotpiao.com/workflow # REST API gateway (CI only)
|
||||
# Direct gRPC in K8s:
|
||||
TEMPORAL_HOSTPORT=temporal-frontend.temporal:7233 # K8s DNS
|
||||
```
|
||||
|
||||
### For CI/CD (Proper Authentication via PAT Token)
|
||||
|
||||
The CI runner uses a PAT (Personal Access Token) for Forgejo authentication. Integration tests gracefully handle Temporal availability:
|
||||
|
||||
1. **Git authentication configured** in CI:
|
||||
- `.gitea/workflows/ci.yaml` uses `${{ secrets.REGISTRY_PAT }}` token
|
||||
- Enables private module access and authenticated requests
|
||||
|
||||
2. **Integration tests behavior**:
|
||||
```bash
|
||||
go test -v ./... # Runs all tests
|
||||
```
|
||||
- If Temporal accessible: ✅ Tests run
|
||||
- If Temporal unavailable: ⏭️ Tests skip gracefully
|
||||
|
||||
3. **Local development** (with Temporal access):
|
||||
```bash
|
||||
go test -v -run TestTemporal ./tests
|
||||
```
|
||||
|
||||
4. **Graceful fallback**:
|
||||
```go
|
||||
// tests/temporal_integration_test.go
|
||||
if err != nil {
|
||||
t.Skipf("skipping: Temporal not accessible - %v", err)
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Rest API Gateway Usage
|
||||
|
||||
### Base URL
|
||||
```
|
||||
https://api.riotpiao.com/workflow
|
||||
```
|
||||
|
||||
### Example: Start a Workflow (from CI)
|
||||
|
||||
Instead of:
|
||||
```go
|
||||
// ❌ This fails in CI (no direct access)
|
||||
c, err := client.Dial(client.Options{
|
||||
HostPort: "127.0.0.1:7233",
|
||||
Namespace: "poimen-harness",
|
||||
})
|
||||
```
|
||||
|
||||
Use HTTP REST calls:
|
||||
```bash
|
||||
curl -X POST https://api.riotpiao.com/workflow \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"action": "START_WORKFLOW",
|
||||
"namespace": "poimen-harness",
|
||||
"payload": {
|
||||
"workflow_id": "test-workflow",
|
||||
"workflow_type": "OrchestratorWorkflow",
|
||||
"task_queue": "poimen-taskqueue",
|
||||
"input": {}
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
### Operations Available
|
||||
|
||||
All standard Temporal operations:
|
||||
- `START_WORKFLOW` - Launch new workflow
|
||||
- `DESCRIBE_WORKFLOW` - Get workflow status
|
||||
- `LIST_WORKFLOWS` - List executions
|
||||
- `GET_WORKFLOW_HISTORY` - View event history
|
||||
- `SIGNAL_WORKFLOW` - Send signals to running workflows
|
||||
- `QUERY_WORKFLOW` - Query workflow state
|
||||
- `TERMINATE_WORKFLOW` - Stop workflow
|
||||
- `CANCEL_WORKFLOW` - Graceful cancellation
|
||||
|
||||
See `~/workplace/homelab-frontend/TEMPORAL_USAGE.md` for full operation reference.
|
||||
|
||||
---
|
||||
|
||||
## Project Structure
|
||||
|
||||
```
|
||||
.
|
||||
├── cmd/
|
||||
│ ├── starter/ - CLI to start workflows (requires Temporal access)
|
||||
│ └── worker/ - Worker that processes tasks
|
||||
├── tests/
|
||||
│ ├── git_test.go - Unit tests (run in CI ✅)
|
||||
│ ├── types_test.go - Unit tests (run in CI ✅)
|
||||
│ └── temporal_integration_test.go - Integration tests (skipped in CI, local only)
|
||||
├── statemachine/
|
||||
│ ├── orchestrator.go - Main workflow definition
|
||||
│ └── taskunit.go - Sub-workflow for tasks
|
||||
└── action/
|
||||
├── git.go - Git operations (activities)
|
||||
├── planner.go - Planning activity
|
||||
├── implementer.go - Implementation activity
|
||||
└── judge.go - Judgment activity
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Running Tests
|
||||
|
||||
### Unit Tests (CI Compatible)
|
||||
```bash
|
||||
go test -v ./tests # ✅ Passes in CI
|
||||
```
|
||||
|
||||
### Integration Tests (Local Only)
|
||||
```bash
|
||||
# Requires TEMPORAL_HOSTPORT to point to accessible Temporal
|
||||
go test -v -run TestTemporal ./tests
|
||||
|
||||
# Or in K8s environment:
|
||||
kubectl exec -it deployment/poimen-worker -- \
|
||||
go test -v ./tests
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Worker Deployment
|
||||
|
||||
### Local Development
|
||||
```bash
|
||||
# Start worker (requires Temporal access)
|
||||
TEMPORAL_HOSTPORT=localhost:7233 go run ./cmd/worker
|
||||
```
|
||||
|
||||
### Kubernetes
|
||||
```bash
|
||||
kubectl apply -k k8s/
|
||||
# Workers connect to temporal-frontend.temporal:7233 (K8s DNS)
|
||||
```
|
||||
|
||||
### Configuration
|
||||
See `k8s/configmap.yaml`:
|
||||
```yaml
|
||||
TEMPORAL_NAMESPACE: "poimen-harness"
|
||||
TEMPORAL_HOSTPORT: "temporal-frontend.temporal:7233"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## CI/CD Pipeline
|
||||
|
||||
The `.gitea/workflows/ci.yaml` runs:
|
||||
|
||||
1. **Git Auth** - Configure Forgejo PAT token for authentication
|
||||
2. **Checkout** - Pull code
|
||||
3. **Dependencies** - `go mod download`
|
||||
4. **Tests** - `go test -v ./...`
|
||||
- Unit tests: ✅ Always pass
|
||||
- Integration tests: ✅ Run if Temporal accessible, ⏭️ skip if not
|
||||
5. **Build** - `go build ./cmd/...`
|
||||
6. **Vet** - `go vet ./...`
|
||||
|
||||
✅ **Always passes** - Proper authentication + graceful test fallback
|
||||
|
||||
---
|
||||
|
||||
## Accessing the Temporal UI
|
||||
|
||||
### Web UI
|
||||
```
|
||||
https://api.riotpiao.com (UI frontend)
|
||||
```
|
||||
|
||||
### Metrics
|
||||
```bash
|
||||
curl https://api.riotpiao.com/workflow/metrics
|
||||
```
|
||||
|
||||
### Health Check
|
||||
```bash
|
||||
curl https://api.riotpiao.com/workflow/health
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Environment Variables Reference
|
||||
|
||||
| Variable | Default | Usage | CI |
|
||||
|----------|---------|-------|----|
|
||||
| `TEMPORAL_NAMESPACE` | `poimen-harness` | Workflow namespace | ✅ |
|
||||
| `TEMPORAL_HOSTPORT` | `localhost:7233` | Server address | ✅ (configurable) |
|
||||
| `ANTHROPIC_API_KEY` | (required) | LLM for AI agents | ✅ (secret) |
|
||||
| `GOPRIVATE` | (empty) | Private module auth | ✅ |
|
||||
| `REGISTRY_PAT` | (required) | Forgejo auth token | ✅ (secret) |
|
||||
|
||||
---
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### "connection refused" in CI
|
||||
✅ **Expected & OK** - Integration tests gracefully skip if Temporal unavailable
|
||||
```bash
|
||||
# Check: integration tests handle connection errors
|
||||
go test -v ./tests
|
||||
# Output: SKIP temporal_integration_test.go:32 (Temporal not accessible)
|
||||
```
|
||||
|
||||
### Tests fail locally with "connection refused"
|
||||
Ensure Temporal is accessible:
|
||||
```bash
|
||||
# Check connectivity
|
||||
curl https://api.riotpiao.com/workflow/health
|
||||
|
||||
# Or for local Temporal:
|
||||
nc -zv localhost 7233
|
||||
```
|
||||
|
||||
### Worker can't reach Temporal in K8s
|
||||
Verify:
|
||||
```bash
|
||||
# Check configmap
|
||||
kubectl get cm poimen-config -o yaml
|
||||
|
||||
# Check pod logs
|
||||
kubectl logs deployment/poimen-worker
|
||||
|
||||
# Verify DNS from pod
|
||||
kubectl exec -it deployment/poimen-worker -- \
|
||||
nslookup temporal-frontend.temporal
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Next Steps
|
||||
|
||||
1. ✅ CI tests pass with proper authentication (PAT token)
|
||||
2. ✅ Integration tests run when Temporal accessible, skip otherwise
|
||||
3. 🔄 Local development: access Temporal for full integration test coverage
|
||||
4. 📦 K8s deployment: workers connect to Temporal service
|
||||
5. 📊 Monitor via REST API: `https://api.riotpiao.com/workflow`
|
||||
|
||||
---
|
||||
|
||||
## References
|
||||
|
||||
- **Full API**: `~/workplace/homelab-frontend/TEMPORAL_USAGE.md`
|
||||
- **K8s Config**: `./k8s/configmap.yaml`
|
||||
- **CI Config**: `.gitea/workflows/ci.yaml`
|
||||
- **Worker Code**: `./cmd/worker/main.go`
|
||||
- **Workflows**: `./statemachine/orchestrator.go`
|
||||
@@ -0,0 +1,296 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// AnalyzeCodeInput is input for AnalyzeCodeActivity
|
||||
type AnalyzeCodeInput struct {
|
||||
Path string `json:"path"`
|
||||
Language string `json:"language,omitempty"`
|
||||
Depth int `json:"depth,omitempty"`
|
||||
}
|
||||
|
||||
// AnalyzeCodeOutput is output from AnalyzeCodeActivity
|
||||
type AnalyzeCodeOutput struct {
|
||||
Quality float64 `json:"quality"`
|
||||
Metrics map[string]interface{} `json:"metrics"`
|
||||
Issues []string `json:"issues"`
|
||||
Summary string `json:"summary"`
|
||||
}
|
||||
|
||||
// AnalyzeCodeActivity analyzes code quality using available tools
|
||||
func AnalyzeCodeActivity(ctx context.Context, in AnalyzeCodeInput) (AnalyzeCodeOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
logger.Info("AnalyzeCodeActivity started", "path", in.Path)
|
||||
|
||||
output := AnalyzeCodeOutput{
|
||||
Quality: 0.0,
|
||||
Metrics: make(map[string]interface{}),
|
||||
Issues: []string{},
|
||||
}
|
||||
|
||||
// Verify path exists
|
||||
if _, err := os.Stat(in.Path); os.IsNotExist(err) {
|
||||
return output, fmt.Errorf("path does not exist: %s", in.Path)
|
||||
}
|
||||
|
||||
// Count files and lines
|
||||
var totalFiles, totalLines int
|
||||
err := filepath.Walk(in.Path, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return nil // Skip errors
|
||||
}
|
||||
if info.IsDir() {
|
||||
// Skip hidden and vendor directories
|
||||
if strings.HasPrefix(info.Name(), ".") || info.Name() == "vendor" || info.Name() == "node_modules" {
|
||||
return filepath.SkipDir
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
ext := filepath.Ext(path)
|
||||
if isCodeFile(ext) {
|
||||
totalFiles++
|
||||
if lines, err := countLines(path); err == nil {
|
||||
totalLines += lines
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
logger.Warn("Error walking path", "error", err)
|
||||
}
|
||||
|
||||
output.Metrics["totalFiles"] = totalFiles
|
||||
output.Metrics["totalLines"] = totalLines
|
||||
|
||||
// Try to run go vet if it's a Go project
|
||||
if _, err := os.Stat(filepath.Join(in.Path, "go.mod")); err == nil {
|
||||
cmd := exec.CommandContext(ctx, "go", "vet", "./...")
|
||||
cmd.Dir = in.Path
|
||||
vetOutput, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
issues := strings.Split(string(vetOutput), "\n")
|
||||
for _, issue := range issues {
|
||||
if strings.TrimSpace(issue) != "" {
|
||||
output.Issues = append(output.Issues, issue)
|
||||
}
|
||||
}
|
||||
}
|
||||
output.Metrics["goVetRan"] = true
|
||||
}
|
||||
|
||||
// Calculate quality score (simple heuristic)
|
||||
issueCount := len(output.Issues)
|
||||
if totalFiles > 0 {
|
||||
issuesPerFile := float64(issueCount) / float64(totalFiles)
|
||||
output.Quality = max(0, 1.0 - (issuesPerFile * 0.1))
|
||||
} else {
|
||||
output.Quality = 0.5
|
||||
}
|
||||
|
||||
output.Summary = fmt.Sprintf("Analyzed %d files (%d lines). Found %d issues. Quality score: %.2f",
|
||||
totalFiles, totalLines, issueCount, output.Quality)
|
||||
|
||||
logger.Info("AnalyzeCodeActivity completed", "quality", output.Quality, "issues", issueCount)
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// SecurityScanInput is input for SecurityScanActivity
|
||||
type SecurityScanInput struct {
|
||||
Path string `json:"path"`
|
||||
Severity string `json:"severity,omitempty"` // low, medium, high, critical
|
||||
}
|
||||
|
||||
// SecurityScanOutput is output from SecurityScanActivity
|
||||
type SecurityScanOutput struct {
|
||||
Vulnerabilities []Vulnerability `json:"vulnerabilities"`
|
||||
SecurityScore float64 `json:"securityScore"`
|
||||
RiskLevel string `json:"riskLevel"`
|
||||
}
|
||||
|
||||
// Vulnerability represents a security issue
|
||||
type Vulnerability struct {
|
||||
ID string `json:"id"`
|
||||
Severity string `json:"severity"`
|
||||
Description string `json:"description"`
|
||||
File string `json:"file,omitempty"`
|
||||
Line int `json:"line,omitempty"`
|
||||
}
|
||||
|
||||
// SecurityScanActivity scans code for security vulnerabilities
|
||||
func SecurityScanActivity(ctx context.Context, in SecurityScanInput) (SecurityScanOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
logger.Info("SecurityScanActivity started", "path", in.Path)
|
||||
|
||||
output := SecurityScanOutput{
|
||||
Vulnerabilities: []Vulnerability{},
|
||||
SecurityScore: 100.0,
|
||||
RiskLevel: "low",
|
||||
}
|
||||
|
||||
// Verify path exists
|
||||
if _, err := os.Stat(in.Path); os.IsNotExist(err) {
|
||||
return output, fmt.Errorf("path does not exist: %s", in.Path)
|
||||
}
|
||||
|
||||
// Check for common security issues
|
||||
// 1. Check for hardcoded secrets
|
||||
secretPatterns := []string{
|
||||
"password=",
|
||||
"secret=",
|
||||
"api_key=",
|
||||
"apikey=",
|
||||
"private_key",
|
||||
"AWS_SECRET",
|
||||
}
|
||||
|
||||
err := filepath.Walk(in.Path, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil || info.IsDir() {
|
||||
return nil
|
||||
}
|
||||
|
||||
if !isCodeFile(filepath.Ext(path)) {
|
||||
return nil
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
contentStr := strings.ToLower(string(content))
|
||||
for _, pattern := range secretPatterns {
|
||||
if strings.Contains(contentStr, pattern) {
|
||||
output.Vulnerabilities = append(output.Vulnerabilities, Vulnerability{
|
||||
ID: fmt.Sprintf("SEC-%d", len(output.Vulnerabilities)+1),
|
||||
Severity: "high",
|
||||
Description: fmt.Sprintf("Possible hardcoded secret: %s", pattern),
|
||||
File: path,
|
||||
})
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
logger.Warn("Error scanning", "error", err)
|
||||
}
|
||||
|
||||
// Try gosec if available and it's a Go project
|
||||
if _, err := os.Stat(filepath.Join(in.Path, "go.mod")); err == nil {
|
||||
if _, err := exec.LookPath("gosec"); err == nil {
|
||||
cmd := exec.CommandContext(ctx, "gosec", "-fmt=json", "-quiet", "./...")
|
||||
cmd.Dir = in.Path
|
||||
// gosec returns non-zero if issues found, so we ignore the error
|
||||
cmd.CombinedOutput()
|
||||
}
|
||||
}
|
||||
|
||||
// Calculate score
|
||||
vulnCount := len(output.Vulnerabilities)
|
||||
if vulnCount == 0 {
|
||||
output.SecurityScore = 100.0
|
||||
output.RiskLevel = "low"
|
||||
} else if vulnCount < 3 {
|
||||
output.SecurityScore = 80.0
|
||||
output.RiskLevel = "medium"
|
||||
} else if vulnCount < 10 {
|
||||
output.SecurityScore = 50.0
|
||||
output.RiskLevel = "high"
|
||||
} else {
|
||||
output.SecurityScore = 20.0
|
||||
output.RiskLevel = "critical"
|
||||
}
|
||||
|
||||
logger.Info("SecurityScanActivity completed", "vulnerabilities", vulnCount, "riskLevel", output.RiskLevel)
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// GenerateReportInput is input for GenerateReportActivity
|
||||
type GenerateReportInput struct {
|
||||
AnalysisResult interface{} `json:"analysisResult"`
|
||||
SecurityResult interface{} `json:"securityResult"`
|
||||
Format string `json:"format,omitempty"` // markdown, html, json
|
||||
}
|
||||
|
||||
// GenerateReportOutput is output from GenerateReportActivity
|
||||
type GenerateReportOutput struct {
|
||||
Report string `json:"report"`
|
||||
ReportPath string `json:"reportPath"`
|
||||
}
|
||||
|
||||
// GenerateReportActivity generates a combined report
|
||||
func GenerateReportActivity(ctx context.Context, in GenerateReportInput) (GenerateReportOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
logger.Info("GenerateReportActivity started")
|
||||
|
||||
format := in.Format
|
||||
if format == "" {
|
||||
format = "markdown"
|
||||
}
|
||||
|
||||
var report strings.Builder
|
||||
timestamp := time.Now().Format(time.RFC3339)
|
||||
|
||||
switch format {
|
||||
case "markdown":
|
||||
report.WriteString("# Analysis Report\n\n")
|
||||
report.WriteString(fmt.Sprintf("Generated: %s\n\n", timestamp))
|
||||
report.WriteString("## Code Analysis\n\n")
|
||||
report.WriteString(fmt.Sprintf("```\n%v\n```\n\n", in.AnalysisResult))
|
||||
report.WriteString("## Security Scan\n\n")
|
||||
report.WriteString(fmt.Sprintf("```\n%v\n```\n\n", in.SecurityResult))
|
||||
case "json":
|
||||
report.WriteString(fmt.Sprintf(`{"timestamp":"%s","analysis":%v,"security":%v}`,
|
||||
timestamp, in.AnalysisResult, in.SecurityResult))
|
||||
default:
|
||||
report.WriteString(fmt.Sprintf("Report generated at %s\n", timestamp))
|
||||
report.WriteString(fmt.Sprintf("Analysis: %v\n", in.AnalysisResult))
|
||||
report.WriteString(fmt.Sprintf("Security: %v\n", in.SecurityResult))
|
||||
}
|
||||
|
||||
// Save to temp file
|
||||
reportPath := filepath.Join(os.TempDir(), fmt.Sprintf("report-%d.%s", time.Now().UnixNano(), format))
|
||||
if err := os.WriteFile(reportPath, []byte(report.String()), 0644); err != nil {
|
||||
logger.Warn("Failed to save report", "error", err)
|
||||
}
|
||||
|
||||
logger.Info("GenerateReportActivity completed", "format", format)
|
||||
return GenerateReportOutput{
|
||||
Report: report.String(),
|
||||
ReportPath: reportPath,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Helper functions
|
||||
|
||||
func isCodeFile(ext string) bool {
|
||||
codeExts := map[string]bool{
|
||||
".go": true, ".py": true, ".js": true, ".ts": true,
|
||||
".java": true, ".c": true, ".cpp": true, ".h": true,
|
||||
".rs": true, ".rb": true, ".php": true, ".swift": true,
|
||||
}
|
||||
return codeExts[ext]
|
||||
}
|
||||
|
||||
func countLines(path string) (int, error) {
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return len(strings.Split(string(content), "\n")), nil
|
||||
}
|
||||
|
||||
func max(a, b float64) float64 {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAnalyzeCodeActivity(t *testing.T) {
|
||||
// Create temp directory with some Go code
|
||||
tmpDir, err := os.MkdirTemp("", "analyze-test")
|
||||
require.NoError(t, err)
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
// Create go.mod
|
||||
goMod := `module test
|
||||
go 1.21
|
||||
`
|
||||
require.NoError(t, os.WriteFile(filepath.Join(tmpDir, "go.mod"), []byte(goMod), 0644))
|
||||
|
||||
// Create a simple Go file
|
||||
goCode := `package main
|
||||
|
||||
func main() {
|
||||
println("hello")
|
||||
}
|
||||
`
|
||||
require.NoError(t, os.WriteFile(filepath.Join(tmpDir, "main.go"), []byte(goCode), 0644))
|
||||
|
||||
// Run activity
|
||||
output, err := AnalyzeCodeActivity(context.Background(), AnalyzeCodeInput{
|
||||
Path: tmpDir,
|
||||
Depth: 3,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify output
|
||||
require.Greater(t, output.Quality, 0.0)
|
||||
require.NotNil(t, output.Metrics)
|
||||
require.NotEmpty(t, output.Summary)
|
||||
|
||||
totalFiles, ok := output.Metrics["totalFiles"].(int)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, 1, totalFiles) // Just main.go (go.mod not counted)
|
||||
}
|
||||
|
||||
func TestAnalyzeCodeActivity_PathNotExist(t *testing.T) {
|
||||
_, err := AnalyzeCodeActivity(context.Background(), AnalyzeCodeInput{
|
||||
Path: "/nonexistent/path",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "does not exist")
|
||||
}
|
||||
|
||||
func TestSecurityScanActivity(t *testing.T) {
|
||||
// Create temp directory
|
||||
tmpDir, err := os.MkdirTemp("", "security-test")
|
||||
require.NoError(t, err)
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
// Create a file with potential secret (matches pattern "password=")
|
||||
code := `package main
|
||||
|
||||
var config = map[string]string{
|
||||
"password=": "supersecret123",
|
||||
"api_key=": "sk-12345",
|
||||
}
|
||||
`
|
||||
require.NoError(t, os.WriteFile(filepath.Join(tmpDir, "main.go"), []byte(code), 0644))
|
||||
|
||||
// Run activity
|
||||
output, err := SecurityScanActivity(context.Background(), SecurityScanInput{
|
||||
Path: tmpDir,
|
||||
Severity: "medium",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Should find the hardcoded password
|
||||
require.Greater(t, len(output.Vulnerabilities), 0)
|
||||
require.Less(t, output.SecurityScore, 100.0)
|
||||
}
|
||||
|
||||
func TestSecurityScanActivity_Clean(t *testing.T) {
|
||||
// Create temp directory with clean code
|
||||
tmpDir, err := os.MkdirTemp("", "security-clean-test")
|
||||
require.NoError(t, err)
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
// Create clean code
|
||||
code := `package main
|
||||
|
||||
func main() {
|
||||
println("hello")
|
||||
}
|
||||
`
|
||||
require.NoError(t, os.WriteFile(filepath.Join(tmpDir, "main.go"), []byte(code), 0644))
|
||||
|
||||
// Run activity
|
||||
output, err := SecurityScanActivity(context.Background(), SecurityScanInput{
|
||||
Path: tmpDir,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Should be clean
|
||||
require.Equal(t, 0, len(output.Vulnerabilities))
|
||||
require.Equal(t, 100.0, output.SecurityScore)
|
||||
require.Equal(t, "low", output.RiskLevel)
|
||||
}
|
||||
|
||||
func TestGenerateReportActivity(t *testing.T) {
|
||||
output, err := GenerateReportActivity(context.Background(), GenerateReportInput{
|
||||
AnalysisResult: map[string]interface{}{"quality": 0.85},
|
||||
SecurityResult: map[string]interface{}{"score": 95.0},
|
||||
Format: "markdown",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Contains(t, output.Report, "# Analysis Report")
|
||||
require.NotEmpty(t, output.ReportPath)
|
||||
|
||||
// Verify file was created
|
||||
_, err = os.Stat(output.ReportPath)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Cleanup
|
||||
os.Remove(output.ReportPath)
|
||||
}
|
||||
|
||||
func TestGenerateReportActivity_JSON(t *testing.T) {
|
||||
output, err := GenerateReportActivity(context.Background(), GenerateReportInput{
|
||||
AnalysisResult: map[string]interface{}{"quality": 0.85},
|
||||
SecurityResult: map[string]interface{}{"score": 95.0},
|
||||
Format: "json",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Contains(t, output.Report, `"timestamp"`)
|
||||
|
||||
// Cleanup
|
||||
os.Remove(output.ReportPath)
|
||||
}
|
||||
@@ -0,0 +1,216 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"time"
|
||||
)
|
||||
|
||||
// AssumeRoleInput is the input to AssumeRoleActivity
|
||||
type AssumeRoleInput struct {
|
||||
// Identity is the user/service identity requesting access
|
||||
Identity string `json:"identity"`
|
||||
|
||||
// ClientID is the OAuth2/OIDC client ID (from vault or env)
|
||||
ClientID string `json:"clientId,omitempty"`
|
||||
|
||||
// ClientSecret is the OAuth2/OIDC client secret (from vault or env)
|
||||
ClientSecret string `json:"clientSecret,omitempty"`
|
||||
|
||||
// Scope defines what APIs this token can access (e.g., "llm:read llm:write")
|
||||
Scope string `json:"scope"`
|
||||
|
||||
// DurationSeconds is how long the token is valid (default: 3600 = 1 hour)
|
||||
DurationSeconds int `json:"durationSeconds,omitempty"`
|
||||
|
||||
// AuthServerURL is the auth server endpoint (from env if not provided)
|
||||
AuthServerURL string `json:"authServerUrl,omitempty"`
|
||||
}
|
||||
|
||||
// AssumeRoleOutput is the output from AssumeRoleActivity
|
||||
type AssumeRoleOutput struct {
|
||||
// Token is the JWT token for calling api.riotpiao.com
|
||||
Token string `json:"token"`
|
||||
|
||||
// ExpiresAt is when the token expires (Unix timestamp)
|
||||
ExpiresAt int64 `json:"expiresAt"`
|
||||
|
||||
// ExpiresIn is the duration in seconds until expiration
|
||||
ExpiresIn int `json:"expiresIn"`
|
||||
|
||||
// TokenType is typically "Bearer"
|
||||
TokenType string `json:"tokenType"`
|
||||
}
|
||||
|
||||
// oauthTokenRequest is sent to the auth server
|
||||
type oauthTokenRequest struct {
|
||||
GrantType string `json:"grant_type"`
|
||||
ClientID string `json:"client_id"`
|
||||
ClientSecret string `json:"client_secret"`
|
||||
Scope string `json:"scope"`
|
||||
Subject string `json:"subject,omitempty"` // The identity being assumed
|
||||
}
|
||||
|
||||
// oauthTokenResponse is returned from the auth server
|
||||
type oauthTokenResponse struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
TokenType string `json:"token_type"`
|
||||
ExpiresIn int `json:"expires_in"`
|
||||
Scope string `json:"scope"`
|
||||
}
|
||||
|
||||
// AssumeRoleActivity requests a temporary JWT token for accessing LLM APIs
|
||||
//
|
||||
// This activity works like AWS AssumeRole:
|
||||
// 1. User provides identity + scope of access needed
|
||||
// 2. Activity exchanges credentials with auth server
|
||||
// 3. Returns JWT token valid for a limited time
|
||||
// 4. Caller uses token in subsequent LLM API calls
|
||||
//
|
||||
// Security: Credentials should come from vault/secrets, never hardcoded
|
||||
func AssumeRoleActivity(ctx context.Context, input *AssumeRoleInput) (*AssumeRoleOutput, error) {
|
||||
// Validate inputs
|
||||
if err := validateAssumeRoleInput(input); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Resolve configuration from input + environment
|
||||
config, err := resolveAssumeRoleConfig(input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Request token from auth server
|
||||
tokenResp, err := requestAuthToken(ctx, config, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Build output
|
||||
return buildAssumeRoleOutput(tokenResp), nil
|
||||
}
|
||||
|
||||
// validateAssumeRoleInput checks required fields
|
||||
func validateAssumeRoleInput(input *AssumeRoleInput) error {
|
||||
if input.Identity == "" {
|
||||
return fmt.Errorf("identity is required")
|
||||
}
|
||||
if input.Scope == "" {
|
||||
return fmt.Errorf("scope is required (e.g., 'llm:read' or 'llm:read llm:write')")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// assumeRoleConfig holds resolved configuration
|
||||
type assumeRoleConfig struct {
|
||||
authServerURL string
|
||||
clientID string
|
||||
clientSecret string
|
||||
duration int
|
||||
}
|
||||
|
||||
// resolveAssumeRoleConfig gets config from input or environment
|
||||
func resolveAssumeRoleConfig(input *AssumeRoleInput) (*assumeRoleConfig, error) {
|
||||
cfg := &assumeRoleConfig{}
|
||||
|
||||
// Helper function to avoid DRY violation
|
||||
getOrEnv := func(val, envKey, fieldName string) (string, error) {
|
||||
if val != "" {
|
||||
return val, nil
|
||||
}
|
||||
if val = os.Getenv(envKey); val != "" {
|
||||
return val, nil
|
||||
}
|
||||
return "", fmt.Errorf("%s not provided and %s not set", fieldName, envKey)
|
||||
}
|
||||
|
||||
var err error
|
||||
if cfg.authServerURL, err = getOrEnv(input.AuthServerURL, "AUTH_SERVER_URL", "authServerUrl"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cfg.clientID, err = getOrEnv(input.ClientID, "OAUTH_CLIENT_ID", "clientId"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cfg.clientSecret, err = getOrEnv(input.ClientSecret, "OAUTH_CLIENT_SECRET", "clientSecret"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Validate and set duration
|
||||
cfg.duration = input.DurationSeconds
|
||||
if cfg.duration == 0 {
|
||||
cfg.duration = 3600 // 1 hour default
|
||||
}
|
||||
if cfg.duration > 86400 {
|
||||
cfg.duration = 86400 // Max 24 hours
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// requestAuthToken calls the auth server and returns the token response
|
||||
func requestAuthToken(ctx context.Context, config *assumeRoleConfig, input *AssumeRoleInput) (*oauthTokenResponse, error) {
|
||||
tokenReq := oauthTokenRequest{
|
||||
GrantType: "client_credentials",
|
||||
ClientID: config.clientID,
|
||||
ClientSecret: config.clientSecret,
|
||||
Scope: input.Scope,
|
||||
Subject: input.Identity,
|
||||
}
|
||||
|
||||
reqBody, err := json.Marshal(tokenReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal token request: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, "POST",
|
||||
fmt.Sprintf("%s/oauth/token", config.authServerURL),
|
||||
bytes.NewReader(reqBody))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := &http.Client{Timeout: 10 * time.Second}
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to call auth server: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("auth server returned status %d: %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
var tokenResp oauthTokenResponse
|
||||
if err := json.Unmarshal(respBody, &tokenResp); err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal token response: %w", err)
|
||||
}
|
||||
|
||||
if tokenResp.AccessToken == "" {
|
||||
return nil, fmt.Errorf("auth server returned empty access token")
|
||||
}
|
||||
|
||||
return &tokenResp, nil
|
||||
}
|
||||
|
||||
// buildAssumeRoleOutput constructs the output from token response
|
||||
func buildAssumeRoleOutput(tokenResp *oauthTokenResponse) *AssumeRoleOutput {
|
||||
expiresAt := time.Now().Add(time.Duration(tokenResp.ExpiresIn) * time.Second).Unix()
|
||||
return &AssumeRoleOutput{
|
||||
Token: tokenResp.AccessToken,
|
||||
ExpiresAt: expiresAt,
|
||||
ExpiresIn: tokenResp.ExpiresIn,
|
||||
TokenType: tokenResp.TokenType,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,295 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/pkg/db"
|
||||
)
|
||||
|
||||
// IncompatibilityWarning explains why two activities can't be connected
|
||||
type IncompatibilityWarning struct {
|
||||
Source string `json:"source"` // Source node ID
|
||||
Target string `json:"target"` // Target node ID
|
||||
Reason string `json:"reason"` // Why they can't connect
|
||||
SourceNeeds string `json:"source_needs"` // What source would need to output
|
||||
TargetNeeds string `json:"target_needs"` // What target requires as input
|
||||
Suggestion string `json:"suggestion"` // Suggestion to make it work
|
||||
}
|
||||
|
||||
// ActivitySchema describes what an activity needs/provides
|
||||
type ActivitySchema struct {
|
||||
ActivityType string `json:"activity_type"`
|
||||
Inputs map[string]InputField `json:"inputs"`
|
||||
Outputs map[string]OutputField `json:"outputs"`
|
||||
}
|
||||
|
||||
type InputField struct {
|
||||
Type string `json:"type"`
|
||||
Description string `json:"description"`
|
||||
Required bool `json:"required"`
|
||||
Enum []string `json:"enum,omitempty"`
|
||||
}
|
||||
|
||||
type OutputField struct {
|
||||
Type string `json:"type"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
// getActivitySchema returns schema from knowledge base
|
||||
func getActivitySchema(activityType string) (*ActivitySchema, error) {
|
||||
kb := knowledgeBaseData()
|
||||
if kb == nil {
|
||||
return nil, fmt.Errorf("knowledge base not loaded")
|
||||
}
|
||||
|
||||
var activities []map[string]interface{}
|
||||
if err := json.Unmarshal([]byte(kb), &activities); err != nil {
|
||||
// Try to extract activities from full KB structure
|
||||
var fullKB map[string]interface{}
|
||||
if err := json.Unmarshal([]byte(kb), &fullKB); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse knowledge base")
|
||||
}
|
||||
if activitiesRaw, ok := fullKB["activities"]; ok {
|
||||
if b, err := json.Marshal(activitiesRaw); err == nil {
|
||||
if err := json.Unmarshal(b, &activities); err != nil {
|
||||
return nil, fmt.Errorf("failed to extract activities from KB")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Find matching activity
|
||||
for _, act := range activities {
|
||||
if name, ok := act["name"].(string); ok {
|
||||
if toActivityName(activityType) == name {
|
||||
// Convert to ActivitySchema
|
||||
schema := &ActivitySchema{
|
||||
ActivityType: activityType,
|
||||
Inputs: make(map[string]InputField),
|
||||
Outputs: make(map[string]OutputField),
|
||||
}
|
||||
|
||||
if inputs, ok := act["inputs"].(map[string]interface{}); ok {
|
||||
for key, val := range inputs {
|
||||
if field, ok := val.(map[string]interface{}); ok {
|
||||
schema.Inputs[key] = parseInputField(field)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if outputs, ok := act["outputs"].(map[string]interface{}); ok {
|
||||
for key, val := range outputs {
|
||||
if field, ok := val.(map[string]interface{}); ok {
|
||||
schema.Outputs[key] = parseOutputField(field)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return schema, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("activity %s not found in knowledge base", activityType)
|
||||
}
|
||||
|
||||
func parseInputField(data map[string]interface{}) InputField {
|
||||
field := InputField{}
|
||||
if t, ok := data["type"].(string); ok {
|
||||
field.Type = t
|
||||
}
|
||||
if d, ok := data["description"].(string); ok {
|
||||
field.Description = d
|
||||
}
|
||||
if r, ok := data["required"].(bool); ok {
|
||||
field.Required = r
|
||||
}
|
||||
return field
|
||||
}
|
||||
|
||||
func parseOutputField(data map[string]interface{}) OutputField {
|
||||
field := OutputField{}
|
||||
if t, ok := data["type"].(string); ok {
|
||||
field.Type = t
|
||||
}
|
||||
if d, ok := data["description"].(string); ok {
|
||||
field.Description = d
|
||||
}
|
||||
return field
|
||||
}
|
||||
|
||||
// CheckConnectionCompatibility validates if source can connect to target
|
||||
func CheckConnectionCompatibility(sourceNode, targetNode db.WorkflowNode) []IncompatibilityWarning {
|
||||
warnings := []IncompatibilityWarning{}
|
||||
|
||||
sourceSchema, err := getActivitySchema(sourceNode.Type)
|
||||
if err != nil {
|
||||
warnings = append(warnings, IncompatibilityWarning{
|
||||
Source: sourceNode.ID,
|
||||
Target: targetNode.ID,
|
||||
Reason: fmt.Sprintf("Source activity schema not found: %v", err),
|
||||
Suggestion: "Ensure source activity type is registered in knowledge base",
|
||||
})
|
||||
return warnings
|
||||
}
|
||||
|
||||
targetSchema, err := getActivitySchema(targetNode.Type)
|
||||
if err != nil {
|
||||
warnings = append(warnings, IncompatibilityWarning{
|
||||
Source: sourceNode.ID,
|
||||
Target: targetNode.ID,
|
||||
Reason: fmt.Sprintf("Target activity schema not found: %v", err),
|
||||
Suggestion: "Ensure target activity type is registered in knowledge base",
|
||||
})
|
||||
return warnings
|
||||
}
|
||||
|
||||
// Check if source produces outputs that target can consume
|
||||
if len(sourceSchema.Outputs) == 0 {
|
||||
warnings = append(warnings, IncompatibilityWarning{
|
||||
Source: sourceNode.ID,
|
||||
Target: targetNode.ID,
|
||||
Reason: fmt.Sprintf("%s produces no outputs", sourceNode.Type),
|
||||
SourceNeeds: "any output",
|
||||
Suggestion: "Source activity must produce outputs",
|
||||
})
|
||||
return warnings
|
||||
}
|
||||
|
||||
if len(targetSchema.Inputs) == 0 {
|
||||
warnings = append(warnings, IncompatibilityWarning{
|
||||
Source: sourceNode.ID,
|
||||
Target: targetNode.ID,
|
||||
Reason: fmt.Sprintf("%s accepts no inputs", targetNode.Type),
|
||||
TargetNeeds: "no input",
|
||||
Suggestion: "Target activity must accept inputs. Check if it's a terminal activity.",
|
||||
})
|
||||
return warnings
|
||||
}
|
||||
|
||||
// Match outputs to inputs
|
||||
sourceOutputs := getOutputNames(sourceSchema.Outputs)
|
||||
targetInputs := getInputNames(targetSchema.Inputs)
|
||||
|
||||
if len(sourceOutputs) == 0 || len(targetInputs) == 0 {
|
||||
warnings = append(warnings, IncompatibilityWarning{
|
||||
Source: sourceNode.ID,
|
||||
Target: targetNode.ID,
|
||||
Reason: "No compatible output/input fields found",
|
||||
SourceNeeds: strings.Join(sourceOutputs, ", "),
|
||||
TargetNeeds: strings.Join(targetInputs, ", "),
|
||||
Suggestion: "Use LLM transformation to map outputs to inputs",
|
||||
})
|
||||
}
|
||||
|
||||
return warnings
|
||||
}
|
||||
|
||||
// CheckCanvasConnectivity analyzes all suggested edges for compatibility
|
||||
func CheckCanvasConnectivity(nodes []db.WorkflowNode, suggestedEdges []db.WorkflowEdge) []IncompatibilityWarning {
|
||||
warnings := []IncompatibilityWarning{}
|
||||
nodeMap := make(map[string]db.WorkflowNode)
|
||||
for _, n := range nodes {
|
||||
nodeMap[n.ID] = n
|
||||
}
|
||||
|
||||
for _, edge := range suggestedEdges {
|
||||
sourceNode, ok := nodeMap[edge.Source]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
targetNode, ok := nodeMap[edge.Target]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
edgeWarnings := CheckConnectionCompatibility(sourceNode, targetNode)
|
||||
warnings = append(warnings, edgeWarnings...)
|
||||
}
|
||||
|
||||
return warnings
|
||||
}
|
||||
|
||||
// IdentifyDisconnectedNodes finds nodes that can't connect to anything
|
||||
func IdentifyDisconnectedNodes(nodes []db.WorkflowNode, suggestedEdges []db.WorkflowEdge) []string {
|
||||
edgeMap := make(map[string]bool)
|
||||
for _, edge := range suggestedEdges {
|
||||
edgeMap[edge.Source] = true
|
||||
edgeMap[edge.Target] = true
|
||||
}
|
||||
|
||||
var disconnected []string
|
||||
for _, node := range nodes {
|
||||
if !edgeMap[node.ID] {
|
||||
disconnected = append(disconnected, node.ID)
|
||||
}
|
||||
}
|
||||
return disconnected
|
||||
}
|
||||
|
||||
// toActivityName converts canvas type to activity name (e.g., "clone-repo" -> "CloneRepoActivity")
|
||||
func toActivityName(canvasType string) string {
|
||||
parts := strings.Split(canvasType, "-")
|
||||
var result string
|
||||
for _, part := range parts {
|
||||
if part != "" {
|
||||
result += strings.ToUpper(part[:1]) + strings.ToLower(part[1:])
|
||||
}
|
||||
}
|
||||
return result + "Activity"
|
||||
}
|
||||
|
||||
// getOutputNames extracts output field names
|
||||
func getOutputNames(outputs map[string]OutputField) []string {
|
||||
var names []string
|
||||
for name := range outputs {
|
||||
names = append(names, name)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// getInputNames extracts input field names (required ones highlighted)
|
||||
func getInputNames(inputs map[string]InputField) []string {
|
||||
var names []string
|
||||
for name, field := range inputs {
|
||||
if field.Required {
|
||||
names = append(names, name+"*")
|
||||
} else {
|
||||
names = append(names, name)
|
||||
}
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// SuggestDataTransformation proposes how to connect incompatible activities
|
||||
func SuggestDataTransformation(sourceNode, targetNode db.WorkflowNode) string {
|
||||
sourceSchema, _ := getActivitySchema(sourceNode.Type)
|
||||
targetSchema, _ := getActivitySchema(targetNode.Type)
|
||||
|
||||
if sourceSchema == nil || targetSchema == nil {
|
||||
return "Cannot analyze compatibility without schemas"
|
||||
}
|
||||
|
||||
sourceOuts := getOutputNames(sourceSchema.Outputs)
|
||||
targetIns := getInputNames(targetSchema.Inputs)
|
||||
|
||||
return fmt.Sprintf(
|
||||
"To connect %s → %s:\n"+
|
||||
" %s outputs: %s\n"+
|
||||
" %s needs: %s\n"+
|
||||
" Solution: Use LLM transformation node to map outputs to inputs",
|
||||
sourceNode.Label, targetNode.Label,
|
||||
sourceNode.Type, strings.Join(sourceOuts, ", "),
|
||||
targetNode.Type, strings.Join(targetIns, ", "),
|
||||
)
|
||||
}
|
||||
|
||||
// knowledgeBaseData returns raw KB JSON (stub - implement with actual KB loading)
|
||||
func knowledgeBaseData() string {
|
||||
// This would load from activity_knowledge_base.json
|
||||
// For now, return empty - real implementation loads from file
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,283 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/action/llm"
|
||||
"github.com/rockliang/poimen/workflows/pkg/db"
|
||||
"github.com/rockliang/poimen/workflows/statemachine"
|
||||
)
|
||||
|
||||
// CanvasReasonerInput infers connections between nodes using LLM reasoning
|
||||
type CanvasReasonerInput struct {
|
||||
Nodes []db.WorkflowNode `json:"nodes"` // Canvas nodes
|
||||
Edges []db.WorkflowEdge `json:"edges"` // Existing edges
|
||||
// If true, only suggest new edges; if false, redesign entire canvas
|
||||
PreserveExisting bool `json:"preserve_existing,omitempty"`
|
||||
AuthToken string `json:"auth_token,omitempty"` // JWT for LLM calls
|
||||
}
|
||||
|
||||
// RelationWording describes semantic meaning of an edge
|
||||
type RelationWording struct {
|
||||
Verb string `json:"verb"` // outputs, inputs, depends-on, etc
|
||||
SourceOutput string `json:"source_output"` // What source produces
|
||||
TargetInput string `json:"target_input"` // What target requires
|
||||
ConnectionType string `json:"connection_type"` // direct-map, requires-transformer, conditional
|
||||
Confidence float64 `json:"confidence"` // 0.0-1.0
|
||||
SemanticMatch string `json:"semantic_match"` // Human-readable explanation
|
||||
TransformerNeeded string `json:"transformer_needed,omitempty"` // If transformation required
|
||||
}
|
||||
|
||||
// EdgeWithWording pairs an edge with its semantic description
|
||||
type EdgeWithWording struct {
|
||||
Source string `json:"source"`
|
||||
Target string `json:"target"`
|
||||
RelationType string `json:"relation_type"` // data-flow, dependency, conditional, parallel
|
||||
RelationLabel string `json:"relation_label"` // e.g., "CloneRepo outputs path → AnalyzeCode requires path"
|
||||
RelationWording RelationWording `json:"relation_wording"`
|
||||
}
|
||||
|
||||
// CanvasReasonerOutput returns suggested edges and reasoning
|
||||
type CanvasReasonerOutput struct {
|
||||
SuggestedEdges []EdgeWithWording `json:"suggested_edges"` // Edges with wording
|
||||
RemovedEdges []db.WorkflowEdge `json:"removed_edges,omitempty"` // Edges to remove
|
||||
Reasoning string `json:"reasoning"` // LLM explanation
|
||||
Confidence float64 `json:"confidence"` // 0.0-1.0
|
||||
IncompatibleEdges []IncompatibilityWarning `json:"incompatible_edges,omitempty"` // Can't connect
|
||||
DisconnectedNodes []string `json:"disconnected_nodes,omitempty"` // No connections
|
||||
UserAlerts []string `json:"user_alerts,omitempty"` // Human-readable warnings
|
||||
}
|
||||
|
||||
// CanvasReasonerActivity uses LLM to infer connections between workflow activities
|
||||
func CanvasReasonerActivity(ctx context.Context, in CanvasReasonerInput) (CanvasReasonerOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
|
||||
output := CanvasReasonerOutput{
|
||||
SuggestedEdges: []db.WorkflowEdge{},
|
||||
}
|
||||
|
||||
if len(in.Nodes) == 0 {
|
||||
return output, fmt.Errorf("no nodes provided")
|
||||
}
|
||||
|
||||
logger.logf("info", "Analyzing canvas with %d nodes, %d edges", len(in.Nodes), len(in.Edges))
|
||||
|
||||
// Build activity descriptions for LLM context
|
||||
nodeDesc := buildNodeDescriptions(in.Nodes)
|
||||
edgeDesc := buildEdgeDescriptions(in.Edges)
|
||||
|
||||
// Create prompt for LLM reasoning with relation wording
|
||||
systemPrompt := `You are a workflow automation expert. Analyze activities and suggest logical connections with semantic descriptions.
|
||||
|
||||
CRITICAL RULES:
|
||||
1. Only suggest edges where outputs→inputs match
|
||||
2. Provide relation wording: verb, source_output, target_input
|
||||
3. Assess connection confidence (0.0-1.0)
|
||||
4. Flag type mismatches that need transformers
|
||||
|
||||
Respond with JSON:
|
||||
{
|
||||
"edges": [
|
||||
{
|
||||
"source": "node-1",
|
||||
"target": "node-2",
|
||||
"relation_type": "data-flow|dependency|conditional|parallel",
|
||||
"relation_label": "Node1 outputs X → Node2 requires X",
|
||||
"relation_wording": {
|
||||
"verb": "outputs|depends-on|triggers|etc",
|
||||
"source_output": "field_name (type): description",
|
||||
"target_input": "field_name (type, required?): description",
|
||||
"connection_type": "direct-map|requires-transformer|conditional",
|
||||
"confidence": 0.95,
|
||||
"semantic_match": "Explanation of why this makes sense"
|
||||
}
|
||||
}
|
||||
],
|
||||
"reasoning": "Overall workflow structure explanation",
|
||||
"confidence": 0.85
|
||||
}`
|
||||
|
||||
userPrompt := fmt.Sprintf(`Canvas Analysis:
|
||||
|
||||
Nodes (including inputs/outputs):
|
||||
%s
|
||||
|
||||
Current Edges:
|
||||
%s
|
||||
|
||||
Task: %s
|
||||
|
||||
KEY RULES:
|
||||
- Preserve existing edges and suggest only NEW edges to add
|
||||
- SKIP any connections where input/output types don't match
|
||||
- If an activity has no outputs, it cannot be a source
|
||||
- If an activity has no inputs, it cannot be a target
|
||||
- Note any activities that are hard to connect (terminal activities, generators, etc)
|
||||
|
||||
Return ONLY valid JSON, no markdown code blocks.`, nodeDesc, edgeDesc, getReasoningTask(in.PreserveExisting))
|
||||
|
||||
logger.logf("info", "Calling LLM reasoning (preserve_existing=%v)", in.PreserveExisting)
|
||||
|
||||
// Call LLM
|
||||
client, err := llm.NewClient()
|
||||
if err != nil {
|
||||
return output, fmt.Errorf("failed to create LLM client: %w", err)
|
||||
}
|
||||
|
||||
response, err := client.CreateMessage(ctx, llm.MessageInput{
|
||||
Model: statemachine.ModelSpec{
|
||||
ModelID: "reasoning", // Use reasoning model for complex analysis
|
||||
},
|
||||
SystemPrompt: systemPrompt,
|
||||
Messages: []llm.MessageParam{
|
||||
{
|
||||
Role: "user",
|
||||
Content: userPrompt,
|
||||
},
|
||||
},
|
||||
AuthToken: in.AuthToken,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return output, fmt.Errorf("LLM reasoning failed: %w", err)
|
||||
}
|
||||
|
||||
// Parse LLM response
|
||||
var reasonerResp struct {
|
||||
Edges []db.WorkflowEdge `json:"edges"`
|
||||
Reasoning string `json:"reasoning"`
|
||||
Confidence float64 `json:"confidence"`
|
||||
}
|
||||
|
||||
if err := json.Unmarshal([]byte(response), &reasonerResp); err != nil {
|
||||
logger.logf("warn", "Failed to parse LLM response as JSON: %v", err)
|
||||
// Try to extract from response text
|
||||
output.Reasoning = response
|
||||
output.Confidence = 0.5
|
||||
return output, fmt.Errorf("failed to parse LLM response: %w", err)
|
||||
}
|
||||
|
||||
// Validate suggested edges
|
||||
nodeMap := make(map[string]bool)
|
||||
for _, n := range in.Nodes {
|
||||
nodeMap[n.ID] = true
|
||||
}
|
||||
|
||||
validEdges := []db.WorkflowEdge{}
|
||||
for _, edge := range reasonerResp.Edges {
|
||||
if !nodeMap[edge.Source] {
|
||||
logger.logf("warn", "Suggested edge references unknown source: %s", edge.Source)
|
||||
continue
|
||||
}
|
||||
if !nodeMap[edge.Target] {
|
||||
logger.logf("warn", "Suggested edge references unknown target: %s", edge.Target)
|
||||
continue
|
||||
}
|
||||
// Don't suggest self-loops
|
||||
if edge.Source == edge.Target {
|
||||
logger.logf("warn", "Skipping self-loop: %s", edge.Source)
|
||||
continue
|
||||
}
|
||||
validEdges = append(validEdges, edge)
|
||||
}
|
||||
|
||||
output.SuggestedEdges = validEdges
|
||||
output.Reasoning = reasonerResp.Reasoning
|
||||
output.Confidence = reasonerResp.Confidence
|
||||
|
||||
// Check compatibility of suggested edges
|
||||
incompatibilities := CheckCanvasConnectivity(in.Nodes, validEdges)
|
||||
if len(incompatibilities) > 0 {
|
||||
output.IncompatibleEdges = incompatibilities
|
||||
logger.logf("warn", "Found %d incompatible edge connections", len(incompatibilities))
|
||||
|
||||
// Generate user-friendly alerts
|
||||
for i, incompat := range incompatibilities {
|
||||
if i < 5 { // Limit to 5 alerts to avoid spam
|
||||
alert := fmt.Sprintf(
|
||||
"⚠️ %s → %s: %s. %s",
|
||||
incompat.Source, incompat.Target, incompat.Reason, incompat.Suggestion,
|
||||
)
|
||||
output.UserAlerts = append(output.UserAlerts, alert)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Identify disconnected nodes
|
||||
disconnected := IdentifyDisconnectedNodes(in.Nodes, validEdges)
|
||||
if len(disconnected) > 0 {
|
||||
output.DisconnectedNodes = disconnected
|
||||
logger.logf("warn", "Found %d disconnected nodes", len(disconnected))
|
||||
|
||||
for _, nodeID := range disconnected {
|
||||
var label string
|
||||
for _, node := range in.Nodes {
|
||||
if node.ID == nodeID {
|
||||
label = node.Label
|
||||
break
|
||||
}
|
||||
}
|
||||
alert := fmt.Sprintf(
|
||||
"🔌 Node '%s' has no connections. Consider adding edges or removing it.",
|
||||
label,
|
||||
)
|
||||
output.UserAlerts = append(output.UserAlerts, alert)
|
||||
}
|
||||
}
|
||||
|
||||
logger.logf("info", "LLM suggested %d edges with confidence %.2f | %d incompatibilities | %d disconnected",
|
||||
len(validEdges), output.Confidence, len(incompatibilities), len(disconnected))
|
||||
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// buildNodeDescriptions creates readable node descriptions for LLM (including schemas)
|
||||
func buildNodeDescriptions(nodes []db.WorkflowNode) string {
|
||||
var desc string
|
||||
for i, node := range nodes {
|
||||
desc += fmt.Sprintf("%d. [%s] %s (type: %s)\n", i+1, node.ID, node.Label, node.Type)
|
||||
|
||||
// Add input/output schema info
|
||||
if schema, err := getActivitySchema(node.Type); err == nil {
|
||||
if len(schema.Inputs) > 0 {
|
||||
desc += fmt.Sprintf(" INPUTS: %v\n", getInputNames(schema.Inputs))
|
||||
} else {
|
||||
desc += fmt.Sprintf(" INPUTS: none (generator/trigger)\n")
|
||||
}
|
||||
if len(schema.Outputs) > 0 {
|
||||
desc += fmt.Sprintf(" OUTPUTS: %v\n", getOutputNames(schema.Outputs))
|
||||
} else {
|
||||
desc += fmt.Sprintf(" OUTPUTS: none (terminal/sink)\n")
|
||||
}
|
||||
}
|
||||
|
||||
if node.Data != nil {
|
||||
if b, err := json.MarshalIndent(node.Data, " ", " "); err == nil {
|
||||
desc += fmt.Sprintf(" CONFIG: %s\n", string(b))
|
||||
}
|
||||
}
|
||||
}
|
||||
return desc
|
||||
}
|
||||
|
||||
// buildEdgeDescriptions creates readable edge descriptions for LLM
|
||||
func buildEdgeDescriptions(edges []db.WorkflowEdge) string {
|
||||
if len(edges) == 0 {
|
||||
return "None"
|
||||
}
|
||||
var desc string
|
||||
for i, edge := range edges {
|
||||
desc += fmt.Sprintf("%d. %s → %s\n", i+1, edge.Source, edge.Target)
|
||||
}
|
||||
return desc
|
||||
}
|
||||
|
||||
// getReasoningTask returns task description based on preservation mode
|
||||
func getReasoningTask(preserveExisting bool) string {
|
||||
if preserveExisting {
|
||||
return "Keep all existing edges and suggest ONLY NEW edges to improve workflow"
|
||||
}
|
||||
return "Design optimal workflow by suggesting all connections and noting any redundant edges"
|
||||
}
|
||||
+164
-20
@@ -1,27 +1,54 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/statemachine"
|
||||
)
|
||||
|
||||
// AnthropicClient is a thin wrapper around the Anthropic API.
|
||||
type AnthropicClient struct {
|
||||
apiKey string
|
||||
var (
|
||||
// LocalLLMBaseURL is the base URL for the local LLM API (OpenAI-compatible)
|
||||
// Can be overridden via LOCAL_LLM_BASE_URL env var (for Kubernetes internal service)
|
||||
LocalLLMBaseURL string
|
||||
)
|
||||
|
||||
func init() {
|
||||
LocalLLMBaseURL = os.Getenv("LOCAL_LLM_BASE_URL")
|
||||
if LocalLLMBaseURL == "" {
|
||||
// Default: external hostname (for local dev)
|
||||
LocalLLMBaseURL = "https://api.riotpiao.com"
|
||||
}
|
||||
}
|
||||
|
||||
// NewClient creates a new AnthropicClient from the ANTHROPIC_API_KEY env var.
|
||||
func NewClient() (*AnthropicClient, error) {
|
||||
apiKey := os.Getenv("ANTHROPIC_API_KEY")
|
||||
if apiKey == "" {
|
||||
return nil, fmt.Errorf("ANTHROPIC_API_KEY environment variable not set")
|
||||
var (
|
||||
// SupportedModels maps local model names to verify they exist
|
||||
SupportedModels = map[string]bool{
|
||||
"reasoning": true, // Reasoning model for planner/judge
|
||||
"ornith:35b": true, // Ornith 35B for implementer
|
||||
"ornith:13b": true, // Alternative Ornith size
|
||||
"qwen2.5:3b": true, // Qwen alternative
|
||||
}
|
||||
)
|
||||
|
||||
return &AnthropicClient{
|
||||
apiKey: apiKey,
|
||||
// OpenAIClient is a wrapper around the local OpenAI-compatible API.
|
||||
type OpenAIClient struct {
|
||||
baseURL string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
// NewClient creates a new OpenAIClient pointing to the local LLM API.
|
||||
func NewClient() (*OpenAIClient, error) {
|
||||
return &OpenAIClient{
|
||||
baseURL: LocalLLMBaseURL,
|
||||
httpClient: &http.Client{
|
||||
Timeout: 0, // No timeout for streaming
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -30,23 +57,140 @@ type MessageInput struct {
|
||||
Model statemachine.ModelSpec
|
||||
SystemPrompt string
|
||||
Messages []MessageParam
|
||||
AuthToken string // Optional JWT token for authenticated endpoints
|
||||
}
|
||||
|
||||
// MessageParam represents a message parameter (simplified).
|
||||
// MessageParam represents a message parameter.
|
||||
type MessageParam struct {
|
||||
Role string
|
||||
Content string
|
||||
}
|
||||
|
||||
// CreateMessage calls the Anthropic API and returns the response text.
|
||||
// Note: This is a stub implementation that would be fully implemented with actual API calls.
|
||||
func (c *AnthropicClient) CreateMessage(ctx context.Context, in MessageInput) (string, error) {
|
||||
if c.apiKey == "" {
|
||||
return "", fmt.Errorf("API key not set")
|
||||
// openaiRequest is the request body for the OpenAI-compatible API.
|
||||
type openaiRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []openaiMessage `json:"messages"`
|
||||
Stream bool `json:"stream"`
|
||||
Temp float64 `json:"temperature,omitempty"`
|
||||
MaxToken int `json:"max_tokens,omitempty"`
|
||||
}
|
||||
|
||||
type openaiMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// openaiResponse is the response from the OpenAI-compatible API.
|
||||
type openaiResponse struct {
|
||||
Choices []struct {
|
||||
Message struct {
|
||||
Content string `json:"content"`
|
||||
} `json:"message"`
|
||||
} `json:"choices"`
|
||||
Usage struct {
|
||||
PromptTokens int `json:"prompt_tokens"`
|
||||
CompletionTokens int `json:"completion_tokens"`
|
||||
TotalTokens int `json:"total_tokens"`
|
||||
} `json:"usage"`
|
||||
}
|
||||
|
||||
// CreateMessage calls the local OpenAI-compatible API and returns the response text.
|
||||
func (c *OpenAIClient) CreateMessage(ctx context.Context, in MessageInput) (string, error) {
|
||||
// Validate model
|
||||
if !SupportedModels[in.Model.ModelID] {
|
||||
return "", fmt.Errorf("unsupported model: %s (supported: reasoning, ornith:35b)", in.Model.ModelID)
|
||||
}
|
||||
|
||||
// Placeholder implementation
|
||||
// In a real implementation, this would call the Anthropic API
|
||||
// For now, we return a mock response to allow testing
|
||||
return fmt.Sprintf("Mock response for model %s: Processing request with %d messages", in.Model.ModelID, len(in.Messages)), nil
|
||||
// Build request
|
||||
messages := []openaiMessage{
|
||||
{
|
||||
Role: "system",
|
||||
Content: in.SystemPrompt,
|
||||
},
|
||||
}
|
||||
for _, msg := range in.Messages {
|
||||
messages = append(messages, openaiMessage{
|
||||
Role: msg.Role,
|
||||
Content: msg.Content,
|
||||
})
|
||||
}
|
||||
|
||||
req := openaiRequest{
|
||||
Model: in.Model.ModelID,
|
||||
Messages: messages,
|
||||
Stream: false,
|
||||
}
|
||||
|
||||
// Marshal request
|
||||
reqBody, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
// Create HTTP request
|
||||
httpReq, err := http.NewRequestWithContext(ctx, "POST",
|
||||
fmt.Sprintf("%s/v1/chat/completions", c.baseURL),
|
||||
bytes.NewReader(reqBody))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to create HTTP request: %w", err)
|
||||
}
|
||||
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
// Add authentication header if token provided
|
||||
if in.AuthToken != "" {
|
||||
httpReq.Header.Set("Authorization", fmt.Sprintf("Bearer %s", in.AuthToken))
|
||||
}
|
||||
|
||||
// Send request
|
||||
resp, err := c.httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to connect to local LLM API at %s: %w (ensure homelab-frontend is running)", c.baseURL, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
// Read response
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to read response body: %w", err)
|
||||
}
|
||||
|
||||
// Check status
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("local LLM API returned status %d: %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
// Unmarshal response
|
||||
var respObj openaiResponse
|
||||
if err := json.Unmarshal(respBody, &respObj); err != nil {
|
||||
return "", fmt.Errorf("failed to unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
// Extract content
|
||||
if len(respObj.Choices) == 0 {
|
||||
return "", fmt.Errorf("no choices in response from local LLM API")
|
||||
}
|
||||
|
||||
return respObj.Choices[0].Message.Content, nil
|
||||
}
|
||||
|
||||
// HealthCheck verifies the local LLM API is reachable and has the required models.
|
||||
func (c *OpenAIClient) HealthCheck(ctx context.Context) error {
|
||||
httpReq, err := http.NewRequestWithContext(ctx, "GET",
|
||||
fmt.Sprintf("%s/readyz", c.baseURL), nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
resp, err := c.httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return fmt.Errorf("local LLM API at %s is unreachable: %w", c.baseURL, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("local LLM API health check failed with status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/statemachine"
|
||||
)
|
||||
|
||||
func TestNewClient(t *testing.T) {
|
||||
client, err := NewClient()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create client: %v", err)
|
||||
}
|
||||
if client == nil {
|
||||
t.Fatal("client is nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHealthCheck(t *testing.T) {
|
||||
client, err := NewClient()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create client: %v", err)
|
||||
}
|
||||
|
||||
// Skip if local LLM API not available
|
||||
err = client.HealthCheck(context.Background())
|
||||
if err != nil {
|
||||
t.Logf("local LLM API not available (expected in test env): %v", err)
|
||||
t.Skip("local LLM API health check failed - skipping integration test")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSupportedModels(t *testing.T) {
|
||||
tests := []struct {
|
||||
model string
|
||||
expected bool
|
||||
}{
|
||||
{"reasoning", true},
|
||||
{"ornith:35b", true},
|
||||
{"ornith:13b", true},
|
||||
{"qwen2.5:3b", true},
|
||||
{"unsupported-model", false},
|
||||
{"", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.model, func(t *testing.T) {
|
||||
if SupportedModels[tt.model] != tt.expected {
|
||||
t.Errorf("model %q: expected %v, got %v", tt.model, tt.expected, SupportedModels[tt.model])
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateMessageValidation(t *testing.T) {
|
||||
client, _ := NewClient()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
modelID string
|
||||
wantErr bool
|
||||
}{
|
||||
{"valid reasoning", "reasoning", true}, // Will fail to connect, but validates model
|
||||
{"valid ornith", "ornith:35b", true}, // Will fail to connect, but validates model
|
||||
{"invalid model", "invalid-model", false}, // Should fail validation
|
||||
{"empty model", "", false}, // Should fail validation
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
in := MessageInput{
|
||||
Model: statemachine.ModelSpec{
|
||||
ModelID: tt.modelID,
|
||||
},
|
||||
SystemPrompt: "test",
|
||||
Messages: []MessageParam{
|
||||
{Role: "user", Content: "test"},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := client.CreateMessage(context.Background(), in)
|
||||
|
||||
hasErr := err != nil
|
||||
if hasErr != tt.wantErr {
|
||||
if tt.wantErr {
|
||||
t.Logf("expected error for model %q (likely API not reachable): %v", tt.modelID, err)
|
||||
} else if !hasErr {
|
||||
t.Errorf("expected error for invalid model %q, but got none", tt.modelID)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalLLMBaseURL(t *testing.T) {
|
||||
if LocalLLMBaseURL != "https://api.riotpiao.com" {
|
||||
t.Errorf("expected base URL https://api.riotpiao.com, got %s", LocalLLMBaseURL)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/action/llm"
|
||||
"github.com/rockliang/poimen/workflows/statemachine"
|
||||
)
|
||||
|
||||
// LLMInferenceInput is input for LLMInferenceActivity
|
||||
type LLMInferenceInput struct {
|
||||
Model string `json:"model"` // Model ID (reasoning, ornith:35b, etc)
|
||||
SystemPrompt string `json:"system_prompt"` // System instruction
|
||||
UserPrompt string `json:"user_prompt"` // User message
|
||||
Temperature float64 `json:"temperature,omitempty"` // LLM temperature (0-1)
|
||||
MaxTokens int `json:"max_tokens,omitempty"` // Max output tokens
|
||||
AuthToken string `json:"auth_token,omitempty"` // JWT token for authenticated endpoints
|
||||
}
|
||||
|
||||
// LLMInferenceOutput is output from LLMInferenceActivity
|
||||
type LLMInferenceOutput struct {
|
||||
Response string `json:"response"` // LLM response text
|
||||
Model string `json:"model"` // Model used
|
||||
StopReason string `json:"stop_reason"` // How inference stopped (stop_sequence, length, etc)
|
||||
TokensUsed int `json:"tokens_used"` // Total tokens consumed
|
||||
ErrorMessage string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// LLMInferenceActivity calls LLM API with given prompt and returns response
|
||||
func LLMInferenceActivity(ctx context.Context, in LLMInferenceInput) (LLMInferenceOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
|
||||
output := LLMInferenceOutput{
|
||||
Model: in.Model,
|
||||
}
|
||||
|
||||
// Validate input
|
||||
if in.Model == "" {
|
||||
return output, fmt.Errorf("model not specified")
|
||||
}
|
||||
|
||||
if in.UserPrompt == "" {
|
||||
return output, fmt.Errorf("user_prompt not specified")
|
||||
}
|
||||
|
||||
logger.logf("info", "Starting LLM inference with model: %s", in.Model)
|
||||
|
||||
// Create LLM client
|
||||
client, err := llm.NewClient()
|
||||
if err != nil {
|
||||
output.ErrorMessage = err.Error()
|
||||
return output, fmt.Errorf("failed to create LLM client: %w", err)
|
||||
}
|
||||
|
||||
// Call LLM
|
||||
logger.logf("info", "Calling LLM API (model=%s, prompt_len=%d, auth=%v)", in.Model, len(in.UserPrompt), in.AuthToken != "")
|
||||
|
||||
response, err := client.CreateMessage(ctx, llm.MessageInput{
|
||||
Model: statemachine.ModelSpec{
|
||||
ModelID: in.Model,
|
||||
},
|
||||
SystemPrompt: in.SystemPrompt,
|
||||
Messages: []llm.MessageParam{
|
||||
{
|
||||
Role: "user",
|
||||
Content: in.UserPrompt,
|
||||
},
|
||||
},
|
||||
AuthToken: in.AuthToken,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
output.ErrorMessage = err.Error()
|
||||
logger.logf("error", "LLM API call failed: %v", err)
|
||||
return output, fmt.Errorf("LLM inference failed: %w", err)
|
||||
}
|
||||
|
||||
output.Response = response
|
||||
output.StopReason = "stop_sequence"
|
||||
|
||||
logger.logf("info", "LLM inference completed (response_len=%d)", len(response))
|
||||
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// LLMBatchInferenceInput is input for batch inference
|
||||
type LLMBatchInferenceInput struct {
|
||||
Model string `json:"model"`
|
||||
SystemPrompt string `json:"system_prompt"`
|
||||
Prompts []string `json:"prompts"` // List of user prompts
|
||||
Temperature float64 `json:"temperature,omitempty"`
|
||||
AuthToken string `json:"auth_token,omitempty"` // JWT token for authenticated endpoints
|
||||
}
|
||||
|
||||
// LLMBatchInferenceOutput is output from batch inference
|
||||
type LLMBatchInferenceOutput struct {
|
||||
Responses []string `json:"responses"` // LLM responses (parallel to input Prompts)
|
||||
Model string `json:"model"`
|
||||
Errors []string `json:"errors,omitempty"`
|
||||
}
|
||||
|
||||
// LLMBatchInferenceActivity calls LLM multiple times in sequence
|
||||
func LLMBatchInferenceActivity(ctx context.Context, in LLMBatchInferenceInput) (LLMBatchInferenceOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
|
||||
output := LLMBatchInferenceOutput{
|
||||
Model: in.Model,
|
||||
Responses: []string{},
|
||||
Errors: []string{},
|
||||
}
|
||||
|
||||
if in.Model == "" {
|
||||
return output, fmt.Errorf("model not specified")
|
||||
}
|
||||
|
||||
if len(in.Prompts) == 0 {
|
||||
return output, fmt.Errorf("no prompts provided")
|
||||
}
|
||||
|
||||
logger.logf("info", "Starting batch LLM inference (model=%s, count=%d)", in.Model, len(in.Prompts))
|
||||
|
||||
// Create LLM client
|
||||
client, err := llm.NewClient()
|
||||
if err != nil {
|
||||
return output, fmt.Errorf("failed to create LLM client: %w", err)
|
||||
}
|
||||
|
||||
// Process each prompt
|
||||
for i, prompt := range in.Prompts {
|
||||
logger.logf("info", "Processing prompt %d/%d", i+1, len(in.Prompts))
|
||||
|
||||
response, err := client.CreateMessage(ctx, llm.MessageInput{
|
||||
Model: statemachine.ModelSpec{
|
||||
ModelID: in.Model,
|
||||
},
|
||||
SystemPrompt: in.SystemPrompt,
|
||||
Messages: []llm.MessageParam{
|
||||
{
|
||||
Role: "user",
|
||||
Content: prompt,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
output.Errors = append(output.Errors, fmt.Sprintf("prompt %d: %v", i, err))
|
||||
output.Responses = append(output.Responses, "")
|
||||
logger.logf("warn", "Failed to process prompt %d: %v", i, err)
|
||||
} else {
|
||||
output.Responses = append(output.Responses, response)
|
||||
}
|
||||
}
|
||||
|
||||
logger.logf("info", "Batch inference completed (responses=%d, errors=%d)",
|
||||
len(output.Responses), len(output.Errors))
|
||||
|
||||
return output, nil
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
|
||||
"go.temporal.io/sdk/activity"
|
||||
)
|
||||
|
||||
// activityLogger provides logging that works both in and outside Temporal context
|
||||
type activityLogger struct {
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
func newActivityLogger(ctx context.Context) *activityLogger {
|
||||
return &activityLogger{ctx: ctx}
|
||||
}
|
||||
|
||||
func (l *activityLogger) Info(msg string, args ...interface{}) {
|
||||
if activity.IsActivity(l.ctx) {
|
||||
activity.GetLogger(l.ctx).Info(msg, args...)
|
||||
} else {
|
||||
log.Printf("INFO: "+msg+" %v", args)
|
||||
}
|
||||
}
|
||||
|
||||
func (l *activityLogger) Warn(msg string, args ...interface{}) {
|
||||
if activity.IsActivity(l.ctx) {
|
||||
activity.GetLogger(l.ctx).Warn(msg, args...)
|
||||
} else {
|
||||
log.Printf("WARN: "+msg+" %v", args)
|
||||
}
|
||||
}
|
||||
|
||||
func (l *activityLogger) Error(msg string, args ...interface{}) {
|
||||
if activity.IsActivity(l.ctx) {
|
||||
activity.GetLogger(l.ctx).Error(msg, args...)
|
||||
} else {
|
||||
log.Printf("ERROR: "+msg+" %v", args)
|
||||
}
|
||||
}
|
||||
|
||||
func (l *activityLogger) Debug(msg string, args ...interface{}) {
|
||||
if activity.IsActivity(l.ctx) {
|
||||
activity.GetLogger(l.ctx).Debug(msg, args...)
|
||||
} else {
|
||||
log.Printf("DEBUG: "+msg+" %v", args)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,255 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/internal/memory"
|
||||
)
|
||||
|
||||
// RetrieveMemoryInput input for RetrieveMemoryActivity
|
||||
type RetrieveMemoryInput struct {
|
||||
// Query semantic search query
|
||||
Query string `json:"query"`
|
||||
|
||||
// Project memory project (default: "poimen")
|
||||
Project string `json:"project,omitempty"`
|
||||
|
||||
// Scope retrieval scope: "skills", "lessons", "all" (default: "all")
|
||||
Scope string `json:"scope,omitempty"`
|
||||
|
||||
// Limit max results (default: 10)
|
||||
Limit int `json:"limit,omitempty"`
|
||||
|
||||
// LevelFilter filter by level: L1, L2, R (reference)
|
||||
LevelFilter []string `json:"levelFilter,omitempty"`
|
||||
|
||||
// Tool tool context for skill matching
|
||||
Tool string `json:"tool,omitempty"`
|
||||
|
||||
// Task task description for context retrieval
|
||||
Task string `json:"task,omitempty"`
|
||||
}
|
||||
|
||||
// RetrieveMemoryOutput output from RetrieveMemoryActivity
|
||||
type RetrieveMemoryOutput struct {
|
||||
// Skills relevant skills found
|
||||
Skills []MemorySkill `json:"skills"`
|
||||
|
||||
// Lessons relevant lessons/knowledge found
|
||||
Lessons []MemoryLesson `json:"lessons"`
|
||||
|
||||
// References reference documents found
|
||||
References []MemoryReference `json:"references"`
|
||||
|
||||
// TotalResults total results found
|
||||
TotalResults int `json:"totalResults"`
|
||||
|
||||
// Budget token budget info
|
||||
Budget MemoryBudget `json:"budget"`
|
||||
}
|
||||
|
||||
// MemorySkill skill from memory
|
||||
type MemorySkill struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Why string `json:"why,omitempty"`
|
||||
}
|
||||
|
||||
// MemoryLesson lesson from memory
|
||||
type MemoryLesson struct {
|
||||
ID string `json:"id"`
|
||||
Text string `json:"text"`
|
||||
Level string `json:"level"`
|
||||
Score float32 `json:"score"`
|
||||
Breadcrumb string `json:"breadcrumb,omitempty"`
|
||||
}
|
||||
|
||||
// MemoryReference reference document from memory
|
||||
type MemoryReference struct {
|
||||
ID string `json:"id"`
|
||||
Text string `json:"text"`
|
||||
Score float32 `json:"score"`
|
||||
Breadcrumb string `json:"breadcrumb,omitempty"`
|
||||
}
|
||||
|
||||
// MemoryBudget token budget tracking
|
||||
type MemoryBudget struct {
|
||||
Requested int `json:"requested"`
|
||||
Used int `json:"used"`
|
||||
}
|
||||
|
||||
// RetrieveMemoryActivity retrieves relevant knowledge from poimen-memory
|
||||
func RetrieveMemoryActivity(ctx context.Context, in RetrieveMemoryInput) (RetrieveMemoryOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
logger.Info("RetrieveMemoryActivity started", "query", in.Query, "scope", in.Scope)
|
||||
|
||||
output := RetrieveMemoryOutput{
|
||||
Skills: []MemorySkill{},
|
||||
Lessons: []MemoryLesson{},
|
||||
References: []MemoryReference{},
|
||||
}
|
||||
|
||||
// Get memory service URL and token
|
||||
baseURL := os.Getenv("POIMEN_MEMORY_URL")
|
||||
if baseURL == "" {
|
||||
baseURL = "http://poimen-memory.poimen.svc.cluster.local:8080"
|
||||
}
|
||||
|
||||
token := os.Getenv("POIMEN_MEMORY_TOKEN")
|
||||
// Token optional for internal cluster access
|
||||
|
||||
// Set defaults
|
||||
project := in.Project
|
||||
if project == "" {
|
||||
project = "poimen"
|
||||
}
|
||||
|
||||
scope := in.Scope
|
||||
if scope == "" {
|
||||
scope = "all"
|
||||
}
|
||||
|
||||
limit := in.Limit
|
||||
if limit == 0 {
|
||||
limit = 10
|
||||
}
|
||||
|
||||
client := memory.NewClient(baseURL, token)
|
||||
|
||||
// If tool/task provided, use Context API for skill matching
|
||||
if in.Tool != "" || in.Task != "" {
|
||||
contextResp, err := client.Context(ctx, &memory.ContextRequest{
|
||||
Project: project,
|
||||
Tool: in.Tool,
|
||||
Task: in.Task,
|
||||
SignatureSource: in.Query,
|
||||
Scope: "tool_context",
|
||||
Budget: 8192,
|
||||
})
|
||||
if err != nil {
|
||||
logger.Warn("Context retrieval failed, falling back to query", "error", err)
|
||||
} else {
|
||||
// Extract skills
|
||||
for _, skill := range contextResp.Skills {
|
||||
output.Skills = append(output.Skills, MemorySkill{
|
||||
Name: skill.Name,
|
||||
Why: skill.Why,
|
||||
})
|
||||
}
|
||||
|
||||
// Extract lessons
|
||||
for i, lesson := range contextResp.Lessons {
|
||||
output.Lessons = append(output.Lessons, MemoryLesson{
|
||||
ID: fmt.Sprintf("ctx-%d", i),
|
||||
Text: lesson.Text,
|
||||
Level: lesson.Level,
|
||||
Score: lesson.Score,
|
||||
Breadcrumb: "",
|
||||
})
|
||||
}
|
||||
|
||||
output.Budget = MemoryBudget{
|
||||
Requested: contextResp.Budget.Requested,
|
||||
Used: contextResp.Budget.Used,
|
||||
}
|
||||
output.TotalResults = len(output.Skills) + len(output.Lessons)
|
||||
}
|
||||
}
|
||||
|
||||
// Also do semantic query for additional context
|
||||
if scope == "all" || scope == "lessons" || scope == "references" {
|
||||
levelFilter := in.LevelFilter
|
||||
if len(levelFilter) == 0 {
|
||||
levelFilter = []string{"L1", "L2"}
|
||||
}
|
||||
|
||||
queryResp, err := client.Query(ctx, &memory.QueryRequest{
|
||||
Project: project,
|
||||
Query: in.Query,
|
||||
LevelFilter: levelFilter,
|
||||
Limit: limit,
|
||||
Scope: "all",
|
||||
})
|
||||
if err != nil {
|
||||
logger.Warn("Query failed", "error", err)
|
||||
} else {
|
||||
for _, result := range queryResp.Results {
|
||||
if result.Level == "R" {
|
||||
output.References = append(output.References, MemoryReference{
|
||||
ID: result.ID,
|
||||
Text: result.Text,
|
||||
Score: result.Score,
|
||||
Breadcrumb: result.Breadcrumb,
|
||||
})
|
||||
} else {
|
||||
// Avoid duplicates from Context call
|
||||
found := false
|
||||
for _, existing := range output.Lessons {
|
||||
if existing.ID == result.ID {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
output.Lessons = append(output.Lessons, MemoryLesson{
|
||||
ID: result.ID,
|
||||
Text: result.Text,
|
||||
Level: result.Level,
|
||||
Score: result.Score,
|
||||
Breadcrumb: result.Breadcrumb,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
output.TotalResults = len(output.Skills) + len(output.Lessons) + len(output.References)
|
||||
}
|
||||
}
|
||||
|
||||
logger.Info("RetrieveMemoryActivity completed",
|
||||
"skills", len(output.Skills),
|
||||
"lessons", len(output.Lessons),
|
||||
"references", len(output.References))
|
||||
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// FormatMemoryForPrompt formats memory output for LLM prompt injection
|
||||
func FormatMemoryForPrompt(mem RetrieveMemoryOutput) string {
|
||||
if mem.TotalResults == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
var result string
|
||||
|
||||
if len(mem.Skills) > 0 {
|
||||
result += "\n## Relevant Skills\n"
|
||||
for _, skill := range mem.Skills {
|
||||
result += fmt.Sprintf("- **%s**: %s\n", skill.Name, skill.Why)
|
||||
}
|
||||
}
|
||||
|
||||
if len(mem.Lessons) > 0 {
|
||||
result += "\n## Relevant Knowledge\n"
|
||||
for _, lesson := range mem.Lessons {
|
||||
result += fmt.Sprintf("- [%s] %s\n", lesson.Level, truncate(lesson.Text, 200))
|
||||
}
|
||||
}
|
||||
|
||||
if len(mem.References) > 0 {
|
||||
result += "\n## Reference Documents\n"
|
||||
for _, ref := range mem.References {
|
||||
result += fmt.Sprintf("- %s\n", truncate(ref.Text, 200))
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func truncate(s string, maxLen int) string {
|
||||
if len(s) <= maxLen {
|
||||
return s
|
||||
}
|
||||
return s[:maxLen] + "..."
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestRetrieveMemoryActivity_Query(t *testing.T) {
|
||||
// Mock memory service
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/memory/query" {
|
||||
resp := map[string]interface{}{
|
||||
"results": []map[string]interface{}{
|
||||
{
|
||||
"id": "doc-1",
|
||||
"level": "L1",
|
||||
"score": 0.95,
|
||||
"text": "Security scanning best practices: always check for hardcoded secrets",
|
||||
},
|
||||
{
|
||||
"id": "doc-2",
|
||||
"level": "L2",
|
||||
"score": 0.85,
|
||||
"text": "Use gosec for Go security analysis",
|
||||
},
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
return
|
||||
}
|
||||
http.NotFound(w, r)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// Set env for test
|
||||
t.Setenv("POIMEN_MEMORY_URL", server.URL)
|
||||
|
||||
output, err := RetrieveMemoryActivity(context.Background(), RetrieveMemoryInput{
|
||||
Query: "security scanning",
|
||||
Project: "poimen",
|
||||
Scope: "lessons",
|
||||
Limit: 5,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, 2, len(output.Lessons))
|
||||
require.Equal(t, "L1", output.Lessons[0].Level)
|
||||
require.Contains(t, output.Lessons[0].Text, "Security scanning")
|
||||
}
|
||||
|
||||
func TestRetrieveMemoryActivity_Context(t *testing.T) {
|
||||
// Mock memory service with context endpoint
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/memory/context" {
|
||||
resp := map[string]interface{}{
|
||||
"tier": 1,
|
||||
"skills": []map[string]interface{}{
|
||||
{
|
||||
"name": "security-analysis",
|
||||
"why": "User is asking about security scanning",
|
||||
},
|
||||
},
|
||||
"lessons": []map[string]interface{}{
|
||||
{
|
||||
"tier": 1,
|
||||
"level": "L1",
|
||||
"score": 0.9,
|
||||
"text": "Always scan dependencies for vulnerabilities",
|
||||
},
|
||||
},
|
||||
"budget": map[string]interface{}{
|
||||
"requested": 8192,
|
||||
"used": 1024,
|
||||
},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
return
|
||||
}
|
||||
if r.URL.Path == "/memory/query" {
|
||||
resp := map[string]interface{}{"results": []interface{}{}}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
return
|
||||
}
|
||||
http.NotFound(w, r)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
t.Setenv("POIMEN_MEMORY_URL", server.URL)
|
||||
|
||||
output, err := RetrieveMemoryActivity(context.Background(), RetrieveMemoryInput{
|
||||
Query: "security scan repo",
|
||||
Tool: "poimen-router",
|
||||
Task: "generate workflow for security scanning",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, 1, len(output.Skills))
|
||||
require.Equal(t, "security-analysis", output.Skills[0].Name)
|
||||
require.Equal(t, 1, len(output.Lessons))
|
||||
require.Equal(t, 8192, output.Budget.Requested)
|
||||
require.Equal(t, 1024, output.Budget.Used)
|
||||
}
|
||||
|
||||
func TestFormatMemoryForPrompt(t *testing.T) {
|
||||
mem := RetrieveMemoryOutput{
|
||||
Skills: []MemorySkill{
|
||||
{Name: "security-scan", Description: "Run security scanner", Why: "Matches user intent"},
|
||||
},
|
||||
Lessons: []MemoryLesson{
|
||||
{ID: "1", Level: "L1", Text: "Always check dependencies"},
|
||||
},
|
||||
TotalResults: 2,
|
||||
}
|
||||
|
||||
result := FormatMemoryForPrompt(mem)
|
||||
|
||||
require.Contains(t, result, "## Relevant Skills")
|
||||
require.Contains(t, result, "security-scan")
|
||||
require.Contains(t, result, "## Relevant Knowledge")
|
||||
require.Contains(t, result, "Always check dependencies")
|
||||
}
|
||||
|
||||
func TestFormatMemoryForPrompt_Empty(t *testing.T) {
|
||||
mem := RetrieveMemoryOutput{
|
||||
TotalResults: 0,
|
||||
}
|
||||
|
||||
result := FormatMemoryForPrompt(mem)
|
||||
require.Equal(t, "", result)
|
||||
}
|
||||
@@ -0,0 +1,343 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"time"
|
||||
)
|
||||
|
||||
// NotifyStatusInput is input for NotifyStatusActivity
|
||||
type NotifyStatusInput struct {
|
||||
Channel string `json:"channel"` // slack, email, webhook
|
||||
Status string `json:"status"` // success, failure, warning
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// NotifyStatusOutput is output from NotifyStatusActivity
|
||||
type NotifyStatusOutput struct {
|
||||
NotificationID string `json:"notificationId"`
|
||||
Timestamp string `json:"timestamp"`
|
||||
}
|
||||
|
||||
// NotifyStatusActivity sends notifications
|
||||
func NotifyStatusActivity(ctx context.Context, in NotifyStatusInput) (NotifyStatusOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
logger.Info("NotifyStatusActivity started", "channel", in.Channel, "status", in.Status)
|
||||
|
||||
timestamp := time.Now().Format(time.RFC3339)
|
||||
notificationID := fmt.Sprintf("notify-%d", time.Now().UnixNano())
|
||||
|
||||
switch in.Channel {
|
||||
case "slack":
|
||||
if err := sendSlackNotification(ctx, in); err != nil {
|
||||
logger.Warn("Slack notification failed", "error", err)
|
||||
// Don't fail the activity, just log
|
||||
}
|
||||
case "webhook":
|
||||
if err := sendWebhookNotification(ctx, in); err != nil {
|
||||
logger.Warn("Webhook notification failed", "error", err)
|
||||
}
|
||||
case "email":
|
||||
// Email would require SMTP setup - log for now
|
||||
logger.Info("Email notification (logged)", "message", in.Message)
|
||||
default:
|
||||
logger.Info("Notification logged", "channel", in.Channel, "message", in.Message)
|
||||
}
|
||||
|
||||
logger.Info("NotifyStatusActivity completed", "notificationId", notificationID)
|
||||
return NotifyStatusOutput{
|
||||
NotificationID: notificationID,
|
||||
Timestamp: timestamp,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func sendSlackNotification(ctx context.Context, in NotifyStatusInput) error {
|
||||
webhookURL := os.Getenv("SLACK_WEBHOOK_URL")
|
||||
if webhookURL == "" {
|
||||
return fmt.Errorf("SLACK_WEBHOOK_URL not set")
|
||||
}
|
||||
|
||||
// Map status to emoji
|
||||
emoji := "ℹ️"
|
||||
switch in.Status {
|
||||
case "success":
|
||||
emoji = "✅"
|
||||
case "failure":
|
||||
emoji = "❌"
|
||||
case "warning":
|
||||
emoji = "⚠️"
|
||||
}
|
||||
|
||||
payload := map[string]string{
|
||||
"text": fmt.Sprintf("%s *%s*: %s", emoji, in.Status, in.Message),
|
||||
}
|
||||
|
||||
body, _ := json.Marshal(payload)
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", webhookURL, bytes.NewBuffer(body))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("slack returned status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func sendWebhookNotification(ctx context.Context, in NotifyStatusInput) error {
|
||||
webhookURL := os.Getenv("NOTIFICATION_WEBHOOK_URL")
|
||||
if webhookURL == "" {
|
||||
return fmt.Errorf("NOTIFICATION_WEBHOOK_URL not set")
|
||||
}
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"channel": in.Channel,
|
||||
"status": in.Status,
|
||||
"message": in.Message,
|
||||
"timestamp": time.Now().Format(time.RFC3339),
|
||||
}
|
||||
|
||||
body, _ := json.Marshal(payload)
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", webhookURL, bytes.NewBuffer(body))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode >= 400 {
|
||||
return fmt.Errorf("webhook returned status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ArchiveResultsInput is input for ArchiveResultsActivity
|
||||
type ArchiveResultsInput struct {
|
||||
ReportPath string `json:"reportPath"`
|
||||
Destination string `json:"destination"` // s3://bucket/path or local path
|
||||
Metadata map[string]interface{} `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
// ArchiveResultsOutput is output from ArchiveResultsActivity
|
||||
type ArchiveResultsOutput struct {
|
||||
ArchiveURL string `json:"archiveUrl"`
|
||||
ArchiveSize int64 `json:"archiveSize"`
|
||||
}
|
||||
|
||||
// ArchiveResultsActivity archives results to storage
|
||||
func ArchiveResultsActivity(ctx context.Context, in ArchiveResultsInput) (ArchiveResultsOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
logger.Info("ArchiveResultsActivity started", "reportPath", in.ReportPath, "destination", in.Destination)
|
||||
|
||||
// Check if source exists
|
||||
info, err := os.Stat(in.ReportPath)
|
||||
if err != nil {
|
||||
return ArchiveResultsOutput{}, fmt.Errorf("report not found: %s", in.ReportPath)
|
||||
}
|
||||
|
||||
// For now, just copy to local destination or log for cloud
|
||||
var archiveURL string
|
||||
var archiveSize int64
|
||||
|
||||
if len(in.Destination) > 5 && in.Destination[:5] == "s3://" {
|
||||
// Would use AWS SDK - for now just log
|
||||
logger.Info("Would upload to S3", "destination", in.Destination)
|
||||
archiveURL = in.Destination
|
||||
archiveSize = info.Size()
|
||||
} else {
|
||||
// Copy to local destination
|
||||
content, err := os.ReadFile(in.ReportPath)
|
||||
if err != nil {
|
||||
return ArchiveResultsOutput{}, fmt.Errorf("failed to read report: %w", err)
|
||||
}
|
||||
|
||||
destPath := in.Destination
|
||||
if destPath == "" {
|
||||
destPath = fmt.Sprintf("/tmp/archive-%d", time.Now().UnixNano())
|
||||
}
|
||||
|
||||
if err := os.WriteFile(destPath, content, 0644); err != nil {
|
||||
return ArchiveResultsOutput{}, fmt.Errorf("failed to write archive: %w", err)
|
||||
}
|
||||
|
||||
archiveURL = destPath
|
||||
archiveSize = int64(len(content))
|
||||
}
|
||||
|
||||
logger.Info("ArchiveResultsActivity completed", "archiveUrl", archiveURL, "size", archiveSize)
|
||||
return ArchiveResultsOutput{
|
||||
ArchiveURL: archiveURL,
|
||||
ArchiveSize: archiveSize,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// DeploymentPreCheckInput is input for DeploymentPreCheckActivity
|
||||
type DeploymentPreCheckInput struct {
|
||||
Path string `json:"path"`
|
||||
CheckType string `json:"checkType,omitempty"` // lint, test, build, all
|
||||
}
|
||||
|
||||
// DeploymentPreCheckOutput is output from DeploymentPreCheckActivity
|
||||
type DeploymentPreCheckOutput struct {
|
||||
Passed bool `json:"passed"`
|
||||
Failures []string `json:"failures"`
|
||||
Warnings []string `json:"warnings"`
|
||||
}
|
||||
|
||||
// checker defines a deployment check
|
||||
type checker struct {
|
||||
name string
|
||||
types []string // which checkTypes trigger this checker
|
||||
isWarning bool // if true, failure goes to warnings not failures
|
||||
run func(ctx context.Context, path string) (string, error)
|
||||
}
|
||||
|
||||
// shouldRun checks if this checker should run for given checkType
|
||||
func (c *checker) shouldRun(checkType string) bool {
|
||||
for _, t := range c.types {
|
||||
if t == checkType {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// goCheckers returns checkers for Go projects
|
||||
func goCheckers() []*checker {
|
||||
return []*checker{
|
||||
{
|
||||
name: "build",
|
||||
types: []string{"all", "build"},
|
||||
run: func(ctx context.Context, path string) (string, error) {
|
||||
cmd := exec.CommandContext(ctx, "go", "build", "./...")
|
||||
cmd.Dir = path
|
||||
out, err := cmd.CombinedOutput()
|
||||
return string(out), err
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "test",
|
||||
types: []string{"all", "test"},
|
||||
run: func(ctx context.Context, path string) (string, error) {
|
||||
cmd := exec.CommandContext(ctx, "go", "test", "-short", "./...")
|
||||
cmd.Dir = path
|
||||
out, err := cmd.CombinedOutput()
|
||||
return string(out), err
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "lint",
|
||||
types: []string{"all", "lint"},
|
||||
isWarning: true,
|
||||
run: func(ctx context.Context, path string) (string, error) {
|
||||
cmd := exec.CommandContext(ctx, "go", "vet", "./...")
|
||||
cmd.Dir = path
|
||||
out, err := cmd.CombinedOutput()
|
||||
return string(out), err
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// detectProjectCheckers returns checkers based on project type
|
||||
func detectProjectCheckers(path string) []*checker {
|
||||
if _, err := os.Stat(fmt.Sprintf("%s/go.mod", path)); err == nil {
|
||||
return goCheckers()
|
||||
}
|
||||
// Add more project types here (Node, Python, etc.)
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeploymentPreCheckActivity validates deployment readiness
|
||||
func DeploymentPreCheckActivity(ctx context.Context, in DeploymentPreCheckInput) (DeploymentPreCheckOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
logger.Info("DeploymentPreCheckActivity started", "path", in.Path, "checkType", in.CheckType)
|
||||
|
||||
output := DeploymentPreCheckOutput{
|
||||
Passed: true,
|
||||
Failures: []string{},
|
||||
Warnings: []string{},
|
||||
}
|
||||
|
||||
checkType := in.CheckType
|
||||
if checkType == "" {
|
||||
checkType = "all"
|
||||
}
|
||||
|
||||
// Check if path exists
|
||||
if _, err := os.Stat(in.Path); os.IsNotExist(err) {
|
||||
output.Passed = false
|
||||
output.Failures = append(output.Failures, fmt.Sprintf("path does not exist: %s", in.Path))
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// Run applicable checkers
|
||||
for _, c := range detectProjectCheckers(in.Path) {
|
||||
if !c.shouldRun(checkType) {
|
||||
continue
|
||||
}
|
||||
if out, err := c.run(ctx, in.Path); err != nil {
|
||||
msg := fmt.Sprintf("%s failed: %s", c.name, out)
|
||||
if c.isWarning {
|
||||
output.Warnings = append(output.Warnings, msg)
|
||||
} else {
|
||||
output.Passed = false
|
||||
output.Failures = append(output.Failures, msg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
logger.Info("DeploymentPreCheckActivity completed", "passed", output.Passed, "failures", len(output.Failures))
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// ApproveWorkflowInput is input for ApproveWorkflowActivity
|
||||
type ApproveWorkflowInput struct {
|
||||
WorkflowID string `json:"workflowId"`
|
||||
RequiredApprovals int `json:"requiredApprovals,omitempty"`
|
||||
TimeoutMinutes int `json:"timeoutMinutes,omitempty"`
|
||||
}
|
||||
|
||||
// ApproveWorkflowOutput is output from ApproveWorkflowActivity
|
||||
type ApproveWorkflowOutput struct {
|
||||
Approved bool `json:"approved"`
|
||||
Approver string `json:"approver,omitempty"`
|
||||
Timestamp string `json:"timestamp"`
|
||||
}
|
||||
|
||||
// ApproveWorkflowActivity handles approval workflow (auto-approves for now)
|
||||
func ApproveWorkflowActivity(ctx context.Context, in ApproveWorkflowInput) (ApproveWorkflowOutput, error) {
|
||||
logger := newActivityLogger(ctx)
|
||||
logger.Info("ApproveWorkflowActivity started", "workflowId", in.WorkflowID)
|
||||
|
||||
// For now, auto-approve
|
||||
// In production, this would wait for human approval via signal or external system
|
||||
|
||||
timestamp := time.Now().Format(time.RFC3339)
|
||||
|
||||
logger.Info("ApproveWorkflowActivity completed (auto-approved)")
|
||||
return ApproveWorkflowOutput{
|
||||
Approved: true,
|
||||
Approver: "system-auto",
|
||||
Timestamp: timestamp,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,130 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNotifyStatusActivity(t *testing.T) {
|
||||
// Test with default channel (logs only)
|
||||
output, err := NotifyStatusActivity(context.Background(), NotifyStatusInput{
|
||||
Channel: "log",
|
||||
Status: "success",
|
||||
Message: "Test notification",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotEmpty(t, output.NotificationID)
|
||||
require.NotEmpty(t, output.Timestamp)
|
||||
}
|
||||
|
||||
func TestNotifyStatusActivity_AllStatuses(t *testing.T) {
|
||||
statuses := []string{"success", "failure", "warning"}
|
||||
|
||||
for _, status := range statuses {
|
||||
t.Run(status, func(t *testing.T) {
|
||||
output, err := NotifyStatusActivity(context.Background(), NotifyStatusInput{
|
||||
Channel: "log",
|
||||
Status: status,
|
||||
Message: "Test " + status,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, output.NotificationID)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestArchiveResultsActivity(t *testing.T) {
|
||||
// Create temp source file
|
||||
tmpDir, err := os.MkdirTemp("", "archive-test")
|
||||
require.NoError(t, err)
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
srcPath := filepath.Join(tmpDir, "report.txt")
|
||||
require.NoError(t, os.WriteFile(srcPath, []byte("test report content"), 0644))
|
||||
|
||||
// Archive to local destination
|
||||
destPath := filepath.Join(tmpDir, "archive.txt")
|
||||
output, err := ArchiveResultsActivity(context.Background(), ArchiveResultsInput{
|
||||
ReportPath: srcPath,
|
||||
Destination: destPath,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, destPath, output.ArchiveURL)
|
||||
require.Greater(t, output.ArchiveSize, int64(0))
|
||||
|
||||
// Verify file was copied
|
||||
content, err := os.ReadFile(destPath)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "test report content", string(content))
|
||||
}
|
||||
|
||||
func TestArchiveResultsActivity_NotFound(t *testing.T) {
|
||||
_, err := ArchiveResultsActivity(context.Background(), ArchiveResultsInput{
|
||||
ReportPath: "/nonexistent/file.txt",
|
||||
Destination: "/tmp/archive.txt",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "not found")
|
||||
}
|
||||
|
||||
func TestDeploymentPreCheckActivity(t *testing.T) {
|
||||
// Create temp directory with valid Go code
|
||||
tmpDir, err := os.MkdirTemp("", "precheck-test")
|
||||
require.NoError(t, err)
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
// Create go.mod
|
||||
goMod := `module test
|
||||
go 1.21
|
||||
`
|
||||
require.NoError(t, os.WriteFile(filepath.Join(tmpDir, "go.mod"), []byte(goMod), 0644))
|
||||
|
||||
// Create valid Go file
|
||||
goCode := `package main
|
||||
|
||||
func main() {
|
||||
println("hello")
|
||||
}
|
||||
`
|
||||
require.NoError(t, os.WriteFile(filepath.Join(tmpDir, "main.go"), []byte(goCode), 0644))
|
||||
|
||||
// Run pre-check
|
||||
output, err := DeploymentPreCheckActivity(context.Background(), DeploymentPreCheckInput{
|
||||
Path: tmpDir,
|
||||
CheckType: "build",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.True(t, output.Passed)
|
||||
require.Empty(t, output.Failures)
|
||||
}
|
||||
|
||||
func TestDeploymentPreCheckActivity_PathNotExist(t *testing.T) {
|
||||
output, err := DeploymentPreCheckActivity(context.Background(), DeploymentPreCheckInput{
|
||||
Path: "/nonexistent/path",
|
||||
})
|
||||
require.NoError(t, err) // Activity doesn't error, just reports failure
|
||||
|
||||
require.False(t, output.Passed)
|
||||
require.NotEmpty(t, output.Failures)
|
||||
}
|
||||
|
||||
func TestApproveWorkflowActivity(t *testing.T) {
|
||||
output, err := ApproveWorkflowActivity(context.Background(), ApproveWorkflowInput{
|
||||
WorkflowID: "test-workflow-123",
|
||||
RequiredApprovals: 1,
|
||||
TimeoutMinutes: 60,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Auto-approved
|
||||
require.True(t, output.Approved)
|
||||
require.Equal(t, "system-auto", output.Approver)
|
||||
require.NotEmpty(t, output.Timestamp)
|
||||
}
|
||||
@@ -0,0 +1,89 @@
|
||||
package action
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/internal/routing"
|
||||
"go.temporal.io/sdk/activity"
|
||||
)
|
||||
|
||||
// LLMRouterActivity is the Temporal activity that routes user requests to workflows
|
||||
func LLMRouterActivity(ctx context.Context, input routing.LLMRouterInput) (*routing.LLMRouterOutput, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
logger.Info("LLMRouterActivity started", "message", input.Message)
|
||||
|
||||
// Load knowledge base
|
||||
kb, err := routing.LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to load knowledge base: %w", err)
|
||||
}
|
||||
|
||||
// Create router
|
||||
router, err := routing.NewLLMRouterDefault(kb)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create router: %w", err)
|
||||
}
|
||||
|
||||
// Route the request
|
||||
output, err := router.Route(ctx, input)
|
||||
if err != nil {
|
||||
logger.Error("LLMRouterActivity failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if output.IsCron {
|
||||
logger.Info("LLMRouterActivity completed (cron)",
|
||||
"workflowName", output.CronSpec.Name,
|
||||
"schedule", output.CronSpec.Schedule,
|
||||
"stateCount", len(output.CronSpec.States))
|
||||
} else {
|
||||
logger.Info("LLMRouterActivity completed",
|
||||
"workflowName", output.Spec.Name,
|
||||
"stateCount", len(output.Spec.States))
|
||||
}
|
||||
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// ValidateWorkflowSpecActivity validates a workflow spec before execution
|
||||
func ValidateWorkflowSpecActivity(ctx context.Context, spec routing.WorkflowSpec) (*routing.ValidationResult, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
logger.Info("ValidateWorkflowSpecActivity started", "workflowName", spec.Name)
|
||||
|
||||
kb, err := routing.LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to load knowledge base: %w", err)
|
||||
}
|
||||
|
||||
validator := routing.NewValidator(kb)
|
||||
result := validator.ValidateWorkflowSpec(&spec)
|
||||
|
||||
logger.Info("ValidateWorkflowSpecActivity completed",
|
||||
"valid", result.Valid,
|
||||
"errorCount", len(result.Errors))
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// ValidateCronWorkflowSpecActivity validates a cron workflow spec before scheduling
|
||||
func ValidateCronWorkflowSpecActivity(ctx context.Context, spec routing.CronWorkflowSpec) (*routing.ValidationResult, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
logger.Info("ValidateCronWorkflowSpecActivity started",
|
||||
"workflowName", spec.Name,
|
||||
"schedule", spec.Schedule)
|
||||
|
||||
kb, err := routing.LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to load knowledge base: %w", err)
|
||||
}
|
||||
|
||||
validator := routing.NewValidator(kb)
|
||||
result := validator.ValidateCronWorkflowSpec(&spec)
|
||||
|
||||
logger.Info("ValidateCronWorkflowSpecActivity completed",
|
||||
"valid", result.Valid,
|
||||
"errorCount", len(result.Errors))
|
||||
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
# Agent Prompts
|
||||
|
||||
LLM agent prompts for Poimen RoutingWorkflow.
|
||||
|
||||
## Agents
|
||||
|
||||
| Agent | Model | Purpose | Prompt File |
|
||||
|-------|-------|---------|-------------|
|
||||
| **Router** | `reasoning` | Natural language → WorkflowSpec | [router/AGENTS.md](router/AGENTS.md) |
|
||||
|
||||
---
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
User Request: "Scan repo X for security issues"
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────┐
|
||||
│ RetrieveMemoryActivity │
|
||||
│ Query poimen-memory for relevant skills/lessons│
|
||||
└─────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────┐
|
||||
│ Router Agent (LLM) │
|
||||
│ Input: message + memory context + activities │
|
||||
│ Output: WorkflowSpec JSON │
|
||||
└─────────────────────────────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────────────────────────────────┐
|
||||
│ RoutingWorkflow │
|
||||
│ Executes WorkflowSpec as state machine │
|
||||
│ Clone → Scan → Report → Notify │
|
||||
└─────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Memory Integration
|
||||
|
||||
Router receives context from `RetrieveMemoryActivity`:
|
||||
|
||||
```
|
||||
1. User: "scan repo for security"
|
||||
2. RetrieveMemoryActivity queries poimen-memory
|
||||
3. Returns: skills, lessons, references
|
||||
4. Injected into Router prompt
|
||||
5. Router generates smarter WorkflowSpec
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## LLM Endpoint
|
||||
|
||||
Default: `https://api.riotpiao.com/v1/chat/completions`
|
||||
|
||||
Override: `LOCAL_LLM_BASE_URL=http://localhost:11434`
|
||||
@@ -0,0 +1,91 @@
|
||||
# Router Agent
|
||||
|
||||
**Purpose**: Intelligent workflow router that analyzes user requests and generates workflow specs.
|
||||
|
||||
**Model**: `reasoning` (via api.riotpiao.com)
|
||||
|
||||
**When Used**: User submits natural language request → Router generates WorkflowSpec/CronWorkflowSpec
|
||||
|
||||
---
|
||||
|
||||
## System Prompt
|
||||
|
||||
```
|
||||
You are an intelligent workflow router. Your job is to:
|
||||
1. Understand what the user wants to accomplish
|
||||
2. Select the appropriate activities from the available list
|
||||
3. Order them correctly based on dependencies
|
||||
4. Extract any parameters mentioned (URLs, branches, etc)
|
||||
5. Detect if user wants scheduled/recurring execution
|
||||
6. Use any relevant knowledge from memory to inform your decisions
|
||||
|
||||
Rules:
|
||||
- Always include CloneRepoActivity first if any analysis activity is needed
|
||||
- Order activities respecting dependencies
|
||||
- If user mentions "daily", "every hour", "weekly", etc → set isCron=true and cronSchedule
|
||||
- Common cron patterns: "0 2 * * *" (2 AM daily), "0 * * * *" (hourly), "0 0 * * 0" (weekly Sunday)
|
||||
- Extract repo URLs, branch names, severity levels from the message
|
||||
- workflowName should be short and descriptive (kebab-case)
|
||||
- If memory context includes relevant skills or lessons, incorporate that knowledge
|
||||
- Skills from memory may suggest specific activity parameters or ordering
|
||||
|
||||
Output ONLY valid JSON.
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## User Prompt Template
|
||||
|
||||
```
|
||||
User request: {{.Message}}
|
||||
{{if .Context}}
|
||||
Provided context: {{.Context}}
|
||||
{{end}}
|
||||
{{if .MemoryContext}}
|
||||
Relevant skills from memory:
|
||||
{{range .MemoryContext.Skills}}- {{.Name}}: {{.Description}} (reason: {{.Why}})
|
||||
{{end}}
|
||||
Relevant knowledge from memory:
|
||||
{{range .MemoryContext.Lessons}}- [{{.Level}}] {{.Text}}
|
||||
{{end}}
|
||||
{{end}}
|
||||
Available activities:
|
||||
{{range .Activities}}- {{.Name}}: {{.Description}} (category: {{.Category}}, timeout: {{.Timeout}}, flaky: {{.IsFlaky}})
|
||||
{{end}}
|
||||
Analyze the request and output JSON with:
|
||||
- activities: ordered list of activity names to execute
|
||||
- parameters: extracted parameters from request (repo URL, branch, etc)
|
||||
- isCron: true if user wants scheduled/recurring execution
|
||||
- cronSchedule: cron expression if scheduled (e.g., "0 2 * * *" for 2 AM daily)
|
||||
- cronTimezone: timezone (default "UTC")
|
||||
- workflowName: short descriptive name
|
||||
- errorHandling: "retry" (default), "fail-fast", or "continue"
|
||||
|
||||
Output ONLY valid JSON, no explanation.
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Expected Output Format
|
||||
|
||||
```json
|
||||
{
|
||||
"activities": ["CloneRepoActivity", "SecurityScanActivity", "GenerateReportActivity"],
|
||||
"parameters": {
|
||||
"repo": "https://github.com/example/repo",
|
||||
"branch": "main",
|
||||
"severity": "high"
|
||||
},
|
||||
"isCron": false,
|
||||
"cronSchedule": "",
|
||||
"cronTimezone": "UTC",
|
||||
"workflowName": "security-scan-example",
|
||||
"errorHandling": "retry"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Source File
|
||||
|
||||
`internal/routing/llm_router.go`
|
||||
@@ -0,0 +1,125 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"flag"
|
||||
"log"
|
||||
"os"
|
||||
"os/signal"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
"go.temporal.io/sdk/client"
|
||||
"go.temporal.io/sdk/worker"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/action"
|
||||
"github.com/rockliang/poimen/workflows/internal/api"
|
||||
"github.com/rockliang/poimen/workflows/internal/config"
|
||||
"github.com/rockliang/poimen/workflows/pkg/db"
|
||||
"github.com/rockliang/poimen/workflows/statemachine"
|
||||
)
|
||||
|
||||
func main() {
|
||||
var (
|
||||
apiPort = flag.Int("port", 8080, "HTTP API port")
|
||||
verbose = flag.Bool("verbose", false, "verbose logging")
|
||||
)
|
||||
flag.Parse()
|
||||
|
||||
logger := log.New(os.Stdout, "[poimen-server] ", log.LstdFlags|log.Lshortfile)
|
||||
|
||||
// Load configuration
|
||||
cfg, err := config.LoadConfig()
|
||||
if err != nil {
|
||||
logger.Fatalf("failed to load config: %v", err)
|
||||
}
|
||||
|
||||
// Connect to database (memory-db via K8s CNPG)
|
||||
logger.Println("connecting to database...")
|
||||
database, err := db.New(os.Getenv("DATABASE_URL"))
|
||||
if err != nil {
|
||||
logger.Fatalf("failed to connect to database: %v", err)
|
||||
}
|
||||
defer database.Close()
|
||||
logger.Println("✓ Connected to database")
|
||||
|
||||
// Connect to Temporal
|
||||
logger.Printf("connecting to Temporal at %s", cfg.Temporal.HostPort)
|
||||
c, err := client.Dial(client.Options{
|
||||
HostPort: cfg.Temporal.HostPort,
|
||||
Namespace: cfg.Temporal.Namespace,
|
||||
})
|
||||
if err != nil {
|
||||
logger.Fatalf("failed to connect to temporal: %v", err)
|
||||
}
|
||||
defer c.Close()
|
||||
|
||||
logger.Println("✓ Connected to Temporal")
|
||||
|
||||
// Create and start Temporal worker
|
||||
w := worker.New(c, "default", worker.Options{})
|
||||
|
||||
// Register RoutingWorkflow
|
||||
w.RegisterWorkflow(statemachine.RoutingWorkflow)
|
||||
|
||||
// Register activities
|
||||
w.RegisterActivity(action.CloneRepoActivity)
|
||||
w.RegisterActivity(action.AnalyzeCodeActivity)
|
||||
w.RegisterActivity(action.SecurityScanActivity)
|
||||
w.RegisterActivity(action.GenerateReportActivity)
|
||||
w.RegisterActivity(action.DeploymentPreCheckActivity)
|
||||
w.RegisterActivity(action.NotifyStatusActivity)
|
||||
w.RegisterActivity(action.ApproveWorkflowActivity)
|
||||
w.RegisterActivity(action.ArchiveResultsActivity)
|
||||
w.RegisterActivity(action.RetrieveMemoryActivity)
|
||||
w.RegisterActivity(action.AssumeRoleActivity)
|
||||
w.RegisterActivity(action.LLMInferenceActivity)
|
||||
w.RegisterActivity(action.LLMBatchInferenceActivity)
|
||||
w.RegisterActivity(action.CanvasReasonerActivity)
|
||||
|
||||
var wg sync.WaitGroup
|
||||
errChan := make(chan error, 2)
|
||||
|
||||
// Start Temporal worker
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
logger.Println("starting Temporal worker...")
|
||||
if err := w.Run(worker.InterruptCh()); err != nil {
|
||||
errChan <- err
|
||||
}
|
||||
}()
|
||||
|
||||
// Start HTTP API server
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
server := api.NewServer(database, c, logger)
|
||||
logger.Printf("starting API server on port %d", *apiPort)
|
||||
if err := server.Start(*apiPort); err != nil {
|
||||
errChan <- err
|
||||
}
|
||||
}()
|
||||
|
||||
// Wait for interrupt signal
|
||||
sigChan := make(chan os.Signal, 1)
|
||||
signal.Notify(sigChan, syscall.SIGINT, syscall.SIGTERM)
|
||||
|
||||
go func() {
|
||||
sig := <-sigChan
|
||||
logger.Printf("received signal: %v", sig)
|
||||
w.Stop()
|
||||
}()
|
||||
|
||||
// Monitor for errors
|
||||
go func() {
|
||||
err := <-errChan
|
||||
if err != nil {
|
||||
logger.Printf("error: %v", err)
|
||||
w.Stop()
|
||||
}
|
||||
}()
|
||||
|
||||
wg.Wait()
|
||||
logger.Println("✓ Server stopped gracefully")
|
||||
}
|
||||
+160
-6
@@ -2,29 +2,40 @@ package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"flag"
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go.temporal.io/sdk/client"
|
||||
"github.com/rockliang/poimen/workflows/action/llm"
|
||||
"github.com/rockliang/poimen/workflows/internal/config"
|
||||
"github.com/rockliang/poimen/workflows/internal/health"
|
||||
"github.com/rockliang/poimen/workflows/internal/logging"
|
||||
"github.com/rockliang/poimen/workflows/internal/routing"
|
||||
"github.com/rockliang/poimen/workflows/statemachine"
|
||||
)
|
||||
|
||||
func main() {
|
||||
var (
|
||||
// Orchestrator flags
|
||||
repoPath = flag.String("repo", "", "target repo path")
|
||||
remoteURL = flag.String("remote", "", "remote URL")
|
||||
milestone = flag.String("milestone", "T0", "milestone ID")
|
||||
dryRun = flag.Bool("dry-run", false, "disable git push/merge")
|
||||
plannerModel = flag.String("planner-model", "ornith", "planner model ID")
|
||||
judgeModel = flag.String("judge-model", "ornith", "judge model ID")
|
||||
implementerModel = flag.String("implementer-model", "claude-sonnet-5", "implementer model ID")
|
||||
dryRun = flag.Bool("dry-run", false, "disable git push/merge (orchestrator) or skip submit (routing)")
|
||||
plannerModel = flag.String("planner-model", "reasoning", "planner model ID (local-llm)")
|
||||
judgeModel = flag.String("judge-model", "reasoning", "judge model ID (local-llm)")
|
||||
implementerModel = flag.String("implementer-model", "ornith:35b", "implementer model ID (local-llm ornith)")
|
||||
piProvider = flag.String("pi-provider", "local-llm", "pi provider name for skills (local-llm)")
|
||||
healthCheck = flag.Bool("health", false, "check health and exit")
|
||||
|
||||
// Routing workflow flags
|
||||
routeMsg = flag.String("route", "", "natural language message for LLM routing")
|
||||
specFile = flag.String("spec", "", "JSON workflow spec file (direct submit, skip LLM)")
|
||||
cronSpec = flag.Bool("cron", false, "treat spec as CronWorkflowSpec")
|
||||
)
|
||||
flag.Parse()
|
||||
|
||||
@@ -64,9 +75,15 @@ func main() {
|
||||
return
|
||||
}
|
||||
|
||||
// Validate required flags for workflow start
|
||||
// Handle routing workflow mode
|
||||
if *routeMsg != "" || *specFile != "" {
|
||||
runRoutingWorkflow(c, *routeMsg, *specFile, *cronSpec, *dryRun)
|
||||
return
|
||||
}
|
||||
|
||||
// Validate required flags for orchestrator workflow
|
||||
if *repoPath == "" || *remoteURL == "" {
|
||||
logging.Fatal("--repo and --remote flags are required")
|
||||
logging.Fatal("--repo and --remote flags are required (or use --route/--spec for routing workflow)")
|
||||
}
|
||||
|
||||
|
||||
@@ -78,6 +95,7 @@ func main() {
|
||||
Milestone: *milestone,
|
||||
DryRun: *dryRun,
|
||||
MaxCyclesBeforeCAN: 100,
|
||||
PiProvider: *piProvider,
|
||||
Config: statemachine.OrchestratorConfig{
|
||||
SystemPrompt: "You are an expert software developer orchestrating multi-agent work.",
|
||||
Skills: []statemachine.SkillRef{},
|
||||
@@ -109,6 +127,17 @@ func main() {
|
||||
},
|
||||
}
|
||||
|
||||
// Health check: verify local LLM API is reachable
|
||||
logging.Info("checking local LLM API connectivity", logging.String("url", "https://api.riotpiao.com"))
|
||||
llmClient, err := llm.NewClient()
|
||||
if err != nil {
|
||||
logging.Fatal("failed to create LLM client", logging.Err(err))
|
||||
}
|
||||
if err := llmClient.HealthCheck(context.Background()); err != nil {
|
||||
logging.Fatal("local LLM API health check failed", logging.Err(err), logging.String("hint", "ensure homelab-frontend gateway is running and accessible"))
|
||||
}
|
||||
logging.Info("local LLM API is reachable", logging.String("planner-model", *plannerModel), logging.String("judge-model", *judgeModel), logging.String("implementer-model", *implementerModel))
|
||||
|
||||
// Start workflow
|
||||
workflowID := "orch-" + strings.ReplaceAll(*repoPath, "/", "-")
|
||||
logging.Info("starting orchestrator workflow", logging.String("workflowID", workflowID), logging.String("repo", *repoPath))
|
||||
@@ -143,3 +172,128 @@ func main() {
|
||||
fmt.Printf("\nWorkflow completed: %+v\n", result)
|
||||
}
|
||||
}
|
||||
|
||||
// runRoutingWorkflow handles --route and --spec flags
|
||||
func runRoutingWorkflow(c client.Client, routeMsg, specFile string, isCron, dryRun bool) {
|
||||
ctx := context.Background()
|
||||
|
||||
var spec *routing.WorkflowSpec
|
||||
var cronSpec *routing.CronWorkflowSpec
|
||||
|
||||
if specFile != "" {
|
||||
// Load spec from file
|
||||
data, err := os.ReadFile(specFile)
|
||||
if err != nil {
|
||||
logging.Fatal("failed to read spec file", logging.Err(err))
|
||||
}
|
||||
|
||||
validator := routing.NewValidator(nil) // nil KB = skip activity validation
|
||||
if isCron {
|
||||
cronSpec = &routing.CronWorkflowSpec{}
|
||||
if err := json.Unmarshal(data, cronSpec); err != nil {
|
||||
logging.Fatal("failed to parse cron spec", logging.Err(err))
|
||||
}
|
||||
// Validate
|
||||
result := validator.ValidateCronWorkflowSpec(cronSpec)
|
||||
if !result.Valid {
|
||||
logging.Fatal("invalid cron spec", logging.String("errors", result.String()))
|
||||
}
|
||||
} else {
|
||||
spec = &routing.WorkflowSpec{}
|
||||
if err := json.Unmarshal(data, spec); err != nil {
|
||||
logging.Fatal("failed to parse spec", logging.Err(err))
|
||||
}
|
||||
// Validate
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if !result.Valid {
|
||||
logging.Fatal("invalid spec", logging.String("errors", result.String()))
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Use LLM router
|
||||
logging.Info("routing message via LLM", logging.String("message", routeMsg))
|
||||
|
||||
kb, err := routing.LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
logging.Fatal("failed to load knowledge base", logging.Err(err))
|
||||
}
|
||||
|
||||
router, err := routing.NewLLMRouterDefault(kb)
|
||||
if err != nil {
|
||||
logging.Fatal("failed to create LLM router", logging.Err(err))
|
||||
}
|
||||
|
||||
output, err := router.Route(ctx, routing.LLMRouterInput{Message: routeMsg})
|
||||
if err != nil {
|
||||
logging.Fatal("LLM routing failed", logging.Err(err))
|
||||
}
|
||||
|
||||
if output.IsCron {
|
||||
cronSpec = output.CronSpec
|
||||
fmt.Printf("\n=== Generated Cron Spec ===\n")
|
||||
fmt.Printf("Name: %s\n", cronSpec.Name)
|
||||
fmt.Printf("Schedule: %s\n", cronSpec.Schedule)
|
||||
fmt.Printf("States: %d\n", len(cronSpec.States))
|
||||
} else {
|
||||
spec = output.Spec
|
||||
fmt.Printf("\n=== Generated Workflow Spec ===\n")
|
||||
fmt.Printf("Name: %s\n", spec.Name)
|
||||
fmt.Printf("States: %d\n", len(spec.States))
|
||||
}
|
||||
}
|
||||
|
||||
if dryRun {
|
||||
fmt.Printf("\n[dry-run] Spec generated but not submitted\n")
|
||||
if spec != nil {
|
||||
data, _ := json.MarshalIndent(spec, "", " ")
|
||||
fmt.Printf("%s\n", data)
|
||||
} else if cronSpec != nil {
|
||||
data, _ := json.MarshalIndent(cronSpec, "", " ")
|
||||
fmt.Printf("%s\n", data)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Submit to Temporal
|
||||
if cronSpec != nil {
|
||||
// For cron, we'd use Temporal's schedule feature
|
||||
// For now, just start as regular workflow (cron scheduling TBD)
|
||||
spec = &routing.WorkflowSpec{
|
||||
Name: cronSpec.Name,
|
||||
Input: cronSpec.Input,
|
||||
States: cronSpec.States,
|
||||
}
|
||||
logging.Warn("cron scheduling not yet implemented, running as one-shot workflow")
|
||||
}
|
||||
|
||||
workflowID := "routing-" + spec.Name + "-" + time.Now().Format("20060102-150405")
|
||||
input := statemachine.RoutingWorkflowInput{Spec: spec}
|
||||
|
||||
run, err := c.ExecuteWorkflow(ctx, client.StartWorkflowOptions{
|
||||
ID: workflowID,
|
||||
TaskQueue: "poimen-taskqueue",
|
||||
}, statemachine.RoutingWorkflow, input)
|
||||
if err != nil {
|
||||
logging.Fatal("failed to start routing workflow", logging.Err(err))
|
||||
}
|
||||
|
||||
fmt.Printf("\n=== Routing Workflow Started ===\n")
|
||||
fmt.Printf("Workflow ID: %s\n", workflowID)
|
||||
fmt.Printf("Run ID: %s\n", run.GetRunID())
|
||||
|
||||
// Wait briefly for result
|
||||
waitCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var result statemachine.RoutingWorkflowOutput
|
||||
if err := run.Get(waitCtx, &result); err != nil {
|
||||
fmt.Printf("\nWorkflow running (check Temporal UI for status)\n")
|
||||
} else {
|
||||
fmt.Printf("\nWorkflow completed: %s\n", result.Status)
|
||||
if len(result.StepResults) > 0 {
|
||||
for step, res := range result.StepResults {
|
||||
fmt.Printf(" %s: %v\n", step, res)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+24
-1
@@ -51,6 +51,7 @@ func main() {
|
||||
w.RegisterWorkflow(statemachine.OrchestratorWorkflow)
|
||||
w.RegisterWorkflow(statemachine.TaskUnitWorkflow)
|
||||
w.RegisterWorkflow(statemachine.TestWorkflow)
|
||||
w.RegisterWorkflow(statemachine.RoutingWorkflow)
|
||||
|
||||
// Register all activities
|
||||
w.RegisterActivity(action.CloneRepoActivity)
|
||||
@@ -64,10 +65,32 @@ func main() {
|
||||
w.RegisterActivity(action.ImplementerActivity)
|
||||
w.RegisterActivity(action.JudgeActivity)
|
||||
// Integration and lessons activities - register when fully tested
|
||||
// w.RegisterActivity(action.RunIntegrationTestActivity)
|
||||
w.RegisterActivity(action.RunIntegrationTestActivity)
|
||||
// w.RegisterActivity(action.UpdateLessonsActivity)
|
||||
// w.RegisterActivity(action.ReadLessonsActivity)
|
||||
|
||||
// Routing workflow activities
|
||||
w.RegisterActivity(action.LLMRouterActivity)
|
||||
w.RegisterActivity(action.ValidateWorkflowSpecActivity)
|
||||
w.RegisterActivity(action.ValidateCronWorkflowSpecActivity)
|
||||
|
||||
// Analysis activities
|
||||
w.RegisterActivity(action.AnalyzeCodeActivity)
|
||||
w.RegisterActivity(action.SecurityScanActivity)
|
||||
w.RegisterActivity(action.GenerateReportActivity)
|
||||
|
||||
// Notification and utility activities
|
||||
w.RegisterActivity(action.NotifyStatusActivity)
|
||||
w.RegisterActivity(action.ArchiveResultsActivity)
|
||||
w.RegisterActivity(action.DeploymentPreCheckActivity)
|
||||
w.RegisterActivity(action.ApproveWorkflowActivity)
|
||||
|
||||
// Authentication activities
|
||||
w.RegisterActivity(action.AssumeRoleActivity)
|
||||
|
||||
// Memory activities
|
||||
w.RegisterActivity(action.RetrieveMemoryActivity)
|
||||
|
||||
// Initialize health checker
|
||||
healthChecker := health.NewChecker(c)
|
||||
healthHandler := health.NewHandler(healthChecker)
|
||||
|
||||
@@ -0,0 +1,322 @@
|
||||
# Canvas Reasoner: Auto-Inferring Workflow Connections
|
||||
|
||||
## Overview
|
||||
|
||||
The **CanvasReasonerActivity** uses LLM reasoning to automatically suggest connections between workflow activities when users drop new nodes onto the canvas. It analyzes input/output compatibility and detects connection problems.
|
||||
|
||||
## Connection Logic
|
||||
|
||||
### How It Works
|
||||
|
||||
1. **Analyze Node Schemas**
|
||||
- Get each activity's input/output fields from knowledge base
|
||||
- Activities are classified as:
|
||||
- **Generators**: No inputs, has outputs (e.g., API call, trigger)
|
||||
- **Processors**: Has inputs and outputs (e.g., analyze code, security scan)
|
||||
- **Sinks/Terminals**: Has inputs, no outputs (e.g., notification, approval)
|
||||
|
||||
2. **LLM Reasoning**
|
||||
- Pass all nodes + their schemas to reasoning model
|
||||
- Ask LLM to suggest edges based on:
|
||||
- Type compatibility (string→string, object→object)
|
||||
- Logical execution order
|
||||
- Data flow requirements
|
||||
- Common workflow patterns
|
||||
|
||||
3. **Validate Suggestions**
|
||||
- Check all suggested edges exist in node map
|
||||
- Skip self-loops
|
||||
- Remove duplicates
|
||||
|
||||
4. **Compatibility Checking**
|
||||
- For each suggested edge: `source → target`
|
||||
- Verify source produces outputs
|
||||
- Verify target accepts inputs
|
||||
- Check output/input type compatibility
|
||||
- Flag incompatible connections
|
||||
|
||||
5. **Identify Issues**
|
||||
- Collect all incompatible edges
|
||||
- Identify disconnected nodes (no edges in/out)
|
||||
- Generate user alerts for problems
|
||||
|
||||
## Connection Impossibility Detection
|
||||
|
||||
### Why Connections Fail
|
||||
|
||||
1. **Missing Outputs**
|
||||
```
|
||||
NotifyStatusActivity → AnalyzeCodeActivity
|
||||
⚠️ NotifyStatusActivity produces no outputs
|
||||
Reason: Notification is terminal activity (sink)
|
||||
Solution: Add an intermediate processor that has outputs
|
||||
```
|
||||
|
||||
2. **Missing Inputs**
|
||||
```
|
||||
CloneRepoActivity → ApproveWorkflowActivity
|
||||
⚠️ ApproveWorkflowActivity accepts no inputs
|
||||
Reason: Approval is a terminal activity (sink)
|
||||
Solution: ApproveWorkflowActivity only works as final step
|
||||
```
|
||||
|
||||
3. **Type Mismatch**
|
||||
```
|
||||
LLMInferenceActivity (output: string) → DeploymentPreCheckActivity (input: object)
|
||||
⚠️ String output cannot satisfy object input requirement
|
||||
Reason: Incompatible data types
|
||||
Solution: Use LLM transformation node to convert string→object
|
||||
```
|
||||
|
||||
4. **Semantic Incompatibility**
|
||||
```
|
||||
NotifyStatusActivity → CloneRepoActivity
|
||||
⚠️ No logical connection between these activities
|
||||
Reason: Notification cannot be input to clone operation
|
||||
Solution: Ensure data flow makes semantic sense
|
||||
```
|
||||
|
||||
### Incompatibility Data Structure
|
||||
|
||||
```json
|
||||
{
|
||||
"incompatible_edges": [
|
||||
{
|
||||
"source": "node-1",
|
||||
"target": "node-2",
|
||||
"reason": "Source activity produces no outputs",
|
||||
"source_needs": "any output",
|
||||
"target_needs": "path, depth",
|
||||
"suggestion": "Use LLM transformation to map outputs to inputs"
|
||||
}
|
||||
],
|
||||
"disconnected_nodes": ["node-5", "node-8"],
|
||||
"user_alerts": [
|
||||
"⚠️ node-1 → node-2: Source activity produces no outputs. Use LLM transformation to map outputs to inputs",
|
||||
"🔌 Node 'NotifyStatus-1' has no connections. Consider adding edges or removing it."
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
## User Alerts
|
||||
|
||||
### Alert Types
|
||||
|
||||
1. **Incompatibility Warnings** (⚠️)
|
||||
```
|
||||
⚠️ source → target: reason. suggestion.
|
||||
```
|
||||
- Highlighted in red on canvas
|
||||
- Shows in error sidebar
|
||||
- Prevents workflow execution until fixed
|
||||
|
||||
2. **Disconnection Warnings** (🔌)
|
||||
```
|
||||
🔌 Node 'label' has no connections. Consider adding edges or removing it.
|
||||
```
|
||||
- Highlighted in yellow
|
||||
- Nodes with no input/output edges
|
||||
- May be valid (first step, last step) or indicate design error
|
||||
|
||||
3. **Type Mismatch Info** (ℹ️)
|
||||
```
|
||||
ℹ️ To connect source → target, use transformer to map: {source_outputs} → {target_inputs}
|
||||
```
|
||||
- Suggestion to use intermediate LLM node
|
||||
- Provides mapping information
|
||||
|
||||
## Frontend Integration
|
||||
|
||||
### Canvas UI Feedback
|
||||
|
||||
When CanvasReasonerActivity returns incompatibilities:
|
||||
|
||||
1. **Visual Markers**
|
||||
- Incompatible suggested edges: ❌ red dashed line (don't auto-add)
|
||||
- Disconnected nodes: ⚠️ yellow border
|
||||
|
||||
2. **Sidebar Alerts**
|
||||
```
|
||||
🚨 Connection Issues (3)
|
||||
|
||||
⚠️ CloneRepo → ApproveWorkflow
|
||||
Reason: ApproveWorkflow is terminal (no outputs)
|
||||
Suggestion: Place ApproveWorkflow at end of workflow
|
||||
|
||||
⚠️ LLMInference → DeploymentPreCheck
|
||||
Reason: Type mismatch (string ≠ object)
|
||||
Suggestion: Add LLM transformation node
|
||||
|
||||
🔌 SecurityScan-1 has no incoming edges
|
||||
Suggestion: Connect CloneRepo → SecurityScan
|
||||
```
|
||||
|
||||
3. **User Actions**
|
||||
- ✅ Accept suggestions (green edges)
|
||||
- ❌ Reject incompatible edges
|
||||
- 🔧 Add transformer nodes
|
||||
- 🗑️ Remove disconnected nodes
|
||||
|
||||
### API Response Example
|
||||
|
||||
```json
|
||||
{
|
||||
"suggested_edges": [
|
||||
{"source": "clone-1", "target": "analyze-1"},
|
||||
{"source": "analyze-1", "target": "security-1"},
|
||||
{"source": "security-1", "target": "report-1"}
|
||||
],
|
||||
"reasoning": "Standard code review workflow: clone → analyze → scan → report",
|
||||
"confidence": 0.92,
|
||||
"incompatible_edges": [
|
||||
{
|
||||
"source": "report-1",
|
||||
"target": "approve-1",
|
||||
"reason": "ReportGenerator has no outputs (terminal activity)",
|
||||
"suggestion": "ApproveWorkflow can only be a final step"
|
||||
}
|
||||
],
|
||||
"disconnected_nodes": [],
|
||||
"user_alerts": [
|
||||
"⚠️ report-1 → approve-1: ReportGenerator has no outputs (terminal activity). ApproveWorkflow can only be a final step"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
## Knowledge Base Schema
|
||||
|
||||
Each activity in `activity_knowledge_base.json` defines:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "CloneRepoActivity",
|
||||
"inputs": {
|
||||
"repo": {"type": "string", "required": true},
|
||||
"branch": {"type": "string", "required": false}
|
||||
},
|
||||
"outputs": {
|
||||
"path": {"type": "string"},
|
||||
"commit": {"type": "string"}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Classification Rules
|
||||
|
||||
- **Generator** (0 inputs): trigger, API call, schedule
|
||||
- **Processor** (1+ inputs, 1+ outputs): analysis, transformation, scan
|
||||
- **Sink** (1+ inputs, 0 outputs): notification, approval, archive
|
||||
- **Bypass** (0 inputs, 0 outputs): rare - usually error
|
||||
|
||||
## Common Patterns
|
||||
|
||||
### ✅ Valid Chains
|
||||
|
||||
```
|
||||
CloneRepo → Analyze → SecurityScan → Report
|
||||
(generator) → (processor) → (processor) → (sink)
|
||||
```
|
||||
|
||||
```
|
||||
Trigger → LLMInference → Decision → (Branch: Notify OR Approve)
|
||||
(gen) → (processor) → (processor) → (sink)
|
||||
```
|
||||
|
||||
### ❌ Invalid Chains
|
||||
|
||||
```
|
||||
Notify → CloneRepo ❌
|
||||
(sink) → (generator) - backward flow
|
||||
|
||||
CloneRepo → CloneRepo → Analyze ❌
|
||||
self-loop - no benefit
|
||||
|
||||
Analyze → Approve → Notify ❌
|
||||
Approve is terminal (sink), can't output to Notify
|
||||
```
|
||||
|
||||
## Edge Cases
|
||||
|
||||
### Multiple Outputs → Single Input
|
||||
```
|
||||
SecurityScan → Report
|
||||
SecurityScan outputs: [issues, metrics, severity]
|
||||
Report inputs: [report_data]
|
||||
|
||||
LLM must infer: bundle all outputs into single report_data object
|
||||
Confidence: 0.7 (requires transformation)
|
||||
```
|
||||
|
||||
### Terminal Activities
|
||||
- **ApproveWorkflowActivity**: Must be last (blocks workflow)
|
||||
- **NotifyStatusActivity**: Can be mid-workflow (async notify)
|
||||
- **ArchiveResultsActivity**: Should be last (persistence)
|
||||
|
||||
### Data Transformation
|
||||
When source outputs don't match target inputs:
|
||||
|
||||
```python
|
||||
# User can insert transformer node:
|
||||
LLMInference → [LLMTransformer] → DeploymentPreCheck
|
||||
|
||||
# Transformer:
|
||||
# - Input: LLMInference.output (string)
|
||||
# - Output: DeploymentPreCheck.requirements (object)
|
||||
# - Action: Call LLM to convert format
|
||||
```
|
||||
|
||||
## Testing Incompatibility Detection
|
||||
|
||||
### Test Case 1: Terminal Activity as Source
|
||||
```go
|
||||
source := db.WorkflowNode{ID: "n1", Type: "notify-status", Label: "Notify"}
|
||||
target := db.WorkflowNode{ID: "n2", Type: "clone-repo", Label: "Clone"}
|
||||
|
||||
warnings := CheckConnectionCompatibility(source, target)
|
||||
// Should warn: NotifyStatusActivity produces no outputs
|
||||
```
|
||||
|
||||
### Test Case 2: Type Mismatch
|
||||
```go
|
||||
source := db.WorkflowNode{ID: "n1", Type: "llm-inference", ...}
|
||||
target := db.WorkflowNode{ID: "n2", Type: "deployment-check", ...}
|
||||
|
||||
warnings := CheckConnectionCompatibility(source, target)
|
||||
// Should warn: string output ≠ object input
|
||||
```
|
||||
|
||||
### Test Case 3: Disconnected Node
|
||||
```go
|
||||
nodes := []db.WorkflowNode{n1, n2, n3}
|
||||
edges := []db.WorkflowEdge{{Source: "n1", Target: "n2"}}
|
||||
|
||||
disconnected := IdentifyDisconnectedNodes(nodes, edges)
|
||||
// Should return ["n3"]
|
||||
```
|
||||
|
||||
## Future Enhancements
|
||||
|
||||
1. **Automatic Transformer Insertion**
|
||||
- Detect incompatibilities
|
||||
- Auto-suggest LLM transformer nodes
|
||||
- Chain transformers if needed
|
||||
|
||||
2. **Confidence Scoring**
|
||||
- Increase when types match perfectly
|
||||
- Decrease for semantic mismatches
|
||||
- Factor in activity dependencies
|
||||
|
||||
3. **Learning from History**
|
||||
- Track successful workflows
|
||||
- Remember user edits to suggestions
|
||||
- Improve LLM prompts over time
|
||||
|
||||
4. **Multi-Path Analysis**
|
||||
- Suggest multiple connection topologies
|
||||
- Show cost/efficiency of each
|
||||
- Rank by execution time/cost
|
||||
|
||||
5. **Dry-Run Validation**
|
||||
- Execute suggested workflow in simulation
|
||||
- Catch runtime errors early
|
||||
- Show data flow through each node
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,712 @@
|
||||
# Workflows + Graph RAG Integration
|
||||
|
||||
Align Poimen Workflows with Poimen Memory System API for versioned relations, semantic queries, and intelligent canvas reasoning.
|
||||
|
||||
**Key Features:**
|
||||
- Versioned workflow relations following `/memory/entities/{id}/versions` pattern
|
||||
- Semantic edge queries via `/workflows/{id}/query/semantic/edges`
|
||||
- Relation wording as Facts (matching Memory System edges schema)
|
||||
- Point-in-time canvas reconstruction via `?as_of=timestamp`
|
||||
- Ranking profiles for relation importance
|
||||
- Automatic Graph RAG indexing
|
||||
|
||||
---
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
┌─────────────────────────────────┐
|
||||
│ Workflow Canvas (React Flow) │
|
||||
│ - Nodes (activities) │
|
||||
│ - Edges (connections) │
|
||||
└──────────────┬──────────────────┘
|
||||
│ PUT /workflows/{id}
|
||||
▼
|
||||
┌─────────────────────────────────┐
|
||||
│ CanvasReasonerActivity │
|
||||
│ - Suggest edges │
|
||||
│ - Validate compatibility │
|
||||
│ - Generate change reasoning │
|
||||
└──────────────┬──────────────────┘
|
||||
│ suggested_edges + reasoning
|
||||
▼
|
||||
┌─────────────────────────────────┐
|
||||
│ Workflows API Server │
|
||||
│ - Update canvas in DB │
|
||||
│ - Create version entry │
|
||||
│ - Store change metadata │
|
||||
└──────────────┬──────────────────┘
|
||||
│ canvas_version, relations
|
||||
▼
|
||||
┌─────────────────────────────────┐
|
||||
│ Graph RAG Backend │
|
||||
│ - Store versioned relations │
|
||||
│ - Index relation wording │
|
||||
│ - Enable semantic queries │
|
||||
└─────────────────────────────────┘
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Data Model
|
||||
|
||||
### Workflow Canvas Version
|
||||
|
||||
```sql
|
||||
-- In memory.workflows_versions (new table)
|
||||
CREATE TABLE workflow_versions (
|
||||
id UUID PRIMARY KEY,
|
||||
workflow_id UUID NOT NULL REFERENCES workflows(id),
|
||||
customer_id UUID NOT NULL,
|
||||
version INT NOT NULL,
|
||||
canvas JSONB NOT NULL, -- {nodes, edges}
|
||||
changed_by UUID,
|
||||
change_reason TEXT,
|
||||
change_type VARCHAR(50), -- 'manual', 'auto_reasoned', 'import'
|
||||
reasoner_confidence FLOAT,
|
||||
reasoning_metadata JSONB, -- LLM reasoning output
|
||||
created_at TIMESTAMP DEFAULT NOW(),
|
||||
|
||||
UNIQUE(workflow_id, version),
|
||||
FOREIGN KEY(workflow_id, customer_id)
|
||||
REFERENCES workflows(id, customer_id)
|
||||
);
|
||||
|
||||
-- In memory.workflow_relations (replaces simple edges)
|
||||
CREATE TABLE workflow_relations (
|
||||
id UUID PRIMARY KEY,
|
||||
workflow_id UUID NOT NULL,
|
||||
version INT NOT NULL,
|
||||
source_node_id VARCHAR(255),
|
||||
target_node_id VARCHAR(255),
|
||||
relation_type VARCHAR(100), -- 'data-flow', 'dependency', 'conditional', 'parallel'
|
||||
relation_label TEXT, -- Human-readable: "SecurityScan outputs issues → Report inputs requirements"
|
||||
relation_wording JSONB, -- {verb, object, context}
|
||||
metadata JSONB, -- {source_output_type, target_input_type, compatibility_score}
|
||||
created_at TIMESTAMP,
|
||||
|
||||
FOREIGN KEY(workflow_id, version)
|
||||
REFERENCES workflow_versions(id, version),
|
||||
INDEX (workflow_id, version)
|
||||
);
|
||||
|
||||
-- In memory.relation_changes (for Graph RAG indexing)
|
||||
CREATE TABLE relation_changes (
|
||||
id UUID PRIMARY KEY,
|
||||
workflow_id UUID,
|
||||
version_from INT,
|
||||
version_to INT,
|
||||
change_type VARCHAR(50), -- 'added', 'removed', 'modified'
|
||||
relation_id UUID REFERENCES workflow_relations(id),
|
||||
source_node_id VARCHAR(255),
|
||||
target_node_id VARCHAR(255),
|
||||
old_wording JSONB,
|
||||
new_wording JSONB,
|
||||
change_reason TEXT,
|
||||
change_timestamp TIMESTAMP,
|
||||
reasoner_confidence FLOAT,
|
||||
|
||||
INDEX (workflow_id, version_to)
|
||||
);
|
||||
```
|
||||
|
||||
### Relation Wording Schema
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "edge-1",
|
||||
"source": "clone-repo-1",
|
||||
"target": "analyze-code-1",
|
||||
"relation_type": "data-flow",
|
||||
|
||||
"relation_label": "CloneRepo outputs path → AnalyzeCode requires path",
|
||||
|
||||
"relation_wording": {
|
||||
"verb": "outputs",
|
||||
"source_output": "path (string): Local filesystem path where repo was cloned",
|
||||
"target_input": "path (string, required): Local filesystem path to analyze",
|
||||
"connection_type": "direct-map",
|
||||
"confidence": 0.98,
|
||||
"notes": "Perfect type match between CloneRepo.path and AnalyzeCode.path"
|
||||
},
|
||||
|
||||
"metadata": {
|
||||
"source_activity": "CloneRepoActivity",
|
||||
"target_activity": "AnalyzeCodeActivity",
|
||||
"output_type": "string",
|
||||
"input_type": "string",
|
||||
"compatibility_score": 0.98,
|
||||
"requires_transformation": false,
|
||||
"semantic_match": "File path passes directly"
|
||||
},
|
||||
|
||||
"change_history": [
|
||||
{
|
||||
"version": 2,
|
||||
"action": "added",
|
||||
"reason": "LLM reasoner suggested data-flow connection",
|
||||
"confidence": 0.98,
|
||||
"timestamp": "2025-09-05T10:00:00Z"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Unified API Endpoints
|
||||
|
||||
All workflow queries follow the unified endpoint pattern from Memory System.
|
||||
|
||||
### 1. Unified Workflow Query
|
||||
|
||||
**Endpoint:** `POST /workflows/{id}/query`
|
||||
|
||||
This is the primary endpoint for all workflow canvas queries (replaces separate search endpoints).
|
||||
|
||||
**Request:**
|
||||
```json
|
||||
{
|
||||
"query": "how does code analysis flow into security scanning",
|
||||
"search_type": "edges|entities|all",
|
||||
"version": 3,
|
||||
"relation_type": "data-flow",
|
||||
"confidence_floor": 0.7,
|
||||
"top_k": 10,
|
||||
"find_paths": true,
|
||||
"target_node_id": "security-scan-1",
|
||||
"max_path_depth": 3,
|
||||
"ranking_profile": "default",
|
||||
"include_reasoning": true
|
||||
}
|
||||
```
|
||||
|
||||
**Response (200 OK):**
|
||||
```json
|
||||
{
|
||||
"workflow_id": "workflow-1",
|
||||
"query": "how does code analysis flow into security scanning",
|
||||
"search_type": "edges",
|
||||
"version": 3,
|
||||
"execution_time_ms": 145,
|
||||
"results": [
|
||||
{
|
||||
"id": "edge_analyze_scan",
|
||||
"source_node_id": "analyze-code-1",
|
||||
"source_name": "AnalyzeCodeActivity",
|
||||
"target_node_id": "security-scan-1",
|
||||
"target_name": "SecurityScanActivity",
|
||||
"relation_type": "data-flow",
|
||||
"relation_label": "AnalyzeCode outputs metrics → SecurityScan requires code structure",
|
||||
"relation_wording": {
|
||||
"verb": "provides-input-for",
|
||||
"source_output": "metrics (object): Code quality and structural metrics",
|
||||
"target_input": "path (string): Directory to scan",
|
||||
"connection_type": "requires-transformer",
|
||||
"confidence": 0.85,
|
||||
"semantic_match": "Analysis metrics can guide security scan prioritization"
|
||||
},
|
||||
"similarity_score": 0.92,
|
||||
"confidence": 0.85,
|
||||
"created_at": "2025-09-05T10:00:00Z",
|
||||
"metadata": {
|
||||
"source": "canvas://workflow-1:v3",
|
||||
"tags": ["code-review", "security"]
|
||||
}
|
||||
}
|
||||
],
|
||||
"paths": [
|
||||
{
|
||||
"source_id": "analyze-code-1",
|
||||
"target_id": "security-scan-1",
|
||||
"path_count": 1,
|
||||
"shortest_distance": 1,
|
||||
"paths_found": [
|
||||
{
|
||||
"node_ids": ["analyze-code-1", "security-scan-1"],
|
||||
"relation_types": ["data-flow"],
|
||||
"distance": 1,
|
||||
"total_confidence": 0.85
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"total_count": 1,
|
||||
"has_more": false,
|
||||
"ranking_profile": "default"
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 2. Update Workflow Canvas (with Versioning)
|
||||
|
||||
**Endpoint:** `PUT /workflows/{id}`
|
||||
|
||||
**Request:**
|
||||
```json
|
||||
{
|
||||
"nodes": [...],
|
||||
"edges": [...],
|
||||
"auto_reason": true,
|
||||
"change_reason": "User connected CloneRepo to AnalyzeCode",
|
||||
"user_id": "uuid"
|
||||
}
|
||||
```
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"id": "workflow-1",
|
||||
"version": 3,
|
||||
"canvas": {
|
||||
"nodes": [...],
|
||||
"edges": [...]
|
||||
},
|
||||
"version_info": {
|
||||
"version_number": 3,
|
||||
"created_at": "2025-09-05T10:05:00Z",
|
||||
"created_by": "user-uuid",
|
||||
"change_reason": "User connected CloneRepo to AnalyzeCode",
|
||||
"change_type": "manual"
|
||||
},
|
||||
"relation_updates": {
|
||||
"added": [
|
||||
{
|
||||
"source": "clone-repo-1",
|
||||
"target": "analyze-code-1",
|
||||
"relation_wording": {
|
||||
"verb": "outputs",
|
||||
"source_output": "path: Local filesystem path where repo was cloned",
|
||||
"target_input": "path (required): Local filesystem path to analyze",
|
||||
"confidence": 0.98
|
||||
}
|
||||
}
|
||||
],
|
||||
"removed": [],
|
||||
"modified": []
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 2. Get Relation Version History (Memory System Pattern)
|
||||
|
||||
**Endpoint:** `GET /workflows/{id}/relations/{edge_id}/versions`
|
||||
|
||||
Follows `/memory/entities/{id}/versions` pattern from Memory System.
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"edge_id": "edge_1",
|
||||
"workflow_id": "workflow-1",
|
||||
"source": "clone-repo-1",
|
||||
"target": "analyze-code-1",
|
||||
"versions": [
|
||||
{
|
||||
"version_num": 1,
|
||||
"operation": "CREATE",
|
||||
"snapshot": {
|
||||
"relation_type": "data-flow",
|
||||
"relation_label": "CloneRepo → AnalyzeCode",
|
||||
"relation_wording": {
|
||||
"verb": "connects-to",
|
||||
"confidence": 0.75
|
||||
}
|
||||
},
|
||||
"changed_at": "2025-09-04T12:00:00Z",
|
||||
"changed_by": "system",
|
||||
"fields_changed": ["relation_type", "relation_wording"]
|
||||
},
|
||||
{
|
||||
"version_num": 2,
|
||||
"operation": "UPDATE",
|
||||
"snapshot": {
|
||||
"relation_type": "data-flow",
|
||||
"relation_label": "CloneRepo outputs path → AnalyzeCode requires path",
|
||||
"relation_wording": {
|
||||
"verb": "outputs",
|
||||
"source_output": "path (string)",
|
||||
"target_input": "path (string, required)",
|
||||
"confidence": 0.98
|
||||
}
|
||||
},
|
||||
"changed_at": "2025-09-05T10:00:00Z",
|
||||
"changed_by": "reasoner-activity",
|
||||
"fields_changed": ["relation_wording", "relation_label"]
|
||||
}
|
||||
],
|
||||
"total_versions": 2,
|
||||
"current_version": 2
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 3. Edge Diff (Versioning Pattern)
|
||||
|
||||
**Endpoint:** `POST /workflows/{id}/relations/diff`
|
||||
|
||||
Follows `/memory/entities/diff` pattern.
|
||||
|
||||
**Request:**
|
||||
```json
|
||||
{
|
||||
"edge_id": "edge_1",
|
||||
"from_version": 1,
|
||||
"to_version": 2
|
||||
}
|
||||
```
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"edge_id": "edge_1",
|
||||
"from_version": 1,
|
||||
"to_version": 2,
|
||||
"source": "clone-repo-1",
|
||||
"target": "analyze-code-1",
|
||||
"diff": {
|
||||
"added_fields": {},
|
||||
"removed_fields": {},
|
||||
"modified_fields": {
|
||||
"relation_wording": {
|
||||
"old": {
|
||||
"verb": "connects-to",
|
||||
"confidence": 0.75
|
||||
},
|
||||
"new": {
|
||||
"verb": "outputs",
|
||||
"source_output": "path (string)",
|
||||
"target_input": "path (string, required)",
|
||||
"confidence": 0.98
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"change_timeline": [
|
||||
{
|
||||
"version": 1,
|
||||
"confidence": 0.75,
|
||||
"changed_at": "2025-09-04T12:00:00Z"
|
||||
},
|
||||
{
|
||||
"version": 2,
|
||||
"confidence": 0.98,
|
||||
"changed_at": "2025-09-05T10:00:00Z"
|
||||
}
|
||||
],
|
||||
"editors_involved": ["system", "reasoner-activity"]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 4. Point-in-Time Canvas (Versioning Pattern)
|
||||
|
||||
**Endpoint:** `GET /workflows/{id}?at_version={v}` or `?as_of=2025-09-05T10:00:00Z`
|
||||
|
||||
Follows `/memory/entities/at` pattern.
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"workflow_id": "workflow-1",
|
||||
"version": 2,
|
||||
"as_of_timestamp": "2025-09-05T10:00:00Z",
|
||||
"canvas": {
|
||||
"nodes": [...],
|
||||
"edges": [...]
|
||||
},
|
||||
"relations": [
|
||||
{
|
||||
"id": "edge_1",
|
||||
"source": "clone-repo-1",
|
||||
"target": "analyze-code-1",
|
||||
"relation_type": "data-flow",
|
||||
"relation_label": "CloneRepo outputs path → AnalyzeCode requires path",
|
||||
"relation_wording": {
|
||||
"verb": "outputs",
|
||||
"source_output": "path (string)",
|
||||
"target_input": "path (string, required)",
|
||||
"confidence": 0.98
|
||||
},
|
||||
"version": 2
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"version_number": 2,
|
||||
"created_at": "2025-09-05T10:00:00Z",
|
||||
"changed_by": "reasoner-activity",
|
||||
"change_reason": "LLM reasoner refined relation wording"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 5. Bulk Import Canvas (with Relations)
|
||||
|
||||
**Endpoint:** `POST /workflows/{id}/import`
|
||||
|
||||
**Request:**
|
||||
```json
|
||||
{
|
||||
"canvas": {
|
||||
"nodes": [...],
|
||||
"edges": [...]
|
||||
},
|
||||
"relations": [
|
||||
{
|
||||
"source": "n1",
|
||||
"target": "n2",
|
||||
"relation_type": "data-flow",
|
||||
"relation_wording": {
|
||||
"verb": "outputs",
|
||||
"source_output": "result (string)",
|
||||
"target_input": "input (string, required)"
|
||||
}
|
||||
}
|
||||
],
|
||||
"change_reason": "Imported from external workflow system"
|
||||
}
|
||||
```
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"workflow_id": "workflow-1",
|
||||
"version": 4,
|
||||
"canvas": {...},
|
||||
"relations": {...},
|
||||
"import_metadata": {
|
||||
"imported_nodes": 5,
|
||||
"imported_edges": 4,
|
||||
"validation_status": "success",
|
||||
"indexing_status": "queued_for_graph_rag"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Integration with CanvasReasonerActivity
|
||||
|
||||
### Flow
|
||||
|
||||
```
|
||||
User edits canvas
|
||||
↓
|
||||
PUT /workflows/{id} with auto_reason=true
|
||||
↓
|
||||
Backend calls CanvasReasonerActivity
|
||||
↓
|
||||
LLM suggests edges + reasoning
|
||||
↓
|
||||
Generate relation_wording from suggestions
|
||||
↓
|
||||
Create workflow_version entry
|
||||
↓
|
||||
Create workflow_relations entries
|
||||
↓
|
||||
Index relations in Graph RAG
|
||||
↓
|
||||
Response includes suggested_edges + relation_wording
|
||||
```
|
||||
|
||||
### Response Structure
|
||||
|
||||
```json
|
||||
{
|
||||
"version": 3,
|
||||
"suggested_edges": [
|
||||
{
|
||||
"source": "clone-1",
|
||||
"target": "analyze-1"
|
||||
}
|
||||
],
|
||||
"reasoning": "Standard code review workflow",
|
||||
"confidence": 0.92,
|
||||
"relation_wordings": [
|
||||
{
|
||||
"source": "clone-1",
|
||||
"target": "analyze-1",
|
||||
"relation_wording": {
|
||||
"verb": "outputs",
|
||||
"source_output": "path (string): Cloned repository path",
|
||||
"target_input": "path (string, required): Directory to analyze",
|
||||
"connection_type": "direct-map",
|
||||
"confidence": 0.98,
|
||||
"semantic_description": "Repository path flows from clone operation to analysis"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Graph RAG Indexing
|
||||
|
||||
### Relations Indexed
|
||||
|
||||
Each workflow relation creates Graph RAG entities:
|
||||
|
||||
```
|
||||
Node: {
|
||||
id: "clone-repo-1",
|
||||
type: "activity",
|
||||
name: "CloneRepoActivity",
|
||||
workflow_id: "workflow-1",
|
||||
version: 3
|
||||
}
|
||||
|
||||
Node: {
|
||||
id: "analyze-code-1",
|
||||
type: "activity",
|
||||
...
|
||||
}
|
||||
|
||||
Edge: {
|
||||
id: "rel-1",
|
||||
source: "clone-repo-1",
|
||||
target: "analyze-code-1",
|
||||
type: "data-flow",
|
||||
label: "outputs path",
|
||||
wording: {...},
|
||||
version: 3,
|
||||
created_at: "2025-09-05T10:00:00Z"
|
||||
}
|
||||
```
|
||||
|
||||
### Semantic Queries
|
||||
|
||||
Users can query like:
|
||||
|
||||
- *"Which activities receive data from CloneRepo?"*
|
||||
- *"What's the data flow from Analyze to Report?"*
|
||||
- *"Which relations were added in version 3?"*
|
||||
- *"Show me all type-mismatched connections"*
|
||||
- *"Find workflows with Security Scan that require approval"*
|
||||
|
||||
---
|
||||
|
||||
## Versioning Strategy
|
||||
|
||||
### Version Numbering
|
||||
|
||||
- Increment on every canvas change
|
||||
- Track change_type: `manual`, `auto_reasoned`, `import`, `rag_query_applied`
|
||||
- Store reasoner_confidence for auto changes
|
||||
|
||||
### Relation Wording Versions
|
||||
|
||||
- Each relation has independent wording history
|
||||
- Confidence scores tracked per version
|
||||
- Sources: user input, LLM reasoner, import, RAG query
|
||||
|
||||
### Changelog
|
||||
|
||||
```json
|
||||
{
|
||||
"workflow_id": "workflow-1",
|
||||
"total_versions": 5,
|
||||
"changes": [
|
||||
{
|
||||
"version": 1,
|
||||
"type": "created",
|
||||
"timestamp": "2025-09-04T10:00:00Z",
|
||||
"user": "system",
|
||||
"reason": "Workflow initialized"
|
||||
},
|
||||
{
|
||||
"version": 2,
|
||||
"type": "auto_reasoned",
|
||||
"timestamp": "2025-09-04T12:00:00Z",
|
||||
"reasoner_confidence": 0.89,
|
||||
"changes": {
|
||||
"edges_added": 4,
|
||||
"edges_modified": 0
|
||||
}
|
||||
},
|
||||
{
|
||||
"version": 3,
|
||||
"type": "manual",
|
||||
"timestamp": "2025-09-05T10:05:00Z",
|
||||
"user": "user-uuid",
|
||||
"reason": "Connected SecurityScan to Report manually"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Data Flow Example
|
||||
|
||||
### Scenario: User drops SecurityScan node, system suggests connection
|
||||
|
||||
1. **User Action:** Drops SecurityScan node onto existing workflow
|
||||
2. **Frontend Call:** `PUT /workflows/{id}` with new nodes + `auto_reason=true`
|
||||
3. **Backend:**
|
||||
- Saves canvas to `workflow_versions` v.3
|
||||
- Calls CanvasReasonerActivity
|
||||
4. **LLM Reasoning:**
|
||||
- Analyzes: AnalyzeCode (outputs: quality, metrics) → SecurityScan (inputs: path, depth)
|
||||
- Suggests: Add transformer node OR use metrics for decision
|
||||
- Confidence: 0.85 (type mismatch, requires transformation)
|
||||
5. **Relation Wording:**
|
||||
```json
|
||||
{
|
||||
"source": "analyze-code-1",
|
||||
"target": "security-scan-1",
|
||||
"relation_wording": {
|
||||
"verb": "provides-context-for",
|
||||
"source_output": "quality (object): Code quality metrics",
|
||||
"target_input": "path (string): Directory to scan",
|
||||
"connection_type": "requires-transformer",
|
||||
"reasoning": "Metrics inform which files to prioritize in scanning"
|
||||
}
|
||||
}
|
||||
```
|
||||
6. **Graph RAG:** Relations indexed automatically
|
||||
7. **Response:** Frontend shows suggestions with natural language descriptions
|
||||
|
||||
---
|
||||
|
||||
## Implementation Checklist
|
||||
|
||||
- [ ] Create `workflow_versions` table in memory DB
|
||||
- [ ] Create `workflow_relations` table with wording schema
|
||||
- [ ] Create `relation_changes` table for tracking modifications
|
||||
- [ ] Update CanvasReasonerActivity to generate relation wordings
|
||||
- [ ] Implement versioned CRUD endpoints (PUT, GET, DIFF)
|
||||
- [ ] Add Graph RAG indexing on relation creation
|
||||
- [ ] Implement semantic query endpoint
|
||||
- [ ] Add changelog view
|
||||
- [ ] Test point-in-time reconstruction
|
||||
- [ ] Document API in homelab-frontend/API.md
|
||||
|
||||
---
|
||||
|
||||
## Relation Wording Language
|
||||
|
||||
### Verbs (relation_type → verb mapping)
|
||||
|
||||
| Type | Verbs | Example |
|
||||
|------|-------|---------|
|
||||
| `data-flow` | outputs, inputs, receives, provides | "CloneRepo outputs path → AnalyzeCode inputs path" |
|
||||
| `dependency` | must-complete-before, depends-on, requires | "SecurityScan depends-on AnalyzeCode completion" |
|
||||
| `conditional` | triggers-if, branches-on, routes-to | "ApproveWorkflow branches-on result approval" |
|
||||
| `parallel` | runs-alongside, concurrent-with, independent-of | "Notify runs-alongside Report generation" |
|
||||
| `transformation` | transforms, converts, maps, adapts | "LLMTransform converts metrics → deployment plan" |
|
||||
|
||||
### Confidence Scoring
|
||||
|
||||
- **0.9-1.0:** Perfect match (type-compatible, direct data flow)
|
||||
- **0.7-0.9:** Good match (semantic fit, minor transformation needed)
|
||||
- **0.5-0.7:** Possible match (requires user confirmation)
|
||||
- **<0.5:** Poor match (suggest removal or transformer)
|
||||
|
||||
@@ -0,0 +1,234 @@
|
||||
// Example: External service integrating with Poimen RoutingWorkflow
|
||||
// Shows how to submit a task and wait for completion
|
||||
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"go.temporal.io/sdk/client"
|
||||
)
|
||||
|
||||
// ---- Types (mirror internal/routing/types.go) ----
|
||||
|
||||
type WorkflowSpec struct {
|
||||
Name string `json:"name"`
|
||||
Input map[string]interface{} `json:"input,omitempty"`
|
||||
States []State `json:"states"`
|
||||
}
|
||||
|
||||
type State struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"` // Task, Pass, Fail
|
||||
Resource string `json:"resource,omitempty"`
|
||||
Parameters map[string]interface{} `json:"parameters,omitempty"`
|
||||
Timeout string `json:"timeout,omitempty"`
|
||||
Next string `json:"next,omitempty"`
|
||||
End bool `json:"end,omitempty"`
|
||||
}
|
||||
|
||||
type RoutingWorkflowInput struct {
|
||||
Spec *WorkflowSpec `json:"spec"`
|
||||
}
|
||||
|
||||
type RoutingWorkflowOutput struct {
|
||||
Status string `json:"status"`
|
||||
StepResults map[string]map[string]interface{} `json:"stepResults"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// ---- Example Service ----
|
||||
|
||||
type TaskService struct {
|
||||
temporalClient client.Client
|
||||
taskQueue string
|
||||
}
|
||||
|
||||
func NewTaskService(temporalHost, namespace, taskQueue string) (*TaskService, error) {
|
||||
c, err := client.Dial(client.Options{
|
||||
HostPort: temporalHost,
|
||||
Namespace: namespace,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to connect to Temporal: %w", err)
|
||||
}
|
||||
|
||||
return &TaskService{
|
||||
temporalClient: c,
|
||||
taskQueue: taskQueue,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *TaskService) Close() {
|
||||
s.temporalClient.Close()
|
||||
}
|
||||
|
||||
// SubmitAndWait submits a workflow spec and waits for completion
|
||||
func (s *TaskService) SubmitAndWait(ctx context.Context, spec *WorkflowSpec, timeout time.Duration) (*RoutingWorkflowOutput, error) {
|
||||
workflowID := fmt.Sprintf("%s-%d", spec.Name, time.Now().UnixNano())
|
||||
|
||||
// Start workflow
|
||||
run, err := s.temporalClient.ExecuteWorkflow(ctx, client.StartWorkflowOptions{
|
||||
ID: workflowID,
|
||||
TaskQueue: s.taskQueue,
|
||||
}, "RoutingWorkflow", RoutingWorkflowInput{Spec: spec})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to start workflow: %w", err)
|
||||
}
|
||||
|
||||
log.Printf("Workflow started: ID=%s, RunID=%s", run.GetID(), run.GetRunID())
|
||||
|
||||
// Wait for completion with timeout
|
||||
waitCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
var result RoutingWorkflowOutput
|
||||
if err := run.Get(waitCtx, &result); err != nil {
|
||||
return nil, fmt.Errorf("workflow failed: %w", err)
|
||||
}
|
||||
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// SubmitAsync submits workflow and returns immediately (fire-and-forget)
|
||||
func (s *TaskService) SubmitAsync(ctx context.Context, spec *WorkflowSpec) (workflowID string, runID string, err error) {
|
||||
workflowID = fmt.Sprintf("%s-%d", spec.Name, time.Now().UnixNano())
|
||||
|
||||
run, err := s.temporalClient.ExecuteWorkflow(ctx, client.StartWorkflowOptions{
|
||||
ID: workflowID,
|
||||
TaskQueue: s.taskQueue,
|
||||
}, "RoutingWorkflow", RoutingWorkflowInput{Spec: spec})
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("failed to start workflow: %w", err)
|
||||
}
|
||||
|
||||
return run.GetID(), run.GetRunID(), nil
|
||||
}
|
||||
|
||||
// WaitForCompletion waits for an existing workflow to complete
|
||||
func (s *TaskService) WaitForCompletion(ctx context.Context, workflowID string, timeout time.Duration) (*RoutingWorkflowOutput, error) {
|
||||
run := s.temporalClient.GetWorkflow(ctx, workflowID, "")
|
||||
|
||||
waitCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
var result RoutingWorkflowOutput
|
||||
if err := run.Get(waitCtx, &result); err != nil {
|
||||
return nil, fmt.Errorf("workflow failed: %w", err)
|
||||
}
|
||||
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// GetStatus gets current workflow status without waiting
|
||||
func (s *TaskService) GetStatus(ctx context.Context, workflowID string) (string, error) {
|
||||
desc, err := s.temporalClient.DescribeWorkflowExecution(ctx, workflowID, "")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return desc.WorkflowExecutionInfo.Status.String(), nil
|
||||
}
|
||||
|
||||
// ---- Example Usage ----
|
||||
|
||||
func main() {
|
||||
// Connect to Temporal
|
||||
svc, err := NewTaskService(
|
||||
"temporal-frontend.temporal:7233",
|
||||
"poimen-harness",
|
||||
"poimen-taskqueue",
|
||||
)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
defer svc.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Example 1: Implement task T0.3 with code analysis
|
||||
fmt.Println("=== Example 1: Submit and Wait ===")
|
||||
spec := &WorkflowSpec{
|
||||
Name: "implement-T0.3",
|
||||
Input: map[string]interface{}{
|
||||
"taskId": "T0.3",
|
||||
"description": "Implement git worktree management",
|
||||
"repo": "https://github.com/rockliang/poimen",
|
||||
},
|
||||
States: []State{
|
||||
{
|
||||
Name: "Clone",
|
||||
Type: "Task",
|
||||
Resource: "CloneRepoActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"repo": "${workflow.input.repo}",
|
||||
"branch": "main",
|
||||
},
|
||||
Timeout: "5m",
|
||||
Next: "Analyze",
|
||||
},
|
||||
{
|
||||
Name: "Analyze",
|
||||
Type: "Task",
|
||||
Resource: "AnalyzeCodeActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"path": "${Clone.output.path}",
|
||||
"depth": 3,
|
||||
},
|
||||
Timeout: "10m",
|
||||
Next: "Report",
|
||||
},
|
||||
{
|
||||
Name: "Report",
|
||||
Type: "Task",
|
||||
Resource: "GenerateReportActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"analysisResult": "${Analyze.output}",
|
||||
"format": "markdown",
|
||||
},
|
||||
Timeout: "2m",
|
||||
End: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result, err := svc.SubmitAndWait(ctx, spec, 30*time.Minute)
|
||||
if err != nil {
|
||||
log.Printf("Error: %v", err)
|
||||
} else {
|
||||
fmt.Printf("Status: %s\n", result.Status)
|
||||
for step, output := range result.StepResults {
|
||||
fmt.Printf(" %s: %v\n", step, output)
|
||||
}
|
||||
}
|
||||
|
||||
// Example 2: Fire and forget, then poll
|
||||
fmt.Println("\n=== Example 2: Async Submit + Poll ===")
|
||||
workflowID, runID, err := svc.SubmitAsync(ctx, spec)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Printf("Submitted: workflowID=%s, runID=%s\n", workflowID, runID)
|
||||
|
||||
// Poll status
|
||||
for i := 0; i < 5; i++ {
|
||||
status, _ := svc.GetStatus(ctx, workflowID)
|
||||
fmt.Printf(" Poll %d: status=%s\n", i+1, status)
|
||||
if status == "WORKFLOW_EXECUTION_STATUS_COMPLETED" {
|
||||
break
|
||||
}
|
||||
time.Sleep(5 * time.Second)
|
||||
}
|
||||
|
||||
// Get final result
|
||||
result, err = svc.WaitForCompletion(ctx, workflowID, 30*time.Minute)
|
||||
if err != nil {
|
||||
log.Printf("Error: %v", err)
|
||||
} else {
|
||||
resultJSON, _ := json.MarshalIndent(result, "", " ")
|
||||
fmt.Printf("Final result:\n%s\n", resultJSON)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
// Example: WaitForTaskComplete pattern
|
||||
// Use case: External service submits implementation task, waits for result
|
||||
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"go.temporal.io/sdk/client"
|
||||
)
|
||||
|
||||
// WaitForTaskComplete - the core pattern
|
||||
//
|
||||
// 1. Build workflow spec for the task
|
||||
// 2. Submit to Temporal
|
||||
// 3. Block until completion or timeout
|
||||
// 4. Return result
|
||||
func WaitForTaskComplete(
|
||||
c client.Client,
|
||||
taskID string,
|
||||
repo string,
|
||||
timeout time.Duration,
|
||||
) (map[string]interface{}, error) {
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Build spec for implementation task
|
||||
spec := map[string]interface{}{
|
||||
"name": fmt.Sprintf("implement-%s", taskID),
|
||||
"input": map[string]interface{}{
|
||||
"taskId": taskID,
|
||||
"repo": repo,
|
||||
},
|
||||
"states": []map[string]interface{}{
|
||||
{
|
||||
"name": "Clone",
|
||||
"type": "Task",
|
||||
"resource": "CloneRepoActivity",
|
||||
"parameters": map[string]interface{}{
|
||||
"repo": repo,
|
||||
},
|
||||
"next": "Analyze",
|
||||
},
|
||||
{
|
||||
"name": "Analyze",
|
||||
"type": "Task",
|
||||
"resource": "AnalyzeCodeActivity",
|
||||
"parameters": map[string]interface{}{
|
||||
"path": "${Clone.output.path}",
|
||||
},
|
||||
"next": "SecurityScan",
|
||||
},
|
||||
{
|
||||
"name": "SecurityScan",
|
||||
"type": "Task",
|
||||
"resource": "SecurityScanActivity",
|
||||
"parameters": map[string]interface{}{
|
||||
"path": "${Clone.output.path}",
|
||||
},
|
||||
"next": "Report",
|
||||
},
|
||||
{
|
||||
"name": "Report",
|
||||
"type": "Task",
|
||||
"resource": "GenerateReportActivity",
|
||||
"parameters": map[string]interface{}{
|
||||
"analysisResult": "${Analyze.output}",
|
||||
"securityResult": "${SecurityScan.output}",
|
||||
},
|
||||
"end": true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Submit workflow
|
||||
workflowID := fmt.Sprintf("%s-%d", taskID, time.Now().UnixNano())
|
||||
run, err := c.ExecuteWorkflow(ctx, client.StartWorkflowOptions{
|
||||
ID: workflowID,
|
||||
TaskQueue: "poimen-taskqueue",
|
||||
}, "RoutingWorkflow", map[string]interface{}{"spec": spec})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("submit failed: %w", err)
|
||||
}
|
||||
|
||||
fmt.Printf("[%s] Workflow started: %s\n", taskID, workflowID)
|
||||
|
||||
// Wait for completion
|
||||
waitCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
var result map[string]interface{}
|
||||
if err := run.Get(waitCtx, &result); err != nil {
|
||||
return nil, fmt.Errorf("workflow failed: %w", err)
|
||||
}
|
||||
|
||||
fmt.Printf("[%s] Workflow completed: %s\n", taskID, result["status"])
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// Example usage in another service:
|
||||
//
|
||||
// func (s *MyService) ImplementTask(taskID string) error {
|
||||
// c, _ := client.Dial(client.Options{
|
||||
// HostPort: "temporal-frontend.temporal:7233",
|
||||
// Namespace: "poimen-harness",
|
||||
// })
|
||||
// defer c.Close()
|
||||
//
|
||||
// result, err := WaitForTaskComplete(c, taskID, "https://github.com/...", 30*time.Minute)
|
||||
// if err != nil {
|
||||
// return err
|
||||
// }
|
||||
//
|
||||
// // Process result
|
||||
// if result["status"] == "COMPLETED" {
|
||||
// report := result["stepResults"].(map[string]interface{})["Report"]
|
||||
// // Use report...
|
||||
// }
|
||||
// return nil
|
||||
// }
|
||||
@@ -1,4 +1,4 @@
|
||||
module github.com/rockliang/poimen/workflows
|
||||
module forgejo.riotpiao.com/rock/poimen-workflows
|
||||
|
||||
go 1.25.4
|
||||
|
||||
@@ -7,6 +7,7 @@ require (
|
||||
github.com/stretchr/testify v1.12.1
|
||||
go.temporal.io/sdk v1.48.0
|
||||
go.uber.org/zap v1.28.0
|
||||
gopkg.in/yaml.v3 v3.0.1
|
||||
)
|
||||
|
||||
require (
|
||||
@@ -18,6 +19,7 @@ require (
|
||||
github.com/google/uuid v1.6.0 // indirect
|
||||
github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.3.2 // indirect
|
||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.22.0 // indirect
|
||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
||||
github.com/nexus-rpc/nexus-proto-annotations v0.1.0 // indirect
|
||||
github.com/nexus-rpc/sdk-go v0.7.0 // indirect
|
||||
@@ -25,6 +27,8 @@ require (
|
||||
github.com/prometheus/common v0.70.1 // indirect
|
||||
github.com/prometheus/procfs v0.21.1 // indirect
|
||||
github.com/robfig/cron v1.2.0 // indirect
|
||||
github.com/spf13/cobra v1.10.2 // indirect
|
||||
github.com/spf13/pflag v1.0.9 // indirect
|
||||
github.com/stretchr/objx v0.5.3 // indirect
|
||||
go.temporal.io/api v1.63.4 // indirect
|
||||
go.uber.org/multierr v1.11.0 // indirect
|
||||
|
||||
@@ -2,6 +2,7 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
||||
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
||||
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
|
||||
github.com/facebookgo/clock v0.0.0-20150410010913-600d898af40a h1:yDWHCSQ40h88yih2JAcL6Ls/kVkSE8GFACTGVnMPruw=
|
||||
github.com/facebookgo/clock v0.0.0-20150410010913-600d898af40a/go.mod h1:7Ga40egUymuWXxAe151lTNnCv97MddSOVsjpPPkityA=
|
||||
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
|
||||
@@ -22,10 +23,16 @@ github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.3.2 h1:sGm2vDRFUrQJO/Veii4h4z
|
||||
github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.3.2/go.mod h1:wd1YpapPLivG6nQgbf7ZkG1hhSOXDhhn4MLTknx2aAc=
|
||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.22.0 h1:asbCHRVmodnJTuQ3qamDwqVOIjwqUPTYmYuemVOx+Ys=
|
||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.22.0/go.mod h1:ggCgvZ2r7uOoQjOyu2Y1NhHmEPPzzuhWgcza5M1Ji1I=
|
||||
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
|
||||
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
|
||||
github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8=
|
||||
github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck=
|
||||
github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk=
|
||||
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
||||
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
|
||||
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
|
||||
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
||||
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
|
||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
||||
@@ -44,6 +51,13 @@ github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+
|
||||
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
|
||||
github.com/robfig/cron v1.2.0 h1:ZjScXvvxeQ63Dbyxy76Fj3AT3Ut0aKsyd2/tl3DTMuQ=
|
||||
github.com/robfig/cron v1.2.0/go.mod h1:JGuDeoQd7Z6yL4zQhZ3OPEVHB7fL6Ka6skscFHfmt2k=
|
||||
github.com/rogpeppe/go-internal v1.11.0 h1:cWPaGQEPrBb5/AsnsZesgZZ9yb1OQ+GOISoDNXVBh4M=
|
||||
github.com/rogpeppe/go-internal v1.11.0/go.mod h1:ddIwULY96R17DhadqLgMfk9H9tvdUzkipdSkR5nkCZA=
|
||||
github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM=
|
||||
github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU=
|
||||
github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4=
|
||||
github.com/spf13/pflag v1.0.9 h1:9exaQaMOCwffKiiiYk6/BndUBv+iRViNW+4lEMi0PvY=
|
||||
github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
|
||||
github.com/stretchr/objx v0.5.3 h1:jmXUvGomnU1o3W/V5h2VEradbpJDwGrzugQQvL0POH4=
|
||||
github.com/stretchr/objx v0.5.3/go.mod h1:rDQraq+vQZU7Fde9LOZLr8Tax6zZvy4kuNKF+QYS+U0=
|
||||
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
|
||||
@@ -75,6 +89,7 @@ go.uber.org/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo=
|
||||
go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q=
|
||||
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
||||
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
|
||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||
@@ -130,3 +145,8 @@ google.golang.org/grpc v1.82.1 h1:NnAxzGRA0677vCa4BUkOAnO5+FfQqVl9iUXeD0IqcGE=
|
||||
google.golang.org/grpc v1.82.1/go.mod h1:yzTZ1TB1Z3SG+LIYaI+WiE8D5+PZ3ArnrSp8zF3+/ZA=
|
||||
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
|
||||
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
|
||||
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
|
||||
@@ -0,0 +1,217 @@
|
||||
package alerting
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// AlertLevel represents alert severity
|
||||
type AlertLevel string
|
||||
|
||||
const (
|
||||
AlertWarning AlertLevel = "warning"
|
||||
AlertError AlertLevel = "error"
|
||||
AlertCritical AlertLevel = "critical"
|
||||
)
|
||||
|
||||
// Alert represents an alert notification
|
||||
type Alert struct {
|
||||
ID string
|
||||
Level AlertLevel
|
||||
Title string
|
||||
Message string
|
||||
Timestamp time.Time
|
||||
Resolved bool
|
||||
Source string
|
||||
}
|
||||
|
||||
// AlertRule represents a rule that triggers alerts
|
||||
type AlertRule struct {
|
||||
ID string
|
||||
Name string
|
||||
Threshold float64
|
||||
Metric string
|
||||
Level AlertLevel
|
||||
}
|
||||
|
||||
// AlertManager manages alert rules and notifications
|
||||
type AlertManager struct {
|
||||
mu sync.RWMutex
|
||||
rules map[string]*AlertRule
|
||||
alerts map[string]*Alert
|
||||
history []*Alert
|
||||
}
|
||||
|
||||
// NewAlertManager creates a new alert manager
|
||||
func NewAlertManager() *AlertManager {
|
||||
return &AlertManager{
|
||||
rules: make(map[string]*AlertRule),
|
||||
alerts: make(map[string]*Alert),
|
||||
history: make([]*Alert, 0),
|
||||
}
|
||||
}
|
||||
|
||||
// AddRule adds an alert rule
|
||||
func (am *AlertManager) AddRule(rule *AlertRule) error {
|
||||
if rule.ID == "" || rule.Name == "" {
|
||||
return fmt.Errorf("rule ID and name required")
|
||||
}
|
||||
|
||||
am.mu.Lock()
|
||||
defer am.mu.Unlock()
|
||||
|
||||
am.rules[rule.ID] = rule
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveRule removes an alert rule
|
||||
func (am *AlertManager) RemoveRule(ruleID string) error {
|
||||
am.mu.Lock()
|
||||
defer am.mu.Unlock()
|
||||
|
||||
if _, exists := am.rules[ruleID]; !exists {
|
||||
return fmt.Errorf("rule not found: %s", ruleID)
|
||||
}
|
||||
|
||||
delete(am.rules, ruleID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// TriggerAlert triggers a new alert
|
||||
func (am *AlertManager) TriggerAlert(title, message string, level AlertLevel) (*Alert, error) {
|
||||
if title == "" {
|
||||
return nil, fmt.Errorf("alert title required")
|
||||
}
|
||||
|
||||
am.mu.Lock()
|
||||
defer am.mu.Unlock()
|
||||
|
||||
alert := &Alert{
|
||||
ID: fmt.Sprintf("alert-%d", len(am.alerts)),
|
||||
Level: level,
|
||||
Title: title,
|
||||
Message: message,
|
||||
Timestamp: time.Now(),
|
||||
Resolved: false,
|
||||
}
|
||||
|
||||
am.alerts[alert.ID] = alert
|
||||
am.history = append(am.history, alert)
|
||||
|
||||
return alert, nil
|
||||
}
|
||||
|
||||
// ResolveAlert marks an alert as resolved
|
||||
func (am *AlertManager) ResolveAlert(alertID string) error {
|
||||
am.mu.Lock()
|
||||
defer am.mu.Unlock()
|
||||
|
||||
alert, exists := am.alerts[alertID]
|
||||
if !exists {
|
||||
return fmt.Errorf("alert not found: %s", alertID)
|
||||
}
|
||||
|
||||
alert.Resolved = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetActiveAlerts returns all unresolved alerts
|
||||
func (am *AlertManager) GetActiveAlerts() []*Alert {
|
||||
am.mu.RLock()
|
||||
defer am.mu.RUnlock()
|
||||
|
||||
result := make([]*Alert, 0)
|
||||
for _, alert := range am.alerts {
|
||||
if !alert.Resolved {
|
||||
result = append(result, alert)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetAlertsByLevel returns alerts by severity level
|
||||
func (am *AlertManager) GetAlertsByLevel(level AlertLevel) []*Alert {
|
||||
am.mu.RLock()
|
||||
defer am.mu.RUnlock()
|
||||
|
||||
result := make([]*Alert, 0)
|
||||
for _, alert := range am.alerts {
|
||||
if alert.Level == level {
|
||||
result = append(result, alert)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetHistory returns alert history
|
||||
func (am *AlertManager) GetHistory() []*Alert {
|
||||
am.mu.RLock()
|
||||
defer am.mu.RUnlock()
|
||||
|
||||
result := make([]*Alert, len(am.history))
|
||||
copy(result, am.history)
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetRules returns all alert rules
|
||||
func (am *AlertManager) GetRules() map[string]*AlertRule {
|
||||
am.mu.RLock()
|
||||
defer am.mu.RUnlock()
|
||||
|
||||
result := make(map[string]*AlertRule)
|
||||
for id, rule := range am.rules {
|
||||
result[id] = rule
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetAlertCount returns total active alert count
|
||||
func (am *AlertManager) GetAlertCount() int {
|
||||
am.mu.RLock()
|
||||
defer am.mu.RUnlock()
|
||||
|
||||
return len(am.alerts)
|
||||
}
|
||||
|
||||
// Clear clears all alerts
|
||||
func (am *AlertManager) Clear() {
|
||||
am.mu.Lock()
|
||||
defer am.mu.Unlock()
|
||||
|
||||
am.alerts = make(map[string]*Alert)
|
||||
}
|
||||
|
||||
// EvaluateRule checks if a metric triggers an alert rule
|
||||
func (am *AlertManager) EvaluateRule(ruleID string, metricValue float64) (*Alert, error) {
|
||||
am.mu.Lock()
|
||||
defer am.mu.Unlock()
|
||||
|
||||
rule, exists := am.rules[ruleID]
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("rule not found: %s", ruleID)
|
||||
}
|
||||
|
||||
if metricValue >= rule.Threshold {
|
||||
alert := &Alert{
|
||||
ID: fmt.Sprintf("alert-%d", len(am.alerts)),
|
||||
Level: rule.Level,
|
||||
Title: rule.Name,
|
||||
Message: fmt.Sprintf("Threshold %.2f exceeded: %.2f", rule.Threshold, metricValue),
|
||||
Timestamp: time.Now(),
|
||||
Resolved: false,
|
||||
Source: ruleID,
|
||||
}
|
||||
|
||||
am.alerts[alert.ID] = alert
|
||||
am.history = append(am.history, alert)
|
||||
|
||||
return alert, nil
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
package alerting
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestAddRule(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
rule := &AlertRule{
|
||||
ID: "rule-1",
|
||||
Name: "High Error Rate",
|
||||
Threshold: 0.1,
|
||||
Metric: "error_rate",
|
||||
Level: AlertError,
|
||||
}
|
||||
|
||||
err := am.AddRule(rule)
|
||||
assert.NoError(t, err)
|
||||
|
||||
rules := am.GetRules()
|
||||
assert.Equal(t, 1, len(rules))
|
||||
}
|
||||
|
||||
func TestTriggerAlert(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
alert, err := am.TriggerAlert("Database Down", "PostgreSQL unavailable", AlertCritical)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, alert)
|
||||
assert.Equal(t, AlertCritical, alert.Level)
|
||||
}
|
||||
|
||||
func TestResolveAlert(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
alert, _ := am.TriggerAlert("Test Alert", "Test", AlertWarning)
|
||||
err := am.ResolveAlert(alert.ID)
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, alert.Resolved)
|
||||
}
|
||||
|
||||
func TestGetActiveAlerts(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
alert1, _ := am.TriggerAlert("Alert 1", "Test", AlertWarning)
|
||||
alert2, _ := am.TriggerAlert("Alert 2", "Test", AlertError)
|
||||
|
||||
am.ResolveAlert(alert1.ID)
|
||||
|
||||
active := am.GetActiveAlerts()
|
||||
assert.Equal(t, 1, len(active))
|
||||
assert.Equal(t, alert2.ID, active[0].ID)
|
||||
}
|
||||
|
||||
func TestGetAlertsByLevel(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
am.TriggerAlert("Alert 1", "Test", AlertWarning)
|
||||
am.TriggerAlert("Alert 2", "Test", AlertError)
|
||||
am.TriggerAlert("Alert 3", "Test", AlertError)
|
||||
|
||||
errors := am.GetAlertsByLevel(AlertError)
|
||||
assert.Equal(t, 2, len(errors))
|
||||
}
|
||||
|
||||
func TestGetHistory(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
am.TriggerAlert("Alert 1", "Test", AlertWarning)
|
||||
am.TriggerAlert("Alert 2", "Test", AlertError)
|
||||
|
||||
history := am.GetHistory()
|
||||
assert.Equal(t, 2, len(history))
|
||||
}
|
||||
|
||||
func TestRemoveRule(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
rule := &AlertRule{
|
||||
ID: "rule-1",
|
||||
Name: "Test Rule",
|
||||
Level: AlertWarning,
|
||||
}
|
||||
|
||||
am.AddRule(rule)
|
||||
err := am.RemoveRule("rule-1")
|
||||
|
||||
assert.NoError(t, err)
|
||||
rules := am.GetRules()
|
||||
assert.Equal(t, 0, len(rules))
|
||||
}
|
||||
|
||||
func TestClear(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
am.TriggerAlert("Alert 1", "Test", AlertWarning)
|
||||
am.TriggerAlert("Alert 2", "Test", AlertError)
|
||||
|
||||
am.Clear()
|
||||
|
||||
assert.Equal(t, 0, am.GetAlertCount())
|
||||
}
|
||||
|
||||
func TestEvaluateRule(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
rule := &AlertRule{
|
||||
ID: "rule-1",
|
||||
Name: "High Error Rate",
|
||||
Threshold: 0.1,
|
||||
Level: AlertError,
|
||||
}
|
||||
|
||||
am.AddRule(rule)
|
||||
|
||||
alert, _ := am.EvaluateRule("rule-1", 0.15)
|
||||
assert.NotNil(t, alert)
|
||||
}
|
||||
|
||||
func TestEvaluateRuleBelowThreshold(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
rule := &AlertRule{
|
||||
ID: "rule-1",
|
||||
Name: "High Error Rate",
|
||||
Threshold: 0.1,
|
||||
Level: AlertError,
|
||||
}
|
||||
|
||||
am.AddRule(rule)
|
||||
|
||||
alert, _ := am.EvaluateRule("rule-1", 0.05)
|
||||
assert.Nil(t, alert)
|
||||
}
|
||||
|
||||
func TestGetAlertCount(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
am.TriggerAlert("Alert 1", "Test", AlertWarning)
|
||||
am.TriggerAlert("Alert 2", "Test", AlertError)
|
||||
|
||||
assert.Equal(t, 2, am.GetAlertCount())
|
||||
}
|
||||
|
||||
func TestAddRuleError(t *testing.T) {
|
||||
am := NewAlertManager()
|
||||
|
||||
rule := &AlertRule{
|
||||
Name: "No ID",
|
||||
}
|
||||
|
||||
err := am.AddRule(rule)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"go.temporal.io/sdk/client"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/pkg/db"
|
||||
)
|
||||
|
||||
// Server handles HTTP routing for workflow APIs
|
||||
type Server struct {
|
||||
api *WorkflowAPI
|
||||
logger *log.Logger
|
||||
}
|
||||
|
||||
// NewServer creates new HTTP server with database connection
|
||||
func NewServer(database *db.DB, temporalClient client.Client, logger *log.Logger) *Server {
|
||||
return &Server{
|
||||
api: NewWorkflowAPI(database, temporalClient, logger),
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
// ServeHTTP dispatches HTTP requests to appropriate handler
|
||||
func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
// Enable CORS
|
||||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||||
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS")
|
||||
w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization")
|
||||
|
||||
if r.Method == http.MethodOptions {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
return
|
||||
}
|
||||
|
||||
path := r.URL.Path
|
||||
method := r.Method
|
||||
|
||||
s.logger.Printf("%s %s", method, path)
|
||||
|
||||
// Route requests
|
||||
switch {
|
||||
// Workflow endpoints
|
||||
case path == "/workflows" && method == http.MethodPost:
|
||||
s.api.CreateWorkflow(w, r)
|
||||
case path == "/workflows" && method == http.MethodGet:
|
||||
s.api.ListWorkflows(w, r)
|
||||
case strings.HasPrefix(path, "/workflows/") && method == http.MethodGet:
|
||||
id := strings.TrimPrefix(path, "/workflows/")
|
||||
// Exclude special paths
|
||||
if !strings.Contains(id, "/") {
|
||||
s.api.GetWorkflow(w, r, id)
|
||||
} else if strings.HasSuffix(id, "/executions") {
|
||||
// GET /workflows/{id}/executions
|
||||
workflowID := strings.TrimSuffix(id, "/executions")
|
||||
s.api.ListExecutions(w, r, workflowID)
|
||||
}
|
||||
case strings.HasPrefix(path, "/workflows/") && method == http.MethodPut:
|
||||
id := extractID(path, "/workflows/")
|
||||
s.api.UpdateWorkflow(w, r, id)
|
||||
case strings.HasPrefix(path, "/workflows/") && method == http.MethodDelete:
|
||||
id := extractID(path, "/workflows/")
|
||||
s.api.DeleteWorkflow(w, r, id)
|
||||
|
||||
// Execute workflow
|
||||
case strings.HasSuffix(path, "/execute") && method == http.MethodPost:
|
||||
// POST /workflows/{id}/execute
|
||||
parts := strings.Split(path, "/")
|
||||
if len(parts) >= 4 && parts[1] == "workflows" && parts[3] == "execute" {
|
||||
s.api.ExecuteWorkflow(w, r, parts[2])
|
||||
}
|
||||
|
||||
// Execution endpoints
|
||||
case strings.HasPrefix(path, "/executions/") && method == http.MethodGet:
|
||||
id := extractID(path, "/executions/")
|
||||
s.api.GetExecution(w, r, id)
|
||||
|
||||
default:
|
||||
http.Error(w, "Not found", http.StatusNotFound)
|
||||
}
|
||||
}
|
||||
|
||||
// extractID extracts resource ID from path
|
||||
func extractID(path, prefix string) string {
|
||||
id := strings.TrimPrefix(path, prefix)
|
||||
if idx := strings.Index(id, "/"); idx != -1 {
|
||||
return id[:idx]
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
// Start starts the HTTP server
|
||||
func (s *Server) Start(port int) error {
|
||||
addr := fmt.Sprintf(":%d", port)
|
||||
s.logger.Printf("Starting API server on %s", addr)
|
||||
return http.ListenAndServe(addr, s)
|
||||
}
|
||||
@@ -0,0 +1,589 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"go.temporal.io/sdk/client"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/internal/routing"
|
||||
"github.com/rockliang/poimen/workflows/pkg/db"
|
||||
)
|
||||
|
||||
// WorkflowNode matches frontend node type
|
||||
type WorkflowNode struct {
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"` // "activity", "start", "end"
|
||||
Position map[string]interface{} `json:"position"`
|
||||
Data struct {
|
||||
Label string `json:"label"`
|
||||
Activity string `json:"activity"`
|
||||
Config map[string]interface{} `json:"config"`
|
||||
} `json:"data"`
|
||||
}
|
||||
|
||||
// WorkflowEdge matches frontend edge type
|
||||
type WorkflowEdge struct {
|
||||
ID string `json:"id"`
|
||||
Source string `json:"source"`
|
||||
Target string `json:"target"`
|
||||
Data map[string]interface{} `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
// WorkflowDef is the request body for creating/updating workflows
|
||||
type WorkflowDef struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Nodes []WorkflowNode `json:"nodes"`
|
||||
Edges []WorkflowEdge `json:"edges"`
|
||||
Status string `json:"status"` // "draft", "active"
|
||||
}
|
||||
|
||||
// WorkflowResponse is the workflow with metadata
|
||||
type WorkflowResponse struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Status string `json:"status"`
|
||||
Version int `json:"version"`
|
||||
Nodes []WorkflowNode `json:"nodes"`
|
||||
Edges []WorkflowEdge `json:"edges"`
|
||||
CreatedAt string `json:"createdAt"`
|
||||
UpdatedAt string `json:"updatedAt"`
|
||||
CreatedBy string `json:"createdBy"`
|
||||
}
|
||||
|
||||
// ExecutionRequest is the request to execute a workflow
|
||||
type ExecutionRequest struct {
|
||||
Inputs map[string]interface{} `json:"inputs"`
|
||||
}
|
||||
|
||||
// ExecutionResponse is the execution result
|
||||
type ExecutionResponse struct {
|
||||
ID string `json:"id"`
|
||||
WorkflowID string `json:"workflowId"`
|
||||
Status string `json:"status"` // "pending", "running", "success", "failed"
|
||||
StartedAt string `json:"startedAt"`
|
||||
CompletedAt string `json:"completedAt,omitempty"`
|
||||
Inputs map[string]interface{} `json:"inputs"`
|
||||
Outputs map[string]interface{} `json:"outputs,omitempty"`
|
||||
Errors []string `json:"errors,omitempty"`
|
||||
Logs []ExecutionLog `json:"logs"`
|
||||
}
|
||||
|
||||
// ExecutionLog is a log entry from execution
|
||||
type ExecutionLog struct {
|
||||
Timestamp string `json:"timestamp"`
|
||||
NodeID string `json:"nodeId"`
|
||||
Level string `json:"level"` // "info", "warn", "error"
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// WorkflowAPI handles workflow endpoints
|
||||
type WorkflowAPI struct {
|
||||
db *db.DB
|
||||
temporalClient client.Client
|
||||
logger *log.Logger
|
||||
customerID string // TODO: Extract from JWT token
|
||||
}
|
||||
|
||||
// NewWorkflowAPI creates new API handler
|
||||
func NewWorkflowAPI(database *db.DB, tc client.Client, logger *log.Logger) *WorkflowAPI {
|
||||
return &WorkflowAPI{
|
||||
db: database,
|
||||
temporalClient: tc,
|
||||
logger: logger,
|
||||
customerID: "default-customer", // TODO: From auth context
|
||||
}
|
||||
}
|
||||
|
||||
// CreateWorkflow handles POST /workflows
|
||||
func (api *WorkflowAPI) CreateWorkflow(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
var req WorkflowDef
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
http.Error(w, fmt.Sprintf("Invalid request: %v", err), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" {
|
||||
http.Error(w, "Workflow name required", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
// Create workflow in database
|
||||
id := uuid.New().String()
|
||||
now := time.Now()
|
||||
|
||||
// Convert nodes and edges to JSONB
|
||||
nodesJSON, err := json.Marshal(req.Nodes)
|
||||
if err != nil {
|
||||
http.Error(w, fmt.Sprintf("Failed to marshal nodes: %v", err), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
edgesJSON, err := json.Marshal(req.Edges)
|
||||
if err != nil {
|
||||
http.Error(w, fmt.Sprintf("Failed to marshal edges: %v", err), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
status := req.Status
|
||||
if status == "" {
|
||||
status = "draft"
|
||||
}
|
||||
|
||||
workflow := &db.Workflow{
|
||||
ID: id,
|
||||
CustomerID: api.customerID,
|
||||
Name: req.Name,
|
||||
Description: req.Description,
|
||||
Status: status,
|
||||
Version: 1,
|
||||
Nodes: nodesJSON,
|
||||
Edges: edgesJSON,
|
||||
CreatedBy: "anonymous", // Use JWT claim in real implementation
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
|
||||
if err := api.db.SaveWorkflow(r.Context(), workflow); err != nil {
|
||||
api.logger.Printf("Failed to save workflow: %v", err)
|
||||
http.Error(w, "Failed to create workflow", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
response := WorkflowResponse{
|
||||
ID: workflow.ID,
|
||||
Name: workflow.Name,
|
||||
Description: workflow.Description,
|
||||
Status: workflow.Status,
|
||||
Version: workflow.Version,
|
||||
Nodes: req.Nodes,
|
||||
Edges: req.Edges,
|
||||
CreatedAt: workflow.CreatedAt.Format(time.RFC3339),
|
||||
UpdatedAt: workflow.UpdatedAt.Format(time.RFC3339),
|
||||
CreatedBy: workflow.CreatedBy,
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(response)
|
||||
}
|
||||
|
||||
// ListWorkflows handles GET /workflows
|
||||
func (api *WorkflowAPI) ListWorkflows(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
page := 1
|
||||
limit := 10
|
||||
// Parse pagination params if needed
|
||||
|
||||
workflows, err := api.db.ListWorkflows(r.Context(), api.customerID, limit, (page-1)*limit)
|
||||
if err != nil {
|
||||
api.logger.Printf("Failed to list workflows: %v", err)
|
||||
http.Error(w, "Failed to list workflows", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
list := make([]WorkflowResponse, 0)
|
||||
for _, wf := range workflows {
|
||||
var nodes []WorkflowNode
|
||||
var edges []WorkflowEdge
|
||||
|
||||
json.Unmarshal(wf.Nodes, &nodes)
|
||||
json.Unmarshal(wf.Edges, &edges)
|
||||
|
||||
list = append(list, WorkflowResponse{
|
||||
ID: wf.ID,
|
||||
Name: wf.Name,
|
||||
Description: wf.Description,
|
||||
Status: wf.Status,
|
||||
Version: wf.Version,
|
||||
Nodes: nodes,
|
||||
Edges: edges,
|
||||
CreatedAt: wf.CreatedAt.Format(time.RFC3339),
|
||||
UpdatedAt: wf.UpdatedAt.Format(time.RFC3339),
|
||||
CreatedBy: wf.CreatedBy,
|
||||
})
|
||||
}
|
||||
|
||||
response := map[string]interface{}{
|
||||
"workflows": list,
|
||||
"total": len(list),
|
||||
"page": page,
|
||||
"limit": limit,
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(response)
|
||||
}
|
||||
|
||||
// GetWorkflow handles GET /workflows/{id}
|
||||
func (api *WorkflowAPI) GetWorkflow(w http.ResponseWriter, r *http.Request, id string) {
|
||||
if r.Method != http.MethodGet {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
workflow, err := api.db.FetchWorkflow(r.Context(), id, api.customerID)
|
||||
if err != nil {
|
||||
http.Error(w, "Workflow not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
var nodes []WorkflowNode
|
||||
var edges []WorkflowEdge
|
||||
|
||||
json.Unmarshal(workflow.Nodes, &nodes)
|
||||
json.Unmarshal(workflow.Edges, &edges)
|
||||
|
||||
response := WorkflowResponse{
|
||||
ID: workflow.ID,
|
||||
Name: workflow.Name,
|
||||
Description: workflow.Description,
|
||||
Status: workflow.Status,
|
||||
Version: workflow.Version,
|
||||
Nodes: nodes,
|
||||
Edges: edges,
|
||||
CreatedAt: workflow.CreatedAt.Format(time.RFC3339),
|
||||
UpdatedAt: workflow.UpdatedAt.Format(time.RFC3339),
|
||||
CreatedBy: workflow.CreatedBy,
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(response)
|
||||
}
|
||||
|
||||
// UpdateWorkflow handles PUT /workflows/{id}
|
||||
func (api *WorkflowAPI) UpdateWorkflow(w http.ResponseWriter, r *http.Request, id string) {
|
||||
if r.Method != http.MethodPut {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
// Fetch existing workflow
|
||||
workflow, err := api.db.FetchWorkflow(r.Context(), id, api.customerID)
|
||||
if err != nil {
|
||||
http.Error(w, "Workflow not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
var req WorkflowDef
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
http.Error(w, fmt.Sprintf("Invalid request: %v", err), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
// Update fields
|
||||
if req.Name != "" {
|
||||
workflow.Name = req.Name
|
||||
}
|
||||
if req.Description != "" {
|
||||
workflow.Description = req.Description
|
||||
}
|
||||
if req.Nodes != nil {
|
||||
nodesJSON, _ := json.Marshal(req.Nodes)
|
||||
workflow.Nodes = nodesJSON
|
||||
}
|
||||
if req.Edges != nil {
|
||||
edgesJSON, _ := json.Marshal(req.Edges)
|
||||
workflow.Edges = edgesJSON
|
||||
}
|
||||
if req.Status != "" {
|
||||
workflow.Status = req.Status
|
||||
}
|
||||
|
||||
workflow.Version++
|
||||
workflow.UpdatedAt = time.Now()
|
||||
|
||||
if err := api.db.SaveWorkflow(r.Context(), workflow); err != nil {
|
||||
api.logger.Printf("Failed to update workflow: %v", err)
|
||||
http.Error(w, "Failed to update workflow", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
var nodes []WorkflowNode
|
||||
var edges []WorkflowEdge
|
||||
|
||||
json.Unmarshal(workflow.Nodes, &nodes)
|
||||
json.Unmarshal(workflow.Edges, &edges)
|
||||
|
||||
response := WorkflowResponse{
|
||||
ID: workflow.ID,
|
||||
Name: workflow.Name,
|
||||
Description: workflow.Description,
|
||||
Status: workflow.Status,
|
||||
Version: workflow.Version,
|
||||
Nodes: nodes,
|
||||
Edges: edges,
|
||||
CreatedAt: workflow.CreatedAt.Format(time.RFC3339),
|
||||
UpdatedAt: workflow.UpdatedAt.Format(time.RFC3339),
|
||||
CreatedBy: workflow.CreatedBy,
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(response)
|
||||
}
|
||||
|
||||
// DeleteWorkflow handles DELETE /workflows/{id}
|
||||
func (api *WorkflowAPI) DeleteWorkflow(w http.ResponseWriter, r *http.Request, id string) {
|
||||
if r.Method != http.MethodDelete {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
if err := api.db.DeleteWorkflow(r.Context(), id, api.customerID); err != nil {
|
||||
http.Error(w, "Workflow not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// ExecuteWorkflow handles POST /workflows/{id}/execute
|
||||
func (api *WorkflowAPI) ExecuteWorkflow(w http.ResponseWriter, r *http.Request, id string) {
|
||||
if r.Method != http.MethodPost {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
workflow, err := api.db.FetchWorkflow(r.Context(), id, api.customerID)
|
||||
if err != nil {
|
||||
http.Error(w, "Workflow not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
var req ExecutionRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
http.Error(w, fmt.Sprintf("Invalid request: %v", err), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
// Unmarshal nodes and edges
|
||||
var nodes []WorkflowNode
|
||||
var edges []WorkflowEdge
|
||||
json.Unmarshal(workflow.Nodes, &nodes)
|
||||
json.Unmarshal(workflow.Edges, &edges)
|
||||
|
||||
// Convert to workflow response for spec conversion
|
||||
workflowResp := &WorkflowResponse{
|
||||
ID: workflow.ID,
|
||||
Name: workflow.Name,
|
||||
Description: workflow.Description,
|
||||
Status: workflow.Status,
|
||||
Version: workflow.Version,
|
||||
Nodes: nodes,
|
||||
Edges: edges,
|
||||
CreatedBy: workflow.CreatedBy,
|
||||
}
|
||||
|
||||
// Convert nodes/edges to WorkflowSpec
|
||||
spec := api.nodesToWorkflowSpec(workflowResp, req.Inputs)
|
||||
|
||||
// Execute via Temporal RoutingWorkflow
|
||||
execID := uuid.New().String()
|
||||
workflowOptions := client.StartWorkflowOptions{
|
||||
ID: execID,
|
||||
TaskQueue: "default",
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
_, err = api.temporalClient.ExecuteWorkflow(ctx, workflowOptions, "RoutingWorkflow", spec)
|
||||
if err != nil {
|
||||
api.logger.Printf("Failed to execute workflow: %v", err)
|
||||
http.Error(w, fmt.Sprintf("Execution failed: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Save execution to database
|
||||
inputsJSON, _ := json.Marshal(req.Inputs)
|
||||
now := time.Now()
|
||||
|
||||
execution := &db.WorkflowExecution{
|
||||
ID: execID,
|
||||
WorkflowID: id,
|
||||
CustomerID: api.customerID,
|
||||
TemporalID: execID,
|
||||
Status: "running",
|
||||
Inputs: inputsJSON,
|
||||
StartedAt: now,
|
||||
}
|
||||
|
||||
if err := api.db.SaveExecution(r.Context(), execution); err != nil {
|
||||
api.logger.Printf("Failed to save execution: %v", err)
|
||||
http.Error(w, "Failed to save execution", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Create execution response
|
||||
execResp := ExecutionResponse{
|
||||
ID: execID,
|
||||
WorkflowID: id,
|
||||
Status: "running",
|
||||
StartedAt: now.Format(time.RFC3339),
|
||||
Inputs: req.Inputs,
|
||||
Outputs: make(map[string]interface{}),
|
||||
Logs: []ExecutionLog{},
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(execResp)
|
||||
}
|
||||
|
||||
// GetExecution handles GET /executions/{id}
|
||||
func (api *WorkflowAPI) GetExecution(w http.ResponseWriter, r *http.Request, id string) {
|
||||
if r.Method != http.MethodGet {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
execution, err := api.db.FetchExecution(r.Context(), id)
|
||||
if err != nil {
|
||||
http.Error(w, "Execution not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
// Get logs from database
|
||||
logs, err := api.db.FetchExecutionLogs(r.Context(), id)
|
||||
if err != nil {
|
||||
api.logger.Printf("Failed to fetch logs: %v", err)
|
||||
}
|
||||
|
||||
execLogs := make([]ExecutionLog, 0)
|
||||
for _, log := range logs {
|
||||
execLogs = append(execLogs, ExecutionLog{
|
||||
Timestamp: log.LoggedAt.Format(time.RFC3339),
|
||||
NodeID: log.NodeID,
|
||||
Level: log.Level,
|
||||
Message: log.Message,
|
||||
})
|
||||
}
|
||||
|
||||
// Parse inputs/outputs
|
||||
var inputs map[string]interface{}
|
||||
var outputs map[string]interface{}
|
||||
json.Unmarshal(execution.Inputs, &inputs)
|
||||
if execution.Outputs != nil {
|
||||
json.Unmarshal(execution.Outputs, &outputs)
|
||||
}
|
||||
|
||||
// Check Temporal workflow status
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
desc, err := api.temporalClient.DescribeWorkflowExecution(ctx, execution.TemporalID, "")
|
||||
status := execution.Status
|
||||
if err == nil && desc != nil {
|
||||
switch desc.Status.String() {
|
||||
case "RUNNING":
|
||||
status = "running"
|
||||
case "COMPLETED":
|
||||
status = "success"
|
||||
case "FAILED":
|
||||
status = "failed"
|
||||
}
|
||||
}
|
||||
|
||||
completedAtStr := ""
|
||||
if execution.CompletedAt != nil {
|
||||
completedAtStr = execution.CompletedAt.Format(time.RFC3339)
|
||||
}
|
||||
|
||||
execResp := ExecutionResponse{
|
||||
ID: execution.ID,
|
||||
WorkflowID: execution.WorkflowID,
|
||||
Status: status,
|
||||
StartedAt: execution.StartedAt.Format(time.RFC3339),
|
||||
CompletedAt: completedAtStr,
|
||||
Inputs: inputs,
|
||||
Outputs: outputs,
|
||||
Logs: execLogs,
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(execResp)
|
||||
}
|
||||
|
||||
// ListExecutions handles GET /workflows/{id}/executions
|
||||
func (api *WorkflowAPI) ListExecutions(w http.ResponseWriter, r *http.Request, workflowID string) {
|
||||
if r.Method != http.MethodGet {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
// TODO: Implement query by workflow_id in database
|
||||
// For now, return empty list (needs DB method for filtering by workflow_id)
|
||||
list := make([]ExecutionResponse, 0)
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(list)
|
||||
}
|
||||
|
||||
// nodesToWorkflowSpec converts frontend nodes/edges to routing.WorkflowSpec
|
||||
func (api *WorkflowAPI) nodesToWorkflowSpec(wf *WorkflowResponse, inputs map[string]interface{}) *routing.WorkflowSpec {
|
||||
spec := &routing.WorkflowSpec{
|
||||
Name: wf.Name,
|
||||
Input: inputs,
|
||||
States: []routing.State{},
|
||||
}
|
||||
|
||||
// Build states from nodes
|
||||
stateMap := make(map[string]*routing.State)
|
||||
|
||||
// Create all states
|
||||
for _, node := range wf.Nodes {
|
||||
if node.Type == "activity" {
|
||||
state := &routing.State{
|
||||
Name: node.ID,
|
||||
Type: routing.StateTypeTask,
|
||||
Resource: node.Data.Activity,
|
||||
Parameters: node.Data.Config,
|
||||
End: false,
|
||||
}
|
||||
stateMap[node.ID] = state
|
||||
spec.States = append(spec.States, *state)
|
||||
}
|
||||
}
|
||||
|
||||
// Wire edges (transitions)
|
||||
for _, edge := range wf.Edges {
|
||||
if state, exists := stateMap[edge.Source]; exists {
|
||||
state.Next = edge.Target
|
||||
}
|
||||
}
|
||||
|
||||
// Mark last state as End
|
||||
if len(spec.States) > 0 {
|
||||
// Find state with no outgoing edge
|
||||
for i := range spec.States {
|
||||
hasNext := false
|
||||
for _, edge := range wf.Edges {
|
||||
if edge.Source == spec.States[i].Name {
|
||||
hasNext = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !hasNext {
|
||||
spec.States[i].End = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return spec
|
||||
}
|
||||
@@ -0,0 +1,342 @@
|
||||
package approval
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ApprovalStatus represents approval state
|
||||
type ApprovalStatus string
|
||||
|
||||
const (
|
||||
StatusPending ApprovalStatus = "pending"
|
||||
StatusApproved ApprovalStatus = "approved"
|
||||
StatusRejected ApprovalStatus = "rejected"
|
||||
StatusExpired ApprovalStatus = "expired"
|
||||
)
|
||||
|
||||
// ApprovalDecision represents an approval decision
|
||||
type ApprovalDecision struct {
|
||||
Status ApprovalStatus `json:"status"`
|
||||
ApprovedBy string `json:"approved_by"`
|
||||
RejectedBy string `json:"rejected_by"`
|
||||
Reason string `json:"reason"`
|
||||
Comments string `json:"comments"`
|
||||
Timestamp time.Time `json:"timestamp"`
|
||||
ExpiresAt time.Time `json:"expires_at"`
|
||||
}
|
||||
|
||||
// ApprovalGate represents a human approval gate
|
||||
type ApprovalGate struct {
|
||||
ID string
|
||||
TaskID string
|
||||
WorkflowID string
|
||||
Description string
|
||||
Decision *ApprovalDecision
|
||||
CreatedAt time.Time
|
||||
TTL time.Duration // Time until gate expires
|
||||
RequiredApprovals int // Number of approvals needed (1 or more)
|
||||
Approvals []string // List of approvers
|
||||
}
|
||||
|
||||
// ApprovalGateManager manages approval gates
|
||||
type ApprovalGateManager struct {
|
||||
mu sync.RWMutex
|
||||
gates map[string]*ApprovalGate
|
||||
decisions map[string]*ApprovalDecision
|
||||
history []*ApprovalRecord
|
||||
stats *ApprovalStats
|
||||
}
|
||||
|
||||
// ApprovalRecord tracks approval history
|
||||
type ApprovalRecord struct {
|
||||
GateID string
|
||||
Decision ApprovalStatus
|
||||
ApprovedBy string
|
||||
RejectedBy string
|
||||
Timestamp time.Time
|
||||
Reason string
|
||||
}
|
||||
|
||||
// ApprovalStats tracks approval statistics
|
||||
type ApprovalStats struct {
|
||||
TotalGates int
|
||||
ApprovedGates int
|
||||
RejectedGates int
|
||||
PendingGates int
|
||||
ExpiredGates int
|
||||
AverageWaitTime time.Duration
|
||||
}
|
||||
|
||||
// NewApprovalGateManager creates a new approval gate manager
|
||||
func NewApprovalGateManager() *ApprovalGateManager {
|
||||
return &ApprovalGateManager{
|
||||
gates: make(map[string]*ApprovalGate),
|
||||
decisions: make(map[string]*ApprovalDecision),
|
||||
history: make([]*ApprovalRecord, 0),
|
||||
stats: &ApprovalStats{
|
||||
AverageWaitTime: 0,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// CreateGate creates a new approval gate
|
||||
func (agm *ApprovalGateManager) CreateGate(taskID, workflowID, description string, ttl time.Duration) *ApprovalGate {
|
||||
if ttl == 0 {
|
||||
ttl = 24 * time.Hour // Default 24 hours
|
||||
}
|
||||
|
||||
gate := &ApprovalGate{
|
||||
ID: fmt.Sprintf("gate-%d", time.Now().UnixNano()),
|
||||
TaskID: taskID,
|
||||
WorkflowID: workflowID,
|
||||
Description: description,
|
||||
CreatedAt: time.Now(),
|
||||
TTL: ttl,
|
||||
RequiredApprovals: 1,
|
||||
Approvals: make([]string, 0),
|
||||
Decision: &ApprovalDecision{
|
||||
Status: StatusPending,
|
||||
ExpiresAt: time.Now().Add(ttl),
|
||||
},
|
||||
}
|
||||
|
||||
agm.mu.Lock()
|
||||
defer agm.mu.Unlock()
|
||||
|
||||
agm.gates[gate.ID] = gate
|
||||
agm.decisions[gate.ID] = gate.Decision
|
||||
agm.stats.TotalGates++
|
||||
agm.stats.PendingGates++
|
||||
|
||||
return gate
|
||||
}
|
||||
|
||||
// ApproveGate approves a gate
|
||||
func (agm *ApprovalGateManager) ApproveGate(gateID, approvedBy, comments string) error {
|
||||
agm.mu.Lock()
|
||||
defer agm.mu.Unlock()
|
||||
|
||||
gate, exists := agm.gates[gateID]
|
||||
if !exists {
|
||||
return fmt.Errorf("gate not found: %s", gateID)
|
||||
}
|
||||
|
||||
if gate.Decision.Status == StatusApproved || gate.Decision.Status == StatusRejected {
|
||||
return fmt.Errorf("gate already has a decision: %s", gate.Decision.Status)
|
||||
}
|
||||
|
||||
if time.Now().After(gate.Decision.ExpiresAt) {
|
||||
gate.Decision.Status = StatusExpired
|
||||
agm.stats.ExpiredGates++
|
||||
agm.stats.PendingGates--
|
||||
return fmt.Errorf("gate has expired")
|
||||
}
|
||||
|
||||
gate.Decision.Status = StatusApproved
|
||||
gate.Decision.ApprovedBy = approvedBy
|
||||
gate.Decision.Comments = comments
|
||||
gate.Decision.Timestamp = time.Now()
|
||||
|
||||
gate.Approvals = append(gate.Approvals, approvedBy)
|
||||
|
||||
// Record in history
|
||||
record := &ApprovalRecord{
|
||||
GateID: gateID,
|
||||
Decision: StatusApproved,
|
||||
ApprovedBy: approvedBy,
|
||||
Timestamp: gate.Decision.Timestamp,
|
||||
Reason: comments,
|
||||
}
|
||||
|
||||
agm.history = append(agm.history, record)
|
||||
agm.stats.ApprovedGates++
|
||||
agm.stats.PendingGates--
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RejectGate rejects a gate
|
||||
func (agm *ApprovalGateManager) RejectGate(gateID, rejectedBy, reason string) error {
|
||||
agm.mu.Lock()
|
||||
defer agm.mu.Unlock()
|
||||
|
||||
gate, exists := agm.gates[gateID]
|
||||
if !exists {
|
||||
return fmt.Errorf("gate not found: %s", gateID)
|
||||
}
|
||||
|
||||
if gate.Decision.Status == StatusApproved || gate.Decision.Status == StatusRejected {
|
||||
return fmt.Errorf("gate already has a decision: %s", gate.Decision.Status)
|
||||
}
|
||||
|
||||
if time.Now().After(gate.Decision.ExpiresAt) {
|
||||
gate.Decision.Status = StatusExpired
|
||||
agm.stats.ExpiredGates++
|
||||
agm.stats.PendingGates--
|
||||
return fmt.Errorf("gate has expired")
|
||||
}
|
||||
|
||||
gate.Decision.Status = StatusRejected
|
||||
gate.Decision.RejectedBy = rejectedBy
|
||||
gate.Decision.Reason = reason
|
||||
gate.Decision.Timestamp = time.Now()
|
||||
|
||||
// Record in history
|
||||
record := &ApprovalRecord{
|
||||
GateID: gateID,
|
||||
Decision: StatusRejected,
|
||||
RejectedBy: rejectedBy,
|
||||
Timestamp: gate.Decision.Timestamp,
|
||||
Reason: reason,
|
||||
}
|
||||
|
||||
agm.history = append(agm.history, record)
|
||||
agm.stats.RejectedGates++
|
||||
agm.stats.PendingGates--
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetGate retrieves a gate
|
||||
func (agm *ApprovalGateManager) GetGate(gateID string) (*ApprovalGate, bool) {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
gate, exists := agm.gates[gateID]
|
||||
return gate, exists
|
||||
}
|
||||
|
||||
// GetDecision retrieves a decision
|
||||
func (agm *ApprovalGateManager) GetDecision(gateID string) (*ApprovalDecision, bool) {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
decision, exists := agm.decisions[gateID]
|
||||
return decision, exists
|
||||
}
|
||||
|
||||
// GetPendingGates returns all pending gates
|
||||
func (agm *ApprovalGateManager) GetPendingGates() []*ApprovalGate {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
pending := make([]*ApprovalGate, 0)
|
||||
for _, gate := range agm.gates {
|
||||
if gate.Decision.Status == StatusPending {
|
||||
// Check if expired
|
||||
if time.Now().After(gate.Decision.ExpiresAt) {
|
||||
gate.Decision.Status = StatusExpired
|
||||
} else {
|
||||
pending = append(pending, gate)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return pending
|
||||
}
|
||||
|
||||
// GetGatesByTask returns all gates for a task
|
||||
func (agm *ApprovalGateManager) GetGatesByTask(taskID string) []*ApprovalGate {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
gates := make([]*ApprovalGate, 0)
|
||||
for _, gate := range agm.gates {
|
||||
if gate.TaskID == taskID {
|
||||
gates = append(gates, gate)
|
||||
}
|
||||
}
|
||||
|
||||
return gates
|
||||
}
|
||||
|
||||
// GetGatesByWorkflow returns all gates for a workflow
|
||||
func (agm *ApprovalGateManager) GetGatesByWorkflow(workflowID string) []*ApprovalGate {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
gates := make([]*ApprovalGate, 0)
|
||||
for _, gate := range agm.gates {
|
||||
if gate.WorkflowID == workflowID {
|
||||
gates = append(gates, gate)
|
||||
}
|
||||
}
|
||||
|
||||
return gates
|
||||
}
|
||||
|
||||
// IsApproved checks if a gate is approved
|
||||
func (agm *ApprovalGateManager) IsApproved(gateID string) bool {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
gate, exists := agm.gates[gateID]
|
||||
if !exists {
|
||||
return false
|
||||
}
|
||||
|
||||
return gate.Decision.Status == StatusApproved
|
||||
}
|
||||
|
||||
// IsRejected checks if a gate is rejected
|
||||
func (agm *ApprovalGateManager) IsRejected(gateID string) bool {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
gate, exists := agm.gates[gateID]
|
||||
if !exists {
|
||||
return false
|
||||
}
|
||||
|
||||
return gate.Decision.Status == StatusRejected
|
||||
}
|
||||
|
||||
// IsPending checks if a gate is still pending
|
||||
func (agm *ApprovalGateManager) IsPending(gateID string) bool {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
gate, exists := agm.gates[gateID]
|
||||
if !exists {
|
||||
return false
|
||||
}
|
||||
|
||||
if time.Now().After(gate.Decision.ExpiresAt) {
|
||||
return false // Expired gates are not pending
|
||||
}
|
||||
|
||||
return gate.Decision.Status == StatusPending
|
||||
}
|
||||
|
||||
// GetStats returns statistics
|
||||
func (agm *ApprovalGateManager) GetStats() *ApprovalStats {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
stats := *agm.stats
|
||||
return &stats
|
||||
}
|
||||
|
||||
// GetHistory returns approval history
|
||||
func (agm *ApprovalGateManager) GetHistory() []*ApprovalRecord {
|
||||
agm.mu.RLock()
|
||||
defer agm.mu.RUnlock()
|
||||
|
||||
result := make([]*ApprovalRecord, len(agm.history))
|
||||
copy(result, agm.history)
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// Clear clears all gates
|
||||
func (agm *ApprovalGateManager) Clear() {
|
||||
agm.mu.Lock()
|
||||
defer agm.mu.Unlock()
|
||||
|
||||
agm.gates = make(map[string]*ApprovalGate)
|
||||
agm.decisions = make(map[string]*ApprovalDecision)
|
||||
agm.history = make([]*ApprovalRecord, 0)
|
||||
agm.stats = &ApprovalStats{}
|
||||
}
|
||||
@@ -0,0 +1,175 @@
|
||||
package approval
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestNewApprovalGateManager(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
assert.NotNil(t, manager)
|
||||
assert.Equal(t, 0, manager.stats.TotalGates)
|
||||
}
|
||||
|
||||
func TestCreateGate(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
assert.NotNil(t, gate)
|
||||
assert.Equal(t, "T0.1", gate.TaskID)
|
||||
assert.Equal(t, StatusPending, gate.Decision.Status)
|
||||
}
|
||||
|
||||
func TestApproveGate(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
err := manager.ApproveGate(gate.ID, "reviewer-1", "Approved")
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, manager.IsApproved(gate.ID))
|
||||
}
|
||||
|
||||
func TestRejectGate(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
err := manager.RejectGate(gate.ID, "reviewer-1", "Rejected")
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, manager.IsRejected(gate.ID))
|
||||
}
|
||||
|
||||
func TestGetPendingGates(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate1 := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
gate2 := manager.CreateGate("T0.2", "workflow-1", "Review", 24*time.Hour)
|
||||
|
||||
manager.ApproveGate(gate1.ID, "reviewer-1", "")
|
||||
|
||||
pending := manager.GetPendingGates()
|
||||
assert.Equal(t, 1, len(pending))
|
||||
if len(pending) > 0 {
|
||||
// Don't rely on map iteration order - just check it's the unapproved gate
|
||||
assert.Equal(t, StatusPending, pending[0].Decision.Status)
|
||||
assert.Equal(t, gate2.ID, pending[0].ID) // gate2 is the only pending one
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetGatesByTaskFiltering(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
|
||||
gates := manager.GetGatesByTask("T0.1")
|
||||
assert.Equal(t, 1, len(gates))
|
||||
assert.Equal(t, gate.ID, gates[0].ID)
|
||||
}
|
||||
|
||||
func TestGetGatesByWorkflowFiltering(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
|
||||
gates := manager.GetGatesByWorkflow("workflow-1")
|
||||
assert.Equal(t, 1, len(gates))
|
||||
assert.Equal(t, gate.ID, gates[0].ID)
|
||||
}
|
||||
|
||||
func TestIsApproved(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
assert.False(t, manager.IsApproved(gate.ID))
|
||||
|
||||
manager.ApproveGate(gate.ID, "reviewer-1", "")
|
||||
assert.True(t, manager.IsApproved(gate.ID))
|
||||
}
|
||||
|
||||
func TestIsRejected(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
assert.False(t, manager.IsRejected(gate.ID))
|
||||
|
||||
manager.RejectGate(gate.ID, "reviewer-1", "")
|
||||
assert.True(t, manager.IsRejected(gate.ID))
|
||||
}
|
||||
|
||||
func TestIsPending(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
assert.True(t, manager.IsPending(gate.ID))
|
||||
|
||||
manager.ApproveGate(gate.ID, "reviewer-1", "")
|
||||
assert.False(t, manager.IsPending(gate.ID))
|
||||
}
|
||||
|
||||
func TestGetStats(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate1 := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
_ = manager.CreateGate("T0.2", "workflow-1", "Review", 24*time.Hour)
|
||||
|
||||
manager.ApproveGate(gate1.ID, "reviewer-1", "")
|
||||
|
||||
stats := manager.GetStats()
|
||||
assert.Equal(t, 2, stats.TotalGates)
|
||||
assert.Equal(t, 1, stats.ApprovedGates)
|
||||
assert.Equal(t, 1, stats.PendingGates)
|
||||
}
|
||||
|
||||
func TestGetHistory(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate1 := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
gate2 := manager.CreateGate("T0.2", "workflow-1", "Review", 24*time.Hour)
|
||||
|
||||
manager.ApproveGate(gate1.ID, "reviewer-1", "")
|
||||
manager.RejectGate(gate2.ID, "reviewer-2", "Needs work")
|
||||
|
||||
history := manager.GetHistory()
|
||||
assert.Greater(t, len(history), 0)
|
||||
}
|
||||
|
||||
func TestClear(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
assert.Equal(t, 1, manager.stats.TotalGates)
|
||||
|
||||
manager.Clear()
|
||||
assert.Equal(t, 0, manager.stats.TotalGates)
|
||||
}
|
||||
|
||||
func TestExpiredGate(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 1*time.Millisecond)
|
||||
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
err := manager.ApproveGate(gate.ID, "reviewer-1", "")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestMultipleApprovals(t *testing.T) {
|
||||
manager := NewApprovalGateManager()
|
||||
gate := manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
gate.RequiredApprovals = 2
|
||||
|
||||
err := manager.ApproveGate(gate.ID, "reviewer-1", "")
|
||||
assert.NoError(t, err)
|
||||
|
||||
retrieved, _ := manager.GetGate(gate.ID)
|
||||
assert.Equal(t, 1, len(retrieved.Approvals))
|
||||
}
|
||||
|
||||
func BenchmarkCreateGate(b *testing.B) {
|
||||
manager := NewApprovalGateManager()
|
||||
for i := 0; i < b.N; i++ {
|
||||
manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkApproveGate(b *testing.B) {
|
||||
manager := NewApprovalGateManager()
|
||||
gates := make([]*ApprovalGate, b.N)
|
||||
for i := 0; i < b.N; i++ {
|
||||
gates[i] = manager.CreateGate("T0.1", "workflow-1", "Review", 24*time.Hour)
|
||||
}
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
manager.ApproveGate(gates[i].ID, "reviewer-1", "")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ImmutableLogEntry represents a tamper-proof audit entry
|
||||
type ImmutableLogEntry struct {
|
||||
Sequence int64 `json:"sequence"`
|
||||
PrevHash string `json:"prev_hash"`
|
||||
Content string `json:"content"`
|
||||
Hash string `json:"hash"`
|
||||
Timestamp time.Time `json:"timestamp"`
|
||||
Signature string `json:"signature,omitempty"`
|
||||
Metadata map[string]interface{} `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
// ImmutableLog maintains a tamper-proof audit trail
|
||||
type ImmutableLog struct {
|
||||
mu sync.RWMutex
|
||||
entries []*ImmutableLogEntry
|
||||
logPath string
|
||||
sequence int64
|
||||
prevHash string
|
||||
workflowKey string
|
||||
}
|
||||
|
||||
// NewImmutableLog creates a new immutable log
|
||||
func NewImmutableLog(logPath string, workflowKey string) *ImmutableLog {
|
||||
return &ImmutableLog{
|
||||
entries: make([]*ImmutableLogEntry, 0),
|
||||
logPath: logPath,
|
||||
sequence: 0,
|
||||
prevHash: "genesis",
|
||||
workflowKey: workflowKey,
|
||||
}
|
||||
}
|
||||
|
||||
// Append adds an entry to the immutable log
|
||||
func (il *ImmutableLog) Append(content string, metadata map[string]interface{}) (*ImmutableLogEntry, error) {
|
||||
il.mu.Lock()
|
||||
defer il.mu.Unlock()
|
||||
|
||||
il.sequence++
|
||||
hash := il.computeHash(il.sequence, il.prevHash, content)
|
||||
|
||||
entry := &ImmutableLogEntry{
|
||||
Sequence: il.sequence,
|
||||
PrevHash: il.prevHash,
|
||||
Content: content,
|
||||
Hash: hash,
|
||||
Timestamp: time.Now(),
|
||||
Metadata: metadata,
|
||||
}
|
||||
|
||||
il.entries = append(il.entries, entry)
|
||||
il.prevHash = hash
|
||||
|
||||
return entry, nil
|
||||
}
|
||||
|
||||
// Verify verifies the integrity of the log
|
||||
func (il *ImmutableLog) Verify() (bool, error) {
|
||||
il.mu.RLock()
|
||||
defer il.mu.RUnlock()
|
||||
|
||||
prevHash := "genesis"
|
||||
|
||||
for _, entry := range il.entries {
|
||||
expectedHash := il.computeHash(entry.Sequence, entry.PrevHash, entry.Content)
|
||||
|
||||
if entry.Hash != expectedHash || entry.PrevHash != prevHash {
|
||||
return false, fmt.Errorf("integrity check failed at sequence %d", entry.Sequence)
|
||||
}
|
||||
|
||||
prevHash = entry.Hash
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// GetEntries returns all entries
|
||||
func (il *ImmutableLog) GetEntries() []*ImmutableLogEntry {
|
||||
il.mu.RLock()
|
||||
defer il.mu.RUnlock()
|
||||
|
||||
result := make([]*ImmutableLogEntry, len(il.entries))
|
||||
copy(result, il.entries)
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetLastHash returns the last hash
|
||||
func (il *ImmutableLog) GetLastHash() string {
|
||||
il.mu.RLock()
|
||||
defer il.mu.RUnlock()
|
||||
|
||||
return il.prevHash
|
||||
}
|
||||
|
||||
// computeHash computes SHA256 hash
|
||||
func (il *ImmutableLog) computeHash(seq int64, prevHash, content string) string {
|
||||
data := fmt.Sprintf("%d:%s:%s:%s", seq, prevHash, content, il.workflowKey)
|
||||
hash := sha256.Sum256([]byte(data))
|
||||
return fmt.Sprintf("%x", hash)
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestAppendEntry(t *testing.T) {
|
||||
log := NewImmutableLog("", "workflow-1")
|
||||
entry, err := log.Append("Decision: approved", map[string]interface{}{})
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, entry)
|
||||
assert.Equal(t, int64(1), entry.Sequence)
|
||||
}
|
||||
|
||||
func TestVerifyIntegrity(t *testing.T) {
|
||||
log := NewImmutableLog("", "workflow-1")
|
||||
|
||||
log.Append("Entry 1", map[string]interface{}{})
|
||||
log.Append("Entry 2", map[string]interface{}{})
|
||||
|
||||
valid, err := log.Verify()
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, valid)
|
||||
}
|
||||
|
||||
func TestGetEntries(t *testing.T) {
|
||||
log := NewImmutableLog("", "workflow-1")
|
||||
|
||||
log.Append("Entry 1", map[string]interface{}{})
|
||||
log.Append("Entry 2", map[string]interface{}{})
|
||||
|
||||
entries := log.GetEntries()
|
||||
assert.Equal(t, 2, len(entries))
|
||||
}
|
||||
|
||||
func TestChainHashes(t *testing.T) {
|
||||
log := NewImmutableLog("", "workflow-1")
|
||||
|
||||
entry1, _ := log.Append("Entry 1", map[string]interface{}{})
|
||||
entry2, _ := log.Append("Entry 2", map[string]interface{}{})
|
||||
|
||||
assert.Equal(t, "genesis", entry1.PrevHash)
|
||||
assert.Equal(t, entry1.Hash, entry2.PrevHash)
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
package clusters
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ClusterInfo represents a Kubernetes cluster
|
||||
type ClusterInfo struct {
|
||||
Name string
|
||||
APIServer string
|
||||
Healthy bool
|
||||
LastCheck time.Time
|
||||
Capacity int // Max concurrent tasks
|
||||
Usage int // Current task count
|
||||
}
|
||||
|
||||
// ClusterManager manages multiple K8s clusters
|
||||
type ClusterManager struct {
|
||||
mu sync.RWMutex
|
||||
clusters map[string]*ClusterInfo
|
||||
}
|
||||
|
||||
// NewClusterManager creates a new cluster manager
|
||||
func NewClusterManager() *ClusterManager {
|
||||
return &ClusterManager{
|
||||
clusters: make(map[string]*ClusterInfo),
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterCluster registers a new cluster
|
||||
func (cm *ClusterManager) RegisterCluster(name, apiServer string, capacity int) error {
|
||||
if name == "" || apiServer == "" {
|
||||
return fmt.Errorf("cluster name and API server required")
|
||||
}
|
||||
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
|
||||
if _, exists := cm.clusters[name]; exists {
|
||||
return fmt.Errorf("cluster already registered: %s", name)
|
||||
}
|
||||
|
||||
cm.clusters[name] = &ClusterInfo{
|
||||
Name: name,
|
||||
APIServer: apiServer,
|
||||
Healthy: true,
|
||||
LastCheck: time.Now(),
|
||||
Capacity: capacity,
|
||||
Usage: 0,
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// UnregisterCluster removes a cluster
|
||||
func (cm *ClusterManager) UnregisterCluster(name string) error {
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
|
||||
if _, exists := cm.clusters[name]; !exists {
|
||||
return fmt.Errorf("cluster not found: %s", name)
|
||||
}
|
||||
|
||||
delete(cm.clusters, name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetCluster retrieves cluster info
|
||||
func (cm *ClusterManager) GetCluster(name string) (*ClusterInfo, bool) {
|
||||
cm.mu.RLock()
|
||||
defer cm.mu.RUnlock()
|
||||
|
||||
cluster, exists := cm.clusters[name]
|
||||
return cluster, exists
|
||||
}
|
||||
|
||||
// ListClusters returns all registered clusters
|
||||
func (cm *ClusterManager) ListClusters() map[string]*ClusterInfo {
|
||||
cm.mu.RLock()
|
||||
defer cm.mu.RUnlock()
|
||||
|
||||
result := make(map[string]*ClusterInfo)
|
||||
for name, cluster := range cm.clusters {
|
||||
result[name] = cluster
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// HealthCheck checks cluster health
|
||||
func (cm *ClusterManager) HealthCheck(name string) error {
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
|
||||
cluster, exists := cm.clusters[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("cluster not found: %s", name)
|
||||
}
|
||||
|
||||
// Simulate health check (in production, query API server)
|
||||
cluster.Healthy = true
|
||||
cluster.LastCheck = time.Now()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarkUnhealthy marks a cluster as unhealthy
|
||||
func (cm *ClusterManager) MarkUnhealthy(name string) error {
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
|
||||
cluster, exists := cm.clusters[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("cluster not found: %s", name)
|
||||
}
|
||||
|
||||
cluster.Healthy = false
|
||||
return nil
|
||||
}
|
||||
|
||||
// AllocateTask allocates a task to a cluster
|
||||
func (cm *ClusterManager) AllocateTask(name string) error {
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
|
||||
cluster, exists := cm.clusters[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("cluster not found: %s", name)
|
||||
}
|
||||
|
||||
if !cluster.Healthy {
|
||||
return fmt.Errorf("cluster not healthy: %s", name)
|
||||
}
|
||||
|
||||
if cluster.Usage >= cluster.Capacity {
|
||||
return fmt.Errorf("cluster at capacity: %s", name)
|
||||
}
|
||||
|
||||
cluster.Usage++
|
||||
return nil
|
||||
}
|
||||
|
||||
// ReleaseTask releases a task from a cluster
|
||||
func (cm *ClusterManager) ReleaseTask(name string) error {
|
||||
cm.mu.Lock()
|
||||
defer cm.mu.Unlock()
|
||||
|
||||
cluster, exists := cm.clusters[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("cluster not found: %s", name)
|
||||
}
|
||||
|
||||
if cluster.Usage > 0 {
|
||||
cluster.Usage--
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// FindBestCluster finds the cluster with most available capacity
|
||||
func (cm *ClusterManager) FindBestCluster() (string, error) {
|
||||
cm.mu.RLock()
|
||||
defer cm.mu.RUnlock()
|
||||
|
||||
var bestCluster string
|
||||
maxCapacity := 0
|
||||
|
||||
for name, cluster := range cm.clusters {
|
||||
if cluster.Healthy {
|
||||
available := cluster.Capacity - cluster.Usage
|
||||
if available > maxCapacity {
|
||||
bestCluster = name
|
||||
maxCapacity = available
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if bestCluster == "" {
|
||||
return "", fmt.Errorf("no healthy clusters available")
|
||||
}
|
||||
|
||||
return bestCluster, nil
|
||||
}
|
||||
|
||||
// GetCapacitySummary returns capacity summary
|
||||
func (cm *ClusterManager) GetCapacitySummary() map[string]interface{} {
|
||||
cm.mu.RLock()
|
||||
defer cm.mu.RUnlock()
|
||||
|
||||
totalCapacity := 0
|
||||
totalUsage := 0
|
||||
healthyCount := 0
|
||||
|
||||
for _, cluster := range cm.clusters {
|
||||
totalCapacity += cluster.Capacity
|
||||
totalUsage += cluster.Usage
|
||||
if cluster.Healthy {
|
||||
healthyCount++
|
||||
}
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"total_capacity": totalCapacity,
|
||||
"total_usage": totalUsage,
|
||||
"healthy_clusters": healthyCount,
|
||||
"total_clusters": len(cm.clusters),
|
||||
}
|
||||
}
|
||||
|
||||
// GetHealthStatus returns health status for all clusters
|
||||
func (cm *ClusterManager) GetHealthStatus() map[string]bool {
|
||||
cm.mu.RLock()
|
||||
defer cm.mu.RUnlock()
|
||||
|
||||
result := make(map[string]bool)
|
||||
for name, cluster := range cm.clusters {
|
||||
result[name] = cluster.Healthy
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package clusters
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestRegisterCluster(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
|
||||
err := cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
assert.NoError(t, err)
|
||||
|
||||
cluster, exists := cm.GetCluster("prod")
|
||||
assert.True(t, exists)
|
||||
assert.Equal(t, "prod", cluster.Name)
|
||||
}
|
||||
|
||||
func TestUnregisterCluster(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
|
||||
err := cm.UnregisterCluster("prod")
|
||||
assert.NoError(t, err)
|
||||
|
||||
_, exists := cm.GetCluster("prod")
|
||||
assert.False(t, exists)
|
||||
}
|
||||
|
||||
func TestListClusters(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
cm.RegisterCluster("staging", "https://k8s-staging.com", 50)
|
||||
|
||||
clusters := cm.ListClusters()
|
||||
assert.Equal(t, 2, len(clusters))
|
||||
}
|
||||
|
||||
func TestHealthCheck(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
|
||||
err := cm.HealthCheck("prod")
|
||||
assert.NoError(t, err)
|
||||
|
||||
cluster, _ := cm.GetCluster("prod")
|
||||
assert.True(t, cluster.Healthy)
|
||||
}
|
||||
|
||||
func TestMarkUnhealthy(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
|
||||
cm.MarkUnhealthy("prod")
|
||||
|
||||
cluster, _ := cm.GetCluster("prod")
|
||||
assert.False(t, cluster.Healthy)
|
||||
}
|
||||
|
||||
func TestAllocateTask(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
|
||||
err := cm.AllocateTask("prod")
|
||||
assert.NoError(t, err)
|
||||
|
||||
cluster, _ := cm.GetCluster("prod")
|
||||
assert.Equal(t, 1, cluster.Usage)
|
||||
}
|
||||
|
||||
func TestAllocateTaskUnhealthy(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
cm.MarkUnhealthy("prod")
|
||||
|
||||
err := cm.AllocateTask("prod")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestAllocateTaskAtCapacity(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 1)
|
||||
|
||||
cm.AllocateTask("prod")
|
||||
err := cm.AllocateTask("prod")
|
||||
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestReleaseTask(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
|
||||
cm.AllocateTask("prod")
|
||||
cm.ReleaseTask("prod")
|
||||
|
||||
cluster, _ := cm.GetCluster("prod")
|
||||
assert.Equal(t, 0, cluster.Usage)
|
||||
}
|
||||
|
||||
func TestFindBestCluster(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
cm.RegisterCluster("staging", "https://k8s-staging.com", 50)
|
||||
|
||||
cm.AllocateTask("staging")
|
||||
cm.AllocateTask("staging")
|
||||
|
||||
best, err := cm.FindBestCluster()
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "prod", best)
|
||||
}
|
||||
|
||||
func TestGetCapacitySummary(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
cm.RegisterCluster("staging", "https://k8s-staging.com", 50)
|
||||
|
||||
cm.AllocateTask("prod")
|
||||
|
||||
summary := cm.GetCapacitySummary()
|
||||
assert.Equal(t, 150, summary["total_capacity"])
|
||||
assert.Equal(t, 1, summary["total_usage"])
|
||||
assert.Equal(t, 2, summary["healthy_clusters"])
|
||||
}
|
||||
|
||||
func TestGetHealthStatus(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
cm.RegisterCluster("staging", "https://k8s-staging.com", 50)
|
||||
|
||||
cm.MarkUnhealthy("staging")
|
||||
|
||||
status := cm.GetHealthStatus()
|
||||
assert.True(t, status["prod"])
|
||||
assert.False(t, status["staging"])
|
||||
}
|
||||
|
||||
func TestRegisterClusterError(t *testing.T) {
|
||||
cm := NewClusterManager()
|
||||
cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
|
||||
err := cm.RegisterCluster("prod", "https://k8s-prod.com", 100)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package composition
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// ChildOrchestrator represents a child orchestrator workflow
|
||||
type ChildOrchestrator struct {
|
||||
ID string
|
||||
ParentTask string
|
||||
Config map[string]interface{}
|
||||
Status string
|
||||
Results map[string]interface{}
|
||||
CreatedAt int64
|
||||
}
|
||||
|
||||
// WorkflowComposer manages nested orchestrator workflows
|
||||
type WorkflowComposer struct {
|
||||
mu sync.RWMutex
|
||||
children map[string]*ChildOrchestrator
|
||||
results map[string]map[string]interface{}
|
||||
}
|
||||
|
||||
// NewWorkflowComposer creates a new workflow composer
|
||||
func NewWorkflowComposer() *WorkflowComposer {
|
||||
return &WorkflowComposer{
|
||||
children: make(map[string]*ChildOrchestrator),
|
||||
results: make(map[string]map[string]interface{}),
|
||||
}
|
||||
}
|
||||
|
||||
// CreateChild creates a child orchestrator
|
||||
func (wc *WorkflowComposer) CreateChild(parentTask string, config map[string]interface{}) (*ChildOrchestrator, error) {
|
||||
if parentTask == "" {
|
||||
return nil, fmt.Errorf("parent task required")
|
||||
}
|
||||
|
||||
wc.mu.Lock()
|
||||
defer wc.mu.Unlock()
|
||||
|
||||
child := &ChildOrchestrator{
|
||||
ID: fmt.Sprintf("child-%s-%d", parentTask, len(wc.children)),
|
||||
ParentTask: parentTask,
|
||||
Config: config,
|
||||
Status: "pending",
|
||||
Results: make(map[string]interface{}),
|
||||
}
|
||||
|
||||
wc.children[child.ID] = child
|
||||
return child, nil
|
||||
}
|
||||
|
||||
// GetChild retrieves a child orchestrator
|
||||
func (wc *WorkflowComposer) GetChild(id string) (*ChildOrchestrator, bool) {
|
||||
wc.mu.RLock()
|
||||
defer wc.mu.RUnlock()
|
||||
|
||||
child, exists := wc.children[id]
|
||||
return child, exists
|
||||
}
|
||||
|
||||
// ListChildren lists all children
|
||||
func (wc *WorkflowComposer) ListChildren() map[string]*ChildOrchestrator {
|
||||
wc.mu.RLock()
|
||||
defer wc.mu.RUnlock()
|
||||
|
||||
result := make(map[string]*ChildOrchestrator)
|
||||
for id, child := range wc.children {
|
||||
result[id] = child
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// SetChildStatus updates child status
|
||||
func (wc *WorkflowComposer) SetChildStatus(id string, status string) error {
|
||||
wc.mu.Lock()
|
||||
defer wc.mu.Unlock()
|
||||
|
||||
child, exists := wc.children[id]
|
||||
if !exists {
|
||||
return fmt.Errorf("child not found: %s", id)
|
||||
}
|
||||
|
||||
child.Status = status
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetHierarchy returns the workflow hierarchy
|
||||
func (wc *WorkflowComposer) GetHierarchy() map[string]interface{} {
|
||||
wc.mu.RLock()
|
||||
defer wc.mu.RUnlock()
|
||||
|
||||
children := make([]map[string]interface{}, 0)
|
||||
for _, child := range wc.children {
|
||||
children = append(children, map[string]interface{}{
|
||||
"id": child.ID,
|
||||
"parent_task": child.ParentTask,
|
||||
"status": child.Status,
|
||||
})
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"children": children,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package composition
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestCreateChild(t *testing.T) {
|
||||
composer := NewWorkflowComposer()
|
||||
config := map[string]interface{}{"tasks": 5}
|
||||
|
||||
child, err := composer.CreateChild("T0.1", config)
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, child)
|
||||
assert.Equal(t, "T0.1", child.ParentTask)
|
||||
}
|
||||
|
||||
func TestGetChild(t *testing.T) {
|
||||
composer := NewWorkflowComposer()
|
||||
|
||||
child, _ := composer.CreateChild("T0.1", map[string]interface{}{})
|
||||
|
||||
retrieved, found := composer.GetChild(child.ID)
|
||||
assert.True(t, found)
|
||||
assert.Equal(t, child.ID, retrieved.ID)
|
||||
}
|
||||
|
||||
func TestListChildren(t *testing.T) {
|
||||
composer := NewWorkflowComposer()
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
composer.CreateChild("T0.1", map[string]interface{}{})
|
||||
}
|
||||
|
||||
children := composer.ListChildren()
|
||||
assert.Equal(t, 3, len(children))
|
||||
}
|
||||
|
||||
func TestSetChildStatus(t *testing.T) {
|
||||
composer := NewWorkflowComposer()
|
||||
|
||||
child, _ := composer.CreateChild("T0.1", map[string]interface{}{})
|
||||
|
||||
err := composer.SetChildStatus(child.ID, "completed")
|
||||
assert.NoError(t, err)
|
||||
|
||||
updated, _ := composer.GetChild(child.ID)
|
||||
assert.Equal(t, "completed", updated.Status)
|
||||
}
|
||||
@@ -0,0 +1,246 @@
|
||||
package cost
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// CostEntry represents a tracked cost
|
||||
type CostEntry struct {
|
||||
ID string
|
||||
Type string // llm, git, compute
|
||||
WorkflowID string
|
||||
TaskID string
|
||||
Amount float64
|
||||
Timestamp time.Time
|
||||
Metadata map[string]interface{}
|
||||
}
|
||||
|
||||
// CostTracker tracks and analyzes workflow costs
|
||||
type CostTracker struct {
|
||||
mu sync.RWMutex
|
||||
entries []*CostEntry
|
||||
rates map[string]float64
|
||||
}
|
||||
|
||||
// NewCostTracker creates a new cost tracker
|
||||
func NewCostTracker() *CostTracker {
|
||||
return &CostTracker{
|
||||
entries: make([]*CostEntry, 0),
|
||||
rates: map[string]float64{
|
||||
"llm_token": 0.0001, // $0.0001 per token
|
||||
"git_push": 0.0, // Free
|
||||
"compute_hour": 0.5, // $0.5 per hour
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// TrackLLMCost tracks LLM API costs
|
||||
func (ct *CostTracker) TrackLLMCost(workflowID, taskID string, tokens int) {
|
||||
cost := float64(tokens) * ct.rates["llm_token"]
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
|
||||
entry := &CostEntry{
|
||||
ID: fmt.Sprintf("llm-%d", len(ct.entries)),
|
||||
Type: "llm",
|
||||
WorkflowID: workflowID,
|
||||
TaskID: taskID,
|
||||
Amount: cost,
|
||||
Timestamp: time.Now(),
|
||||
Metadata: map[string]interface{}{
|
||||
"tokens": tokens,
|
||||
},
|
||||
}
|
||||
|
||||
ct.entries = append(ct.entries, entry)
|
||||
}
|
||||
|
||||
// TrackGitCost tracks git operation costs
|
||||
func (ct *CostTracker) TrackGitCost(workflowID string, operations int) {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
|
||||
entry := &CostEntry{
|
||||
ID: fmt.Sprintf("git-%d", len(ct.entries)),
|
||||
Type: "git",
|
||||
WorkflowID: workflowID,
|
||||
Amount: 0,
|
||||
Timestamp: time.Now(),
|
||||
Metadata: map[string]interface{}{
|
||||
"operations": operations,
|
||||
},
|
||||
}
|
||||
|
||||
ct.entries = append(ct.entries, entry)
|
||||
}
|
||||
|
||||
// TrackComputeCost tracks compute resource costs (duration in milliseconds)
|
||||
func (ct *CostTracker) TrackComputeCost(workflowID, taskID string, durationMs float64) {
|
||||
// Convert milliseconds to hours
|
||||
durationHours := durationMs / (1000.0 * 3600.0)
|
||||
cost := durationHours * ct.rates["compute_hour"]
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
|
||||
entry := &CostEntry{
|
||||
ID: fmt.Sprintf("compute-%d", len(ct.entries)),
|
||||
Type: "compute",
|
||||
WorkflowID: workflowID,
|
||||
TaskID: taskID,
|
||||
Amount: cost,
|
||||
Timestamp: time.Now(),
|
||||
Metadata: map[string]interface{}{
|
||||
"duration_ms": durationMs,
|
||||
},
|
||||
}
|
||||
|
||||
ct.entries = append(ct.entries, entry)
|
||||
}
|
||||
|
||||
// GetTotalCost returns total cost for all workflows
|
||||
func (ct *CostTracker) GetTotalCost() float64 {
|
||||
ct.mu.RLock()
|
||||
defer ct.mu.RUnlock()
|
||||
|
||||
total := 0.0
|
||||
for _, entry := range ct.entries {
|
||||
total += entry.Amount
|
||||
}
|
||||
|
||||
return total
|
||||
}
|
||||
|
||||
// GetWorkflowCost returns total cost for a specific workflow
|
||||
func (ct *CostTracker) GetWorkflowCost(workflowID string) float64 {
|
||||
ct.mu.RLock()
|
||||
defer ct.mu.RUnlock()
|
||||
|
||||
total := 0.0
|
||||
for _, entry := range ct.entries {
|
||||
if entry.WorkflowID == workflowID {
|
||||
total += entry.Amount
|
||||
}
|
||||
}
|
||||
|
||||
return total
|
||||
}
|
||||
|
||||
// GetCostByType returns total cost by type
|
||||
func (ct *CostTracker) GetCostByType(costType string) float64 {
|
||||
ct.mu.RLock()
|
||||
defer ct.mu.RUnlock()
|
||||
|
||||
total := 0.0
|
||||
for _, entry := range ct.entries {
|
||||
if entry.Type == costType {
|
||||
total += entry.Amount
|
||||
}
|
||||
}
|
||||
|
||||
return total
|
||||
}
|
||||
|
||||
// GetAverageCostPerTask returns average cost per task
|
||||
func (ct *CostTracker) GetAverageCostPerTask(workflowID string) float64 {
|
||||
ct.mu.RLock()
|
||||
defer ct.mu.RUnlock()
|
||||
|
||||
total := 0.0
|
||||
count := 0
|
||||
|
||||
for _, entry := range ct.entries {
|
||||
if entry.WorkflowID == workflowID {
|
||||
total += entry.Amount
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
if count == 0 {
|
||||
return 0
|
||||
}
|
||||
|
||||
return total / float64(count)
|
||||
}
|
||||
|
||||
// GetOptimizationSuggestions returns cost optimization recommendations
|
||||
func (ct *CostTracker) GetOptimizationSuggestions(workflowID string) []string {
|
||||
suggestions := make([]string, 0)
|
||||
|
||||
ct.mu.RLock()
|
||||
defer ct.mu.RUnlock()
|
||||
|
||||
llmCost := 0.0
|
||||
computeCost := 0.0
|
||||
|
||||
for _, entry := range ct.entries {
|
||||
if entry.WorkflowID == workflowID {
|
||||
if entry.Type == "llm" {
|
||||
llmCost += entry.Amount
|
||||
} else if entry.Type == "compute" {
|
||||
computeCost += entry.Amount
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if llmCost > computeCost*2 {
|
||||
suggestions = append(suggestions, "Consider caching LLM results to reduce API calls")
|
||||
}
|
||||
|
||||
if computeCost > llmCost*2 {
|
||||
suggestions = append(suggestions, "Consider parallelizing compute tasks")
|
||||
}
|
||||
|
||||
return suggestions
|
||||
}
|
||||
|
||||
// GetEntries returns all cost entries
|
||||
func (ct *CostTracker) GetEntries() []*CostEntry {
|
||||
ct.mu.RLock()
|
||||
defer ct.mu.RUnlock()
|
||||
|
||||
result := make([]*CostEntry, len(ct.entries))
|
||||
copy(result, ct.entries)
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetEntriesForWorkflow returns cost entries for a workflow
|
||||
func (ct *CostTracker) GetEntriesForWorkflow(workflowID string) []*CostEntry {
|
||||
ct.mu.RLock()
|
||||
defer ct.mu.RUnlock()
|
||||
|
||||
result := make([]*CostEntry, 0)
|
||||
for _, entry := range ct.entries {
|
||||
if entry.WorkflowID == workflowID {
|
||||
result = append(result, entry)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// SetRate sets the cost rate for a type
|
||||
func (ct *CostTracker) SetRate(costType string, rate float64) {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
|
||||
ct.rates[costType] = rate
|
||||
}
|
||||
|
||||
// GetRate gets the cost rate for a type
|
||||
func (ct *CostTracker) GetRate(costType string) float64 {
|
||||
ct.mu.RLock()
|
||||
defer ct.mu.RUnlock()
|
||||
|
||||
return ct.rates[costType]
|
||||
}
|
||||
|
||||
// Clear clears all cost entries
|
||||
func (ct *CostTracker) Clear() {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
|
||||
ct.entries = make([]*CostEntry, 0)
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
package cost
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestTrackLLMCost(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 1000)
|
||||
|
||||
entries := tracker.GetEntries()
|
||||
assert.Equal(t, 1, len(entries))
|
||||
assert.Equal(t, "llm", entries[0].Type)
|
||||
assert.Equal(t, 0.1, entries[0].Amount) // 1000 tokens * 0.0001
|
||||
}
|
||||
|
||||
func TestTrackGitCost(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
tracker.TrackGitCost("wf-1", 5)
|
||||
|
||||
entries := tracker.GetEntries()
|
||||
assert.Equal(t, 1, len(entries))
|
||||
assert.Equal(t, "git", entries[0].Type)
|
||||
}
|
||||
|
||||
func TestTrackComputeCost(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
tracker.TrackComputeCost("wf-1", "task-1", 3600000) // 1 hour in ms
|
||||
|
||||
entries := tracker.GetEntries()
|
||||
assert.Equal(t, 1, len(entries))
|
||||
assert.Equal(t, "compute", entries[0].Type)
|
||||
assert.Equal(t, 0.5, entries[0].Amount) // 1 hour * $0.5/hour
|
||||
}
|
||||
|
||||
func TestGetTotalCost(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 1000)
|
||||
tracker.TrackComputeCost("wf-1", "task-1", 3600000)
|
||||
|
||||
total := tracker.GetTotalCost()
|
||||
assert.Equal(t, 0.6, total) // 0.1 + 0.5
|
||||
}
|
||||
|
||||
func TestGetWorkflowCost(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 1000)
|
||||
tracker.TrackLLMCost("wf-2", "task-1", 2000)
|
||||
|
||||
cost := tracker.GetWorkflowCost("wf-1")
|
||||
assert.Equal(t, 0.1, cost)
|
||||
|
||||
cost = tracker.GetWorkflowCost("wf-2")
|
||||
assert.Equal(t, 0.2, cost)
|
||||
}
|
||||
|
||||
func TestGetCostByType(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 1000)
|
||||
tracker.TrackLLMCost("wf-1", "task-2", 1000)
|
||||
tracker.TrackComputeCost("wf-1", "task-3", 3600000)
|
||||
|
||||
llmCost := tracker.GetCostByType("llm")
|
||||
assert.Equal(t, 0.2, llmCost)
|
||||
|
||||
computeCost := tracker.GetCostByType("compute")
|
||||
assert.Equal(t, 0.5, computeCost)
|
||||
}
|
||||
|
||||
func TestGetAverageCostPerTask(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 1000)
|
||||
tracker.TrackLLMCost("wf-1", "task-2", 1000)
|
||||
|
||||
avg := tracker.GetAverageCostPerTask("wf-1")
|
||||
assert.Equal(t, 0.1, avg)
|
||||
}
|
||||
|
||||
func TestGetOptimizationSuggestions(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
// High LLM cost
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 10000)
|
||||
tracker.TrackLLMCost("wf-1", "task-2", 10000)
|
||||
tracker.TrackComputeCost("wf-1", "task-3", 360000) // 0.1 seconds
|
||||
|
||||
suggestions := tracker.GetOptimizationSuggestions("wf-1")
|
||||
// Just verify it returns without error - suggestions depend on cost ratios
|
||||
assert.NotNil(t, suggestions)
|
||||
}
|
||||
|
||||
func TestGetEntriesForWorkflow(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 1000)
|
||||
tracker.TrackLLMCost("wf-2", "task-1", 1000)
|
||||
|
||||
entries := tracker.GetEntriesForWorkflow("wf-1")
|
||||
assert.Equal(t, 1, len(entries))
|
||||
}
|
||||
|
||||
func TestSetAndGetRate(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
tracker.SetRate("custom", 0.5)
|
||||
rate := tracker.GetRate("custom")
|
||||
|
||||
assert.Equal(t, 0.5, rate)
|
||||
}
|
||||
|
||||
func TestClear(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 1000)
|
||||
tracker.Clear()
|
||||
|
||||
entries := tracker.GetEntries()
|
||||
assert.Equal(t, 0, len(entries))
|
||||
}
|
||||
|
||||
func TestMultipleCosts(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 1000)
|
||||
tracker.TrackGitCost("wf-1", 5)
|
||||
tracker.TrackComputeCost("wf-1", "task-1", 1800000) // 30 min
|
||||
|
||||
total := tracker.GetTotalCost()
|
||||
assert.True(t, total > 0.2)
|
||||
}
|
||||
|
||||
func TestZeroCost(t *testing.T) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
cost := tracker.GetWorkflowCost("nonexistent")
|
||||
assert.Equal(t, 0.0, cost)
|
||||
}
|
||||
|
||||
func BenchmarkTrackLLMCost(b *testing.B) {
|
||||
tracker := NewCostTracker()
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
tracker.TrackLLMCost("wf-1", "task-1", 1000)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,203 @@
|
||||
package dashboard
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// MetricSnapshot represents a point-in-time metric value
|
||||
type MetricSnapshot struct {
|
||||
Timestamp time.Time
|
||||
Value float64
|
||||
Name string
|
||||
}
|
||||
|
||||
// MetricsAggregator aggregates Prometheus metrics for dashboard display
|
||||
type MetricsAggregator struct {
|
||||
mu sync.RWMutex
|
||||
metrics map[string][]MetricSnapshot
|
||||
ttl time.Duration
|
||||
maxSize int
|
||||
}
|
||||
|
||||
// NewMetricsAggregator creates a new metrics aggregator
|
||||
func NewMetricsAggregator(ttl time.Duration, maxSize int) *MetricsAggregator {
|
||||
return &MetricsAggregator{
|
||||
metrics: make(map[string][]MetricSnapshot),
|
||||
ttl: ttl,
|
||||
maxSize: maxSize,
|
||||
}
|
||||
}
|
||||
|
||||
// Record records a metric value
|
||||
func (ma *MetricsAggregator) Record(name string, value float64) {
|
||||
ma.mu.Lock()
|
||||
defer ma.mu.Unlock()
|
||||
|
||||
snapshot := MetricSnapshot{
|
||||
Timestamp: time.Now(),
|
||||
Value: value,
|
||||
Name: name,
|
||||
}
|
||||
|
||||
ma.metrics[name] = append(ma.metrics[name], snapshot)
|
||||
|
||||
// Trim old entries
|
||||
if len(ma.metrics[name]) > ma.maxSize {
|
||||
ma.metrics[name] = ma.metrics[name][1:]
|
||||
}
|
||||
}
|
||||
|
||||
// GetTimeSeries retrieves metric time series
|
||||
func (ma *MetricsAggregator) GetTimeSeries(name string) []MetricSnapshot {
|
||||
ma.mu.RLock()
|
||||
defer ma.mu.RUnlock()
|
||||
|
||||
snapshots, exists := ma.metrics[name]
|
||||
if !exists {
|
||||
return []MetricSnapshot{}
|
||||
}
|
||||
|
||||
result := make([]MetricSnapshot, len(snapshots))
|
||||
copy(result, snapshots)
|
||||
return result
|
||||
}
|
||||
|
||||
// GetPercentile calculates percentile for a metric
|
||||
func (ma *MetricsAggregator) GetPercentile(name string, percentile float64) (float64, error) {
|
||||
ma.mu.RLock()
|
||||
defer ma.mu.RUnlock()
|
||||
|
||||
snapshots, exists := ma.metrics[name]
|
||||
if !exists || len(snapshots) == 0 {
|
||||
return 0, fmt.Errorf("metric not found: %s", name)
|
||||
}
|
||||
|
||||
values := make([]float64, len(snapshots))
|
||||
for i, s := range snapshots {
|
||||
values[i] = s.Value
|
||||
}
|
||||
|
||||
sort.Float64s(values)
|
||||
|
||||
index := int(float64(len(values)) * percentile / 100)
|
||||
if index >= len(values) {
|
||||
index = len(values) - 1
|
||||
}
|
||||
|
||||
return values[index], nil
|
||||
}
|
||||
|
||||
// GetAverage calculates average for a metric
|
||||
func (ma *MetricsAggregator) GetAverage(name string) (float64, error) {
|
||||
ma.mu.RLock()
|
||||
defer ma.mu.RUnlock()
|
||||
|
||||
snapshots, exists := ma.metrics[name]
|
||||
if !exists || len(snapshots) == 0 {
|
||||
return 0, fmt.Errorf("metric not found: %s", name)
|
||||
}
|
||||
|
||||
sum := 0.0
|
||||
for _, s := range snapshots {
|
||||
sum += s.Value
|
||||
}
|
||||
|
||||
return sum / float64(len(snapshots)), nil
|
||||
}
|
||||
|
||||
// GetMax returns maximum value for a metric
|
||||
func (ma *MetricsAggregator) GetMax(name string) (float64, error) {
|
||||
ma.mu.RLock()
|
||||
defer ma.mu.RUnlock()
|
||||
|
||||
snapshots, exists := ma.metrics[name]
|
||||
if !exists || len(snapshots) == 0 {
|
||||
return 0, fmt.Errorf("metric not found: %s", name)
|
||||
}
|
||||
|
||||
max := snapshots[0].Value
|
||||
for _, s := range snapshots {
|
||||
if s.Value > max {
|
||||
max = s.Value
|
||||
}
|
||||
}
|
||||
|
||||
return max, nil
|
||||
}
|
||||
|
||||
// GetMin returns minimum value for a metric
|
||||
func (ma *MetricsAggregator) GetMin(name string) (float64, error) {
|
||||
ma.mu.RLock()
|
||||
defer ma.mu.RUnlock()
|
||||
|
||||
snapshots, exists := ma.metrics[name]
|
||||
if !exists || len(snapshots) == 0 {
|
||||
return 0, fmt.Errorf("metric not found: %s", name)
|
||||
}
|
||||
|
||||
min := snapshots[0].Value
|
||||
for _, s := range snapshots {
|
||||
if s.Value < min {
|
||||
min = s.Value
|
||||
}
|
||||
}
|
||||
|
||||
return min, nil
|
||||
}
|
||||
|
||||
// GetMetricNames returns all recorded metric names
|
||||
func (ma *MetricsAggregator) GetMetricNames() []string {
|
||||
ma.mu.RLock()
|
||||
defer ma.mu.RUnlock()
|
||||
|
||||
names := make([]string, 0, len(ma.metrics))
|
||||
for name := range ma.metrics {
|
||||
names = append(names, name)
|
||||
}
|
||||
|
||||
return names
|
||||
}
|
||||
|
||||
// GetLatest returns the latest snapshot for a metric
|
||||
func (ma *MetricsAggregator) GetLatest(name string) (MetricSnapshot, error) {
|
||||
ma.mu.RLock()
|
||||
defer ma.mu.RUnlock()
|
||||
|
||||
snapshots, exists := ma.metrics[name]
|
||||
if !exists || len(snapshots) == 0 {
|
||||
return MetricSnapshot{}, fmt.Errorf("metric not found: %s", name)
|
||||
}
|
||||
|
||||
return snapshots[len(snapshots)-1], nil
|
||||
}
|
||||
|
||||
// Clear clears all metrics
|
||||
func (ma *MetricsAggregator) Clear() {
|
||||
ma.mu.Lock()
|
||||
defer ma.mu.Unlock()
|
||||
|
||||
ma.metrics = make(map[string][]MetricSnapshot)
|
||||
}
|
||||
|
||||
// GetCountInRange returns count of metrics within a time range
|
||||
func (ma *MetricsAggregator) GetCountInRange(name string, start, end time.Time) (int, error) {
|
||||
ma.mu.RLock()
|
||||
defer ma.mu.RUnlock()
|
||||
|
||||
snapshots, exists := ma.metrics[name]
|
||||
if !exists {
|
||||
return 0, fmt.Errorf("metric not found: %s", name)
|
||||
}
|
||||
|
||||
count := 0
|
||||
for _, s := range snapshots {
|
||||
if s.Timestamp.After(start) && s.Timestamp.Before(end) {
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
package dashboard
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestRecord(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
agg.Record("request_latency", 150.5)
|
||||
|
||||
names := agg.GetMetricNames()
|
||||
assert.Equal(t, 1, len(names))
|
||||
assert.Equal(t, "request_latency", names[0])
|
||||
}
|
||||
|
||||
func TestGetTimeSeries(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
agg.Record("latency", 100)
|
||||
agg.Record("latency", 200)
|
||||
agg.Record("latency", 150)
|
||||
|
||||
series := agg.GetTimeSeries("latency")
|
||||
assert.Equal(t, 3, len(series))
|
||||
assert.Equal(t, 100.0, series[0].Value)
|
||||
assert.Equal(t, 200.0, series[1].Value)
|
||||
assert.Equal(t, 150.0, series[2].Value)
|
||||
}
|
||||
|
||||
func TestGetPercentile(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
for i := 1; i <= 100; i++ {
|
||||
agg.Record("latency", float64(i))
|
||||
}
|
||||
|
||||
p50, _ := agg.GetPercentile("latency", 50)
|
||||
p95, _ := agg.GetPercentile("latency", 95)
|
||||
p99, _ := agg.GetPercentile("latency", 99)
|
||||
|
||||
assert.True(t, p50 > 40 && p50 < 60)
|
||||
assert.True(t, p95 > 90)
|
||||
assert.True(t, p99 > 95)
|
||||
}
|
||||
|
||||
func TestGetAverage(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
agg.Record("latency", 100)
|
||||
agg.Record("latency", 200)
|
||||
agg.Record("latency", 300)
|
||||
|
||||
avg, _ := agg.GetAverage("latency")
|
||||
assert.Equal(t, 200.0, avg)
|
||||
}
|
||||
|
||||
func TestGetMax(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
agg.Record("latency", 100)
|
||||
agg.Record("latency", 500)
|
||||
agg.Record("latency", 300)
|
||||
|
||||
max, _ := agg.GetMax("latency")
|
||||
assert.Equal(t, 500.0, max)
|
||||
}
|
||||
|
||||
func TestGetMin(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
agg.Record("latency", 100)
|
||||
agg.Record("latency", 500)
|
||||
agg.Record("latency", 50)
|
||||
|
||||
min, _ := agg.GetMin("latency")
|
||||
assert.Equal(t, 50.0, min)
|
||||
}
|
||||
|
||||
func TestGetLatest(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
agg.Record("latency", 100)
|
||||
agg.Record("latency", 200)
|
||||
|
||||
latest, _ := agg.GetLatest("latency")
|
||||
assert.Equal(t, 200.0, latest.Value)
|
||||
}
|
||||
|
||||
func TestMultipleMetrics(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
agg.Record("latency", 100)
|
||||
agg.Record("errors", 5)
|
||||
agg.Record("throughput", 1000)
|
||||
|
||||
names := agg.GetMetricNames()
|
||||
assert.Equal(t, 3, len(names))
|
||||
}
|
||||
|
||||
func TestClear(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
agg.Record("latency", 100)
|
||||
agg.Clear()
|
||||
|
||||
names := agg.GetMetricNames()
|
||||
assert.Equal(t, 0, len(names))
|
||||
}
|
||||
|
||||
func TestGetCountInRange(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
now := time.Now()
|
||||
agg.Record("latency", 100)
|
||||
agg.Record("latency", 200)
|
||||
|
||||
count, _ := agg.GetCountInRange("latency", now.Add(-1*time.Minute), now.Add(1*time.Minute))
|
||||
assert.Equal(t, 2, count)
|
||||
}
|
||||
|
||||
func TestNotFoundError(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 100)
|
||||
|
||||
_, err := agg.GetPercentile("nonexistent", 50)
|
||||
assert.Error(t, err)
|
||||
|
||||
_, err = agg.GetAverage("nonexistent")
|
||||
assert.Error(t, err)
|
||||
|
||||
_, err = agg.GetLatest("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestMaxSize(t *testing.T) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 5)
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
agg.Record("latency", float64(i))
|
||||
}
|
||||
|
||||
series := agg.GetTimeSeries("latency")
|
||||
assert.Equal(t, 5, len(series))
|
||||
}
|
||||
|
||||
func BenchmarkRecord(b *testing.B) {
|
||||
agg := NewMetricsAggregator(1*time.Hour, 1000)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
agg.Record("latency", float64(i))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,235 @@
|
||||
package deployment
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// DeploymentStatus represents deployment status
|
||||
type DeploymentStatus string
|
||||
|
||||
const (
|
||||
StatusPending DeploymentStatus = "pending"
|
||||
StatusBuilding DeploymentStatus = "building"
|
||||
StatusPushing DeploymentStatus = "pushing"
|
||||
StatusApplying DeploymentStatus = "applying"
|
||||
StatusSuccess DeploymentStatus = "success"
|
||||
StatusFailed DeploymentStatus = "failed"
|
||||
)
|
||||
|
||||
// DeploymentInfo represents a deployment attempt
|
||||
type DeploymentInfo struct {
|
||||
ID string
|
||||
Version string
|
||||
Status DeploymentStatus
|
||||
StartedAt time.Time
|
||||
CompletedAt time.Time
|
||||
Container string
|
||||
Registry string
|
||||
Manifest string
|
||||
}
|
||||
|
||||
// SelfDeployer handles orchestrator self-deployment
|
||||
type SelfDeployer struct {
|
||||
mu sync.RWMutex
|
||||
deployments map[string]*DeploymentInfo
|
||||
currentVersion string
|
||||
registry string
|
||||
kubeConfig string
|
||||
}
|
||||
|
||||
// NewSelfDeployer creates a new self deployer
|
||||
func NewSelfDeployer(registry, kubeConfig string) *SelfDeployer {
|
||||
return &SelfDeployer{
|
||||
deployments: make(map[string]*DeploymentInfo),
|
||||
currentVersion: "1.0.0",
|
||||
registry: registry,
|
||||
kubeConfig: kubeConfig,
|
||||
}
|
||||
}
|
||||
|
||||
// BuildContainer builds a Docker container image
|
||||
func (sd *SelfDeployer) BuildContainer(version string) (string, error) {
|
||||
if version == "" {
|
||||
return "", fmt.Errorf("version required")
|
||||
}
|
||||
|
||||
sd.mu.Lock()
|
||||
defer sd.mu.Unlock()
|
||||
|
||||
deploymentID := fmt.Sprintf("deploy-%s-%d", version, len(sd.deployments))
|
||||
|
||||
deployment := &DeploymentInfo{
|
||||
ID: deploymentID,
|
||||
Version: version,
|
||||
Status: StatusBuilding,
|
||||
StartedAt: time.Now(),
|
||||
Container: fmt.Sprintf("%s/orchestrator:%s", sd.registry, version),
|
||||
Registry: sd.registry,
|
||||
}
|
||||
|
||||
sd.deployments[deploymentID] = deployment
|
||||
|
||||
// Simulate build
|
||||
deployment.Status = StatusPushing
|
||||
|
||||
return deploymentID, nil
|
||||
}
|
||||
|
||||
// PushImage pushes the container image to registry
|
||||
func (sd *SelfDeployer) PushImage(deploymentID string) error {
|
||||
sd.mu.Lock()
|
||||
defer sd.mu.Unlock()
|
||||
|
||||
deployment, exists := sd.deployments[deploymentID]
|
||||
if !exists {
|
||||
return fmt.Errorf("deployment not found: %s", deploymentID)
|
||||
}
|
||||
|
||||
if deployment.Status != StatusPushing {
|
||||
return fmt.Errorf("invalid status for push: %s", deployment.Status)
|
||||
}
|
||||
|
||||
// Simulate push
|
||||
deployment.Status = StatusApplying
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GenerateManifest generates K8s manifests
|
||||
func (sd *SelfDeployer) GenerateManifest(deploymentID string, replicas int) (string, error) {
|
||||
sd.mu.Lock()
|
||||
defer sd.mu.Unlock()
|
||||
|
||||
deployment, exists := sd.deployments[deploymentID]
|
||||
if !exists {
|
||||
return "", fmt.Errorf("deployment not found: %s", deploymentID)
|
||||
}
|
||||
|
||||
manifest := fmt.Sprintf(`
|
||||
apiVersion: apps/v1
|
||||
kind: Deployment
|
||||
metadata:
|
||||
name: poimen-orchestrator
|
||||
spec:
|
||||
replicas: %d
|
||||
selector:
|
||||
matchLabels:
|
||||
app: poimen-orchestrator
|
||||
template:
|
||||
metadata:
|
||||
labels:
|
||||
app: poimen-orchestrator
|
||||
spec:
|
||||
containers:
|
||||
- name: orchestrator
|
||||
image: %s
|
||||
ports:
|
||||
- containerPort: 7233
|
||||
- containerPort: 8081
|
||||
`, replicas, deployment.Container)
|
||||
|
||||
deployment.Manifest = manifest
|
||||
|
||||
return manifest, nil
|
||||
}
|
||||
|
||||
// Deploy applies the deployment to K8s
|
||||
func (sd *SelfDeployer) Deploy(deploymentID string) error {
|
||||
sd.mu.Lock()
|
||||
defer sd.mu.Unlock()
|
||||
|
||||
deployment, exists := sd.deployments[deploymentID]
|
||||
if !exists {
|
||||
return fmt.Errorf("deployment not found: %s", deploymentID)
|
||||
}
|
||||
|
||||
if deployment.Status != StatusApplying {
|
||||
return fmt.Errorf("invalid status for deploy: %s", deployment.Status)
|
||||
}
|
||||
|
||||
// Simulate deployment
|
||||
deployment.Status = StatusSuccess
|
||||
deployment.CompletedAt = time.Now()
|
||||
|
||||
// Update current version
|
||||
sd.currentVersion = deployment.Version
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Rollback rolls back to previous version
|
||||
func (sd *SelfDeployer) Rollback(previousVersion string) error {
|
||||
sd.mu.Lock()
|
||||
defer sd.mu.Unlock()
|
||||
|
||||
// Create a new deployment for rollback
|
||||
deploymentID := fmt.Sprintf("rollback-%s-%d", previousVersion, len(sd.deployments))
|
||||
|
||||
deployment := &DeploymentInfo{
|
||||
ID: deploymentID,
|
||||
Version: previousVersion,
|
||||
Status: StatusSuccess,
|
||||
StartedAt: time.Now(),
|
||||
CompletedAt: time.Now(),
|
||||
Container: fmt.Sprintf("%s/orchestrator:%s", sd.registry, previousVersion),
|
||||
}
|
||||
|
||||
sd.deployments[deploymentID] = deployment
|
||||
sd.currentVersion = previousVersion
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetDeploymentInfo retrieves deployment info
|
||||
func (sd *SelfDeployer) GetDeploymentInfo(deploymentID string) (*DeploymentInfo, bool) {
|
||||
sd.mu.RLock()
|
||||
defer sd.mu.RUnlock()
|
||||
|
||||
deployment, exists := sd.deployments[deploymentID]
|
||||
return deployment, exists
|
||||
}
|
||||
|
||||
// GetCurrentVersion returns the current orchestrator version
|
||||
func (sd *SelfDeployer) GetCurrentVersion() string {
|
||||
sd.mu.RLock()
|
||||
defer sd.mu.RUnlock()
|
||||
|
||||
return sd.currentVersion
|
||||
}
|
||||
|
||||
// ListDeployments returns all deployments
|
||||
func (sd *SelfDeployer) ListDeployments() map[string]*DeploymentInfo {
|
||||
sd.mu.RLock()
|
||||
defer sd.mu.RUnlock()
|
||||
|
||||
result := make(map[string]*DeploymentInfo)
|
||||
for id, deployment := range sd.deployments {
|
||||
result[id] = deployment
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// HealthCheck checks if the deployed orchestrator is healthy
|
||||
func (sd *SelfDeployer) HealthCheck(deploymentID string) (bool, error) {
|
||||
sd.mu.RLock()
|
||||
defer sd.mu.RUnlock()
|
||||
|
||||
deployment, exists := sd.deployments[deploymentID]
|
||||
if !exists {
|
||||
return false, fmt.Errorf("deployment not found: %s", deploymentID)
|
||||
}
|
||||
|
||||
// Simulate health check
|
||||
return deployment.Status == StatusSuccess, nil
|
||||
}
|
||||
|
||||
// SetVersion sets the target version
|
||||
func (sd *SelfDeployer) SetVersion(version string) {
|
||||
sd.mu.Lock()
|
||||
defer sd.mu.Unlock()
|
||||
|
||||
sd.currentVersion = version
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
package deployment
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestNewSelfDeployer(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
assert.NotNil(t, deployer)
|
||||
assert.Equal(t, "1.0.0", deployer.GetCurrentVersion())
|
||||
}
|
||||
|
||||
func TestBuildContainer(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
deploymentID, err := deployer.BuildContainer("2.0.0")
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, deploymentID)
|
||||
|
||||
deployment, exists := deployer.GetDeploymentInfo(deploymentID)
|
||||
assert.True(t, exists)
|
||||
assert.Equal(t, "2.0.0", deployment.Version)
|
||||
}
|
||||
|
||||
func TestPushImage(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
deploymentID, _ := deployer.BuildContainer("2.0.0")
|
||||
err := deployer.PushImage(deploymentID)
|
||||
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestGenerateManifest(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
deploymentID, _ := deployer.BuildContainer("2.0.0")
|
||||
manifest, err := deployer.GenerateManifest(deploymentID, 3)
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, manifest)
|
||||
assert.Contains(t, manifest, "replicas: 3")
|
||||
}
|
||||
|
||||
func TestDeploy(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
deploymentID, _ := deployer.BuildContainer("2.0.0")
|
||||
deployer.PushImage(deploymentID)
|
||||
deployer.GenerateManifest(deploymentID, 3)
|
||||
err := deployer.Deploy(deploymentID)
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "2.0.0", deployer.GetCurrentVersion())
|
||||
}
|
||||
|
||||
func TestRollback(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
deployer.SetVersion("2.0.0")
|
||||
err := deployer.Rollback("1.0.0")
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "1.0.0", deployer.GetCurrentVersion())
|
||||
}
|
||||
|
||||
func TestHealthCheck(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
deploymentID, _ := deployer.BuildContainer("2.0.0")
|
||||
deployer.PushImage(deploymentID)
|
||||
deployer.GenerateManifest(deploymentID, 3)
|
||||
deployer.Deploy(deploymentID)
|
||||
|
||||
healthy, err := deployer.HealthCheck(deploymentID)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, healthy)
|
||||
}
|
||||
|
||||
func TestListDeployments(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
deployer.BuildContainer("2.0.0")
|
||||
deployer.BuildContainer("2.0.1")
|
||||
|
||||
deployments := deployer.ListDeployments()
|
||||
assert.Equal(t, 2, len(deployments))
|
||||
}
|
||||
|
||||
func TestBuildContainerError(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
_, err := deployer.BuildContainer("")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestPushImageError(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
err := deployer.PushImage("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestFullDeploymentCycle(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
// Build
|
||||
deploymentID, _ := deployer.BuildContainer("2.0.0")
|
||||
|
||||
// Push
|
||||
deployer.PushImage(deploymentID)
|
||||
|
||||
// Generate manifest
|
||||
deployer.GenerateManifest(deploymentID, 3)
|
||||
|
||||
// Deploy
|
||||
err := deployer.Deploy(deploymentID)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify
|
||||
assert.Equal(t, "2.0.0", deployer.GetCurrentVersion())
|
||||
|
||||
// Health check
|
||||
healthy, _ := deployer.HealthCheck(deploymentID)
|
||||
assert.True(t, healthy)
|
||||
}
|
||||
|
||||
func TestSetVersion(t *testing.T) {
|
||||
deployer := NewSelfDeployer("docker.io", "/etc/kubernetes/config")
|
||||
|
||||
deployer.SetVersion("3.0.0")
|
||||
assert.Equal(t, "3.0.0", deployer.GetCurrentVersion())
|
||||
}
|
||||
Vendored
+108
@@ -0,0 +1,108 @@
|
||||
package external
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// ExternalTask represents an imported task from external systems
|
||||
type ExternalTask struct {
|
||||
ID string
|
||||
Source string // "github", "linear", "jira"
|
||||
ExternalID string
|
||||
Title string
|
||||
Status string
|
||||
Body string
|
||||
Labels []string
|
||||
Assignee string
|
||||
}
|
||||
|
||||
// TaskImporter imports tasks from external systems
|
||||
type TaskImporter struct {
|
||||
mu sync.RWMutex
|
||||
tasks map[string]*ExternalTask
|
||||
}
|
||||
|
||||
// NewTaskImporter creates a new task importer
|
||||
func NewTaskImporter() *TaskImporter {
|
||||
return &TaskImporter{
|
||||
tasks: make(map[string]*ExternalTask),
|
||||
}
|
||||
}
|
||||
|
||||
// Import imports a task from external source
|
||||
func (ti *TaskImporter) Import(task *ExternalTask) error {
|
||||
if task.ID == "" {
|
||||
return fmt.Errorf("task ID required")
|
||||
}
|
||||
|
||||
ti.mu.Lock()
|
||||
defer ti.mu.Unlock()
|
||||
|
||||
ti.tasks[task.ID] = task
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetTask retrieves an imported task
|
||||
func (ti *TaskImporter) GetTask(id string) (*ExternalTask, bool) {
|
||||
ti.mu.RLock()
|
||||
defer ti.mu.RUnlock()
|
||||
|
||||
task, exists := ti.tasks[id]
|
||||
return task, exists
|
||||
}
|
||||
|
||||
// ListTasks lists all imported tasks
|
||||
func (ti *TaskImporter) ListTasks() map[string]*ExternalTask {
|
||||
ti.mu.RLock()
|
||||
defer ti.mu.RUnlock()
|
||||
|
||||
result := make(map[string]*ExternalTask)
|
||||
for id, task := range ti.tasks {
|
||||
result[id] = task
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// UpdateStatus updates task status
|
||||
func (ti *TaskImporter) UpdateStatus(id string, status string) error {
|
||||
ti.mu.Lock()
|
||||
defer ti.mu.Unlock()
|
||||
|
||||
task, exists := ti.tasks[id]
|
||||
if !exists {
|
||||
return fmt.Errorf("task not found: %s", id)
|
||||
}
|
||||
|
||||
task.Status = status
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetBySource lists tasks from a specific source
|
||||
func (ti *TaskImporter) GetBySource(source string) []*ExternalTask {
|
||||
ti.mu.RLock()
|
||||
defer ti.mu.RUnlock()
|
||||
|
||||
result := make([]*ExternalTask, 0)
|
||||
for _, task := range ti.tasks {
|
||||
if task.Source == source {
|
||||
result = append(result, task)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// Remove removes a task
|
||||
func (ti *TaskImporter) Remove(id string) error {
|
||||
ti.mu.Lock()
|
||||
defer ti.mu.Unlock()
|
||||
|
||||
if _, exists := ti.tasks[id]; !exists {
|
||||
return fmt.Errorf("task not found: %s", id)
|
||||
}
|
||||
|
||||
delete(ti.tasks, id)
|
||||
return nil
|
||||
}
|
||||
+79
@@ -0,0 +1,79 @@
|
||||
package external
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestImport(t *testing.T) {
|
||||
importer := NewTaskImporter()
|
||||
|
||||
task := &ExternalTask{
|
||||
ID: "github-123",
|
||||
Source: "github",
|
||||
ExternalID: "123",
|
||||
Title: "Add feature",
|
||||
}
|
||||
|
||||
err := importer.Import(task)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestGetTask(t *testing.T) {
|
||||
importer := NewTaskImporter()
|
||||
|
||||
task := &ExternalTask{
|
||||
ID: "github-123",
|
||||
Source: "github",
|
||||
ExternalID: "123",
|
||||
Title: "Add feature",
|
||||
}
|
||||
|
||||
importer.Import(task)
|
||||
|
||||
retrieved, found := importer.GetTask("github-123")
|
||||
assert.True(t, found)
|
||||
assert.Equal(t, "Add feature", retrieved.Title)
|
||||
}
|
||||
|
||||
func TestListTasks(t *testing.T) {
|
||||
importer := NewTaskImporter()
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
importer.Import(&ExternalTask{
|
||||
ID: "task-" + string(rune(48+i)),
|
||||
Source: "github",
|
||||
})
|
||||
}
|
||||
|
||||
tasks := importer.ListTasks()
|
||||
assert.Equal(t, 3, len(tasks))
|
||||
}
|
||||
|
||||
func TestUpdateStatus(t *testing.T) {
|
||||
importer := NewTaskImporter()
|
||||
|
||||
task := &ExternalTask{
|
||||
ID: "github-123",
|
||||
Source: "github",
|
||||
Status: "open",
|
||||
}
|
||||
|
||||
importer.Import(task)
|
||||
importer.UpdateStatus("github-123", "closed")
|
||||
|
||||
updated, _ := importer.GetTask("github-123")
|
||||
assert.Equal(t, "closed", updated.Status)
|
||||
}
|
||||
|
||||
func TestGetBySource(t *testing.T) {
|
||||
importer := NewTaskImporter()
|
||||
|
||||
importer.Import(&ExternalTask{ID: "gh-1", Source: "github"})
|
||||
importer.Import(&ExternalTask{ID: "gh-2", Source: "github"})
|
||||
importer.Import(&ExternalTask{ID: "jira-1", Source: "jira"})
|
||||
|
||||
github := importer.GetBySource("github")
|
||||
assert.Equal(t, 2, len(github))
|
||||
}
|
||||
@@ -0,0 +1,404 @@
|
||||
package graph
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// Task represents a node in the dependency graph
|
||||
type Task struct {
|
||||
ID string
|
||||
Title string
|
||||
Status string // pending, ready, running, completed, failed
|
||||
DependsOn []string
|
||||
Metadata map[string]interface{}
|
||||
}
|
||||
|
||||
// DependencyGraph manages task dependencies
|
||||
type DependencyGraph struct {
|
||||
mu sync.RWMutex
|
||||
tasks map[string]*Task
|
||||
adjacencyList map[string][]string // task -> dependent tasks
|
||||
reverseList map[string][]string // task -> dependencies
|
||||
topologicalOrder []string
|
||||
cycleDetected bool
|
||||
status map[string]string // task -> status
|
||||
}
|
||||
|
||||
// NewDependencyGraph creates a new dependency graph
|
||||
func NewDependencyGraph() *DependencyGraph {
|
||||
return &DependencyGraph{
|
||||
tasks: make(map[string]*Task),
|
||||
adjacencyList: make(map[string][]string),
|
||||
reverseList: make(map[string][]string),
|
||||
topologicalOrder: make([]string, 0),
|
||||
status: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
// AddTask adds a task to the graph
|
||||
func (dg *DependencyGraph) AddTask(task *Task) error {
|
||||
if task == nil || task.ID == "" {
|
||||
return fmt.Errorf("task cannot be nil and must have an ID")
|
||||
}
|
||||
|
||||
dg.mu.Lock()
|
||||
defer dg.mu.Unlock()
|
||||
|
||||
if _, exists := dg.tasks[task.ID]; exists {
|
||||
return fmt.Errorf("task already exists: %s", task.ID)
|
||||
}
|
||||
|
||||
dg.tasks[task.ID] = task
|
||||
dg.status[task.ID] = "pending"
|
||||
|
||||
// Initialize adjacency lists
|
||||
if _, exists := dg.adjacencyList[task.ID]; !exists {
|
||||
dg.adjacencyList[task.ID] = make([]string, 0)
|
||||
}
|
||||
if _, exists := dg.reverseList[task.ID]; !exists {
|
||||
dg.reverseList[task.ID] = make([]string, 0)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddDependency adds a dependency: dependent depends on prerequisite
|
||||
func (dg *DependencyGraph) AddDependency(dependent, prerequisite string) error {
|
||||
dg.mu.Lock()
|
||||
defer dg.mu.Unlock()
|
||||
|
||||
if _, exists := dg.tasks[dependent]; !exists {
|
||||
return fmt.Errorf("dependent task not found: %s", dependent)
|
||||
}
|
||||
|
||||
if _, exists := dg.tasks[prerequisite]; !exists {
|
||||
return fmt.Errorf("prerequisite task not found: %s", prerequisite)
|
||||
}
|
||||
|
||||
// Check for duplicate
|
||||
for _, dep := range dg.reverseList[dependent] {
|
||||
if dep == prerequisite {
|
||||
return fmt.Errorf("dependency already exists: %s -> %s", dependent, prerequisite)
|
||||
}
|
||||
}
|
||||
|
||||
dg.reverseList[dependent] = append(dg.reverseList[dependent], prerequisite)
|
||||
dg.adjacencyList[prerequisite] = append(dg.adjacencyList[prerequisite], dependent)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateGraph checks for cycles and structural integrity
|
||||
func (dg *DependencyGraph) ValidateGraph() error {
|
||||
dg.mu.Lock()
|
||||
defer dg.mu.Unlock()
|
||||
|
||||
// Check for cycles using DFS
|
||||
visited := make(map[string]bool)
|
||||
recStack := make(map[string]bool)
|
||||
|
||||
for taskID := range dg.tasks {
|
||||
if !visited[taskID] {
|
||||
if dg.hasCycleLocked(taskID, visited, recStack) {
|
||||
dg.cycleDetected = true
|
||||
return fmt.Errorf("cycle detected in dependency graph")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// hasCycleLocked detects cycles using DFS (must be called with lock held)
|
||||
func (dg *DependencyGraph) hasCycleLocked(node string, visited, recStack map[string]bool) bool {
|
||||
visited[node] = true
|
||||
recStack[node] = true
|
||||
|
||||
for _, dep := range dg.reverseList[node] {
|
||||
if !visited[dep] {
|
||||
if dg.hasCycleLocked(dep, visited, recStack) {
|
||||
return true
|
||||
}
|
||||
} else if recStack[dep] {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
recStack[node] = false
|
||||
return false
|
||||
}
|
||||
|
||||
// GetTopologicalOrder returns tasks in execution order
|
||||
func (dg *DependencyGraph) GetTopologicalOrder() ([]string, error) {
|
||||
dg.mu.Lock()
|
||||
defer dg.mu.Unlock()
|
||||
|
||||
if dg.cycleDetected {
|
||||
return nil, fmt.Errorf("graph contains cycles")
|
||||
}
|
||||
|
||||
// Kahn's algorithm
|
||||
inDegree := make(map[string]int)
|
||||
for taskID := range dg.tasks {
|
||||
inDegree[taskID] = len(dg.reverseList[taskID])
|
||||
}
|
||||
|
||||
queue := make([]string, 0)
|
||||
for taskID, degree := range inDegree {
|
||||
if degree == 0 {
|
||||
queue = append(queue, taskID)
|
||||
}
|
||||
}
|
||||
|
||||
topOrder := make([]string, 0)
|
||||
for len(queue) > 0 {
|
||||
current := queue[0]
|
||||
queue = queue[1:]
|
||||
topOrder = append(topOrder, current)
|
||||
|
||||
for _, dependent := range dg.adjacencyList[current] {
|
||||
inDegree[dependent]--
|
||||
if inDegree[dependent] == 0 {
|
||||
queue = append(queue, dependent)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(topOrder) != len(dg.tasks) {
|
||||
return nil, fmt.Errorf("topological sort failed - graph may have cycles")
|
||||
}
|
||||
|
||||
dg.topologicalOrder = topOrder
|
||||
return topOrder, nil
|
||||
}
|
||||
|
||||
// GetReadyTasks returns tasks that have no remaining dependencies
|
||||
func (dg *DependencyGraph) GetReadyTasks() []string {
|
||||
dg.mu.RLock()
|
||||
defer dg.mu.RUnlock()
|
||||
|
||||
ready := make([]string, 0)
|
||||
|
||||
for taskID, deps := range dg.reverseList {
|
||||
allDepsComplete := true
|
||||
for _, dep := range deps {
|
||||
if dg.status[dep] != "completed" {
|
||||
allDepsComplete = false
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if allDepsComplete && dg.status[taskID] == "pending" {
|
||||
ready = append(ready, taskID)
|
||||
}
|
||||
}
|
||||
|
||||
return ready
|
||||
}
|
||||
|
||||
// MarkCompleted marks a task as completed and updates dependents
|
||||
func (dg *DependencyGraph) MarkCompleted(taskID string) error {
|
||||
dg.mu.Lock()
|
||||
defer dg.mu.Unlock()
|
||||
|
||||
if _, exists := dg.tasks[taskID]; !exists {
|
||||
return fmt.Errorf("task not found: %s", taskID)
|
||||
}
|
||||
|
||||
dg.status[taskID] = "completed"
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarkFailed marks a task as failed
|
||||
func (dg *DependencyGraph) MarkFailed(taskID string) error {
|
||||
dg.mu.Lock()
|
||||
defer dg.mu.Unlock()
|
||||
|
||||
if _, exists := dg.tasks[taskID]; !exists {
|
||||
return fmt.Errorf("task not found: %s", taskID)
|
||||
}
|
||||
|
||||
dg.status[taskID] = "failed"
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetTaskStatus returns the status of a task
|
||||
func (dg *DependencyGraph) GetTaskStatus(taskID string) (string, error) {
|
||||
dg.mu.RLock()
|
||||
defer dg.mu.RUnlock()
|
||||
|
||||
status, exists := dg.status[taskID]
|
||||
if !exists {
|
||||
return "", fmt.Errorf("task not found: %s", taskID)
|
||||
}
|
||||
|
||||
return status, nil
|
||||
}
|
||||
|
||||
// GetDependencies returns all dependencies of a task
|
||||
func (dg *DependencyGraph) GetDependencies(taskID string) ([]string, error) {
|
||||
dg.mu.RLock()
|
||||
defer dg.mu.RUnlock()
|
||||
|
||||
deps, exists := dg.reverseList[taskID]
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("task not found: %s", taskID)
|
||||
}
|
||||
|
||||
result := make([]string, len(deps))
|
||||
copy(result, deps)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// GetDependents returns all tasks that depend on this task
|
||||
func (dg *DependencyGraph) GetDependents(taskID string) ([]string, error) {
|
||||
dg.mu.RLock()
|
||||
defer dg.mu.RUnlock()
|
||||
|
||||
deps, exists := dg.adjacencyList[taskID]
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("task not found: %s", taskID)
|
||||
}
|
||||
|
||||
result := make([]string, len(deps))
|
||||
copy(result, deps)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// GetTask returns a task by ID
|
||||
func (dg *DependencyGraph) GetTask(taskID string) (*Task, bool) {
|
||||
dg.mu.RLock()
|
||||
defer dg.mu.RUnlock()
|
||||
|
||||
task, exists := dg.tasks[taskID]
|
||||
return task, exists
|
||||
}
|
||||
|
||||
// GetAllTasks returns all tasks
|
||||
func (dg *DependencyGraph) GetAllTasks() map[string]*Task {
|
||||
dg.mu.RLock()
|
||||
defer dg.mu.RUnlock()
|
||||
|
||||
result := make(map[string]*Task)
|
||||
for id, task := range dg.tasks {
|
||||
result[id] = task
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetGraphStats returns statistics about the graph
|
||||
func (dg *DependencyGraph) GetGraphStats() map[string]interface{} {
|
||||
dg.mu.RLock()
|
||||
defer dg.mu.RUnlock()
|
||||
|
||||
pending := 0
|
||||
completed := 0
|
||||
failed := 0
|
||||
|
||||
for _, status := range dg.status {
|
||||
switch status {
|
||||
case "pending":
|
||||
pending++
|
||||
case "completed":
|
||||
completed++
|
||||
case "failed":
|
||||
failed++
|
||||
}
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"total_tasks": len(dg.tasks),
|
||||
"pending_tasks": pending,
|
||||
"completed_tasks": completed,
|
||||
"failed_tasks": failed,
|
||||
"cycle_detected": dg.cycleDetected,
|
||||
"total_edges": dg.countEdgesLocked(),
|
||||
}
|
||||
}
|
||||
|
||||
// countEdgesLocked counts total dependencies (must be called with lock held)
|
||||
func (dg *DependencyGraph) countEdgesLocked() int {
|
||||
count := 0
|
||||
for _, deps := range dg.reverseList {
|
||||
count += len(deps)
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
// Clear clears all tasks and dependencies
|
||||
func (dg *DependencyGraph) Clear() {
|
||||
dg.mu.Lock()
|
||||
defer dg.mu.Unlock()
|
||||
|
||||
dg.tasks = make(map[string]*Task)
|
||||
dg.adjacencyList = make(map[string][]string)
|
||||
dg.reverseList = make(map[string][]string)
|
||||
dg.topologicalOrder = make([]string, 0)
|
||||
dg.status = make(map[string]string)
|
||||
dg.cycleDetected = false
|
||||
}
|
||||
|
||||
// CanExecuteTask checks if a task can be executed (all deps complete)
|
||||
func (dg *DependencyGraph) CanExecuteTask(taskID string) bool {
|
||||
dg.mu.RLock()
|
||||
defer dg.mu.RUnlock()
|
||||
|
||||
deps, exists := dg.reverseList[taskID]
|
||||
if !exists {
|
||||
return false
|
||||
}
|
||||
|
||||
for _, dep := range deps {
|
||||
if dg.status[dep] != "completed" {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// GetCriticalPath returns the longest path through the graph
|
||||
func (dg *DependencyGraph) GetCriticalPath() []string {
|
||||
dg.mu.RLock()
|
||||
defer dg.mu.RUnlock()
|
||||
|
||||
// Use longest path algorithm
|
||||
distances := make(map[string]int)
|
||||
parent := make(map[string]string)
|
||||
|
||||
for taskID := range dg.tasks {
|
||||
distances[taskID] = 0
|
||||
}
|
||||
|
||||
// Process in topological order
|
||||
for _, taskID := range dg.topologicalOrder {
|
||||
for _, dependent := range dg.adjacencyList[taskID] {
|
||||
if distances[dependent] < distances[taskID]+1 {
|
||||
distances[dependent] = distances[taskID] + 1
|
||||
parent[dependent] = taskID
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Find task with maximum distance
|
||||
maxDist := 0
|
||||
endTask := ""
|
||||
for taskID, dist := range distances {
|
||||
if dist > maxDist {
|
||||
maxDist = dist
|
||||
endTask = taskID
|
||||
}
|
||||
}
|
||||
|
||||
// Reconstruct path
|
||||
path := make([]string, 0)
|
||||
current := endTask
|
||||
for current != "" {
|
||||
path = append([]string{current}, path...)
|
||||
current = parent[current]
|
||||
}
|
||||
|
||||
return path
|
||||
}
|
||||
@@ -0,0 +1,372 @@
|
||||
package graph
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestNewDependencyGraph(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
assert.NotNil(t, graph)
|
||||
assert.Equal(t, 0, len(graph.GetAllTasks()))
|
||||
}
|
||||
|
||||
func TestAddTask(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
task := &Task{ID: "T1", Title: "Task 1"}
|
||||
err := graph.AddTask(task)
|
||||
|
||||
assert.NoError(t, err)
|
||||
retrieved, exists := graph.GetTask("T1")
|
||||
assert.True(t, exists)
|
||||
assert.Equal(t, "T1", retrieved.ID)
|
||||
}
|
||||
|
||||
func TestAddTaskNil(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
err := graph.AddTask(nil)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestAddTaskDuplicate(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
task := &Task{ID: "T1", Title: "Task 1"}
|
||||
graph.AddTask(task)
|
||||
|
||||
err := graph.AddTask(task)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestAddDependency(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
graph.AddTask(&Task{ID: "T2", Title: "Task 2"})
|
||||
|
||||
err := graph.AddDependency("T2", "T1")
|
||||
assert.NoError(t, err)
|
||||
|
||||
deps, _ := graph.GetDependencies("T2")
|
||||
assert.Equal(t, 1, len(deps))
|
||||
assert.Equal(t, "T1", deps[0])
|
||||
}
|
||||
|
||||
func TestAddDependencyNotFound(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
|
||||
err := graph.AddDependency("T2", "T1")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestValidateGraphNoCycles(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
graph.AddTask(&Task{ID: "T2", Title: "Task 2"})
|
||||
graph.AddTask(&Task{ID: "T3", Title: "Task 3"})
|
||||
|
||||
graph.AddDependency("T2", "T1")
|
||||
graph.AddDependency("T3", "T2")
|
||||
|
||||
err := graph.ValidateGraph()
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestValidateGraphWithCycle(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
graph.AddTask(&Task{ID: "T2", Title: "Task 2"})
|
||||
graph.AddTask(&Task{ID: "T3", Title: "Task 3"})
|
||||
|
||||
graph.AddDependency("T2", "T1")
|
||||
graph.AddDependency("T3", "T2")
|
||||
graph.AddDependency("T1", "T3") // Creates cycle
|
||||
|
||||
err := graph.ValidateGraph()
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestGetTopologicalOrder(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
graph.AddTask(&Task{ID: "T2", Title: "Task 2"})
|
||||
graph.AddTask(&Task{ID: "T3", Title: "Task 3"})
|
||||
|
||||
graph.AddDependency("T2", "T1")
|
||||
graph.AddDependency("T3", "T2")
|
||||
|
||||
graph.ValidateGraph()
|
||||
order, err := graph.GetTopologicalOrder()
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 3, len(order))
|
||||
assert.Equal(t, "T1", order[0])
|
||||
assert.Equal(t, "T2", order[1])
|
||||
assert.Equal(t, "T3", order[2])
|
||||
}
|
||||
|
||||
func TestGetReadyTasks(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
graph.AddTask(&Task{ID: "T2", Title: "Task 2"})
|
||||
graph.AddTask(&Task{ID: "T3", Title: "Task 3"})
|
||||
|
||||
graph.AddDependency("T2", "T1")
|
||||
graph.AddDependency("T3", "T1")
|
||||
|
||||
ready := graph.GetReadyTasks()
|
||||
assert.Equal(t, 1, len(ready))
|
||||
assert.Equal(t, "T1", ready[0])
|
||||
|
||||
graph.MarkCompleted("T1")
|
||||
ready = graph.GetReadyTasks()
|
||||
assert.Equal(t, 2, len(ready))
|
||||
}
|
||||
|
||||
func TestMarkCompleted(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
err := graph.MarkCompleted("T1")
|
||||
|
||||
assert.NoError(t, err)
|
||||
status, _ := graph.GetTaskStatus("T1")
|
||||
assert.Equal(t, "completed", status)
|
||||
}
|
||||
|
||||
func TestMarkFailed(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
err := graph.MarkFailed("T1")
|
||||
|
||||
assert.NoError(t, err)
|
||||
status, _ := graph.GetTaskStatus("T1")
|
||||
assert.Equal(t, "failed", status)
|
||||
}
|
||||
|
||||
func TestGetDependencies(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
graph.AddTask(&Task{ID: "T2", Title: "Task 2"})
|
||||
graph.AddTask(&Task{ID: "T3", Title: "Task 3"})
|
||||
|
||||
graph.AddDependency("T3", "T1")
|
||||
graph.AddDependency("T3", "T2")
|
||||
|
||||
deps, _ := graph.GetDependencies("T3")
|
||||
assert.Equal(t, 2, len(deps))
|
||||
}
|
||||
|
||||
func TestGetDependents(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
graph.AddTask(&Task{ID: "T2", Title: "Task 2"})
|
||||
graph.AddTask(&Task{ID: "T3", Title: "Task 3"})
|
||||
|
||||
graph.AddDependency("T2", "T1")
|
||||
graph.AddDependency("T3", "T1")
|
||||
|
||||
dependents, _ := graph.GetDependents("T1")
|
||||
assert.Equal(t, 2, len(dependents))
|
||||
}
|
||||
|
||||
func TestCanExecuteTask(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
graph.AddTask(&Task{ID: "T2", Title: "Task 2"})
|
||||
|
||||
graph.AddDependency("T2", "T1")
|
||||
|
||||
assert.False(t, graph.CanExecuteTask("T2"))
|
||||
|
||||
graph.MarkCompleted("T1")
|
||||
assert.True(t, graph.CanExecuteTask("T2"))
|
||||
}
|
||||
|
||||
func TestGetGraphStats(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
graph.AddTask(&Task{ID: "T2", Title: "Task 2"})
|
||||
|
||||
graph.MarkCompleted("T1")
|
||||
|
||||
stats := graph.GetGraphStats()
|
||||
assert.Equal(t, 2, stats["total_tasks"])
|
||||
assert.Equal(t, 1, stats["completed_tasks"])
|
||||
assert.Equal(t, 1, stats["pending_tasks"])
|
||||
}
|
||||
|
||||
func TestClear(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1", Title: "Task 1"})
|
||||
assert.Equal(t, 1, len(graph.GetAllTasks()))
|
||||
|
||||
graph.Clear()
|
||||
assert.Equal(t, 0, len(graph.GetAllTasks()))
|
||||
}
|
||||
|
||||
func TestMultipleDependencies(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
for i := 1; i <= 5; i++ {
|
||||
id := string(rune(48 + i))
|
||||
graph.AddTask(&Task{ID: "T" + id, Title: "Task " + id})
|
||||
}
|
||||
|
||||
// Chain: T1 -> T2 -> T3 -> T4 -> T5
|
||||
for i := 2; i <= 5; i++ {
|
||||
graph.AddDependency("T"+string(rune(48+i)), "T"+string(rune(48+i-1)))
|
||||
}
|
||||
|
||||
ready := graph.GetReadyTasks()
|
||||
assert.Equal(t, 1, len(ready))
|
||||
assert.Equal(t, "T1", ready[0])
|
||||
}
|
||||
|
||||
func TestDiamondDependency(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
// Diamond: T1 -> (T2, T3) -> T4
|
||||
graph.AddTask(&Task{ID: "T1"})
|
||||
graph.AddTask(&Task{ID: "T2"})
|
||||
graph.AddTask(&Task{ID: "T3"})
|
||||
graph.AddTask(&Task{ID: "T4"})
|
||||
|
||||
graph.AddDependency("T2", "T1")
|
||||
graph.AddDependency("T3", "T1")
|
||||
graph.AddDependency("T4", "T2")
|
||||
graph.AddDependency("T4", "T3")
|
||||
|
||||
graph.ValidateGraph()
|
||||
order, _ := graph.GetTopologicalOrder()
|
||||
|
||||
assert.Equal(t, 4, len(order))
|
||||
assert.Equal(t, "T1", order[0])
|
||||
assert.Equal(t, "T4", order[3])
|
||||
}
|
||||
|
||||
func TestGetCriticalPath(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
graph.AddTask(&Task{ID: "T1"})
|
||||
graph.AddTask(&Task{ID: "T2"})
|
||||
graph.AddTask(&Task{ID: "T3"})
|
||||
graph.AddTask(&Task{ID: "T4"})
|
||||
|
||||
graph.AddDependency("T2", "T1")
|
||||
graph.AddDependency("T3", "T2")
|
||||
graph.AddDependency("T4", "T3")
|
||||
|
||||
graph.ValidateGraph()
|
||||
graph.GetTopologicalOrder()
|
||||
|
||||
path := graph.GetCriticalPath()
|
||||
assert.Greater(t, len(path), 0)
|
||||
assert.Equal(t, "T1", path[0])
|
||||
}
|
||||
|
||||
func TestComplexGraph(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
// Create 10 tasks with complex dependencies
|
||||
for i := 1; i <= 10; i++ {
|
||||
id := string(rune(48 + i%10))
|
||||
if i >= 10 {
|
||||
id = "T" + id
|
||||
} else {
|
||||
id = "T0" + id
|
||||
}
|
||||
graph.AddTask(&Task{ID: id})
|
||||
}
|
||||
|
||||
// Add various dependencies
|
||||
graph.AddDependency("T02", "T01")
|
||||
graph.AddDependency("T03", "T01")
|
||||
graph.AddDependency("T04", "T02")
|
||||
graph.AddDependency("T04", "T03")
|
||||
|
||||
graph.ValidateGraph()
|
||||
order, _ := graph.GetTopologicalOrder()
|
||||
assert.Equal(t, 10, len(order))
|
||||
}
|
||||
|
||||
func TestTaskWithMetadata(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
task := &Task{
|
||||
ID: "T1",
|
||||
Title: "Task 1",
|
||||
Metadata: map[string]interface{}{
|
||||
"priority": "high",
|
||||
"owner": "team-a",
|
||||
},
|
||||
}
|
||||
|
||||
graph.AddTask(task)
|
||||
retrieved, _ := graph.GetTask("T1")
|
||||
assert.Equal(t, "high", retrieved.Metadata["priority"])
|
||||
}
|
||||
|
||||
func TestGetAllTasks(t *testing.T) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
for i := 1; i <= 5; i++ {
|
||||
id := string(rune(48 + i))
|
||||
graph.AddTask(&Task{ID: "T" + id})
|
||||
}
|
||||
|
||||
all := graph.GetAllTasks()
|
||||
assert.Equal(t, 5, len(all))
|
||||
}
|
||||
|
||||
func BenchmarkAddTask(b *testing.B) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
for i := 0; i < b.N; i++ {
|
||||
id := string(rune(48 + i%100))
|
||||
graph.AddTask(&Task{ID: "T" + id})
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkAddDependency(b *testing.B) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
for i := 0; i < 1000; i++ {
|
||||
graph.AddTask(&Task{ID: "T" + string(rune(48+i%100))})
|
||||
}
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
from := "T" + string(rune(48+i%100))
|
||||
to := "T" + string(rune(48+(i+1)%100))
|
||||
graph.AddDependency(from, to)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkGetReadyTasks(b *testing.B) {
|
||||
graph := NewDependencyGraph()
|
||||
|
||||
for i := 1; i <= 100; i++ {
|
||||
id := "T" + string(rune(48+i%100))
|
||||
graph.AddTask(&Task{ID: id})
|
||||
}
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
graph.GetReadyTasks()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
package judge
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// Judge represents the interface for custom judge implementations
|
||||
type Judge interface {
|
||||
// Name returns the judge name
|
||||
Name() string
|
||||
// Judge evaluates a task implementation
|
||||
Judge(taskID string, input map[string]interface{}) (map[string]interface{}, error)
|
||||
// Validate checks judge configuration
|
||||
Validate() error
|
||||
}
|
||||
|
||||
// CustomJudgeRegistry manages custom judge implementations
|
||||
type CustomJudgeRegistry struct {
|
||||
mu sync.RWMutex
|
||||
judges map[string]Judge
|
||||
defaultJudge Judge
|
||||
}
|
||||
|
||||
// NewCustomJudgeRegistry creates a new custom judge registry
|
||||
func NewCustomJudgeRegistry() *CustomJudgeRegistry {
|
||||
return &CustomJudgeRegistry{
|
||||
judges: make(map[string]Judge),
|
||||
}
|
||||
}
|
||||
|
||||
// Register registers a custom judge
|
||||
func (cjr *CustomJudgeRegistry) Register(name string, judge Judge) error {
|
||||
if name == "" || judge == nil {
|
||||
return fmt.Errorf("name and judge cannot be empty")
|
||||
}
|
||||
|
||||
if err := judge.Validate(); err != nil {
|
||||
return fmt.Errorf("judge validation failed: %w", err)
|
||||
}
|
||||
|
||||
cjr.mu.Lock()
|
||||
defer cjr.mu.Unlock()
|
||||
|
||||
if _, exists := cjr.judges[name]; exists {
|
||||
return fmt.Errorf("judge already registered: %s", name)
|
||||
}
|
||||
|
||||
cjr.judges[name] = judge
|
||||
return nil
|
||||
}
|
||||
|
||||
// Unregister removes a judge
|
||||
func (cjr *CustomJudgeRegistry) Unregister(name string) error {
|
||||
cjr.mu.Lock()
|
||||
defer cjr.mu.Unlock()
|
||||
|
||||
if _, exists := cjr.judges[name]; !exists {
|
||||
return fmt.Errorf("judge not found: %s", name)
|
||||
}
|
||||
|
||||
delete(cjr.judges, name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get retrieves a judge by name
|
||||
func (cjr *CustomJudgeRegistry) Get(name string) (Judge, bool) {
|
||||
cjr.mu.RLock()
|
||||
defer cjr.mu.RUnlock()
|
||||
|
||||
judge, exists := cjr.judges[name]
|
||||
return judge, exists
|
||||
}
|
||||
|
||||
// Judge executes judgment with custom judge
|
||||
func (cjr *CustomJudgeRegistry) Judge(name string, taskID string, input map[string]interface{}) (map[string]interface{}, error) {
|
||||
judge, exists := cjr.Get(name)
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("judge not found: %s", name)
|
||||
}
|
||||
|
||||
return judge.Judge(taskID, input)
|
||||
}
|
||||
|
||||
// SetDefaultJudge sets the default judge
|
||||
func (cjr *CustomJudgeRegistry) SetDefaultJudge(judge Judge) error {
|
||||
if err := judge.Validate(); err != nil {
|
||||
return fmt.Errorf("judge validation failed: %w", err)
|
||||
}
|
||||
|
||||
cjr.mu.Lock()
|
||||
defer cjr.mu.Unlock()
|
||||
|
||||
cjr.defaultJudge = judge
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetDefaultJudge gets the default judge
|
||||
func (cjr *CustomJudgeRegistry) GetDefaultJudge() Judge {
|
||||
cjr.mu.RLock()
|
||||
defer cjr.mu.RUnlock()
|
||||
|
||||
return cjr.defaultJudge
|
||||
}
|
||||
|
||||
// ListJudges returns all registered judges
|
||||
func (cjr *CustomJudgeRegistry) ListJudges() map[string]Judge {
|
||||
cjr.mu.RLock()
|
||||
defer cjr.mu.RUnlock()
|
||||
|
||||
result := make(map[string]Judge)
|
||||
for name, judge := range cjr.judges {
|
||||
result[name] = judge
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package judge
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
type MockJudge struct {
|
||||
name string
|
||||
}
|
||||
|
||||
func (mj *MockJudge) Name() string {
|
||||
return mj.name
|
||||
}
|
||||
|
||||
func (mj *MockJudge) Judge(taskID string, input map[string]interface{}) (map[string]interface{}, error) {
|
||||
return map[string]interface{}{"approved": true}, nil
|
||||
}
|
||||
|
||||
func (mj *MockJudge) Validate() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestRegisterJudge(t *testing.T) {
|
||||
registry := NewCustomJudgeRegistry()
|
||||
judge := &MockJudge{name: "security-auditor"}
|
||||
|
||||
err := registry.Register("security-auditor", judge)
|
||||
assert.NoError(t, err)
|
||||
|
||||
retrieved, exists := registry.Get("security-auditor")
|
||||
assert.True(t, exists)
|
||||
assert.Equal(t, "security-auditor", retrieved.Name())
|
||||
}
|
||||
|
||||
func TestJudge(t *testing.T) {
|
||||
registry := NewCustomJudgeRegistry()
|
||||
judge := &MockJudge{name: "security-auditor"}
|
||||
|
||||
registry.Register("security-auditor", judge)
|
||||
result, err := registry.Judge("security-auditor", "T0.1", map[string]interface{}{})
|
||||
|
||||
assert.NoError(t, err)
|
||||
approved, ok := result["approved"].(bool)
|
||||
assert.True(t, ok)
|
||||
assert.True(t, approved)
|
||||
}
|
||||
|
||||
func TestListJudges(t *testing.T) {
|
||||
registry := NewCustomJudgeRegistry()
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
registry.Register("judge-"+string(rune(48+i)), &MockJudge{})
|
||||
}
|
||||
|
||||
judges := registry.ListJudges()
|
||||
assert.Equal(t, 3, len(judges))
|
||||
}
|
||||
|
||||
func TestSetDefault(t *testing.T) {
|
||||
registry := NewCustomJudgeRegistry()
|
||||
judge := &MockJudge{name: "default"}
|
||||
|
||||
registry.SetDefaultJudge(judge)
|
||||
assert.NotNil(t, registry.GetDefaultJudge())
|
||||
}
|
||||
|
||||
func TestUnregister(t *testing.T) {
|
||||
registry := NewCustomJudgeRegistry()
|
||||
judge := &MockJudge{name: "test"}
|
||||
|
||||
registry.Register("test", judge)
|
||||
err := registry.Unregister("test")
|
||||
|
||||
assert.NoError(t, err)
|
||||
_, exists := registry.Get("test")
|
||||
assert.False(t, exists)
|
||||
}
|
||||
@@ -0,0 +1,337 @@
|
||||
package locking
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// LockBackend defines the interface for lock backends
|
||||
type LockBackend interface {
|
||||
// Acquire attempts to acquire a lock
|
||||
Acquire(key string, ttl time.Duration) (string, error)
|
||||
// Release releases a lock
|
||||
Release(key string, token string) error
|
||||
// Renew renews a lock's TTL
|
||||
Renew(key string, token string, ttl time.Duration) error
|
||||
// IsLocked checks if a lock is held
|
||||
IsLocked(key string) (bool, error)
|
||||
}
|
||||
|
||||
// LocalLockBackend is a fallback in-memory lock backend
|
||||
type LocalLockBackend struct {
|
||||
mu sync.RWMutex
|
||||
locks map[string]string
|
||||
}
|
||||
|
||||
// NewLocalLockBackend creates a new local lock backend
|
||||
func NewLocalLockBackend() *LocalLockBackend {
|
||||
return &LocalLockBackend{
|
||||
locks: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
// Acquire acquires a lock locally
|
||||
func (lb *LocalLockBackend) Acquire(key string, ttl time.Duration) (string, error) {
|
||||
lb.mu.Lock()
|
||||
defer lb.mu.Unlock()
|
||||
|
||||
if _, exists := lb.locks[key]; exists {
|
||||
return "", fmt.Errorf("lock already held")
|
||||
}
|
||||
|
||||
token := generateToken()
|
||||
lb.locks[key] = token
|
||||
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// Release releases a lock locally
|
||||
func (lb *LocalLockBackend) Release(key string, token string) error {
|
||||
lb.mu.Lock()
|
||||
defer lb.mu.Unlock()
|
||||
|
||||
if held, exists := lb.locks[key]; !exists || held != token {
|
||||
return fmt.Errorf("lock not held by token")
|
||||
}
|
||||
|
||||
delete(lb.locks, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Renew renews a lock locally (no-op for local backend)
|
||||
func (lb *LocalLockBackend) Renew(key string, token string, ttl time.Duration) error {
|
||||
lb.mu.RLock()
|
||||
defer lb.mu.RUnlock()
|
||||
|
||||
if held, exists := lb.locks[key]; !exists || held != token {
|
||||
return fmt.Errorf("lock not held by token")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// IsLocked checks if a lock is held locally
|
||||
func (lb *LocalLockBackend) IsLocked(key string) (bool, error) {
|
||||
lb.mu.RLock()
|
||||
defer lb.mu.RUnlock()
|
||||
|
||||
_, exists := lb.locks[key]
|
||||
return exists, nil
|
||||
}
|
||||
|
||||
// DistributedLock represents a distributed lock
|
||||
type DistributedLock struct {
|
||||
key string
|
||||
token string
|
||||
backend LockBackend
|
||||
mu sync.RWMutex
|
||||
acquired bool
|
||||
acquiredAt time.Time
|
||||
ttl time.Duration
|
||||
}
|
||||
|
||||
// NewDistributedLock creates a new distributed lock
|
||||
func NewDistributedLock(key string, backend LockBackend, ttl time.Duration) *DistributedLock {
|
||||
if ttl == 0 {
|
||||
ttl = 30 * time.Second // Default TTL
|
||||
}
|
||||
|
||||
return &DistributedLock{
|
||||
key: key,
|
||||
backend: backend,
|
||||
ttl: ttl,
|
||||
}
|
||||
}
|
||||
|
||||
// Acquire acquires the lock with timeout
|
||||
func (dl *DistributedLock) Acquire(timeout time.Duration) error {
|
||||
if timeout == 0 {
|
||||
timeout = 5 * time.Second // Default timeout
|
||||
}
|
||||
|
||||
deadline := time.Now().Add(timeout)
|
||||
|
||||
for {
|
||||
token, err := dl.backend.Acquire(dl.key, dl.ttl)
|
||||
if err == nil {
|
||||
dl.mu.Lock()
|
||||
dl.token = token
|
||||
dl.acquired = true
|
||||
dl.acquiredAt = time.Now()
|
||||
dl.mu.Unlock()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
if time.Now().After(deadline) {
|
||||
return fmt.Errorf("lock acquisition timeout")
|
||||
}
|
||||
|
||||
// Back off before retrying
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
|
||||
// Release releases the lock
|
||||
func (dl *DistributedLock) Release() error {
|
||||
dl.mu.Lock()
|
||||
defer dl.mu.Unlock()
|
||||
|
||||
if !dl.acquired {
|
||||
return fmt.Errorf("lock not acquired")
|
||||
}
|
||||
|
||||
err := dl.backend.Release(dl.key, dl.token)
|
||||
if err == nil {
|
||||
dl.acquired = false
|
||||
dl.token = ""
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// Renew renews the lock's TTL
|
||||
func (dl *DistributedLock) Renew() error {
|
||||
dl.mu.RLock()
|
||||
defer dl.mu.RUnlock()
|
||||
|
||||
if !dl.acquired {
|
||||
return fmt.Errorf("lock not acquired")
|
||||
}
|
||||
|
||||
return dl.backend.Renew(dl.key, dl.token, dl.ttl)
|
||||
}
|
||||
|
||||
// IsAcquired checks if the lock is currently acquired
|
||||
func (dl *DistributedLock) IsAcquired() bool {
|
||||
dl.mu.RLock()
|
||||
defer dl.mu.RUnlock()
|
||||
|
||||
return dl.acquired
|
||||
}
|
||||
|
||||
// GetAcquiredAt returns when the lock was acquired
|
||||
func (dl *DistributedLock) GetAcquiredAt() time.Time {
|
||||
dl.mu.RLock()
|
||||
defer dl.mu.RUnlock()
|
||||
|
||||
return dl.acquiredAt
|
||||
}
|
||||
|
||||
// GetHoldDuration returns how long the lock has been held
|
||||
func (dl *DistributedLock) GetHoldDuration() time.Duration {
|
||||
dl.mu.RLock()
|
||||
defer dl.mu.RUnlock()
|
||||
|
||||
if !dl.acquired {
|
||||
return 0
|
||||
}
|
||||
|
||||
return time.Since(dl.acquiredAt)
|
||||
}
|
||||
|
||||
// LockManager manages multiple distributed locks
|
||||
type LockManager struct {
|
||||
mu sync.RWMutex
|
||||
backend LockBackend
|
||||
locks map[string]*DistributedLock
|
||||
lockTTL time.Duration
|
||||
stats *LockStats
|
||||
}
|
||||
|
||||
// LockStats tracks lock statistics
|
||||
type LockStats struct {
|
||||
TotalAcquisitions int
|
||||
TotalReleases int
|
||||
FailedAcquisitions int
|
||||
ActiveLocks int
|
||||
AverageLockTime time.Duration
|
||||
}
|
||||
|
||||
// NewLockManager creates a new lock manager
|
||||
func NewLockManager(backend LockBackend, lockTTL time.Duration) *LockManager {
|
||||
if lockTTL == 0 {
|
||||
lockTTL = 30 * time.Second
|
||||
}
|
||||
|
||||
return &LockManager{
|
||||
backend: backend,
|
||||
locks: make(map[string]*DistributedLock),
|
||||
lockTTL: lockTTL,
|
||||
stats: &LockStats{
|
||||
TotalAcquisitions: 0,
|
||||
TotalReleases: 0,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// AcquireLock acquires or retrieves an existing lock
|
||||
func (lm *LockManager) AcquireLock(key string, timeout time.Duration) error {
|
||||
lm.mu.Lock()
|
||||
defer lm.mu.Unlock()
|
||||
|
||||
// Check if lock already exists and is acquired
|
||||
if lock, exists := lm.locks[key]; exists && lock.IsAcquired() {
|
||||
return fmt.Errorf("lock already acquired by this manager")
|
||||
}
|
||||
|
||||
lock := NewDistributedLock(key, lm.backend, lm.lockTTL)
|
||||
err := lock.Acquire(timeout)
|
||||
if err != nil {
|
||||
lm.stats.FailedAcquisitions++
|
||||
return err
|
||||
}
|
||||
|
||||
lm.locks[key] = lock
|
||||
lm.stats.TotalAcquisitions++
|
||||
lm.stats.ActiveLocks++
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ReleaseLock releases a lock
|
||||
func (lm *LockManager) ReleaseLock(key string) error {
|
||||
lm.mu.Lock()
|
||||
defer lm.mu.Unlock()
|
||||
|
||||
lock, exists := lm.locks[key]
|
||||
if !exists {
|
||||
return fmt.Errorf("lock not found")
|
||||
}
|
||||
|
||||
err := lock.Release()
|
||||
if err == nil {
|
||||
lm.stats.TotalReleases++
|
||||
lm.stats.ActiveLocks--
|
||||
delete(lm.locks, key)
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// RenewLock renews a lock
|
||||
func (lm *LockManager) RenewLock(key string) error {
|
||||
lm.mu.RLock()
|
||||
defer lm.mu.RUnlock()
|
||||
|
||||
lock, exists := lm.locks[key]
|
||||
if !exists {
|
||||
return fmt.Errorf("lock not found")
|
||||
}
|
||||
|
||||
return lock.Renew()
|
||||
}
|
||||
|
||||
// GetLockStats returns lock statistics
|
||||
func (lm *LockManager) GetLockStats() *LockStats {
|
||||
lm.mu.RLock()
|
||||
defer lm.mu.RUnlock()
|
||||
|
||||
stats := *lm.stats
|
||||
return &stats
|
||||
}
|
||||
|
||||
// GetActiveLocks returns list of active lock keys
|
||||
func (lm *LockManager) GetActiveLocks() []string {
|
||||
lm.mu.RLock()
|
||||
defer lm.mu.RUnlock()
|
||||
|
||||
keys := make([]string, 0, len(lm.locks))
|
||||
for key, lock := range lm.locks {
|
||||
if lock.IsAcquired() {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
}
|
||||
|
||||
return keys
|
||||
}
|
||||
|
||||
// ReleaseAll releases all locks
|
||||
func (lm *LockManager) ReleaseAll() error {
|
||||
lm.mu.Lock()
|
||||
defer lm.mu.Unlock()
|
||||
|
||||
var lastErr error
|
||||
for key, lock := range lm.locks {
|
||||
if lock.IsAcquired() {
|
||||
err := lock.Release()
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
}
|
||||
delete(lm.locks, key)
|
||||
}
|
||||
}
|
||||
|
||||
lm.stats.ActiveLocks = 0
|
||||
return lastErr
|
||||
}
|
||||
|
||||
// generateToken generates a random token for lock identification
|
||||
func generateToken() string {
|
||||
b := make([]byte, 16)
|
||||
rand.Read(b)
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
@@ -0,0 +1,312 @@
|
||||
package locking
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestLocalLockBackend(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
|
||||
token1, err := backend.Acquire("test-lock", 30*time.Second)
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, token1)
|
||||
|
||||
// Try to acquire again (should fail)
|
||||
_, err = backend.Acquire("test-lock", 30*time.Second)
|
||||
assert.Error(t, err)
|
||||
|
||||
// Release
|
||||
err = backend.Release("test-lock", token1)
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Acquire again (should succeed)
|
||||
_, err = backend.Acquire("test-lock", 30*time.Second)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestLocalLockReleaseWrongToken(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
|
||||
_, _ = backend.Acquire("test-lock", 30*time.Second)
|
||||
err := backend.Release("test-lock", "wrong-token")
|
||||
assert.Error(t, err)
|
||||
|
||||
// Lock should still be held
|
||||
locked, _ := backend.IsLocked("test-lock")
|
||||
assert.True(t, locked)
|
||||
}
|
||||
|
||||
func TestLocalLockIsLocked(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
|
||||
locked, _ := backend.IsLocked("test-lock")
|
||||
assert.False(t, locked)
|
||||
|
||||
backend.Acquire("test-lock", 30*time.Second)
|
||||
locked, _ = backend.IsLocked("test-lock")
|
||||
assert.True(t, locked)
|
||||
}
|
||||
|
||||
func TestDistributedLockAcquire(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock := NewDistributedLock("test-lock", backend, 30*time.Second)
|
||||
|
||||
err := lock.Acquire(5 * time.Second)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, lock.IsAcquired())
|
||||
}
|
||||
|
||||
func TestDistributedLockRelease(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock := NewDistributedLock("test-lock", backend, 30*time.Second)
|
||||
|
||||
lock.Acquire(5 * time.Second)
|
||||
err := lock.Release()
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, lock.IsAcquired())
|
||||
}
|
||||
|
||||
func TestDistributedLockTimeout(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock1 := NewDistributedLock("test-lock", backend, 30*time.Second)
|
||||
lock2 := NewDistributedLock("test-lock", backend, 30*time.Second)
|
||||
|
||||
lock1.Acquire(5 * time.Second)
|
||||
|
||||
// Try to acquire with very short timeout
|
||||
start := time.Now()
|
||||
err := lock2.Acquire(100 * time.Millisecond)
|
||||
elapsed := time.Since(start)
|
||||
|
||||
assert.Error(t, err)
|
||||
assert.Greater(t, elapsed, 50*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestDistributedLockGetAcquiredAt(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock := NewDistributedLock("test-lock", backend, 30*time.Second)
|
||||
|
||||
lock.Acquire(5 * time.Second)
|
||||
acquiredAt := lock.GetAcquiredAt()
|
||||
|
||||
assert.NotZero(t, acquiredAt)
|
||||
assert.True(t, acquiredAt.Before(time.Now()))
|
||||
}
|
||||
|
||||
func TestDistributedLockGetHoldDuration(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock := NewDistributedLock("test-lock", backend, 30*time.Second)
|
||||
|
||||
lock.Acquire(5 * time.Second)
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
duration := lock.GetHoldDuration()
|
||||
|
||||
assert.Greater(t, duration, 50*time.Millisecond)
|
||||
assert.Less(t, duration, 200*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestDistributedLockRenew(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock := NewDistributedLock("test-lock", backend, 30*time.Second)
|
||||
|
||||
lock.Acquire(5 * time.Second)
|
||||
err := lock.Renew()
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, lock.IsAcquired())
|
||||
}
|
||||
|
||||
func TestLockManagerAcquire(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
|
||||
err := manager.AcquireLock("lock-1", 5*time.Second)
|
||||
assert.NoError(t, err)
|
||||
|
||||
stats := manager.GetLockStats()
|
||||
assert.Equal(t, 1, stats.TotalAcquisitions)
|
||||
assert.Equal(t, 1, stats.ActiveLocks)
|
||||
}
|
||||
|
||||
func TestLockManagerRelease(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
|
||||
manager.AcquireLock("lock-1", 5*time.Second)
|
||||
err := manager.ReleaseLock("lock-1")
|
||||
assert.NoError(t, err)
|
||||
|
||||
stats := manager.GetLockStats()
|
||||
assert.Equal(t, 1, stats.TotalReleases)
|
||||
assert.Equal(t, 0, stats.ActiveLocks)
|
||||
}
|
||||
|
||||
func TestLockManagerMultipleLocks(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
key := "lock-" + string(rune(48+i))
|
||||
err := manager.AcquireLock(key, 5*time.Second)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
stats := manager.GetLockStats()
|
||||
assert.Equal(t, 5, stats.ActiveLocks)
|
||||
|
||||
activeLocks := manager.GetActiveLocks()
|
||||
assert.Equal(t, 5, len(activeLocks))
|
||||
}
|
||||
|
||||
func TestLockManagerReleaseAll(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
key := "lock-" + string(rune(48+i))
|
||||
manager.AcquireLock(key, 5*time.Second)
|
||||
}
|
||||
|
||||
assert.Equal(t, 5, manager.GetLockStats().ActiveLocks)
|
||||
|
||||
manager.ReleaseAll()
|
||||
|
||||
assert.Equal(t, 0, manager.GetLockStats().ActiveLocks)
|
||||
assert.Equal(t, 0, len(manager.GetActiveLocks()))
|
||||
}
|
||||
|
||||
func TestLockManagerRenew(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
|
||||
manager.AcquireLock("lock-1", 5*time.Second)
|
||||
err := manager.RenewLock("lock-1")
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestLockManagerGetStats(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
|
||||
manager.AcquireLock("lock-1", 5*time.Second)
|
||||
manager.AcquireLock("lock-2", 5*time.Second)
|
||||
|
||||
manager.ReleaseLock("lock-1")
|
||||
|
||||
stats := manager.GetLockStats()
|
||||
assert.Equal(t, 2, stats.TotalAcquisitions)
|
||||
assert.Equal(t, 1, stats.TotalReleases)
|
||||
assert.Equal(t, 1, stats.ActiveLocks)
|
||||
}
|
||||
|
||||
func TestLockManagerFailedAcquisition(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
lock1 := NewDistributedLock("lock-1", backend, 30*time.Second)
|
||||
|
||||
// Acquire from outside manager
|
||||
lock1.Acquire(5 * time.Second)
|
||||
|
||||
// Try to acquire from manager
|
||||
err := manager.AcquireLock("lock-1", 100*time.Millisecond)
|
||||
assert.Error(t, err)
|
||||
|
||||
stats := manager.GetLockStats()
|
||||
assert.Equal(t, 1, stats.FailedAcquisitions)
|
||||
}
|
||||
|
||||
func TestDistributedLockDifferentKeys(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock1 := NewDistributedLock("lock-1", backend, 30*time.Second)
|
||||
lock2 := NewDistributedLock("lock-2", backend, 30*time.Second)
|
||||
|
||||
lock1.Acquire(5 * time.Second)
|
||||
// lock2 should acquire without blocking
|
||||
err := lock2.Acquire(100 * time.Millisecond)
|
||||
assert.NoError(t, err)
|
||||
|
||||
assert.True(t, lock1.IsAcquired())
|
||||
assert.True(t, lock2.IsAcquired())
|
||||
}
|
||||
|
||||
func TestLockManagerDuplicateAcquisition(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
|
||||
manager.AcquireLock("lock-1", 5*time.Second)
|
||||
err := manager.AcquireLock("lock-1", 5*time.Second)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestLockManagerReleaseMissing(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
|
||||
err := manager.ReleaseLock("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestDistributedLockReleaseNotAcquired(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock := NewDistributedLock("test-lock", backend, 30*time.Second)
|
||||
|
||||
err := lock.Release()
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestConcurrentLockAcquisition(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock := NewDistributedLock("shared-lock", backend, 30*time.Second)
|
||||
|
||||
acquired := false
|
||||
lock.Acquire(5 * time.Second)
|
||||
|
||||
// Simulate another goroutine trying to acquire
|
||||
go func() {
|
||||
lock2 := NewDistributedLock("shared-lock", backend, 30*time.Second)
|
||||
err := lock2.Acquire(100 * time.Millisecond)
|
||||
if err == nil {
|
||||
acquired = true
|
||||
}
|
||||
}()
|
||||
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
assert.False(t, acquired)
|
||||
}
|
||||
|
||||
func TestDefaultLockTTL(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
lock := NewDistributedLock("test-lock", backend, 0)
|
||||
|
||||
assert.Equal(t, 30*time.Second, lock.ttl)
|
||||
}
|
||||
|
||||
func TestDefaultLockManagerTTL(t *testing.T) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 0)
|
||||
|
||||
assert.Equal(t, 30*time.Second, manager.lockTTL)
|
||||
}
|
||||
|
||||
func BenchmarkLockAcquisition(b *testing.B) {
|
||||
backend := NewLocalLockBackend()
|
||||
|
||||
for i := 0; i < b.N; i++ {
|
||||
lock := NewDistributedLock("test-lock", backend, 30*time.Second)
|
||||
lock.Acquire(5 * time.Second)
|
||||
lock.Release()
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLockManagerAcquisition(b *testing.B) {
|
||||
backend := NewLocalLockBackend()
|
||||
manager := NewLockManager(backend, 30*time.Second)
|
||||
|
||||
for i := 0; i < b.N; i++ {
|
||||
key := "lock-" + string(rune(48+i%100))
|
||||
manager.AcquireLock(key, 5*time.Second)
|
||||
manager.ReleaseLock(key)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,288 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"go.temporal.io/sdk/activity"
|
||||
)
|
||||
|
||||
// Activities memory service activities for Temporal workflows
|
||||
type Activities struct {
|
||||
service *Service
|
||||
}
|
||||
|
||||
// NewActivities creates memory service activities
|
||||
func NewActivities(service *Service) *Activities {
|
||||
return &Activities{
|
||||
service: service,
|
||||
}
|
||||
}
|
||||
|
||||
// CreateKnowledgeActivity creates knowledge record from workflow execution
|
||||
func (a *Activities) CreateKnowledgeActivity(ctx context.Context, record *KnowledgeRecord) (string, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Creating knowledge", "title", record.Title)
|
||||
|
||||
id, err := a.service.CreateKnowledge(ctx, record)
|
||||
if err != nil {
|
||||
logger.Error("Failed to create knowledge", "error", err)
|
||||
return "", err
|
||||
}
|
||||
|
||||
logger.Info("Knowledge created", "id", id)
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// UpdateKnowledgeActivity updates existing knowledge record
|
||||
func (a *Activities) UpdateKnowledgeActivity(ctx context.Context, record *KnowledgeRecord) (string, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Updating knowledge", "id", record.ID)
|
||||
|
||||
id, err := a.service.UpdateKnowledge(ctx, record)
|
||||
if err != nil {
|
||||
logger.Error("Failed to update knowledge", "error", err)
|
||||
return "", err
|
||||
}
|
||||
|
||||
logger.Info("Knowledge updated", "id", id)
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// SearchKnowledgeActivity searches knowledge base
|
||||
func (a *Activities) SearchKnowledgeActivity(ctx context.Context, query string, opts *RetrievalOptions) ([]KnowledgeRecord, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Searching knowledge", "query", query)
|
||||
|
||||
records, err := a.service.RetrieveKnowledge(ctx, query, opts)
|
||||
if err != nil {
|
||||
logger.Error("Search failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
logger.Info("Found records", "count", len(records))
|
||||
return records, nil
|
||||
}
|
||||
|
||||
// GetContextActivity retrieves context for tool/task (three-tier retrieval)
|
||||
func (a *Activities) GetContextActivity(ctx context.Context, tool, task string, budget int) (*ServiceContext, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Getting context", "tool", tool, "task", task)
|
||||
|
||||
svcCtx, err := a.service.RetrieveContext(ctx, tool, task, budget)
|
||||
if err != nil {
|
||||
logger.Error("Get context failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
logger.Info("Retrieved context", "tier", svcCtx.Tier, "lessons", len(svcCtx.Lessons))
|
||||
return svcCtx, nil
|
||||
}
|
||||
|
||||
// GetVaultActivity lists vault files
|
||||
func (a *Activities) GetVaultActivity(ctx context.Context) ([]VaultInfo, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Fetching vault")
|
||||
|
||||
files, err := a.service.GetVault(ctx)
|
||||
if err != nil {
|
||||
logger.Error("Get vault failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
logger.Info("Vault files", "count", len(files))
|
||||
return files, nil
|
||||
}
|
||||
|
||||
// HealthCheckActivity checks memory service health
|
||||
func (a *Activities) HealthCheckActivity(ctx context.Context) (bool, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Checking memory service health")
|
||||
|
||||
if !a.service.IsHealthy(ctx) {
|
||||
logger.Warn("Memory service is unhealthy")
|
||||
return false, fmt.Errorf("memory service unhealthy")
|
||||
}
|
||||
|
||||
logger.Info("Memory service is healthy")
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// LearnFromExecutionActivity learns from task execution result
|
||||
func (a *Activities) LearnFromExecutionActivity(ctx context.Context, taskID string, result string, tags []string) (string, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Learning from task execution", "taskID", taskID)
|
||||
|
||||
metadata := map[string]interface{}{
|
||||
"task_id": taskID,
|
||||
"type": "execution_result",
|
||||
}
|
||||
|
||||
if len(tags) > 0 {
|
||||
metadata["tags"] = tags
|
||||
}
|
||||
|
||||
id, err := a.service.CreateKnowledge(ctx, &KnowledgeRecord{
|
||||
Level: "L1",
|
||||
Title: fmt.Sprintf("Task Execution: %s", taskID),
|
||||
Content: result,
|
||||
Source: fmt.Sprintf("workflow://task/%s", taskID),
|
||||
Metadata: metadata,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
logger.Error("Failed to learn from execution", "error", err)
|
||||
return "", err
|
||||
}
|
||||
|
||||
logger.Info("Learned from execution", "id", id)
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// DiagnoseIssueActivity diagnoses issue using memory context
|
||||
func (a *Activities) DiagnoseIssueActivity(ctx context.Context, tool, issue string) ([]string, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Diagnosing issue", "tool", tool, "issue", issue)
|
||||
|
||||
svcCtx, err := a.service.RetrieveContext(ctx, tool, issue, 8192)
|
||||
if err != nil {
|
||||
logger.Error("Diagnosis failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Extract recommendations
|
||||
recommendations := make([]string, 0)
|
||||
|
||||
// Add tier-1 lessons (highest confidence)
|
||||
for _, lesson := range svcCtx.Lessons {
|
||||
if lesson.Tier == 1 {
|
||||
recommendations = append(recommendations, fmt.Sprintf("[Tier 1] %s", lesson.Text))
|
||||
}
|
||||
}
|
||||
|
||||
// Add skills
|
||||
for _, skill := range svcCtx.Skills {
|
||||
recommendations = append(recommendations, fmt.Sprintf("[Skill] %s: %s", skill.Name, skill.Why))
|
||||
}
|
||||
|
||||
// Add tier-2 lessons if no tier-1
|
||||
if len(recommendations) == 0 {
|
||||
for _, lesson := range svcCtx.Lessons {
|
||||
if lesson.Tier == 2 {
|
||||
recommendations = append(recommendations, fmt.Sprintf("[Tier 2] %s", lesson.Text))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
logger.Info("Generated recommendations", "count", len(recommendations))
|
||||
return recommendations, nil
|
||||
}
|
||||
|
||||
// AnalyzeErrorActivity analyzes error and retrieves relevant knowledge
|
||||
func (a *Activities) AnalyzeErrorActivity(ctx context.Context, errorMsg string) ([]KnowledgeRecord, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Analyzing error")
|
||||
|
||||
// Search for relevant knowledge
|
||||
records, err := a.service.RetrieveKnowledge(ctx, errorMsg, &RetrievalOptions{
|
||||
Limit: 10,
|
||||
LevelFilter: []string{"L1", "L2"},
|
||||
Floor: 0.6,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
logger.Error("Error analysis failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
logger.Info("Found relevant records for error", "count", len(records))
|
||||
return records, nil
|
||||
}
|
||||
|
||||
// DocumentDecisionActivity documents workflow decision in knowledge base
|
||||
func (a *Activities) DocumentDecisionActivity(ctx context.Context, decisionType string, decision string, reasoning string) (string, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Documenting decision", "type", decisionType)
|
||||
|
||||
content := fmt.Sprintf("Decision: %s\n\nReasoning: %s", decision, reasoning)
|
||||
|
||||
id, err := a.service.CreateKnowledge(ctx, &KnowledgeRecord{
|
||||
Level: "L2",
|
||||
Title: fmt.Sprintf("Decision: %s", decisionType),
|
||||
Content: content,
|
||||
Source: fmt.Sprintf("workflow://decision/%s", decisionType),
|
||||
Metadata: map[string]interface{}{
|
||||
"decision_type": decisionType,
|
||||
"type": "workflow_decision",
|
||||
},
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
logger.Error("Failed to document decision", "error", err)
|
||||
return "", err
|
||||
}
|
||||
|
||||
logger.Info("Decision documented", "id", id)
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// SearchAndApplyActivity searches knowledge and applies it
|
||||
func (a *Activities) SearchAndApplyActivity(ctx context.Context, query string, selector func(record *KnowledgeRecord) bool) ([]string, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Searching and applying", "query", query)
|
||||
|
||||
records, err := a.service.RetrieveKnowledge(ctx, query, &RetrievalOptions{
|
||||
Limit: 10,
|
||||
Floor: 0.7,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
logger.Error("Search and apply failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
applied := make([]string, 0)
|
||||
for _, record := range records {
|
||||
if selector == nil || selector(&record) {
|
||||
applied = append(applied, record.Content)
|
||||
logger.Info("Applied knowledge", "id", record.ID)
|
||||
}
|
||||
}
|
||||
|
||||
logger.Info("Applied knowledge records", "count", len(applied))
|
||||
return applied, nil
|
||||
}
|
||||
|
||||
// RefreshMemoryActivity refreshes memory context (periodic activity)
|
||||
func (a *Activities) RefreshMemoryActivity(ctx context.Context) (map[string]interface{}, error) {
|
||||
logger := activity.GetLogger(ctx)
|
||||
|
||||
logger.Info("Refreshing memory context")
|
||||
|
||||
vault, err := a.service.GetVault(ctx)
|
||||
if err != nil {
|
||||
logger.Error("Memory refresh failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
healthy := a.service.IsHealthy(ctx)
|
||||
|
||||
result := map[string]interface{}{
|
||||
"vault_files": len(vault),
|
||||
"healthy": healthy,
|
||||
}
|
||||
|
||||
logger.Info("Memory refreshed", "vault_files", len(vault), "healthy", healthy)
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,351 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"go.temporal.io/sdk/testsuite"
|
||||
)
|
||||
|
||||
func TestActivityCreateKnowledge(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(IngestResponse{ID: "chunk-123"})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
suite := &testsuite.WorkflowTestSuite{}
|
||||
env := suite.NewTestActivityEnvironment()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
activities := NewActivities(svc)
|
||||
|
||||
env.RegisterActivity(activities.CreateKnowledgeActivity)
|
||||
|
||||
record := &KnowledgeRecord{
|
||||
Level: "L1",
|
||||
Content: "test",
|
||||
}
|
||||
|
||||
result, err := env.ExecuteActivity(activities.CreateKnowledgeActivity, record)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("activity failed: %v", err)
|
||||
}
|
||||
|
||||
var id string
|
||||
if err := result.Get(&id); err != nil {
|
||||
t.Fatalf("get result failed: %v", err)
|
||||
}
|
||||
|
||||
if id != "chunk-123" {
|
||||
t.Errorf("expected chunk-123, got %s", id)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivitySearchKnowledge(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(QueryResponse{
|
||||
Results: []QueryResult{
|
||||
{
|
||||
ID: "chunk-123",
|
||||
Text: "matching knowledge",
|
||||
},
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
suite := &testsuite.WorkflowTestSuite{}
|
||||
env := suite.NewTestActivityEnvironment()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
activities := NewActivities(svc)
|
||||
|
||||
env.RegisterActivity(activities.SearchKnowledgeActivity)
|
||||
|
||||
result, err := env.ExecuteActivity(activities.SearchKnowledgeActivity, "test query", nil)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("activity failed: %v", err)
|
||||
}
|
||||
|
||||
var records []KnowledgeRecord
|
||||
if err := result.Get(&records); err != nil {
|
||||
t.Fatalf("get result failed: %v", err)
|
||||
}
|
||||
|
||||
if len(records) != 1 {
|
||||
t.Errorf("expected 1 record, got %d", len(records))
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityGetContext(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(ContextResponse{
|
||||
Tier: 1,
|
||||
Lessons: []ContextLesson{
|
||||
{
|
||||
Tier: 1,
|
||||
Text: "lesson text",
|
||||
},
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
suite := &testsuite.WorkflowTestSuite{}
|
||||
env := suite.NewTestActivityEnvironment()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
activities := NewActivities(svc)
|
||||
|
||||
env.RegisterActivity(activities.GetContextActivity)
|
||||
|
||||
result, err := env.ExecuteActivity(activities.GetContextActivity, "kubectl", "debug", 8192)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("activity failed: %v", err)
|
||||
}
|
||||
|
||||
var ctx *ServiceContext
|
||||
if err := result.Get(&ctx); err != nil {
|
||||
t.Fatalf("get result failed: %v", err)
|
||||
}
|
||||
|
||||
if ctx.Tier != 1 {
|
||||
t.Errorf("expected tier 1, got %d", ctx.Tier)
|
||||
}
|
||||
|
||||
if len(ctx.Lessons) != 1 {
|
||||
t.Errorf("expected 1 lesson, got %d", len(ctx.Lessons))
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityDiagnoseIssue(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(ContextResponse{
|
||||
Tier: 1,
|
||||
Lessons: []ContextLesson{
|
||||
{
|
||||
Tier: 1,
|
||||
Text: "diagnosis: check logs",
|
||||
},
|
||||
},
|
||||
Skills: []ContextSkill{
|
||||
{
|
||||
Name: "debug-skill",
|
||||
Why: "matched",
|
||||
},
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
suite := &testsuite.WorkflowTestSuite{}
|
||||
env := suite.NewTestActivityEnvironment()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
activities := NewActivities(svc)
|
||||
|
||||
env.RegisterActivity(activities.DiagnoseIssueActivity)
|
||||
|
||||
result, err := env.ExecuteActivity(activities.DiagnoseIssueActivity, "kubectl", "pod-crash")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("activity failed: %v", err)
|
||||
}
|
||||
|
||||
var recommendations []string
|
||||
if err := result.Get(&recommendations); err != nil {
|
||||
t.Fatalf("get result failed: %v", err)
|
||||
}
|
||||
|
||||
if len(recommendations) == 0 {
|
||||
t.Error("expected recommendations")
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityAnalyzeError(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(QueryResponse{
|
||||
Results: []QueryResult{
|
||||
{
|
||||
ID: "chunk-123",
|
||||
Level: "L1",
|
||||
Text: "solution: restart pod",
|
||||
},
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
suite := &testsuite.WorkflowTestSuite{}
|
||||
env := suite.NewTestActivityEnvironment()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
activities := NewActivities(svc)
|
||||
|
||||
env.RegisterActivity(activities.AnalyzeErrorActivity)
|
||||
|
||||
result, err := env.ExecuteActivity(activities.AnalyzeErrorActivity, "CrashLoopBackOff")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("activity failed: %v", err)
|
||||
}
|
||||
|
||||
var records []KnowledgeRecord
|
||||
if err := result.Get(&records); err != nil {
|
||||
t.Fatalf("get result failed: %v", err)
|
||||
}
|
||||
|
||||
if len(records) != 1 {
|
||||
t.Errorf("expected 1 record, got %d", len(records))
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityHealthCheck(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
suite := &testsuite.WorkflowTestSuite{}
|
||||
env := suite.NewTestActivityEnvironment()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
activities := NewActivities(svc)
|
||||
|
||||
env.RegisterActivity(activities.HealthCheckActivity)
|
||||
|
||||
result, err := env.ExecuteActivity(activities.HealthCheckActivity)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("activity failed: %v", err)
|
||||
}
|
||||
|
||||
var healthy bool
|
||||
if err := result.Get(&healthy); err != nil {
|
||||
t.Fatalf("get result failed: %v", err)
|
||||
}
|
||||
|
||||
if !healthy {
|
||||
t.Error("expected healthy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityLearnFromExecution(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(IngestResponse{ID: "chunk-456"})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
suite := &testsuite.WorkflowTestSuite{}
|
||||
env := suite.NewTestActivityEnvironment()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
activities := NewActivities(svc)
|
||||
|
||||
env.RegisterActivity(activities.LearnFromExecutionActivity)
|
||||
|
||||
result, err := env.ExecuteActivity(activities.LearnFromExecutionActivity, "task-123", "success", []string{"tag1"})
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("activity failed: %v", err)
|
||||
}
|
||||
|
||||
var id string
|
||||
if err := result.Get(&id); err != nil {
|
||||
t.Fatalf("get result failed: %v", err)
|
||||
}
|
||||
|
||||
if id != "chunk-456" {
|
||||
t.Errorf("expected chunk-456, got %s", id)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityDocumentDecision(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(IngestResponse{ID: "chunk-789"})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
suite := &testsuite.WorkflowTestSuite{}
|
||||
env := suite.NewTestActivityEnvironment()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
activities := NewActivities(svc)
|
||||
|
||||
env.RegisterActivity(activities.DocumentDecisionActivity)
|
||||
|
||||
result, err := env.ExecuteActivity(
|
||||
activities.DocumentDecisionActivity,
|
||||
"scaling",
|
||||
"scale to 5 replicas",
|
||||
"high CPU usage",
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("activity failed: %v", err)
|
||||
}
|
||||
|
||||
var id string
|
||||
if err := result.Get(&id); err != nil {
|
||||
t.Fatalf("get result failed: %v", err)
|
||||
}
|
||||
|
||||
if id != "chunk-789" {
|
||||
t.Errorf("expected chunk-789, got %s", id)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityOptions(t *testing.T) {
|
||||
opts := DefaultActivityOptions()
|
||||
|
||||
if opts.RetryAttempts != 3 {
|
||||
t.Errorf("expected 3 retry attempts, got %d", opts.RetryAttempts)
|
||||
}
|
||||
|
||||
if opts.StartTimeout == 0 {
|
||||
t.Error("expected non-zero start timeout")
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityError(t *testing.T) {
|
||||
err := &MemoryActivityError{
|
||||
ActivityName: "test-activity",
|
||||
Attempt: 2,
|
||||
Err: context.Canceled,
|
||||
}
|
||||
|
||||
msg := err.Error()
|
||||
if msg == "" {
|
||||
t.Error("expected error message")
|
||||
}
|
||||
|
||||
if !contains(msg, "test-activity") {
|
||||
t.Error("expected activity name in error")
|
||||
}
|
||||
|
||||
if !contains(msg, "attempt 2") {
|
||||
t.Error("expected attempt number in error")
|
||||
}
|
||||
}
|
||||
|
||||
func contains(s, substr string) bool {
|
||||
for i := 0; i < len(s)-len(substr)+1; i++ {
|
||||
if s[i:i+len(substr)] == substr {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,296 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Client memory service client with JWT auth
|
||||
type Client struct {
|
||||
baseURL string
|
||||
httpClient *http.Client
|
||||
token string
|
||||
}
|
||||
|
||||
// NewClient creates memory service client
|
||||
func NewClient(baseURL, token string) *Client {
|
||||
return &Client{
|
||||
baseURL: baseURL,
|
||||
httpClient: &http.Client{
|
||||
Timeout: 10 * time.Second,
|
||||
},
|
||||
token: token,
|
||||
}
|
||||
}
|
||||
|
||||
// IngestRequest ingest knowledge record
|
||||
type IngestRequest struct {
|
||||
Project string `json:"project"`
|
||||
Source string `json:"source"`
|
||||
Kind string `json:"kind"` // L1|L2|reference
|
||||
Text string `json:"text"`
|
||||
Metadata map[string]interface{} `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
// IngestResponse ingest response
|
||||
type IngestResponse struct {
|
||||
ID string `json:"id"`
|
||||
SHA256 string `json:"sha256"`
|
||||
QueueStatus string `json:"queue_status"`
|
||||
IdempotencyID string `json:"idempotency_key"`
|
||||
}
|
||||
|
||||
// Ingest creates knowledge record
|
||||
func (c *Client) Ingest(ctx context.Context, req *IngestRequest) (*IngestResponse, error) {
|
||||
body, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal ingest request: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, "POST", c.baseURL+"/memory/ingest", bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
|
||||
c.setAuthHeader(httpReq)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := c.httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ingest request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusCreated && resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("ingest failed (%d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var result IngestResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||
return nil, fmt.Errorf("decode ingest response: %w", err)
|
||||
}
|
||||
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// QueryRequest query memory
|
||||
type QueryRequest struct {
|
||||
Project string `json:"project"`
|
||||
Query string `json:"query"`
|
||||
LevelFilter []string `json:"level_filter,omitempty"` // L1, L2, R
|
||||
Floor float32 `json:"floor,omitempty"`
|
||||
Limit int `json:"limit,omitempty"`
|
||||
Scope string `json:"scope,omitempty"` // learned|reference|all
|
||||
}
|
||||
|
||||
// QueryResult single search result
|
||||
type QueryResult struct {
|
||||
ID string `json:"id"`
|
||||
Level string `json:"level"`
|
||||
Score float32 `json:"score"`
|
||||
SemanticScore float32 `json:"semantic_score"`
|
||||
LexicalScore float32 `json:"lexical_score"`
|
||||
Text string `json:"text"`
|
||||
Breadcrumb string `json:"breadcrumb"`
|
||||
Source string `json:"source"`
|
||||
}
|
||||
|
||||
// QueryResponse query response
|
||||
type QueryResponse struct {
|
||||
Query string `json:"query"`
|
||||
Results []QueryResult `json:"results"`
|
||||
TotalHits int `json:"total_hits"`
|
||||
SearchTimeMS int `json:"search_time_ms"`
|
||||
}
|
||||
|
||||
// Query searches knowledge
|
||||
func (c *Client) Query(ctx context.Context, req *QueryRequest) (*QueryResponse, error) {
|
||||
if req.Limit == 0 {
|
||||
req.Limit = 10
|
||||
}
|
||||
|
||||
body, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal query request: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, "POST", c.baseURL+"/memory/query", bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
|
||||
c.setAuthHeader(httpReq)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := c.httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("query request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("query failed (%d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var result QueryResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||
return nil, fmt.Errorf("decode query response: %w", err)
|
||||
}
|
||||
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// ContextRequest retrieve context (three-tier)
|
||||
type ContextRequest struct {
|
||||
Project string `json:"project"`
|
||||
Tool string `json:"tool"`
|
||||
Task string `json:"task"`
|
||||
SignatureSource string `json:"signature_source"`
|
||||
Scope string `json:"scope,omitempty"` // tool_context
|
||||
Budget int `json:"budget,omitempty"`
|
||||
}
|
||||
|
||||
// ContextLesson lesson from context
|
||||
type ContextLesson struct {
|
||||
Tier int `json:"tier"`
|
||||
Level string `json:"level"`
|
||||
Score float32 `json:"score"`
|
||||
Text string `json:"text"`
|
||||
MatchedKind string `json:"matched_kind,omitempty"`
|
||||
SeenCount int `json:"seen_count,omitempty"`
|
||||
LastSeen string `json:"last_seen,omitempty"`
|
||||
}
|
||||
|
||||
// ContextSkill skill suggestion
|
||||
type ContextSkill struct {
|
||||
Name string `json:"name"`
|
||||
Why string `json:"why"`
|
||||
}
|
||||
|
||||
// ContextBudget budget tracking
|
||||
type ContextBudget struct {
|
||||
Requested int `json:"requested"`
|
||||
Used int `json:"used"`
|
||||
Dropped int `json:"dropped"`
|
||||
Degradation *string `json:"degradation"`
|
||||
}
|
||||
|
||||
// ContextResponse context response
|
||||
type ContextResponse struct {
|
||||
Tier int `json:"tier"`
|
||||
Lessons []ContextLesson `json:"lessons"`
|
||||
Skills []ContextSkill `json:"skills"`
|
||||
Budget ContextBudget `json:"budget"`
|
||||
}
|
||||
|
||||
// Context retrieves context (three-tier retrieval)
|
||||
func (c *Client) Context(ctx context.Context, req *ContextRequest) (*ContextResponse, error) {
|
||||
if req.Budget == 0 {
|
||||
req.Budget = 8192
|
||||
}
|
||||
|
||||
body, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal context request: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, "POST", c.baseURL+"/memory/context", bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
|
||||
c.setAuthHeader(httpReq)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := c.httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("context request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("context failed (%d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var result ContextResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||
return nil, fmt.Errorf("decode context response: %w", err)
|
||||
}
|
||||
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// VaultFile file in vault
|
||||
type VaultFile struct {
|
||||
Path string `json:"path"`
|
||||
Title string `json:"title"`
|
||||
Level string `json:"level"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
RecordCount int `json:"record_count"`
|
||||
}
|
||||
|
||||
// VaultResponse vault browse response
|
||||
type VaultResponse struct {
|
||||
Project string `json:"project"`
|
||||
Files []VaultFile `json:"files"`
|
||||
TotalRecords int `json:"total_records"`
|
||||
}
|
||||
|
||||
// Vault browses vault files
|
||||
func (c *Client) Vault(ctx context.Context, project string) (*VaultResponse, error) {
|
||||
httpReq, err := http.NewRequestWithContext(ctx, "GET", fmt.Sprintf("%s/memory/vault?project=%s", c.baseURL, project), nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
|
||||
c.setAuthHeader(httpReq)
|
||||
|
||||
resp, err := c.httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("vault request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("vault failed (%d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var result VaultResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||
return nil, fmt.Errorf("decode vault response: %w", err)
|
||||
}
|
||||
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// setAuthHeader sets JWT Bearer token
|
||||
func (c *Client) setAuthHeader(req *http.Request) {
|
||||
if c.token != "" {
|
||||
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", c.token))
|
||||
}
|
||||
}
|
||||
|
||||
// Health checks memory service
|
||||
func (c *Client) Health(ctx context.Context) (bool, error) {
|
||||
httpReq, err := http.NewRequestWithContext(ctx, "GET", c.baseURL+"/health", nil)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
resp, err := c.httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
return resp.StatusCode == http.StatusOK, nil
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestClientIngest(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/memory/ingest" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
if r.Header.Get("Authorization") == "" {
|
||||
t.Error("missing Authorization header")
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(IngestResponse{
|
||||
ID: "chunk-123",
|
||||
SHA256: "abc123",
|
||||
QueueStatus: "pending",
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewClient(server.URL, "test-token")
|
||||
resp, err := client.Ingest(context.Background(), &IngestRequest{
|
||||
Project: "poimen",
|
||||
Source: "test",
|
||||
Kind: "L1",
|
||||
Text: "test content",
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("ingest failed: %v", err)
|
||||
}
|
||||
|
||||
if resp.ID != "chunk-123" {
|
||||
t.Errorf("expected ID chunk-123, got %s", resp.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientQuery(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/memory/query" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(QueryResponse{
|
||||
Query: "test query",
|
||||
TotalHits: 1,
|
||||
SearchTimeMS: 100,
|
||||
Results: []QueryResult{
|
||||
{
|
||||
ID: "chunk-123",
|
||||
Level: "L1",
|
||||
Score: 0.95,
|
||||
Text: "matching result",
|
||||
},
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewClient(server.URL, "test-token")
|
||||
resp, err := client.Query(context.Background(), &QueryRequest{
|
||||
Project: "poimen",
|
||||
Query: "test query",
|
||||
Limit: 10,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("query failed: %v", err)
|
||||
}
|
||||
|
||||
if len(resp.Results) != 1 {
|
||||
t.Errorf("expected 1 result, got %d", len(resp.Results))
|
||||
}
|
||||
|
||||
if resp.Results[0].Text != "matching result" {
|
||||
t.Errorf("unexpected result text")
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientContext(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/memory/context" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(ContextResponse{
|
||||
Tier: 1,
|
||||
Lessons: []ContextLesson{
|
||||
{
|
||||
Tier: 1,
|
||||
Level: "L1",
|
||||
Score: 1.0,
|
||||
Text: "tier-1 lesson",
|
||||
},
|
||||
},
|
||||
Budget: ContextBudget{
|
||||
Requested: 8192,
|
||||
Used: 100,
|
||||
Dropped: 0,
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewClient(server.URL, "test-token")
|
||||
resp, err := client.Context(context.Background(), &ContextRequest{
|
||||
Project: "poimen",
|
||||
Tool: "kubectl",
|
||||
Task: "debug",
|
||||
SignatureSource: "log",
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("context failed: %v", err)
|
||||
}
|
||||
|
||||
if resp.Tier != 1 {
|
||||
t.Errorf("expected tier 1, got %d", resp.Tier)
|
||||
}
|
||||
|
||||
if len(resp.Lessons) != 1 {
|
||||
t.Errorf("expected 1 lesson, got %d", len(resp.Lessons))
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientVault(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/memory/vault" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(VaultResponse{
|
||||
Project: "poimen",
|
||||
TotalRecords: 42,
|
||||
Files: []VaultFile{
|
||||
{
|
||||
Path: "test.md",
|
||||
Title: "Test",
|
||||
Level: "L1",
|
||||
RecordCount: 5,
|
||||
},
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewClient(server.URL, "test-token")
|
||||
resp, err := client.Vault(context.Background(), "poimen")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("vault failed: %v", err)
|
||||
}
|
||||
|
||||
if len(resp.Files) != 1 {
|
||||
t.Errorf("expected 1 file, got %d", len(resp.Files))
|
||||
}
|
||||
|
||||
if resp.TotalRecords != 42 {
|
||||
t.Errorf("expected 42 records, got %d", resp.TotalRecords)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientHealth(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/health" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(map[string]string{"status": "ok"})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewClient(server.URL, "test-token")
|
||||
ok, err := client.Health(context.Background())
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("health check failed: %v", err)
|
||||
}
|
||||
|
||||
if !ok {
|
||||
t.Error("expected health check to pass")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// ExampleActivity demonstrates memory service integration in workflows
|
||||
// This can be used as a template for workflow activities
|
||||
|
||||
// LearnTaskActivity learns from task execution
|
||||
func LearnTaskActivity(ctx context.Context, service *Service, task string, result string) error {
|
||||
// Create knowledge from task result
|
||||
knowledgeID, err := service.CreateKnowledge(ctx, &KnowledgeRecord{
|
||||
Level: "L1",
|
||||
Title: fmt.Sprintf("Task: %s", task),
|
||||
Content: result,
|
||||
Source: fmt.Sprintf("workflow://task/%s", task),
|
||||
Metadata: map[string]interface{}{
|
||||
"task": task,
|
||||
"type": "execution_result",
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("learn task: %w", err)
|
||||
}
|
||||
|
||||
fmt.Printf("Learned: %s\n", knowledgeID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// DiagnosticActivity retrieves context for problem diagnosis
|
||||
func DiagnosticActivity(ctx context.Context, service *Service, tool string, issue string) ([]string, error) {
|
||||
// Retrieve context for tool/issue
|
||||
svcCtx, err := service.RetrieveContext(ctx, tool, issue, 8192)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("diagnose: %w", err)
|
||||
}
|
||||
|
||||
// Extract lessons
|
||||
diagnostics := make([]string, 0)
|
||||
for _, lesson := range svcCtx.Lessons {
|
||||
if lesson.Tier == 1 {
|
||||
diagnostics = append(diagnostics, lesson.Text)
|
||||
}
|
||||
}
|
||||
|
||||
// Extract skills
|
||||
for _, skill := range svcCtx.Skills {
|
||||
diagnostics = append(diagnostics, fmt.Sprintf("Skill: %s (%s)", skill.Name, skill.Why))
|
||||
}
|
||||
|
||||
return diagnostics, nil
|
||||
}
|
||||
|
||||
// DocumentationActivity searches vault for relevant docs
|
||||
func DocumentationActivity(ctx context.Context, service *Service, topic string) ([]string, error) {
|
||||
// Retrieve vault files
|
||||
files, err := service.GetVault(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get vault: %w", err)
|
||||
}
|
||||
|
||||
// Filter by topic
|
||||
results := make([]string, 0)
|
||||
for _, file := range files {
|
||||
if file.Level == "R" { // Reference docs
|
||||
results = append(results, fmt.Sprintf("%s: %s", file.Title, file.Path))
|
||||
}
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
// SearchKnowledgeActivity searches knowledge base
|
||||
func SearchKnowledgeActivity(ctx context.Context, service *Service, query string) ([]string, error) {
|
||||
records, err := service.RetrieveKnowledge(ctx, query, &RetrievalOptions{
|
||||
LevelFilter: []string{"L1", "L2"},
|
||||
Limit: 5,
|
||||
Floor: 0.7,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("search: %w", err)
|
||||
}
|
||||
|
||||
results := make([]string, 0)
|
||||
for _, record := range records {
|
||||
results = append(results, fmt.Sprintf("[%s] %s", record.Level, record.Content))
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
// UpdateLessonActivity updates learned facts
|
||||
func UpdateLessonActivity(ctx context.Context, service *Service, id string, newContent string) error {
|
||||
_, err := service.UpdateKnowledge(ctx, &KnowledgeRecord{
|
||||
ID: id,
|
||||
Level: "L2",
|
||||
Content: newContent,
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
// HealthCheckActivity checks memory service health
|
||||
func HealthCheckActivity(ctx context.Context, service *Service) (bool, error) {
|
||||
return service.IsHealthy(ctx), nil
|
||||
}
|
||||
|
||||
// Example workflow structure using memory service
|
||||
type WorkflowWithMemory struct {
|
||||
MemoryService *Service
|
||||
}
|
||||
|
||||
// ExecuteWithLearning executes task and learns from it
|
||||
func (w *WorkflowWithMemory) ExecuteWithLearning(ctx context.Context, task string, executor func() (string, error)) error {
|
||||
// Execute task
|
||||
result, err := executor()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Learn from result
|
||||
return LearnTaskActivity(ctx, w.MemoryService, task, result)
|
||||
}
|
||||
|
||||
// DiagnoseWithContext diagnoses issue using memory context
|
||||
func (w *WorkflowWithMemory) DiagnoseWithContext(ctx context.Context, tool string, issue string) ([]string, error) {
|
||||
return DiagnosticActivity(ctx, w.MemoryService, tool, issue)
|
||||
}
|
||||
|
||||
// SearchKnowledge searches knowledge
|
||||
func (w *WorkflowWithMemory) SearchKnowledge(ctx context.Context, query string) ([]string, error) {
|
||||
return SearchKnowledgeActivity(ctx, w.MemoryService, query)
|
||||
}
|
||||
@@ -0,0 +1,234 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// Service memory service manager
|
||||
type Service struct {
|
||||
client *Client
|
||||
project string
|
||||
}
|
||||
|
||||
// NewService creates memory service manager
|
||||
func NewService(baseURL, token, project string) *Service {
|
||||
return &Service{
|
||||
client: NewClient(baseURL, token),
|
||||
project: project,
|
||||
}
|
||||
}
|
||||
|
||||
// KnowledgeRecord high-level knowledge record
|
||||
type KnowledgeRecord struct {
|
||||
ID string
|
||||
Level string // L1|L2|reference
|
||||
Title string
|
||||
Content string
|
||||
Source string
|
||||
Metadata map[string]interface{}
|
||||
SHA256 string
|
||||
}
|
||||
|
||||
// CreateKnowledge creates knowledge record
|
||||
func (s *Service) CreateKnowledge(ctx context.Context, record *KnowledgeRecord) (string, error) {
|
||||
if record.Level == "" {
|
||||
record.Level = "L1"
|
||||
}
|
||||
if record.Source == "" {
|
||||
record.Source = "workflow"
|
||||
}
|
||||
|
||||
req := &IngestRequest{
|
||||
Project: s.project,
|
||||
Source: record.Source,
|
||||
Kind: record.Level,
|
||||
Text: record.Content,
|
||||
Metadata: record.Metadata,
|
||||
}
|
||||
|
||||
resp, err := s.client.Ingest(ctx, req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("create knowledge: %w", err)
|
||||
}
|
||||
|
||||
return resp.ID, nil
|
||||
}
|
||||
|
||||
// UpdateKnowledge updates existing knowledge (re-ingest)
|
||||
func (s *Service) UpdateKnowledge(ctx context.Context, record *KnowledgeRecord) (string, error) {
|
||||
// Update done by re-ingesting with same signature/source
|
||||
// Memory service deduplicates based on idempotency key
|
||||
if record.Metadata == nil {
|
||||
record.Metadata = make(map[string]interface{})
|
||||
}
|
||||
|
||||
// Use ID as session_id for idempotency
|
||||
record.Metadata["session_id"] = record.ID
|
||||
|
||||
return s.CreateKnowledge(ctx, record)
|
||||
}
|
||||
|
||||
// RetrievalOptions search options
|
||||
type RetrievalOptions struct {
|
||||
LevelFilter []string // L1, L2, R
|
||||
Floor float32 // minimum relevance
|
||||
Limit int // default 10
|
||||
Scope string // learned|reference|all
|
||||
}
|
||||
|
||||
// RetrieveKnowledge searches knowledge
|
||||
func (s *Service) RetrieveKnowledge(ctx context.Context, query string, opts *RetrievalOptions) ([]KnowledgeRecord, error) {
|
||||
if opts == nil {
|
||||
opts = &RetrievalOptions{}
|
||||
}
|
||||
|
||||
if opts.Limit == 0 {
|
||||
opts.Limit = 10
|
||||
}
|
||||
|
||||
req := &QueryRequest{
|
||||
Project: s.project,
|
||||
Query: query,
|
||||
LevelFilter: opts.LevelFilter,
|
||||
Floor: opts.Floor,
|
||||
Limit: opts.Limit,
|
||||
Scope: opts.Scope,
|
||||
}
|
||||
|
||||
resp, err := s.client.Query(ctx, req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("retrieve knowledge: %w", err)
|
||||
}
|
||||
|
||||
records := make([]KnowledgeRecord, len(resp.Results))
|
||||
for i, r := range resp.Results {
|
||||
records[i] = KnowledgeRecord{
|
||||
ID: r.ID,
|
||||
Level: r.Level,
|
||||
Content: r.Text,
|
||||
Source: r.Source,
|
||||
SHA256: "", // Not in response
|
||||
Metadata: map[string]interface{}{
|
||||
"score": r.Score,
|
||||
"semantic_score": r.SemanticScore,
|
||||
"lexical_score": r.LexicalScore,
|
||||
"breadcrumb": r.Breadcrumb,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
return records, nil
|
||||
}
|
||||
|
||||
// ServiceContext tool/task context
|
||||
type ServiceContext struct {
|
||||
Tier int
|
||||
Lessons []Lesson
|
||||
Skills []Skill
|
||||
BudgetUsed int
|
||||
BudgetMax int
|
||||
}
|
||||
|
||||
// Lesson learned fact or reference
|
||||
type Lesson struct {
|
||||
Tier int
|
||||
Level string
|
||||
Score float32
|
||||
Text string
|
||||
MatchedKind string
|
||||
SeenCount int
|
||||
LastSeen string
|
||||
}
|
||||
|
||||
// Skill recommended action
|
||||
type Skill struct {
|
||||
Name string
|
||||
Why string
|
||||
}
|
||||
|
||||
// RetrieveContext retrieves context for tool/task (three-tier)
|
||||
func (s *Service) RetrieveContext(ctx context.Context, tool, task string, budget int) (*ServiceContext, error) {
|
||||
if budget == 0 {
|
||||
budget = 8192
|
||||
}
|
||||
|
||||
req := &ContextRequest{
|
||||
Project: s.project,
|
||||
Tool: tool,
|
||||
Task: task,
|
||||
SignatureSource: fmt.Sprintf("%s:%s", tool, task),
|
||||
Scope: "tool_context",
|
||||
Budget: budget,
|
||||
}
|
||||
|
||||
resp, err := s.client.Context(ctx, req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("retrieve context: %w", err)
|
||||
}
|
||||
|
||||
lessons := make([]Lesson, len(resp.Lessons))
|
||||
for i, l := range resp.Lessons {
|
||||
lessons[i] = Lesson{
|
||||
Tier: l.Tier,
|
||||
Level: l.Level,
|
||||
Score: l.Score,
|
||||
Text: l.Text,
|
||||
MatchedKind: l.MatchedKind,
|
||||
SeenCount: l.SeenCount,
|
||||
LastSeen: l.LastSeen,
|
||||
}
|
||||
}
|
||||
|
||||
skills := make([]Skill, len(resp.Skills))
|
||||
for i, sk := range resp.Skills {
|
||||
skills[i] = Skill{
|
||||
Name: sk.Name,
|
||||
Why: sk.Why,
|
||||
}
|
||||
}
|
||||
|
||||
return &ServiceContext{
|
||||
Tier: resp.Tier,
|
||||
Lessons: lessons,
|
||||
Skills: skills,
|
||||
BudgetUsed: resp.Budget.Used,
|
||||
BudgetMax: resp.Budget.Requested,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// VaultInfo vault browsing
|
||||
type VaultInfo struct {
|
||||
Path string
|
||||
Title string
|
||||
Level string
|
||||
UpdatedAt string
|
||||
RecordCount int
|
||||
}
|
||||
|
||||
// GetVault lists vault files
|
||||
func (s *Service) GetVault(ctx context.Context) ([]VaultInfo, error) {
|
||||
resp, err := s.client.Vault(ctx, s.project)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get vault: %w", err)
|
||||
}
|
||||
|
||||
files := make([]VaultInfo, len(resp.Files))
|
||||
for i, f := range resp.Files {
|
||||
files[i] = VaultInfo{
|
||||
Path: f.Path,
|
||||
Title: f.Title,
|
||||
Level: f.Level,
|
||||
UpdatedAt: f.UpdatedAt,
|
||||
RecordCount: f.RecordCount,
|
||||
}
|
||||
}
|
||||
|
||||
return files, nil
|
||||
}
|
||||
|
||||
// IsHealthy checks service health
|
||||
func (s *Service) IsHealthy(ctx context.Context) bool {
|
||||
ok, err := s.client.Health(ctx)
|
||||
return ok && err == nil
|
||||
}
|
||||
@@ -0,0 +1,254 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestServiceCreateKnowledge(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/memory/ingest" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
var req IngestRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
t.Fatalf("decode request: %v", err)
|
||||
}
|
||||
|
||||
if req.Project != "poimen" {
|
||||
t.Errorf("expected project poimen, got %s", req.Project)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(IngestResponse{
|
||||
ID: "chunk-123",
|
||||
QueueStatus: "pending",
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
id, err := svc.CreateKnowledge(context.Background(), &KnowledgeRecord{
|
||||
Content: "test knowledge",
|
||||
Level: "L1",
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("create knowledge failed: %v", err)
|
||||
}
|
||||
|
||||
if id != "chunk-123" {
|
||||
t.Errorf("expected ID chunk-123, got %s", id)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceRetrieveKnowledge(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/memory/query" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(QueryResponse{
|
||||
Query: "test",
|
||||
Results: []QueryResult{
|
||||
{
|
||||
ID: "chunk-123",
|
||||
Level: "L1",
|
||||
Score: 0.95,
|
||||
SemanticScore: 0.96,
|
||||
LexicalScore: 0.94,
|
||||
Text: "knowledge content",
|
||||
Breadcrumb: "path > to > doc",
|
||||
Source: "test",
|
||||
},
|
||||
},
|
||||
TotalHits: 1,
|
||||
SearchTimeMS: 50,
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
records, err := svc.RetrieveKnowledge(context.Background(), "test", nil)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("retrieve knowledge failed: %v", err)
|
||||
}
|
||||
|
||||
if len(records) != 1 {
|
||||
t.Errorf("expected 1 record, got %d", len(records))
|
||||
}
|
||||
|
||||
if records[0].Content != "knowledge content" {
|
||||
t.Errorf("unexpected content")
|
||||
}
|
||||
|
||||
meta := records[0].Metadata
|
||||
if score, ok := meta["score"].(float32); ok {
|
||||
if score != 0.95 {
|
||||
t.Errorf("expected score 0.95, got %f", score)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceRetrieveContext(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/memory/context" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(ContextResponse{
|
||||
Tier: 1,
|
||||
Lessons: []ContextLesson{
|
||||
{
|
||||
Tier: 1,
|
||||
Level: "L1",
|
||||
Score: 1.0,
|
||||
Text: "first lesson",
|
||||
MatchedKind: "signature",
|
||||
},
|
||||
{
|
||||
Tier: 2,
|
||||
Level: "L2",
|
||||
Score: 0.87,
|
||||
Text: "second lesson",
|
||||
},
|
||||
},
|
||||
Skills: []ContextSkill{
|
||||
{
|
||||
Name: "debug-skill",
|
||||
Why: "tier 1 matched",
|
||||
},
|
||||
},
|
||||
Budget: ContextBudget{
|
||||
Requested: 8192,
|
||||
Used: 2048,
|
||||
Dropped: 0,
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
ctx, err := svc.RetrieveContext(context.Background(), "kubectl", "debug-pod", 8192)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("retrieve context failed: %v", err)
|
||||
}
|
||||
|
||||
if ctx.Tier != 1 {
|
||||
t.Errorf("expected tier 1, got %d", ctx.Tier)
|
||||
}
|
||||
|
||||
if len(ctx.Lessons) != 2 {
|
||||
t.Errorf("expected 2 lessons, got %d", len(ctx.Lessons))
|
||||
}
|
||||
|
||||
if len(ctx.Skills) != 1 {
|
||||
t.Errorf("expected 1 skill, got %d", len(ctx.Skills))
|
||||
}
|
||||
|
||||
if ctx.BudgetUsed != 2048 {
|
||||
t.Errorf("expected budget used 2048, got %d", ctx.BudgetUsed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceGetVault(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/memory/vault" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
json.NewEncoder(w).Encode(VaultResponse{
|
||||
Project: "poimen",
|
||||
TotalRecords: 100,
|
||||
Files: []VaultFile{
|
||||
{
|
||||
Path: "docs/guide.md",
|
||||
Title: "Guide",
|
||||
Level: "L1",
|
||||
RecordCount: 25,
|
||||
},
|
||||
{
|
||||
Path: "reference/api.md",
|
||||
Title: "API",
|
||||
Level: "R",
|
||||
RecordCount: 75,
|
||||
},
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
files, err := svc.GetVault(context.Background())
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("get vault failed: %v", err)
|
||||
}
|
||||
|
||||
if len(files) != 2 {
|
||||
t.Errorf("expected 2 files, got %d", len(files))
|
||||
}
|
||||
|
||||
if files[0].Title != "Guide" {
|
||||
t.Errorf("expected title Guide, got %s", files[0].Title)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceIsHealthy(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
if !svc.IsHealthy(context.Background()) {
|
||||
t.Error("expected service to be healthy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceUpdateKnowledge(t *testing.T) {
|
||||
callCount := 0
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
callCount++
|
||||
if r.URL.Path != "/memory/ingest" {
|
||||
t.Errorf("unexpected path: %s", r.URL.Path)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(IngestResponse{
|
||||
ID: "chunk-123",
|
||||
QueueStatus: "pending",
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
svc := NewService(server.URL, "test-token", "poimen")
|
||||
record := &KnowledgeRecord{
|
||||
ID: "chunk-123",
|
||||
Content: "updated knowledge",
|
||||
Level: "L1",
|
||||
}
|
||||
|
||||
id, err := svc.UpdateKnowledge(context.Background(), record)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("update knowledge failed: %v", err)
|
||||
}
|
||||
|
||||
if id != "chunk-123" {
|
||||
t.Errorf("expected ID chunk-123, got %s", id)
|
||||
}
|
||||
|
||||
if callCount != 1 {
|
||||
t.Errorf("expected 1 call, got %d", callCount)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,333 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"go.temporal.io/sdk/activity"
|
||||
"go.temporal.io/sdk/temporal"
|
||||
"go.temporal.io/sdk/worker"
|
||||
"go.temporal.io/sdk/workflow"
|
||||
)
|
||||
|
||||
// RegisterMemoryActivities registers all memory service activities with worker
|
||||
func RegisterMemoryActivities(w worker.Worker, service *Service) {
|
||||
activities := NewActivities(service)
|
||||
|
||||
// Register activities (activity name = "ActivityName" → "activityName")
|
||||
w.RegisterActivity(activities.CreateKnowledgeActivity)
|
||||
w.RegisterActivity(activities.UpdateKnowledgeActivity)
|
||||
w.RegisterActivity(activities.SearchKnowledgeActivity)
|
||||
w.RegisterActivity(activities.GetContextActivity)
|
||||
w.RegisterActivity(activities.GetVaultActivity)
|
||||
w.RegisterActivity(activities.HealthCheckActivity)
|
||||
w.RegisterActivity(activities.LearnFromExecutionActivity)
|
||||
w.RegisterActivity(activities.DiagnoseIssueActivity)
|
||||
w.RegisterActivity(activities.AnalyzeErrorActivity)
|
||||
w.RegisterActivity(activities.DocumentDecisionActivity)
|
||||
w.RegisterActivity(activities.SearchAndApplyActivity)
|
||||
w.RegisterActivity(activities.RefreshMemoryActivity)
|
||||
}
|
||||
|
||||
// ActivityOptions memory service activity options
|
||||
type ActivityOptions struct {
|
||||
RetryAttempts int
|
||||
RetryBackoff time.Duration
|
||||
StartTimeout time.Duration
|
||||
HeartbeatRate time.Duration
|
||||
}
|
||||
|
||||
// DefaultActivityOptions returns sensible defaults
|
||||
func DefaultActivityOptions() *ActivityOptions {
|
||||
return &ActivityOptions{
|
||||
RetryAttempts: 3,
|
||||
RetryBackoff: time.Second,
|
||||
StartTimeout: 30 * time.Second,
|
||||
HeartbeatRate: 10 * time.Second,
|
||||
}
|
||||
}
|
||||
|
||||
// ExecuteCreateKnowledge wrapper for CreateKnowledgeActivity
|
||||
func ExecuteCreateKnowledge(
|
||||
ctx workflow.Context,
|
||||
record *KnowledgeRecord,
|
||||
opts *ActivityOptions,
|
||||
) (string, error) {
|
||||
if opts == nil {
|
||||
opts = DefaultActivityOptions()
|
||||
}
|
||||
|
||||
activityCtx := workflow.WithActivityOptions(ctx, workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: 2 * time.Minute,
|
||||
StartToCloseTimeout: time.Minute,
|
||||
RetryPolicy: &temporal.RetryPolicy{
|
||||
InitialInterval: opts.RetryBackoff,
|
||||
BackoffCoefficient: 2.0,
|
||||
MaximumInterval: 30 * time.Second,
|
||||
MaximumAttempts: int32(opts.RetryAttempts),
|
||||
NonRetryableErrorTypes: []string{},
|
||||
},
|
||||
})
|
||||
|
||||
var result string
|
||||
err := workflow.ExecuteActivity(activityCtx, "CreateKnowledgeActivity", record).Get(activityCtx, &result)
|
||||
return result, err
|
||||
}
|
||||
|
||||
// ExecuteSearchKnowledge wrapper for SearchKnowledgeActivity
|
||||
func ExecuteSearchKnowledge(
|
||||
ctx workflow.Context,
|
||||
query string,
|
||||
opts *RetrievalOptions,
|
||||
activityOpts *ActivityOptions,
|
||||
) ([]KnowledgeRecord, error) {
|
||||
if activityOpts == nil {
|
||||
activityOpts = DefaultActivityOptions()
|
||||
}
|
||||
|
||||
activityCtx := workflow.WithActivityOptions(ctx, workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: 3 * time.Minute,
|
||||
StartToCloseTimeout: 2 * time.Minute,
|
||||
RetryPolicy: &temporal.RetryPolicy{
|
||||
InitialInterval: activityOpts.RetryBackoff,
|
||||
BackoffCoefficient: 2.0,
|
||||
MaximumInterval: 30 * time.Second,
|
||||
MaximumAttempts: int32(activityOpts.RetryAttempts),
|
||||
},
|
||||
})
|
||||
|
||||
var result []KnowledgeRecord
|
||||
err := workflow.ExecuteActivity(activityCtx, "SearchKnowledgeActivity", query, opts).Get(activityCtx, &result)
|
||||
return result, err
|
||||
}
|
||||
|
||||
// ExecuteGetContext wrapper for GetContextActivity
|
||||
func ExecuteGetContext(
|
||||
ctx workflow.Context,
|
||||
tool, task string,
|
||||
budget int,
|
||||
opts *ActivityOptions,
|
||||
) (*ServiceContext, error) {
|
||||
if opts == nil {
|
||||
opts = DefaultActivityOptions()
|
||||
}
|
||||
|
||||
activityCtx := workflow.WithActivityOptions(ctx, workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: 3 * time.Minute,
|
||||
StartToCloseTimeout: 2 * time.Minute,
|
||||
RetryPolicy: &temporal.RetryPolicy{
|
||||
InitialInterval: opts.RetryBackoff,
|
||||
BackoffCoefficient: 2.0,
|
||||
MaximumInterval: 30 * time.Second,
|
||||
MaximumAttempts: int32(opts.RetryAttempts),
|
||||
},
|
||||
})
|
||||
|
||||
var result *ServiceContext
|
||||
err := workflow.ExecuteActivity(activityCtx, "GetContextActivity", tool, task, budget).Get(activityCtx, &result)
|
||||
return result, err
|
||||
}
|
||||
|
||||
// ExecuteDiagnoseIssue wrapper for DiagnoseIssueActivity
|
||||
func ExecuteDiagnoseIssue(
|
||||
ctx workflow.Context,
|
||||
tool, issue string,
|
||||
opts *ActivityOptions,
|
||||
) ([]string, error) {
|
||||
if opts == nil {
|
||||
opts = DefaultActivityOptions()
|
||||
}
|
||||
|
||||
activityCtx := workflow.WithActivityOptions(ctx, workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: 2 * time.Minute,
|
||||
StartToCloseTimeout: time.Minute,
|
||||
RetryPolicy: &temporal.RetryPolicy{
|
||||
InitialInterval: opts.RetryBackoff,
|
||||
BackoffCoefficient: 2.0,
|
||||
MaximumInterval: 30 * time.Second,
|
||||
MaximumAttempts: int32(opts.RetryAttempts),
|
||||
},
|
||||
})
|
||||
|
||||
var result []string
|
||||
err := workflow.ExecuteActivity(activityCtx, "DiagnoseIssueActivity", tool, issue).Get(activityCtx, &result)
|
||||
return result, err
|
||||
}
|
||||
|
||||
// ExecuteAnalyzeError wrapper for AnalyzeErrorActivity
|
||||
func ExecuteAnalyzeError(
|
||||
ctx workflow.Context,
|
||||
errorMsg string,
|
||||
opts *ActivityOptions,
|
||||
) ([]KnowledgeRecord, error) {
|
||||
if opts == nil {
|
||||
opts = DefaultActivityOptions()
|
||||
}
|
||||
|
||||
activityCtx := workflow.WithActivityOptions(ctx, workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: 2 * time.Minute,
|
||||
StartToCloseTimeout: time.Minute,
|
||||
RetryPolicy: &temporal.RetryPolicy{
|
||||
InitialInterval: opts.RetryBackoff,
|
||||
BackoffCoefficient: 2.0,
|
||||
MaximumInterval: 30 * time.Second,
|
||||
MaximumAttempts: int32(opts.RetryAttempts),
|
||||
},
|
||||
})
|
||||
|
||||
var result []KnowledgeRecord
|
||||
err := workflow.ExecuteActivity(activityCtx, "AnalyzeErrorActivity", errorMsg).Get(activityCtx, &result)
|
||||
return result, err
|
||||
}
|
||||
|
||||
// ExecuteHealthCheck wrapper for HealthCheckActivity
|
||||
func ExecuteHealthCheck(
|
||||
ctx workflow.Context,
|
||||
opts *ActivityOptions,
|
||||
) (bool, error) {
|
||||
if opts == nil {
|
||||
opts = DefaultActivityOptions()
|
||||
}
|
||||
|
||||
activityCtx := workflow.WithActivityOptions(ctx, workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: 1 * time.Minute,
|
||||
StartToCloseTimeout: 30 * time.Second,
|
||||
RetryPolicy: &temporal.RetryPolicy{
|
||||
InitialInterval: opts.RetryBackoff,
|
||||
BackoffCoefficient: 2.0,
|
||||
MaximumInterval: 15 * time.Second,
|
||||
MaximumAttempts: int32(opts.RetryAttempts),
|
||||
},
|
||||
})
|
||||
|
||||
var result bool
|
||||
err := workflow.ExecuteActivity(activityCtx, "HealthCheckActivity").Get(activityCtx, &result)
|
||||
return result, err
|
||||
}
|
||||
|
||||
// ExecuteLearnFromExecution wrapper for LearnFromExecutionActivity
|
||||
func ExecuteLearnFromExecution(
|
||||
ctx workflow.Context,
|
||||
taskID, result string,
|
||||
tags []string,
|
||||
opts *ActivityOptions,
|
||||
) (string, error) {
|
||||
if opts == nil {
|
||||
opts = DefaultActivityOptions()
|
||||
}
|
||||
|
||||
activityCtx := workflow.WithActivityOptions(ctx, workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: 2 * time.Minute,
|
||||
StartToCloseTimeout: time.Minute,
|
||||
RetryPolicy: &temporal.RetryPolicy{
|
||||
InitialInterval: opts.RetryBackoff,
|
||||
BackoffCoefficient: 2.0,
|
||||
MaximumInterval: 30 * time.Second,
|
||||
MaximumAttempts: int32(opts.RetryAttempts),
|
||||
},
|
||||
})
|
||||
|
||||
var recordID string
|
||||
err := workflow.ExecuteActivity(activityCtx, "LearnFromExecutionActivity", taskID, result, tags).Get(activityCtx, &recordID)
|
||||
return recordID, err
|
||||
}
|
||||
|
||||
// ExecuteDocumentDecision wrapper for DocumentDecisionActivity
|
||||
func ExecuteDocumentDecision(
|
||||
ctx workflow.Context,
|
||||
decisionType, decision, reasoning string,
|
||||
opts *ActivityOptions,
|
||||
) (string, error) {
|
||||
if opts == nil {
|
||||
opts = DefaultActivityOptions()
|
||||
}
|
||||
|
||||
activityCtx := workflow.WithActivityOptions(ctx, workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: 2 * time.Minute,
|
||||
StartToCloseTimeout: time.Minute,
|
||||
RetryPolicy: &temporal.RetryPolicy{
|
||||
InitialInterval: opts.RetryBackoff,
|
||||
BackoffCoefficient: 2.0,
|
||||
MaximumInterval: 30 * time.Second,
|
||||
MaximumAttempts: int32(opts.RetryAttempts),
|
||||
},
|
||||
})
|
||||
|
||||
var recordID string
|
||||
err := workflow.ExecuteActivity(activityCtx, "DocumentDecisionActivity", decisionType, decision, reasoning).Get(activityCtx, &recordID)
|
||||
return recordID, err
|
||||
}
|
||||
|
||||
// ExecuteRefreshMemory wrapper for RefreshMemoryActivity
|
||||
func ExecuteRefreshMemory(
|
||||
ctx workflow.Context,
|
||||
opts *ActivityOptions,
|
||||
) (map[string]interface{}, error) {
|
||||
if opts == nil {
|
||||
opts = DefaultActivityOptions()
|
||||
}
|
||||
|
||||
activityCtx := workflow.WithActivityOptions(ctx, workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: 2 * time.Minute,
|
||||
StartToCloseTimeout: time.Minute,
|
||||
RetryPolicy: &temporal.RetryPolicy{
|
||||
InitialInterval: opts.RetryBackoff,
|
||||
BackoffCoefficient: 2.0,
|
||||
MaximumInterval: 30 * time.Second,
|
||||
MaximumAttempts: int32(opts.RetryAttempts),
|
||||
},
|
||||
})
|
||||
|
||||
var result map[string]interface{}
|
||||
err := workflow.ExecuteActivity(activityCtx, "RefreshMemoryActivity").Get(activityCtx, &result)
|
||||
return result, err
|
||||
}
|
||||
|
||||
// HeartbeatMemoryActivity sends heartbeat every N seconds
|
||||
// Usage: Long-running memory operations
|
||||
func HeartbeatMemoryActivity(ctx context.Context, maxDuration time.Duration) error {
|
||||
ticker := time.NewTicker(10 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
deadline := time.Now().Add(maxDuration)
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
activity.RecordHeartbeat(ctx, time.Now())
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
|
||||
if time.Now().After(deadline) {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// MemoryActivityError wraps errors with activity context
|
||||
type MemoryActivityError struct {
|
||||
ActivityName string
|
||||
Attempt int
|
||||
Err error
|
||||
}
|
||||
|
||||
func (e *MemoryActivityError) Error() string {
|
||||
return fmt.Sprintf("memory activity %s (attempt %d): %v", e.ActivityName, e.Attempt, e.Err)
|
||||
}
|
||||
|
||||
// CaptureActivityError captures activity execution errors
|
||||
func CaptureActivityError(activityName string, err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &MemoryActivityError{
|
||||
ActivityName: activityName,
|
||||
Attempt: 1,
|
||||
Err: err,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,303 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"go.temporal.io/sdk/workflow"
|
||||
)
|
||||
|
||||
// LearningWorkflow learns from task execution
|
||||
// Pattern: Execute → Learn → Document
|
||||
func LearningWorkflow(ctx workflow.Context, taskID string, executor string) (string, error) {
|
||||
// Execute task (placeholder - replace with actual activity)
|
||||
taskResult := fmt.Sprintf("Task %s executed by %s", taskID, executor)
|
||||
|
||||
// Learn from execution
|
||||
knowledgeID, err := ExecuteLearnFromExecution(
|
||||
ctx,
|
||||
taskID,
|
||||
taskResult,
|
||||
[]string{"execution", "learning"},
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("learn from execution: %w", err)
|
||||
}
|
||||
|
||||
return knowledgeID, nil
|
||||
}
|
||||
|
||||
// DiagnosticWorkflow diagnoses issue using memory service
|
||||
// Pattern: Get Context → Extract Recommendations → Apply
|
||||
func DiagnosticWorkflow(ctx workflow.Context, tool, issue string) ([]string, error) {
|
||||
// Get recommendations
|
||||
recommendations, err := ExecuteDiagnoseIssue(
|
||||
ctx,
|
||||
tool,
|
||||
issue,
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("diagnose: %w", err)
|
||||
}
|
||||
|
||||
return recommendations, nil
|
||||
}
|
||||
|
||||
// SearchAndApplyWorkflow searches knowledge and applies it
|
||||
// Pattern: Search → Filter → Apply
|
||||
func SearchAndApplyWorkflow(ctx workflow.Context, query string) ([]KnowledgeRecord, error) {
|
||||
// Search knowledge
|
||||
records, err := ExecuteSearchKnowledge(
|
||||
ctx,
|
||||
query,
|
||||
&RetrievalOptions{
|
||||
Limit: 10,
|
||||
LevelFilter: []string{"L1", "L2"},
|
||||
Floor: 0.7,
|
||||
},
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("search: %w", err)
|
||||
}
|
||||
|
||||
return records, nil
|
||||
}
|
||||
|
||||
// ContextualDecisionWorkflow makes decisions with memory context
|
||||
// Pattern: Get Context → Make Decision → Document Decision
|
||||
func ContextualDecisionWorkflow(ctx workflow.Context, tool, task string, decision string) (string, error) {
|
||||
// Get context
|
||||
svcCtx, err := ExecuteGetContext(
|
||||
ctx,
|
||||
tool,
|
||||
task,
|
||||
8192,
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("get context: %w", err)
|
||||
}
|
||||
|
||||
// Build reasoning from lessons
|
||||
reasoning := fmt.Sprintf("Based on %d lessons from memory service (tier %d)", len(svcCtx.Lessons), svcCtx.Tier)
|
||||
|
||||
// Document decision
|
||||
docID, err := ExecuteDocumentDecision(
|
||||
ctx,
|
||||
tool,
|
||||
decision,
|
||||
reasoning,
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("document decision: %w", err)
|
||||
}
|
||||
|
||||
return docID, nil
|
||||
}
|
||||
|
||||
// ErrorRecoveryWorkflow analyzes error and searches for recovery
|
||||
// Pattern: Error → Analyze → Search → Recover
|
||||
func ErrorRecoveryWorkflow(ctx workflow.Context, errorMsg string) ([]string, error) {
|
||||
// Analyze error
|
||||
records, err := ExecuteAnalyzeError(
|
||||
ctx,
|
||||
errorMsg,
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("analyze error: %w", err)
|
||||
}
|
||||
|
||||
// Extract recovery recommendations
|
||||
recommendations := make([]string, 0)
|
||||
for _, record := range records {
|
||||
if record.Level == "L1" { // High confidence
|
||||
recommendations = append(recommendations, record.Content)
|
||||
}
|
||||
}
|
||||
|
||||
return recommendations, nil
|
||||
}
|
||||
|
||||
// HealthAwareWorkflow checks health before proceeding
|
||||
// Pattern: HealthCheck → Conditional Proceed
|
||||
func HealthAwareWorkflow(ctx workflow.Context, taskID string) (bool, error) {
|
||||
// Check health
|
||||
healthy, err := ExecuteHealthCheck(ctx, DefaultActivityOptions())
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("health check: %w", err)
|
||||
}
|
||||
|
||||
if !healthy {
|
||||
return false, fmt.Errorf("memory service unhealthy, skipping task %s", taskID)
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// IterativeLearnWorkflow learns iteratively
|
||||
// Pattern: Execute → Learn → Refine → Learn Again
|
||||
func IterativeLearnWorkflow(ctx workflow.Context, topic string, iterations int) ([]string, error) {
|
||||
knowledgeIDs := make([]string, 0)
|
||||
|
||||
for i := 0; i < iterations; i++ {
|
||||
// Learn current iteration
|
||||
id, err := ExecuteLearnFromExecution(
|
||||
ctx,
|
||||
fmt.Sprintf("%s-iteration-%d", topic, i+1),
|
||||
fmt.Sprintf("Iteration %d: %s", i+1, topic),
|
||||
[]string{"iteration", fmt.Sprintf("iteration-%d", i+1)},
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("learn iteration %d: %w", i+1, err)
|
||||
}
|
||||
|
||||
knowledgeIDs = append(knowledgeIDs, id)
|
||||
|
||||
// Search for related knowledge
|
||||
records, err := ExecuteSearchKnowledge(
|
||||
ctx,
|
||||
topic,
|
||||
&RetrievalOptions{Limit: 5, Floor: 0.6},
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("search iteration %d: %w", i+1, err)
|
||||
}
|
||||
|
||||
// Log found records
|
||||
if len(records) > 0 {
|
||||
workflow.GetLogger(ctx).Info("Iteration found related records", "iteration", i+1, "records", len(records))
|
||||
}
|
||||
}
|
||||
|
||||
return knowledgeIDs, nil
|
||||
}
|
||||
|
||||
// ConditionalLearningWorkflow learns only on success
|
||||
// Pattern: Execute → If Success → Learn
|
||||
func ConditionalLearningWorkflow(ctx workflow.Context, taskID string, shouldSucceed bool) (string, error) {
|
||||
if !shouldSucceed {
|
||||
return "", fmt.Errorf("task failed, skipping learning")
|
||||
}
|
||||
|
||||
// Only learn on success
|
||||
result := fmt.Sprintf("Task %s succeeded", taskID)
|
||||
|
||||
id, err := ExecuteLearnFromExecution(
|
||||
ctx,
|
||||
taskID,
|
||||
result,
|
||||
[]string{"success"},
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("learn from success: %w", err)
|
||||
}
|
||||
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// MultiStepWorkflow performs multiple memory operations
|
||||
// Pattern: Create → Search → Context → Document
|
||||
func MultiStepWorkflow(ctx workflow.Context, topic string) (map[string]interface{}, error) {
|
||||
results := make(map[string]interface{})
|
||||
|
||||
// Step 1: Create knowledge
|
||||
createID, err := ExecuteCreateKnowledge(
|
||||
ctx,
|
||||
&KnowledgeRecord{
|
||||
Level: "L1",
|
||||
Title: fmt.Sprintf("Initial: %s", topic),
|
||||
Content: fmt.Sprintf("Starting workflow for %s", topic),
|
||||
Source: "workflow://multi-step",
|
||||
},
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create knowledge: %w", err)
|
||||
}
|
||||
results["created"] = createID
|
||||
|
||||
// Step 2: Search knowledge
|
||||
searchRecords, err := ExecuteSearchKnowledge(
|
||||
ctx,
|
||||
topic,
|
||||
&RetrievalOptions{Limit: 5},
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("search knowledge: %w", err)
|
||||
}
|
||||
results["found"] = len(searchRecords)
|
||||
|
||||
// Step 3: Get context
|
||||
svcCtx, err := ExecuteGetContext(
|
||||
ctx,
|
||||
"workflow",
|
||||
topic,
|
||||
8192,
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get context: %w", err)
|
||||
}
|
||||
results["context_tier"] = svcCtx.Tier
|
||||
results["lessons"] = len(svcCtx.Lessons)
|
||||
results["skills"] = len(svcCtx.Skills)
|
||||
|
||||
// Step 4: Document completion
|
||||
docID, err := ExecuteDocumentDecision(
|
||||
ctx,
|
||||
"workflow_completion",
|
||||
fmt.Sprintf("Completed multi-step workflow for %s", topic),
|
||||
fmt.Sprintf("Found %d records, tier %d context", len(searchRecords), svcCtx.Tier),
|
||||
DefaultActivityOptions(),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("document completion: %w", err)
|
||||
}
|
||||
results["documented"] = docID
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
// ParallelLearnWorkflow learns from multiple sources in parallel
|
||||
// Pattern: Execute Multiple Tasks in Parallel → Learn from Each
|
||||
func ParallelLearnWorkflow(ctx workflow.Context, taskIDs []string) ([]string, error) {
|
||||
// Create parallel activities
|
||||
futures := make([]workflow.Future, len(taskIDs))
|
||||
|
||||
for i, taskID := range taskIDs {
|
||||
// Execute each task in parallel
|
||||
future := workflow.ExecuteActivity(
|
||||
workflow.WithActivityOptions(
|
||||
ctx,
|
||||
workflow.ActivityOptions{
|
||||
ScheduleToCloseTimeout: DefaultActivityOptions().RetryBackoff * 60,
|
||||
StartToCloseTimeout: DefaultActivityOptions().RetryBackoff * 30,
|
||||
},
|
||||
),
|
||||
"LearnFromExecutionActivity",
|
||||
taskID,
|
||||
fmt.Sprintf("Result from %s", taskID),
|
||||
[]string{"parallel", taskID},
|
||||
)
|
||||
futures[i] = future
|
||||
}
|
||||
|
||||
// Collect results
|
||||
results := make([]string, len(futures))
|
||||
for i, future := range futures {
|
||||
err := future.Get(ctx, &results[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parallel learn task %d: %w", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
@@ -0,0 +1,267 @@
|
||||
package plugins
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// PluginLoader loads and manages plugin lifecycle
|
||||
type PluginLoader struct {
|
||||
mu sync.RWMutex
|
||||
registry *PluginRegistry
|
||||
pluginPath string
|
||||
loadedPlugins map[string]*LoadedPlugin
|
||||
loadTime map[string]time.Time
|
||||
failedLoads map[string]error
|
||||
}
|
||||
|
||||
// LoadedPlugin represents a loaded plugin with additional metadata
|
||||
type LoadedPlugin struct {
|
||||
Plugin SkillPlugin
|
||||
LoadedAt time.Time
|
||||
ReloadCount int
|
||||
LastError error
|
||||
}
|
||||
|
||||
// PluginConfig represents plugin configuration from file
|
||||
type PluginConfig struct {
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
Author string `json:"author"`
|
||||
Description string `json:"description"`
|
||||
Path string `json:"path"`
|
||||
Config map[string]interface{} `json:"config,omitempty"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
// NewPluginLoader creates a new plugin loader
|
||||
func NewPluginLoader(registry *PluginRegistry, pluginPath string) *PluginLoader {
|
||||
return &PluginLoader{
|
||||
registry: registry,
|
||||
pluginPath: pluginPath,
|
||||
loadedPlugins: make(map[string]*LoadedPlugin),
|
||||
loadTime: make(map[string]time.Time),
|
||||
failedLoads: make(map[string]error),
|
||||
}
|
||||
}
|
||||
|
||||
// LoadPlugin loads a plugin from URL
|
||||
func (pl *PluginLoader) LoadPlugin(url string) error {
|
||||
pl.mu.Lock()
|
||||
defer pl.mu.Unlock()
|
||||
|
||||
if IsPluginURL(url) {
|
||||
// Already loaded - no need to load from file
|
||||
return nil
|
||||
}
|
||||
|
||||
// Try to load from file
|
||||
return pl.loadFromFileLocked(url)
|
||||
}
|
||||
|
||||
// loadFromFileLocked loads a plugin from a file (must be called with lock held)
|
||||
func (pl *PluginLoader) loadFromFileLocked(path string) error {
|
||||
// Read config file
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
pl.failedLoads[path] = err
|
||||
return fmt.Errorf("failed to read plugin config: %w", err)
|
||||
}
|
||||
|
||||
var config PluginConfig
|
||||
if err := json.Unmarshal(data, &config); err != nil {
|
||||
pl.failedLoads[path] = err
|
||||
return fmt.Errorf("failed to parse plugin config: %w", err)
|
||||
}
|
||||
|
||||
// For now, return a placeholder plugin load
|
||||
// In a real implementation, this would use reflection or plugin packages
|
||||
// to dynamically load compiled plugins
|
||||
pl.loadTime[config.Name] = time.Now()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadPluginDirectory loads all plugins from a directory
|
||||
func (pl *PluginLoader) LoadPluginDirectory(directory string) error {
|
||||
entries, err := os.ReadDir(directory)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read plugin directory: %w", err)
|
||||
}
|
||||
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
if filepath.Ext(entry.Name()) == ".json" {
|
||||
pluginPath := filepath.Join(directory, entry.Name())
|
||||
if err := pl.LoadPlugin(pluginPath); err != nil {
|
||||
// Log error but continue loading other plugins
|
||||
pl.mu.Lock()
|
||||
pl.failedLoads[pluginPath] = err
|
||||
pl.mu.Unlock()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RegisterLoadedPlugin registers a loaded plugin with the registry
|
||||
func (pl *PluginLoader) RegisterLoadedPlugin(plugin SkillPlugin, author string, config map[string]interface{}) error {
|
||||
pl.mu.Lock()
|
||||
defer pl.mu.Unlock()
|
||||
|
||||
if err := pl.registry.Register(plugin, author, config); err != nil {
|
||||
pl.failedLoads[plugin.Name()] = err
|
||||
return err
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
pl.loadedPlugins[plugin.Name()] = &LoadedPlugin{
|
||||
Plugin: plugin,
|
||||
LoadedAt: now,
|
||||
LastError: nil,
|
||||
}
|
||||
pl.loadTime[plugin.Name()] = now
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// UnloadPlugin unloads a plugin
|
||||
func (pl *PluginLoader) UnloadPlugin(name string) error {
|
||||
pl.mu.Lock()
|
||||
defer pl.mu.Unlock()
|
||||
|
||||
if _, exists := pl.loadedPlugins[name]; !exists {
|
||||
return fmt.Errorf("plugin not loaded: %s", name)
|
||||
}
|
||||
|
||||
err := pl.registry.Unregister(name)
|
||||
if err == nil {
|
||||
delete(pl.loadedPlugins, name)
|
||||
delete(pl.loadTime, name)
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// ReloadPlugin reloads a plugin
|
||||
func (pl *PluginLoader) ReloadPlugin(name string) error {
|
||||
pl.mu.Lock()
|
||||
defer pl.mu.Unlock()
|
||||
|
||||
loadedPlugin, exists := pl.loadedPlugins[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("plugin not loaded: %s", name)
|
||||
}
|
||||
|
||||
// Re-validate plugin
|
||||
if err := loadedPlugin.Plugin.Validate(); err != nil {
|
||||
pl.failedLoads[name] = err
|
||||
loadedPlugin.LastError = err
|
||||
return fmt.Errorf("plugin validation failed: %w", err)
|
||||
}
|
||||
|
||||
loadedPlugin.ReloadCount++
|
||||
loadedPlugin.LastError = nil
|
||||
pl.loadTime[name] = time.Now()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetLoadedPlugins returns all loaded plugins
|
||||
func (pl *PluginLoader) GetLoadedPlugins() map[string]*LoadedPlugin {
|
||||
pl.mu.RLock()
|
||||
defer pl.mu.RUnlock()
|
||||
|
||||
result := make(map[string]*LoadedPlugin)
|
||||
for name, plugin := range pl.loadedPlugins {
|
||||
result[name] = plugin
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetFailedLoads returns all failed plugin loads
|
||||
func (pl *PluginLoader) GetFailedLoads() map[string]error {
|
||||
pl.mu.RLock()
|
||||
defer pl.mu.RUnlock()
|
||||
|
||||
result := make(map[string]error)
|
||||
for path, err := range pl.failedLoads {
|
||||
result[path] = err
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetLoadTime returns when a plugin was loaded
|
||||
func (pl *PluginLoader) GetLoadTime(name string) (time.Time, bool) {
|
||||
pl.mu.RLock()
|
||||
defer pl.mu.RUnlock()
|
||||
|
||||
t, exists := pl.loadTime[name]
|
||||
return t, exists
|
||||
}
|
||||
|
||||
// IsPluginLoaded checks if a plugin is loaded
|
||||
func (pl *PluginLoader) IsPluginLoaded(name string) bool {
|
||||
pl.mu.RLock()
|
||||
defer pl.mu.RUnlock()
|
||||
|
||||
_, exists := pl.loadedPlugins[name]
|
||||
return exists
|
||||
}
|
||||
|
||||
// ExecutePlugin executes a loaded plugin
|
||||
func (pl *PluginLoader) ExecutePlugin(name string, input map[string]interface{}) (map[string]interface{}, error) {
|
||||
pl.mu.RLock()
|
||||
if _, exists := pl.loadedPlugins[name]; !exists {
|
||||
pl.mu.RUnlock()
|
||||
return nil, fmt.Errorf("plugin not loaded: %s", name)
|
||||
}
|
||||
pl.mu.RUnlock()
|
||||
|
||||
return pl.registry.Execute(name, input)
|
||||
}
|
||||
|
||||
// GetPluginStats returns stats for a loaded plugin
|
||||
func (pl *PluginLoader) GetPluginStats(name string) (*PluginStats, error) {
|
||||
pl.mu.RLock()
|
||||
defer pl.mu.RUnlock()
|
||||
|
||||
if _, exists := pl.loadedPlugins[name]; !exists {
|
||||
return nil, fmt.Errorf("plugin not loaded: %s", name)
|
||||
}
|
||||
|
||||
return pl.registry.GetStats(), nil
|
||||
}
|
||||
|
||||
// Close closes the plugin loader and unloads all plugins
|
||||
func (pl *PluginLoader) Close() error {
|
||||
pl.mu.Lock()
|
||||
defer pl.mu.Unlock()
|
||||
|
||||
var lastErr error
|
||||
for name := range pl.loadedPlugins {
|
||||
if err := pl.registry.Unregister(name); err != nil {
|
||||
lastErr = err
|
||||
}
|
||||
}
|
||||
|
||||
pl.loadedPlugins = make(map[string]*LoadedPlugin)
|
||||
pl.loadTime = make(map[string]time.Time)
|
||||
|
||||
return lastErr
|
||||
}
|
||||
|
||||
// GetPluginRegistry returns the underlying registry
|
||||
func (pl *PluginLoader) GetPluginRegistry() *PluginRegistry {
|
||||
return pl.registry
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
package plugins
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestNewPluginLoader(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
assert.NotNil(t, loader)
|
||||
assert.Equal(t, registry, loader.registry)
|
||||
}
|
||||
|
||||
func TestRegisterLoadedPlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
err := loader.RegisterLoadedPlugin(plugin, "test-author", nil)
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, loader.IsPluginLoaded("test-plugin"))
|
||||
}
|
||||
|
||||
func TestUnloadPlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
loader.RegisterLoadedPlugin(plugin, "test-author", nil)
|
||||
|
||||
assert.True(t, loader.IsPluginLoaded("test-plugin"))
|
||||
|
||||
err := loader.UnloadPlugin("test-plugin")
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, loader.IsPluginLoaded("test-plugin"))
|
||||
}
|
||||
|
||||
func TestUnloadPluginNotFound(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
err := loader.UnloadPlugin("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestReloadPlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
loader.RegisterLoadedPlugin(plugin, "test-author", nil)
|
||||
|
||||
_, _ = loader.GetLoadTime("test-plugin")
|
||||
|
||||
err := loader.ReloadPlugin("test-plugin")
|
||||
assert.NoError(t, err)
|
||||
|
||||
loadedPlugins := loader.GetLoadedPlugins()
|
||||
assert.Equal(t, 1, loadedPlugins["test-plugin"].ReloadCount)
|
||||
}
|
||||
|
||||
func TestReloadPluginNotFound(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
err := loader.ReloadPlugin("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestGetLoadedPlugins(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
plugin := &MockPlugin{
|
||||
name: string(rune(48 + i)) + "-plugin",
|
||||
version: "1.0.0",
|
||||
}
|
||||
loader.RegisterLoadedPlugin(plugin, "author", nil)
|
||||
}
|
||||
|
||||
loaded := loader.GetLoadedPlugins()
|
||||
assert.Equal(t, 3, len(loaded))
|
||||
}
|
||||
|
||||
func TestGetLoadTime(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
loader.RegisterLoadedPlugin(plugin, "author", nil)
|
||||
|
||||
loadTime, exists := loader.GetLoadTime("test-plugin")
|
||||
assert.True(t, exists)
|
||||
assert.NotZero(t, loadTime)
|
||||
}
|
||||
|
||||
func TestIsPluginLoaded(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
assert.False(t, loader.IsPluginLoaded("test-plugin"))
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
loader.RegisterLoadedPlugin(plugin, "author", nil)
|
||||
|
||||
assert.True(t, loader.IsPluginLoaded("test-plugin"))
|
||||
}
|
||||
|
||||
func TestLoaderExecutePlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
loader.RegisterLoadedPlugin(plugin, "author", nil)
|
||||
|
||||
output, err := loader.ExecutePlugin("test-plugin", map[string]interface{}{})
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, output)
|
||||
}
|
||||
|
||||
func TestLoaderExecutePluginNotLoaded(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
_, err := loader.ExecutePlugin("nonexistent", map[string]interface{}{})
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestGetPluginStats(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
loader.RegisterLoadedPlugin(plugin, "author", nil)
|
||||
|
||||
stats, err := loader.GetPluginStats("test-plugin")
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, stats)
|
||||
}
|
||||
|
||||
func TestGetPluginStatsNotLoaded(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
_, err := loader.GetPluginStats("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestClose(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
plugin := &MockPlugin{
|
||||
name: string(rune(48+i)) + "-plugin",
|
||||
version: "1.0.0",
|
||||
}
|
||||
loader.RegisterLoadedPlugin(plugin, "author", nil)
|
||||
}
|
||||
|
||||
assert.Equal(t, 3, len(loader.GetLoadedPlugins()))
|
||||
|
||||
err := loader.Close()
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 0, len(loader.GetLoadedPlugins()))
|
||||
}
|
||||
|
||||
func TestGetPluginRegistry(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
retrieved := loader.GetPluginRegistry()
|
||||
assert.Equal(t, registry, retrieved)
|
||||
}
|
||||
|
||||
func TestLoadPlugin(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, tmpDir)
|
||||
|
||||
// For now, test with plugin:// URL (no file loading needed)
|
||||
err := loader.LoadPlugin("plugin://test-plugin")
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestLoadPluginDirectory(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, tmpDir)
|
||||
|
||||
// Create some mock config files
|
||||
configContent := `{
|
||||
"name": "test-plugin",
|
||||
"version": "1.0.0",
|
||||
"author": "test-author",
|
||||
"path": "test-plugin"
|
||||
}`
|
||||
|
||||
configFile := filepath.Join(tmpDir, "test-plugin.json")
|
||||
os.WriteFile(configFile, []byte(configContent), 0644)
|
||||
|
||||
// Load from directory (won't actually load plugins without more setup)
|
||||
err := loader.LoadPluginDirectory(tmpDir)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestLoadPluginDirectoryNotFound(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
err := loader.LoadPluginDirectory("/nonexistent/directory")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestGetFailedLoads(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, tmpDir)
|
||||
|
||||
// Try to load from nonexistent file
|
||||
loader.LoadPlugin(filepath.Join(tmpDir, "nonexistent.json"))
|
||||
|
||||
failed := loader.GetFailedLoads()
|
||||
assert.Greater(t, len(failed), 0)
|
||||
}
|
||||
|
||||
func TestMultiplePluginLifecycle(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
// Load plugins
|
||||
for i := 0; i < 5; i++ {
|
||||
plugin := &MockPlugin{
|
||||
name: string(rune(48+i)) + "-plugin",
|
||||
version: "1.0.0",
|
||||
}
|
||||
err := loader.RegisterLoadedPlugin(plugin, "author", nil)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
assert.Equal(t, 5, len(loader.GetLoadedPlugins()))
|
||||
|
||||
// Execute plugins
|
||||
for i := 0; i < 5; i++ {
|
||||
name := string(rune(48+i)) + "-plugin"
|
||||
output, err := loader.ExecutePlugin(name, map[string]interface{}{})
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, output)
|
||||
}
|
||||
|
||||
// Unload plugins
|
||||
for i := 0; i < 5; i++ {
|
||||
name := string(rune(48+i)) + "-plugin"
|
||||
err := loader.UnloadPlugin(name)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
assert.Equal(t, 0, len(loader.GetLoadedPlugins()))
|
||||
}
|
||||
|
||||
func BenchmarkRegisterLoadedPlugin(b *testing.B) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
for i := 0; i < b.N; i++ {
|
||||
plugin := &MockPlugin{
|
||||
name: string(rune(48+i%100)) + "-plugin",
|
||||
version: "1.0.0",
|
||||
}
|
||||
loader.RegisterLoadedPlugin(plugin, "author", nil)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkExecuteLoadedPlugin(b *testing.B) {
|
||||
registry := NewPluginRegistry()
|
||||
loader := NewPluginLoader(registry, "/tmp/plugins")
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
loader.RegisterLoadedPlugin(plugin, "author", nil)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
loader.ExecutePlugin("test-plugin", map[string]interface{}{})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,322 @@
|
||||
package plugins
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// SkillPlugin represents a custom skill plugin
|
||||
type SkillPlugin interface {
|
||||
// Name returns the plugin name
|
||||
Name() string
|
||||
// Version returns the plugin version
|
||||
Version() string
|
||||
// Execute executes the plugin with the given input
|
||||
Execute(input map[string]interface{}) (map[string]interface{}, error)
|
||||
// Validate validates the plugin configuration
|
||||
Validate() error
|
||||
// Description returns a human-readable description
|
||||
Description() string
|
||||
}
|
||||
|
||||
// PluginMetadata holds metadata about a plugin
|
||||
type PluginMetadata struct {
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
Author string `json:"author"`
|
||||
Description string `json:"description"`
|
||||
URL string `json:"url"`
|
||||
Config map[string]interface{} `json:"config,omitempty"`
|
||||
LoadedAt time.Time `json:"loaded_at"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
// PluginRegistry manages custom skill plugins
|
||||
type PluginRegistry struct {
|
||||
mu sync.RWMutex
|
||||
plugins map[string]SkillPlugin
|
||||
metadata map[string]*PluginMetadata
|
||||
executionLog map[string][]*ExecutionRecord
|
||||
stats *PluginStats
|
||||
}
|
||||
|
||||
// ExecutionRecord tracks plugin execution
|
||||
type ExecutionRecord struct {
|
||||
PluginName string
|
||||
Timestamp time.Time
|
||||
Duration time.Duration
|
||||
Input map[string]interface{}
|
||||
Output map[string]interface{}
|
||||
Error error
|
||||
Success bool
|
||||
}
|
||||
|
||||
// PluginStats tracks plugin statistics
|
||||
type PluginStats struct {
|
||||
TotalExecutions int
|
||||
SuccessfulExecutions int
|
||||
FailedExecutions int
|
||||
TotalPlugins int
|
||||
EnabledPlugins int
|
||||
AverageExecutionTime time.Duration
|
||||
}
|
||||
|
||||
// NewPluginRegistry creates a new plugin registry
|
||||
func NewPluginRegistry() *PluginRegistry {
|
||||
return &PluginRegistry{
|
||||
plugins: make(map[string]SkillPlugin),
|
||||
metadata: make(map[string]*PluginMetadata),
|
||||
executionLog: make(map[string][]*ExecutionRecord),
|
||||
stats: &PluginStats{
|
||||
TotalExecutions: 0,
|
||||
SuccessfulExecutions: 0,
|
||||
FailedExecutions: 0,
|
||||
TotalPlugins: 0,
|
||||
EnabledPlugins: 0,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Register registers a new plugin
|
||||
func (pr *PluginRegistry) Register(plugin SkillPlugin, author string, config map[string]interface{}) error {
|
||||
if plugin == nil {
|
||||
return fmt.Errorf("plugin cannot be nil")
|
||||
}
|
||||
|
||||
// Validate plugin
|
||||
if err := plugin.Validate(); err != nil {
|
||||
return fmt.Errorf("plugin validation failed: %w", err)
|
||||
}
|
||||
|
||||
pr.mu.Lock()
|
||||
defer pr.mu.Unlock()
|
||||
|
||||
name := plugin.Name()
|
||||
if _, exists := pr.plugins[name]; exists {
|
||||
return fmt.Errorf("plugin already registered: %s", name)
|
||||
}
|
||||
|
||||
pr.plugins[name] = plugin
|
||||
pr.metadata[name] = &PluginMetadata{
|
||||
Name: name,
|
||||
Version: plugin.Version(),
|
||||
Author: author,
|
||||
Description: plugin.Description(),
|
||||
URL: fmt.Sprintf("plugin://%s", name),
|
||||
Config: config,
|
||||
LoadedAt: time.Now(),
|
||||
Enabled: true,
|
||||
}
|
||||
|
||||
pr.stats.TotalPlugins++
|
||||
pr.stats.EnabledPlugins++
|
||||
pr.executionLog[name] = make([]*ExecutionRecord, 0)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Unregister unregisters a plugin
|
||||
func (pr *PluginRegistry) Unregister(name string) error {
|
||||
pr.mu.Lock()
|
||||
defer pr.mu.Unlock()
|
||||
|
||||
if _, exists := pr.plugins[name]; !exists {
|
||||
return fmt.Errorf("plugin not found: %s", name)
|
||||
}
|
||||
|
||||
delete(pr.plugins, name)
|
||||
if pr.metadata[name].Enabled {
|
||||
pr.stats.EnabledPlugins--
|
||||
}
|
||||
pr.stats.TotalPlugins--
|
||||
delete(pr.metadata, name)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Execute executes a plugin by name
|
||||
func (pr *PluginRegistry) Execute(name string, input map[string]interface{}) (map[string]interface{}, error) {
|
||||
pr.mu.RLock()
|
||||
plugin, exists := pr.plugins[name]
|
||||
metadata, metaExists := pr.metadata[name]
|
||||
pr.mu.RUnlock()
|
||||
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("plugin not found: %s", name)
|
||||
}
|
||||
|
||||
if !metaExists || !metadata.Enabled {
|
||||
return nil, fmt.Errorf("plugin is disabled: %s", name)
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
output, err := plugin.Execute(input)
|
||||
duration := time.Since(start)
|
||||
|
||||
// Record execution
|
||||
record := &ExecutionRecord{
|
||||
PluginName: name,
|
||||
Timestamp: start,
|
||||
Duration: duration,
|
||||
Input: input,
|
||||
Output: output,
|
||||
Error: err,
|
||||
Success: err == nil,
|
||||
}
|
||||
|
||||
pr.mu.Lock()
|
||||
pr.executionLog[name] = append(pr.executionLog[name], record)
|
||||
pr.stats.TotalExecutions++
|
||||
if err == nil {
|
||||
pr.stats.SuccessfulExecutions++
|
||||
} else {
|
||||
pr.stats.FailedExecutions++
|
||||
}
|
||||
pr.mu.Unlock()
|
||||
|
||||
return output, err
|
||||
}
|
||||
|
||||
// Get retrieves a plugin by name
|
||||
func (pr *PluginRegistry) Get(name string) (SkillPlugin, bool) {
|
||||
pr.mu.RLock()
|
||||
defer pr.mu.RUnlock()
|
||||
|
||||
plugin, exists := pr.plugins[name]
|
||||
return plugin, exists
|
||||
}
|
||||
|
||||
// GetMetadata retrieves plugin metadata
|
||||
func (pr *PluginRegistry) GetMetadata(name string) (*PluginMetadata, bool) {
|
||||
pr.mu.RLock()
|
||||
defer pr.mu.RUnlock()
|
||||
|
||||
meta, exists := pr.metadata[name]
|
||||
return meta, exists
|
||||
}
|
||||
|
||||
// ListPlugins returns all registered plugins
|
||||
func (pr *PluginRegistry) ListPlugins() map[string]*PluginMetadata {
|
||||
pr.mu.RLock()
|
||||
defer pr.mu.RUnlock()
|
||||
|
||||
result := make(map[string]*PluginMetadata)
|
||||
for name, meta := range pr.metadata {
|
||||
result[name] = meta
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// EnablePlugin enables a plugin
|
||||
func (pr *PluginRegistry) EnablePlugin(name string) error {
|
||||
pr.mu.Lock()
|
||||
defer pr.mu.Unlock()
|
||||
|
||||
meta, exists := pr.metadata[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("plugin not found: %s", name)
|
||||
}
|
||||
|
||||
if meta.Enabled {
|
||||
return fmt.Errorf("plugin already enabled: %s", name)
|
||||
}
|
||||
|
||||
meta.Enabled = true
|
||||
pr.stats.EnabledPlugins++
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DisablePlugin disables a plugin
|
||||
func (pr *PluginRegistry) DisablePlugin(name string) error {
|
||||
pr.mu.Lock()
|
||||
defer pr.mu.Unlock()
|
||||
|
||||
meta, exists := pr.metadata[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("plugin not found: %s", name)
|
||||
}
|
||||
|
||||
if !meta.Enabled {
|
||||
return fmt.Errorf("plugin already disabled: %s", name)
|
||||
}
|
||||
|
||||
meta.Enabled = false
|
||||
pr.stats.EnabledPlugins--
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetExecutionLog returns execution history for a plugin
|
||||
func (pr *PluginRegistry) GetExecutionLog(name string) []*ExecutionRecord {
|
||||
pr.mu.RLock()
|
||||
defer pr.mu.RUnlock()
|
||||
|
||||
if log, exists := pr.executionLog[name]; exists {
|
||||
result := make([]*ExecutionRecord, len(log))
|
||||
copy(result, log)
|
||||
return result
|
||||
}
|
||||
|
||||
return make([]*ExecutionRecord, 0)
|
||||
}
|
||||
|
||||
// GetStats returns registry statistics
|
||||
func (pr *PluginRegistry) GetStats() *PluginStats {
|
||||
pr.mu.RLock()
|
||||
defer pr.mu.RUnlock()
|
||||
|
||||
stats := *pr.stats
|
||||
if stats.TotalExecutions > 0 {
|
||||
totalDuration := time.Duration(0)
|
||||
for _, log := range pr.executionLog {
|
||||
for _, record := range log {
|
||||
totalDuration += record.Duration
|
||||
}
|
||||
}
|
||||
stats.AverageExecutionTime = totalDuration / time.Duration(stats.TotalExecutions)
|
||||
}
|
||||
|
||||
return &stats
|
||||
}
|
||||
|
||||
// ResolvePluginURL resolves a plugin:// URL
|
||||
func (pr *PluginRegistry) ResolvePluginURL(url string) (SkillPlugin, error) {
|
||||
if len(url) < 9 || url[:9] != "plugin://" {
|
||||
return nil, fmt.Errorf("invalid plugin URL: %s", url)
|
||||
}
|
||||
|
||||
name := url[9:] // Remove "plugin://" prefix
|
||||
plugin, exists := pr.Get(name)
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("plugin not found: %s", name)
|
||||
}
|
||||
|
||||
return plugin, nil
|
||||
}
|
||||
|
||||
// Clear clears all plugins
|
||||
func (pr *PluginRegistry) Clear() {
|
||||
pr.mu.Lock()
|
||||
defer pr.mu.Unlock()
|
||||
|
||||
pr.plugins = make(map[string]SkillPlugin)
|
||||
pr.metadata = make(map[string]*PluginMetadata)
|
||||
pr.executionLog = make(map[string][]*ExecutionRecord)
|
||||
pr.stats = &PluginStats{}
|
||||
}
|
||||
|
||||
// IsPluginURL checks if a URL is a plugin URL
|
||||
func IsPluginURL(url string) bool {
|
||||
return len(url) > 9 && url[:9] == "plugin://"
|
||||
}
|
||||
|
||||
// ExtractPluginName extracts plugin name from plugin URL
|
||||
func ExtractPluginName(url string) string {
|
||||
if IsPluginURL(url) {
|
||||
return url[9:]
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,408 @@
|
||||
package plugins
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
// MockPlugin is a test plugin implementation
|
||||
type MockPlugin struct {
|
||||
name string
|
||||
version string
|
||||
description string
|
||||
shouldFail bool
|
||||
shouldWait time.Duration
|
||||
}
|
||||
|
||||
func (mp *MockPlugin) Name() string {
|
||||
return mp.name
|
||||
}
|
||||
|
||||
func (mp *MockPlugin) Version() string {
|
||||
return mp.version
|
||||
}
|
||||
|
||||
func (mp *MockPlugin) Description() string {
|
||||
return mp.description
|
||||
}
|
||||
|
||||
func (mp *MockPlugin) Execute(input map[string]interface{}) (map[string]interface{}, error) {
|
||||
if mp.shouldWait > 0 {
|
||||
time.Sleep(mp.shouldWait)
|
||||
}
|
||||
|
||||
if mp.shouldFail {
|
||||
return nil, fmt.Errorf("plugin execution failed")
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"result": "success",
|
||||
"input": input,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (mp *MockPlugin) Validate() error {
|
||||
if mp.name == "" {
|
||||
return fmt.Errorf("plugin name cannot be empty")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestNewPluginRegistry(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
assert.NotNil(t, registry)
|
||||
assert.Equal(t, 0, registry.stats.TotalPlugins)
|
||||
}
|
||||
|
||||
func TestRegisterPlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{
|
||||
name: "test-plugin",
|
||||
version: "1.0.0",
|
||||
description: "Test plugin",
|
||||
}
|
||||
|
||||
err := registry.Register(plugin, "test-author", nil)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 1, registry.stats.TotalPlugins)
|
||||
}
|
||||
|
||||
func TestRegisterPluginNil(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
err := registry.Register(nil, "test-author", nil)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestRegisterDuplicatePlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{
|
||||
name: "test-plugin",
|
||||
version: "1.0.0",
|
||||
}
|
||||
|
||||
registry.Register(plugin, "author", nil)
|
||||
err := registry.Register(plugin, "author", nil)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestUnregisterPlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
assert.Equal(t, 1, registry.stats.TotalPlugins)
|
||||
|
||||
err := registry.Unregister("test-plugin")
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 0, registry.stats.TotalPlugins)
|
||||
}
|
||||
|
||||
func TestExecutePlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
input := map[string]interface{}{"key": "value"}
|
||||
output, err := registry.Execute("test-plugin", input)
|
||||
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, output)
|
||||
assert.Equal(t, "success", output["result"])
|
||||
}
|
||||
|
||||
func TestExecutePluginNotFound(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
_, err := registry.Execute("nonexistent", map[string]interface{}{})
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestExecutePluginDisabled(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
registry.DisablePlugin("test-plugin")
|
||||
|
||||
_, err := registry.Execute("test-plugin", map[string]interface{}{})
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestExecutePluginFailure(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{
|
||||
name: "test-plugin",
|
||||
version: "1.0.0",
|
||||
shouldFail: true,
|
||||
}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
_, err := registry.Execute("test-plugin", map[string]interface{}{})
|
||||
assert.Error(t, err)
|
||||
assert.Equal(t, 1, registry.stats.FailedExecutions)
|
||||
}
|
||||
|
||||
func TestGetPlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
retrieved, found := registry.Get("test-plugin")
|
||||
assert.True(t, found)
|
||||
assert.Equal(t, "test-plugin", retrieved.Name())
|
||||
}
|
||||
|
||||
func TestGetPluginNotFound(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
_, found := registry.Get("nonexistent")
|
||||
assert.False(t, found)
|
||||
}
|
||||
|
||||
func TestGetMetadata(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{
|
||||
name: "test-plugin",
|
||||
version: "1.0.0",
|
||||
description: "Test description",
|
||||
}
|
||||
registry.Register(plugin, "test-author", nil)
|
||||
|
||||
meta, found := registry.GetMetadata("test-plugin")
|
||||
assert.True(t, found)
|
||||
assert.Equal(t, "test-plugin", meta.Name)
|
||||
assert.Equal(t, "1.0.0", meta.Version)
|
||||
assert.Equal(t, "test-author", meta.Author)
|
||||
assert.Equal(t, "plugin://test-plugin", meta.URL)
|
||||
assert.True(t, meta.Enabled)
|
||||
}
|
||||
|
||||
func TestListPlugins(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
plugin := &MockPlugin{
|
||||
name: fmt.Sprintf("plugin-%d", i),
|
||||
version: "1.0.0",
|
||||
}
|
||||
registry.Register(plugin, "author", nil)
|
||||
}
|
||||
|
||||
plugins := registry.ListPlugins()
|
||||
assert.Equal(t, 3, len(plugins))
|
||||
}
|
||||
|
||||
func TestEnableDisablePlugin(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
assert.Equal(t, 1, registry.stats.EnabledPlugins)
|
||||
|
||||
registry.DisablePlugin("test-plugin")
|
||||
assert.Equal(t, 0, registry.stats.EnabledPlugins)
|
||||
|
||||
registry.EnablePlugin("test-plugin")
|
||||
assert.Equal(t, 1, registry.stats.EnabledPlugins)
|
||||
}
|
||||
|
||||
func TestGetExecutionLog(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
registry.Execute("test-plugin", map[string]interface{}{})
|
||||
registry.Execute("test-plugin", map[string]interface{}{})
|
||||
|
||||
log := registry.GetExecutionLog("test-plugin")
|
||||
assert.Equal(t, 2, len(log))
|
||||
}
|
||||
|
||||
func TestExecutionLogSuccess(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
registry.Execute("test-plugin", map[string]interface{}{})
|
||||
|
||||
log := registry.GetExecutionLog("test-plugin")
|
||||
assert.Equal(t, 1, len(log))
|
||||
assert.True(t, log[0].Success)
|
||||
assert.Nil(t, log[0].Error)
|
||||
}
|
||||
|
||||
func TestExecutionLogFailure(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{
|
||||
name: "test-plugin",
|
||||
version: "1.0.0",
|
||||
shouldFail: true,
|
||||
}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
registry.Execute("test-plugin", map[string]interface{}{})
|
||||
|
||||
log := registry.GetExecutionLog("test-plugin")
|
||||
assert.Equal(t, 1, len(log))
|
||||
assert.False(t, log[0].Success)
|
||||
assert.NotNil(t, log[0].Error)
|
||||
}
|
||||
|
||||
func TestGetStats(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
registry.Execute("test-plugin", map[string]interface{}{})
|
||||
registry.Execute("test-plugin", map[string]interface{}{})
|
||||
|
||||
stats := registry.GetStats()
|
||||
assert.Equal(t, 1, stats.TotalPlugins)
|
||||
assert.Equal(t, 2, stats.TotalExecutions)
|
||||
assert.Equal(t, 2, stats.SuccessfulExecutions)
|
||||
}
|
||||
|
||||
func TestResolvePluginURL(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
resolved, err := registry.ResolvePluginURL("plugin://test-plugin")
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "test-plugin", resolved.Name())
|
||||
}
|
||||
|
||||
func TestResolvePluginURLInvalid(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
_, err := registry.ResolvePluginURL("http://example.com")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestResolvePluginURLNotFound(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
_, err := registry.ResolvePluginURL("plugin://nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestClear(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
assert.Equal(t, 1, registry.stats.TotalPlugins)
|
||||
|
||||
registry.Clear()
|
||||
assert.Equal(t, 0, registry.stats.TotalPlugins)
|
||||
}
|
||||
|
||||
func TestIsPluginURL(t *testing.T) {
|
||||
assert.True(t, IsPluginURL("plugin://test"))
|
||||
assert.False(t, IsPluginURL("http://test"))
|
||||
assert.False(t, IsPluginURL("file://test"))
|
||||
}
|
||||
|
||||
func TestExtractPluginName(t *testing.T) {
|
||||
name := ExtractPluginName("plugin://test-plugin")
|
||||
assert.Equal(t, "test-plugin", name)
|
||||
|
||||
name = ExtractPluginName("http://test")
|
||||
assert.Equal(t, "", name)
|
||||
}
|
||||
|
||||
func TestExecutionTiming(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{
|
||||
name: "test-plugin",
|
||||
version: "1.0.0",
|
||||
shouldWait: 50 * time.Millisecond,
|
||||
}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
registry.Execute("test-plugin", map[string]interface{}{})
|
||||
|
||||
log := registry.GetExecutionLog("test-plugin")
|
||||
assert.Greater(t, log[0].Duration, 40*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestMultiplePlugins(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
plugin := &MockPlugin{
|
||||
name: fmt.Sprintf("plugin-%d", i),
|
||||
version: "1.0.0",
|
||||
}
|
||||
registry.Register(plugin, "author", nil)
|
||||
}
|
||||
|
||||
assert.Equal(t, 5, registry.stats.TotalPlugins)
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
registry.Execute(fmt.Sprintf("plugin-%d", i), map[string]interface{}{})
|
||||
}
|
||||
|
||||
stats := registry.GetStats()
|
||||
assert.Equal(t, 5, stats.TotalExecutions)
|
||||
}
|
||||
|
||||
func TestPluginConfig(t *testing.T) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
config := map[string]interface{}{
|
||||
"setting1": "value1",
|
||||
"setting2": 42,
|
||||
}
|
||||
|
||||
registry.Register(plugin, "author", config)
|
||||
|
||||
meta, _ := registry.GetMetadata("test-plugin")
|
||||
assert.NotNil(t, meta.Config)
|
||||
assert.Equal(t, "value1", meta.Config["setting1"])
|
||||
assert.Equal(t, 42, meta.Config["setting2"])
|
||||
}
|
||||
|
||||
func BenchmarkRegisterPlugin(b *testing.B) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
for i := 0; i < b.N; i++ {
|
||||
plugin := &MockPlugin{
|
||||
name: fmt.Sprintf("plugin-%d", i),
|
||||
version: "1.0.0",
|
||||
}
|
||||
registry.Register(plugin, "author", nil)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkExecutePlugin(b *testing.B) {
|
||||
registry := NewPluginRegistry()
|
||||
|
||||
plugin := &MockPlugin{name: "test-plugin", version: "1.0.0"}
|
||||
registry.Register(plugin, "author", nil)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
registry.Execute("test-plugin", map[string]interface{}{})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,208 @@
|
||||
package profiling
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// TaskProfile represents profiling data for a task
|
||||
type TaskProfile struct {
|
||||
TaskID string
|
||||
Duration float64
|
||||
CPUUsage float64
|
||||
MemUsage float64
|
||||
Throughput float64
|
||||
}
|
||||
|
||||
// WorkflowProfile represents profiling for an entire workflow
|
||||
type WorkflowProfile struct {
|
||||
WorkflowID string
|
||||
Tasks map[string]*TaskProfile
|
||||
TotalDuration float64
|
||||
CriticalPath []string
|
||||
}
|
||||
|
||||
// WorkflowProfiler profiles workflow execution
|
||||
type WorkflowProfiler struct {
|
||||
mu sync.RWMutex
|
||||
profiles map[string]*WorkflowProfile
|
||||
}
|
||||
|
||||
// NewWorkflowProfiler creates a new workflow profiler
|
||||
func NewWorkflowProfiler() *WorkflowProfiler {
|
||||
return &WorkflowProfiler{
|
||||
profiles: make(map[string]*WorkflowProfile),
|
||||
}
|
||||
}
|
||||
|
||||
// RecordTaskExecution records task execution metrics
|
||||
func (wp *WorkflowProfiler) RecordTaskExecution(workflowID, taskID string, duration, cpu, mem float64) {
|
||||
wp.mu.Lock()
|
||||
defer wp.mu.Unlock()
|
||||
|
||||
if _, exists := wp.profiles[workflowID]; !exists {
|
||||
wp.profiles[workflowID] = &WorkflowProfile{
|
||||
WorkflowID: workflowID,
|
||||
Tasks: make(map[string]*TaskProfile),
|
||||
CriticalPath: make([]string, 0),
|
||||
}
|
||||
}
|
||||
|
||||
profile := wp.profiles[workflowID]
|
||||
profile.Tasks[taskID] = &TaskProfile{
|
||||
TaskID: taskID,
|
||||
Duration: duration,
|
||||
CPUUsage: cpu,
|
||||
MemUsage: mem,
|
||||
Throughput: 1000.0 / duration, // Tasks per second
|
||||
}
|
||||
|
||||
// Recalculate total duration
|
||||
total := 0.0
|
||||
for _, tp := range profile.Tasks {
|
||||
if tp.Duration > total {
|
||||
total = tp.Duration
|
||||
}
|
||||
}
|
||||
profile.TotalDuration = total
|
||||
}
|
||||
|
||||
// GetSlowTasks returns tasks sorted by duration (slowest first)
|
||||
func (wp *WorkflowProfiler) GetSlowTasks(workflowID string, limit int) []string {
|
||||
wp.mu.RLock()
|
||||
defer wp.mu.RUnlock()
|
||||
|
||||
profile, exists := wp.profiles[workflowID]
|
||||
if !exists {
|
||||
return []string{}
|
||||
}
|
||||
|
||||
// Sort tasks by duration
|
||||
type taskDuration struct {
|
||||
taskID string
|
||||
duration float64
|
||||
}
|
||||
|
||||
tasks := make([]taskDuration, 0)
|
||||
for taskID, tp := range profile.Tasks {
|
||||
tasks = append(tasks, taskDuration{taskID, tp.Duration})
|
||||
}
|
||||
|
||||
sort.Slice(tasks, func(i, j int) bool {
|
||||
return tasks[i].duration > tasks[j].duration
|
||||
})
|
||||
|
||||
result := make([]string, 0)
|
||||
for i, t := range tasks {
|
||||
if i >= limit {
|
||||
break
|
||||
}
|
||||
result = append(result, t.taskID)
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetHighCPUTasks returns tasks with high CPU usage
|
||||
func (wp *WorkflowProfiler) GetHighCPUTasks(workflowID string, threshold float64) []string {
|
||||
wp.mu.RLock()
|
||||
defer wp.mu.RUnlock()
|
||||
|
||||
profile, exists := wp.profiles[workflowID]
|
||||
if !exists {
|
||||
return []string{}
|
||||
}
|
||||
|
||||
result := make([]string, 0)
|
||||
for taskID, tp := range profile.Tasks {
|
||||
if tp.CPUUsage > threshold {
|
||||
result = append(result, taskID)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetHighMemTasks returns tasks with high memory usage
|
||||
func (wp *WorkflowProfiler) GetHighMemTasks(workflowID string, threshold float64) []string {
|
||||
wp.mu.RLock()
|
||||
defer wp.mu.RUnlock()
|
||||
|
||||
profile, exists := wp.profiles[workflowID]
|
||||
if !exists {
|
||||
return []string{}
|
||||
}
|
||||
|
||||
result := make([]string, 0)
|
||||
for taskID, tp := range profile.Tasks {
|
||||
if tp.MemUsage > threshold {
|
||||
result = append(result, taskID)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetOptimizationSuggestions returns optimization recommendations
|
||||
func (wp *WorkflowProfiler) GetOptimizationSuggestions(workflowID string) []string {
|
||||
wp.mu.RLock()
|
||||
defer wp.mu.RUnlock()
|
||||
|
||||
profile, exists := wp.profiles[workflowID]
|
||||
if !exists {
|
||||
return []string{}
|
||||
}
|
||||
|
||||
suggestions := make([]string, 0)
|
||||
|
||||
// Check for slow tasks
|
||||
for taskID, tp := range profile.Tasks {
|
||||
if tp.Duration > profile.TotalDuration*0.5 {
|
||||
suggestions = append(suggestions, fmt.Sprintf("Task %s takes 50%% of total time, consider optimizing", taskID))
|
||||
}
|
||||
if tp.CPUUsage > 0.8 {
|
||||
suggestions = append(suggestions, fmt.Sprintf("Task %s has high CPU usage (%.2f), consider parallelizing", taskID, tp.CPUUsage))
|
||||
}
|
||||
if tp.MemUsage > 0.8 {
|
||||
suggestions = append(suggestions, fmt.Sprintf("Task %s has high memory usage (%.2f), consider reducing payload", taskID, tp.MemUsage))
|
||||
}
|
||||
}
|
||||
|
||||
return suggestions
|
||||
}
|
||||
|
||||
// GetProfile retrieves profiling data for a workflow
|
||||
func (wp *WorkflowProfiler) GetProfile(workflowID string) (*WorkflowProfile, bool) {
|
||||
wp.mu.RLock()
|
||||
defer wp.mu.RUnlock()
|
||||
|
||||
profile, exists := wp.profiles[workflowID]
|
||||
return profile, exists
|
||||
}
|
||||
|
||||
// GetTaskProfile retrieves profiling data for a specific task
|
||||
func (wp *WorkflowProfiler) GetTaskProfile(workflowID, taskID string) (*TaskProfile, error) {
|
||||
wp.mu.RLock()
|
||||
defer wp.mu.RUnlock()
|
||||
|
||||
profile, exists := wp.profiles[workflowID]
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("workflow not found: %s", workflowID)
|
||||
}
|
||||
|
||||
taskProfile, exists := profile.Tasks[taskID]
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("task not found: %s", taskID)
|
||||
}
|
||||
|
||||
return taskProfile, nil
|
||||
}
|
||||
|
||||
// Clear clears all profiles
|
||||
func (wp *WorkflowProfiler) Clear() {
|
||||
wp.mu.Lock()
|
||||
defer wp.mu.Unlock()
|
||||
|
||||
wp.profiles = make(map[string]*WorkflowProfile)
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package profiling
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestRecordTaskExecution(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 100, 0.5, 0.3)
|
||||
|
||||
profile, exists := profiler.GetProfile("wf-1")
|
||||
assert.True(t, exists)
|
||||
assert.Equal(t, 1, len(profile.Tasks))
|
||||
}
|
||||
|
||||
func TestGetSlowTasks(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 100, 0.5, 0.3)
|
||||
profiler.RecordTaskExecution("wf-1", "task-2", 500, 0.8, 0.6)
|
||||
profiler.RecordTaskExecution("wf-1", "task-3", 200, 0.4, 0.2)
|
||||
|
||||
slow := profiler.GetSlowTasks("wf-1", 2)
|
||||
assert.Equal(t, 2, len(slow))
|
||||
assert.Equal(t, "task-2", slow[0])
|
||||
}
|
||||
|
||||
func TestGetHighCPUTasks(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 100, 0.5, 0.3)
|
||||
profiler.RecordTaskExecution("wf-1", "task-2", 200, 0.9, 0.6)
|
||||
|
||||
highCPU := profiler.GetHighCPUTasks("wf-1", 0.7)
|
||||
assert.Equal(t, 1, len(highCPU))
|
||||
assert.Equal(t, "task-2", highCPU[0])
|
||||
}
|
||||
|
||||
func TestGetHighMemTasks(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 100, 0.5, 0.3)
|
||||
profiler.RecordTaskExecution("wf-1", "task-2", 200, 0.6, 0.9)
|
||||
|
||||
highMem := profiler.GetHighMemTasks("wf-1", 0.7)
|
||||
assert.Equal(t, 1, len(highMem))
|
||||
assert.Equal(t, "task-2", highMem[0])
|
||||
}
|
||||
|
||||
func TestGetOptimizationSuggestions(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 100, 0.5, 0.3)
|
||||
profiler.RecordTaskExecution("wf-1", "task-2", 200, 0.9, 0.9)
|
||||
|
||||
suggestions := profiler.GetOptimizationSuggestions("wf-1")
|
||||
assert.Greater(t, len(suggestions), 0)
|
||||
}
|
||||
|
||||
func TestGetProfile(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 100, 0.5, 0.3)
|
||||
|
||||
profile, exists := profiler.GetProfile("wf-1")
|
||||
assert.True(t, exists)
|
||||
assert.Equal(t, "wf-1", profile.WorkflowID)
|
||||
}
|
||||
|
||||
func TestGetTaskProfile(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 100, 0.5, 0.3)
|
||||
|
||||
taskProfile, err := profiler.GetTaskProfile("wf-1", "task-1")
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 100.0, taskProfile.Duration)
|
||||
}
|
||||
|
||||
func TestTaskNotFound(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
_, err := profiler.GetTaskProfile("wf-1", "task-999")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestClear(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 100, 0.5, 0.3)
|
||||
profiler.Clear()
|
||||
|
||||
profile, exists := profiler.GetProfile("wf-1")
|
||||
assert.False(t, exists)
|
||||
assert.Nil(t, profile)
|
||||
}
|
||||
|
||||
func TestThroughputCalculation(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 1000, 0.5, 0.3)
|
||||
|
||||
profile, _ := profiler.GetProfile("wf-1")
|
||||
taskProfile := profile.Tasks["task-1"]
|
||||
|
||||
assert.Equal(t, 1.0, taskProfile.Throughput) // 1000ms = 1 task per second
|
||||
}
|
||||
|
||||
func TestMultipleWorkflows(t *testing.T) {
|
||||
profiler := NewWorkflowProfiler()
|
||||
|
||||
profiler.RecordTaskExecution("wf-1", "task-1", 100, 0.5, 0.3)
|
||||
profiler.RecordTaskExecution("wf-2", "task-1", 200, 0.6, 0.4)
|
||||
|
||||
profile1, exists1 := profiler.GetProfile("wf-1")
|
||||
profile2, exists2 := profiler.GetProfile("wf-2")
|
||||
|
||||
assert.True(t, exists1)
|
||||
assert.True(t, exists2)
|
||||
assert.Equal(t, 100.0, profile1.TotalDuration)
|
||||
assert.Equal(t, 200.0, profile2.TotalDuration)
|
||||
}
|
||||
@@ -0,0 +1,684 @@
|
||||
{
|
||||
"version": "1.0",
|
||||
"activities": [
|
||||
{
|
||||
"name": "CloneRepoActivity",
|
||||
"description": "Clone a Git repository to the worker filesystem",
|
||||
"category": "repository",
|
||||
"inputs": {
|
||||
"repo": {
|
||||
"type": "string",
|
||||
"description": "Git repository URL",
|
||||
"required": true
|
||||
},
|
||||
"branch": {
|
||||
"type": "string",
|
||||
"description": "Git branch to clone (default: main)",
|
||||
"required": false,
|
||||
"default": "main"
|
||||
},
|
||||
"depth": {
|
||||
"type": "integer",
|
||||
"description": "Shallow clone depth (optional)",
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Local filesystem path where repo was cloned"
|
||||
},
|
||||
"commit": {
|
||||
"type": "string",
|
||||
"description": "Current commit hash"
|
||||
},
|
||||
"branch": {
|
||||
"type": "string",
|
||||
"description": "Current branch name"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "5m",
|
||||
"isFlaky": false,
|
||||
"recommendedRetries": 1,
|
||||
"retryBackoff": 1.5,
|
||||
"dependencies": [],
|
||||
"notes": "Network-dependent, may timeout on slow connections"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "AnalyzeCodeActivity",
|
||||
"description": "Analyze code quality, structure, and metrics using ast-grep and pi CLI",
|
||||
"category": "analysis",
|
||||
"inputs": {
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Local filesystem path to analyze",
|
||||
"required": true
|
||||
},
|
||||
"language": {
|
||||
"type": "string",
|
||||
"description": "Programming language (go, python, javascript, etc)",
|
||||
"required": false
|
||||
},
|
||||
"depth": {
|
||||
"type": "integer",
|
||||
"description": "Analysis depth (1=shallow, 5=deep)",
|
||||
"required": false,
|
||||
"default": 3
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"quality": {
|
||||
"type": "number",
|
||||
"description": "Quality score 0-1.0"
|
||||
},
|
||||
"metrics": {
|
||||
"type": "object",
|
||||
"description": "Code metrics (LOC, complexity, etc)"
|
||||
},
|
||||
"issues": {
|
||||
"type": "array",
|
||||
"description": "List of identified issues"
|
||||
},
|
||||
"summary": {
|
||||
"type": "string",
|
||||
"description": "Human-readable analysis summary"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "10m",
|
||||
"isFlaky": true,
|
||||
"recommendedRetries": 3,
|
||||
"retryBackoff": 2.0,
|
||||
"dependencies": ["CloneRepoActivity"],
|
||||
"notes": "CPU-intensive, can timeout on large repos. Flaky on memory pressure."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "SecurityScanActivity",
|
||||
"description": "Run security scanning (SAST) on codebase",
|
||||
"category": "security",
|
||||
"inputs": {
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Local filesystem path to scan",
|
||||
"required": true
|
||||
},
|
||||
"severity": {
|
||||
"type": "string",
|
||||
"description": "Minimum severity level (low, medium, high, critical)",
|
||||
"required": false,
|
||||
"default": "medium"
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"vulnerabilities": {
|
||||
"type": "array",
|
||||
"description": "List of vulnerabilities found"
|
||||
},
|
||||
"securityScore": {
|
||||
"type": "number",
|
||||
"description": "Security score 0-100"
|
||||
},
|
||||
"riskLevel": {
|
||||
"type": "string",
|
||||
"description": "Risk level (low, medium, high, critical)"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "8m",
|
||||
"isFlaky": false,
|
||||
"recommendedRetries": 2,
|
||||
"retryBackoff": 1.5,
|
||||
"dependencies": ["CloneRepoActivity"],
|
||||
"notes": "Network calls for vulnerability databases may timeout"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "GenerateReportActivity",
|
||||
"description": "Generate comprehensive report from analysis and scan results",
|
||||
"category": "reporting",
|
||||
"inputs": {
|
||||
"analysisResult": {
|
||||
"type": "object",
|
||||
"description": "Output from AnalyzeCodeActivity",
|
||||
"required": true
|
||||
},
|
||||
"securityResult": {
|
||||
"type": "object",
|
||||
"description": "Output from SecurityScanActivity",
|
||||
"required": true
|
||||
},
|
||||
"format": {
|
||||
"type": "string",
|
||||
"description": "Report format (markdown, html, json)",
|
||||
"required": false,
|
||||
"default": "markdown"
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"report": {
|
||||
"type": "string",
|
||||
"description": "Generated report content"
|
||||
},
|
||||
"reportPath": {
|
||||
"type": "string",
|
||||
"description": "Path to saved report file"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "2m",
|
||||
"isFlaky": false,
|
||||
"recommendedRetries": 1,
|
||||
"retryBackoff": 1.0,
|
||||
"dependencies": ["AnalyzeCodeActivity", "SecurityScanActivity"],
|
||||
"notes": "CPU-light, reliable. Depends on upstream results."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "DeploymentPreCheckActivity",
|
||||
"description": "Validate readiness for deployment (linting, tests, etc)",
|
||||
"category": "deployment",
|
||||
"inputs": {
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Local filesystem path to check",
|
||||
"required": true
|
||||
},
|
||||
"checkType": {
|
||||
"type": "string",
|
||||
"description": "Type of check (lint, test, build, all)",
|
||||
"required": false,
|
||||
"default": "all"
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"passed": {
|
||||
"type": "boolean",
|
||||
"description": "Whether all checks passed"
|
||||
},
|
||||
"failures": {
|
||||
"type": "array",
|
||||
"description": "List of failed checks"
|
||||
},
|
||||
"warnings": {
|
||||
"type": "array",
|
||||
"description": "List of warnings"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "15m",
|
||||
"isFlaky": true,
|
||||
"recommendedRetries": 2,
|
||||
"retryBackoff": 2.0,
|
||||
"dependencies": ["CloneRepoActivity"],
|
||||
"notes": "Very flaky - tests are non-deterministic, network issues, race conditions. Retry 2x."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "NotifyStatusActivity",
|
||||
"description": "Send notifications to Slack, email, or webhook",
|
||||
"category": "notification",
|
||||
"inputs": {
|
||||
"channel": {
|
||||
"type": "string",
|
||||
"description": "Target channel or email",
|
||||
"required": true
|
||||
},
|
||||
"status": {
|
||||
"type": "string",
|
||||
"description": "Status to report (success, failure, warning)",
|
||||
"required": true
|
||||
},
|
||||
"message": {
|
||||
"type": "string",
|
||||
"description": "Message body",
|
||||
"required": true
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"notificationId": {
|
||||
"type": "string",
|
||||
"description": "ID of sent notification"
|
||||
},
|
||||
"timestamp": {
|
||||
"type": "string",
|
||||
"description": "When notification was sent"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "3m",
|
||||
"isFlaky": true,
|
||||
"recommendedRetries": 3,
|
||||
"retryBackoff": 1.5,
|
||||
"dependencies": [],
|
||||
"notes": "Network-dependent, may fail due to network or external service issues. Retry 3x."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "ApproveWorkflowActivity",
|
||||
"description": "Human approval step or automated policy check",
|
||||
"category": "approval",
|
||||
"inputs": {
|
||||
"workflowId": {
|
||||
"type": "string",
|
||||
"description": "ID of workflow awaiting approval",
|
||||
"required": true
|
||||
},
|
||||
"requiredApprovals": {
|
||||
"type": "integer",
|
||||
"description": "Number of approvals needed (default 1)",
|
||||
"required": false,
|
||||
"default": 1
|
||||
},
|
||||
"timeoutMinutes": {
|
||||
"type": "integer",
|
||||
"description": "Minutes to wait for approval",
|
||||
"required": false,
|
||||
"default": 60
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"approved": {
|
||||
"type": "boolean",
|
||||
"description": "Whether approved"
|
||||
},
|
||||
"approver": {
|
||||
"type": "string",
|
||||
"description": "Who approved (if approved)"
|
||||
},
|
||||
"timestamp": {
|
||||
"type": "string",
|
||||
"description": "When approval was given"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "120m",
|
||||
"isFlaky": false,
|
||||
"recommendedRetries": 1,
|
||||
"retryBackoff": 1.0,
|
||||
"dependencies": [],
|
||||
"notes": "Waits for human input. Long timeout. Cannot retry (user input is irrevocable)."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "ArchiveResultsActivity",
|
||||
"description": "Archive analysis results to cloud storage (S3, GCS)",
|
||||
"category": "storage",
|
||||
"inputs": {
|
||||
"reportPath": {
|
||||
"type": "string",
|
||||
"description": "Path to report to archive",
|
||||
"required": true
|
||||
},
|
||||
"destination": {
|
||||
"type": "string",
|
||||
"description": "Cloud destination (s3://bucket/path or gcs://bucket/path)",
|
||||
"required": true
|
||||
},
|
||||
"metadata": {
|
||||
"type": "object",
|
||||
"description": "Optional metadata tags",
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"archiveUrl": {
|
||||
"type": "string",
|
||||
"description": "URL of archived file"
|
||||
},
|
||||
"archiveSize": {
|
||||
"type": "integer",
|
||||
"description": "Size of archived file in bytes"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "5m",
|
||||
"isFlaky": true,
|
||||
"recommendedRetries": 2,
|
||||
"retryBackoff": 1.5,
|
||||
"dependencies": [],
|
||||
"notes": "Network-dependent. May fail on network issues or service throttling. Retry 2x."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "RetrieveMemoryActivity",
|
||||
"description": "Retrieve relevant knowledge, skills, and lessons from poimen-memory semantic search",
|
||||
"category": "memory",
|
||||
"inputs": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "Semantic search query",
|
||||
"required": true
|
||||
},
|
||||
"project": {
|
||||
"type": "string",
|
||||
"description": "Memory project (default: poimen)",
|
||||
"required": false,
|
||||
"default": "poimen"
|
||||
},
|
||||
"scope": {
|
||||
"type": "string",
|
||||
"description": "Retrieval scope: skills, lessons, references, all",
|
||||
"required": false,
|
||||
"default": "all"
|
||||
},
|
||||
"limit": {
|
||||
"type": "integer",
|
||||
"description": "Max results to return",
|
||||
"required": false,
|
||||
"default": 10
|
||||
},
|
||||
"tool": {
|
||||
"type": "string",
|
||||
"description": "Tool context for skill matching",
|
||||
"required": false
|
||||
},
|
||||
"task": {
|
||||
"type": "string",
|
||||
"description": "Task description for context retrieval",
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"skills": {
|
||||
"type": "array",
|
||||
"description": "Relevant skills found"
|
||||
},
|
||||
"lessons": {
|
||||
"type": "array",
|
||||
"description": "Relevant lessons/knowledge found"
|
||||
},
|
||||
"references": {
|
||||
"type": "array",
|
||||
"description": "Reference documents found"
|
||||
},
|
||||
"totalResults": {
|
||||
"type": "integer",
|
||||
"description": "Total results found"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "30s",
|
||||
"isFlaky": true,
|
||||
"recommendedRetries": 2,
|
||||
"retryBackoff": 1.5,
|
||||
"dependencies": [],
|
||||
"notes": "Network-dependent. First activity to run for context-aware routing. Fast timeout."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "AssumeRoleActivity",
|
||||
"description": "Request temporary JWT token for accessing LLM APIs (like AWS AssumeRole)",
|
||||
"category": "authentication",
|
||||
"inputs": {
|
||||
"identity": {
|
||||
"type": "string",
|
||||
"description": "User/service identity requesting access",
|
||||
"required": true,
|
||||
"examples": ["[email protected]", "service:poimen-worker"]
|
||||
},
|
||||
"clientId": {
|
||||
"type": "string",
|
||||
"description": "OAuth2 client ID (from vault if not provided)",
|
||||
"required": false
|
||||
},
|
||||
"clientSecret": {
|
||||
"type": "string",
|
||||
"description": "OAuth2 client secret (from vault if not provided)",
|
||||
"required": false
|
||||
},
|
||||
"scope": {
|
||||
"type": "string",
|
||||
"description": "Scope of access (e.g., 'llm:read' or 'llm:read llm:write')",
|
||||
"required": true,
|
||||
"examples": ["llm:read", "llm:read llm:write", "llm:admin"]
|
||||
},
|
||||
"durationSeconds": {
|
||||
"type": "integer",
|
||||
"description": "Token validity duration in seconds (default: 3600, max: 86400)",
|
||||
"required": false,
|
||||
"default": 3600
|
||||
},
|
||||
"authServerUrl": {
|
||||
"type": "string",
|
||||
"description": "Auth server URL (from AUTH_SERVER_URL env if not provided)",
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"token": {
|
||||
"type": "string",
|
||||
"description": "JWT token for calling api.riotpiao.com"
|
||||
},
|
||||
"expiresAt": {
|
||||
"type": "integer",
|
||||
"description": "Token expiration time (Unix timestamp)"
|
||||
},
|
||||
"expiresIn": {
|
||||
"type": "integer",
|
||||
"description": "Seconds until token expires"
|
||||
},
|
||||
"tokenType": {
|
||||
"type": "string",
|
||||
"description": "Token type (typically 'Bearer')"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "30s",
|
||||
"isFlaky": false,
|
||||
"recommendedRetries": 2,
|
||||
"retryBackoff": 1.5,
|
||||
"dependencies": [],
|
||||
"notes": "Must run before LLM Router to provide auth token. Call early in workflow."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "LLMInferenceActivity",
|
||||
"description": "Call LLM API with custom prompt and get response text",
|
||||
"category": "llm",
|
||||
"inputs": {
|
||||
"model": {
|
||||
"type": "string",
|
||||
"description": "Model ID (reasoning, ornith:35b, ornith:13b, qwen2.5:3b)",
|
||||
"required": true,
|
||||
"examples": ["reasoning", "ornith:35b"]
|
||||
},
|
||||
"system_prompt": {
|
||||
"type": "string",
|
||||
"description": "System instruction for the model",
|
||||
"required": false,
|
||||
"default": ""
|
||||
},
|
||||
"user_prompt": {
|
||||
"type": "string",
|
||||
"description": "User message to send to the model",
|
||||
"required": true
|
||||
},
|
||||
"temperature": {
|
||||
"type": "number",
|
||||
"description": "Sampling temperature (0.0-1.0, higher=more creative)",
|
||||
"required": false,
|
||||
"default": 0.7
|
||||
},
|
||||
"max_tokens": {
|
||||
"type": "integer",
|
||||
"description": "Maximum tokens in response",
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"response": {
|
||||
"type": "string",
|
||||
"description": "LLM response text"
|
||||
},
|
||||
"model": {
|
||||
"type": "string",
|
||||
"description": "Model used for inference"
|
||||
},
|
||||
"stop_reason": {
|
||||
"type": "string",
|
||||
"description": "Why inference stopped (stop_sequence, length, etc)"
|
||||
},
|
||||
"tokens_used": {
|
||||
"type": "integer",
|
||||
"description": "Total tokens consumed"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "120s",
|
||||
"isFlaky": true,
|
||||
"recommendedRetries": 2,
|
||||
"retryBackoff": 2.0,
|
||||
"dependencies": [],
|
||||
"notes": "API-dependent. Network flaky. Use for single prompts. See LLMBatchInferenceActivity for multiple."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "LLMBatchInferenceActivity",
|
||||
"description": "Call LLM API multiple times sequentially with different prompts",
|
||||
"category": "llm",
|
||||
"inputs": {
|
||||
"model": {
|
||||
"type": "string",
|
||||
"description": "Model ID (reasoning, ornith:35b, ornith:13b, qwen2.5:3b)",
|
||||
"required": true
|
||||
},
|
||||
"system_prompt": {
|
||||
"type": "string",
|
||||
"description": "System instruction (same for all prompts)",
|
||||
"required": false
|
||||
},
|
||||
"prompts": {
|
||||
"type": "array",
|
||||
"description": "List of user prompts to process",
|
||||
"required": true,
|
||||
"items": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"temperature": {
|
||||
"type": "number",
|
||||
"description": "Sampling temperature (0.0-1.0)",
|
||||
"required": false,
|
||||
"default": 0.7
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"responses": {
|
||||
"type": "array",
|
||||
"description": "List of LLM responses (parallel to input prompts)",
|
||||
"items": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"model": {
|
||||
"type": "string",
|
||||
"description": "Model used"
|
||||
},
|
||||
"errors": {
|
||||
"type": "array",
|
||||
"description": "Error messages for failed prompts",
|
||||
"items": {
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "600s",
|
||||
"isFlaky": true,
|
||||
"recommendedRetries": 1,
|
||||
"retryBackoff": 2.0,
|
||||
"dependencies": [],
|
||||
"notes": "Sequential processing of multiple prompts. Use for batch analysis, summarization, etc."
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "CanvasReasonerActivity",
|
||||
"description": "Use LLM reasoning to infer and suggest connections between workflow activities",
|
||||
"category": "workflow",
|
||||
"inputs": {
|
||||
"nodes": {
|
||||
"type": "array",
|
||||
"description": "Canvas workflow nodes to analyze",
|
||||
"required": true,
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"id": {"type": "string"},
|
||||
"type": {"type": "string"},
|
||||
"label": {"type": "string"}
|
||||
}
|
||||
}
|
||||
},
|
||||
"edges": {
|
||||
"type": "array",
|
||||
"description": "Existing edges in the workflow",
|
||||
"required": false,
|
||||
"items": {
|
||||
"type": "object"
|
||||
}
|
||||
},
|
||||
"preserve_existing": {
|
||||
"type": "boolean",
|
||||
"description": "If true, only suggest new edges; if false, redesign entire workflow",
|
||||
"required": false,
|
||||
"default": true
|
||||
},
|
||||
"auth_token": {
|
||||
"type": "string",
|
||||
"description": "JWT token for authenticated LLM calls",
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"outputs": {
|
||||
"suggested_edges": {
|
||||
"type": "array",
|
||||
"description": "Edges suggested by LLM reasoning",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"source": {"type": "string"},
|
||||
"target": {"type": "string"}
|
||||
}
|
||||
}
|
||||
},
|
||||
"reasoning": {
|
||||
"type": "string",
|
||||
"description": "LLM explanation of suggested connections"
|
||||
},
|
||||
"confidence": {
|
||||
"type": "number",
|
||||
"description": "Confidence score (0.0-1.0) of the suggestions"
|
||||
}
|
||||
},
|
||||
"constraints": {
|
||||
"defaultTimeout": "120s",
|
||||
"isFlaky": true,
|
||||
"recommendedRetries": 2,
|
||||
"retryBackoff": 2.0,
|
||||
"dependencies": [],
|
||||
"notes": "Uses reasoning model to analyze workflow logic. Good for understanding data flow and connections between activities."
|
||||
}
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"totalActivities": 13,
|
||||
"lastUpdated": "2025-09-05T00:00:00Z",
|
||||
"categories": {
|
||||
"repository": 1,
|
||||
"analysis": 1,
|
||||
"security": 1,
|
||||
"reporting": 1,
|
||||
"deployment": 1,
|
||||
"notification": 1,
|
||||
"approval": 1,
|
||||
"storage": 1,
|
||||
"memory": 1,
|
||||
"authentication": 1,
|
||||
"llm": 2,
|
||||
"workflow": 1
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,179 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/pkg/db"
|
||||
)
|
||||
|
||||
// CanvasConverter converts visual canvas to executable WorkflowSpec
|
||||
type CanvasConverter struct {
|
||||
validator *CanvasValidator
|
||||
}
|
||||
|
||||
// NewCanvasConverter creates a converter
|
||||
func NewCanvasConverter() *CanvasConverter {
|
||||
return &CanvasConverter{
|
||||
validator: NewCanvasValidator(),
|
||||
}
|
||||
}
|
||||
|
||||
// CanvasToWorkflowSpec converts canvas to WorkflowSpec
|
||||
func (cc *CanvasConverter) CanvasToWorkflowSpec(canvas *db.Canvas) (*WorkflowSpec, error) {
|
||||
// Validate first
|
||||
if err := cc.validator.ValidateCanvas(canvas); err != nil {
|
||||
return nil, fmt.Errorf("canvas validation failed: %w", err)
|
||||
}
|
||||
|
||||
// Get topological order
|
||||
sortedNodes, err := cc.validator.TopoSort(canvas.Nodes, canvas.Edges)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("topological sort failed: %w", err)
|
||||
}
|
||||
|
||||
// Build states from sorted nodes
|
||||
states := []State{}
|
||||
nodeToState := make(map[string]int) // node ID to state index
|
||||
|
||||
for i, node := range sortedNodes {
|
||||
state := cc.nodeToState(node, canvas.Edges)
|
||||
states = append(states, state)
|
||||
nodeToState[node.ID] = i
|
||||
}
|
||||
|
||||
// Wire up transitions
|
||||
for i, node := range sortedNodes {
|
||||
outgoing := cc.getOutgoingEdges(node.ID, canvas.Edges)
|
||||
|
||||
if len(outgoing) == 0 {
|
||||
// Last state - no transitions
|
||||
continue
|
||||
}
|
||||
|
||||
if len(outgoing) == 1 {
|
||||
// Single outgoing edge
|
||||
targetNode := outgoing[0]
|
||||
targetIdx := nodeToState[targetNode]
|
||||
if targetIdx > i {
|
||||
states[i].Next = states[targetIdx].Name
|
||||
}
|
||||
} else {
|
||||
// Multiple outgoing edges - parallel
|
||||
states[i].Type = "Parallel"
|
||||
branches := []interface{}{}
|
||||
for _, targetNode := range outgoing {
|
||||
branches = append(branches, map[string]string{
|
||||
"state": states[nodeToState[targetNode]].Name,
|
||||
})
|
||||
}
|
||||
if states[i].Branches == nil {
|
||||
states[i].Branches = branches
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: canvas.Name,
|
||||
Input: map[string]interface{}{},
|
||||
States: states,
|
||||
}
|
||||
|
||||
return spec, nil
|
||||
}
|
||||
|
||||
// nodeToState converts a canvas node to a workflow state
|
||||
func (cc *CanvasConverter) nodeToState(node db.WorkflowNode, edges []db.WorkflowEdge) State {
|
||||
// Map node type to activity name
|
||||
activityName := cc.mapActivityType(node.Type)
|
||||
|
||||
state := State{
|
||||
Name: node.ID,
|
||||
Type: TaskActivity,
|
||||
Activity: activityName,
|
||||
Retry: &RetryPolicy{MaxAttempts: 3, BackoffSeconds: 2},
|
||||
Timeout: "300s",
|
||||
Parameters: node.Data,
|
||||
}
|
||||
|
||||
return state
|
||||
}
|
||||
|
||||
// mapActivityType maps canvas activity type to Poimen activity
|
||||
func (cc *CanvasConverter) mapActivityType(canvasType string) string {
|
||||
typeMap := map[string]string{
|
||||
"clone-repo": "CloneRepoActivity",
|
||||
"analyze-code": "AnalyzeCodeActivity",
|
||||
"security-scan": "SecurityScanActivity",
|
||||
"generate-report": "GenerateReportActivity",
|
||||
"deployment-precheck": "DeploymentPreCheckActivity",
|
||||
"notify-status": "NotifyStatusActivity",
|
||||
"approve-workflow": "ApproveWorkflowActivity",
|
||||
"archive-results": "ArchiveResultsActivity",
|
||||
"retrieve-memory": "RetrieveMemoryActivity",
|
||||
"assume-role": "AssumeRoleActivity",
|
||||
"llm-inference": "LLMInferenceActivity",
|
||||
"llm-batch-inference": "LLMBatchInferenceActivity",
|
||||
"canvas-reasoner": "CanvasReasonerActivity",
|
||||
}
|
||||
|
||||
if mapped, ok := typeMap[canvasType]; ok {
|
||||
return mapped
|
||||
}
|
||||
|
||||
return canvasType // fallback to type as-is
|
||||
}
|
||||
|
||||
// getOutgoingEdges returns target node IDs for a given source node
|
||||
func (cc *CanvasConverter) getOutgoingEdges(nodeID string, edges []db.WorkflowEdge) []string {
|
||||
targets := []string{}
|
||||
seen := make(map[string]bool)
|
||||
|
||||
for _, edge := range edges {
|
||||
if edge.Source == nodeID && !seen[edge.Target] {
|
||||
targets = append(targets, edge.Target)
|
||||
seen[edge.Target] = true
|
||||
}
|
||||
}
|
||||
|
||||
return targets
|
||||
}
|
||||
|
||||
// CanvasToExecutionPlan converts canvas to sequential activity list
|
||||
func (cc *CanvasConverter) CanvasToExecutionPlan(canvas *db.Canvas) ([]ExecutionStep, error) {
|
||||
// Validate first
|
||||
if err := cc.validator.ValidateCanvas(canvas); err != nil {
|
||||
return nil, fmt.Errorf("canvas validation failed: %w", err)
|
||||
}
|
||||
|
||||
// Get topological order
|
||||
sortedNodes, err := cc.validator.TopoSort(canvas.Nodes, canvas.Edges)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("topological sort failed: %w", err)
|
||||
}
|
||||
|
||||
steps := []ExecutionStep{}
|
||||
for i, node := range sortedNodes {
|
||||
step := ExecutionStep{
|
||||
Index: i,
|
||||
NodeID: node.ID,
|
||||
ActivityName: cc.mapActivityType(node.Type),
|
||||
Label: node.Label,
|
||||
Parameters: node.Data,
|
||||
Timeout: "300s",
|
||||
}
|
||||
steps = append(steps, step)
|
||||
}
|
||||
|
||||
return steps, nil
|
||||
}
|
||||
|
||||
// ExecutionStep represents one activity in execution plan
|
||||
type ExecutionStep struct {
|
||||
Index int `json:"index"`
|
||||
NodeID string `json:"node_id"`
|
||||
ActivityName string `json:"activity_name"`
|
||||
Label string `json:"label"`
|
||||
Parameters map[string]interface{} `json:"parameters"`
|
||||
Timeout string `json:"timeout"`
|
||||
DependsOn []int `json:"depends_on,omitempty"` // Indices of predecessor steps
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/rockliang/poimen/workflows/pkg/db"
|
||||
)
|
||||
|
||||
// CanvasValidator validates React Flow canvas (nodes + edges)
|
||||
type CanvasValidator struct {
|
||||
activityRegistry map[string]bool
|
||||
}
|
||||
|
||||
// NewCanvasValidator creates validator with activity registry
|
||||
func NewCanvasValidator() *CanvasValidator {
|
||||
return &CanvasValidator{
|
||||
activityRegistry: map[string]bool{
|
||||
"clone-repo": true,
|
||||
"analyze-code": true,
|
||||
"security-scan": true,
|
||||
"generate-report": true,
|
||||
"deployment-precheck": true,
|
||||
"notify-status": true,
|
||||
"approve-workflow": true,
|
||||
"archive-results": true,
|
||||
"retrieve-memory": true,
|
||||
"assume-role": true,
|
||||
"llm-inference": true,
|
||||
"llm-batch-inference": true,
|
||||
"canvas-reasoner": true,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateCanvas checks canvas structure, connectivity, and DAG
|
||||
func (cv *CanvasValidator) ValidateCanvas(canvas *db.Canvas) error {
|
||||
if canvas == nil {
|
||||
return fmt.Errorf("canvas is nil")
|
||||
}
|
||||
|
||||
if len(canvas.Nodes) == 0 {
|
||||
return fmt.Errorf("canvas has no nodes")
|
||||
}
|
||||
|
||||
// Step 1: Validate nodes
|
||||
if err := cv.validateNodes(canvas.Nodes); err != nil {
|
||||
return fmt.Errorf("node validation failed: %w", err)
|
||||
}
|
||||
|
||||
// Step 2: Validate edges
|
||||
if err := cv.validateEdges(canvas.Nodes, canvas.Edges); err != nil {
|
||||
return fmt.Errorf("edge validation failed: %w", err)
|
||||
}
|
||||
|
||||
// Step 3: Check for cycles (must be DAG)
|
||||
if err := cv.detectCycles(canvas.Nodes, canvas.Edges); err != nil {
|
||||
return fmt.Errorf("cycle detected: %w", err)
|
||||
}
|
||||
|
||||
// Step 4: Check connectivity (all nodes reachable from start)
|
||||
if err := cv.validateConnectivity(canvas.Nodes, canvas.Edges); err != nil {
|
||||
return fmt.Errorf("connectivity check failed: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateNodes checks each node has required fields and valid type
|
||||
func (cv *CanvasValidator) validateNodes(nodes []db.WorkflowNode) error {
|
||||
if len(nodes) == 0 {
|
||||
return fmt.Errorf("no nodes in canvas")
|
||||
}
|
||||
|
||||
nodeIds := make(map[string]bool)
|
||||
|
||||
for i, node := range nodes {
|
||||
// Check required fields
|
||||
if node.ID == "" {
|
||||
return fmt.Errorf("node[%d] has empty ID", i)
|
||||
}
|
||||
|
||||
if nodeIds[node.ID] {
|
||||
return fmt.Errorf("node[%d] has duplicate ID: %s", i, node.ID)
|
||||
}
|
||||
nodeIds[node.ID] = true
|
||||
|
||||
if node.Label == "" {
|
||||
return fmt.Errorf("node[%d] (%s) has empty label", i, node.ID)
|
||||
}
|
||||
|
||||
if node.Position == nil {
|
||||
return fmt.Errorf("node[%d] (%s) has no position", i, node.ID)
|
||||
}
|
||||
|
||||
// Check activity type (if present)
|
||||
if node.Type != "" && !cv.activityRegistry[strings.ToLower(node.Type)] {
|
||||
return fmt.Errorf("node[%d] (%s) has unknown activity type: %s", i, node.ID, node.Type)
|
||||
}
|
||||
|
||||
// Check data structure
|
||||
if node.Data == nil {
|
||||
return fmt.Errorf("node[%d] (%s) has no data", i, node.ID)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateEdges checks edges reference valid nodes
|
||||
func (cv *CanvasValidator) validateEdges(nodes []db.WorkflowNode, edges []db.WorkflowEdge) error {
|
||||
nodeIds := make(map[string]bool)
|
||||
for _, node := range nodes {
|
||||
nodeIds[node.ID] = true
|
||||
}
|
||||
|
||||
for i, edge := range edges {
|
||||
// Check required fields
|
||||
if edge.Source == "" {
|
||||
return fmt.Errorf("edge[%d] has empty source", i)
|
||||
}
|
||||
|
||||
if edge.Target == "" {
|
||||
return fmt.Errorf("edge[%d] has empty target", i)
|
||||
}
|
||||
|
||||
// Check source node exists
|
||||
if !nodeIds[edge.Source] {
|
||||
return fmt.Errorf("edge[%d] references unknown source node: %s", i, edge.Source)
|
||||
}
|
||||
|
||||
// Check target node exists
|
||||
if !nodeIds[edge.Target] {
|
||||
return fmt.Errorf("edge[%d] references unknown target node: %s", i, edge.Target)
|
||||
}
|
||||
|
||||
// Check self-loops (discouraged but allow for now)
|
||||
if edge.Source == edge.Target {
|
||||
// Could warn here but not fail
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// detectCycles checks for cycles in the DAG (must be acyclic)
|
||||
func (cv *CanvasValidator) detectCycles(nodes []db.WorkflowNode, edges []db.WorkflowEdge) error {
|
||||
// Build adjacency list
|
||||
graph := make(map[string][]string)
|
||||
inDegree := make(map[string]int)
|
||||
|
||||
for _, node := range nodes {
|
||||
graph[node.ID] = []string{}
|
||||
inDegree[node.ID] = 0
|
||||
}
|
||||
|
||||
for _, edge := range edges {
|
||||
graph[edge.Source] = append(graph[edge.Source], edge.Target)
|
||||
inDegree[edge.Target]++
|
||||
}
|
||||
|
||||
// Kahn's algorithm: topological sort
|
||||
queue := []string{}
|
||||
for _, node := range nodes {
|
||||
if inDegree[node.ID] == 0 {
|
||||
queue = append(queue, node.ID)
|
||||
}
|
||||
}
|
||||
|
||||
processed := 0
|
||||
for len(queue) > 0 {
|
||||
// Dequeue
|
||||
current := queue[0]
|
||||
queue = queue[1:]
|
||||
processed++
|
||||
|
||||
// Visit neighbors
|
||||
for _, neighbor := range graph[current] {
|
||||
inDegree[neighbor]--
|
||||
if inDegree[neighbor] == 0 {
|
||||
queue = append(queue, neighbor)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If we didn't process all nodes, there's a cycle
|
||||
if processed != len(nodes) {
|
||||
return fmt.Errorf("graph has cycle (processed %d/%d nodes)", processed, len(nodes))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateConnectivity checks all nodes are reachable from start nodes
|
||||
func (cv *CanvasValidator) validateConnectivity(nodes []db.WorkflowNode, edges []db.WorkflowEdge) error {
|
||||
if len(nodes) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Build adjacency list
|
||||
graph := make(map[string][]string)
|
||||
inDegree := make(map[string]int)
|
||||
|
||||
for _, node := range nodes {
|
||||
graph[node.ID] = []string{}
|
||||
inDegree[node.ID] = 0
|
||||
}
|
||||
|
||||
for _, edge := range edges {
|
||||
graph[edge.Source] = append(graph[edge.Source], edge.Target)
|
||||
inDegree[edge.Target]++
|
||||
}
|
||||
|
||||
// Find start nodes (in-degree 0)
|
||||
startNodes := []string{}
|
||||
for _, node := range nodes {
|
||||
if inDegree[node.ID] == 0 {
|
||||
startNodes = append(startNodes, node.ID)
|
||||
}
|
||||
}
|
||||
|
||||
if len(startNodes) == 0 {
|
||||
return fmt.Errorf("no start nodes found (all nodes have incoming edges)")
|
||||
}
|
||||
|
||||
// BFS from all start nodes
|
||||
visited := make(map[string]bool)
|
||||
queue := startNodes
|
||||
|
||||
for len(queue) > 0 {
|
||||
// Dequeue
|
||||
current := queue[0]
|
||||
queue = queue[1:]
|
||||
|
||||
if visited[current] {
|
||||
continue
|
||||
}
|
||||
visited[current] = true
|
||||
|
||||
// Visit neighbors
|
||||
for _, neighbor := range graph[current] {
|
||||
if !visited[neighbor] {
|
||||
queue = append(queue, neighbor)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check all nodes were visited
|
||||
if len(visited) != len(nodes) {
|
||||
unreached := []string{}
|
||||
for _, node := range nodes {
|
||||
if !visited[node.ID] {
|
||||
unreached = append(unreached, node.ID)
|
||||
}
|
||||
}
|
||||
return fmt.Errorf("unreachable nodes: %v", unreached)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// TopoSort returns nodes in topological order (execution order)
|
||||
func (cv *CanvasValidator) TopoSort(nodes []db.WorkflowNode, edges []db.WorkflowEdge) ([]db.WorkflowNode, error) {
|
||||
if len(nodes) == 0 {
|
||||
return []db.WorkflowNode{}, nil
|
||||
}
|
||||
|
||||
// Build adjacency list and in-degree map
|
||||
graph := make(map[string][]string)
|
||||
inDegree := make(map[string]int)
|
||||
nodeMap := make(map[string]db.WorkflowNode)
|
||||
|
||||
for _, node := range nodes {
|
||||
graph[node.ID] = []string{}
|
||||
inDegree[node.ID] = 0
|
||||
nodeMap[node.ID] = node
|
||||
}
|
||||
|
||||
for _, edge := range edges {
|
||||
graph[edge.Source] = append(graph[edge.Source], edge.Target)
|
||||
inDegree[edge.Target]++
|
||||
}
|
||||
|
||||
// Kahn's algorithm
|
||||
queue := []string{}
|
||||
for _, node := range nodes {
|
||||
if inDegree[node.ID] == 0 {
|
||||
queue = append(queue, node.ID)
|
||||
}
|
||||
}
|
||||
|
||||
result := []db.WorkflowNode{}
|
||||
processed := make(map[string]bool)
|
||||
|
||||
for len(queue) > 0 {
|
||||
// Dequeue
|
||||
current := queue[0]
|
||||
queue = queue[1:]
|
||||
|
||||
result = append(result, nodeMap[current])
|
||||
processed[current] = true
|
||||
|
||||
// Visit neighbors
|
||||
for _, neighbor := range graph[current] {
|
||||
inDegree[neighbor]--
|
||||
if inDegree[neighbor] == 0 {
|
||||
queue = append(queue, neighbor)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(result) != len(nodes) {
|
||||
return nil, fmt.Errorf("topological sort failed: graph has cycle")
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,224 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// ActivityExecutor defines how to execute an activity
|
||||
type ActivityExecutor interface {
|
||||
// Execute runs the activity with given parameters
|
||||
Execute(ctx context.Context, activityName string, params map[string]interface{}) (interface{}, error)
|
||||
}
|
||||
|
||||
// TemporalActivityExecutor executes activities via Temporal
|
||||
type TemporalActivityExecutor struct {
|
||||
// This would be implemented by workflow context
|
||||
executor func(context.Context, string, interface{}) error
|
||||
}
|
||||
|
||||
// StateTransitioner defines state machine transitions
|
||||
type StateTransitioner interface {
|
||||
// CanTransition checks if transition is allowed
|
||||
CanTransition(from, to *State) bool
|
||||
// Transit performs the transition
|
||||
Transit(from, to *State) error
|
||||
}
|
||||
|
||||
// DefaultStateTransitioner implements basic transitions
|
||||
type DefaultStateTransitioner struct {
|
||||
validators []TransitionValidator
|
||||
}
|
||||
|
||||
// TransitionValidator validates a specific transition
|
||||
type TransitionValidator interface {
|
||||
Validate(from, to *State) error
|
||||
}
|
||||
|
||||
// NewDefaultStateTransitioner creates a new transitioner
|
||||
func NewDefaultStateTransitioner() *DefaultStateTransitioner {
|
||||
return &DefaultStateTransitioner{
|
||||
validators: []TransitionValidator{
|
||||
&StateTypeValidator{},
|
||||
&OutputMatchValidator{},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// CanTransition checks if transition is valid
|
||||
func (dst *DefaultStateTransitioner) CanTransition(from, to *State) bool {
|
||||
return dst.Transit(from, to) == nil
|
||||
}
|
||||
|
||||
// Transit validates and performs transition
|
||||
func (dst *DefaultStateTransitioner) Transit(from, to *State) error {
|
||||
for _, v := range dst.validators {
|
||||
if err := v.Validate(from, to); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// StateTypeValidator checks state type compatibility
|
||||
type StateTypeValidator struct{}
|
||||
|
||||
func (stv *StateTypeValidator) Validate(from, to *State) error {
|
||||
if from == nil {
|
||||
return nil // Initial transition
|
||||
}
|
||||
|
||||
// Can't transition from terminal states
|
||||
if from.Type == StateTypeFail {
|
||||
return fmt.Errorf("cannot transition from Fail state")
|
||||
}
|
||||
if from.End && from.Type != StateTypePass {
|
||||
return fmt.Errorf("cannot transition from end state")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// OutputMatchValidator checks output-input binding
|
||||
type OutputMatchValidator struct{}
|
||||
|
||||
func (omv *OutputMatchValidator) Validate(from, to *State) error {
|
||||
// Could validate that outputs from previous state match inputs needed
|
||||
return nil
|
||||
}
|
||||
|
||||
// ParameterBinder resolves parameters from context
|
||||
type ParameterBinder interface {
|
||||
// Bind resolves all parameters for a state
|
||||
Bind(state *State, context *ExecutionContext) (map[string]interface{}, error)
|
||||
}
|
||||
|
||||
// DefaultParameterBinder implements parameter resolution
|
||||
type DefaultParameterBinder struct {
|
||||
resolver *JSONPathResolver
|
||||
}
|
||||
|
||||
// NewDefaultParameterBinder creates a new binder
|
||||
func NewDefaultParameterBinder() *DefaultParameterBinder {
|
||||
return &DefaultParameterBinder{
|
||||
resolver: NewJSONPathResolver(nil, nil),
|
||||
}
|
||||
}
|
||||
|
||||
// Bind resolves all parameters
|
||||
func (dpb *DefaultParameterBinder) Bind(state *State, context *ExecutionContext) (map[string]interface{}, error) {
|
||||
dpb.resolver.input = context.Input
|
||||
dpb.resolver.stepResults = context.StepResults
|
||||
|
||||
resolved, err := dpb.resolver.ResolvePaths(state.Parameters)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parameter binding failed: %w", err)
|
||||
}
|
||||
|
||||
return resolved, nil
|
||||
}
|
||||
|
||||
// WorkflowValidator validates workflow specs
|
||||
type WorkflowValidator interface {
|
||||
// Validate checks if workflow is valid
|
||||
Validate(spec *WorkflowSpec) error
|
||||
}
|
||||
|
||||
// CompositeValidator combines multiple validators
|
||||
type CompositeValidator struct {
|
||||
validators []WorkflowValidator
|
||||
}
|
||||
|
||||
// NewCompositeValidator creates a composite validator
|
||||
func NewCompositeValidator(validators ...WorkflowValidator) *CompositeValidator {
|
||||
return &CompositeValidator{validators: validators}
|
||||
}
|
||||
|
||||
// Validate runs all validators
|
||||
func (cv *CompositeValidator) Validate(spec *WorkflowSpec) error {
|
||||
for _, v := range cv.validators {
|
||||
if err := v.Validate(spec); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// StateGraphValidator validates state graph structure
|
||||
type StateGraphValidator struct{}
|
||||
|
||||
func (sgv *StateGraphValidator) Validate(spec *WorkflowSpec) error {
|
||||
if spec == nil {
|
||||
return fmt.Errorf("workflow spec is nil")
|
||||
}
|
||||
if len(spec.States) == 0 {
|
||||
return fmt.Errorf("workflow has no states")
|
||||
}
|
||||
|
||||
stateMap := make(map[string]*State)
|
||||
for i := range spec.States {
|
||||
stateMap[spec.States[i].Name] = &spec.States[i]
|
||||
}
|
||||
|
||||
// Check all transitions point to valid states
|
||||
for _, state := range spec.States {
|
||||
if state.Type == StateTypeTask && !state.End {
|
||||
if state.Next == "" {
|
||||
return fmt.Errorf("state %s has no next state and is not end", state.Name)
|
||||
}
|
||||
if _, exists := stateMap[state.Next]; !exists {
|
||||
return fmt.Errorf("state %s references non-existent next state %s", state.Name, state.Next)
|
||||
}
|
||||
}
|
||||
|
||||
// Validate catch clauses
|
||||
for _, catch := range state.Catch {
|
||||
if _, exists := stateMap[catch.Next]; !exists {
|
||||
return fmt.Errorf("catch handler in %s references non-existent state %s", state.Name, catch.Next)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ActivityAvailabilityValidator validates activities exist
|
||||
type ActivityAvailabilityValidator struct {
|
||||
kb *KnowledgeBase
|
||||
}
|
||||
|
||||
// NewActivityAvailabilityValidator creates a new validator
|
||||
func NewActivityAvailabilityValidator(kb *KnowledgeBase) *ActivityAvailabilityValidator {
|
||||
return &ActivityAvailabilityValidator{kb: kb}
|
||||
}
|
||||
|
||||
// Validate checks all activities are available
|
||||
func (aav *ActivityAvailabilityValidator) Validate(spec *WorkflowSpec) error {
|
||||
for _, state := range spec.States {
|
||||
if state.Type == StateTypeTask {
|
||||
if !aav.kb.HasActivity(state.Resource) {
|
||||
return fmt.Errorf("activity %s not found in knowledge base", state.Resource)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// TimeoutValidator validates timeouts
|
||||
type TimeoutValidator struct{}
|
||||
|
||||
func (tv *TimeoutValidator) Validate(spec *WorkflowSpec) error {
|
||||
for _, state := range spec.States {
|
||||
if state.Timeout != "" {
|
||||
if _, err := parseDuration(state.Timeout); err != nil {
|
||||
return fmt.Errorf("invalid timeout in state %s: %w", state.Name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseDuration(d string) (interface{}, error) {
|
||||
// Placeholder for duration parsing
|
||||
return nil, nil
|
||||
}
|
||||
@@ -0,0 +1,230 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// JSONPathResolver resolves JSONPath expressions like ${input.repo}, ${Clone.output.path}
|
||||
type JSONPathResolver struct {
|
||||
input map[string]interface{}
|
||||
stepResults map[string]interface{}
|
||||
}
|
||||
|
||||
// NewJSONPathResolver creates a new resolver with input and step results
|
||||
func NewJSONPathResolver(input map[string]interface{}, stepResults map[string]interface{}) *JSONPathResolver {
|
||||
return &JSONPathResolver{
|
||||
input: input,
|
||||
stepResults: stepResults,
|
||||
}
|
||||
}
|
||||
|
||||
// Resolve resolves a single JSONPath expression
|
||||
// Supports: ${input.field}, ${StepName.output.field}, ${StepName.output.nested.field}
|
||||
func (r *JSONPathResolver) Resolve(expr string) (interface{}, error) {
|
||||
if expr == "" {
|
||||
return nil, fmt.Errorf("expression cannot be empty")
|
||||
}
|
||||
|
||||
// Check if it's a template expression (starts with ${ and ends with })
|
||||
if !strings.HasPrefix(expr, "${") || !strings.HasSuffix(expr, "}") {
|
||||
// Return as-is if not a template
|
||||
return expr, nil
|
||||
}
|
||||
|
||||
// Extract the path from ${...}
|
||||
path := strings.TrimPrefix(expr, "${")
|
||||
path = strings.TrimSuffix(path, "}")
|
||||
|
||||
return r.resolvePath(path)
|
||||
}
|
||||
|
||||
// ResolveString resolves a string that may contain multiple JSONPath expressions
|
||||
// Example: "Analysis at ${Clone.output.path} completed"
|
||||
func (r *JSONPathResolver) ResolveString(str string) (string, error) {
|
||||
// Find all ${...} patterns
|
||||
pattern := regexp.MustCompile(`\$\{[^}]+\}`)
|
||||
|
||||
result := str
|
||||
matches := pattern.FindAllString(str, -1)
|
||||
|
||||
for _, match := range matches {
|
||||
value, err := r.Resolve(match)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// Convert value to string
|
||||
strValue := fmt.Sprintf("%v", value)
|
||||
result = strings.ReplaceAll(result, match, strValue)
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// ResolvePaths resolves all JSONPath expressions in a map recursively
|
||||
func (r *JSONPathResolver) ResolvePaths(data map[string]interface{}) (map[string]interface{}, error) {
|
||||
result := make(map[string]interface{})
|
||||
|
||||
for key, value := range data {
|
||||
resolved, err := r.resolveValue(value)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to resolve key '%s': %w", key, err)
|
||||
}
|
||||
result[key] = resolved
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// resolvePath resolves a dot-separated path
|
||||
// Paths can be: input.field, StepName.output.field, etc.
|
||||
func (r *JSONPathResolver) resolvePath(path string) (interface{}, error) {
|
||||
parts := strings.Split(path, ".")
|
||||
if len(parts) == 0 {
|
||||
return nil, fmt.Errorf("invalid path: %s", path)
|
||||
}
|
||||
|
||||
// Check if first part is "input"
|
||||
if parts[0] == "input" {
|
||||
return r.resolveFromInput(parts[1:])
|
||||
}
|
||||
|
||||
// Otherwise, assume it's a step name
|
||||
stepName := parts[0]
|
||||
stepData, ok := r.stepResults[stepName]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("step '%s' not found in results", stepName)
|
||||
}
|
||||
|
||||
// Navigate through remaining parts
|
||||
return r.navigateObject(stepData, parts[1:])
|
||||
}
|
||||
|
||||
// resolveFromInput resolves path from input data
|
||||
func (r *JSONPathResolver) resolveFromInput(parts []string) (interface{}, error) {
|
||||
if len(parts) == 0 {
|
||||
return r.input, nil
|
||||
}
|
||||
|
||||
return r.navigateObject(r.input, parts)
|
||||
}
|
||||
|
||||
// navigateObject navigates through an object using path parts
|
||||
func (r *JSONPathResolver) navigateObject(obj interface{}, parts []string) (interface{}, error) {
|
||||
current := obj
|
||||
|
||||
for i, part := range parts {
|
||||
if current == nil {
|
||||
return nil, fmt.Errorf("cannot navigate through nil at part %d (%s)", i, part)
|
||||
}
|
||||
|
||||
// Handle map
|
||||
if mapObj, ok := current.(map[string]interface{}); ok {
|
||||
value, exists := mapObj[part]
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("key '%s' not found in object", part)
|
||||
}
|
||||
current = value
|
||||
continue
|
||||
}
|
||||
|
||||
// Handle map[string]string
|
||||
if strMap, ok := current.(map[string]string); ok {
|
||||
value, exists := strMap[part]
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("key '%s' not found in string map", part)
|
||||
}
|
||||
current = value
|
||||
continue
|
||||
}
|
||||
|
||||
// Cannot navigate further
|
||||
return nil, fmt.Errorf("cannot navigate through non-object type at part %d (%s)", i, part)
|
||||
}
|
||||
|
||||
return current, nil
|
||||
}
|
||||
|
||||
// resolveValue recursively resolves values (strings, maps, slices)
|
||||
func (r *JSONPathResolver) resolveValue(value interface{}) (interface{}, error) {
|
||||
if value == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
// Try to resolve as JSONPath
|
||||
if strings.Contains(v, "${") {
|
||||
// Check if it's a pure template (only one expression filling the whole string)
|
||||
if strings.HasPrefix(v, "${") && strings.HasSuffix(v, "}") && strings.Count(v, "${") == 1 {
|
||||
// Pure template - resolve as object
|
||||
resolved, err := r.Resolve(v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return resolved, nil
|
||||
}
|
||||
// String with embedded expressions - resolve as string
|
||||
resolved, err := r.ResolveString(v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return resolved, nil
|
||||
}
|
||||
return v, nil
|
||||
|
||||
case map[string]interface{}:
|
||||
// Recursively resolve map
|
||||
return r.ResolvePaths(v)
|
||||
|
||||
case []interface{}:
|
||||
// Recursively resolve slice
|
||||
result := make([]interface{}, len(v))
|
||||
for i, item := range v {
|
||||
resolved, err := r.resolveValue(item)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result[i] = resolved
|
||||
}
|
||||
return result, nil
|
||||
|
||||
default:
|
||||
// Return as-is for other types
|
||||
return v, nil
|
||||
}
|
||||
}
|
||||
|
||||
// ValidatePath checks if a path is valid (doesn't guarantee it resolves)
|
||||
func (r *JSONPathResolver) ValidatePath(path string) error {
|
||||
if !strings.Contains(path, ".") && path != "input" {
|
||||
return fmt.Errorf("invalid path: must contain '.' or be 'input'")
|
||||
}
|
||||
|
||||
parts := strings.Split(path, ".")
|
||||
if len(parts) == 0 {
|
||||
return fmt.Errorf("invalid path: no parts")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetAvailableSteps returns list of available steps in step results
|
||||
func (r *JSONPathResolver) GetAvailableSteps() []string {
|
||||
steps := make([]string, 0, len(r.stepResults))
|
||||
for step := range r.stepResults {
|
||||
steps = append(steps, step)
|
||||
}
|
||||
return steps
|
||||
}
|
||||
|
||||
// GetInputFields returns list of available input fields
|
||||
func (r *JSONPathResolver) GetInputFields() []string {
|
||||
fields := make([]string, 0, len(r.input))
|
||||
for field := range r.input {
|
||||
fields = append(fields, field)
|
||||
}
|
||||
return fields
|
||||
}
|
||||
@@ -0,0 +1,437 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestResolveInputField(t *testing.T) {
|
||||
input := map[string]interface{}{
|
||||
"repo": "https://github.com/test/repo",
|
||||
"branch": "main",
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, map[string]interface{}{})
|
||||
|
||||
// Test resolving input field
|
||||
value, err := resolver.Resolve("${input.repo}")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve: %v", err)
|
||||
}
|
||||
|
||||
if value != "https://github.com/test/repo" {
|
||||
t.Errorf("Expected 'https://github.com/test/repo', got %v", value)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveNestedField(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
stepResults := map[string]interface{}{
|
||||
"Clone": map[string]interface{}{
|
||||
"output": map[string]interface{}{
|
||||
"path": "/tmp/repo",
|
||||
"commit": "abc123",
|
||||
},
|
||||
},
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, stepResults)
|
||||
|
||||
// Test resolving nested field
|
||||
value, err := resolver.Resolve("${Clone.output.path}")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve: %v", err)
|
||||
}
|
||||
|
||||
if value != "/tmp/repo" {
|
||||
t.Errorf("Expected '/tmp/repo', got %v", value)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveDeepNesting(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
stepResults := map[string]interface{}{
|
||||
"Analyze": map[string]interface{}{
|
||||
"output": map[string]interface{}{
|
||||
"metrics": map[string]interface{}{
|
||||
"quality": map[string]interface{}{
|
||||
"score": 0.95,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, stepResults)
|
||||
|
||||
value, err := resolver.Resolve("${Analyze.output.metrics.quality.score}")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve: %v", err)
|
||||
}
|
||||
|
||||
score, ok := value.(float64)
|
||||
if !ok {
|
||||
t.Fatalf("Expected float64, got %T", value)
|
||||
}
|
||||
|
||||
if score != 0.95 {
|
||||
t.Errorf("Expected 0.95, got %v", score)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveNonTemplate(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
resolver := NewJSONPathResolver(input, map[string]interface{}{})
|
||||
|
||||
// Non-template strings should be returned as-is
|
||||
value, err := resolver.Resolve("plain string")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve: %v", err)
|
||||
}
|
||||
|
||||
if value != "plain string" {
|
||||
t.Errorf("Expected 'plain string', got %v", value)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveMissingStep(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
resolver := NewJSONPathResolver(input, map[string]interface{}{})
|
||||
|
||||
// Should error on missing step
|
||||
_, err := resolver.Resolve("${NonExistentStep.output.field}")
|
||||
if err == nil {
|
||||
t.Error("Expected error for missing step")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveMissingField(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
stepResults := map[string]interface{}{
|
||||
"Clone": map[string]interface{}{
|
||||
"output": map[string]interface{}{
|
||||
"path": "/tmp/repo",
|
||||
},
|
||||
},
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, stepResults)
|
||||
|
||||
// Should error on missing field
|
||||
_, err := resolver.Resolve("${Clone.output.nonexistent}")
|
||||
if err == nil {
|
||||
t.Error("Expected error for missing field")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveString(t *testing.T) {
|
||||
input := map[string]interface{}{
|
||||
"repo": "https://github.com/test/repo",
|
||||
}
|
||||
stepResults := map[string]interface{}{
|
||||
"Clone": map[string]interface{}{
|
||||
"output": map[string]interface{}{
|
||||
"path": "/tmp/repo",
|
||||
},
|
||||
},
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, stepResults)
|
||||
|
||||
// Resolve string with multiple expressions
|
||||
result, err := resolver.ResolveString("Repository at ${input.repo} cloned to ${Clone.output.path}")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve string: %v", err)
|
||||
}
|
||||
|
||||
expected := "Repository at https://github.com/test/repo cloned to /tmp/repo"
|
||||
if result != expected {
|
||||
t.Errorf("Expected '%s', got '%s'", expected, result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveStringNoExpressions(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
resolver := NewJSONPathResolver(input, map[string]interface{}{})
|
||||
|
||||
// String without expressions should be returned unchanged
|
||||
result, err := resolver.ResolveString("plain string")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve string: %v", err)
|
||||
}
|
||||
|
||||
if result != "plain string" {
|
||||
t.Errorf("Expected 'plain string', got '%s'", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvePaths(t *testing.T) {
|
||||
input := map[string]interface{}{
|
||||
"repo": "https://github.com/test/repo",
|
||||
}
|
||||
stepResults := map[string]interface{}{
|
||||
"Clone": map[string]interface{}{
|
||||
"output": map[string]interface{}{
|
||||
"path": "/tmp/repo",
|
||||
},
|
||||
},
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, stepResults)
|
||||
|
||||
// Resolve a map with JSONPath values
|
||||
data := map[string]interface{}{
|
||||
"repository": "${input.repo}",
|
||||
"path": "${Clone.output.path}",
|
||||
"literal": "just a string",
|
||||
}
|
||||
|
||||
result, err := resolver.ResolvePaths(data)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve paths: %v", err)
|
||||
}
|
||||
|
||||
if result["repository"] != "https://github.com/test/repo" {
|
||||
t.Errorf("repository mismatch: %v", result["repository"])
|
||||
}
|
||||
|
||||
if result["path"] != "/tmp/repo" {
|
||||
t.Errorf("path mismatch: %v", result["path"])
|
||||
}
|
||||
|
||||
if result["literal"] != "just a string" {
|
||||
t.Errorf("literal mismatch: %v", result["literal"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveNestedMap(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
stepResults := map[string]interface{}{
|
||||
"Analyze": map[string]interface{}{
|
||||
"output": map[string]interface{}{
|
||||
"score": 0.95,
|
||||
},
|
||||
},
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, stepResults)
|
||||
|
||||
// Resolve nested map
|
||||
data := map[string]interface{}{
|
||||
"analysis": map[string]interface{}{
|
||||
"quality": "${Analyze.output.score}",
|
||||
},
|
||||
}
|
||||
|
||||
result, err := resolver.ResolvePaths(data)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve nested map: %v", err)
|
||||
}
|
||||
|
||||
analysisMap := result["analysis"].(map[string]interface{})
|
||||
if analysisMap["quality"] != 0.95 {
|
||||
t.Errorf("Expected 0.95, got %v", analysisMap["quality"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveSlice(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
stepResults := map[string]interface{}{
|
||||
"Scan": map[string]interface{}{
|
||||
"output": map[string]interface{}{
|
||||
"vulnerabilities": []map[string]interface{}{
|
||||
{"cve": "CVE-001"},
|
||||
{"cve": "CVE-002"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, stepResults)
|
||||
|
||||
// Resolve slice
|
||||
data := map[string]interface{}{
|
||||
"issues": "${Scan.output.vulnerabilities}",
|
||||
}
|
||||
|
||||
result, err := resolver.ResolvePaths(data)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve slice: %v", err)
|
||||
}
|
||||
|
||||
issues := result["issues"].([]map[string]interface{})
|
||||
if len(issues) != 2 {
|
||||
t.Errorf("Expected 2 issues, got %d", len(issues))
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidatePath(t *testing.T) {
|
||||
resolver := NewJSONPathResolver(map[string]interface{}{}, map[string]interface{}{})
|
||||
|
||||
// Valid paths
|
||||
validPaths := []string{
|
||||
"input.repo",
|
||||
"Clone.output.path",
|
||||
"Analyze.output.metrics.quality.score",
|
||||
}
|
||||
|
||||
for _, path := range validPaths {
|
||||
if err := resolver.ValidatePath(path); err != nil {
|
||||
t.Errorf("Path '%s' should be valid: %v", path, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Invalid paths
|
||||
invalidPaths := []string{
|
||||
"",
|
||||
"singleword",
|
||||
}
|
||||
|
||||
for _, path := range invalidPaths {
|
||||
if err := resolver.ValidatePath(path); err == nil {
|
||||
t.Errorf("Path '%s' should be invalid", path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetAvailableSteps(t *testing.T) {
|
||||
stepResults := map[string]interface{}{
|
||||
"Clone": map[string]interface{}{},
|
||||
"Analyze": map[string]interface{}{},
|
||||
"Scan": map[string]interface{}{},
|
||||
}
|
||||
resolver := NewJSONPathResolver(map[string]interface{}{}, stepResults)
|
||||
|
||||
steps := resolver.GetAvailableSteps()
|
||||
if len(steps) != 3 {
|
||||
t.Errorf("Expected 3 steps, got %d", len(steps))
|
||||
}
|
||||
|
||||
// Check all steps are present
|
||||
stepMap := make(map[string]bool)
|
||||
for _, step := range steps {
|
||||
stepMap[step] = true
|
||||
}
|
||||
|
||||
if !stepMap["Clone"] || !stepMap["Analyze"] || !stepMap["Scan"] {
|
||||
t.Error("Missing expected steps")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetInputFields(t *testing.T) {
|
||||
input := map[string]interface{}{
|
||||
"repo": "test",
|
||||
"branch": "main",
|
||||
"path": "/tmp",
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, map[string]interface{}{})
|
||||
|
||||
fields := resolver.GetInputFields()
|
||||
if len(fields) != 3 {
|
||||
t.Errorf("Expected 3 fields, got %d", len(fields))
|
||||
}
|
||||
|
||||
// Check all fields are present
|
||||
fieldMap := make(map[string]bool)
|
||||
for _, field := range fields {
|
||||
fieldMap[field] = true
|
||||
}
|
||||
|
||||
if !fieldMap["repo"] || !fieldMap["branch"] || !fieldMap["path"] {
|
||||
t.Error("Missing expected input fields")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveWithStringMap(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
stepResults := map[string]interface{}{
|
||||
"Config": map[string]string{
|
||||
"url": "https://example.com",
|
||||
"port": "8080",
|
||||
},
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, stepResults)
|
||||
|
||||
// Resolve from string map
|
||||
value, err := resolver.Resolve("${Config.url}")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve: %v", err)
|
||||
}
|
||||
|
||||
if value != "https://example.com" {
|
||||
t.Errorf("Expected 'https://example.com', got %v", value)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveComplexWorkflow(t *testing.T) {
|
||||
input := map[string]interface{}{
|
||||
"repo": "https://github.com/test/repo",
|
||||
"branch": "feature/new",
|
||||
}
|
||||
stepResults := map[string]interface{}{
|
||||
"Clone": map[string]interface{}{
|
||||
"output": map[string]interface{}{
|
||||
"path": "/tmp/repo",
|
||||
"commit": "abc123def456",
|
||||
},
|
||||
},
|
||||
"Analyze": map[string]interface{}{
|
||||
"output": map[string]interface{}{
|
||||
"quality": 0.92,
|
||||
"issues": []string{"issue1", "issue2"},
|
||||
},
|
||||
},
|
||||
}
|
||||
resolver := NewJSONPathResolver(input, stepResults)
|
||||
|
||||
// Complex workflow parameters
|
||||
params := map[string]interface{}{
|
||||
"source_repo": "${input.repo}",
|
||||
"target_branch": "${input.branch}",
|
||||
"cloned_path": "${Clone.output.path}",
|
||||
"commit_hash": "${Clone.output.commit}",
|
||||
"quality_score": "${Analyze.output.quality}",
|
||||
"issues_found": "${Analyze.output.issues}",
|
||||
"report": "Quality score is ${Analyze.output.quality} for commit ${Clone.output.commit}",
|
||||
}
|
||||
|
||||
resolved, err := resolver.ResolvePaths(params)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve workflow: %v", err)
|
||||
}
|
||||
|
||||
if resolved["source_repo"] != "https://github.com/test/repo" {
|
||||
t.Error("source_repo mismatch")
|
||||
}
|
||||
|
||||
if resolved["target_branch"] != "feature/new" {
|
||||
t.Error("target_branch mismatch")
|
||||
}
|
||||
|
||||
if resolved["cloned_path"] != "/tmp/repo" {
|
||||
t.Error("cloned_path mismatch")
|
||||
}
|
||||
|
||||
if resolved["commit_hash"] != "abc123def456" {
|
||||
t.Error("commit_hash mismatch")
|
||||
}
|
||||
|
||||
if resolved["quality_score"] != 0.92 {
|
||||
t.Error("quality_score mismatch")
|
||||
}
|
||||
|
||||
// Check report string resolution
|
||||
report := resolved["report"].(string)
|
||||
if !strings.Contains(report, "0.92") || !strings.Contains(report, "abc123def456") {
|
||||
t.Errorf("Report not properly resolved: %s", report)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveEmptyInput(t *testing.T) {
|
||||
input := map[string]interface{}{}
|
||||
resolver := NewJSONPathResolver(input, map[string]interface{}{})
|
||||
|
||||
// Should resolve to just input when accessing input
|
||||
value, err := resolver.Resolve("${input}")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve: %v", err)
|
||||
}
|
||||
|
||||
// Should be empty map
|
||||
inputMap, ok := value.(map[string]interface{})
|
||||
if !ok || len(inputMap) != 0 {
|
||||
t.Error("Expected empty input map")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,275 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io/ioutil"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
)
|
||||
|
||||
// KnowledgeBase represents the activity knowledge base
|
||||
type KnowledgeBase struct {
|
||||
Version string `json:"version"`
|
||||
Activities []ActivityMetadata `json:"activities"`
|
||||
Metadata KnowledgeBaseMetadata `json:"metadata"`
|
||||
|
||||
// Index for fast lookups
|
||||
byName map[string]*ActivityMetadata
|
||||
}
|
||||
|
||||
// KnowledgeBaseMetadata tracks KB metadata
|
||||
type KnowledgeBaseMetadata struct {
|
||||
TotalActivities int `json:"totalActivities"`
|
||||
LastUpdated string `json:"lastUpdated"`
|
||||
Categories map[string]int `json:"categories"`
|
||||
}
|
||||
|
||||
// LoadKnowledgeBase loads the activity knowledge base from a JSON file
|
||||
func LoadKnowledgeBase(filePath string) (*KnowledgeBase, error) {
|
||||
// Read file
|
||||
data, err := ioutil.ReadFile(filePath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read knowledge base file: %w", err)
|
||||
}
|
||||
|
||||
// Parse JSON
|
||||
var kb KnowledgeBase
|
||||
err = json.Unmarshal(data, &kb)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse knowledge base JSON: %w", err)
|
||||
}
|
||||
|
||||
// Build index
|
||||
kb.byName = make(map[string]*ActivityMetadata)
|
||||
for i := range kb.Activities {
|
||||
kb.byName[kb.Activities[i].Name] = &kb.Activities[i]
|
||||
}
|
||||
|
||||
return &kb, nil
|
||||
}
|
||||
|
||||
// LoadKnowledgeBaseFromDefaultPath loads KB from default location
|
||||
// Looks for activity_knowledge_base.json in same directory as caller
|
||||
func LoadKnowledgeBaseFromDefaultPath() (*KnowledgeBase, error) {
|
||||
// Try to find from package directory
|
||||
execDir, err := os.Executable()
|
||||
if err == nil {
|
||||
// Try in same directory as binary
|
||||
path := filepath.Join(filepath.Dir(execDir), "activity_knowledge_base.json")
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
return LoadKnowledgeBase(path)
|
||||
}
|
||||
}
|
||||
|
||||
// Try from current working directory
|
||||
if _, err := os.Stat("activity_knowledge_base.json"); err == nil {
|
||||
return LoadKnowledgeBase("activity_knowledge_base.json")
|
||||
}
|
||||
|
||||
// Try from internal/routing directory relative to cwd
|
||||
if _, err := os.Stat("internal/routing/activity_knowledge_base.json"); err == nil {
|
||||
return LoadKnowledgeBase("internal/routing/activity_knowledge_base.json")
|
||||
}
|
||||
|
||||
// Try from parent directory (for tests running from tests/ dir)
|
||||
if _, err := os.Stat("../internal/routing/activity_knowledge_base.json"); err == nil {
|
||||
return LoadKnowledgeBase("../internal/routing/activity_knowledge_base.json")
|
||||
}
|
||||
|
||||
// Try using runtime to find package directory
|
||||
_, filename, _, ok := runtime.Caller(0)
|
||||
if ok {
|
||||
pkgDir := filepath.Dir(filename)
|
||||
path := filepath.Join(pkgDir, "activity_knowledge_base.json")
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
return LoadKnowledgeBase(path)
|
||||
}
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("activity_knowledge_base.json not found in any expected location")
|
||||
}
|
||||
|
||||
// GetActivity returns metadata for a specific activity
|
||||
func (kb *KnowledgeBase) GetActivity(name string) *ActivityMetadata {
|
||||
return kb.byName[name]
|
||||
}
|
||||
|
||||
// ListActivities returns all activities
|
||||
func (kb *KnowledgeBase) ListActivities() []ActivityMetadata {
|
||||
return kb.Activities
|
||||
}
|
||||
|
||||
// ListActivitiesByCategory returns all activities in a category
|
||||
func (kb *KnowledgeBase) ListActivitiesByCategory(category string) []ActivityMetadata {
|
||||
var result []ActivityMetadata
|
||||
for _, activity := range kb.Activities {
|
||||
if activity.Category == category {
|
||||
result = append(result, activity)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// GetActivityNames returns all activity names
|
||||
func (kb *KnowledgeBase) GetActivityNames() []string {
|
||||
names := make([]string, len(kb.Activities))
|
||||
for i, activity := range kb.Activities {
|
||||
names[i] = activity.Name
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// HasActivity checks if an activity exists
|
||||
func (kb *KnowledgeBase) HasActivity(name string) bool {
|
||||
_, exists := kb.byName[name]
|
||||
return exists
|
||||
}
|
||||
|
||||
// GetDependencies returns all dependencies for an activity
|
||||
func (kb *KnowledgeBase) GetDependencies(activityName string) []string {
|
||||
activity := kb.GetActivity(activityName)
|
||||
if activity == nil {
|
||||
return []string{}
|
||||
}
|
||||
return activity.Constraints.Dependencies
|
||||
}
|
||||
|
||||
// GetTimeoutForActivity returns the timeout for an activity
|
||||
func (kb *KnowledgeBase) GetTimeoutForActivity(activityName string) string {
|
||||
activity := kb.GetActivity(activityName)
|
||||
if activity == nil {
|
||||
return "5m" // Default timeout
|
||||
}
|
||||
return activity.Constraints.DefaultTimeout
|
||||
}
|
||||
|
||||
// GetRetryPolicyForActivity returns retry configuration for an activity
|
||||
func (kb *KnowledgeBase) GetRetryPolicyForActivity(activityName string) *RetryPolicy {
|
||||
activity := kb.GetActivity(activityName)
|
||||
if activity == nil {
|
||||
return &RetryPolicy{
|
||||
MaxAttempts: 1,
|
||||
BackoffRate: 1.0,
|
||||
InitialInterval: "1s",
|
||||
}
|
||||
}
|
||||
|
||||
return &RetryPolicy{
|
||||
MaxAttempts: int32(activity.Constraints.RecommendedRetries),
|
||||
BackoffRate: activity.Constraints.RetryBackoff,
|
||||
InitialInterval: "1s",
|
||||
MaxInterval: "30s",
|
||||
}
|
||||
}
|
||||
|
||||
// IsFlaky returns whether an activity is marked as flaky
|
||||
func (kb *KnowledgeBase) IsFlaky(activityName string) bool {
|
||||
activity := kb.GetActivity(activityName)
|
||||
if activity == nil {
|
||||
return false
|
||||
}
|
||||
return activity.Constraints.IsFlaky
|
||||
}
|
||||
|
||||
// GetNotes returns implementation notes for an activity
|
||||
func (kb *KnowledgeBase) GetNotes(activityName string) string {
|
||||
activity := kb.GetActivity(activityName)
|
||||
if activity == nil {
|
||||
return ""
|
||||
}
|
||||
return activity.Constraints.Notes
|
||||
}
|
||||
|
||||
// Validate checks the knowledge base for consistency
|
||||
func (kb *KnowledgeBase) Validate() error {
|
||||
// Check for circular dependencies
|
||||
visited := make(map[string]bool)
|
||||
for _, activity := range kb.Activities {
|
||||
if err := kb.checkDependencies(activity.Name, visited, []string{}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Check that all dependencies exist
|
||||
for _, activity := range kb.Activities {
|
||||
for _, dep := range activity.Constraints.Dependencies {
|
||||
if !kb.HasActivity(dep) {
|
||||
return fmt.Errorf("activity %s depends on non-existent activity %s", activity.Name, dep)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// checkDependencies validates activity dependencies for cycles
|
||||
func (kb *KnowledgeBase) checkDependencies(activityName string, visited map[string]bool, path []string) error {
|
||||
// Check for cycles
|
||||
for _, p := range path {
|
||||
if p == activityName {
|
||||
cycleStr := ""
|
||||
found := false
|
||||
for _, n := range path {
|
||||
if found {
|
||||
cycleStr += " -> " + n
|
||||
}
|
||||
if n == activityName {
|
||||
found = true
|
||||
cycleStr += n
|
||||
}
|
||||
}
|
||||
cycleStr += " -> " + activityName
|
||||
return fmt.Errorf("circular dependency detected: %s", cycleStr)
|
||||
}
|
||||
}
|
||||
|
||||
if visited[activityName] {
|
||||
return nil // Already checked this branch
|
||||
}
|
||||
|
||||
visited[activityName] = true
|
||||
newPath := append(path, activityName)
|
||||
|
||||
activity := kb.GetActivity(activityName)
|
||||
if activity == nil {
|
||||
return nil // Non-existent activity will be caught elsewhere
|
||||
}
|
||||
|
||||
for _, dep := range activity.Constraints.Dependencies {
|
||||
if err := kb.checkDependencies(dep, visited, newPath); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns a human-readable description of the knowledge base
|
||||
func (kb *KnowledgeBase) String() string {
|
||||
return fmt.Sprintf("KnowledgeBase(v%s, %d activities)", kb.Version, kb.Metadata.TotalActivities)
|
||||
}
|
||||
|
||||
// PrintSummary prints a summary of available activities
|
||||
func (kb *KnowledgeBase) PrintSummary() string {
|
||||
summary := fmt.Sprintf("=== Activity Knowledge Base ===\nVersion: %s\nTotal Activities: %d\n\n", kb.Version, kb.Metadata.TotalActivities)
|
||||
|
||||
summary += "Activities by Category:\n"
|
||||
for category, count := range kb.Metadata.Categories {
|
||||
summary += fmt.Sprintf(" %s: %d\n", category, count)
|
||||
}
|
||||
|
||||
summary += "\nActivity Details:\n"
|
||||
for _, activity := range kb.Activities {
|
||||
summary += fmt.Sprintf("\n[%s] %s\n", activity.Name, activity.Description)
|
||||
summary += fmt.Sprintf(" Category: %s\n", activity.Category)
|
||||
summary += fmt.Sprintf(" Timeout: %s\n", activity.Constraints.DefaultTimeout)
|
||||
summary += fmt.Sprintf(" Flaky: %v (Retries: %d)\n", activity.Constraints.IsFlaky, activity.Constraints.RecommendedRetries)
|
||||
if len(activity.Constraints.Dependencies) > 0 {
|
||||
summary += fmt.Sprintf(" Dependencies: %v\n", activity.Constraints.Dependencies)
|
||||
}
|
||||
}
|
||||
|
||||
return summary
|
||||
}
|
||||
@@ -0,0 +1,371 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func getKBPath() string {
|
||||
// Try direct name (when running from this directory)
|
||||
if _, err := os.Stat("activity_knowledge_base.json"); err == nil {
|
||||
return "activity_knowledge_base.json"
|
||||
}
|
||||
// Try relative path
|
||||
if _, err := os.Stat("./internal/routing/activity_knowledge_base.json"); err == nil {
|
||||
return "./internal/routing/activity_knowledge_base.json"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func TestLoadKnowledgeBase(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
if kb == nil {
|
||||
t.Error("Knowledge base is nil")
|
||||
}
|
||||
|
||||
if kb.Version == "" {
|
||||
t.Error("Knowledge base version is empty")
|
||||
}
|
||||
|
||||
if len(kb.Activities) == 0 {
|
||||
t.Error("Knowledge base has no activities")
|
||||
}
|
||||
|
||||
if len(kb.byName) != len(kb.Activities) {
|
||||
t.Errorf("Index size (%d) doesn't match activities (%d)", len(kb.byName), len(kb.Activities))
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetActivity(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
activity := kb.GetActivity("CloneRepoActivity")
|
||||
if activity == nil {
|
||||
t.Error("CloneRepoActivity not found")
|
||||
} else {
|
||||
if activity.Name != "CloneRepoActivity" {
|
||||
t.Errorf("Activity name mismatch: %s", activity.Name)
|
||||
}
|
||||
if activity.Description == "" {
|
||||
t.Error("Activity description is empty")
|
||||
}
|
||||
}
|
||||
|
||||
missing := kb.GetActivity("NonExistentActivity")
|
||||
if missing != nil {
|
||||
t.Error("NonExistentActivity should be nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHasActivity(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
if !kb.HasActivity("CloneRepoActivity") {
|
||||
t.Error("CloneRepoActivity should exist")
|
||||
}
|
||||
|
||||
if kb.HasActivity("NonExistentActivity") {
|
||||
t.Error("NonExistentActivity should not exist")
|
||||
}
|
||||
}
|
||||
|
||||
func TestListActivities(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
activities := kb.ListActivities()
|
||||
if len(activities) == 0 {
|
||||
t.Error("ListActivities returned empty list")
|
||||
}
|
||||
|
||||
for _, activity := range activities {
|
||||
if activity.Name == "" {
|
||||
t.Error("Activity name is empty")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestListActivitiesByCategory(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
repoActivities := kb.ListActivitiesByCategory("repository")
|
||||
if len(repoActivities) == 0 {
|
||||
t.Error("No repository activities found")
|
||||
}
|
||||
|
||||
for _, activity := range repoActivities {
|
||||
if activity.Category != "repository" {
|
||||
t.Errorf("Activity %s has wrong category: %s", activity.Name, activity.Category)
|
||||
}
|
||||
}
|
||||
|
||||
unknown := kb.ListActivitiesByCategory("unknown")
|
||||
if len(unknown) != 0 {
|
||||
t.Error("Unknown category should return empty list")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetActivityNames(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
names := kb.GetActivityNames()
|
||||
if len(names) == 0 {
|
||||
t.Error("GetActivityNames returned empty list")
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, name := range names {
|
||||
if name == "CloneRepoActivity" {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Error("CloneRepoActivity not found in activity names")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDependencies(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
deps := kb.GetDependencies("AnalyzeCodeActivity")
|
||||
if len(deps) == 0 {
|
||||
t.Error("AnalyzeCodeActivity should have dependencies")
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, dep := range deps {
|
||||
if dep == "CloneRepoActivity" {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Error("CloneRepoActivity should be a dependency of AnalyzeCodeActivity")
|
||||
}
|
||||
|
||||
noDeps := kb.GetDependencies("CloneRepoActivity")
|
||||
if len(noDeps) != 0 {
|
||||
t.Error("CloneRepoActivity should have no dependencies")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetTimeout(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
timeout := kb.GetTimeoutForActivity("CloneRepoActivity")
|
||||
if timeout == "" {
|
||||
t.Error("Timeout should not be empty")
|
||||
}
|
||||
|
||||
defaultTimeout := kb.GetTimeoutForActivity("NonExistent")
|
||||
if defaultTimeout != "5m" {
|
||||
t.Errorf("Default timeout should be 5m, got %s", defaultTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetRetryPolicy(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
policy := kb.GetRetryPolicyForActivity("AnalyzeCodeActivity")
|
||||
if policy == nil {
|
||||
t.Error("Retry policy should not be nil")
|
||||
} else {
|
||||
if policy.MaxAttempts < 2 {
|
||||
t.Errorf("Flaky activity should have multiple retries, got %d", policy.MaxAttempts)
|
||||
}
|
||||
if policy.BackoffRate == 0 {
|
||||
t.Error("Backoff rate should be set")
|
||||
}
|
||||
}
|
||||
|
||||
stablePolicy := kb.GetRetryPolicyForActivity("CloneRepoActivity")
|
||||
if stablePolicy == nil {
|
||||
t.Error("Retry policy should not be nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsFlaky(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
if !kb.IsFlaky("AnalyzeCodeActivity") {
|
||||
t.Error("AnalyzeCodeActivity should be marked as flaky")
|
||||
}
|
||||
|
||||
if kb.IsFlaky("CloneRepoActivity") {
|
||||
t.Error("CloneRepoActivity should not be marked as flaky")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetNotes(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
notes := kb.GetNotes("AnalyzeCodeActivity")
|
||||
if notes == "" {
|
||||
t.Error("Notes should not be empty")
|
||||
}
|
||||
|
||||
missingNotes := kb.GetNotes("NonExistent")
|
||||
if missingNotes != "" {
|
||||
t.Error("Non-existent activity should have empty notes")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidate(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
if err := kb.Validate(); err != nil {
|
||||
t.Fatalf("Knowledge base validation failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestString(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
str := kb.String()
|
||||
if str == "" {
|
||||
t.Error("String() returned empty string")
|
||||
}
|
||||
|
||||
if !contains(str, "KnowledgeBase") {
|
||||
t.Error("String should contain 'KnowledgeBase'")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintSummary(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base file not found, skipping test")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
summary := kb.PrintSummary()
|
||||
if summary == "" {
|
||||
t.Error("PrintSummary() returned empty string")
|
||||
}
|
||||
|
||||
if !contains(summary, "Activity Knowledge Base") {
|
||||
t.Error("Summary should contain 'Activity Knowledge Base'")
|
||||
}
|
||||
|
||||
if !contains(summary, "CloneRepoActivity") {
|
||||
t.Error("Summary should list activities")
|
||||
}
|
||||
}
|
||||
|
||||
func contains(str, substr string) bool {
|
||||
for i := 0; i < len(str)-len(substr)+1; i++ {
|
||||
if str[i:i+len(substr)] == substr {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var (
|
||||
// llmBaseURL is the base URL for the LLM API
|
||||
llmBaseURL string
|
||||
)
|
||||
|
||||
func init() {
|
||||
llmBaseURL = os.Getenv("LOCAL_LLM_BASE_URL")
|
||||
if llmBaseURL == "" {
|
||||
llmBaseURL = "https://api.riotpiao.com"
|
||||
}
|
||||
}
|
||||
|
||||
// AuthType specifies the authentication mechanism
|
||||
type AuthType string
|
||||
|
||||
const (
|
||||
// AuthTypeNone - no authentication
|
||||
AuthTypeNone AuthType = "none"
|
||||
// AuthTypeBearer - Bearer token (JWT, OAuth2)
|
||||
AuthTypeBearer AuthType = "bearer"
|
||||
// AuthTypeAPIKey - API Key authentication
|
||||
AuthTypeAPIKey AuthType = "api-key"
|
||||
// AuthTypeCustom - Custom header-based authentication
|
||||
AuthTypeCustom AuthType = "custom"
|
||||
)
|
||||
|
||||
// LLMAuth configures authentication for LLM API
|
||||
type LLMAuth struct {
|
||||
// Type of authentication
|
||||
Type AuthType `json:"type"`
|
||||
|
||||
// Token is the JWT/OAuth2 token for Bearer auth
|
||||
Token string `json:"token,omitempty"`
|
||||
|
||||
// APIKey is the API key for API Key auth
|
||||
APIKey string `json:"apiKey,omitempty"`
|
||||
|
||||
// HeaderName is the custom header name for Custom auth
|
||||
HeaderName string `json:"headerName,omitempty"`
|
||||
|
||||
// HeaderValue is the custom header value for Custom auth
|
||||
HeaderValue string `json:"headerValue,omitempty"`
|
||||
}
|
||||
|
||||
// LLMClient is a simple LLM client for routing
|
||||
type LLMClient struct {
|
||||
baseURL string
|
||||
httpClient *http.Client
|
||||
auth *LLMAuth
|
||||
}
|
||||
|
||||
// NewLLMClient creates a new LLM client with default (no) auth
|
||||
func NewLLMClient() *LLMClient {
|
||||
return &LLMClient{
|
||||
baseURL: llmBaseURL,
|
||||
httpClient: &http.Client{},
|
||||
auth: nil,
|
||||
}
|
||||
}
|
||||
|
||||
// NewLLMClientWithAuth creates a new LLM client with authentication
|
||||
func NewLLMClientWithAuth(auth *LLMAuth) *LLMClient {
|
||||
return &LLMClient{
|
||||
baseURL: llmBaseURL,
|
||||
httpClient: &http.Client{},
|
||||
auth: auth,
|
||||
}
|
||||
}
|
||||
|
||||
// Name returns the provider name
|
||||
func (c *LLMClient) Name() string {
|
||||
return "riotpiao"
|
||||
}
|
||||
|
||||
// IsAvailable checks if the LLM service is available
|
||||
func (c *LLMClient) IsAvailable(ctx context.Context) error {
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", c.baseURL, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := c.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("LLM service unavailable: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode >= 500 {
|
||||
return fmt.Errorf("LLM service error: %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// llmRequest is the request body for the OpenAI-compatible API
|
||||
type llmRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []llmMessage `json:"messages"`
|
||||
Stream bool `json:"stream"`
|
||||
}
|
||||
|
||||
type llmMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// llmResponse is the response from the OpenAI-compatible API
|
||||
type llmResponse struct {
|
||||
Choices []struct {
|
||||
Message struct {
|
||||
Content string `json:"content"`
|
||||
} `json:"message"`
|
||||
} `json:"choices"`
|
||||
}
|
||||
|
||||
// Chat sends a chat completion request
|
||||
func (c *LLMClient) Chat(ctx context.Context, systemPrompt, userMessage string) (string, error) {
|
||||
req := llmRequest{
|
||||
Model: "reasoning",
|
||||
Messages: []llmMessage{
|
||||
{Role: "system", Content: systemPrompt},
|
||||
{Role: "user", Content: userMessage},
|
||||
},
|
||||
Stream: false,
|
||||
}
|
||||
|
||||
reqBody, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, "POST",
|
||||
fmt.Sprintf("%s/v1/chat/completions", c.baseURL),
|
||||
bytes.NewReader(reqBody))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to create HTTP request: %w", err)
|
||||
}
|
||||
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
// Apply authentication headers
|
||||
if err := c.applyAuth(httpReq); err != nil {
|
||||
return "", fmt.Errorf("failed to apply authentication: %w", err)
|
||||
}
|
||||
|
||||
resp, err := c.httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to connect to LLM API at %s: %w", c.baseURL, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to read response body: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("LLM API returned status %d: %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
var respObj llmResponse
|
||||
if err := json.Unmarshal(respBody, &respObj); err != nil {
|
||||
return "", fmt.Errorf("failed to unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
if len(respObj.Choices) == 0 {
|
||||
return "", fmt.Errorf("no choices in response from LLM API")
|
||||
}
|
||||
|
||||
return respObj.Choices[0].Message.Content, nil
|
||||
}
|
||||
|
||||
// applyAuth applies authentication to the HTTP request based on config
|
||||
func (c *LLMClient) applyAuth(req *http.Request) error {
|
||||
if c.auth == nil || c.auth.Type == AuthTypeNone {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch c.auth.Type {
|
||||
case AuthTypeBearer:
|
||||
if c.auth.Token == "" {
|
||||
return fmt.Errorf("bearer token is required but not provided")
|
||||
}
|
||||
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", c.auth.Token))
|
||||
|
||||
case AuthTypeAPIKey:
|
||||
if c.auth.APIKey == "" {
|
||||
return fmt.Errorf("API key is required but not provided")
|
||||
}
|
||||
// Common API key header names: X-API-Key, api-key, Authorization
|
||||
req.Header.Set("X-API-Key", c.auth.APIKey)
|
||||
|
||||
case AuthTypeCustom:
|
||||
if c.auth.HeaderName == "" || c.auth.HeaderValue == "" {
|
||||
return fmt.Errorf("custom header name and value are required but not provided")
|
||||
}
|
||||
req.Header.Set(c.auth.HeaderName, c.auth.HeaderValue)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateAuth updates the authentication config at runtime
|
||||
func (c *LLMClient) UpdateAuth(auth *LLMAuth) error {
|
||||
if auth == nil {
|
||||
return fmt.Errorf("auth config cannot be nil")
|
||||
}
|
||||
c.auth = auth
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetAuth returns the current authentication config
|
||||
func (c *LLMClient) GetAuth() *LLMAuth {
|
||||
return c.auth
|
||||
}
|
||||
@@ -0,0 +1,552 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// LLMRouterInput is input to the llm-router activity
|
||||
type LLMRouterInput struct {
|
||||
Message string `json:"message"`
|
||||
Context map[string]interface{} `json:"context,omitempty"` // Optional context (repo, branch, etc)
|
||||
MemoryContext *MemoryContext `json:"memoryContext,omitempty"` // Optional memory retrieval results
|
||||
UseMemory bool `json:"useMemory,omitempty"` // Enable memory retrieval (default: false)
|
||||
}
|
||||
|
||||
// MemoryContext holds retrieved memory for prompt injection
|
||||
type MemoryContext struct {
|
||||
Skills []MemorySkill `json:"skills"`
|
||||
Lessons []MemoryLesson `json:"lessons"`
|
||||
References []MemoryReference `json:"references"`
|
||||
}
|
||||
|
||||
// MemorySkill from memory service
|
||||
type MemorySkill struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Why string `json:"why,omitempty"`
|
||||
}
|
||||
|
||||
// MemoryLesson from memory service
|
||||
type MemoryLesson struct {
|
||||
ID string `json:"id"`
|
||||
Text string `json:"text"`
|
||||
Level string `json:"level"`
|
||||
}
|
||||
|
||||
// MemoryReference from memory service
|
||||
type MemoryReference struct {
|
||||
ID string `json:"id"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
// LLMRouterOutput is output from the llm-router activity
|
||||
type LLMRouterOutput struct {
|
||||
Spec *WorkflowSpec `json:"spec,omitempty"`
|
||||
CronSpec *CronWorkflowSpec `json:"cronSpec,omitempty"`
|
||||
IsCron bool `json:"isCron"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// LLMRouter orchestrates intent analysis and spec generation
|
||||
type LLMRouter struct {
|
||||
provider LLMProvider
|
||||
knowledgeBase *KnowledgeBase
|
||||
specBuilder SpecBuilder
|
||||
validators []WorkflowValidator
|
||||
paramBinder ParameterBinder
|
||||
promptTemplate PromptTemplate
|
||||
}
|
||||
|
||||
// LLMRouterConfig configures the router
|
||||
type LLMRouterConfig struct {
|
||||
Provider LLMProvider
|
||||
KnowledgeBase *KnowledgeBase
|
||||
SpecBuilder SpecBuilder
|
||||
Validators []WorkflowValidator
|
||||
ParamBinder ParameterBinder
|
||||
Auth *LLMAuth // Authentication config for LLM API
|
||||
}
|
||||
|
||||
// NewLLMRouter creates a new LLM router with custom config
|
||||
func NewLLMRouter(config LLMRouterConfig) (*LLMRouter, error) {
|
||||
if config.Provider == nil {
|
||||
return nil, fmt.Errorf("provider is required")
|
||||
}
|
||||
if config.KnowledgeBase == nil {
|
||||
return nil, fmt.Errorf("knowledge base is required")
|
||||
}
|
||||
|
||||
router := &LLMRouter{
|
||||
provider: config.Provider,
|
||||
knowledgeBase: config.KnowledgeBase,
|
||||
}
|
||||
|
||||
// Set defaults
|
||||
if config.SpecBuilder == nil {
|
||||
router.specBuilder = NewDefaultSpecBuilder(config.KnowledgeBase)
|
||||
} else {
|
||||
router.specBuilder = config.SpecBuilder
|
||||
}
|
||||
|
||||
if config.ParamBinder == nil {
|
||||
router.paramBinder = NewDefaultParameterBinder()
|
||||
} else {
|
||||
router.paramBinder = config.ParamBinder
|
||||
}
|
||||
|
||||
router.validators = config.Validators
|
||||
if len(router.validators) == 0 {
|
||||
router.validators = []WorkflowValidator{
|
||||
&StateGraphValidator{},
|
||||
NewActivityAvailabilityValidator(config.KnowledgeBase),
|
||||
&TimeoutValidator{},
|
||||
}
|
||||
}
|
||||
|
||||
return router, nil
|
||||
}
|
||||
|
||||
// NewLLMRouterDefault creates router with default HTTP provider
|
||||
func NewLLMRouterDefault(kb *KnowledgeBase) (*LLMRouter, error) {
|
||||
client := NewLLMClient()
|
||||
return NewLLMRouter(LLMRouterConfig{
|
||||
Provider: client,
|
||||
KnowledgeBase: kb,
|
||||
})
|
||||
}
|
||||
|
||||
// Route analyzes user message and generates appropriate workflow spec
|
||||
func (r *LLMRouter) Route(ctx context.Context, input LLMRouterInput) (*LLMRouterOutput, error) {
|
||||
// 1. Analyze intent using LLM provider
|
||||
intent, err := r.analyzeIntent(ctx, input)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("intent analysis failed: %w", err)
|
||||
}
|
||||
|
||||
// 2. Build workflow spec based on intent
|
||||
metadata := &BuildMetadata{
|
||||
KnowledgeBase: r.knowledgeBase,
|
||||
Context: input.Context,
|
||||
Validators: r.validators,
|
||||
}
|
||||
|
||||
spec, err := r.specBuilder.FromIntent(intent, metadata)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("spec build failed: %w", err)
|
||||
}
|
||||
|
||||
return &LLMRouterOutput{
|
||||
Spec: spec,
|
||||
IsCron: intent.IsCron,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Intent represents analyzed user intent
|
||||
type Intent struct {
|
||||
Activities []string `json:"activities"` // Selected activity names
|
||||
Parameters map[string]interface{} `json:"parameters"` // Extracted parameters
|
||||
IsCron bool `json:"isCron"` // Is scheduled workflow?
|
||||
CronSchedule string `json:"cronSchedule"` // Cron expression if scheduled
|
||||
CronTimezone string `json:"cronTimezone"` // Timezone for cron
|
||||
WorkflowName string `json:"workflowName"` // Generated workflow name
|
||||
ErrorHandling string `json:"errorHandling"` // "retry", "fail-fast", "continue"
|
||||
}
|
||||
|
||||
// analyzeIntent uses LLM provider to understand user request
|
||||
func (r *LLMRouter) analyzeIntent(ctx context.Context, input LLMRouterInput) (*Intent, error) {
|
||||
// Build prompt with knowledge base context
|
||||
userPrompt := r.buildIntentPrompt(input)
|
||||
|
||||
// Call LLM provider
|
||||
response, err := r.provider.Chat(ctx, intentSystemPrompt, userPrompt)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("LLM provider failed: %w", err)
|
||||
}
|
||||
|
||||
// Parse LLM response
|
||||
intent, err := parseIntentResponse(response)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse intent: %w", err)
|
||||
}
|
||||
|
||||
// Validate activities exist
|
||||
for _, actName := range intent.Activities {
|
||||
if !r.knowledgeBase.HasActivity(actName) {
|
||||
return nil, fmt.Errorf("unknown activity: %s", actName)
|
||||
}
|
||||
}
|
||||
|
||||
return intent, nil
|
||||
}
|
||||
|
||||
// buildIntentPrompt creates the prompt for intent analysis
|
||||
func (r *LLMRouter) buildIntentPrompt(input LLMRouterInput) string {
|
||||
// Get activity summaries
|
||||
var activityList strings.Builder
|
||||
for _, act := range r.knowledgeBase.Activities {
|
||||
activityList.WriteString(fmt.Sprintf("- %s: %s (category: %s, timeout: %s, flaky: %v)\n",
|
||||
act.Name, act.Description, act.Category,
|
||||
act.Constraints.DefaultTimeout, act.Constraints.IsFlaky))
|
||||
}
|
||||
|
||||
// Build context string
|
||||
contextStr := ""
|
||||
if len(input.Context) > 0 {
|
||||
ctxBytes, _ := json.Marshal(input.Context)
|
||||
contextStr = fmt.Sprintf("\nProvided context: %s", string(ctxBytes))
|
||||
}
|
||||
|
||||
// Build memory context string
|
||||
memoryStr := r.formatMemoryContext(input.MemoryContext)
|
||||
|
||||
return fmt.Sprintf(`User request: %s
|
||||
%s%s
|
||||
Available activities:
|
||||
%s
|
||||
Analyze the request and output JSON with:
|
||||
- activities: ordered list of activity names to execute
|
||||
- parameters: extracted parameters from request (repo URL, branch, etc)
|
||||
- isCron: true if user wants scheduled/recurring execution
|
||||
- cronSchedule: cron expression if scheduled (e.g., "0 2 * * *" for 2 AM daily)
|
||||
- cronTimezone: timezone (default "UTC")
|
||||
- workflowName: short descriptive name
|
||||
- errorHandling: "retry" (default), "fail-fast", or "continue"
|
||||
|
||||
Output ONLY valid JSON, no explanation.`, input.Message, contextStr, memoryStr, activityList.String())
|
||||
}
|
||||
|
||||
// formatMemoryContext formats memory context for prompt injection
|
||||
func (r *LLMRouter) formatMemoryContext(mem *MemoryContext) string {
|
||||
if mem == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
var sb strings.Builder
|
||||
|
||||
if len(mem.Skills) > 0 {
|
||||
sb.WriteString("\n\nRelevant skills from memory:\n")
|
||||
for _, skill := range mem.Skills {
|
||||
if skill.Why != "" {
|
||||
sb.WriteString(fmt.Sprintf("- %s: %s (reason: %s)\n", skill.Name, skill.Description, skill.Why))
|
||||
} else {
|
||||
sb.WriteString(fmt.Sprintf("- %s: %s\n", skill.Name, skill.Description))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(mem.Lessons) > 0 {
|
||||
sb.WriteString("\nRelevant knowledge from memory:\n")
|
||||
for _, lesson := range mem.Lessons {
|
||||
text := lesson.Text
|
||||
if len(text) > 300 {
|
||||
text = text[:300] + "..."
|
||||
}
|
||||
sb.WriteString(fmt.Sprintf("- [%s] %s\n", lesson.Level, text))
|
||||
}
|
||||
}
|
||||
|
||||
if len(mem.References) > 0 {
|
||||
sb.WriteString("\nReference documents:\n")
|
||||
for _, ref := range mem.References {
|
||||
text := ref.Text
|
||||
if len(text) > 200 {
|
||||
text = text[:200] + "..."
|
||||
}
|
||||
sb.WriteString(fmt.Sprintf("- %s\n", text))
|
||||
}
|
||||
}
|
||||
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
// parseIntentResponse extracts Intent from LLM response
|
||||
func parseIntentResponse(response string) (*Intent, error) {
|
||||
// Try to extract JSON from response
|
||||
response = strings.TrimSpace(response)
|
||||
|
||||
// Handle markdown code blocks
|
||||
if strings.HasPrefix(response, "```") {
|
||||
re := regexp.MustCompile("```(?:json)?\\s*([\\s\\S]*?)```")
|
||||
matches := re.FindStringSubmatch(response)
|
||||
if len(matches) > 1 {
|
||||
response = strings.TrimSpace(matches[1])
|
||||
}
|
||||
}
|
||||
|
||||
var intent Intent
|
||||
if err := json.Unmarshal([]byte(response), &intent); err != nil {
|
||||
return nil, fmt.Errorf("invalid JSON from LLM: %w\nResponse: %s", err, response)
|
||||
}
|
||||
|
||||
// Set defaults
|
||||
if intent.CronTimezone == "" {
|
||||
intent.CronTimezone = "UTC"
|
||||
}
|
||||
if intent.ErrorHandling == "" {
|
||||
intent.ErrorHandling = "retry"
|
||||
}
|
||||
if intent.WorkflowName == "" {
|
||||
intent.WorkflowName = "generated-workflow"
|
||||
}
|
||||
|
||||
return &intent, nil
|
||||
}
|
||||
|
||||
// buildSpec creates WorkflowSpec from intent
|
||||
func (r *LLMRouter) buildSpec(intent *Intent, input LLMRouterInput) (*WorkflowSpec, error) {
|
||||
if len(intent.Activities) == 0 {
|
||||
return nil, fmt.Errorf("no activities selected")
|
||||
}
|
||||
|
||||
states := make([]State, 0, len(intent.Activities)+1)
|
||||
|
||||
// Build states for each activity
|
||||
for i, actName := range intent.Activities {
|
||||
act := r.knowledgeBase.GetActivity(actName)
|
||||
|
||||
state := State{
|
||||
Name: actName,
|
||||
Type: StateTypeTask,
|
||||
Resource: actName,
|
||||
Parameters: r.buildParameters(act, intent, i),
|
||||
Timeout: act.Constraints.DefaultTimeout,
|
||||
Retry: r.buildRetryPolicy(act, intent),
|
||||
}
|
||||
|
||||
// Set next state or end
|
||||
if i < len(intent.Activities)-1 {
|
||||
state.Next = intent.Activities[i+1]
|
||||
} else {
|
||||
state.End = true
|
||||
}
|
||||
|
||||
// Add catch clause for flaky activities
|
||||
if act.Constraints.IsFlaky && intent.ErrorHandling != "fail-fast" {
|
||||
state.Catch = []CatchClause{
|
||||
{
|
||||
ErrorEquals: []string{"ActivityError", "TimeoutError"},
|
||||
Next: "HandleError",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
states = append(states, state)
|
||||
}
|
||||
|
||||
// Add error handler if needed
|
||||
hasFlaky := false
|
||||
for _, actName := range intent.Activities {
|
||||
act := r.knowledgeBase.GetActivity(actName)
|
||||
if act != nil && act.Constraints.IsFlaky {
|
||||
hasFlaky = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if hasFlaky && intent.ErrorHandling != "fail-fast" {
|
||||
states = append(states, State{
|
||||
Name: "HandleError",
|
||||
Type: StateTypeFail,
|
||||
Error: "WorkflowError",
|
||||
Cause: "Activity failed after retries",
|
||||
})
|
||||
}
|
||||
|
||||
// Build input map
|
||||
inputMap := make(map[string]interface{})
|
||||
for k, v := range intent.Parameters {
|
||||
inputMap[k] = v
|
||||
}
|
||||
for k, v := range input.Context {
|
||||
if _, exists := inputMap[k]; !exists {
|
||||
inputMap[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
return &WorkflowSpec{
|
||||
Name: intent.WorkflowName,
|
||||
Input: inputMap,
|
||||
States: states,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// getStringFromMap safely extracts a string from a map
|
||||
func getStringFromMap(m map[string]interface{}, key string) string {
|
||||
if m == nil {
|
||||
return ""
|
||||
}
|
||||
if v, ok := m[key].(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// firstNonEmpty returns the first non-empty string from the list
|
||||
func firstNonEmpty(values ...string) string {
|
||||
for _, v := range values {
|
||||
if v != "" {
|
||||
return v
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// buildCronSpec creates CronWorkflowSpec from intent
|
||||
func (r *LLMRouter) buildCronSpec(intent *Intent, input LLMRouterInput) (*CronWorkflowSpec, error) {
|
||||
spec, err := r.buildSpec(intent, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Extract schedule from multiple sources
|
||||
schedule := firstNonEmpty(
|
||||
intent.CronSchedule,
|
||||
getStringFromMap(intent.Parameters, "cronSchedule"),
|
||||
getStringFromMap(spec.Input, "cronSchedule"),
|
||||
)
|
||||
|
||||
// Extract timezone from multiple sources, default to UTC
|
||||
timezone := firstNonEmpty(
|
||||
intent.CronTimezone,
|
||||
getStringFromMap(intent.Parameters, "cronTimezone"),
|
||||
getStringFromMap(spec.Input, "cronTimezone"),
|
||||
"UTC",
|
||||
)
|
||||
|
||||
// Clean cron fields from input
|
||||
delete(spec.Input, "cronSchedule")
|
||||
delete(spec.Input, "cronTimezone")
|
||||
|
||||
return &CronWorkflowSpec{
|
||||
Name: spec.Name,
|
||||
Type: "CronWorkflow",
|
||||
Schedule: schedule,
|
||||
Timezone: timezone,
|
||||
Input: spec.Input,
|
||||
States: spec.States,
|
||||
MaxConcurrent: 1,
|
||||
Timeout: "1h",
|
||||
EnableHistory: true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// isCommonInputField checks if field name is a common workflow input
|
||||
func isCommonInputField(name string) bool {
|
||||
switch name {
|
||||
case "repo", "path", "branch":
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// paramResolver resolves activity parameters from multiple sources
|
||||
type paramResolver struct {
|
||||
intent *Intent
|
||||
kb *KnowledgeBase
|
||||
prevState string
|
||||
}
|
||||
|
||||
// resolve finds parameter value from intent, previous output, default, or input ref
|
||||
func (r *paramResolver) resolve(inputName string, inputDef InputField) interface{} {
|
||||
// 1. From intent parameters
|
||||
if val, ok := r.intent.Parameters[inputName]; ok {
|
||||
return val
|
||||
}
|
||||
|
||||
// 2. From previous state output
|
||||
if val := r.fromPrevOutput(inputName); val != nil {
|
||||
return val
|
||||
}
|
||||
|
||||
// 3. Default value
|
||||
if inputDef.Default != nil {
|
||||
return inputDef.Default
|
||||
}
|
||||
|
||||
// 4. Input reference for common fields
|
||||
if isCommonInputField(inputName) {
|
||||
return fmt.Sprintf("${input.%s}", inputName)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// fromPrevOutput checks if previous activity has matching output
|
||||
func (r *paramResolver) fromPrevOutput(inputName string) interface{} {
|
||||
if r.prevState == "" {
|
||||
return nil
|
||||
}
|
||||
prevActDef := r.kb.GetActivity(r.prevState)
|
||||
if prevActDef == nil {
|
||||
return nil
|
||||
}
|
||||
for outName := range prevActDef.Outputs {
|
||||
if outName == inputName || strings.EqualFold(outName, inputName) {
|
||||
return fmt.Sprintf("${%s.output.%s}", r.prevState, outName)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildParameters creates parameter map for activity
|
||||
func (r *LLMRouter) buildParameters(act *ActivityMetadata, intent *Intent, stateIndex int) map[string]interface{} {
|
||||
var prevState string
|
||||
if stateIndex > 0 {
|
||||
prevState = intent.Activities[stateIndex-1]
|
||||
}
|
||||
|
||||
resolver := ¶mResolver{
|
||||
intent: intent,
|
||||
kb: r.knowledgeBase,
|
||||
prevState: prevState,
|
||||
}
|
||||
|
||||
params := make(map[string]interface{})
|
||||
for name, def := range act.Inputs {
|
||||
if val := resolver.resolve(name, def); val != nil {
|
||||
params[name] = val
|
||||
}
|
||||
}
|
||||
return params
|
||||
}
|
||||
|
||||
// buildRetryPolicy creates retry policy based on activity constraints
|
||||
func (r *LLMRouter) buildRetryPolicy(act *ActivityMetadata, intent *Intent) *RetryPolicy {
|
||||
if intent.ErrorHandling == "fail-fast" {
|
||||
return &RetryPolicy{
|
||||
MaxAttempts: 1,
|
||||
BackoffRate: 1.0,
|
||||
InitialInterval: "1s",
|
||||
}
|
||||
}
|
||||
|
||||
return &RetryPolicy{
|
||||
MaxAttempts: int32(act.Constraints.RecommendedRetries),
|
||||
BackoffRate: act.Constraints.RetryBackoff,
|
||||
InitialInterval: "1s",
|
||||
MaxInterval: "30s",
|
||||
}
|
||||
}
|
||||
|
||||
const intentSystemPrompt = `You are an intelligent workflow router. Your job is to:
|
||||
1. Understand what the user wants to accomplish
|
||||
2. Select the appropriate activities from the available list
|
||||
3. Order them correctly based on dependencies
|
||||
4. Extract any parameters mentioned (URLs, branches, etc)
|
||||
5. Detect if user wants scheduled/recurring execution
|
||||
6. Use any relevant knowledge from memory to inform your decisions
|
||||
|
||||
Rules:
|
||||
- Always include CloneRepoActivity first if any analysis activity is needed
|
||||
- Order activities respecting dependencies
|
||||
- If user mentions "daily", "every hour", "weekly", etc → set isCron=true and cronSchedule
|
||||
- Common cron patterns: "0 2 * * *" (2 AM daily), "0 * * * *" (hourly), "0 0 * * 0" (weekly Sunday)
|
||||
- Extract repo URLs, branch names, severity levels from the message
|
||||
- workflowName should be short and descriptive (kebab-case)
|
||||
- If memory context includes relevant skills or lessons, incorporate that knowledge
|
||||
- Skills from memory may suggest specific activity parameters or ordering
|
||||
|
||||
Output ONLY valid JSON.`
|
||||
@@ -0,0 +1,103 @@
|
||||
// +build integration
|
||||
|
||||
package routing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// TestLLMRouterIntegration tests against real api.riotpiao.com
|
||||
// Run with: go test -tags=integration -v -run TestLLMRouterIntegration
|
||||
func TestLLMRouterIntegration(t *testing.T) {
|
||||
// Skip if not explicitly enabled
|
||||
if os.Getenv("RUN_INTEGRATION_TESTS") != "1" {
|
||||
t.Skip("Skipping integration test. Set RUN_INTEGRATION_TESTS=1 to run.")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
router, err := NewLLMRouter(kb)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create router: %v", err)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
input LLMRouterInput
|
||||
validate func(*testing.T, *LLMRouterOutput)
|
||||
}{
|
||||
{
|
||||
name: "analyze repo request",
|
||||
input: LLMRouterInput{
|
||||
Message: "Analyze the GitHub repo https://github.com/rockliang/poimen for code quality and security issues",
|
||||
Context: map[string]interface{}{
|
||||
"branch": "main",
|
||||
},
|
||||
},
|
||||
validate: func(t *testing.T, output *LLMRouterOutput) {
|
||||
if output.IsCron {
|
||||
t.Error("expected one-time workflow, not cron")
|
||||
}
|
||||
if output.Spec == nil {
|
||||
t.Fatal("expected spec, got nil")
|
||||
}
|
||||
if len(output.Spec.States) < 2 {
|
||||
t.Errorf("expected at least 2 states, got %d", len(output.Spec.States))
|
||||
}
|
||||
// Should start with CloneRepoActivity
|
||||
if output.Spec.States[0].Resource != "CloneRepoActivity" {
|
||||
t.Errorf("expected first activity to be CloneRepoActivity, got %s", output.Spec.States[0].Resource)
|
||||
}
|
||||
t.Logf("Generated workflow: %s with %d states", output.Spec.Name, len(output.Spec.States))
|
||||
for i, state := range output.Spec.States {
|
||||
t.Logf(" State %d: %s (%s)", i, state.Name, state.Resource)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "daily security scan (cron)",
|
||||
input: LLMRouterInput{
|
||||
Message: "Run a security scan on https://github.com/rockliang/poimen every day at 3 AM UTC",
|
||||
},
|
||||
validate: func(t *testing.T, output *LLMRouterOutput) {
|
||||
if !output.IsCron {
|
||||
t.Error("expected cron workflow")
|
||||
}
|
||||
if output.CronSpec == nil {
|
||||
t.Fatal("expected cron spec, got nil")
|
||||
}
|
||||
if output.CronSpec.Schedule == "" {
|
||||
t.Error("expected cron schedule")
|
||||
}
|
||||
t.Logf("Generated cron workflow: %s, schedule: %s", output.CronSpec.Name, output.CronSpec.Schedule)
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
|
||||
defer cancel()
|
||||
|
||||
output, err := router.Route(ctx, tt.input)
|
||||
if err != nil {
|
||||
t.Fatalf("Route failed: %v", err)
|
||||
}
|
||||
|
||||
// Pretty print output
|
||||
jsonOut, _ := json.MarshalIndent(output, "", " ")
|
||||
t.Logf("Output:\n%s", string(jsonOut))
|
||||
|
||||
if tt.validate != nil {
|
||||
tt.validate(t, output)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,492 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseIntentResponse(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
response string
|
||||
wantErr bool
|
||||
validate func(*testing.T, *Intent)
|
||||
}{
|
||||
{
|
||||
name: "basic intent",
|
||||
response: `{
|
||||
"activities": ["CloneRepoActivity", "AnalyzeCodeActivity"],
|
||||
"parameters": {"repo": "https://github.com/test/repo"},
|
||||
"isCron": false,
|
||||
"workflowName": "analyze-repo"
|
||||
}`,
|
||||
wantErr: false,
|
||||
validate: func(t *testing.T, intent *Intent) {
|
||||
if len(intent.Activities) != 2 {
|
||||
t.Errorf("expected 2 activities, got %d", len(intent.Activities))
|
||||
}
|
||||
if intent.Activities[0] != "CloneRepoActivity" {
|
||||
t.Errorf("expected CloneRepoActivity first, got %s", intent.Activities[0])
|
||||
}
|
||||
if intent.IsCron {
|
||||
t.Error("expected isCron=false")
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "cron intent",
|
||||
response: `{
|
||||
"activities": ["CloneRepoActivity", "SecurityScanActivity"],
|
||||
"parameters": {"repo": "https://github.com/test/repo"},
|
||||
"isCron": true,
|
||||
"cronSchedule": "0 2 * * *",
|
||||
"cronTimezone": "America/New_York",
|
||||
"workflowName": "daily-security-scan"
|
||||
}`,
|
||||
wantErr: false,
|
||||
validate: func(t *testing.T, intent *Intent) {
|
||||
if !intent.IsCron {
|
||||
t.Error("expected isCron=true")
|
||||
}
|
||||
if intent.CronSchedule != "0 2 * * *" {
|
||||
t.Errorf("expected cron schedule '0 2 * * *', got %s", intent.CronSchedule)
|
||||
}
|
||||
if intent.CronTimezone != "America/New_York" {
|
||||
t.Errorf("expected timezone 'America/New_York', got %s", intent.CronTimezone)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "with markdown code block",
|
||||
response: "```json\n{\"activities\": [\"CloneRepoActivity\"], \"parameters\": {}, \"isCron\": false}\n```",
|
||||
wantErr: false,
|
||||
validate: func(t *testing.T, intent *Intent) {
|
||||
if len(intent.Activities) != 1 {
|
||||
t.Errorf("expected 1 activity, got %d", len(intent.Activities))
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "defaults applied",
|
||||
response: `{"activities": ["CloneRepoActivity"], "parameters": {}}`,
|
||||
wantErr: false,
|
||||
validate: func(t *testing.T, intent *Intent) {
|
||||
if intent.CronTimezone != "UTC" {
|
||||
t.Errorf("expected default timezone UTC, got %s", intent.CronTimezone)
|
||||
}
|
||||
if intent.ErrorHandling != "retry" {
|
||||
t.Errorf("expected default errorHandling 'retry', got %s", intent.ErrorHandling)
|
||||
}
|
||||
if intent.WorkflowName != "generated-workflow" {
|
||||
t.Errorf("expected default workflowName, got %s", intent.WorkflowName)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "invalid json",
|
||||
response: "this is not json",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
intent, err := parseIntentResponse(tt.response)
|
||||
if tt.wantErr {
|
||||
if err == nil {
|
||||
t.Error("expected error, got nil")
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Errorf("unexpected error: %v", err)
|
||||
return
|
||||
}
|
||||
if tt.validate != nil {
|
||||
tt.validate(t, intent)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildSpec(t *testing.T) {
|
||||
// Load knowledge base
|
||||
kb, err := LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
router := &LLMRouter{
|
||||
knowledgeBase: kb,
|
||||
}
|
||||
|
||||
intent := &Intent{
|
||||
Activities: []string{"CloneRepoActivity", "AnalyzeCodeActivity", "SecurityScanActivity"},
|
||||
Parameters: map[string]interface{}{"repo": "https://github.com/test/repo", "branch": "main"},
|
||||
WorkflowName: "test-workflow",
|
||||
ErrorHandling: "retry",
|
||||
}
|
||||
|
||||
input := LLMRouterInput{
|
||||
Message: "Analyze repo for security",
|
||||
Context: map[string]interface{}{},
|
||||
}
|
||||
|
||||
spec, err := router.buildSpec(intent, input)
|
||||
if err != nil {
|
||||
t.Fatalf("buildSpec failed: %v", err)
|
||||
}
|
||||
|
||||
// Validate spec
|
||||
if spec.Name != "test-workflow" {
|
||||
t.Errorf("expected name 'test-workflow', got %s", spec.Name)
|
||||
}
|
||||
|
||||
if len(spec.States) < 3 {
|
||||
t.Errorf("expected at least 3 states, got %d", len(spec.States))
|
||||
}
|
||||
|
||||
// First state should be CloneRepoActivity
|
||||
if spec.States[0].Resource != "CloneRepoActivity" {
|
||||
t.Errorf("expected first state to be CloneRepoActivity, got %s", spec.States[0].Resource)
|
||||
}
|
||||
|
||||
// Last activity state should have End=true
|
||||
lastActivityIdx := len(spec.States) - 1
|
||||
if spec.States[lastActivityIdx].Type == StateTypeFail {
|
||||
lastActivityIdx--
|
||||
}
|
||||
if !spec.States[lastActivityIdx].End {
|
||||
t.Error("expected last activity state to have End=true")
|
||||
}
|
||||
|
||||
// Check retry policy on flaky activity (AnalyzeCodeActivity)
|
||||
for _, state := range spec.States {
|
||||
if state.Resource == "AnalyzeCodeActivity" {
|
||||
if state.Retry == nil {
|
||||
t.Error("expected retry policy on flaky activity")
|
||||
} else if state.Retry.MaxAttempts != 3 {
|
||||
t.Errorf("expected 3 max attempts for flaky activity, got %d", state.Retry.MaxAttempts)
|
||||
}
|
||||
if len(state.Catch) == 0 {
|
||||
t.Error("expected catch clause on flaky activity")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildCronSpec(t *testing.T) {
|
||||
kb, err := LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
router := &LLMRouter{
|
||||
knowledgeBase: kb,
|
||||
}
|
||||
|
||||
intent := &Intent{
|
||||
Activities: []string{"CloneRepoActivity", "SecurityScanActivity"},
|
||||
Parameters: map[string]interface{}{"repo": "https://github.com/test/repo"},
|
||||
IsCron: true,
|
||||
CronSchedule: "0 2 * * *",
|
||||
CronTimezone: "UTC",
|
||||
WorkflowName: "daily-scan",
|
||||
}
|
||||
|
||||
input := LLMRouterInput{
|
||||
Message: "Run security scan daily at 2 AM",
|
||||
}
|
||||
|
||||
cronSpec, err := router.buildCronSpec(intent, input)
|
||||
if err != nil {
|
||||
t.Fatalf("buildCronSpec failed: %v", err)
|
||||
}
|
||||
|
||||
if cronSpec.Type != "CronWorkflow" {
|
||||
t.Errorf("expected type 'CronWorkflow', got %s", cronSpec.Type)
|
||||
}
|
||||
if cronSpec.Schedule != "0 2 * * *" {
|
||||
t.Errorf("expected schedule '0 2 * * *', got %s", cronSpec.Schedule)
|
||||
}
|
||||
if cronSpec.Timezone != "UTC" {
|
||||
t.Errorf("expected timezone 'UTC', got %s", cronSpec.Timezone)
|
||||
}
|
||||
if !cronSpec.EnableHistory {
|
||||
t.Error("expected EnableHistory=true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildParameters(t *testing.T) {
|
||||
kb, err := LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
router := &LLMRouter{
|
||||
knowledgeBase: kb,
|
||||
}
|
||||
|
||||
// Test first activity (CloneRepoActivity) - should use input references
|
||||
cloneAct := kb.GetActivity("CloneRepoActivity")
|
||||
intent := &Intent{
|
||||
Activities: []string{"CloneRepoActivity", "AnalyzeCodeActivity"},
|
||||
Parameters: map[string]interface{}{"repo": "https://github.com/test/repo"},
|
||||
}
|
||||
|
||||
params := router.buildParameters(cloneAct, intent, 0)
|
||||
if params["repo"] != "https://github.com/test/repo" {
|
||||
t.Errorf("expected repo from parameters, got %v", params["repo"])
|
||||
}
|
||||
|
||||
// Test second activity (AnalyzeCodeActivity) - should reference previous output
|
||||
analyzeAct := kb.GetActivity("AnalyzeCodeActivity")
|
||||
params = router.buildParameters(analyzeAct, intent, 1)
|
||||
if params["path"] != "${CloneRepoActivity.output.path}" {
|
||||
t.Errorf("expected JSONPath reference to CloneRepoActivity.output.path, got %v", params["path"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRetryPolicy(t *testing.T) {
|
||||
kb, err := LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
router := &LLMRouter{
|
||||
knowledgeBase: kb,
|
||||
}
|
||||
|
||||
// Flaky activity with retry error handling
|
||||
analyzeAct := kb.GetActivity("AnalyzeCodeActivity")
|
||||
intent := &Intent{ErrorHandling: "retry"}
|
||||
policy := router.buildRetryPolicy(analyzeAct, intent)
|
||||
|
||||
if policy.MaxAttempts != 3 {
|
||||
t.Errorf("expected 3 max attempts for flaky activity, got %d", policy.MaxAttempts)
|
||||
}
|
||||
if policy.BackoffRate != 2.0 {
|
||||
t.Errorf("expected backoff rate 2.0, got %f", policy.BackoffRate)
|
||||
}
|
||||
|
||||
// Fail-fast error handling
|
||||
intent = &Intent{ErrorHandling: "fail-fast"}
|
||||
policy = router.buildRetryPolicy(analyzeAct, intent)
|
||||
|
||||
if policy.MaxAttempts != 1 {
|
||||
t.Errorf("expected 1 max attempt for fail-fast, got %d", policy.MaxAttempts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetStringFromMap(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
m map[string]interface{}
|
||||
key string
|
||||
expected string
|
||||
}{
|
||||
{"nil map", nil, "key", ""},
|
||||
{"missing key", map[string]interface{}{"a": "b"}, "key", ""},
|
||||
{"found string", map[string]interface{}{"key": "value"}, "key", "value"},
|
||||
{"non-string value", map[string]interface{}{"key": 123}, "key", ""},
|
||||
{"empty string", map[string]interface{}{"key": ""}, "key", ""},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := getStringFromMap(tt.m, tt.key)
|
||||
if got != tt.expected {
|
||||
t.Errorf("getStringFromMap() = %q, want %q", got, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFirstNonEmpty(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
values []string
|
||||
expected string
|
||||
}{
|
||||
{"all empty", []string{"", "", ""}, ""},
|
||||
{"first non-empty", []string{"first", "second"}, "first"},
|
||||
{"second non-empty", []string{"", "second", "third"}, "second"},
|
||||
{"last non-empty", []string{"", "", "last"}, "last"},
|
||||
{"no values", []string{}, ""},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := firstNonEmpty(tt.values...)
|
||||
if got != tt.expected {
|
||||
t.Errorf("firstNonEmpty() = %q, want %q", got, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsCommonInputField(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
expected bool
|
||||
}{
|
||||
{"repo", true},
|
||||
{"path", true},
|
||||
{"branch", true},
|
||||
{"unknown", false},
|
||||
{"Repository", false}, // case-sensitive
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := isCommonInputField(tt.name)
|
||||
if got != tt.expected {
|
||||
t.Errorf("isCommonInputField(%q) = %v, want %v", tt.name, got, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildCronSpecScheduleSources(t *testing.T) {
|
||||
kb, err := LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
router := &LLMRouter{knowledgeBase: kb}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
intentSchedule string
|
||||
intentTimezone string
|
||||
paramSchedule string
|
||||
paramTimezone string
|
||||
wantSchedule string
|
||||
wantTimezone string
|
||||
}{
|
||||
{
|
||||
name: "from intent",
|
||||
intentSchedule: "0 2 * * *",
|
||||
intentTimezone: "PST",
|
||||
wantSchedule: "0 2 * * *",
|
||||
wantTimezone: "PST",
|
||||
},
|
||||
{
|
||||
name: "from params",
|
||||
paramSchedule: "0 3 * * *",
|
||||
paramTimezone: "EST",
|
||||
wantSchedule: "0 3 * * *",
|
||||
wantTimezone: "EST",
|
||||
},
|
||||
{
|
||||
name: "default UTC",
|
||||
wantSchedule: "",
|
||||
wantTimezone: "UTC",
|
||||
},
|
||||
{
|
||||
name: "intent priority",
|
||||
intentSchedule: "0 1 * * *",
|
||||
intentTimezone: "UTC",
|
||||
paramSchedule: "0 2 * * *",
|
||||
paramTimezone: "PST",
|
||||
wantSchedule: "0 1 * * *",
|
||||
wantTimezone: "UTC",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
intent := &Intent{
|
||||
Activities: []string{"CloneRepoActivity"},
|
||||
Parameters: map[string]interface{}{},
|
||||
IsCron: true,
|
||||
CronSchedule: tt.intentSchedule,
|
||||
CronTimezone: tt.intentTimezone,
|
||||
WorkflowName: "test",
|
||||
}
|
||||
if tt.paramSchedule != "" {
|
||||
intent.Parameters["cronSchedule"] = tt.paramSchedule
|
||||
}
|
||||
if tt.paramTimezone != "" {
|
||||
intent.Parameters["cronTimezone"] = tt.paramTimezone
|
||||
}
|
||||
|
||||
input := LLMRouterInput{Message: "test"}
|
||||
spec, err := router.buildCronSpec(intent, input)
|
||||
if err != nil {
|
||||
t.Fatalf("buildCronSpec failed: %v", err)
|
||||
}
|
||||
|
||||
if spec.Schedule != tt.wantSchedule {
|
||||
t.Errorf("schedule = %q, want %q", spec.Schedule, tt.wantSchedule)
|
||||
}
|
||||
if spec.Timezone != tt.wantTimezone {
|
||||
t.Errorf("timezone = %q, want %q", spec.Timezone, tt.wantTimezone)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParamResolverFromPrevOutput(t *testing.T) {
|
||||
kb, err := LoadKnowledgeBaseFromDefaultPath()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load knowledge base: %v", err)
|
||||
}
|
||||
|
||||
resolver := ¶mResolver{
|
||||
intent: &Intent{
|
||||
Activities: []string{"CloneRepoActivity", "AnalyzeCodeActivity"},
|
||||
Parameters: map[string]interface{}{},
|
||||
},
|
||||
kb: kb,
|
||||
prevState: "CloneRepoActivity",
|
||||
}
|
||||
|
||||
// Should find matching output from CloneRepoActivity
|
||||
got := resolver.fromPrevOutput("path")
|
||||
if got == nil {
|
||||
t.Error("expected to find path from prev output")
|
||||
}
|
||||
if got != "${CloneRepoActivity.output.path}" {
|
||||
t.Errorf("got %v, want ${CloneRepoActivity.output.path}", got)
|
||||
}
|
||||
|
||||
// Should not find non-existent output
|
||||
got = resolver.fromPrevOutput("nonexistent")
|
||||
if got != nil {
|
||||
t.Errorf("expected nil for nonexistent, got %v", got)
|
||||
}
|
||||
|
||||
// Empty prevState should return nil
|
||||
resolver.prevState = ""
|
||||
got = resolver.fromPrevOutput("path")
|
||||
if got != nil {
|
||||
t.Errorf("expected nil for empty prevState, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIntentJSONMarshal(t *testing.T) {
|
||||
intent := &Intent{
|
||||
Activities: []string{"CloneRepoActivity"},
|
||||
Parameters: map[string]interface{}{"repo": "https://test"},
|
||||
IsCron: true,
|
||||
CronSchedule: "0 * * * *",
|
||||
CronTimezone: "UTC",
|
||||
WorkflowName: "test",
|
||||
ErrorHandling: "retry",
|
||||
}
|
||||
|
||||
data, err := json.Marshal(intent)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal failed: %v", err)
|
||||
}
|
||||
|
||||
var decoded Intent
|
||||
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||
t.Fatalf("unmarshal failed: %v", err)
|
||||
}
|
||||
|
||||
if decoded.CronSchedule != intent.CronSchedule {
|
||||
t.Errorf("expected schedule %s, got %s", intent.CronSchedule, decoded.CronSchedule)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,252 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// LLMProvider defines interface for LLM services
|
||||
type LLMProvider interface {
|
||||
// Name returns provider name (e.g., "openai", "claude", "local")
|
||||
Name() string
|
||||
|
||||
// Chat sends a message and returns response
|
||||
Chat(ctx context.Context, systemPrompt, userPrompt string) (string, error)
|
||||
|
||||
// IsAvailable checks if provider is configured and reachable
|
||||
IsAvailable(ctx context.Context) error
|
||||
}
|
||||
|
||||
// ProviderRegistry manages available LLM providers
|
||||
type ProviderRegistry struct {
|
||||
providers map[string]LLMProvider
|
||||
default_ string
|
||||
}
|
||||
|
||||
// NewProviderRegistry creates a new registry
|
||||
func NewProviderRegistry() *ProviderRegistry {
|
||||
return &ProviderRegistry{
|
||||
providers: make(map[string]LLMProvider),
|
||||
}
|
||||
}
|
||||
|
||||
// Register adds a provider
|
||||
func (pr *ProviderRegistry) Register(provider LLMProvider) error {
|
||||
if provider.Name() == "" {
|
||||
return fmt.Errorf("provider name cannot be empty")
|
||||
}
|
||||
pr.providers[provider.Name()] = provider
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetDefault sets the default provider
|
||||
func (pr *ProviderRegistry) SetDefault(name string) error {
|
||||
if _, exists := pr.providers[name]; !exists {
|
||||
return fmt.Errorf("provider %s not registered", name)
|
||||
}
|
||||
pr.default_ = name
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get retrieves a provider by name
|
||||
func (pr *ProviderRegistry) Get(name string) (LLMProvider, error) {
|
||||
if name == "" {
|
||||
name = pr.default_
|
||||
}
|
||||
provider, exists := pr.providers[name]
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("provider %s not found", name)
|
||||
}
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
// GetDefault returns the default provider
|
||||
func (pr *ProviderRegistry) GetDefault() (LLMProvider, error) {
|
||||
if pr.default_ == "" {
|
||||
return nil, fmt.Errorf("no default provider set")
|
||||
}
|
||||
return pr.Get(pr.default_)
|
||||
}
|
||||
|
||||
// RoutingProviderLLM routes between multiple LLM providers with fallback
|
||||
type RoutingProviderLLM struct {
|
||||
registry *ProviderRegistry
|
||||
fallbackOrder []string
|
||||
}
|
||||
|
||||
// NewRoutingProviderLLM creates a routing LLM
|
||||
func NewRoutingProviderLLM(registry *ProviderRegistry, order ...string) *RoutingProviderLLM {
|
||||
return &RoutingProviderLLM{
|
||||
registry: registry,
|
||||
fallbackOrder: order,
|
||||
}
|
||||
}
|
||||
|
||||
// Chat tries providers in order
|
||||
func (rp *RoutingProviderLLM) Chat(ctx context.Context, systemPrompt, userPrompt string) (string, error) {
|
||||
for _, providerName := range rp.fallbackOrder {
|
||||
provider, err := rp.registry.Get(providerName)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
if err := provider.IsAvailable(ctx); err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
response, err := provider.Chat(ctx, systemPrompt, userPrompt)
|
||||
if err == nil {
|
||||
return response, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("all LLM providers failed")
|
||||
}
|
||||
|
||||
// CachingLLMProvider wraps a provider with caching
|
||||
type CachingLLMProvider struct {
|
||||
provider LLMProvider
|
||||
cache map[string]string
|
||||
}
|
||||
|
||||
// NewCachingLLMProvider creates a cached provider
|
||||
func NewCachingLLMProvider(provider LLMProvider) *CachingLLMProvider {
|
||||
return &CachingLLMProvider{
|
||||
provider: provider,
|
||||
cache: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
// Chat returns cached response if available
|
||||
func (clp *CachingLLMProvider) Chat(ctx context.Context, systemPrompt, userPrompt string) (string, error) {
|
||||
key := systemPrompt + "|" + userPrompt
|
||||
|
||||
if cached, exists := clp.cache[key]; exists {
|
||||
return cached, nil
|
||||
}
|
||||
|
||||
response, err := clp.provider.Chat(ctx, systemPrompt, userPrompt)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
clp.cache[key] = response
|
||||
return response, nil
|
||||
}
|
||||
|
||||
// IsAvailable delegates to wrapped provider
|
||||
func (clp *CachingLLMProvider) IsAvailable(ctx context.Context) error {
|
||||
return clp.provider.IsAvailable(ctx)
|
||||
}
|
||||
|
||||
// Name delegates to wrapped provider
|
||||
func (clp *CachingLLMProvider) Name() string {
|
||||
return clp.provider.Name() + "-cached"
|
||||
}
|
||||
|
||||
// RetryingLLMProvider wraps a provider with retry logic
|
||||
type RetryingLLMProvider struct {
|
||||
provider LLMProvider
|
||||
maxRetries int
|
||||
backoffFunc func(attempt int) interface{}
|
||||
}
|
||||
|
||||
// NewRetryingLLMProvider creates a retrying provider
|
||||
func NewRetryingLLMProvider(provider LLMProvider, maxRetries int) *RetryingLLMProvider {
|
||||
return &RetryingLLMProvider{
|
||||
provider: provider,
|
||||
maxRetries: maxRetries,
|
||||
backoffFunc: func(attempt int) interface{} {
|
||||
// Exponential backoff: 1s, 2s, 4s...
|
||||
return 1 << uint(attempt)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Chat retries on failure
|
||||
func (rlp *RetryingLLMProvider) Chat(ctx context.Context, systemPrompt, userPrompt string) (string, error) {
|
||||
var lastErr error
|
||||
|
||||
for attempt := 0; attempt <= rlp.maxRetries; attempt++ {
|
||||
response, err := rlp.provider.Chat(ctx, systemPrompt, userPrompt)
|
||||
if err == nil {
|
||||
return response, nil
|
||||
}
|
||||
lastErr = err
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("failed after %d retries: %w", rlp.maxRetries, lastErr)
|
||||
}
|
||||
|
||||
// IsAvailable delegates to wrapped provider
|
||||
func (rlp *RetryingLLMProvider) IsAvailable(ctx context.Context) error {
|
||||
return rlp.provider.IsAvailable(ctx)
|
||||
}
|
||||
|
||||
// Name delegates to wrapped provider
|
||||
func (rlp *RetryingLLMProvider) Name() string {
|
||||
return rlp.provider.Name() + "-retrying"
|
||||
}
|
||||
|
||||
// PromptTemplate defines a reusable prompt structure
|
||||
type PromptTemplate interface {
|
||||
// Render creates a prompt from values
|
||||
Render(values map[string]interface{}) (string, error)
|
||||
}
|
||||
|
||||
// SimplePromptTemplate uses Go text/template syntax
|
||||
type SimplePromptTemplate struct {
|
||||
template string
|
||||
}
|
||||
|
||||
// NewSimplePromptTemplate creates a simple template
|
||||
func NewSimplePromptTemplate(template string) *SimplePromptTemplate {
|
||||
return &SimplePromptTemplate{template: template}
|
||||
}
|
||||
|
||||
// Render renders the template (placeholder implementation)
|
||||
func (spt *SimplePromptTemplate) Render(values map[string]interface{}) (string, error) {
|
||||
// In real implementation, use text/template
|
||||
return spt.template, nil
|
||||
}
|
||||
|
||||
// PromptBuilder builds prompts from components
|
||||
type PromptBuilder struct {
|
||||
system string
|
||||
sections []string
|
||||
}
|
||||
|
||||
// NewPromptBuilder creates a new builder
|
||||
func NewPromptBuilder() *PromptBuilder {
|
||||
return &PromptBuilder{
|
||||
sections: []string{},
|
||||
}
|
||||
}
|
||||
|
||||
// System sets the system prompt
|
||||
func (pb *PromptBuilder) System(prompt string) *PromptBuilder {
|
||||
pb.system = prompt
|
||||
return pb
|
||||
}
|
||||
|
||||
// AddSection adds a prompt section
|
||||
func (pb *PromptBuilder) AddSection(title, content string) *PromptBuilder {
|
||||
if title != "" {
|
||||
pb.sections = append(pb.sections, fmt.Sprintf("## %s\n%s", title, content))
|
||||
} else {
|
||||
pb.sections = append(pb.sections, content)
|
||||
}
|
||||
return pb
|
||||
}
|
||||
|
||||
// Build returns the complete prompt
|
||||
func (pb *PromptBuilder) Build() (system, user string) {
|
||||
user = ""
|
||||
for i, section := range pb.sections {
|
||||
if i > 0 {
|
||||
user += "\n\n"
|
||||
}
|
||||
user += section
|
||||
}
|
||||
return pb.system, user
|
||||
}
|
||||
@@ -0,0 +1,280 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// SpecBuilder defines interface for building workflow specs
|
||||
type SpecBuilder interface {
|
||||
// FromIntent builds spec from an analyzed intent
|
||||
FromIntent(intent *Intent, metadata *BuildMetadata) (*WorkflowSpec, error)
|
||||
|
||||
// Validate checks if builder can build this intent
|
||||
Validate(intent *Intent) error
|
||||
}
|
||||
|
||||
// BuildMetadata contains metadata for spec building
|
||||
type BuildMetadata struct {
|
||||
KnowledgeBase *KnowledgeBase
|
||||
Context map[string]interface{}
|
||||
Validators []WorkflowValidator
|
||||
}
|
||||
|
||||
// DefaultSpecBuilder implements basic spec building
|
||||
type DefaultSpecBuilder struct {
|
||||
kb *KnowledgeBase
|
||||
}
|
||||
|
||||
// NewDefaultSpecBuilder creates a builder
|
||||
func NewDefaultSpecBuilder(kb *KnowledgeBase) *DefaultSpecBuilder {
|
||||
return &DefaultSpecBuilder{kb: kb}
|
||||
}
|
||||
|
||||
// Validate checks if intent is buildable
|
||||
func (dsb *DefaultSpecBuilder) Validate(intent *Intent) error {
|
||||
if len(intent.Activities) == 0 {
|
||||
return fmt.Errorf("intent has no activities")
|
||||
}
|
||||
|
||||
for _, actName := range intent.Activities {
|
||||
if !dsb.kb.HasActivity(actName) {
|
||||
return fmt.Errorf("activity %s not found", actName)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// FromIntent builds spec from intent
|
||||
func (dsb *DefaultSpecBuilder) FromIntent(intent *Intent, metadata *BuildMetadata) (*WorkflowSpec, error) {
|
||||
if err := dsb.Validate(intent); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
states := make([]State, 0, len(intent.Activities)+1)
|
||||
|
||||
// Build states for each activity
|
||||
for i, actName := range intent.Activities {
|
||||
act := dsb.kb.GetActivity(actName)
|
||||
|
||||
state := State{
|
||||
Name: actName,
|
||||
Type: StateTypeTask,
|
||||
Resource: actName,
|
||||
Parameters: dsb.buildParameters(act, intent, i),
|
||||
Timeout: act.Constraints.DefaultTimeout,
|
||||
Retry: dsb.buildRetryPolicy(act, intent),
|
||||
}
|
||||
|
||||
// Set next state
|
||||
if i < len(intent.Activities)-1 {
|
||||
state.Next = intent.Activities[i+1]
|
||||
} else {
|
||||
state.End = true
|
||||
}
|
||||
|
||||
// Add catch for flaky activities
|
||||
if act.Constraints.IsFlaky && intent.ErrorHandling != "fail-fast" {
|
||||
state.Catch = []CatchClause{
|
||||
{
|
||||
ErrorEquals: []string{"ActivityError", "TimeoutError"},
|
||||
Next: "HandleError",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
states = append(states, state)
|
||||
}
|
||||
|
||||
// Add error handler if needed
|
||||
if dsb.hasFlaky(intent) && intent.ErrorHandling != "fail-fast" {
|
||||
states = append(states, State{
|
||||
Name: "HandleError",
|
||||
Type: StateTypeFail,
|
||||
Error: "ActivityFailed",
|
||||
Cause: "One or more activities failed",
|
||||
})
|
||||
}
|
||||
|
||||
// Build input map
|
||||
inputMap := make(map[string]interface{})
|
||||
for k, v := range intent.Parameters {
|
||||
inputMap[k] = v
|
||||
}
|
||||
if metadata != nil && metadata.Context != nil {
|
||||
for k, v := range metadata.Context {
|
||||
if _, exists := inputMap[k]; !exists {
|
||||
inputMap[k] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: intent.WorkflowName,
|
||||
Input: inputMap,
|
||||
States: states,
|
||||
}
|
||||
|
||||
// Validate if validators provided
|
||||
if metadata != nil && len(metadata.Validators) > 0 {
|
||||
for _, v := range metadata.Validators {
|
||||
if err := v.Validate(spec); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return spec, nil
|
||||
}
|
||||
|
||||
// buildParameters creates parameters for activity
|
||||
func (dsb *DefaultSpecBuilder) buildParameters(act *ActivityMetadata, intent *Intent, stateIndex int) map[string]interface{} {
|
||||
params := make(map[string]interface{})
|
||||
resolver := ¶mResolver{
|
||||
intent: intent,
|
||||
kb: dsb.kb,
|
||||
}
|
||||
|
||||
if stateIndex > 0 {
|
||||
resolver.prevState = intent.Activities[stateIndex-1]
|
||||
}
|
||||
|
||||
for name, def := range act.Inputs {
|
||||
if val := resolver.resolve(name, def); val != nil {
|
||||
params[name] = val
|
||||
}
|
||||
}
|
||||
|
||||
return params
|
||||
}
|
||||
|
||||
// buildRetryPolicy creates retry policy
|
||||
func (dsb *DefaultSpecBuilder) buildRetryPolicy(act *ActivityMetadata, intent *Intent) *RetryPolicy {
|
||||
if intent.ErrorHandling == "fail-fast" {
|
||||
return &RetryPolicy{
|
||||
MaxAttempts: 1,
|
||||
BackoffRate: 1.0,
|
||||
InitialInterval: "1s",
|
||||
}
|
||||
}
|
||||
|
||||
return &RetryPolicy{
|
||||
MaxAttempts: int32(act.Constraints.RecommendedRetries),
|
||||
BackoffRate: act.Constraints.RetryBackoff,
|
||||
InitialInterval: "1s",
|
||||
MaxInterval: "30s",
|
||||
}
|
||||
}
|
||||
|
||||
// hasFlaky checks if any activity is flaky
|
||||
func (dsb *DefaultSpecBuilder) hasFlaky(intent *Intent) bool {
|
||||
for _, actName := range intent.Activities {
|
||||
act := dsb.kb.GetActivity(actName)
|
||||
if act != nil && act.Constraints.IsFlaky {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// CronSpecBuilder builds cron workflow specs
|
||||
type CronSpecBuilder struct {
|
||||
regularBuilder SpecBuilder
|
||||
}
|
||||
|
||||
// NewCronSpecBuilder creates a cron builder
|
||||
func NewCronSpecBuilder(regularBuilder SpecBuilder) *CronSpecBuilder {
|
||||
return &CronSpecBuilder{
|
||||
regularBuilder: regularBuilder,
|
||||
}
|
||||
}
|
||||
|
||||
// Validate checks if intent is valid for cron
|
||||
func (csb *CronSpecBuilder) Validate(intent *Intent) error {
|
||||
if !intent.IsCron {
|
||||
return fmt.Errorf("intent is not marked as cron")
|
||||
}
|
||||
|
||||
if intent.CronSchedule == "" {
|
||||
return fmt.Errorf("cron schedule is empty")
|
||||
}
|
||||
|
||||
return csb.regularBuilder.Validate(intent)
|
||||
}
|
||||
|
||||
// FromIntent builds cron spec
|
||||
func (csb *CronSpecBuilder) FromIntent(intent *Intent, metadata *BuildMetadata) (*WorkflowSpec, error) {
|
||||
spec, err := csb.regularBuilder.FromIntent(intent, metadata)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// In production, wrap with cron metadata
|
||||
// For now, just return the regular spec
|
||||
return spec, nil
|
||||
}
|
||||
|
||||
// SpecBuilderFactory creates appropriate spec builders
|
||||
type SpecBuilderFactory struct {
|
||||
kb *KnowledgeBase
|
||||
}
|
||||
|
||||
// NewSpecBuilderFactory creates a factory
|
||||
func NewSpecBuilderFactory(kb *KnowledgeBase) *SpecBuilderFactory {
|
||||
return &SpecBuilderFactory{kb: kb}
|
||||
}
|
||||
|
||||
// CreateBuilder creates appropriate builder for intent
|
||||
func (sbf *SpecBuilderFactory) CreateBuilder(intent *Intent) (SpecBuilder, error) {
|
||||
if intent.IsCron {
|
||||
return NewCronSpecBuilder(NewDefaultSpecBuilder(sbf.kb)), nil
|
||||
}
|
||||
|
||||
return NewDefaultSpecBuilder(sbf.kb), nil
|
||||
}
|
||||
|
||||
// CompositeSpecBuilder combines multiple builders with fallback
|
||||
type CompositeSpecBuilder struct {
|
||||
builders []SpecBuilder
|
||||
}
|
||||
|
||||
// NewCompositeSpecBuilder creates a composite builder
|
||||
func NewCompositeSpecBuilder(builders ...SpecBuilder) *CompositeSpecBuilder {
|
||||
return &CompositeSpecBuilder{builders: builders}
|
||||
}
|
||||
|
||||
// Validate tries each builder
|
||||
func (csb *CompositeSpecBuilder) Validate(intent *Intent) error {
|
||||
var lastErr error
|
||||
|
||||
for _, builder := range csb.builders {
|
||||
if err := builder.Validate(intent); err == nil {
|
||||
return nil
|
||||
} else {
|
||||
lastErr = err
|
||||
}
|
||||
}
|
||||
|
||||
if lastErr != nil {
|
||||
return lastErr
|
||||
}
|
||||
return fmt.Errorf("no builder could validate intent")
|
||||
}
|
||||
|
||||
// FromIntent tries each builder
|
||||
func (csb *CompositeSpecBuilder) FromIntent(intent *Intent, metadata *BuildMetadata) (*WorkflowSpec, error) {
|
||||
var lastErr error
|
||||
|
||||
for _, builder := range csb.builders {
|
||||
if spec, err := builder.FromIntent(intent, metadata); err == nil {
|
||||
return spec, nil
|
||||
} else {
|
||||
lastErr = err
|
||||
}
|
||||
}
|
||||
|
||||
if lastErr != nil {
|
||||
return nil, lastErr
|
||||
}
|
||||
return nil, fmt.Errorf("no builder could create spec")
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
package routing
|
||||
|
||||
import "time"
|
||||
|
||||
// WorkflowSpec is generated by llm-router (one-time execution)
|
||||
type WorkflowSpec struct {
|
||||
Name string `json:"name"`
|
||||
Input map[string]interface{} `json:"input"`
|
||||
States []State `json:"states"`
|
||||
}
|
||||
|
||||
// CronWorkflowSpec is generated by llm-router (scheduled execution)
|
||||
type CronWorkflowSpec struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"` // "CronWorkflow"
|
||||
Schedule string `json:"schedule"` // Cron expression (e.g., "0 2 * * *")
|
||||
Timezone string `json:"timezone"` // "UTC", "America/New_York", etc
|
||||
Input map[string]interface{} `json:"input"` // Fixed input for each run
|
||||
States []State `json:"states"` // Workflow states
|
||||
MaxConcurrent int `json:"maxConcurrent,omitempty"` // Max parallel runs (default 1)
|
||||
Timeout string `json:"timeout,omitempty"` // Overall timeout per run
|
||||
EnableHistory bool `json:"enableHistory,omitempty"` // Keep execution history
|
||||
}
|
||||
|
||||
// State is a step in the workflow
|
||||
type State struct {
|
||||
Name string
|
||||
Type StateType `json:"type"`
|
||||
|
||||
// Task fields
|
||||
Resource string `json:"resource,omitempty"`
|
||||
Parameters map[string]interface{} `json:"parameters,omitempty"`
|
||||
Timeout string `json:"timeout,omitempty"`
|
||||
Retry *RetryPolicy `json:"retry,omitempty"`
|
||||
Catch []CatchClause `json:"catch,omitempty"`
|
||||
|
||||
// Pass fields
|
||||
Result interface{} `json:"result,omitempty"`
|
||||
|
||||
// Fail fields
|
||||
Error string `json:"error,omitempty"`
|
||||
Cause string `json:"cause,omitempty"`
|
||||
|
||||
// Transition
|
||||
Next string `json:"next,omitempty"`
|
||||
End bool `json:"end,omitempty"`
|
||||
}
|
||||
|
||||
// StateType defines valid state types
|
||||
type StateType string
|
||||
|
||||
const (
|
||||
StateTypeTask StateType = "Task"
|
||||
StateTypePass StateType = "Pass"
|
||||
StateTypeFail StateType = "Fail"
|
||||
)
|
||||
|
||||
// RetryPolicy defines retry behavior for activities
|
||||
type RetryPolicy struct {
|
||||
MaxAttempts int32 `json:"maxAttempts"`
|
||||
BackoffRate float64 `json:"backoffRate"`
|
||||
InitialInterval string `json:"initialInterval"`
|
||||
MaxInterval string `json:"maxInterval,omitempty"`
|
||||
}
|
||||
|
||||
// CatchClause defines error handling
|
||||
type CatchClause struct {
|
||||
ErrorEquals []string `json:"errorEquals"`
|
||||
ResultPath *string `json:"resultPath,omitempty"`
|
||||
Next string `json:"next"`
|
||||
}
|
||||
|
||||
// ExecutionContext tracks state during workflow execution
|
||||
type ExecutionContext struct {
|
||||
Input map[string]interface{}
|
||||
StepResults map[string]interface{}
|
||||
CurrentState string
|
||||
History []ExecutionEvent
|
||||
}
|
||||
|
||||
// ExecutionEvent tracks individual state execution
|
||||
type ExecutionEvent struct {
|
||||
Timestamp time.Time
|
||||
State string
|
||||
Type string // "Started", "Completed", "Failed", "Retried"
|
||||
Result interface{}
|
||||
Error error
|
||||
}
|
||||
|
||||
// PollParams for AwaitTaskComplete states
|
||||
type PollParams struct {
|
||||
QueueName string
|
||||
CorrelationID string
|
||||
PollInterval time.Duration
|
||||
Timeout time.Duration
|
||||
}
|
||||
|
||||
// PollResult is the result of polling
|
||||
type PollResult struct {
|
||||
Result interface{}
|
||||
Status string
|
||||
}
|
||||
|
||||
// Result is the final workflow output
|
||||
type Result struct {
|
||||
FinalOutput interface{}
|
||||
Status string // "COMPLETED", "FAILED"
|
||||
Error error
|
||||
}
|
||||
|
||||
// Heartbeat contains state for polling activities
|
||||
type Heartbeat struct {
|
||||
CorrelationID string
|
||||
Queue string
|
||||
Attempt int
|
||||
Elapsed time.Duration
|
||||
LastCheck time.Time
|
||||
}
|
||||
|
||||
// ActivityMetadata describes an activity's capabilities and constraints
|
||||
type ActivityMetadata struct {
|
||||
Name string
|
||||
Description string
|
||||
Category string
|
||||
Inputs map[string]InputField
|
||||
Outputs map[string]OutputField
|
||||
Constraints Constraints
|
||||
}
|
||||
|
||||
// InputField describes an activity input parameter
|
||||
type InputField struct {
|
||||
Type string `json:"type"`
|
||||
Description string `json:"description"`
|
||||
Required bool `json:"required"`
|
||||
Default interface{} `json:"default,omitempty"`
|
||||
}
|
||||
|
||||
// OutputField describes an activity output field
|
||||
type OutputField struct {
|
||||
Type string `json:"type"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
// Constraints describes activity execution constraints
|
||||
type Constraints struct {
|
||||
DefaultTimeout string
|
||||
IsFlaky bool
|
||||
RecommendedRetries int
|
||||
RetryBackoff float64
|
||||
Dependencies []string
|
||||
Notes string
|
||||
}
|
||||
@@ -0,0 +1,327 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestWorkflowSpecMarshal(t *testing.T) {
|
||||
spec := WorkflowSpec{
|
||||
Name: "test-workflow",
|
||||
Input: map[string]interface{}{
|
||||
"repo": "https://github.com/test/repo",
|
||||
},
|
||||
States: []State{
|
||||
{
|
||||
Name: "Clone",
|
||||
Type: StateTypeTask,
|
||||
Resource: "CloneRepoActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"repo": "${input.repo}",
|
||||
},
|
||||
Timeout: "5m",
|
||||
Next: "Analyze",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Marshal to JSON
|
||||
data, err := json.Marshal(spec)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to marshal: %v", err)
|
||||
}
|
||||
|
||||
// Unmarshal back
|
||||
var spec2 WorkflowSpec
|
||||
err = json.Unmarshal(data, &spec2)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unmarshal: %v", err)
|
||||
}
|
||||
|
||||
// Verify
|
||||
if spec2.Name != spec.Name {
|
||||
t.Errorf("Name mismatch: %s != %s", spec2.Name, spec.Name)
|
||||
}
|
||||
if len(spec2.States) != len(spec.States) {
|
||||
t.Errorf("State count mismatch: %d != %d", len(spec2.States), len(spec.States))
|
||||
}
|
||||
}
|
||||
|
||||
func TestCronWorkflowSpecMarshal(t *testing.T) {
|
||||
spec := CronWorkflowSpec{
|
||||
Name: "daily-scan",
|
||||
Type: "CronWorkflow",
|
||||
Schedule: "0 2 * * *",
|
||||
Timezone: "UTC",
|
||||
MaxConcurrent: 1,
|
||||
EnableHistory: true,
|
||||
Input: map[string]interface{}{
|
||||
"repos": []string{"repo1", "repo2"},
|
||||
},
|
||||
}
|
||||
|
||||
// Marshal to JSON
|
||||
data, err := json.Marshal(spec)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to marshal: %v", err)
|
||||
}
|
||||
|
||||
// Unmarshal back
|
||||
var spec2 CronWorkflowSpec
|
||||
err = json.Unmarshal(data, &spec2)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unmarshal: %v", err)
|
||||
}
|
||||
|
||||
// Verify
|
||||
if spec2.Schedule != spec.Schedule {
|
||||
t.Errorf("Schedule mismatch: %s != %s", spec2.Schedule, spec.Schedule)
|
||||
}
|
||||
if spec2.Timezone != spec.Timezone {
|
||||
t.Errorf("Timezone mismatch: %s != %s", spec2.Timezone, spec.Timezone)
|
||||
}
|
||||
if spec2.EnableHistory != spec.EnableHistory {
|
||||
t.Errorf("EnableHistory mismatch: %v != %v", spec2.EnableHistory, spec.EnableHistory)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetryPolicyMarshal(t *testing.T) {
|
||||
policy := RetryPolicy{
|
||||
MaxAttempts: 3,
|
||||
BackoffRate: 2.0,
|
||||
InitialInterval: "1s",
|
||||
MaxInterval: "1m",
|
||||
}
|
||||
|
||||
data, err := json.Marshal(policy)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to marshal: %v", err)
|
||||
}
|
||||
|
||||
var policy2 RetryPolicy
|
||||
err = json.Unmarshal(data, &policy2)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if policy2.MaxAttempts != policy.MaxAttempts {
|
||||
t.Errorf("MaxAttempts mismatch: %d != %d", policy2.MaxAttempts, policy.MaxAttempts)
|
||||
}
|
||||
if policy2.BackoffRate != policy.BackoffRate {
|
||||
t.Errorf("BackoffRate mismatch: %f != %f", policy2.BackoffRate, policy.BackoffRate)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStateMarshal(t *testing.T) {
|
||||
state := State{
|
||||
Name: "Analyze",
|
||||
Type: StateTypeTask,
|
||||
Resource: "AnalyzeCodeActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"path": "${Clone.output.path}",
|
||||
},
|
||||
Timeout: "10m",
|
||||
Retry: &RetryPolicy{
|
||||
MaxAttempts: 3,
|
||||
BackoffRate: 2.0,
|
||||
InitialInterval: "1s",
|
||||
},
|
||||
Catch: []CatchClause{
|
||||
{
|
||||
ErrorEquals: []string{"Timeout"},
|
||||
Next: "HandleTimeout",
|
||||
},
|
||||
},
|
||||
Next: "Judge",
|
||||
}
|
||||
|
||||
data, err := json.Marshal(state)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to marshal: %v", err)
|
||||
}
|
||||
|
||||
var state2 State
|
||||
err = json.Unmarshal(data, &state2)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if state2.Name != state.Name {
|
||||
t.Errorf("Name mismatch: %s != %s", state2.Name, state.Name)
|
||||
}
|
||||
if state2.Type != state.Type {
|
||||
t.Errorf("Type mismatch: %s != %s", state2.Type, state.Type)
|
||||
}
|
||||
if len(state2.Catch) != len(state.Catch) {
|
||||
t.Errorf("Catch count mismatch: %d != %d", len(state2.Catch), len(state.Catch))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPassState(t *testing.T) {
|
||||
state := State{
|
||||
Name: "SetSuccess",
|
||||
Type: StateTypePass,
|
||||
Result: map[string]interface{}{"status": "success"},
|
||||
End: true,
|
||||
}
|
||||
|
||||
data, err := json.Marshal(state)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to marshal: %v", err)
|
||||
}
|
||||
|
||||
var state2 State
|
||||
err = json.Unmarshal(data, &state2)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if state2.Type != StateTypePass {
|
||||
t.Errorf("Type should be Pass, got: %s", state2.Type)
|
||||
}
|
||||
if !state2.End {
|
||||
t.Error("End should be true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFailState(t *testing.T) {
|
||||
state := State{
|
||||
Name: "HandleError",
|
||||
Type: StateTypeFail,
|
||||
Error: "InvalidInput",
|
||||
Cause: "Repository URL is invalid",
|
||||
}
|
||||
|
||||
data, err := json.Marshal(state)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to marshal: %v", err)
|
||||
}
|
||||
|
||||
var state2 State
|
||||
err = json.Unmarshal(data, &state2)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if state2.Type != StateTypeFail {
|
||||
t.Errorf("Type should be Fail, got: %s", state2.Type)
|
||||
}
|
||||
if state2.Error != "InvalidInput" {
|
||||
t.Errorf("Error mismatch: %s != InvalidInput", state2.Error)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecutionContextInit(t *testing.T) {
|
||||
ec := &ExecutionContext{
|
||||
Input: map[string]interface{}{"repo": "test"},
|
||||
StepResults: make(map[string]interface{}),
|
||||
History: make([]ExecutionEvent, 0),
|
||||
}
|
||||
|
||||
if ec.Input == nil {
|
||||
t.Error("Input should not be nil")
|
||||
}
|
||||
if ec.StepResults == nil {
|
||||
t.Error("StepResults should not be nil")
|
||||
}
|
||||
if ec.History == nil {
|
||||
t.Error("History should not be nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestComplexWorkflowSpec(t *testing.T) {
|
||||
// Test a realistic workflow spec
|
||||
spec := WorkflowSpec{
|
||||
Name: "code-review",
|
||||
Input: map[string]interface{}{
|
||||
"repo": "https://github.com/rockliang/poimen",
|
||||
"branch": "feature/x",
|
||||
},
|
||||
States: []State{
|
||||
{
|
||||
Name: "Clone",
|
||||
Type: StateTypeTask,
|
||||
Resource: "CloneRepoActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"repo": "${input.repo}",
|
||||
"branch": "${input.branch}",
|
||||
},
|
||||
Timeout: "5m",
|
||||
Retry: &RetryPolicy{
|
||||
MaxAttempts: 2,
|
||||
BackoffRate: 1.5,
|
||||
InitialInterval: "1s",
|
||||
},
|
||||
Next: "Analyze",
|
||||
},
|
||||
{
|
||||
Name: "Analyze",
|
||||
Type: StateTypeTask,
|
||||
Resource: "AnalyzeCodeActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"path": "${Clone.output.path}",
|
||||
},
|
||||
Timeout: "10m",
|
||||
Retry: &RetryPolicy{
|
||||
MaxAttempts: 3,
|
||||
BackoffRate: 2.0,
|
||||
InitialInterval: "1s",
|
||||
},
|
||||
Catch: []CatchClause{
|
||||
{
|
||||
ErrorEquals: []string{"Timeout"},
|
||||
Next: "HandleTimeout",
|
||||
},
|
||||
},
|
||||
Next: "Judge",
|
||||
},
|
||||
{
|
||||
Name: "Judge",
|
||||
Type: StateTypeTask,
|
||||
Resource: "JudgeActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"quality": "${Analyze.output.quality}",
|
||||
},
|
||||
Timeout: "5m",
|
||||
End: true,
|
||||
},
|
||||
{
|
||||
Name: "HandleTimeout",
|
||||
Type: StateTypeFail,
|
||||
Error: "AnalysisTimeout",
|
||||
Cause: "Code analysis timed out",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Marshal
|
||||
data, err := json.Marshal(spec)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to marshal: %v", err)
|
||||
}
|
||||
|
||||
// Unmarshal
|
||||
var spec2 WorkflowSpec
|
||||
err = json.Unmarshal(data, &spec2)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unmarshal: %v", err)
|
||||
}
|
||||
|
||||
// Verify
|
||||
if spec2.Name != "code-review" {
|
||||
t.Errorf("Name mismatch")
|
||||
}
|
||||
if len(spec2.States) != 4 {
|
||||
t.Errorf("Expected 4 states, got %d", len(spec2.States))
|
||||
}
|
||||
|
||||
// Check first state
|
||||
if spec2.States[0].Resource != "CloneRepoActivity" {
|
||||
t.Errorf("First state resource mismatch")
|
||||
}
|
||||
|
||||
// Check error catching
|
||||
if len(spec2.States[1].Catch) == 0 {
|
||||
t.Error("Analyze state should have catch clauses")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,346 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ValidationError represents a single validation error
|
||||
type ValidationError struct {
|
||||
Path string // JSONPath where error occurred
|
||||
Message string // Error message
|
||||
}
|
||||
|
||||
// ValidationResult contains all validation errors
|
||||
type ValidationResult struct {
|
||||
Valid bool
|
||||
Errors []ValidationError
|
||||
}
|
||||
|
||||
// Validator validates WorkflowSpec and CronWorkflowSpec
|
||||
type Validator struct {
|
||||
kb *KnowledgeBase
|
||||
}
|
||||
|
||||
// NewValidator creates a new validator with knowledge base
|
||||
func NewValidator(kb *KnowledgeBase) *Validator {
|
||||
return &Validator{kb: kb}
|
||||
}
|
||||
|
||||
// ValidateWorkflowSpec validates a one-time workflow spec
|
||||
func (v *Validator) ValidateWorkflowSpec(spec *WorkflowSpec) *ValidationResult {
|
||||
result := &ValidationResult{
|
||||
Valid: true,
|
||||
Errors: []ValidationError{},
|
||||
}
|
||||
|
||||
if spec == nil {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec",
|
||||
Message: "workflow spec cannot be nil",
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// Validate name
|
||||
var errs []ValidationError
|
||||
if spec.Name == "" {
|
||||
result.Valid = false
|
||||
errs = append(errs, ValidationError{
|
||||
Path: "spec.name",
|
||||
Message: "workflow name is required",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate states
|
||||
if len(spec.States) == 0 {
|
||||
result.Valid = false
|
||||
errs = append(errs, ValidationError{
|
||||
Path: "spec.states",
|
||||
Message: "at least one state is required",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate each state
|
||||
stateNames := make(map[string]bool)
|
||||
|
||||
for i, state := range spec.States {
|
||||
path := fmt.Sprintf("spec.states[%d]", i)
|
||||
|
||||
if state.Name == "" {
|
||||
result.Valid = false
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + ".name",
|
||||
Message: "state name is required",
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
if stateNames[state.Name] {
|
||||
result.Valid = false
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + ".name",
|
||||
Message: fmt.Sprintf("duplicate state name: %s", state.Name),
|
||||
})
|
||||
}
|
||||
stateNames[state.Name] = true
|
||||
|
||||
// Validate state type
|
||||
stateErrs := v.validateState(state, path)
|
||||
if len(stateErrs) > 0 {
|
||||
result.Valid = false
|
||||
errs = append(errs, stateErrs...)
|
||||
}
|
||||
}
|
||||
|
||||
// Validate state transitions
|
||||
for i, state := range spec.States {
|
||||
path := fmt.Sprintf("spec.states[%d]", i)
|
||||
|
||||
// Check that Next state exists (if specified)
|
||||
if state.Next != "" && !stateNames[state.Next] {
|
||||
result.Valid = false
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + ".next",
|
||||
Message: fmt.Sprintf("state '%s' does not exist", state.Next),
|
||||
})
|
||||
}
|
||||
|
||||
// Check that Catch targets exist
|
||||
for j, catchClause := range state.Catch {
|
||||
if catchClause.Next != "" && !stateNames[catchClause.Next] {
|
||||
result.Valid = false
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + fmt.Sprintf(".catch[%d].next", j),
|
||||
Message: fmt.Sprintf("state '%s' does not exist", catchClause.Next),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result.Errors = errs
|
||||
return result
|
||||
}
|
||||
|
||||
// ValidateCronWorkflowSpec validates a scheduled workflow spec
|
||||
func (v *Validator) ValidateCronWorkflowSpec(spec *CronWorkflowSpec) *ValidationResult {
|
||||
result := &ValidationResult{
|
||||
Valid: true,
|
||||
Errors: []ValidationError{},
|
||||
}
|
||||
|
||||
if spec == nil {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec",
|
||||
Message: "cron workflow spec cannot be nil",
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// Validate name
|
||||
if spec.Name == "" {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec.name",
|
||||
Message: "workflow name is required",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate type
|
||||
if spec.Type != "CronWorkflow" {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec.type",
|
||||
Message: "type must be 'CronWorkflow'",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate cron expression
|
||||
if spec.Schedule == "" {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec.schedule",
|
||||
Message: "cron schedule is required",
|
||||
})
|
||||
} else {
|
||||
if err := validateCronExpression(spec.Schedule); err != nil {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec.schedule",
|
||||
Message: err.Error(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Validate timezone
|
||||
if spec.Timezone == "" {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec.timezone",
|
||||
Message: "timezone is required",
|
||||
})
|
||||
} else {
|
||||
if _, err := time.LoadLocation(spec.Timezone); err != nil {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec.timezone",
|
||||
Message: fmt.Sprintf("invalid timezone: %s", spec.Timezone),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Validate maxConcurrent
|
||||
if spec.MaxConcurrent < 1 {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec.maxConcurrent",
|
||||
Message: "maxConcurrent must be >= 1",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate states (same as one-time workflow)
|
||||
if len(spec.States) == 0 {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: "spec.states",
|
||||
Message: "at least one state is required",
|
||||
})
|
||||
}
|
||||
|
||||
stateNames := make(map[string]bool)
|
||||
for i, state := range spec.States {
|
||||
path := fmt.Sprintf("spec.states[%d]", i)
|
||||
|
||||
if state.Name == "" {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: path + ".name",
|
||||
Message: "state name is required",
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
if stateNames[state.Name] {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: path + ".name",
|
||||
Message: fmt.Sprintf("duplicate state name: %s", state.Name),
|
||||
})
|
||||
}
|
||||
stateNames[state.Name] = true
|
||||
|
||||
stateErrs := v.validateState(state, path)
|
||||
if len(stateErrs) > 0 {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, stateErrs...)
|
||||
}
|
||||
}
|
||||
|
||||
// Validate state transitions
|
||||
for i, state := range spec.States {
|
||||
path := fmt.Sprintf("spec.states[%d]", i)
|
||||
|
||||
if state.Next != "" && !stateNames[state.Next] {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: path + ".next",
|
||||
Message: fmt.Sprintf("state '%s' does not exist", state.Next),
|
||||
})
|
||||
}
|
||||
|
||||
for j, catchClause := range state.Catch {
|
||||
if catchClause.Next != "" && !stateNames[catchClause.Next] {
|
||||
result.Valid = false
|
||||
result.Errors = append(result.Errors, ValidationError{
|
||||
Path: path + fmt.Sprintf(".catch[%d].next", j),
|
||||
Message: fmt.Sprintf("state '%s' does not exist", catchClause.Next),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// validateState validates a single state
|
||||
func (v *Validator) validateState(state State, path string) []ValidationError {
|
||||
var errs []ValidationError
|
||||
|
||||
switch state.Type {
|
||||
case StateTypeTask:
|
||||
if state.Resource == "" {
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + ".resource",
|
||||
Message: "resource is required for Task state",
|
||||
})
|
||||
} else {
|
||||
// Check if activity exists in knowledge base
|
||||
if v.kb != nil && !v.kb.HasActivity(state.Resource) {
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + ".resource",
|
||||
Message: fmt.Sprintf("activity '%s' not found in knowledge base", state.Resource),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Validate timeout format
|
||||
if state.Timeout != "" {
|
||||
if err := validateDuration(state.Timeout); err != nil {
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + ".timeout",
|
||||
Message: err.Error(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
case StateTypePass:
|
||||
// Pass state is valid with just a result
|
||||
if state.Result == nil {
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + ".result",
|
||||
Message: "result is required for Pass state",
|
||||
})
|
||||
}
|
||||
|
||||
case StateTypeFail:
|
||||
// Fail state requires error
|
||||
if state.Error == "" {
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + ".error",
|
||||
Message: "error is required for Fail state",
|
||||
})
|
||||
}
|
||||
|
||||
default:
|
||||
errs = append(errs, ValidationError{
|
||||
Path: path + ".type",
|
||||
Message: fmt.Sprintf("invalid state type: %s", state.Type),
|
||||
})
|
||||
}
|
||||
|
||||
return errs
|
||||
}
|
||||
|
||||
// validateDuration validates Go duration string (e.g., "5m", "30s")
|
||||
func validateDuration(dur string) error {
|
||||
_, err := time.ParseDuration(dur)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid duration: %s", dur)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns a human-readable validation result
|
||||
func (vr *ValidationResult) String() string {
|
||||
if vr.Valid {
|
||||
return "Valid ✓"
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("Invalid ✗ (%d errors):\n", len(vr.Errors))
|
||||
for _, err := range vr.Errors {
|
||||
msg += fmt.Sprintf(" %s: %s\n", err.Path, err.Message)
|
||||
}
|
||||
return msg
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// validateCronExpression validates a cron expression (simplified)
|
||||
// Supports standard 5-field cron: minute hour day month weekday
|
||||
// Does NOT validate all possible edge cases - just basic format
|
||||
func validateCronExpression(expr string) error {
|
||||
fields := strings.Fields(expr)
|
||||
if len(fields) != 5 {
|
||||
return fmt.Errorf("cron expression must have 5 fields (minute hour day month weekday), got %d", len(fields))
|
||||
}
|
||||
|
||||
// Validate field ranges
|
||||
ranges := []struct {
|
||||
name string
|
||||
min int
|
||||
max int
|
||||
}{
|
||||
{"minute", 0, 59},
|
||||
{"hour", 0, 23},
|
||||
{"day", 1, 31},
|
||||
{"month", 1, 12},
|
||||
{"weekday", 0, 6},
|
||||
}
|
||||
|
||||
// Basic pattern: * or */n or n or n-m or n,m or n-m/p
|
||||
// This is simplified and doesn't validate all edge cases
|
||||
fieldRegex := regexp.MustCompile(`^(\*|(\d+)(,(\d+))*(\/\d+)?|(\d+)-(\d+)(\/\d+)?|\*\/\d+)$`)
|
||||
|
||||
for i, field := range fields {
|
||||
if field == "*" {
|
||||
continue
|
||||
}
|
||||
|
||||
// Check basic format
|
||||
if !fieldRegex.MatchString(field) {
|
||||
return fmt.Errorf("invalid %s field: %s", ranges[i].name, field)
|
||||
}
|
||||
|
||||
// Validate simple number values
|
||||
if !strings.ContainsAny(field, "*,-/") {
|
||||
var val int
|
||||
_, _ = fmt.Sscanf(field, "%d", &val)
|
||||
if val < ranges[i].min || val > ranges[i].max {
|
||||
return fmt.Errorf("invalid %s value %d (range %d-%d)", ranges[i].name, val, ranges[i].min, ranges[i].max)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,540 @@
|
||||
package routing
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateValidWorkflowSpec(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "test-workflow",
|
||||
Input: map[string]interface{}{
|
||||
"repo": "https://github.com/test/repo",
|
||||
},
|
||||
States: []State{
|
||||
{
|
||||
Name: "Clone",
|
||||
Type: StateTypeTask,
|
||||
Resource: "CloneRepoActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"repo": "${input.repo}",
|
||||
},
|
||||
Timeout: "5m",
|
||||
Next: "Analyze",
|
||||
},
|
||||
{
|
||||
Name: "Analyze",
|
||||
Type: StateTypeTask,
|
||||
Resource: "AnalyzeCodeActivity",
|
||||
Parameters: map[string]interface{}{
|
||||
"path": "${Clone.output.path}",
|
||||
},
|
||||
Timeout: "10m",
|
||||
End: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if !result.Valid {
|
||||
t.Errorf("Spec should be valid. Errors: %v", result.Errors)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateWorkflowSpecNilName(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "",
|
||||
States: []State{
|
||||
{
|
||||
Name: "Test",
|
||||
Type: StateTypeTask,
|
||||
Resource: "CloneRepoActivity",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Spec with empty name should be invalid")
|
||||
}
|
||||
|
||||
if len(result.Errors) == 0 {
|
||||
t.Error("Should have validation errors")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateWorkflowSpecNoStates(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "test",
|
||||
States: []State{},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Spec with no states should be invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateInvalidStateTransition(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "test",
|
||||
States: []State{
|
||||
{
|
||||
Name: "A",
|
||||
Type: StateTypeTask,
|
||||
Resource: "CloneRepoActivity",
|
||||
Next: "NonExistent",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Spec with invalid state transition should be invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateDuplicateStateName(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "test",
|
||||
States: []State{
|
||||
{
|
||||
Name: "A",
|
||||
Type: StateTypeTask,
|
||||
Resource: "CloneRepoActivity",
|
||||
},
|
||||
{
|
||||
Name: "A", // Duplicate!
|
||||
Type: StateTypeTask,
|
||||
Resource: "CloneRepoActivity",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Spec with duplicate state names should be invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateUnknownActivity(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "test",
|
||||
States: []State{
|
||||
{
|
||||
Name: "A",
|
||||
Type: StateTypeTask,
|
||||
Resource: "UnknownActivity",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Spec with unknown activity should be invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidatePassState(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "test",
|
||||
States: []State{
|
||||
{
|
||||
Name: "Success",
|
||||
Type: StateTypePass,
|
||||
Result: map[string]interface{}{"status": "ok"},
|
||||
End: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if !result.Valid {
|
||||
t.Errorf("Pass state spec should be valid: %v", result.Errors)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidatePassStateNoResult(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "test",
|
||||
States: []State{
|
||||
{
|
||||
Name: "Success",
|
||||
Type: StateTypePass,
|
||||
End: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Pass state without result should be invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateFailState(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "test",
|
||||
States: []State{
|
||||
{
|
||||
Name: "Fail",
|
||||
Type: StateTypeFail,
|
||||
Error: "TestError",
|
||||
Cause: "For testing",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if !result.Valid {
|
||||
t.Errorf("Fail state spec should be valid: %v", result.Errors)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateInvalidDuration(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &WorkflowSpec{
|
||||
Name: "test",
|
||||
States: []State{
|
||||
{
|
||||
Name: "A",
|
||||
Type: StateTypeTask,
|
||||
Resource: "CloneRepoActivity",
|
||||
Timeout: "invalid",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Spec with invalid timeout should be invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateCronWorkflowSpec(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &CronWorkflowSpec{
|
||||
Name: "daily-scan",
|
||||
Type: "CronWorkflow",
|
||||
Schedule: "0 2 * * *",
|
||||
Timezone: "UTC",
|
||||
MaxConcurrent: 1,
|
||||
States: []State{
|
||||
{
|
||||
Name: "Scan",
|
||||
Type: StateTypeTask,
|
||||
Resource: "SecurityScanActivity",
|
||||
End: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateCronWorkflowSpec(spec)
|
||||
if !result.Valid {
|
||||
t.Errorf("Valid cron spec should pass: %v", result.Errors)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateCronWorkflowSpecInvalidTimezone(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &CronWorkflowSpec{
|
||||
Name: "daily-scan",
|
||||
Type: "CronWorkflow",
|
||||
Schedule: "0 2 * * *",
|
||||
Timezone: "InvalidTimezone",
|
||||
MaxConcurrent: 1,
|
||||
States: []State{
|
||||
{
|
||||
Name: "Scan",
|
||||
Type: StateTypeTask,
|
||||
Resource: "SecurityScanActivity",
|
||||
End: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateCronWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Cron spec with invalid timezone should be invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateCronWorkflowSpecInvalidSchedule(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &CronWorkflowSpec{
|
||||
Name: "daily-scan",
|
||||
Type: "CronWorkflow",
|
||||
Schedule: "invalid cron",
|
||||
Timezone: "UTC",
|
||||
MaxConcurrent: 1,
|
||||
States: []State{
|
||||
{
|
||||
Name: "Scan",
|
||||
Type: StateTypeTask,
|
||||
Resource: "SecurityScanActivity",
|
||||
End: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateCronWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Cron spec with invalid schedule should be invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateCronWorkflowSpecInvalidType(t *testing.T) {
|
||||
kbPath := getKBPath()
|
||||
if kbPath == "" {
|
||||
t.Skip("Knowledge base not found")
|
||||
}
|
||||
|
||||
kb, err := LoadKnowledgeBase(kbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load KB: %v", err)
|
||||
}
|
||||
|
||||
validator := NewValidator(kb)
|
||||
|
||||
spec := &CronWorkflowSpec{
|
||||
Name: "daily-scan",
|
||||
Type: "WrongType",
|
||||
Schedule: "0 2 * * *",
|
||||
Timezone: "UTC",
|
||||
MaxConcurrent: 1,
|
||||
States: []State{
|
||||
{
|
||||
Name: "Scan",
|
||||
Type: StateTypeTask,
|
||||
Resource: "SecurityScanActivity",
|
||||
End: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result := validator.ValidateCronWorkflowSpec(spec)
|
||||
if result.Valid {
|
||||
t.Error("Cron spec with wrong type should be invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateDuration(t *testing.T) {
|
||||
validDurations := []string{"1s", "5m", "1h", "100ms"}
|
||||
for _, dur := range validDurations {
|
||||
if err := validateDuration(dur); err != nil {
|
||||
t.Errorf("Duration %s should be valid: %v", dur, err)
|
||||
}
|
||||
}
|
||||
|
||||
invalidDurations := []string{"invalid", "5x", ""}
|
||||
for _, dur := range invalidDurations {
|
||||
if err := validateDuration(dur); err == nil {
|
||||
t.Errorf("Duration %s should be invalid", dur)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateCronExpression(t *testing.T) {
|
||||
validCrons := []string{
|
||||
"0 2 * * *", // 2 AM daily
|
||||
"*/5 * * * *", // Every 5 minutes
|
||||
"0 0 1 * *", // First of month
|
||||
"0 12 * * 1", // Noon on Mondays
|
||||
"30 15 * * *", // 3:30 PM daily
|
||||
}
|
||||
|
||||
for _, cron := range validCrons {
|
||||
if err := validateCronExpression(cron); err != nil {
|
||||
t.Errorf("Cron %s should be valid: %v", cron, err)
|
||||
}
|
||||
}
|
||||
|
||||
invalidCrons := []string{
|
||||
"invalid", // Too few fields
|
||||
"0 2 * * * *", // Too many fields
|
||||
"60 * * * *", // Invalid minute
|
||||
"* 25 * * *", // Invalid hour
|
||||
}
|
||||
|
||||
for _, cron := range invalidCrons {
|
||||
if err := validateCronExpression(cron); err == nil {
|
||||
t.Errorf("Cron %s should be invalid", cron)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidationResultString(t *testing.T) {
|
||||
result := &ValidationResult{
|
||||
Valid: true,
|
||||
Errors: []ValidationError{},
|
||||
}
|
||||
|
||||
str := result.String()
|
||||
if str != "Valid ✓" {
|
||||
t.Errorf("Valid result string should be 'Valid ✓', got: %s", str)
|
||||
}
|
||||
|
||||
result.Valid = false
|
||||
result.Errors = []ValidationError{
|
||||
{
|
||||
Path: "spec.name",
|
||||
Message: "name is required",
|
||||
},
|
||||
}
|
||||
|
||||
str = result.String()
|
||||
if !contains(str, "Invalid") {
|
||||
t.Error("Invalid result string should contain 'Invalid'")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,234 @@
|
||||
package search
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// WorkflowEntry represents an indexed workflow
|
||||
type WorkflowEntry struct {
|
||||
ID string
|
||||
Name string
|
||||
Status string
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
Tags []string
|
||||
Content string
|
||||
Assignee string
|
||||
}
|
||||
|
||||
// WorkflowSearch provides full-text search and filtering
|
||||
type WorkflowSearch struct {
|
||||
mu sync.RWMutex
|
||||
entries map[string]*WorkflowEntry
|
||||
index map[string][]string // word -> workflow IDs
|
||||
filters map[string]interface{}
|
||||
}
|
||||
|
||||
// NewWorkflowSearch creates a new workflow search index
|
||||
func NewWorkflowSearch() *WorkflowSearch {
|
||||
return &WorkflowSearch{
|
||||
entries: make(map[string]*WorkflowEntry),
|
||||
index: make(map[string][]string),
|
||||
filters: make(map[string]interface{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Index adds a workflow to the search index
|
||||
func (ws *WorkflowSearch) Index(entry *WorkflowEntry) error {
|
||||
if entry.ID == "" {
|
||||
return fmt.Errorf("workflow ID required")
|
||||
}
|
||||
|
||||
ws.mu.Lock()
|
||||
defer ws.mu.Unlock()
|
||||
|
||||
ws.entries[entry.ID] = entry
|
||||
|
||||
// Index content
|
||||
words := strings.Fields(strings.ToLower(entry.Content + " " + entry.Name))
|
||||
for _, word := range words {
|
||||
// Remove punctuation
|
||||
clean := strings.Trim(word, ".,!?;:")
|
||||
if clean != "" {
|
||||
ws.index[clean] = append(ws.index[clean], entry.ID)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Search performs full-text search
|
||||
func (ws *WorkflowSearch) Search(query string) []*WorkflowEntry {
|
||||
ws.mu.RLock()
|
||||
defer ws.mu.RUnlock()
|
||||
|
||||
query = strings.ToLower(query)
|
||||
matches := make(map[string]int)
|
||||
|
||||
words := strings.Fields(query)
|
||||
for _, word := range words {
|
||||
if ids, exists := ws.index[word]; exists {
|
||||
for _, id := range ids {
|
||||
matches[id]++
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Sort by match count
|
||||
result := make([]*WorkflowEntry, 0)
|
||||
for id := range matches {
|
||||
if entry, exists := ws.entries[id]; exists {
|
||||
result = append(result, entry)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// FilterByStatus filters workflows by status
|
||||
func (ws *WorkflowSearch) FilterByStatus(status string) []*WorkflowEntry {
|
||||
ws.mu.RLock()
|
||||
defer ws.mu.RUnlock()
|
||||
|
||||
result := make([]*WorkflowEntry, 0)
|
||||
for _, entry := range ws.entries {
|
||||
if entry.Status == status {
|
||||
result = append(result, entry)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// FilterByAssignee filters workflows by assignee
|
||||
func (ws *WorkflowSearch) FilterByAssignee(assignee string) []*WorkflowEntry {
|
||||
ws.mu.RLock()
|
||||
defer ws.mu.RUnlock()
|
||||
|
||||
result := make([]*WorkflowEntry, 0)
|
||||
for _, entry := range ws.entries {
|
||||
if entry.Assignee == assignee {
|
||||
result = append(result, entry)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// FilterByTag filters workflows by tag
|
||||
func (ws *WorkflowSearch) FilterByTag(tag string) []*WorkflowEntry {
|
||||
ws.mu.RLock()
|
||||
defer ws.mu.RUnlock()
|
||||
|
||||
result := make([]*WorkflowEntry, 0)
|
||||
for _, entry := range ws.entries {
|
||||
for _, t := range entry.Tags {
|
||||
if t == tag {
|
||||
result = append(result, entry)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// FilterByDateRange filters workflows by date range
|
||||
func (ws *WorkflowSearch) FilterByDateRange(start, end time.Time) []*WorkflowEntry {
|
||||
ws.mu.RLock()
|
||||
defer ws.mu.RUnlock()
|
||||
|
||||
result := make([]*WorkflowEntry, 0)
|
||||
for _, entry := range ws.entries {
|
||||
if entry.CreatedAt.After(start) && entry.CreatedAt.Before(end) {
|
||||
result = append(result, entry)
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// SearchRegex performs regex search on content
|
||||
func (ws *WorkflowSearch) SearchRegex(pattern string) ([]*WorkflowEntry, error) {
|
||||
re, err := regexp.Compile(pattern)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ws.mu.RLock()
|
||||
defer ws.mu.RUnlock()
|
||||
|
||||
result := make([]*WorkflowEntry, 0)
|
||||
for _, entry := range ws.entries {
|
||||
if re.MatchString(entry.Content) || re.MatchString(entry.Name) {
|
||||
result = append(result, entry)
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// SaveFilter saves a named filter
|
||||
func (ws *WorkflowSearch) SaveFilter(name string, filter interface{}) {
|
||||
ws.mu.Lock()
|
||||
defer ws.mu.Unlock()
|
||||
|
||||
ws.filters[name] = filter
|
||||
}
|
||||
|
||||
// GetFilter retrieves a saved filter
|
||||
func (ws *WorkflowSearch) GetFilter(name string) (interface{}, bool) {
|
||||
ws.mu.RLock()
|
||||
defer ws.mu.RUnlock()
|
||||
|
||||
filter, exists := ws.filters[name]
|
||||
return filter, exists
|
||||
}
|
||||
|
||||
// GetAll returns all workflows
|
||||
func (ws *WorkflowSearch) GetAll() []*WorkflowEntry {
|
||||
ws.mu.RLock()
|
||||
defer ws.mu.RUnlock()
|
||||
|
||||
result := make([]*WorkflowEntry, 0, len(ws.entries))
|
||||
for _, entry := range ws.entries {
|
||||
result = append(result, entry)
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetByID retrieves a workflow by ID
|
||||
func (ws *WorkflowSearch) GetByID(id string) (*WorkflowEntry, bool) {
|
||||
ws.mu.RLock()
|
||||
defer ws.mu.RUnlock()
|
||||
|
||||
entry, exists := ws.entries[id]
|
||||
return entry, exists
|
||||
}
|
||||
|
||||
// Delete removes a workflow from the index
|
||||
func (ws *WorkflowSearch) Delete(id string) error {
|
||||
ws.mu.Lock()
|
||||
defer ws.mu.Unlock()
|
||||
|
||||
if _, exists := ws.entries[id]; !exists {
|
||||
return fmt.Errorf("workflow not found: %s", id)
|
||||
}
|
||||
|
||||
delete(ws.entries, id)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Clear clears the entire index
|
||||
func (ws *WorkflowSearch) Clear() {
|
||||
ws.mu.Lock()
|
||||
defer ws.mu.Unlock()
|
||||
|
||||
ws.entries = make(map[string]*WorkflowEntry)
|
||||
ws.index = make(map[string][]string)
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
package search
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestIndex(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
entry := &WorkflowEntry{
|
||||
ID: "wf-1",
|
||||
Name: "Deploy Service",
|
||||
Status: "completed",
|
||||
Content: "deployment task",
|
||||
}
|
||||
|
||||
err := ws.Index(entry)
|
||||
assert.NoError(t, err)
|
||||
|
||||
retrieved, exists := ws.GetByID("wf-1")
|
||||
assert.True(t, exists)
|
||||
assert.Equal(t, "Deploy Service", retrieved.Name)
|
||||
}
|
||||
|
||||
func TestSearch(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-1",
|
||||
Name: "Deploy Service",
|
||||
Content: "deployment production",
|
||||
})
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-2",
|
||||
Name: "Build Docker",
|
||||
Content: "docker image",
|
||||
})
|
||||
|
||||
results := ws.Search("deployment")
|
||||
assert.Equal(t, 1, len(results))
|
||||
assert.Equal(t, "wf-1", results[0].ID)
|
||||
}
|
||||
|
||||
func TestFilterByStatus(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{ID: "wf-1", Status: "completed"})
|
||||
ws.Index(&WorkflowEntry{ID: "wf-2", Status: "running"})
|
||||
ws.Index(&WorkflowEntry{ID: "wf-3", Status: "completed"})
|
||||
|
||||
results := ws.FilterByStatus("completed")
|
||||
assert.Equal(t, 2, len(results))
|
||||
}
|
||||
|
||||
func TestFilterByAssignee(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{ID: "wf-1", Assignee: "alice"})
|
||||
ws.Index(&WorkflowEntry{ID: "wf-2", Assignee: "bob"})
|
||||
ws.Index(&WorkflowEntry{ID: "wf-3", Assignee: "alice"})
|
||||
|
||||
results := ws.FilterByAssignee("alice")
|
||||
assert.Equal(t, 2, len(results))
|
||||
}
|
||||
|
||||
func TestFilterByTag(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-1",
|
||||
Tags: []string{"production", "critical"},
|
||||
})
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-2",
|
||||
Tags: []string{"staging"},
|
||||
})
|
||||
|
||||
results := ws.FilterByTag("production")
|
||||
assert.Equal(t, 1, len(results))
|
||||
}
|
||||
|
||||
func TestFilterByDateRange(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
now := time.Now()
|
||||
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-1",
|
||||
CreatedAt: now.Add(-1 * time.Hour),
|
||||
})
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-2",
|
||||
CreatedAt: now.Add(-24 * time.Hour),
|
||||
})
|
||||
|
||||
results := ws.FilterByDateRange(now.Add(-2*time.Hour), now)
|
||||
assert.Equal(t, 1, len(results))
|
||||
}
|
||||
|
||||
func TestSearchRegex(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-1",
|
||||
Content: "error 404 not found",
|
||||
})
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-2",
|
||||
Content: "success 200 ok",
|
||||
})
|
||||
|
||||
results, err := ws.SearchRegex("error.*404")
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 1, len(results))
|
||||
}
|
||||
|
||||
func TestSaveAndGetFilter(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
filter := map[string]interface{}{"status": "completed"}
|
||||
ws.SaveFilter("completed-only", filter)
|
||||
|
||||
retrieved, exists := ws.GetFilter("completed-only")
|
||||
assert.True(t, exists)
|
||||
assert.NotNil(t, retrieved)
|
||||
}
|
||||
|
||||
func TestGetAll(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{ID: "wf-1"})
|
||||
ws.Index(&WorkflowEntry{ID: "wf-2"})
|
||||
ws.Index(&WorkflowEntry{ID: "wf-3"})
|
||||
|
||||
all := ws.GetAll()
|
||||
assert.Equal(t, 3, len(all))
|
||||
}
|
||||
|
||||
func TestDelete(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{ID: "wf-1"})
|
||||
ws.Delete("wf-1")
|
||||
|
||||
_, exists := ws.GetByID("wf-1")
|
||||
assert.False(t, exists)
|
||||
}
|
||||
|
||||
func TestClear(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{ID: "wf-1"})
|
||||
ws.Index(&WorkflowEntry{ID: "wf-2"})
|
||||
ws.Clear()
|
||||
|
||||
all := ws.GetAll()
|
||||
assert.Equal(t, 0, len(all))
|
||||
}
|
||||
|
||||
func TestMultiwordSearch(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-1",
|
||||
Content: "deploy service to production",
|
||||
})
|
||||
|
||||
results := ws.Search("deploy service")
|
||||
assert.Equal(t, 1, len(results))
|
||||
}
|
||||
|
||||
func TestCaseSensitivity(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
ws.Index(&WorkflowEntry{
|
||||
ID: "wf-1",
|
||||
Content: "Deploy Service Production",
|
||||
})
|
||||
|
||||
results := ws.Search("deploy")
|
||||
assert.Equal(t, 1, len(results))
|
||||
}
|
||||
|
||||
func TestIndexError(t *testing.T) {
|
||||
ws := NewWorkflowSearch()
|
||||
|
||||
entry := &WorkflowEntry{
|
||||
Name: "No ID",
|
||||
}
|
||||
|
||||
err := ws.Index(entry)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
@@ -0,0 +1,315 @@
|
||||
package templates
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// WorkflowTemplate defines the structure of a workflow template
|
||||
type WorkflowTemplate struct {
|
||||
// Metadata
|
||||
Name string `yaml:"name" json:"name"`
|
||||
Version string `yaml:"version" json:"version"`
|
||||
Description string `yaml:"description" json:"description"`
|
||||
Author string `yaml:"author" json:"author"`
|
||||
CreatedAt time.Time `yaml:"created_at" json:"created_at"`
|
||||
UpdatedAt time.Time `yaml:"updated_at" json:"updated_at"`
|
||||
|
||||
// Orchestrator Configuration
|
||||
Orchestrator OrchestratorConfig `yaml:"orchestrator" json:"orchestrator"`
|
||||
|
||||
// Task Configuration
|
||||
Tasks []TaskConfig `yaml:"tasks" json:"tasks"`
|
||||
|
||||
// Variable Defaults
|
||||
Defaults map[string]interface{} `yaml:"defaults" json:"defaults,omitempty"`
|
||||
|
||||
// Tags and Metadata
|
||||
Tags map[string]string `yaml:"tags" json:"tags,omitempty"`
|
||||
}
|
||||
|
||||
// OrchestratorConfig defines orchestrator settings
|
||||
type OrchestratorConfig struct {
|
||||
PlannerURL string `yaml:"planner_url" json:"planner_url"`
|
||||
JudgeURL string `yaml:"judge_url" json:"judge_url"`
|
||||
ImplementerURL string `yaml:"implementer_url" json:"implementer_url"`
|
||||
TimeoutSeconds int `yaml:"timeout_seconds" json:"timeout_seconds"`
|
||||
RetryPolicy string `yaml:"retry_policy" json:"retry_policy"`
|
||||
MaxConcurrency int `yaml:"max_concurrency" json:"max_concurrency"`
|
||||
Variables map[string]interface{} `yaml:"variables" json:"variables,omitempty"`
|
||||
}
|
||||
|
||||
// TaskConfig defines task template
|
||||
type TaskConfig struct {
|
||||
ID string `yaml:"id" json:"id"`
|
||||
Title string `yaml:"title" json:"title"`
|
||||
Description string `yaml:"description" json:"description"`
|
||||
Type string `yaml:"type" json:"type"` // feature, bugfix, refactor, etc.
|
||||
Priority string `yaml:"priority" json:"priority"` // high, medium, low
|
||||
DependsOn []string `yaml:"depends_on" json:"depends_on,omitempty"`
|
||||
Config map[string]interface{} `yaml:"config" json:"config,omitempty"`
|
||||
}
|
||||
|
||||
// WorkflowTemplateManager manages workflow templates
|
||||
type WorkflowTemplateManager struct {
|
||||
mu sync.RWMutex
|
||||
templates map[string]*WorkflowTemplate
|
||||
templatePath string
|
||||
loadedFrom map[string]string // template name -> file path
|
||||
lastModified map[string]time.Time
|
||||
stats *TemplateStats
|
||||
}
|
||||
|
||||
// TemplateStats tracks template usage
|
||||
type TemplateStats struct {
|
||||
TotalTemplates int
|
||||
LoadedTemplates int
|
||||
ExecutedTemplates int
|
||||
LastUsed map[string]time.Time
|
||||
}
|
||||
|
||||
// NewWorkflowTemplateManager creates a new template manager
|
||||
func NewWorkflowTemplateManager(templatePath string) *WorkflowTemplateManager {
|
||||
return &WorkflowTemplateManager{
|
||||
templates: make(map[string]*WorkflowTemplate),
|
||||
templatePath: templatePath,
|
||||
loadedFrom: make(map[string]string),
|
||||
lastModified: make(map[string]time.Time),
|
||||
stats: &TemplateStats{
|
||||
LastUsed: make(map[string]time.Time),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// SaveTemplate saves a template to file
|
||||
func (wtm *WorkflowTemplateManager) SaveTemplate(template *WorkflowTemplate) error {
|
||||
if template == nil {
|
||||
return fmt.Errorf("template cannot be nil")
|
||||
}
|
||||
|
||||
if template.Name == "" {
|
||||
return fmt.Errorf("template name cannot be empty")
|
||||
}
|
||||
|
||||
wtm.mu.Lock()
|
||||
defer wtm.mu.Unlock()
|
||||
|
||||
// Update timestamps
|
||||
now := time.Now()
|
||||
template.UpdatedAt = now
|
||||
if template.CreatedAt.IsZero() {
|
||||
template.CreatedAt = now
|
||||
}
|
||||
|
||||
// Create directory if needed
|
||||
if err := os.MkdirAll(wtm.templatePath, 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Save to YAML file
|
||||
filePath := filepath.Join(wtm.templatePath, template.Name+".yaml")
|
||||
data, err := yaml.Marshal(template)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := os.WriteFile(filePath, data, 0644); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Update tracking
|
||||
wtm.templates[template.Name] = template
|
||||
wtm.loadedFrom[template.Name] = filePath
|
||||
wtm.lastModified[template.Name] = now
|
||||
wtm.stats.TotalTemplates++
|
||||
wtm.stats.LoadedTemplates++
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadTemplate loads a template from file
|
||||
func (wtm *WorkflowTemplateManager) LoadTemplate(name string) (*WorkflowTemplate, error) {
|
||||
wtm.mu.Lock()
|
||||
defer wtm.mu.Unlock()
|
||||
|
||||
// Check if already loaded
|
||||
if template, exists := wtm.templates[name]; exists {
|
||||
return template, nil
|
||||
}
|
||||
|
||||
// Try to load from file
|
||||
filePath := filepath.Join(wtm.templatePath, name+".yaml")
|
||||
data, err := os.ReadFile(filePath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read template: %w", err)
|
||||
}
|
||||
|
||||
var template WorkflowTemplate
|
||||
if err := yaml.Unmarshal(data, &template); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse template: %w", err)
|
||||
}
|
||||
|
||||
// Cache the template
|
||||
wtm.templates[name] = &template
|
||||
wtm.loadedFrom[name] = filePath
|
||||
wtm.lastModified[name] = time.Now()
|
||||
wtm.stats.LoadedTemplates++
|
||||
wtm.stats.TotalTemplates++
|
||||
|
||||
return &template, nil
|
||||
}
|
||||
|
||||
// LoadTemplateDirectory loads all templates from a directory
|
||||
func (wtm *WorkflowTemplateManager) LoadTemplateDirectory() error {
|
||||
entries, err := os.ReadDir(wtm.templatePath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil // Directory doesn't exist yet
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() && filepath.Ext(entry.Name()) == ".yaml" {
|
||||
name := entry.Name()[:len(entry.Name())-5] // Remove .yaml
|
||||
_, _ = wtm.LoadTemplate(name)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetTemplate retrieves a cached template
|
||||
func (wtm *WorkflowTemplateManager) GetTemplate(name string) (*WorkflowTemplate, bool) {
|
||||
wtm.mu.RLock()
|
||||
defer wtm.mu.RUnlock()
|
||||
|
||||
template, exists := wtm.templates[name]
|
||||
return template, exists
|
||||
}
|
||||
|
||||
// ListTemplates returns all loaded templates
|
||||
func (wtm *WorkflowTemplateManager) ListTemplates() map[string]*WorkflowTemplate {
|
||||
wtm.mu.RLock()
|
||||
defer wtm.mu.RUnlock()
|
||||
|
||||
result := make(map[string]*WorkflowTemplate)
|
||||
for name, template := range wtm.templates {
|
||||
result[name] = template
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// DeleteTemplate deletes a template
|
||||
func (wtm *WorkflowTemplateManager) DeleteTemplate(name string) error {
|
||||
wtm.mu.Lock()
|
||||
defer wtm.mu.Unlock()
|
||||
|
||||
filePath, exists := wtm.loadedFrom[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("template not found: %s", name)
|
||||
}
|
||||
|
||||
if err := os.Remove(filePath); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
|
||||
delete(wtm.templates, name)
|
||||
delete(wtm.loadedFrom, name)
|
||||
delete(wtm.lastModified, name)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExportTemplateJSON exports a template as JSON
|
||||
func (wtm *WorkflowTemplateManager) ExportTemplateJSON(name string) (string, error) {
|
||||
wtm.mu.RLock()
|
||||
template, exists := wtm.templates[name]
|
||||
wtm.mu.RUnlock()
|
||||
|
||||
if !exists {
|
||||
return "", fmt.Errorf("template not found: %s", name)
|
||||
}
|
||||
|
||||
data, err := json.MarshalIndent(template, "", " ")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
// ValidateTemplate validates a template
|
||||
func (wtm *WorkflowTemplateManager) ValidateTemplate(template *WorkflowTemplate) error {
|
||||
if template.Name == "" {
|
||||
return fmt.Errorf("template name is required")
|
||||
}
|
||||
|
||||
if template.Orchestrator.PlannerURL == "" {
|
||||
return fmt.Errorf("planner_url is required")
|
||||
}
|
||||
|
||||
if len(template.Tasks) == 0 {
|
||||
return fmt.Errorf("at least one task is required")
|
||||
}
|
||||
|
||||
// Validate task dependencies
|
||||
taskIDs := make(map[string]bool)
|
||||
for _, task := range template.Tasks {
|
||||
if task.ID == "" {
|
||||
return fmt.Errorf("task ID is required")
|
||||
}
|
||||
taskIDs[task.ID] = true
|
||||
}
|
||||
|
||||
for _, task := range template.Tasks {
|
||||
for _, dep := range task.DependsOn {
|
||||
if !taskIDs[dep] {
|
||||
return fmt.Errorf("task %s depends on non-existent task %s", task.ID, dep)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RecordUsage records template usage
|
||||
func (wtm *WorkflowTemplateManager) RecordUsage(name string) error {
|
||||
wtm.mu.Lock()
|
||||
defer wtm.mu.Unlock()
|
||||
|
||||
if _, exists := wtm.templates[name]; !exists {
|
||||
return fmt.Errorf("template not found: %s", name)
|
||||
}
|
||||
|
||||
wtm.stats.ExecutedTemplates++
|
||||
wtm.stats.LastUsed[name] = time.Now()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetStats returns template manager statistics
|
||||
func (wtm *WorkflowTemplateManager) GetStats() *TemplateStats {
|
||||
wtm.mu.RLock()
|
||||
defer wtm.mu.RUnlock()
|
||||
|
||||
stats := *wtm.stats
|
||||
return &stats
|
||||
}
|
||||
|
||||
// ClearCache clears all cached templates
|
||||
func (wtm *WorkflowTemplateManager) ClearCache() {
|
||||
wtm.mu.Lock()
|
||||
defer wtm.mu.Unlock()
|
||||
|
||||
wtm.templates = make(map[string]*WorkflowTemplate)
|
||||
wtm.loadedFrom = make(map[string]string)
|
||||
wtm.lastModified = make(map[string]time.Time)
|
||||
}
|
||||
@@ -0,0 +1,357 @@
|
||||
package templates
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func createTestTemplate() *WorkflowTemplate {
|
||||
return &WorkflowTemplate{
|
||||
Name: "golang-project",
|
||||
Version: "1.0.0",
|
||||
Description: "Template for Go projects",
|
||||
Author: "test-author",
|
||||
Orchestrator: OrchestratorConfig{
|
||||
PlannerURL: "http://planner:8000",
|
||||
JudgeURL: "http://judge:8000",
|
||||
ImplementerURL: "http://implementer:8000",
|
||||
TimeoutSeconds: 300,
|
||||
RetryPolicy: "exponential",
|
||||
MaxConcurrency: 10,
|
||||
},
|
||||
Tasks: []TaskConfig{
|
||||
{
|
||||
ID: "T0.1",
|
||||
Title: "Analyze Requirements",
|
||||
Description: "Analyze project requirements",
|
||||
Type: "feature",
|
||||
Priority: "high",
|
||||
},
|
||||
{
|
||||
ID: "T0.2",
|
||||
Title: "Implement Solution",
|
||||
Description: "Implement the solution",
|
||||
Type: "feature",
|
||||
Priority: "high",
|
||||
DependsOn: []string{"T0.1"},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewWorkflowTemplateManager(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
assert.NotNil(t, manager)
|
||||
assert.Equal(t, tmpDir, manager.templatePath)
|
||||
}
|
||||
|
||||
func TestSaveTemplate(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
err := manager.SaveTemplate(template)
|
||||
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Check file was created
|
||||
filePath := filepath.Join(tmpDir, "golang-project.yaml")
|
||||
_, err = os.Stat(filePath)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestSaveTemplateNil(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
err := manager.SaveTemplate(nil)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestLoadTemplate(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
loaded, err := manager.LoadTemplate("golang-project")
|
||||
assert.NoError(t, err)
|
||||
assert.NotNil(t, loaded)
|
||||
assert.Equal(t, "golang-project", loaded.Name)
|
||||
}
|
||||
|
||||
func TestLoadTemplateNotFound(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
_, err := manager.LoadTemplate("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestGetTemplate(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
retrieved, exists := manager.GetTemplate("golang-project")
|
||||
assert.True(t, exists)
|
||||
assert.Equal(t, "golang-project", retrieved.Name)
|
||||
}
|
||||
|
||||
func TestGetTemplateNotFound(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
_, exists := manager.GetTemplate("nonexistent")
|
||||
assert.False(t, exists)
|
||||
}
|
||||
|
||||
func TestListTemplates(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
template := createTestTemplate()
|
||||
template.Name = "template-" + string(rune(48+i))
|
||||
manager.SaveTemplate(template)
|
||||
}
|
||||
|
||||
templates := manager.ListTemplates()
|
||||
assert.Equal(t, 3, len(templates))
|
||||
}
|
||||
|
||||
func TestDeleteTemplate(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
err := manager.DeleteTemplate("golang-project")
|
||||
assert.NoError(t, err)
|
||||
|
||||
_, exists := manager.GetTemplate("golang-project")
|
||||
assert.False(t, exists)
|
||||
}
|
||||
|
||||
func TestDeleteTemplateNotFound(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
err := manager.DeleteTemplate("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestValidateTemplateValid(t *testing.T) {
|
||||
template := createTestTemplate()
|
||||
err := NewWorkflowTemplateManager("/tmp").ValidateTemplate(template)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestValidateTemplateEmptyName(t *testing.T) {
|
||||
template := createTestTemplate()
|
||||
template.Name = ""
|
||||
err := NewWorkflowTemplateManager("/tmp").ValidateTemplate(template)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestValidateTemplateNoTasks(t *testing.T) {
|
||||
template := createTestTemplate()
|
||||
template.Tasks = make([]TaskConfig, 0)
|
||||
err := NewWorkflowTemplateManager("/tmp").ValidateTemplate(template)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestValidateTemplateInvalidDependency(t *testing.T) {
|
||||
template := createTestTemplate()
|
||||
template.Tasks[1].DependsOn = []string{"nonexistent"}
|
||||
err := NewWorkflowTemplateManager("/tmp").ValidateTemplate(template)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestExportTemplateJSON(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
json, err := manager.ExportTemplateJSON("golang-project")
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, json)
|
||||
assert.Contains(t, json, "golang-project")
|
||||
}
|
||||
|
||||
func TestExportTemplateJSONNotFound(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
_, err := manager.ExportTemplateJSON("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestRecordUsage(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
err := manager.RecordUsage("golang-project")
|
||||
assert.NoError(t, err)
|
||||
|
||||
stats := manager.GetStats()
|
||||
assert.Equal(t, 1, stats.ExecutedTemplates)
|
||||
}
|
||||
|
||||
func TestRecordUsageNotFound(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
err := manager.RecordUsage("nonexistent")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestTemplateGetStats(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
stats := manager.GetStats()
|
||||
assert.Equal(t, 1, stats.TotalTemplates)
|
||||
assert.Equal(t, 1, stats.LoadedTemplates)
|
||||
}
|
||||
|
||||
func TestLoadTemplateDirectory(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
template := createTestTemplate()
|
||||
template.Name = "template-" + string(rune(48+i))
|
||||
manager.SaveTemplate(template)
|
||||
}
|
||||
|
||||
manager.ClearCache()
|
||||
err := manager.LoadTemplateDirectory()
|
||||
assert.NoError(t, err)
|
||||
|
||||
templates := manager.ListTemplates()
|
||||
assert.Equal(t, 3, len(templates))
|
||||
}
|
||||
|
||||
func TestClearCache(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
assert.Equal(t, 1, len(manager.ListTemplates()))
|
||||
|
||||
manager.ClearCache()
|
||||
assert.Equal(t, 0, len(manager.ListTemplates()))
|
||||
}
|
||||
|
||||
func TestTemplateTimestamps(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
retrieved, _ := manager.GetTemplate("golang-project")
|
||||
assert.False(t, retrieved.CreatedAt.IsZero())
|
||||
assert.False(t, retrieved.UpdatedAt.IsZero())
|
||||
assert.True(t, retrieved.UpdatedAt.After(retrieved.CreatedAt) || retrieved.UpdatedAt.Equal(retrieved.CreatedAt))
|
||||
}
|
||||
|
||||
func TestTemplateWithDefaults(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
template.Defaults = map[string]interface{}{
|
||||
"language": "go",
|
||||
"version": "1.20",
|
||||
}
|
||||
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
retrieved, _ := manager.GetTemplate("golang-project")
|
||||
assert.NotNil(t, retrieved.Defaults)
|
||||
assert.Equal(t, "go", retrieved.Defaults["language"])
|
||||
}
|
||||
|
||||
func TestTemplateWithTags(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
template.Tags = map[string]string{
|
||||
"environment": "production",
|
||||
"team": "backend",
|
||||
}
|
||||
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
retrieved, _ := manager.GetTemplate("golang-project")
|
||||
assert.NotNil(t, retrieved.Tags)
|
||||
assert.Equal(t, "production", retrieved.Tags["environment"])
|
||||
}
|
||||
|
||||
func TestMultipleTemplates(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
template := createTestTemplate()
|
||||
template.Name = "template-" + string(rune(48+i))
|
||||
manager.SaveTemplate(template)
|
||||
}
|
||||
|
||||
templates := manager.ListTemplates()
|
||||
assert.Equal(t, 5, len(templates))
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
name := "template-" + string(rune(48+i))
|
||||
manager.RecordUsage(name)
|
||||
}
|
||||
|
||||
stats := manager.GetStats()
|
||||
assert.Equal(t, 5, stats.ExecutedTemplates)
|
||||
}
|
||||
|
||||
func BenchmarkSaveTemplate(b *testing.B) {
|
||||
tmpDir := b.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
for i := 0; i < b.N; i++ {
|
||||
template := createTestTemplate()
|
||||
template.Name = "template-" + string(rune(48+i%100))
|
||||
manager.SaveTemplate(template)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLoadTemplate(b *testing.B) {
|
||||
tmpDir := b.TempDir()
|
||||
manager := NewWorkflowTemplateManager(tmpDir)
|
||||
|
||||
template := createTestTemplate()
|
||||
manager.SaveTemplate(template)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
manager.LoadTemplate("golang-project")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,237 @@
|
||||
package visualization
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// TaskNode represents a task in the DAG
|
||||
type TaskNode struct {
|
||||
ID string
|
||||
Status string // pending, running, completed, failed
|
||||
Duration float64
|
||||
Critical bool
|
||||
}
|
||||
|
||||
// DAGRenderer renders workflow dependency graphs
|
||||
type DAGRenderer struct {
|
||||
nodes map[string]*TaskNode
|
||||
edges map[string][]string
|
||||
}
|
||||
|
||||
// NewDAGRenderer creates a new DAG renderer
|
||||
func NewDAGRenderer() *DAGRenderer {
|
||||
return &DAGRenderer{
|
||||
nodes: make(map[string]*TaskNode),
|
||||
edges: make(map[string][]string),
|
||||
}
|
||||
}
|
||||
|
||||
// AddNode adds a task node
|
||||
func (dr *DAGRenderer) AddNode(id, status string, duration float64) {
|
||||
dr.nodes[id] = &TaskNode{
|
||||
ID: id,
|
||||
Status: status,
|
||||
Duration: duration,
|
||||
}
|
||||
}
|
||||
|
||||
// AddEdge adds a dependency edge
|
||||
func (dr *DAGRenderer) AddEdge(from, to string) error {
|
||||
if _, exists := dr.nodes[from]; !exists {
|
||||
return fmt.Errorf("source node not found: %s", from)
|
||||
}
|
||||
if _, exists := dr.nodes[to]; !exists {
|
||||
return fmt.Errorf("target node not found: %s", to)
|
||||
}
|
||||
|
||||
dr.edges[from] = append(dr.edges[from], to)
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarkCriticalPath marks nodes on the critical path
|
||||
func (dr *DAGRenderer) MarkCriticalPath(nodes []string) error {
|
||||
for _, nodeID := range nodes {
|
||||
if node, exists := dr.nodes[nodeID]; exists {
|
||||
node.Critical = true
|
||||
} else {
|
||||
return fmt.Errorf("node not found: %s", nodeID)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RenderDOT generates DOT format for Graphviz
|
||||
func (dr *DAGRenderer) RenderDOT() string {
|
||||
var buf strings.Builder
|
||||
|
||||
buf.WriteString("digraph WorkflowDAG {\n")
|
||||
buf.WriteString(" rankdir=LR;\n")
|
||||
buf.WriteString(" node [shape=box];\n\n")
|
||||
|
||||
// Render nodes
|
||||
for _, node := range dr.nodes {
|
||||
color := "lightgray"
|
||||
if node.Critical {
|
||||
color = "red"
|
||||
} else if node.Status == "completed" {
|
||||
color = "lightgreen"
|
||||
} else if node.Status == "failed" {
|
||||
color = "lightcoral"
|
||||
} else if node.Status == "running" {
|
||||
color = "lightyellow"
|
||||
}
|
||||
|
||||
label := fmt.Sprintf("%s\\n%.0fms", node.ID, node.Duration)
|
||||
buf.WriteString(fmt.Sprintf(" \"%s\" [label=\"%s\", fillcolor=%s, style=filled];\n",
|
||||
node.ID, label, color))
|
||||
}
|
||||
|
||||
buf.WriteString("\n")
|
||||
|
||||
// Render edges
|
||||
for from, tos := range dr.edges {
|
||||
for _, to := range tos {
|
||||
buf.WriteString(fmt.Sprintf(" \"%s\" -> \"%s\";\n", from, to))
|
||||
}
|
||||
}
|
||||
|
||||
buf.WriteString("}\n")
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
// RenderHTML generates a simple HTML visualization
|
||||
func (dr *DAGRenderer) RenderHTML() string {
|
||||
var buf strings.Builder
|
||||
|
||||
buf.WriteString("<html><body>\n")
|
||||
buf.WriteString("<h1>Workflow DAG</h1>\n")
|
||||
buf.WriteString("<table border='1'>\n")
|
||||
buf.WriteString("<tr><th>Task ID</th><th>Status</th><th>Duration (ms)</th><th>Critical Path</th></tr>\n")
|
||||
|
||||
for _, node := range dr.nodes {
|
||||
critical := "No"
|
||||
if node.Critical {
|
||||
critical = "Yes"
|
||||
}
|
||||
buf.WriteString(fmt.Sprintf("<tr><td>%s</td><td>%s</td><td>%.0f</td><td>%s</td></tr>\n",
|
||||
node.ID, node.Status, node.Duration, critical))
|
||||
}
|
||||
|
||||
buf.WriteString("</table>\n")
|
||||
buf.WriteString("</body></html>\n")
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
// GetTopologicalSort returns tasks in topological order
|
||||
func (dr *DAGRenderer) GetTopologicalSort() ([]string, error) {
|
||||
// Simple topological sort using DFS
|
||||
visited := make(map[string]bool)
|
||||
result := make([]string, 0)
|
||||
|
||||
var visit func(string) error
|
||||
visit = func(nodeID string) error {
|
||||
if visited[nodeID] {
|
||||
return nil
|
||||
}
|
||||
visited[nodeID] = true
|
||||
|
||||
// Visit dependencies first
|
||||
for _, dep := range dr.edges[nodeID] {
|
||||
if err := visit(dep); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
result = append(result, nodeID)
|
||||
return nil
|
||||
}
|
||||
|
||||
for nodeID := range dr.nodes {
|
||||
if err := visit(nodeID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// GetParallel returns groups of tasks that can run in parallel
|
||||
func (dr *DAGRenderer) GetParallel() map[int][]string {
|
||||
levels := make(map[int][]string)
|
||||
inDegree := make(map[string]int)
|
||||
|
||||
// Calculate in-degree
|
||||
for _, node := range dr.nodes {
|
||||
inDegree[node.ID] = 0
|
||||
}
|
||||
|
||||
for _, tos := range dr.edges {
|
||||
for _, to := range tos {
|
||||
inDegree[to]++
|
||||
}
|
||||
}
|
||||
|
||||
// Find nodes by level
|
||||
processed := make(map[string]bool)
|
||||
level := 0
|
||||
|
||||
for len(processed) < len(dr.nodes) {
|
||||
var current []string
|
||||
|
||||
for _, node := range dr.nodes {
|
||||
if !processed[node.ID] && inDegree[node.ID] == 0 {
|
||||
current = append(current, node.ID)
|
||||
}
|
||||
}
|
||||
|
||||
if len(current) == 0 {
|
||||
break
|
||||
}
|
||||
|
||||
levels[level] = current
|
||||
|
||||
// Update in-degrees
|
||||
for _, nodeID := range current {
|
||||
processed[nodeID] = true
|
||||
for _, to := range dr.edges[nodeID] {
|
||||
inDegree[to]--
|
||||
}
|
||||
}
|
||||
|
||||
level++
|
||||
}
|
||||
|
||||
return levels
|
||||
}
|
||||
|
||||
// GetStats returns statistics about the DAG
|
||||
func (dr *DAGRenderer) GetStats() map[string]interface{} {
|
||||
totalDuration := 0.0
|
||||
maxDuration := 0.0
|
||||
criticalCount := 0
|
||||
edgeCount := 0
|
||||
|
||||
for _, node := range dr.nodes {
|
||||
totalDuration += node.Duration
|
||||
if node.Duration > maxDuration {
|
||||
maxDuration = node.Duration
|
||||
}
|
||||
if node.Critical {
|
||||
criticalCount++
|
||||
}
|
||||
}
|
||||
|
||||
// Count total edges
|
||||
for _, tos := range dr.edges {
|
||||
edgeCount += len(tos)
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"node_count": len(dr.nodes),
|
||||
"edge_count": edgeCount,
|
||||
"total_duration": totalDuration,
|
||||
"max_duration": maxDuration,
|
||||
"critical_count": criticalCount,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,135 @@
|
||||
package visualization
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestAddNode(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
|
||||
assert.NotNil(t, dag.nodes["T1"])
|
||||
assert.Equal(t, "completed", dag.nodes["T1"].Status)
|
||||
}
|
||||
|
||||
func TestAddEdge(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
dag.AddNode("T2", "completed", 200)
|
||||
|
||||
err := dag.AddEdge("T1", "T2")
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 1, len(dag.edges["T1"]))
|
||||
}
|
||||
|
||||
func TestAddEdgeNotFound(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
|
||||
err := dag.AddEdge("T1", "T2")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestMarkCriticalPath(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
dag.AddNode("T2", "completed", 200)
|
||||
|
||||
err := dag.MarkCriticalPath([]string{"T1", "T2"})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, dag.nodes["T1"].Critical)
|
||||
assert.True(t, dag.nodes["T2"].Critical)
|
||||
}
|
||||
|
||||
func TestRenderDOT(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
dag.AddNode("T2", "running", 200)
|
||||
dag.AddEdge("T1", "T2")
|
||||
|
||||
dot := dag.RenderDOT()
|
||||
assert.True(t, strings.Contains(dot, "digraph WorkflowDAG"))
|
||||
assert.True(t, strings.Contains(dot, "T1"))
|
||||
assert.True(t, strings.Contains(dot, "T2"))
|
||||
assert.True(t, strings.Contains(dot, "->"))
|
||||
}
|
||||
|
||||
func TestRenderHTML(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
|
||||
html := dag.RenderHTML()
|
||||
assert.True(t, strings.Contains(html, "<html>"))
|
||||
assert.True(t, strings.Contains(html, "Workflow DAG"))
|
||||
assert.True(t, strings.Contains(html, "T1"))
|
||||
}
|
||||
|
||||
func TestGetTopologicalSort(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
dag.AddNode("T2", "completed", 200)
|
||||
dag.AddNode("T3", "completed", 150)
|
||||
|
||||
dag.AddEdge("T1", "T2")
|
||||
dag.AddEdge("T2", "T3")
|
||||
|
||||
sorted, err := dag.GetTopologicalSort()
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 3, len(sorted))
|
||||
}
|
||||
|
||||
func TestGetParallel(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
dag.AddNode("T2", "completed", 200)
|
||||
dag.AddNode("T3", "completed", 150)
|
||||
|
||||
dag.AddEdge("T1", "T3")
|
||||
|
||||
parallel := dag.GetParallel()
|
||||
assert.True(t, len(parallel) > 0)
|
||||
}
|
||||
|
||||
func TestGetStats(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
dag.AddNode("T2", "completed", 200)
|
||||
dag.AddEdge("T1", "T2")
|
||||
|
||||
stats := dag.GetStats()
|
||||
assert.Equal(t, 2, stats["node_count"])
|
||||
assert.Equal(t, 1, stats["edge_count"])
|
||||
assert.Equal(t, 300.0, stats["total_duration"])
|
||||
}
|
||||
|
||||
func TestCriticalPathHighlighting(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
dag.AddNode("T1", "completed", 100)
|
||||
dag.AddNode("T2", "completed", 200)
|
||||
dag.MarkCriticalPath([]string{"T1", "T2"})
|
||||
|
||||
dot := dag.RenderDOT()
|
||||
assert.True(t, strings.Contains(dot, "fillcolor=red"))
|
||||
}
|
||||
|
||||
func TestComplexDAG(t *testing.T) {
|
||||
dag := NewDAGRenderer()
|
||||
|
||||
// Create a diamond-shaped DAG
|
||||
dag.AddNode("Start", "completed", 50)
|
||||
dag.AddNode("A", "completed", 100)
|
||||
dag.AddNode("B", "completed", 150)
|
||||
dag.AddNode("End", "completed", 50)
|
||||
|
||||
dag.AddEdge("Start", "A")
|
||||
dag.AddEdge("Start", "B")
|
||||
dag.AddEdge("A", "End")
|
||||
dag.AddEdge("B", "End")
|
||||
|
||||
stats := dag.GetStats()
|
||||
assert.Equal(t, 4, stats["node_count"])
|
||||
assert.Equal(t, 4, stats["edge_count"])
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user