feat(llmclient): support OPENAI_API_KEY, openrouter wins

This commit is contained in:
pj committed 2026-06-12 10:52:14 +05:30
1 parent 2bd40cc4ee
commit 8eabaaa842
2 files changed
+80 -21

No files matched your search

+166
View File
@@ -0,0 +1,166 @@
// Package llmclient is a minimal client for the OpenAI-compatible
// chat-completions API, covering only what the LLM action backend needs: a
// single multimodal (text + one image) request per step with strict
// json_schema structured output. No streaming, tools, or other extras.
//
// The provider comes from the environment: OPENROUTER_API_KEY routes to
// OpenRouter, OPENAI_API_KEY to OpenAI; OpenRouter wins when both are set.
// Both speak the same wire format, so there is no provider-specific code.
package llmclient
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"os"
"time"
)
const (
openRouterBaseURL = "https://openrouter.ai/api/v1"
openAIBaseURL = "https://api.openai.com/v1"
)
// requestTimeout bounds a single chat-completion round-trip. Vision + strict
// structured output is slower than a plain text call, so this is generous; the
// runner also passes a context the caller can cancel.
const requestTimeout = 60 * time.Second
// Client talks to an OpenAI-compatible chat-completions endpoint.
type Client struct {
httpClient *http.Client
apiKey string
baseURL string
}
// New builds a Client from the environment. OPENROUTER_API_KEY selects
// OpenRouter, OPENAI_API_KEY selects OpenAI; OpenRouter wins when both are
// set. OPENROUTER_BASE_URL / OPENAI_BASE_URL override the chosen provider's
// endpoint (tests, local OpenAI-compatible servers).
func New() (*Client, error) {
apiKey := os.Getenv("OPENROUTER_API_KEY")
baseURL := openRouterBaseURL
override := os.Getenv("OPENROUTER_BASE_URL")
if apiKey == "" {
apiKey = os.Getenv("OPENAI_API_KEY")
baseURL = openAIBaseURL
override = os.Getenv("OPENAI_BASE_URL")
}
if apiKey == "" {
return nil, errors.New("llmclient: neither OPENROUTER_API_KEY nor OPENAI_API_KEY is set")
}
if override != "" {
baseURL = override
}
return &Client{
httpClient: &http.Client{Timeout: requestTimeout},
apiKey: apiKey,
baseURL: baseURL,
}, nil
}
// Request is a chat-completions request body. Only the fields the action
// backend sets are modeled.
type Request struct {
Model string `json:"model"`
Messages []Message `json:"messages"`
ResponseFormat *ResponseFormat `json:"response_format,omitempty"`
}
// Message is one chat message. Content is always the array form (a list of
// parts), which OpenRouter accepts for every role.
type Message struct {
Role string `json:"role"`
Content []ContentPart `json:"content"`
}
// ContentPart is one piece of a message: either a text run or an image given
// as a data URL.
type ContentPart struct {
Type string `json:"type"`
Text string `json:"text,omitempty"`
ImageURL *ImageURL `json:"image_url,omitempty"`
}
// ImageURL carries an image as a (typically data:) URL.
type ImageURL struct {
URL string `json:"url"`
}
// TextPart builds a text content part.
func TextPart(text string) ContentPart {
return ContentPart{Type: "text", Text: text}
}
// ImagePart builds an image content part from a data URL.
func ImagePart(dataURL string) ContentPart {
return ContentPart{Type: "image_url", ImageURL: &ImageURL{URL: dataURL}}
}
// ResponseFormat pins the model to a strict JSON schema.
type ResponseFormat struct {
Type string `json:"type"`
JSONSchema JSONSchema `json:"json_schema"`
}
// JSONSchema is the strict structured-output schema.
type JSONSchema struct {
Name string `json:"name"`
Strict bool `json:"strict"`
Schema map[string]any `json:"schema"`
}
// Response is the slice of a chat-completions response we read.
type Response struct {
Choices []Choice `json:"choices"`
}
// Choice is one completion choice.
type Choice struct {
Message ResponseMessage `json:"message"`
}
// ResponseMessage carries the assistant's content. With json_schema output the
// content is a JSON string matching the schema.
type ResponseMessage struct {
Content string `json:"content"`
}
// ChatCompletion POSTs req to /chat/completions and decodes the response. A
// non-2xx status is returned as an error carrying the response body.
func (c *Client) ChatCompletion(ctx context.Context, req Request) (*Response, error) {
body, err := json.Marshal(req)
if err != nil {
return nil, fmt.Errorf("llmclient: marshal request: %w", err)
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+"/chat/completions", bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("llmclient: build request: %w", err)
}
httpReq.Header.Set("Authorization", "Bearer "+c.apiKey)
httpReq.Header.Set("Content-Type", "application/json")
resp, err := c.httpClient.Do(httpReq)
if err != nil {
return nil, fmt.Errorf("llmclient: request failed: %w", err)
}
defer resp.Body.Close()
responseBody, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("llmclient: read response: %w", err)
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("llmclient: status %d: %s", resp.StatusCode, string(responseBody))
}
var out Response
if err := json.Unmarshal(responseBody, &out); err != nil {
return nil, fmt.Errorf("llmclient: decode response: %w", err)
}
return &out, nil
}
+177
View File
@@ -0,0 +1,177 @@
package llmclient
import (
"context"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestChatCompletionRequestShapeAndParse(t *testing.T) {
var captured map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/chat/completions" {
t.Errorf("path = %q, want /chat/completions", r.URL.Path)
}
if got := r.Header.Get("Authorization"); got != "Bearer test-key" {
t.Errorf("Authorization = %q, want Bearer test-key", got)
}
body, _ := io.ReadAll(r.Body)
if err := json.Unmarshal(body, &captured); err != nil {
t.Fatalf("unmarshal request: %v", err)
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"reasoning\":\"tap login\",\"ranked\":[2,0]}"}}]}`))
}))
defer server.Close()
t.Setenv("OPENROUTER_API_KEY", "test-key")
t.Setenv("OPENROUTER_BASE_URL", server.URL)
client, err := New()
if err != nil {
t.Fatalf("New: %v", err)
}
resp, err := client.ChatCompletion(context.Background(), Request{
Model: "vendor/model",
Messages: []Message{
{Role: "system", Content: []ContentPart{TextPart("system")}},
{Role: "user", Content: []ContentPart{
TextPart("candidates"),
ImagePart("data:image/png;base64,AAAA"),
}},
},
ResponseFormat: &ResponseFormat{
Type: "json_schema",
JSONSchema: JSONSchema{
Name: "ranked_actions",
Strict: true,
Schema: map[string]any{"type": "object"},
},
},
})
if err != nil {
t.Fatalf("ChatCompletion: %v", err)
}
// Request carried the model.
if captured["model"] != "vendor/model" {
t.Errorf("model = %v, want vendor/model", captured["model"])
}
// Request carried an image_url content part.
messages := captured["messages"].([]any)
user := messages[1].(map[string]any)
parts := user["content"].([]any)
foundImage := false
for _, part := range parts {
if part.(map[string]any)["type"] == "image_url" {
foundImage = true
image := part.(map[string]any)["image_url"].(map[string]any)
if !strings.HasPrefix(image["url"].(string), "data:image/png;base64,") {
t.Errorf("image url = %v, want data URL", image["url"])
}
}
}
if !foundImage {
t.Error("request carried no image_url content part")
}
// Request carried the strict json_schema response_format.
rf := captured["response_format"].(map[string]any)
if rf["type"] != "json_schema" {
t.Errorf("response_format.type = %v, want json_schema", rf["type"])
}
if schema := rf["json_schema"].(map[string]any); schema["strict"] != true {
t.Errorf("json_schema.strict = %v, want true", schema["strict"])
}
// Response parsed into the ranked-index content.
if len(resp.Choices) != 1 {
t.Fatalf("choices = %d, want 1", len(resp.Choices))
}
var content struct {
Reasoning string `json:"reasoning"`
Ranked []int `json:"ranked"`
}
if err := json.Unmarshal([]byte(resp.Choices[0].Message.Content), &content); err != nil {
t.Fatalf("unmarshal content: %v", err)
}
if content.Reasoning != "tap login" {
t.Errorf("reasoning = %q, want tap login", content.Reasoning)
}
if len(content.Ranked) != 2 || content.Ranked[0] != 2 || content.Ranked[1] != 0 {
t.Errorf("ranked = %v, want [2 0]", content.Ranked)
}
}
func TestNewRequiresAPIKey(t *testing.T) {
t.Setenv("OPENROUTER_API_KEY", "")
t.Setenv("OPENAI_API_KEY", "")
if _, err := New(); err == nil {
t.Fatal("expected error when neither API key is set")
}
}
func TestNewFallsBackToOpenAIKey(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if got := r.Header.Get("Authorization"); got != "Bearer openai-key" {
t.Errorf("Authorization = %q, want Bearer openai-key", got)
}
_, _ = w.Write([]byte(`{"choices":[]}`))
}))
defer server.Close()
t.Setenv("OPENROUTER_API_KEY", "")
t.Setenv("OPENAI_API_KEY", "openai-key")
t.Setenv("OPENAI_BASE_URL", server.URL)
client, err := New()
if err != nil {
t.Fatalf("New: %v", err)
}
if _, err := client.ChatCompletion(context.Background(), Request{Model: "m"}); err != nil {
t.Fatalf("ChatCompletion: %v", err)
}
}
func TestNewPrefersOpenRouterOverOpenAI(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if got := r.Header.Get("Authorization"); got != "Bearer router-key" {
t.Errorf("Authorization = %q, want Bearer router-key (OpenRouter must win)", got)
}
_, _ = w.Write([]byte(`{"choices":[]}`))
}))
defer server.Close()
t.Setenv("OPENROUTER_API_KEY", "router-key")
t.Setenv("OPENROUTER_BASE_URL", server.URL)
t.Setenv("OPENAI_API_KEY", "openai-key")
t.Setenv("OPENAI_BASE_URL", "http://127.0.0.1:1") // unreachable; must not be used
client, err := New()
if err != nil {
t.Fatalf("New: %v", err)
}
if _, err := client.ChatCompletion(context.Background(), Request{Model: "m"}); err != nil {
t.Fatalf("ChatCompletion: %v", err)
}
}
func TestChatCompletionSurfacesHTTPError(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusTooManyRequests)
_, _ = w.Write([]byte(`{"error":"rate limited"}`))
}))
defer server.Close()
t.Setenv("OPENROUTER_API_KEY", "test-key")
t.Setenv("OPENROUTER_BASE_URL", server.URL)
client, err := New()
if err != nil {
t.Fatalf("New: %v", err)
}
_, err = client.ChatCompletion(context.Background(), Request{Model: "m"})
if err == nil || !strings.Contains(err.Error(), "429") {
t.Fatalf("expected 429 error, got %v", err)
}
}