Parent directory

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}