131 lines
4.9 KiB
Go
131 lines
4.9 KiB
Go
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)
|
|
}
|
|
}
|