Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
194 changes: 173 additions & 21 deletions internal/cli/cli_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -241,11 +241,8 @@ func TestSend(t *testing.T) {
if err != nil {
t.Fatalf("runCMD(%q) error = %v", strings.Join(tt.args(mode.url), " "), err)
}
var task a2a.Task
if err := json.Unmarshal([]byte(out), &task); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err)
}
if text := testutil.AllArtifactText(&task); text != tt.wantText {
task := mustDecodeTask(t, out)
if text := testutil.AllArtifactText(task); text != tt.wantText {
t.Fatalf("allArtifactText() = %q, want %q", text, tt.wantText)
}
})
Expand Down Expand Up @@ -285,11 +282,8 @@ func TestSend_AgentCardFromFile(t *testing.T) {
if err != nil {
t.Fatalf("runCMD(%q) error = %v", strings.Join(tt.args, " "), err)
}
var task a2a.Task
if err := json.Unmarshal([]byte(out), &task); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err)
}
if text := testutil.AllArtifactText(&task); text != tt.wantText {
task := mustDecodeTask(t, out)
if text := testutil.AllArtifactText(task); text != tt.wantText {
t.Fatalf("allArtifactText() = %q, want %q", text, tt.wantText)
}
})
Expand All @@ -306,11 +300,8 @@ func TestSendDataPart(t *testing.T) {
}

out := mustRunCMD(t, "send", "-a", url, "-o", "json", "--data-part", path)
var task a2a.Task
if err := json.Unmarshal([]byte(out), &task); err != nil {
t.Fatalf("json.Unmarshal(send --data-part output) error = %v", err)
}
if got := testutil.AllArtifactText(&task); got != `{"hello":"world"}` {
task := mustDecodeTask(t, out)
if got := testutil.AllArtifactText(task); got != `{"hello":"world"}` {
t.Fatalf("allArtifactText() = %q, want %q", got, `{"hello":"world"}`)
}
}
Expand All @@ -325,11 +316,8 @@ func TestSendRequestPayloadFile(t *testing.T) {
}

out := mustRunCMD(t, "send", "-a", url, "-o", "json", "--request-payload", path)
var task a2a.Task
if err := json.Unmarshal([]byte(out), &task); err != nil {
t.Fatalf("json.Unmarshal(send --request-payload output) error = %v", err)
}
if got := testutil.AllArtifactText(&task); got != "from file" {
task := mustDecodeTask(t, out)
if got := testutil.AllArtifactText(task); got != "from file" {
t.Fatalf("allArtifactText() = %q, want %q", got, "from file")
}
}
Expand Down Expand Up @@ -490,12 +478,61 @@ func TestSendStreaming(t *testing.T) {
}
}

func TestSendStreamJSONL(t *testing.T) {
t.Parallel()
url := startTestServer(t)

testCases := []struct {
name string
flags []string
wantObjectPerLine bool
}{
{
name: "compact one object per line",
flags: []string{"-a", url, "-o", "json", "--stream"},
wantObjectPerLine: true,
},
{
name: "indented objects when pretty",
flags: []string{"-a", url, "-o", "json", "--stream", "--pretty"},
wantObjectPerLine: false,
},
}

for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
command := append([]string{"send", "stream me"}, tc.flags...)
out := mustRunCMD(t, command...)
lines := strings.Split(strings.TrimRight(out, "\n"), "\n")
if len(lines) == 0 {
t.Fatalf("send --stream produced no JSONL lines")
}
objectPerLine := true
for i, line := range lines {
var sr a2a.StreamResponse
if err := json.Unmarshal([]byte(line), &sr); err != nil {
if tc.wantObjectPerLine {
t.Fatalf("JSONL line %d is not an independently parseable object: %v\nline: %s", i, err, line)
}
objectPerLine = false
break
}
}
if objectPerLine && !tc.wantObjectPerLine {
t.Fatalf("all outputs lines contained a well-formed a2a.StreamResponse:\n%s", out)
}
})
}

}

func TestSendStreamingFallbackUsesDefaultPoller(t *testing.T) {
t.Parallel()
nonStreamingURL := startTestServerWith(t, a2a.AgentCapabilities{Streaming: false})

out, err := runCMDWithConfig(t, deps{cfgLoader: clicfg.LoadEmpty},
"send", "-a", nonStreamingURL, "-o", "json", "--stream", "stream me", "--polling-interval", "5ms")
"send", "-a", nonStreamingURL, "-o", "json", "--stream", "stream me", "--poll-interval", "5ms")
if err != nil {
t.Fatalf("runCMDWithConfig() error = %v", err)
}
Expand All @@ -514,6 +551,108 @@ func TestSendStreamingFallbackUsesDefaultPoller(t *testing.T) {
}
}

