client.go
4619 bytes
1// Package llm is an OpenAI-compatible chat-completions client with the Qwen3
2// no-thinking request options, bounded conversation history, and strict parsers
3// for the NPC reply and judge output contracts. It is an HTTP client only and
4// never starts or manages any model service.
5package llm
6
7import (
8 "bufio"
9 "bytes"
10 "context"
11 "encoding/json"
12 "fmt"
13 "git.theedgeofrage.com/TheEdgeOfRage/kaiwari/internal/config"
14 "io"
15 "net/http"
16 "strings"
17)
18
19type Role string
20
21const (
22 RoleSystem Role = "system"
23 RoleUser Role = "user"
24 RoleAssistant Role = "assistant"
25)
26
27type Message struct {
28 Role Role `json:"role"`
29 Content string `json:"content"`
30}
31
32type Client struct {
33 baseURL string
34 httpClient *http.Client
35 Temperature float64
36 MaxTokens int
37 EnableThinking bool
38 // Slot pins every request (including warmup) to a specific llama-server
39 // slot via the id_slot field. -1 lets the server choose.
40 Slot int
41}
42
43func NewClient(cfg *config.Config, hc *http.Client) *Client {
44 if hc == nil {
45 hc = &http.Client{}
46 }
47 return &Client{
48 httpClient: hc,
49 baseURL: cfg.LLMConfig.BaseURL,
50 Temperature: cfg.LLMConfig.Temperature,
51 MaxTokens: cfg.LLMConfig.MaxTokens,
52 EnableThinking: cfg.LLMConfig.EnableThinking,
53 }
54}
55
56type chatRequest struct {
57 Messages []Message `json:"messages"`
58 ChatTemplateKwargs map[string]any `json:"chat_template_kwargs"`
59 Temperature float64 `json:"temperature"`
60 MaxTokens int `json:"max_tokens"`
61 Stream bool `json:"stream"`
62 CachePrompt bool `json:"cache_prompt"`
63 IDSlot int `json:"id_slot"`
64}
65
66type sseDelta struct {
67 Choices []struct {
68 Delta struct {
69 Content string `json:"content"`
70 } `json:"delta"`
71 } `json:"choices"`
72}
73
74// Generate sends one streaming chat request and returns the full assistant
75// content, consuming SSE data events until [DONE]. The context bounds the whole
76// call, including the read.
77func (c *Client) Generate(ctx context.Context, msgs []Message) (string, error) {
78 return c.chat(ctx, msgs, c.MaxTokens)
79}
80
81// Warmup sends each system prompt once with a minimal user line and one output
82// token so the server keeps those prompt prefixes in its KV cache before the
83// first real turn.
84func (c *Client) Warmup(ctx context.Context, systems []string) error {
85 for _, s := range systems {
86 msgs := []Message{{Role: RoleSystem, Content: s}, {Role: RoleUser, Content: "Ready."}}
87 if _, err := c.chat(ctx, msgs, 1); err != nil {
88 return fmt.Errorf("llm: warmup: %w", err)
89 }
90 }
91 return nil
92}
93
94func (c *Client) chat(ctx context.Context, msgs []Message, maxTokens int) (string, error) {
95 if c.baseURL == "" {
96 return "", fmt.Errorf("llm: no base URL configured")
97 }
98 body := chatRequest{
99 Messages: msgs,
100 ChatTemplateKwargs: map[string]any{"enable_thinking": c.EnableThinking},
101 Temperature: c.Temperature,
102 MaxTokens: maxTokens,
103 Stream: true,
104 CachePrompt: true,
105 IDSlot: c.Slot,
106 }
107 payload, err := json.Marshal(body)
108 if err != nil {
109 return "", fmt.Errorf("llm: encode request: %w", err)
110 }
111
112 url := strings.TrimRight(c.baseURL, "/") + "/v1/chat/completions"
113 req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
114 if err != nil {
115 return "", fmt.Errorf("llm: build request: %w", err)
116 }
117 req.Header.Set("Content-Type", "application/json")
118
119 resp, err := c.httpClient.Do(req)
120 if err != nil {
121 return "", fmt.Errorf("llm: request to %s: %w", url, err)
122 }
123 defer func() { _ = resp.Body.Close() }()
124
125 if resp.StatusCode < 200 || resp.StatusCode >= 300 {
126 detail, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
127 return "", fmt.Errorf("llm: %s returned HTTP %d: %s", url, resp.StatusCode, strings.TrimSpace(string(detail)))
128 }
129
130 text, err := readSSE(resp.Body)
131 if err != nil {
132 return "", fmt.Errorf("llm: read stream from %s: %w", url, err)
133 }
134 return text, nil
135}
136
137func readSSE(r io.Reader) (string, error) {
138 var sb strings.Builder
139 scanner := bufio.NewScanner(r)
140 for scanner.Scan() {
141 line := scanner.Text()
142 if !strings.HasPrefix(line, "data:") {
143 continue
144 }
145 data := strings.TrimSpace(strings.TrimPrefix(line, "data:"))
146 if data == "[DONE]" {
147 break
148 }
149 var ev sseDelta
150 if err := json.Unmarshal([]byte(data), &ev); err != nil {
151 return "", fmt.Errorf("decode SSE event %q: %w", data, err)
152 }
153 for _, ch := range ev.Choices {
154 sb.WriteString(ch.Delta.Content)
155 }
156 }
157 if err := scanner.Err(); err != nil {
158 return "", err
159 }
160 return sb.String(), nil
161}