78 lines
3.1 KiB
Go
78 lines
3.1 KiB
Go
package ai
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"sync/atomic"
|
|
"testing"
|
|
)
|
|
|
|
func TestLLMSDKPreservesExplicitParametersWithoutRetry(t *testing.T) {
|
|
for _, code := range []int{http.StatusOK, http.StatusInternalServerError} {
|
|
t.Run(http.StatusText(code), func(t *testing.T) {
|
|
var calls atomic.Int32
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
calls.Add(1)
|
|
if r.Method != "POST" || r.URL.Path != "/v1/chat/completions" {
|
|
t.Errorf("unexpected request %s %s", r.Method, r.URL.Path)
|
|
}
|
|
var body struct {
|
|
Model string `json:"model"`
|
|
Temperature *float64 `json:"temperature"`
|
|
MaxTokens int `json:"max_tokens"`
|
|
Messages []struct{ Role, Content string } `json:"messages"`
|
|
}
|
|
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
|
t.Error(err)
|
|
}
|
|
if body.Model != "approved-model" || body.Temperature == nil || *body.Temperature != 0 || body.MaxTokens != 17 || len(body.Messages) != 2 {
|
|
t.Errorf("explicit configuration not transmitted: %+v", body)
|
|
}
|
|
if len(body.Messages) == 2 && (body.Messages[0].Role != "system" || body.Messages[0].Content != "approved system" || body.Messages[1].Content != "synthetic input") {
|
|
t.Error("message content changed")
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(code)
|
|
if code == http.StatusOK {
|
|
_, _ = w.Write([]byte(`{"id":"local","object":"chat.completion","choices":[{"index":0,"message":{"role":"assistant","content":" local reply "},"finish_reason":"stop"}]}`))
|
|
} else {
|
|
_, _ = w.Write([]byte(`{"error":{"message":"injected failure","type":"server_error"}}`))
|
|
}
|
|
}))
|
|
defer server.Close()
|
|
pipeline := &ProviderPipeline{cfg: ProviderPipelineConfig{BailianAPIKey: "isolated-test-key", BailianBaseURL: server.URL + "/v1"}}
|
|
result, err := pipeline.complete(context.Background(), "approved-model", 0, 17, "approved system", "synthetic input")
|
|
if code == http.StatusOK && (err != nil || result != "local reply") {
|
|
t.Fatalf("response %q: %v", result, err)
|
|
}
|
|
if code != http.StatusOK && err == nil {
|
|
t.Fatal("provider failure hidden")
|
|
}
|
|
if calls.Load() != 1 {
|
|
t.Fatalf("automatic request retry: %d", calls.Load())
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestLLMSDKRejectsEmptyChoices(t *testing.T) {
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
var body map[string]any
|
|
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
|
t.Error(err)
|
|
}
|
|
if len(body["messages"].([]any)) != 1 {
|
|
t.Error("invented system message")
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{"id":"local","object":"chat.completion","choices":[]}`))
|
|
}))
|
|
defer server.Close()
|
|
pipeline := &ProviderPipeline{cfg: ProviderPipelineConfig{BailianAPIKey: "isolated-test-key", BailianBaseURL: server.URL + "/v1"}}
|
|
if _, err := pipeline.complete(context.Background(), "approved-model", 0, 17, "", "synthetic input"); err == nil {
|
|
t.Fatal("missing choices reported success")
|
|
}
|
|
}
|