func TestSend_ResumeHintForInputRequiredTask(t *testing.T) {
t.Parallel()

var taskID a2a.TaskID
server := httptest.NewServer(a2asrv.NewRESTHandler(a2asrv.NewHandler(
a2asrv.AgentExecutorFunc(func(ctx context.Context, ec *a2asrv.ExecutorContext) iter.Seq2[a2a.Event, error] {
return func(yield func(a2a.Event, error) bool) {
taskID = ec.TaskID
task := &a2a.Task{
ID: ec.TaskID,
ContextID: ec.ContextID,
Status: a2a.TaskStatus{State: a2a.TaskStateInputRequired},
}
yield(task, nil)
}
}),
)))
t.Cleanup(server.Close)

out := mustRunCMD(t, "send", "-e", server.URL, "--transport", "rest", "hello")
if !strings.Contains(out, "a2a send --task-id "+string(taskID)) {
t.Fatalf("send text output missing the resume hint:\n%s", out)
}
}

func TestSendWithVersionSelector(t *testing.T) {
t.Parallel()
url := startTestServer(t)
legacyURL := startLegacyTestServer(t)

testCases := []struct {
name string
connect []string
version string
wantErr bool
}{
{
name: "new server success",
connect: []string{"-a", url},
version: "1.0",
},
{
name: "old server success",
connect: []string{"-a", legacyURL},
version: "0.3",
},
{
name: "new server direct success",
connect: []string{"-e", url, "--transport", "rest"},
version: "1.0",
},
{
name: "old server direct success",
connect: []string{"-e", legacyURL, "--transport", "jsonrpc"},
version: "0.3",
},
{
name: "new server failure",
connect: []string{"-a", url},
version: "0.3",
wantErr: true,
},
{
name: "new server direct failure",
connect: []string{"-e", url, "--transport", "rest"},
version: "0.3",
wantErr: true,
},
{
name: "old server failure",
connect: []string{"-a", legacyURL},
version: "1.0",
wantErr: true,
},
{
name: "old server direct failure",
connect: []string{"-e", legacyURL, "--transport", "jsonrpc"},
version: "1.0",
wantErr: true,
},
{
name: "unknown version failure",
connect: []string{"-e", url},
version: "3.0",
wantErr: true,
},
}
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
command := []string{"send", "--a2a-version", tc.version, "-o", "json", "hi"}
command = append(command, tc.connect...)
_, err := runCMD(t, command...)
if err != nil && !tc.wantErr {
t.Fatalf("send error = %v", err)
}
if err == nil && tc.wantErr {
t.Fatal("send error = nil, wanted a failure")
}
})
}
}

