Files
doormile_backend/internal/ai/playground/openai_compat_test.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)
}
}