package playground import ( "context" "encoding/json" "errors" "io" "net/http" "net/http/httptest" "strings" "testing" ) // A fake OpenAI-compatible server: first call asks for a tool, second answers. func fakeProvider(t *testing.T, seen *[]map[string]any) *httptest.Server { t.Helper() calls := 0 return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/openai/v1/chat/completions" { t.Errorf("path = %s", r.URL.Path) } if r.Header.Get("Authorization") != "Bearer test-key" { t.Errorf("auth header = %q", r.Header.Get("Authorization")) } b, _ := io.ReadAll(r.Body) var body map[string]any _ = json.Unmarshal(b, &body) *seen = append(*seen, body) calls++ w.Header().Set("Content-Type", "application/json") if calls == 1 { _, _ = io.WriteString(w, `{"choices":[{"message":{"role":"assistant","content":null,"tool_calls":[ {"id":"call_1","type":"function","function":{"name":"get_booking_cache","arguments":"{\"booking_id\":5}"}}, {"id":"call_2","type":"function","function":{"name":"reassign_booking","arguments":"not json"}}]}, "finish_reason":"tool_calls"}],"usage":{"prompt_tokens":100,"completion_tokens":20}}`) return } _, _ = io.WriteString(w, `{"choices":[{"message":{"role":"assistant","content":"Proposed a reassign."},"finish_reason":"stop"}], "usage":{"prompt_tokens":150,"completion_tokens":10}}`) })) } func TestOpenAICompatRunsTheToolLoop(t *testing.T) { var seen []map[string]any srv := fakeProvider(t, &seen) defer srv.Close() m := NewOpenAICompat(srv.URL+"/openai/v1/", "test-key", "openai/gpt-oss-120b") plan := mustPlan(t, "EXCEPTION_AGENT", "stall_response") plan.Model = m.ModelName() execs := map[string]Executor{ "get_booking_cache": func(context.Context, json.RawMessage) (any, error) { return map[string]any{"bookingid": 5, "customername": "Ravi"}, nil }, } tr, err := Run(context.Background(), m, plan, "booking 5 is stuck", execs) if err != nil { t.Fatal(err) } if tr.Final != "Proposed a reassign." || tr.Turns != 2 || tr.Inputtokens != 250 || tr.Outputtokens != 30 { t.Fatalf("trace = %+v", tr) } if tr.Model != "openai/gpt-oss-120b" { t.Fatalf("model = %s", tr.Model) } if tr.Steps[0].Outcome != OutcomeExecuted || tr.Steps[1].Outcome != OutcomeProposed || string(tr.Steps[1].Input) != "{}" { t.Fatalf("steps = %+v", tr.Steps) } // First request: system + user, the model from config, tools as functions // with object schemas. first := seen[0] if first["model"] != "openai/gpt-oss-120b" || first["max_completion_tokens"].(float64) != MaxTokens { t.Fatalf("first request = %v", first) } msgs := first["messages"].([]any) if msgs[0].(map[string]any)["role"] != "system" || msgs[1].(map[string]any)["role"] != "user" { t.Fatalf("messages = %v", msgs) } tool := first["tools"].([]any)[0].(map[string]any) params := tool["function"].(map[string]any)["parameters"].(map[string]any) if tool["type"] != "function" || params["type"] != "object" || params["properties"] == nil { t.Fatalf("tool = %v", tool) } // Second request: the assistant's tool_calls echoed, then one tool message // per call, redacted. msgs = seen[1]["messages"].([]any) asst := msgs[2].(map[string]any) if asst["role"] != "assistant" || len(asst["tool_calls"].([]any)) != 2 { t.Fatalf("assistant echo = %v", asst) } res1, res2 := msgs[3].(map[string]any), msgs[4].(map[string]any) if res1["role"] != "tool" || res1["tool_call_id"] != "call_1" || res2["tool_call_id"] != "call_2" { t.Fatalf("tool results = %v %v", res1, res2) } if strings.Contains(res1["content"].(string), "Ravi") { t.Fatal("personal data reached the provider") } if !strings.Contains(res2["content"].(string), `"executed":false`) { t.Fatalf("write tool was not answered as a proposal: %v", res2["content"]) } } func TestOpenAICompatReportsProviderErrors(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusTooManyRequests) _, _ = io.WriteString(w, `{"error":{"message":"Rate limit reached for model"}}`) })) defer srv.Close() _, err := NewOpenAICompat(srv.URL, "secret-key-value", "m").Next(context.Background(), Request{Turns: []Turn{{Role: "user", Text: "hi"}}}) var pe *ProviderError if !errors.As(err, &pe) || pe.Status != 429 || !strings.Contains(pe.Message, "Rate limit") { t.Fatalf("err = %v", err) } if strings.Contains(err.Error(), "secret-key-value") { t.Fatal("the API key leaked into the error") } } func TestOpenAICompatStopWithToolCallsStillLoops(t *testing.T) { r, err := fromWire(oaResponse{Choices: []struct { Message oaMessage `json:"message"` FinishReason string `json:"finish_reason"` }{{Message: oaMessage{ToolCalls: []oaToolCall{{ID: "c", Function: oaFunctionCall{Name: "x", Arguments: "{}"}}}}, FinishReason: "stop"}}}) if err != nil || r.StopReason != "tool_use" { t.Fatalf("reply = %+v, err = %v", r, err) } }