ollie/cmd/olliesrv/internal/backend/backend.go

439 lines
16 KiB
Go

// Package backend defines the Backend interface and shared types for LLM
// providers. All backends speak the same canonical types; provider-specific
// wire formats are handled inside each implementation.
package backend
import (
"bufio"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"time"
)
// ContentBlock is a single block within a multi-modal message.
// When a Message has ContentBlocks set, it takes precedence over Content.
type ContentBlock struct {
Type string `json:"type"` // "text" | "image"
Text string `json:"text,omitempty"` // for text blocks
ImageSource *ImageSource `json:"source,omitempty"` // for image blocks
}
// ImageSource describes a base64-encoded image.
type ImageSource struct {
Type string `json:"type"` // "base64"
MediaType string `json:"media_type"` // "image/png", "image/jpeg", "image/gif", "image/webp"
Data string `json:"data"` // base64-encoded bytes
}
// Message is a single conversation turn.
type Message struct {
ID string `json:"id,omitempty"` // stable local ID; not sent to providers
Role string `json:"role"` // "system" | "user" | "assistant" | "tool"
Content string `json:"content"`
Reasoning string `json:"reasoning,omitempty"` // thinking/reasoning content (DeepSeek, OpenAI o-series)
ContentBlocks []ContentBlock `json:"content_blocks,omitempty"` // when set, overrides Content
ToolCalls []ToolCall `json:"tool_calls,omitempty"` // set by assistant when calling tools
ToolCallID string `json:"tool_call_id,omitempty"` // set on role=tool replies (required by OpenAI)
ToolName string `json:"tool_name,omitempty"` // local tool identity for history/cache semantics
SummaryApplied bool `json:"summary_applied,omitempty"` // local history state
}
// Tool describes a callable function exposed to the model.
type Tool struct {
Name string `json:"name"`
Description string `json:"description"`
Parameters json.RawMessage `json:"parameters"` // JSON Schema object
}
// ToolCall is the model's request to invoke a function.
type ToolCall struct {
ID string `json:"id,omitempty"` // provider-assigned; may be empty (Ollama)
Name string `json:"name"`
Arguments json.RawMessage `json:"arguments"` // always a JSON object
}
// Usage holds token counts and optional cost for a single Chat call.
// CostUSD is non-zero only when the backend reports it directly (e.g. OpenRouter).
// CachedInputTokens are tokens served from the prompt cache; they are NOT
// included in InputTokens (Anthropic normalises this way; we do the same for
// OpenAI by subtracting cached_tokens from prompt_tokens before setting
// InputTokens).
// CacheCreationTokens are tokens written to the cache (Anthropic only).
type Usage struct {
InputTokens int
CachedInputTokens int
CacheCreationTokens int
OutputTokens int
CostUSD float64
}
// ModelPricing holds per-token costs for a model. All values are in USD per token.
// Zero values mean pricing is not available or not applicable.
type ModelPricing struct {
Input float64 // cost per input token
Output float64 // cost per output token
CacheRead float64 // cost per cached input token (if different from Input)
CacheWrite float64 // cost per cache creation token (if applicable)
Estimated bool // true if pricing is estimated from static table, not API
}
// ModelInfo holds extended information about a model.
type ModelInfo struct {
ID string
ContextLength int
Pricing *ModelPricing // nil if pricing not available
}
// StreamEvent is a single increment from a streaming chat call.
// Content is an incremental text delta (append, not replace).
// Reasoning is an incremental thinking/reasoning delta (append, not replace).
// ToolCalls accumulates complete calls; they may arrive on any event.
// The final event has Done==true.
type StreamEvent struct {
Content string // incremental text delta (may be "")
Reasoning string // incremental thinking delta (may be ""); Anthropic extended thinking, OpenAI reasoning
ToolCalls []ToolCall // complete tool calls assembled so far
Done bool
StopReason string // meaningful when Done==true
Usage Usage // meaningful only when Done==true
}
// RateLimitError is returned when the backend responds with HTTP 429.
// RetryAfter is the suggested wait duration; zero means no hint was given.
type RateLimitError struct {
RetryAfter time.Duration
Message string
}
func (e *RateLimitError) Error() string {
if e.RetryAfter > 0 {
return fmt.Sprintf("rate limited (retry after %v): %s", e.RetryAfter, e.Message)
}
return fmt.Sprintf("rate limited: %s", e.Message)
}
// TransientError is returned for retryable backend failures: 5xx HTTP responses
// and network-level errors (connection reset, DNS failure, etc.).
type TransientError struct {
Message string
}
func (e *TransientError) Error() string { return e.Message }
// ContextOverflowError is returned when the request exceeds the model's
// context window. The caller should compact history and retry.
type ContextOverflowError struct {
Message string
}
func (e *ContextOverflowError) Error() string { return e.Message }
// ToolUnsupportedError is returned when the model or backend does not support
// tool/function calling. The request should not be retried with the same model.
type ToolUnsupportedError struct {
Message string
}
func (e *ToolUnsupportedError) Error() string { return e.Message }
// GenerationParams controls sampling behaviour for a single ChatStream call.
// Zero values mean "use the API default".
type GenerationParams struct {
MaxTokens int `json:"maxTokens,omitempty"`
MaxCompletionTokens int `json:"maxCompletionTokens,omitempty"`
Temperature *float64 `json:"temperature,omitempty"`
TopP *float64 `json:"topP,omitempty"`
TopK *int `json:"topK,omitempty"`
MinP *float64 `json:"minP,omitempty"`
TopA *float64 `json:"topA,omitempty"`
FrequencyPenalty *float64 `json:"frequencyPenalty,omitempty"`
PresencePenalty *float64 `json:"presencePenalty,omitempty"`
RepetitionPenalty *float64 `json:"repetitionPenalty,omitempty"`
ThinkingBudget int `json:"thinkingBudget,omitempty"`
ReasoningEffort string `json:"reasoningEffort,omitempty"`
IncludeReasoning *bool `json:"includeReasoning,omitempty"`
ResponseFormat string `json:"responseFormat,omitempty"`
Stop []string `json:"stop,omitempty"`
// Environment state for backends that support it (e.g., Kiro).
// These are passed through to the backend API but not used for sampling.
CWD string `json:"-"` // Current working directory
}
// Reasoning effort levels. ReasoningEffort is the primary, human-facing knob
// for controlling how much a reasoning-capable model deliberates before
// answering. Backends that speak a discrete effort level (OpenAI-family
// reasoning_effort) send these strings verbatim. Backends that require a
// numeric thinking-token budget (Anthropic extended thinking) translate the
// level via EffortThinkingBudget.
const (
EffortLow = "low"
EffortMedium = "medium"
EffortHigh = "high"
)
// Thinking-token budgets that each reasoning effort level maps to, for backends
// that require a numeric budget rather than a discrete level. These are the
// single documented source of truth for the thresholds; keep data/agents/README.md
// in sync.
const (
thinkingBudgetLow = 4096
thinkingBudgetMedium = 8192
thinkingBudgetHigh = 16384
)
// ValidReasoningEffort reports whether s is a recognised effort level. The
// empty string is valid and means "unset" (use the API default).
func ValidReasoningEffort(s string) bool {
switch s {
case "", EffortLow, EffortMedium, EffortHigh:
return true
}
return false
}
// EffortThinkingBudget maps a reasoning effort level to a thinking-token budget.
// Returns 0 for an unset or unrecognised level, meaning "no explicit budget".
func EffortThinkingBudget(effort string) int {
switch effort {
case EffortLow:
return thinkingBudgetLow
case EffortMedium:
return thinkingBudgetMedium
case EffortHigh:
return thinkingBudgetHigh
}
return 0
}
// ResolvedThinkingBudget returns the thinking-token budget to use for backends
// that require a numeric budget. An explicit ThinkingBudget always wins; when
// it is unset (0), the budget is derived from ReasoningEffort. Returns 0 when
// neither is set, meaning the backend should not enable extended thinking.
func (p GenerationParams) ResolvedThinkingBudget() int {
if p.ThinkingBudget > 0 {
return p.ThinkingBudget
}
return EffortThinkingBudget(p.ReasoningEffort)
}
// Backend is the interface all LLM providers must implement.
// Streaming is the only supported mode; backends that wrap blocking APIs
// should implement ChatStream as a single-event stream.
type Backend interface {
ChatStream(ctx context.Context, messages []Message, tools []Tool, params GenerationParams) (<-chan StreamEvent, error)
// Name returns a short human-readable label for this backend (e.g. "ollama", "openrouter").
Name() string
// DefaultModel returns a reasonable default model name for this backend.
DefaultModel() string
// Model returns the currently active model name.
Model() string
// SetModel changes the active model.
SetModel(model string)
// ContextLength returns the context window size in tokens for the active model.
// Returns 0 if unknown. May make an API call; results should be cached by the implementation.
ContextLength(ctx context.Context) int
// Models returns the list of available model IDs from the provider.
Models(ctx context.Context) []string
}
// ModelLister is an optional interface backends can implement to provide
// extended model information including pricing. The ModelCache uses this
// when available.
type ModelLister interface {
// ModelsInfo returns extended information about available models.
// Backends that support pricing should return it in the Pricing field.
ModelsInfo(ctx context.Context) []ModelInfo
}
// staticPriceEntry holds USD per 1M input/output tokens for a model family.
// Source: public API pricing pages. More specific prefixes must appear before general ones.
type staticPriceEntry struct {
prefix string
inputPer1M float64
outputPer1M float64
cacheRead float64 // discount rate (0.10 = 10% of input), 0 means use default
cacheWrite float64 // premium rate (1.25 = 125% of input), 0 means use default
}
// staticPrices maps case-insensitive model name prefixes to USD/1M token prices.
// Used for estimated pricing when backends don't provide it via API.
var staticPrices = []staticPriceEntry{
// Anthropic (cache: read=10%, write=125%)
{"claude-opus-4", 15.00, 75.00, 0.10, 1.25},
{"claude-sonnet-4", 3.00, 15.00, 0.10, 1.25},
{"claude-haiku-4", 0.80, 4.00, 0.10, 1.25},
{"claude-3-opus", 15.00, 75.00, 0.10, 1.25},
{"claude-3-5-sonnet", 3.00, 15.00, 0.10, 1.25},
{"claude-3-5-haiku", 1.00, 5.00, 0.10, 1.25},
{"claude-3-sonnet", 3.00, 15.00, 0.10, 1.25},
{"claude-3-haiku", 0.25, 1.25, 0.10, 1.25},
// OpenAI — specific before general (cache: read=50%, write not applicable)
{"gpt-4o-mini", 0.15, 0.60, 0.50, 0},
{"gpt-4o", 2.50, 10.00, 0.50, 0},
{"gpt-4-turbo", 10.00, 30.00, 0.50, 0},
{"gpt-4.1", 2.00, 8.00, 0.50, 0},
{"gpt-4.1-mini", 0.40, 1.60, 0.50, 0},
{"gpt-4.1-nano", 0.10, 0.40, 0.50, 0},
{"o4-mini", 1.10, 4.40, 0.275, 0},
{"o3-mini", 1.10, 4.40, 0.275, 0},
{"o3", 10.00, 40.00, 0.25, 0},
{"o1-mini", 1.10, 4.40, 0.275, 0},
{"o1", 15.00, 60.00, 0.375, 0},
// Google Gemini (cache: read=25%)
{"gemini-2.5-pro", 1.25, 10.00, 0.25, 0},
{"gemini-2.5-flash", 0.15, 0.60, 0.25, 0},
{"gemini-2.0-flash", 0.10, 0.40, 0.25, 0},
{"gemini-1.5-pro", 1.25, 5.00, 0.25, 0},
{"gemini-1.5-flash", 0.075, 0.30, 0.25, 0},
// DeepSeek
{"deepseek-chat", 0.27, 1.10, 0.035, 0},
{"deepseek-reasoner", 0.55, 2.19, 0.14, 0},
{"deepseek-v3", 0.27, 1.10, 0.035, 0},
}
// LookupStaticPricing returns estimated pricing for a model ID based on the static table.
// Returns nil if no matching prefix is found. Pricing is marked as Estimated.
func LookupStaticPricing(modelID string) *ModelPricing {
lower := strings.ToLower(modelID)
for _, e := range staticPrices {
if strings.Contains(lower, e.prefix) {
p := &ModelPricing{
Input: e.inputPer1M / 1_000_000,
Output: e.outputPer1M / 1_000_000,
Estimated: true,
}
if e.cacheRead > 0 {
p.CacheRead = p.Input * e.cacheRead
}
if e.cacheWrite > 0 {
p.CacheWrite = p.Input * e.cacheWrite
}
return p
}
}
return nil
}
// streamRequest executes req via client, handles status errors, and spawns a
// goroutine that feeds resp.Body through parseFn into the returned channel.
// label is used in error messages (e.g. "openai", "ollama").
// Returns *RateLimitError for 429, *TransientError for 5xx and network errors,
// and a plain error for other non-200 responses.
func streamRequest(client *http.Client, req *http.Request, label string, parseFn func(context.Context, io.Reader, chan<- StreamEvent)) (<-chan StreamEvent, error) {
resp, err := client.Do(req)
if err != nil {
return nil, &TransientError{Message: err.Error()}
}
if resp.StatusCode == http.StatusTooManyRequests {
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
return nil, &RateLimitError{RetryAfter: parseRetryAfter(resp.Header.Get("Retry-After")), Message: string(body)}
}
if resp.StatusCode >= 500 {
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
return nil, &TransientError{Message: fmt.Sprintf("%s HTTP %d: %s", label, resp.StatusCode, string(body))}
}
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
msg := string(body)
if resp.StatusCode == http.StatusBadRequest && isContextOverflow(msg) {
return nil, &ContextOverflowError{Message: fmt.Sprintf("%s HTTP %d: %s", label, resp.StatusCode, msg)}
}
if (resp.StatusCode == http.StatusBadRequest || resp.StatusCode == http.StatusUnprocessableEntity || resp.StatusCode == http.StatusNotFound) && isToolUnsupported(msg) {
return nil, &ToolUnsupportedError{Message: fmt.Sprintf("%s HTTP %d: %s", label, resp.StatusCode, msg)}
}
return nil, fmt.Errorf("%s HTTP %d: %s", label, resp.StatusCode, msg)
}
ch := make(chan StreamEvent, 8)
go func() {
defer close(ch)
defer resp.Body.Close()
parseFn(req.Context(), resp.Body, ch)
}()
return ch, nil
}
func sendStreamEvent(ctx context.Context, ch chan<- StreamEvent, ev StreamEvent) bool {
select {
case ch <- ev:
return true
case <-ctx.Done():
return false
}
}
// model does not support tool/function calling.
func isToolUnsupported(body string) bool {
markers := []string{
"tool use is not supported",
"tools is not supported",
"function calling is not supported",
"does not support tools",
"does not support function",
"tool_use is not enabled",
"not support tool",
"no endpoints found that support tool",
}
lower := strings.ToLower(body)
for _, m := range markers {
if strings.Contains(lower, m) {
return true
}
}
return false
}
// isContextOverflow returns true when a 400 response body indicates the request
// exceeded the model's context window.
func isContextOverflow(body string) bool {
markers := []string{
"context_length_exceeded", // OpenAI/OpenRouter error code
"prompt is too long", // Anthropic
"Please reduce the length", // OpenAI prose
"too many tokens",
}
lower := strings.ToLower(body)
for _, m := range markers {
if strings.Contains(lower, strings.ToLower(m)) {
return true
}
}
return false
}
// parseRetryAfter parses the Retry-After header value, which may be an integer
// number of seconds or an HTTP-date. Returns zero if the header is absent or
// unparseable.
func parseRetryAfter(header string) time.Duration {
if header == "" {
return 0
}
header = strings.TrimSpace(header)
if secs, err := strconv.Atoi(header); err == nil {
return time.Duration(secs) * time.Second
}
if t, err := http.ParseTime(header); err == nil {
if d := time.Until(t); d > 0 {
return d
}
}
return 0
}
// newStreamScanner returns a bufio.Scanner suitable for reading SSE/NDJSON
// streaming responses. The buffer is sized for typical LLM chunks (64K initial,
// 4MB max).
func newStreamScanner(r io.Reader) *bufio.Scanner {
s := bufio.NewScanner(r)
s.Buffer(make([]byte, 0, 64*1024), 4*1024*1024)
return s
}