func TestGetTask(t *testing.T) {
t.Parallel()
url := startTestServer(t)
Expand Down Expand Up @@ -695,6 +834,19 @@ func startLegacyTestServer(t *testing.T) string {
return server.URL
}

func mustDecodeTask(t *testing.T, out string) *a2a.Task {
t.Helper()
var resp a2a.StreamResponse
if err := json.Unmarshal([]byte(out), &resp); err != nil {
t.Fatalf("json.Unmarshal() error = %v\noutput: %s", err, out)
}
task, ok := resp.Event.(*a2a.Task)
if !ok {
t.Fatalf("send output has no task wrapper: %s", out)
}
return task
}

func sendTestMessage(t *testing.T, url, text string) a2a.TaskID {
t.Helper()
ctx := t.Context()
Expand Down
30 changes: 21 additions & 9 deletions internal/cli/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,9 @@ func newClientFromEndpoint(ctx context.Context, cfg *globalConfig, ref string, e
cfg.logf("connecting directly to %s via %s (skipping card resolution)", endpointURL, protocol)

endpoint := a2a.NewAgentInterface(endpointURL, protocol)
if cfg.a2aVersion != "" {
endpoint.ProtocolVersion = a2a.ProtocolVersion(cfg.a2aVersion)
}
client, err := a2aclient.NewFromEndpoints(ctx, []*a2a.AgentInterface{endpoint}, append(clientFactoryOpts(cfg), extraOpts...)...)
return client, hintInsecure(err)
}
Expand Down Expand Up @@ -108,19 +111,28 @@ func hintInsecure(err error) error {
}

func clientFactoryOpts(cfg *globalConfig) []a2aclient.FactoryOption {
factoryOpts := []a2aclient.FactoryOption{
a2av0.WithRESTTransport(a2av0.RESTTransportConfig{}),
a2av0.WithJSONRPCTransport(a2av0.JSONRPCTransportConfig{}),
}
var grpcOpts []grpc.DialOption
if cfg.insecureGRPC {
grpcOpts = append(grpcOpts, grpc.WithTransportCredentials(insecure.NewCredentials()))
}
factoryOpts = append(factoryOpts,
a2agrpcv0.WithGRPCTransport(grpcOpts...),
a2agrpc.WithGRPCTransport(grpcOpts...),
)
return factoryOpts
opts := []a2aclient.FactoryOption{a2aclient.WithDefaultsDisabled()}
if cfg.a2aVersion == "" || cfg.a2aVersion == "1.0" {
opts = append(
opts,
a2aclient.WithRESTTransport(nil),
a2aclient.WithJSONRPCTransport(nil),
a2agrpc.WithGRPCTransport(grpcOpts...),
)
}
if cfg.a2aVersion == "" || cfg.a2aVersion == "0.3" {
opts = append(
opts,
a2av0.WithRESTTransport(a2av0.RESTTransportConfig{}),
a2av0.WithJSONRPCTransport(a2av0.JSONRPCTransportConfig{}),
a2agrpcv0.WithGRPCTransport(grpcOpts...),
)
}
return opts
}

func stripHTTPScheme(raw string) string {
Expand Down
5 changes: 5 additions & 0 deletions internal/cli/root.go
Original file line number Diff line number Diff line change
Expand Up @@ -50,10 +50,12 @@ type globalConfig struct {
url string
transports []string
svcParams *flagparse.ServiceParams
a2aVersion string
tenant string
timeout time.Duration
verbose bool
insecureGRPC bool
pretty bool
configPath string

bindings []clicfg.FlagBinding
Expand Down Expand Up @@ -107,6 +109,7 @@ func newRootCmd(cfg *globalConfig, deps deps) *cobra.Command {
default:
return fmt.Errorf("invalid --output %q (want text or json)", cfg.output)
}
cfg.Printer.PrettyJSONL = cfg.pretty
return nil
},
}
Expand All @@ -116,11 +119,13 @@ func newRootCmd(cfg *globalConfig, deps deps) *cobra.Command {
pf.VarP(&cfg.agentCard, "agent-card", "a", "Agent Card reference: host/origin, full card URL, or local file path")
pf.StringVarP(&cfg.url, "endpoint", "e", "", "Agent interface URL for a direct connection; skips card resolution and requires a single --transport flag")
pf.StringArrayVar(&cfg.transports, "transport", nil, "Transport preference: rest, jsonrpc, grpc (repeatable, highest preference first)")
pf.StringVar(&cfg.a2aVersion, "a2a-version", "", "Controls which a2a-protocol version client will advertise to the server.")
cfg.svcParams.Attach(pf)
pf.StringVar(&cfg.tenant, "tenant", "", "Tenant identifier")
pf.DurationVar(&cfg.timeout, "timeout", 30*time.Second, "Request timeout")
pf.BoolVarP(&cfg.verbose, "verbose", "v", false, "Verbose output to stderr")
pf.BoolVar(&cfg.insecureGRPC, "insecure", false, "Use insecure (plaintext) gRPC transport credentials")
pf.BoolVar(&cfg.pretty, "pretty", false, "Pretty-print (indent) streamed JSONL records instead of emitting one compact object per line")
pf.StringVar(&cfg.configPath, "config", "", "Load configuration from an explicit .env file in place of the local .env")

cmd.AddCommand(
Expand Down
24 changes: 12 additions & 12 deletions internal/cli/send.go
Original file line number Diff line number Diff line change
Expand Up @@ -31,15 +31,15 @@ import (
)

type sendFlags struct {
stream bool
async bool
payload string
taskID string
contextID string
history int
pollingInterval time.Duration
parts flagparse.Parts
meta flagparse.Metadata
stream bool
async bool
payload string
taskID string
contextID string
history int
pollInterval time.Duration
parts flagparse.Parts
meta flagparse.Metadata
}

type pollerFunc func(ctx context.Context, client *a2aclient.Client, req *a2a.SendMessageRequest, interval time.Duration) iter.Seq2[a2a.Event, error]
Expand Down Expand Up @@ -81,9 +81,9 @@ func newSendCmd(cfg *globalConfig, poller pollerFunc) *cobra.Command {
return utils.UnpackCause(ctx, err)
}

cfg.logf("falling back to polling (%v): %v", flags.pollingInterval, err)
cfg.logf("falling back to polling (%v): %v", flags.pollInterval, err)

for event, err := range poller(ctx, client, req, flags.pollingInterval) {
for event, err := range poller(ctx, client, req, flags.pollInterval) {
debounceTimeout()
if err := handleStreamEntry(cfg, event, err); err != nil {
return utils.UnpackCause(ctx, err)
Expand Down Expand Up @@ -113,7 +113,7 @@ func newSendCmd(cfg *globalConfig, poller pollerFunc) *cobra.Command {
f.StringVar(&flags.taskID, "task-id", "", "Task ID to continue an existing task")
f.StringVar(&flags.contextID, "context-id", "", "Context ID to group this turn under")
f.IntVar(&flags.history, "history", 0, "Request n history messages in the response")
f.DurationVar(&flags.pollingInterval, "polling-interval", 5*time.Second, "Duration between GetTask requests in polling fallback mode.")
f.DurationVar(&flags.pollInterval, "poll-interval", 2*time.Second, "Duration between GetTask requests in polling fallback mode.")
flags.parts.Attach(f)
flags.meta.Attach(f, "metadata", "Attach request metadata as a JSON object (repeatable)")

Expand Down
Loading