yao/agent/test/input_source.go
Max 6ea982f027 Implement Agent-Driven Input and Dry-Run Mode in Test Framework
- Added support for generating test cases using agents and scripts through the new `LoadFromAgent` and `LoadFromScript` methods in the loader.
- Enhanced the `RunTests` method to handle different input sources, including agent-driven and script-based test case generation.
- Introduced a `--dry-run` flag to allow users to preview generated test cases without executing them.
- Updated relevant documentation in DESIGN_V2.md and TODO_V2.md to reflect the new features and implementation status.
2025-12-26 10:53:11 +08:00

392 lines
10 KiB
Go

package test
import (
"fmt"
"net/url"
"strconv"
"strings"
jsoniter "github.com/json-iterator/go"
goutext "github.com/yaoapp/gou/text"
"github.com/yaoapp/yao/agent/assistant"
"github.com/yaoapp/yao/agent/context"
)
// InputSourceType represents the type of input source
type InputSourceType string
const (
// InputSourceFile indicates input from a JSONL file
InputSourceFile InputSourceType = "file"
// InputSourceMessage indicates input from a direct message string
InputSourceMessage InputSourceType = "message"
// InputSourceScript indicates script test mode
InputSourceScript InputSourceType = "script"
// InputSourceAgent indicates input generated by an agent
InputSourceAgent InputSourceType = "agent"
)
// InputSource represents a parsed input source
type InputSource struct {
Type InputSourceType // file, message, script, agent
Value string // path, message, script ref, or agent ID
Params map[string]interface{} // query parameters (for agent source)
}
// ParseInputSource parses the -i flag value into an InputSource
// Supported formats:
// - "agents:workers.test.generator" - Agent-generated test cases
// - "agents:workers.test.generator?count=10&focus=edge-cases" - With parameters
// - "scripts.tests.gen" - Script-generated test cases
// - "./tests/inputs.jsonl" - JSONL file
// - "Hello, how are you?" - Direct message
func ParseInputSource(input string) *InputSource {
// Check for agents: prefix
if strings.HasPrefix(input, "agents:") {
return parseAgentSource(strings.TrimPrefix(input, "agents:"))
}
// Check for scripts: prefix (for generator scripts)
if strings.HasPrefix(input, "scripts:") {
return &InputSource{
Type: InputSourceScript,
Value: strings.TrimPrefix(input, "scripts:"),
}
}
// Check for script test mode (scripts.xxx format without prefix)
if strings.HasPrefix(input, "scripts.") {
return &InputSource{
Type: InputSourceScript,
Value: input,
}
}
// Check for file extension
if strings.HasSuffix(input, ".jsonl") || strings.HasSuffix(input, ".json") {
return &InputSource{
Type: InputSourceFile,
Value: input,
}
}
// Check if it looks like a file path
if strings.Contains(input, "/") || strings.Contains(input, "\\") {
return &InputSource{
Type: InputSourceFile,
Value: input,
}
}
// Default to message
return &InputSource{
Type: InputSourceMessage,
Value: input,
}
}
// parseAgentSource parses an agent source string with optional query parameters
// Format: "agent.id" or "agent.id?count=10&focus=edge-cases"
func parseAgentSource(input string) *InputSource {
source := &InputSource{
Type: InputSourceAgent,
Params: make(map[string]interface{}),
}
// Check for query parameters
if idx := strings.Index(input, "?"); idx >= 0 {
source.Value = input[:idx]
queryStr := input[idx+1:]
// Parse query parameters
values, err := url.ParseQuery(queryStr)
if err == nil {
for key, vals := range values {
if len(vals) > 0 {
// Try to parse as number
if num, err := strconv.Atoi(vals[0]); err == nil {
source.Params[key] = num
} else if num, err := strconv.ParseFloat(vals[0], 64); err == nil {
source.Params[key] = num
} else if vals[0] == "true" {
source.Params[key] = true
} else if vals[0] == "false" {
source.Params[key] = false
} else {
source.Params[key] = vals[0]
}
}
}
}
} else {
source.Value = input
}
return source
}
// GeneratorInput represents the input sent to a generator agent
type GeneratorInput struct {
TargetAgent *TargetAgentInfo `json:"target_agent"`
Count int `json:"count,omitempty"`
Focus string `json:"focus,omitempty"`
Extra map[string]interface{} `json:"extra,omitempty"`
}
// TargetAgentInfo contains information about the agent being tested
type TargetAgentInfo struct {
ID string `json:"id"`
Description string `json:"description,omitempty"`
Tools []map[string]interface{} `json:"tools,omitempty"`
}
// GenerateTestCases generates test cases using a generator agent
func GenerateTestCases(agentID string, targetInfo *TargetAgentInfo, params map[string]interface{}) ([]*Case, error) {
// Get generator assistant
ast, err := assistant.Get(agentID)
if err != nil {
return nil, fmt.Errorf("failed to get generator agent %s: %w", agentID, err)
}
// Build generation request
genInput := &GeneratorInput{
TargetAgent: targetInfo,
Count: 5, // Default count
}
// Apply parameters
if params != nil {
if count, ok := params["count"].(int); ok {
genInput.Count = count
}
if focus, ok := params["focus"].(string); ok {
genInput.Focus = focus
}
// Store extra parameters
genInput.Extra = make(map[string]interface{})
for k, v := range params {
if k != "count" && k != "focus" {
genInput.Extra[k] = v
}
}
}
// Create context
env := NewEnvironment("", "")
ctx := NewTestContext("generator", agentID, env)
defer ctx.Release()
// Build options - skip history and trace for efficiency
opts := &context.Options{
Skip: &context.Skip{
History: true,
Trace: true,
Output: true,
},
Metadata: map[string]interface{}{
"test_mode": "generator",
},
}
// Build message
inputJSON, err := jsoniter.Marshal(genInput)
if err != nil {
return nil, fmt.Errorf("failed to marshal generator input: %w", err)
}
messages := []context.Message{{
Role: context.RoleUser,
Content: string(inputJSON),
}}
// Call generator agent
response, err := ast.Stream(ctx, messages, opts)
if err != nil {
return nil, fmt.Errorf("generator agent error: %w", err)
}
// Extract and parse response
return parseGeneratedCases(response)
}
// parseGeneratedCases parses the generator agent's response into test cases
func parseGeneratedCases(response *context.Response) ([]*Case, error) {
if response == nil || response.Completion == nil {
return nil, fmt.Errorf("empty response from generator agent")
}
// Extract content
content := response.Completion.Content
if content == nil {
return nil, fmt.Errorf("no content in generator response")
}
// Convert content to string
var text string
switch v := content.(type) {
case string:
text = v
default:
data, err := jsoniter.Marshal(content)
if err != nil {
return nil, fmt.Errorf("failed to marshal content: %w", err)
}
text = string(data)
}
// Use goutext.ExtractJSON for fault-tolerant parsing
parsed := goutext.ExtractJSON(text)
if parsed == nil {
return nil, fmt.Errorf("failed to parse generator response as JSON: %s", truncateOutput(text, 200))
}
// Convert to []*Case
return convertToCases(parsed)
}
// convertToCases converts parsed JSON to test cases
func convertToCases(parsed interface{}) ([]*Case, error) {
// Handle array of cases
arr, ok := parsed.([]interface{})
if !ok {
// Maybe it's a single case wrapped in an object
if obj, ok := parsed.(map[string]interface{}); ok {
if cases, ok := obj["cases"].([]interface{}); ok {
arr = cases
} else if testCases, ok := obj["test_cases"].([]interface{}); ok {
arr = testCases
} else {
// Single case
arr = []interface{}{obj}
}
} else {
return nil, fmt.Errorf("expected array of test cases, got %T", parsed)
}
}
cases := make([]*Case, 0, len(arr))
for i, item := range arr {
caseMap, ok := item.(map[string]interface{})
if !ok {
return nil, fmt.Errorf("test case %d is not an object", i)
}
tc, err := mapToCase(caseMap)
if err != nil {
return nil, fmt.Errorf("failed to parse test case %d: %w", i, err)
}
cases = append(cases, tc)
}
return cases, nil
}
// mapToCase converts a map to a Case struct
func mapToCase(m map[string]interface{}) (*Case, error) {
tc := &Case{}
// Required: id
if id, ok := m["id"].(string); ok {
tc.ID = id
} else {
return nil, fmt.Errorf("missing required field 'id'")
}
// Required: input
if input, ok := m["input"]; ok {
tc.Input = input
} else {
return nil, fmt.Errorf("missing required field 'input'")
}
// Optional: assertions/assert
if assertions, ok := m["assertions"]; ok {
tc.Assert = assertions
} else if assert, ok := m["assert"]; ok {
tc.Assert = assert
}
// Optional: options - convert map to CaseOptions
if options, ok := m["options"].(map[string]interface{}); ok {
tc.Options = mapToCaseOptions(options)
}
// Optional: before/after
if before, ok := m["before"].(string); ok {
tc.Before = before
}
if after, ok := m["after"].(string); ok {
tc.After = after
}
// Optional: timeout
if timeout, ok := m["timeout"].(string); ok {
tc.Timeout = timeout
}
return tc, nil
}
// ToInputMode converts InputSourceType to InputMode for backward compatibility
func (s *InputSource) ToInputMode() InputMode {
switch s.Type {
case InputSourceFile:
return InputModeFile
case InputSourceMessage:
return InputModeMessage
case InputSourceScript:
return InputModeScript
case InputSourceAgent:
// Agent source generates cases, then runs in file mode
return InputModeFile
default:
return InputModeMessage
}
}
// mapToCaseOptions converts a map to CaseOptions
func mapToCaseOptions(m map[string]interface{}) *CaseOptions {
opts := &CaseOptions{}
if connector, ok := m["connector"].(string); ok {
opts.Connector = connector
}
if mode, ok := m["mode"].(string); ok {
opts.Mode = mode
}
if disableGlobalPrompts, ok := m["disable_global_prompts"].(bool); ok {
opts.DisableGlobalPrompts = disableGlobalPrompts
}
if search, ok := m["search"].(bool); ok {
opts.Search = &search
}
if metadata, ok := m["metadata"].(map[string]interface{}); ok {
opts.Metadata = metadata
}
if skip, ok := m["skip"].(map[string]interface{}); ok {
opts.Skip = &CaseSkipOptions{}
if history, ok := skip["history"].(bool); ok {
opts.Skip.History = history
}
if trace, ok := skip["trace"].(bool); ok {
opts.Skip.Trace = trace
}
if output, ok := skip["output"].(bool); ok {
opts.Skip.Output = output
}
if keyword, ok := skip["keyword"].(bool); ok {
opts.Skip.Keyword = keyword
}
if searchSkip, ok := skip["search"].(bool); ok {
opts.Skip.Search = searchSkip
}
}
return opts
}