572 lines
19 KiB
Go
572 lines
19 KiB
Go
package llm
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"gitea.maximumdirect.net/eric/audita/internal/framework/contracts"
|
|
"gitea.maximumdirect.net/eric/audita/internal/framework/responseschema"
|
|
)
|
|
|
|
func TestNewOpenAICompatibleClientValidation(t *testing.T) {
|
|
_, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: " ",
|
|
Model: "model",
|
|
})
|
|
if err == nil || !strings.Contains(err.Error(), "base URL") {
|
|
t.Fatalf("expected base URL validation error, got %v", err)
|
|
}
|
|
|
|
_, err = NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: "https://example.test/v1",
|
|
Model: " ",
|
|
})
|
|
if err == nil || !strings.Contains(err.Error(), "model") {
|
|
t.Fatalf("expected model validation error, got %v", err)
|
|
}
|
|
|
|
_, err = NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: "https://example.test/v1",
|
|
Model: "model",
|
|
MaxRetries: -1,
|
|
})
|
|
if err == nil || !strings.Contains(err.Error(), "max retries") {
|
|
t.Fatalf("expected max retries validation error, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientRequestShapeAndDecode(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
var seenPath string
|
|
var seenAuthorization string
|
|
var seenReq map[string]any
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
seenPath = r.URL.Path
|
|
seenAuthorization = r.Header.Get("Authorization")
|
|
if err := json.NewDecoder(r.Body).Decode(&seenReq); err != nil {
|
|
t.Fatalf("decode request: %v", err)
|
|
}
|
|
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{
|
|
"model":"provider-model",
|
|
"choices":[{"message":{"content":"{\"name\":\"Robby\",\"age\":22}"}}],
|
|
"usage":{"prompt_tokens":11,"completion_tokens":7,"total_tokens":18}
|
|
}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL + "/v1",
|
|
Model: "test-model",
|
|
APIKey: "secret-key",
|
|
MaxRetries: 1,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
type person struct {
|
|
Name string `json:"name"`
|
|
Age int `json:"age"`
|
|
}
|
|
var out person
|
|
resp, err := client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err != nil {
|
|
t.Fatalf("CompleteStructured: %v", err)
|
|
}
|
|
|
|
if seenPath != "/v1/chat/completions" {
|
|
t.Fatalf("unexpected request path: %q", seenPath)
|
|
}
|
|
if seenAuthorization != "Bearer secret-key" {
|
|
t.Fatalf("unexpected authorization header: %q", seenAuthorization)
|
|
}
|
|
|
|
responseFormat, ok := seenReq["response_format"].(map[string]any)
|
|
if !ok {
|
|
t.Fatalf("expected response_format object, got %T", seenReq["response_format"])
|
|
}
|
|
if responseFormat["type"] != "json_schema" {
|
|
t.Fatalf("unexpected response_format.type: %v", responseFormat["type"])
|
|
}
|
|
jsonSchema, ok := responseFormat["json_schema"].(map[string]any)
|
|
if !ok {
|
|
t.Fatalf("expected response_format.json_schema object, got %T", responseFormat["json_schema"])
|
|
}
|
|
if jsonSchema["name"] != schema.Name {
|
|
t.Fatalf("unexpected response schema name: %v", jsonSchema["name"])
|
|
}
|
|
if jsonSchema["strict"] != true {
|
|
t.Fatalf("expected strict=true, got %v", jsonSchema["strict"])
|
|
}
|
|
if _, ok := jsonSchema["schema"].(map[string]any); !ok {
|
|
t.Fatalf("expected embedded JSON schema object, got %T", jsonSchema["schema"])
|
|
}
|
|
|
|
if out.Name != "Robby" || out.Age != 22 {
|
|
t.Fatalf("unexpected decoded output: %+v", out)
|
|
}
|
|
if resp.Provider != "openai-compatible" {
|
|
t.Fatalf("unexpected provider metadata: %q", resp.Provider)
|
|
}
|
|
if resp.Model != "provider-model" {
|
|
t.Fatalf("unexpected model metadata: %q", resp.Model)
|
|
}
|
|
if resp.PromptTokens != 11 || resp.CompletionTokens != 7 || resp.TotalTokens != 18 {
|
|
t.Fatalf("unexpected token metadata: %+v", resp)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientDecodesCorrectionSetStructuredResponse(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{
|
|
"choices":[{"message":{"content":"{\"corrections\":[{\"id\":1,\"original_text\":\"teh\",\"corrected_text\":\"the\",\"confidence\":0.9}]}"}}]
|
|
}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
type correction struct {
|
|
TargetSegmentID int `json:"id"`
|
|
OriginalText string `json:"original_text"`
|
|
CorrectedText string `json:"corrected_text"`
|
|
Confidence float64 `json:"confidence"`
|
|
}
|
|
type correctionSet struct {
|
|
Corrections []correction `json:"corrections"`
|
|
}
|
|
var out correctionSet
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err != nil {
|
|
t.Fatalf("CompleteStructured: %v", err)
|
|
}
|
|
if len(out.Corrections) != 1 {
|
|
t.Fatalf("expected one correction, got %+v", out.Corrections)
|
|
}
|
|
if out.Corrections[0].TargetSegmentID != 1 || out.Corrections[0].CorrectedText != "the" {
|
|
t.Fatalf("unexpected correction payload: %+v", out.Corrections[0])
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientDecodesValidatorDecisionSetStructuredResponse(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.ValidatorDecisionSetKey)
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{
|
|
"choices":[{"message":{"content":"{\"validations\":[{\"correction_index\":0,\"approved\":true,\"confidence\":0.95,\"reason\":\"ok\"}]}"}}]
|
|
}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
type validationDecision struct {
|
|
CorrectionIndex int `json:"correction_index"`
|
|
Approved bool `json:"approved"`
|
|
Confidence float64 `json:"confidence"`
|
|
Reason string `json:"reason"`
|
|
}
|
|
type validationResponse struct {
|
|
Validations []validationDecision `json:"validations"`
|
|
}
|
|
var out validationResponse
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err != nil {
|
|
t.Fatalf("CompleteStructured: %v", err)
|
|
}
|
|
if len(out.Validations) != 1 {
|
|
t.Fatalf("expected one validation decision, got %+v", out.Validations)
|
|
}
|
|
if out.Validations[0].CorrectionIndex != 0 || !out.Validations[0].Approved {
|
|
t.Fatalf("unexpected validation payload: %+v", out.Validations[0])
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientNoAuthorizationHeaderWithoutAPIKey(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
var seenAuthorization string
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
seenAuthorization = r.Header.Get("Authorization")
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"ok\":true}"}}]}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
var out map[string]any
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err != nil {
|
|
t.Fatalf("CompleteStructured: %v", err)
|
|
}
|
|
if seenAuthorization != "" {
|
|
t.Fatalf("expected empty Authorization header, got %q", seenAuthorization)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientMalformedJSONFailsSafely(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{"}}]}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
MaxRetries: 0,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
var out map[string]any
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err == nil || !strings.Contains(err.Error(), "decode structured output") {
|
|
t.Fatalf("expected decode error, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientMissingRequiredFieldsFailsSafely(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{"model":"x","choices":[]}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
MaxRetries: 0,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
var out map[string]any
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err == nil || !strings.Contains(err.Error(), "missing choices") {
|
|
t.Fatalf("expected missing-field error, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientUnknownExtraFieldsFollowLocalDecoderPolicy(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{
|
|
"choices":[{"message":{"content":"{\"corrections\":[{\"id\":1,\"original_text\":\"teh\",\"corrected_text\":\"the\",\"confidence\":0.9,\"extra\":\"ignored\"}],\"top_extra\":true}"}}]
|
|
}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
type correction struct {
|
|
TargetSegmentID int `json:"id"`
|
|
OriginalText string `json:"original_text"`
|
|
CorrectedText string `json:"corrected_text"`
|
|
Confidence float64 `json:"confidence"`
|
|
}
|
|
type correctionSet struct {
|
|
Corrections []correction `json:"corrections"`
|
|
}
|
|
var out correctionSet
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err != nil {
|
|
t.Fatalf("expected unknown extra fields to be ignored by local decoder, got %v", err)
|
|
}
|
|
if len(out.Corrections) != 1 || out.Corrections[0].CorrectedText != "the" {
|
|
t.Fatalf("unexpected decoded payload: %+v", out.Corrections)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientProviderErrorRedactsSecret(t *testing.T) {
|
|
secret := "super-secret-key"
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.WriteHeader(http.StatusUnauthorized)
|
|
_, _ = w.Write([]byte(`{"error":{"message":"Authorization failed for Bearer super-secret-key"}}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
APIKey: secret,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
var out map[string]any
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err == nil {
|
|
t.Fatalf("expected provider error")
|
|
}
|
|
if strings.Contains(err.Error(), secret) {
|
|
t.Fatalf("error leaked secret: %v", err)
|
|
}
|
|
if !strings.Contains(err.Error(), "[REDACTED]") {
|
|
t.Fatalf("expected redaction marker in error: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientRequestErrorRedactsSecret(t *testing.T) {
|
|
secret := "super-secret-key"
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: "https://example.test/v1",
|
|
Model: "test-model",
|
|
APIKey: secret,
|
|
HTTPClient: &http.Client{
|
|
Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) {
|
|
_ = r
|
|
return nil, fmt.Errorf("request failed for Authorization: Bearer %s", secret)
|
|
}),
|
|
},
|
|
MaxRetries: 0,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
var out map[string]any
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err == nil {
|
|
t.Fatalf("expected request error")
|
|
}
|
|
if strings.Contains(err.Error(), secret) {
|
|
t.Fatalf("error leaked secret: %v", err)
|
|
}
|
|
if !strings.Contains(err.Error(), "[REDACTED]") {
|
|
t.Fatalf("expected redaction marker in error: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientCancellationAndTimeout(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
select {
|
|
case <-r.Context().Done():
|
|
return
|
|
case <-time.After(200 * time.Millisecond):
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"ok\":true}"}}]}`))
|
|
}
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
RequestTimeout: 20 * time.Millisecond,
|
|
MaxRetries: 0,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
var out map[string]any
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err == nil {
|
|
t.Fatalf("expected timeout-related error")
|
|
}
|
|
if !strings.Contains(strings.ToLower(err.Error()), "context deadline") {
|
|
t.Fatalf("expected context deadline in error, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientRetryBehavior(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
var attempts int32
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
current := atomic.AddInt32(&attempts, 1)
|
|
if current == 1 {
|
|
w.WriteHeader(http.StatusInternalServerError)
|
|
_, _ = w.Write([]byte(`{"error":{"message":"temporary failure"}}`))
|
|
return
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"ok\":true}"}}]}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
MaxRetries: 1,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
var out map[string]any
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err != nil {
|
|
t.Fatalf("CompleteStructured: %v", err)
|
|
}
|
|
if atomic.LoadInt32(&attempts) != 2 {
|
|
t.Fatalf("expected 2 attempts, got %d", attempts)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientRetryOnMalformedStructuredOutput(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
var attempts int32
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
current := atomic.AddInt32(&attempts, 1)
|
|
w.Header().Set("Content-Type", "application/json")
|
|
if current == 1 {
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{"}}]}`))
|
|
return
|
|
}
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"ok\":true}"}}]}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
MaxRetries: 1,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
var out map[string]any
|
|
_, err = client.CompleteStructured(context.Background(), contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err != nil {
|
|
t.Fatalf("CompleteStructured: %v", err)
|
|
}
|
|
if atomic.LoadInt32(&attempts) != 2 {
|
|
t.Fatalf("expected 2 attempts, got %d", attempts)
|
|
}
|
|
}
|
|
|
|
func TestOpenAICompatibleClientHonorsCancelledContextWithoutRetry(t *testing.T) {
|
|
schema := responseschema.MustLookup(responseschema.CorrectionSetKey)
|
|
var attempts int32
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
atomic.AddInt32(&attempts, 1)
|
|
<-r.Context().Done()
|
|
}))
|
|
defer server.Close()
|
|
|
|
client, err := NewOpenAICompatibleClient(OpenAICompatibleClientConfig{
|
|
BaseURL: server.URL,
|
|
Model: "test-model",
|
|
MaxRetries: 3,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("NewOpenAICompatibleClient: %v", err)
|
|
}
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
cancel()
|
|
|
|
var out map[string]any
|
|
_, err = client.CompleteStructured(ctx, contracts.StructuredCompletionRequest{
|
|
Messages: []contracts.LLMMessage{{Role: "user", Content: "extract"}},
|
|
ResponseSchema: &schema,
|
|
}, &out)
|
|
if err == nil {
|
|
t.Fatalf("expected cancellation error")
|
|
}
|
|
if !errors.Is(err, context.Canceled) && !strings.Contains(strings.ToLower(err.Error()), "canceled") {
|
|
t.Fatalf("expected cancellation-related error, got %v", err)
|
|
}
|
|
if atomic.LoadInt32(&attempts) > 1 {
|
|
t.Fatalf("expected no retry after cancellation, got attempts=%d", attempts)
|
|
}
|
|
}
|
|
|
|
type roundTripFunc func(*http.Request) (*http.Response, error)
|
|
|
|
func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) {
|
|
return f(r)
|
|
}
|
|
|
|
func newJSONHTTPResponse(status int, body string) *http.Response {
|
|
return &http.Response{
|
|
StatusCode: status,
|
|
Header: http.Header{"Content-Type": []string{"application/json"}},
|
|
Body: io.NopCloser(strings.NewReader(body)),
|
|
}
|
|
}
|