diff --git a/a2aclient/client.go b/a2aclient/client.go index 3d50a976..2ba41968 100644 --- a/a2aclient/client.go +++ b/a2aclient/client.go @@ -16,10 +16,12 @@ package a2aclient import ( "context" + "fmt" "iter" "sync/atomic" "github.com/a2aproject/a2a-go/a2a" + "github.com/a2aproject/a2a-go/a2aext" "github.com/a2aproject/a2a-go/internal/utils" ) @@ -270,6 +272,26 @@ func (c *Client) GetAgentCard(ctx context.Context) (*a2a.AgentCard, error) { return resp, err } +func Invoke[Req, Resp any](ctx context.Context, client *Client, method *a2aext.UnaryClientMethod[Req, Resp], arg *Req) (*Resp, error) { + ctx, err := client.interceptBefore(ctx, method.Name(), arg) + if err != nil { + return nil, err + } + + resp, err := client.transport.Invoke(ctx, method, arg) + if errOverride := client.interceptAfter(ctx, method.Name(), resp, err); errOverride != nil { + return nil, errOverride + } + + typedResp, ok := resp.(*Resp) + if !ok { + var want Resp + return nil, fmt.Errorf("unexpected extension method result type %T, want pointer to %T", resp, want) + } + + return typedResp, err +} + func (c *Client) Destroy() error { return c.transport.Destroy() } diff --git a/a2aclient/grpc.go b/a2aclient/grpc.go index 2c92bf84..54f485a4 100644 --- a/a2aclient/grpc.go +++ b/a2aclient/grpc.go @@ -16,6 +16,7 @@ package a2aclient import ( "context" + "fmt" "io" "iter" "strings" @@ -24,10 +25,19 @@ import ( "google.golang.org/grpc/metadata" "github.com/a2aproject/a2a-go/a2a" + "github.com/a2aproject/a2a-go/a2aext" "github.com/a2aproject/a2a-go/a2apb" "github.com/a2aproject/a2a-go/a2apb/pbconv" ) +type grpcExtensionBinding struct { + makeCall func(ctx context.Context, req any) (any, error) +} + +func (g *grpcExtensionBinding) Protocol() a2a.TransportProtocol { + return a2a.TransportProtocolGRPC +} + // WithGRPCTransport create a gRPC transport implementation which will use the provided [grpc.DialOption]s during connection establishment. func WithGRPCTransport(opts ...grpc.DialOption) FactoryOption { return WithTransport( @@ -238,6 +248,20 @@ func (c *grpcTransport) GetAgentCard(ctx context.Context) (*a2a.AgentCard, error return pbconv.FromProtoAgentCard(pCard) } +func (t *grpcTransport) Invoke(ctx context.Context, method a2aext.Method, req any) (any, error) { + binding, ok := method.Binding(a2a.TransportProtocolGRPC) + if !ok { + return nil, fmt.Errorf("method %s is not bound to gRPC", method.Name()) + } + + typedBinding, ok := binding.(*grpcExtensionBinding) + if !ok { + return nil, fmt.Errorf("method %s is not bound to JSON-RPC", method.Name()) + } + + return typedBinding.makeCall(ctx, req) +} + func (c *grpcTransport) Destroy() error { return c.closeConnFn() } diff --git a/a2aclient/jsonrpc.go b/a2aclient/jsonrpc.go index bdf8bbe0..2a39f575 100644 --- a/a2aclient/jsonrpc.go +++ b/a2aclient/jsonrpc.go @@ -25,6 +25,7 @@ import ( "time" "github.com/a2aproject/a2a-go/a2a" + "github.com/a2aproject/a2a-go/a2aext" "github.com/a2aproject/a2a-go/internal/jsonrpc" "github.com/a2aproject/a2a-go/internal/sse" "github.com/a2aproject/a2a-go/log" @@ -47,6 +48,28 @@ type jsonrpcResponse struct { Error *jsonrpc.Error `json:"error,omitempty"` } +type jsonrpcExtensionBinding struct { + Method string + ResponseFromBytes func(resp []byte) (any, error) +} + +func NewJSONRPCExtensionBinding[Resp any](name string) a2aext.Binding { + return &jsonrpcExtensionBinding{ + Method: name, + ResponseFromBytes: func(resp []byte) (any, error) { + var zero Resp + if err := json.Unmarshal(resp, &zero); err != nil { + return nil, err + } + return &zero, nil + }, + } +} + +func (*jsonrpcExtensionBinding) Protocol() a2a.TransportProtocol { + return a2a.TransportProtocolJSONRPC +} + // JSONRPCOption configures optional parameters for the JSONRPC transport. // Options are applied during NewJSONRPCTransport initialization. type JSONRPCOption func(*jsonrpcTransport) @@ -351,6 +374,29 @@ func (t *jsonrpcTransport) DeleteTaskPushConfig(ctx context.Context, params *a2a return err } +func (t *jsonrpcTransport) Invoke(ctx context.Context, method a2aext.Method, req any) (any, error) { + binding, ok := method.Binding(a2a.TransportProtocolJSONRPC) + if !ok { + return nil, fmt.Errorf("method %s is not bound to JSON-RPC", method.Name()) + } + + typedBinding, ok := binding.(*jsonrpcExtensionBinding) + if !ok { + return nil, fmt.Errorf("method %s is not bound to JSON-RPC", method.Name()) + } + + result, err := t.sendRequest(ctx, typedBinding.Method, req) + if err != nil { + return nil, err + } + + response, err := typedBinding.ResponseFromBytes(result) + if err != nil { + return nil, fmt.Errorf("result violates A2A spec - could not determine type: %w; data: %s", err, string(result)) + } + return response, nil +} + // GetAgentCard retrieves the agent's card. func (t *jsonrpcTransport) GetAgentCard(ctx context.Context) (*a2a.AgentCard, error) { result, err := t.sendRequest(ctx, jsonrpc.MethodGetExtendedAgentCard, nil) diff --git a/a2aclient/transport.go b/a2aclient/transport.go index 41b69912..6e0525c8 100644 --- a/a2aclient/transport.go +++ b/a2aclient/transport.go @@ -20,6 +20,7 @@ import ( "iter" "github.com/a2aproject/a2a-go/a2a" + "github.com/a2aproject/a2a-go/a2aext" ) // A2AClient defines a transport-agnostic interface for making A2A requests. @@ -56,6 +57,9 @@ type Transport interface { // If extended card is supported calls the 'agent/getAuthenticatedExtendedCard' protocol method. GetAgentCard(ctx context.Context) (*a2a.AgentCard, error) + // Invoke calls the provided extension method. + Invoke(ctx context.Context, client a2aext.Method, arg any) (any, error) + // Clean up resources associated with the transport (eg. close a gRPC channel). Destroy() error } @@ -120,6 +124,10 @@ func (unimplementedTransport) GetAgentCard(ctx context.Context) (*a2a.AgentCard, return nil, errNotImplemented } +func (unimplementedTransport) Invoke(ctx context.Context, client a2aext.Method, arg any) (any, error) { + return nil, errNotImplemented +} + func (unimplementedTransport) Destroy() error { return nil } diff --git a/a2aext/method.go b/a2aext/method.go new file mode 100644 index 00000000..f55492cb --- /dev/null +++ b/a2aext/method.go @@ -0,0 +1,92 @@ +package a2aext + +import ( + "context" + "errors" + "iter" + + "github.com/a2aproject/a2a-go/a2a" +) + +type Binding interface { + Protocol() a2a.TransportProtocol +} + +type Method interface { + Name() string + Binding(a2a.TransportProtocol) (Binding, bool) + Streaming() bool +} + +type ServerMethod interface { + Method + + InvokeUnary(ctx context.Context, arg any) (any, error) + InvokeStreaming(ctx context.Context, arg any) iter.Seq2[any, error] +} + +type methodDescriptor struct { + name string + bindings map[a2a.TransportProtocol]Binding + streaming bool +} + +func (d *methodDescriptor) Name() string { return d.name } + +func (d *methodDescriptor) Binding(protocol a2a.TransportProtocol) (Binding, bool) { + b, ok := d.bindings[protocol] + return b, ok +} + +func (d *methodDescriptor) Streaming() bool { return d.streaming } + +type UnaryClientMethod[Req, Resp any] struct { + methodDescriptor +} + +var _ Method = (*UnaryClientMethod[any, any])(nil) + +// NewUnaryClientMethod creates a new unary extension method definition. +func NewUnaryClientMethod[Arg, Res any]( + name string, + bindings ...Binding, +) *UnaryClientMethod[Arg, Res] { + bmap := map[a2a.TransportProtocol]Binding{} + for _, b := range bindings { + bmap[b.Protocol()] = b + } + return &UnaryClientMethod[Arg, Res]{ + methodDescriptor: methodDescriptor{name: name, bindings: bmap, streaming: false}, + } +} + +type unaryServerMethod[Arg, Res any] struct { + methodDescriptor + call func(context.Context, *Arg) (*Res, error) +} + +// NewUnaryMethod creates a new unary extension method definition. +func NewUnaryServerMethod[Arg, Res any]( + name string, + call func(context.Context, *Arg) (*Res, error), + bindings ...Binding, +) ServerMethod { + bmap := map[a2a.TransportProtocol]Binding{} + for _, b := range bindings { + bmap[b.Protocol()] = b + } + return &unaryServerMethod[Arg, Res]{ + methodDescriptor: methodDescriptor{name: name, bindings: bmap, streaming: false}, + call: call, + } +} + +func (m *unaryServerMethod[Arg, Res]) InvokeUnary(ctx context.Context, arg any) (any, error) { + return m.call(ctx, arg.(*Arg)) +} + +func (m *unaryServerMethod[Arg, Res]) InvokeStreaming(ctx context.Context, arg any) iter.Seq2[any, error] { + return func(yield func(any, error) bool) { + yield(nil, errors.New("unary method invoked as streaming")) + } +} diff --git a/a2asrv/handler.go b/a2asrv/handler.go index 4b32c796..a4123f74 100644 --- a/a2asrv/handler.go +++ b/a2asrv/handler.go @@ -21,6 +21,7 @@ import ( "log/slog" "github.com/a2aproject/a2a-go/a2a" + "github.com/a2aproject/a2a-go/a2aext" "github.com/a2aproject/a2a-go/a2asrv/eventqueue" "github.com/a2aproject/a2a-go/a2asrv/limiter" "github.com/a2aproject/a2a-go/a2asrv/push" @@ -59,6 +60,15 @@ type RequestHandler interface { // GetAgentCard returns an extended a2a.AgentCard if configured. OnGetExtendedAgentCard(ctx context.Context) (*a2a.AgentCard, error) + + // GetAgentCard returns an extended a2a.AgentCard if configured. + GetExtensionMethods() []a2aext.ServerMethod + + // Invoke calls a registered extension method. + Invoke(ctx context.Context, method a2aext.ServerMethod, arg any) (any, error) + + // InvokeStreaming calls a registered streaming extension method. + InvokeStreaming(ctx context.Context, method a2aext.ServerMethod, arg any) iter.Seq2[any, error] } // Implements a2asrv.RequestHandler. @@ -75,6 +85,8 @@ type defaultRequestHandler struct { reqContextInterceptors []RequestContextInterceptor authenticatedCardProducer AgentCardProducer + + extensionMethods []a2aext.ServerMethod } // RequestHandlerOption can be used to customize the default [RequestHandler] implementation behavior. @@ -104,12 +116,20 @@ func WithConcurrencyConfig(config limiter.ConcurrencyConfig) RequestHandlerOptio } } +// WithConcurrencyConfig allows to set limits on the number of concurrent executions. +func WithExtensionMethod(method a2aext.ServerMethod) RequestHandlerOption { + return func(ih *InterceptedHandler, h *defaultRequestHandler) { + h.extensionMethods = append(h.extensionMethods, method) + } +} + // NewHandler creates a new request handler. func NewHandler(executor AgentExecutor, options ...RequestHandlerOption) RequestHandler { h := &defaultRequestHandler{ - agentExecutor: executor, - queueManager: eventqueue.NewInMemoryManager(), - taskStore: taskstore.NewMem(), + agentExecutor: executor, + extensionMethods: []a2aext.ServerMethod{}, + queueManager: eventqueue.NewInMemoryManager(), + taskStore: taskstore.NewMem(), // push notifications are not supported by default } ih := &InterceptedHandler{Handler: h, Logger: slog.Default()} @@ -307,6 +327,18 @@ func (h *defaultRequestHandler) OnGetExtendedAgentCard(ctx context.Context) (*a2 return h.authenticatedCardProducer.Card(ctx) } +func (h *defaultRequestHandler) GetExtensionMethods() []a2aext.ServerMethod { + return h.extensionMethods +} + +func (h *defaultRequestHandler) Invoke(ctx context.Context, method a2aext.ServerMethod, arg any) (any, error) { + return method.InvokeUnary(ctx, arg) +} + +func (h *defaultRequestHandler) InvokeStreaming(ctx context.Context, method a2aext.ServerMethod, arg any) iter.Seq2[any, error] { + return method.InvokeStreaming(ctx, arg) +} + func shouldInterruptNonStreaming(params *a2a.MessageSendParams, event a2a.Event) (a2a.TaskID, bool) { // Non-blocking clients receive a result on the first task event, default Blocking to TRUE if params.Config != nil && params.Config.Blocking != nil && !(*params.Config.Blocking) { diff --git a/a2asrv/intercepted_handler.go b/a2asrv/intercepted_handler.go index 70471795..2cc1ad30 100644 --- a/a2asrv/intercepted_handler.go +++ b/a2asrv/intercepted_handler.go @@ -22,6 +22,7 @@ import ( "github.com/google/uuid" "github.com/a2aproject/a2a-go/a2a" + "github.com/a2aproject/a2a-go/a2aext" "github.com/a2aproject/a2a-go/log" ) @@ -225,6 +226,45 @@ func (h *InterceptedHandler) OnGetExtendedAgentCard(ctx context.Context) (*a2a.A return response, err } +func (h *InterceptedHandler) GetExtensionMethods() []a2aext.ServerMethod { + return h.Handler.GetExtensionMethods() +} + +func (h *InterceptedHandler) Invoke(ctx context.Context, method a2aext.ServerMethod, arg any) (any, error) { + ctx, callCtx := withMethodCallContext(ctx, method.Name()) + ctx = h.withLoggerContext(ctx) + ctx, err := h.interceptBefore(ctx, callCtx, arg) + if err != nil { + return nil, err + } + response, err := h.Handler.Invoke(ctx, method, arg) + if errOverride := h.interceptAfter(ctx, callCtx, response, err); errOverride != nil { + return nil, errOverride + } + return response, err +} + +func (h *InterceptedHandler) InvokeStreaming(ctx context.Context, method a2aext.ServerMethod, arg any) iter.Seq2[any, error] { + return func(yield func(any, error) bool) { + ctx, callCtx := withMethodCallContext(ctx, method.Name()) + ctx = h.withLoggerContext(ctx) + ctx, err := h.interceptBefore(ctx, callCtx, arg) + if err != nil { + yield(nil, err) + return + } + for event, err := range h.Handler.InvokeStreaming(ctx, method, arg) { + if errOverride := h.interceptAfter(ctx, callCtx, event, err); errOverride != nil { + yield(nil, errOverride) + return + } + if !yield(event, err) { + return + } + } + } +} + func (h *InterceptedHandler) interceptBefore(ctx context.Context, callCtx *CallContext, payload any) (context.Context, error) { request := &Request{Payload: payload} diff --git a/a2asrv/jsonrpc.go b/a2asrv/jsonrpc.go index 3d5d9a35..a7a2e2b5 100644 --- a/a2asrv/jsonrpc.go +++ b/a2asrv/jsonrpc.go @@ -22,6 +22,7 @@ import ( "net/http" "github.com/a2aproject/a2a-go/a2a" + "github.com/a2aproject/a2a-go/a2aext" "github.com/a2aproject/a2a-go/internal/jsonrpc" "github.com/a2aproject/a2a-go/internal/sse" "github.com/a2aproject/a2a-go/log" @@ -44,13 +45,58 @@ type jsonrpcResponse struct { Error *jsonrpc.Error `json:"error,omitempty"` } +type jsonrpcExtensionBinding struct { + method string + mapRequest func(json.RawMessage) (any, error) + handler RequestHandler + serverMethod a2aext.ServerMethod +} + +func (b *jsonrpcExtensionBinding) Protocol() a2a.TransportProtocol { + return a2a.TransportProtocolJSONRPC +} + +func (b *jsonrpcExtensionBinding) init(h RequestHandler, method a2aext.ServerMethod) { + b.handler = h + b.serverMethod = method +} + +func NewJSONRPCMethodBinding[Req any](method string) a2aext.Binding { + return &jsonrpcExtensionBinding{ + method: method, + mapRequest: func(params json.RawMessage) (any, error) { + var req Req + if err := json.Unmarshal(params, &req); err != nil { + return nil, err + } + return &req, nil + }, + } +} + type jsonrpcHandler struct { - handler RequestHandler + handler RequestHandler + methodBinding map[string]*jsonrpcExtensionBinding } // NewJSONRPCHandler creates an [http.Handler] implementation for serving A2A-protocol over JSONRPC. func NewJSONRPCHandler(handler RequestHandler) http.Handler { - return &jsonrpcHandler{handler: handler} + transport := &jsonrpcHandler{handler: handler, methodBinding: make(map[string]*jsonrpcExtensionBinding)} + + for _, method := range handler.GetExtensionMethods() { + binding, ok := method.Binding(a2a.TransportProtocolJSONRPC) + if !ok { + continue + } + typedBinding, ok := binding.(*jsonrpcExtensionBinding) + if !ok { // fail? + continue + } + typedBinding.init(handler, method) + transport.methodBinding[typedBinding.method] = typedBinding + } + + return transport } func (h *jsonrpcHandler) ServeHTTP(rw http.ResponseWriter, req *http.Request) { @@ -106,7 +152,7 @@ func (h *jsonrpcHandler) handleRequest(ctx context.Context, rw http.ResponseWrit case jsonrpc.MethodGetExtendedAgentCard: result, err = h.onGetAgentCard(ctx) default: - err = a2a.ErrMethodNotFound + result, err = h.onUnaryExtensionMethod(ctx, req.Method, req.Params) } if err != nil { @@ -295,6 +341,18 @@ func (h *jsonrpcHandler) onGetAgentCard(ctx context.Context) (*a2a.AgentCard, er return h.handler.OnGetExtendedAgentCard(ctx) } +func (h *jsonrpcHandler) onUnaryExtensionMethod(ctx context.Context, method string, raw json.RawMessage) (any, error) { + binding, ok := h.methodBinding[method] + if !ok { + return nil, a2a.ErrMethodNotFound + } + mapped, err := binding.mapRequest(raw) + if err != nil { + return nil, newParseError(err) + } + return binding.handler.Invoke(ctx, binding.serverMethod, mapped) +} + func newParseError(cause error) error { return fmt.Errorf("%w: %w", a2a.ErrParseError, cause) } diff --git a/examples/tasksearch/client/main.go b/examples/tasksearch/client/main.go new file mode 100644 index 00000000..978dc590 --- /dev/null +++ b/examples/tasksearch/client/main.go @@ -0,0 +1,64 @@ +// Copyright 2025 The A2A Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package main + +import ( + "context" + "encoding/json" + "flag" + "log" + + "github.com/a2aproject/a2a-go/a2aclient" + "github.com/a2aproject/a2a-go/a2aclient/agentcard" + "github.com/a2aproject/a2a-go/examples/tasksearch/extension" + + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" +) + +var ( + cardURL = flag.String("card-url", "http://127.0.0.1:9001", "Base URL of AgentCard server.") + query = flag.String("query", "", "Query to search for tasks.") +) + +func main() { + flag.Parse() + ctx := context.Background() + card, err := agentcard.DefaultResolver.Resolve(ctx, *cardURL) + if err != nil { + log.Fatalf("Failed to resolve an AgentCard: %v", err) + } + + withInsecureGRPC := a2aclient.WithGRPCTransport(grpc.WithTransportCredentials(insecure.NewCredentials())) + client, err := a2aclient.NewFromCard(ctx, card, withInsecureGRPC) + if err != nil { + log.Fatalf("Failed to create a client: %v", err) + } + + search, err := a2aclient.Invoke( + ctx, + client, + tasksearchext.ClientTaskSearch, + &tasksearchext.Request{Query: *query}, + ) + if err != nil { + log.Fatalf("Failed to invoke a method: %v", err) + } + str, err := json.MarshalIndent(search, "", " ") + if err != nil { + log.Fatalf("Failed to marshal a response: %v", err) + } + log.Printf("Server responded with: %s", str) +} diff --git a/examples/tasksearch/extension/client.go b/examples/tasksearch/extension/client.go new file mode 100644 index 00000000..76078881 --- /dev/null +++ b/examples/tasksearch/extension/client.go @@ -0,0 +1,11 @@ +package tasksearchext + +import ( + "github.com/a2aproject/a2a-go/a2aclient" + "github.com/a2aproject/a2a-go/a2aext" +) + +var ClientTaskSearch = a2aext.NewUnaryClientMethod[Request, Response]( + MethodName, + a2aclient.NewJSONRPCExtensionBinding[Response](JSONRPCMethodName), +) diff --git a/examples/tasksearch/extension/common.go b/examples/tasksearch/extension/common.go new file mode 100644 index 00000000..0cb52313 --- /dev/null +++ b/examples/tasksearch/extension/common.go @@ -0,0 +1,23 @@ +package tasksearchext + +import "github.com/a2aproject/a2a-go/a2a" + +var URI = "https://v1.tasksearchext.example.com" + +var MethodName = "SearchTasks" + +var JSONRPCMethodName = "SearchTasks" + +type Request struct { + Query string `json:"query"` +} + +type Response struct { + Tasks []*a2a.Task `json:"tasks"` +} + +var Definition = a2a.AgentExtension{ + URI: URI, + Description: "Example method extension, helps to search for tasks on the A2A server.", + Required: false, +} diff --git a/examples/tasksearch/extension/server.go b/examples/tasksearch/extension/server.go new file mode 100644 index 00000000..d6eaaffc --- /dev/null +++ b/examples/tasksearch/extension/server.go @@ -0,0 +1,59 @@ +package tasksearchext + +import ( + "context" + "fmt" + "strings" + + "github.com/a2aproject/a2a-go/a2a" + "github.com/a2aproject/a2a-go/a2aext" + "github.com/a2aproject/a2a-go/a2asrv" +) + +type ServerExt struct { + Method a2aext.ServerMethod +} + +func NewForServer(store a2asrv.TaskStore) *ServerExt { + method := a2aext.NewUnaryServerMethod( + MethodName, + func(ctx context.Context, req *Request) (*Response, error) { + tasks, err := store.List(ctx, &a2a.ListTasksRequest{}) + if err != nil { + return nil, fmt.Errorf("failed to list tasks: %w", err) + } + + var filteredTasks []*a2a.Task + + tasksScan: + for _, task := range tasks.Tasks { + if task.Status.Message != nil && partsContain(task.Status.Message.Parts, req.Query) { + filteredTasks = append(filteredTasks, task) + continue + } + for _, artifact := range task.Artifacts { + if partsContain(artifact.Parts, req.Query) { + filteredTasks = append(filteredTasks, task) + continue tasksScan + } + } + } + + return &Response{Tasks: filteredTasks}, nil + }, + a2asrv.NewJSONRPCMethodBinding[Request](JSONRPCMethodName), + ) + return &ServerExt{Method: method} +} + +func partsContain(parts []a2a.Part, text string) bool { + lowerText := strings.ToLower(text) + for _, part := range parts { + if tp, ok := part.(a2a.TextPart); ok { + if strings.Contains(strings.ToLower(tp.Text), lowerText) { + return true + } + } + } + return false +} diff --git a/examples/tasksearch/server/main.go b/examples/tasksearch/server/main.go new file mode 100644 index 00000000..c194c7cd --- /dev/null +++ b/examples/tasksearch/server/main.go @@ -0,0 +1,118 @@ +// Copyright 2025 The A2A Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package main + +import ( + "context" + "flag" + "fmt" + "log" + "net" + "net/http" + + "github.com/a2aproject/a2a-go/a2a" + "github.com/a2aproject/a2a-go/a2asrv" + "github.com/a2aproject/a2a-go/a2asrv/eventqueue" + tasksearchext "github.com/a2aproject/a2a-go/examples/tasksearch/extension" + "github.com/a2aproject/a2a-go/internal/taskstore" + a2alog "github.com/a2aproject/a2a-go/log" +) + +var ( + port = flag.Int("port", 9001, "Port for an A2A server to listen on.") +) + +var userName = "user" + +type agentExecutor struct{} + +func (*agentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.RequestContext, q eventqueue.Queue) error { + return q.Write(ctx, a2a.NewMessage(a2a.MessageRoleAgent, a2a.TextPart{Text: "Hello, world!"})) +} + +func (*agentExecutor) Cancel(ctx context.Context, reqCtx *a2asrv.RequestContext, q eventqueue.Queue) error { + return nil +} + +type authInterceptor struct { +} + +func (i *authInterceptor) Before(ctx context.Context, callCtx *a2asrv.CallContext, req *a2asrv.Request) (context.Context, error) { + callCtx.User = &a2asrv.AuthenticatedUser{UserName: userName} + return ctx, nil +} + +func (i *authInterceptor) After(ctx context.Context, callCtx *a2asrv.CallContext, resp *a2asrv.Response) error { + a2alog.Info(ctx, "request served", "method", callCtx.Method()) + return nil +} + +func main() { + flag.Parse() + + agentCard := &a2a.AgentCard{ + Name: "Task Search Extension Host", + URL: fmt.Sprintf("http://127.0.0.1:%d/invoke", *port), + PreferredTransport: a2a.TransportProtocolJSONRPC, + } + + listener, err := net.Listen("tcp", fmt.Sprintf(":%d", *port)) + if err != nil { + log.Fatalf("Failed to bind to a port: %v", err) + } + log.Printf("Starting a JSONRPC server on 127.0.0.1:%d", *port) + + store := setupTaskStore(context.Background()) + taskSearchExt := tasksearchext.NewForServer(store) + requestHandler := a2asrv.NewHandler( + &agentExecutor{}, + a2asrv.WithTaskStore(store), + a2asrv.WithCallInterceptor(&authInterceptor{}), + a2asrv.WithExtensionMethod(taskSearchExt.Method), + ) + + mux := http.NewServeMux() + mux.Handle("/invoke", a2asrv.NewJSONRPCHandler(requestHandler)) + mux.Handle(a2asrv.WellKnownAgentCardPath, a2asrv.NewStaticAgentCardHandler(agentCard)) + + err = http.Serve(listener, mux) + + log.Printf("Server stopped: %v", err) +} + +func setupTaskStore(ctx context.Context) a2asrv.TaskStore { + store := taskstore.NewMem(taskstore.WithAuthenticator(func(ctx context.Context) (taskstore.UserName, bool) { + if callCtx, ok := a2asrv.CallContextFrom(ctx); ok { + return taskstore.UserName(callCtx.User.Name()), true + } + return "", false + })) + authenticatedCtx, callCtx := a2asrv.WithCallContext(ctx, nil) + callCtx.User = &a2asrv.AuthenticatedUser{UserName: userName} + for _, status := range []string{"Hello, world!", "Foo", "FooBar"} { + task := &a2a.Task{ + ID: a2a.NewTaskID(), + ContextID: a2a.NewContextID(), + Status: a2a.TaskStatus{ + State: a2a.TaskStateCompleted, + Message: a2a.NewMessage(a2a.MessageRoleAgent, a2a.TextPart{Text: status}), + }, + } + if err := store.Save(authenticatedCtx, task); err != nil { + log.Fatalf("Failed to save a task: %v", err) + } + } + return store +}