Skip to content

Commit efafba9

Browse files
authored
feat: support requesting a single tool (#56)
Adds support for getting a ToolsManifest with a single tool when a GET `/tools/$toolname` request is sent.
1 parent ed15418 commit efafba9

2 files changed

Lines changed: 129 additions & 4 deletions

File tree

‎internal/server/api.go‎

Lines changed: 24 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,19 +21,22 @@ import (
2121
"github.com/go-chi/chi/v5"
2222
"github.com/go-chi/chi/v5/middleware"
2323
"github.com/go-chi/render"
24+
"github.com/googleapis/genai-toolbox/internal/tools"
2425
)
2526

2627
// apiRouter creates a router that represents the routes under /api
2728
func apiRouter(s *Server) (chi.Router, error) {
2829
r := chi.NewRouter()
2930

3031
r.Use(middleware.AllowContentType("application/json"))
32+
r.Use(middleware.StripSlashes)
3133
r.Use(render.SetContentType(render.ContentTypeJSON))
3234

33-
r.Get("/toolset/", func(w http.ResponseWriter, r *http.Request) { toolsetHandler(s, w, r) })
35+
r.Get("/toolset", func(w http.ResponseWriter, r *http.Request) { toolsetHandler(s, w, r) })
3436
r.Get("/toolset/{toolsetName}", func(w http.ResponseWriter, r *http.Request) { toolsetHandler(s, w, r) })
3537

3638
r.Route("/tool/{toolName}", func(r chi.Router) {
39+
r.Get("/", func(w http.ResponseWriter, r *http.Request) { toolGetHandler(s, w, r) })
3740
r.Post("/invoke", func(w http.ResponseWriter, r *http.Request) { toolInvokeHandler(s, w, r) })
3841
})
3942

@@ -51,6 +54,26 @@ func toolsetHandler(s *Server, w http.ResponseWriter, r *http.Request) {
5154
render.JSON(w, r, toolset.Manifest)
5255
}
5356

57+
// toolGetHandler handles requests for a single Tool.
58+
func toolGetHandler(s *Server, w http.ResponseWriter, r *http.Request) {
59+
toolName := chi.URLParam(r, "toolName")
60+
tool, ok := s.tools[toolName]
61+
if !ok {
62+
err := fmt.Errorf("invalid tool name: tool with name %q does not exist", toolName)
63+
_ = render.Render(w, r, newErrResponse(err, http.StatusNotFound))
64+
return
65+
}
66+
// TODO: this can be optimized later with some caching
67+
m := tools.ToolsetManifest{
68+
ServerVersion: s.conf.Version,
69+
ToolsManifest: map[string]tools.Manifest{
70+
toolName: tool.Manifest(),
71+
},
72+
}
73+
74+
render.JSON(w, r, m)
75+
}
76+
5477
// toolInvokeHandler handles the API request to invoke a specific Tool.
5578
func toolInvokeHandler(s *Server, w http.ResponseWriter, r *http.Request) {
5679
toolName := chi.URLParam(r, "toolName")

‎internal/server/api_test.go‎

Lines changed: 105 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -64,15 +64,15 @@ func TestToolsetEndpoint(t *testing.T) {
6464
toolsets[name] = m
6565
}
6666

67-
server := Server{tools: toolsMap, toolsets: toolsets}
67+
server := Server{conf: ServerConfig{}, tools: toolsMap, toolsets: toolsets}
6868
r, err := apiRouter(&server)
6969
if err != nil {
70-
t.Fatalf("unable to initalize router: %s", err)
70+
t.Fatalf("unable to initialize router: %s", err)
7171
}
7272
ts := httptest.NewServer(r)
7373
defer ts.Close()
7474

75-
// wantRepsonse is a struct for checks against test cases
75+
// wantResponse is a struct for checks against test cases
7676
type wantResponse struct {
7777
statusCode int
7878
isErr bool
@@ -160,6 +160,108 @@ func TestToolsetEndpoint(t *testing.T) {
160160
})
161161
}
162162
}
163+
func TestToolGetEndpoint(t *testing.T) {
164+
// Set up resources to test against
165+
tool1 := MockTool{
166+
Name: "no_params",
167+
Params: []tools.Parameter{},
168+
}
169+
tool2 := MockTool{
170+
Name: "some_params",
171+
Params: tools.Parameters{
172+
tools.NewIntParameter("param1", "This is the first parameter."),
173+
tools.NewIntParameter("param2", "This is the second parameter."),
174+
},
175+
}
176+
toolsMap := map[string]tools.Tool{tool1.Name: tool1, tool2.Name: tool2}
177+
178+
server := Server{conf: ServerConfig{Version: "0.0.0"}, tools: toolsMap}
179+
r, err := apiRouter(&server)
180+
if err != nil {
181+
t.Fatalf("unable to initialize router: %s", err)
182+
}
183+
ts := httptest.NewServer(r)
184+
defer ts.Close()
185+
186+
// wantResponse is a struct for checks against test cases
187+
type wantResponse struct {
188+
statusCode int
189+
isErr bool
190+
version string
191+
tools []string
192+
}
193+
194+
testCases := []struct {
195+
name string
196+
toolName string
197+
want wantResponse
198+
}{
199+
{
200+
name: "tool1",
201+
toolName: tool1.Name,
202+
want: wantResponse{
203+
statusCode: http.StatusOK,
204+
version: "0.0.0",
205+
tools: []string{tool1.Name},
206+
},
207+
},
208+
{
209+
name: "tool2",
210+
toolName: tool2.Name,
211+
want: wantResponse{
212+
statusCode: http.StatusOK,
213+
version: "0.0.0",
214+
tools: []string{tool2.Name},
215+
},
216+
},
217+
{
218+
name: "invalid tool",
219+
toolName: "some_imaginary_tool",
220+
want: wantResponse{
221+
statusCode: http.StatusNotFound,
222+
isErr: true,
223+
},
224+
},
225+
}
226+
227+
for _, tc := range testCases {
228+
t.Run(tc.name, func(t *testing.T) {
229+
resp, body, err := testRequest(ts, http.MethodGet, fmt.Sprintf("/tool/%s", tc.toolName), nil)
230+
if err != nil {
231+
t.Fatalf("unexpected error during request: %s", err)
232+
}
233+
234+
if contentType := resp.Header.Get("Content-type"); contentType != "application/json" {
235+
t.Fatalf("unexpected content-type header: want %s, got %s", "application/json", contentType)
236+
}
237+
238+
if resp.StatusCode != tc.want.statusCode {
239+
t.Logf("response body: %s", body)
240+
t.Fatalf("unexpected status code: want %d, got %d", tc.want.statusCode, resp.StatusCode)
241+
}
242+
if tc.want.isErr {
243+
// skip the rest of the checks if this is an error case
244+
return
245+
}
246+
var m tools.ToolsetManifest
247+
err = json.Unmarshal(body, &m)
248+
if err != nil {
249+
t.Fatalf("unable to parse ToolsetManifest: %s", err)
250+
}
251+
// Check the version is correct
252+
if m.ServerVersion != tc.want.version {
253+
t.Fatalf("unexpected ServerVersion: want %q, got %q", tc.want.version, m.ServerVersion)
254+
}
255+
// validate that the tools in the toolset are correct
256+
for _, name := range tc.want.tools {
257+
_, ok := m.ToolsManifest[name]
258+
if !ok {
259+
t.Errorf("%q tool not found in manfiest", name)
260+
}
261+
}
262+
})
263+
}
264+
}
163265

164266
func testRequest(ts *httptest.Server, method, path string, body io.Reader) (*http.Response, []byte, error) {
165267
req, err := http.NewRequest(method, ts.URL+path, body)

0 commit comments

Comments
 (0)