286 lines
9.3 KiB
Go
286 lines
9.3 KiB
Go
package utils
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"net/http"
|
|
"strings"
|
|
"time"
|
|
|
|
"nearle/config"
|
|
)
|
|
|
|
// The model gateway.
|
|
//
|
|
// A sibling of `embedding.go`, and deliberately the same shape: one small
|
|
// interface, a provider switch behind a factory, the shared `postJSON`, and a
|
|
// `Model()` so what produced an answer can be recorded. No framework — the
|
|
// backend talks to models over plain HTTP already, and the OpenAI chat
|
|
// completions format is what Groq, Ollama, vLLM, LM Studio, Together and Azure
|
|
// all accept, so one client reaches all of them.
|
|
//
|
|
// ── Agents name a tier, never a model ───────────────────────────────────────
|
|
//
|
|
// A tier is a promise about how much thinking a question deserves. Wiring a
|
|
// model name into an agent means editing every agent to change provider, and
|
|
// means an agent's YAML going stale the day a model is retired. The mapping
|
|
// lives in config and nowhere else.
|
|
//
|
|
// ── Degrades rather than fails ──────────────────────────────────────────────
|
|
//
|
|
// `ErrChatNotConfigured` is what a deployment with no provider gets, and the
|
|
// assistant answers "typed questions are not switched on here" instead of a
|
|
// 500. The tools keep working either way: they are ordinary Go functions, and
|
|
// only the part that turns a sentence into a tool call is missing.
|
|
|
|
// Tier is how much thinking a question deserves.
|
|
const (
|
|
// TierFast is a chip answer: one tool, no reasoning.
|
|
TierFast = "fast"
|
|
// TierBalanced is the default, and what most typed questions get.
|
|
TierBalanced = "balanced"
|
|
// TierDeep is multi-step reasoning across several tools.
|
|
TierDeep = "deep"
|
|
)
|
|
|
|
// Roles in a conversation, as the wire format spells them.
|
|
const (
|
|
RoleSystem = "system"
|
|
RoleUser = "user"
|
|
RoleAssistant = "assistant"
|
|
RoleTool = "tool"
|
|
)
|
|
|
|
// ErrChatNotConfigured is what callers see when no provider is set.
|
|
var ErrChatNotConfigured = errors.New("no assistant model is configured")
|
|
|
|
// ToolCall is the model asking for a tool to be run.
|
|
//
|
|
// `Arguments` is decoded here rather than passed along as the raw string the
|
|
// wire carries, so a model that emits malformed JSON is caught at the edge —
|
|
// one error, in one place, instead of every caller parsing it again.
|
|
type ToolCall struct {
|
|
ID string
|
|
Name string
|
|
Arguments map[string]any
|
|
}
|
|
|
|
// Message is one turn.
|
|
type Message struct {
|
|
Role string
|
|
Content string
|
|
// Set on an assistant turn that asked for tools.
|
|
ToolCalls []ToolCall
|
|
// Set on a tool turn, naming the call it answers. Without it the model
|
|
// cannot tell which result belongs to which request when it asked for two.
|
|
ToolCallID string
|
|
// The tool's name, which some providers want on the tool turn as well.
|
|
Name string
|
|
}
|
|
|
|
// ChatRequest is one round trip.
|
|
type ChatRequest struct {
|
|
Tier string
|
|
Messages []Message
|
|
// Tool definitions as the registry describes them: name, description,
|
|
// input_schema. Converted to the provider's shape here so nothing above
|
|
// this file knows what that shape is.
|
|
Tools []map[string]any
|
|
MaxTokens int
|
|
}
|
|
|
|
// ChatReply is what came back.
|
|
type ChatReply struct {
|
|
Content string
|
|
ToolCalls []ToolCall
|
|
// Which model actually answered. Recorded on every audit row: an answer
|
|
// nobody can attribute to a model cannot be reproduced when it is wrong.
|
|
Model string
|
|
// The provider's own word for why it stopped — `stop`, `tool_calls`,
|
|
// `length`. `length` is the one that matters: a truncated answer reads as a
|
|
// complete one unless somebody looks.
|
|
StopReason string
|
|
}
|
|
|
|
// Chat is the whole interface. One method, like the embedder.
|
|
type Chat interface {
|
|
Complete(ctx context.Context, req ChatRequest) (ChatReply, error)
|
|
// ModelFor names the model a tier resolves to, for the audit trail.
|
|
ModelFor(tier string) string
|
|
}
|
|
|
|
const chatTimeout = 60 * time.Second
|
|
|
|
// NewChat builds the gateway, or nil when none is configured.
|
|
//
|
|
// Nil rather than an error for the unconfigured case, matching `NewEmbedder`:
|
|
// "no model here" is a deployment choice, not a fault, and the caller checks
|
|
// for nil exactly as it does for the embedder.
|
|
func NewChat(cfg config.AssistantConfig) (Chat, error) {
|
|
if !cfg.Enabled() {
|
|
return nil, nil
|
|
}
|
|
switch cfg.Provider {
|
|
case "openai", "groq", "ollama", "together", "compatible":
|
|
base := strings.TrimRight(cfg.BaseURL, "/")
|
|
if base == "" {
|
|
base = "https://api.openai.com/v1"
|
|
}
|
|
return &openAIChat{cfg: cfg, base: base, client: &http.Client{Timeout: chatTimeout}}, nil
|
|
}
|
|
return nil, fmt.Errorf("assistant provider %q is not supported", cfg.Provider)
|
|
}
|
|
|
|
// ── OpenAI-compatible chat completions ──────────────────────────────────────
|
|
|
|
type openAIChat struct {
|
|
cfg config.AssistantConfig
|
|
base string
|
|
client *http.Client
|
|
}
|
|
|
|
func (c *openAIChat) ModelFor(tier string) string { return c.cfg.ModelFor(tier) }
|
|
|
|
// wire types, kept unexported: nothing above this file should know that a tool
|
|
// call arrives with its arguments as a string.
|
|
type wireToolCall struct {
|
|
ID string `json:"id"`
|
|
Type string `json:"type"`
|
|
Function struct {
|
|
Name string `json:"name"`
|
|
Arguments string `json:"arguments"`
|
|
} `json:"function"`
|
|
}
|
|
|
|
type wireMessage struct {
|
|
Role string `json:"role"`
|
|
Content string `json:"content,omitempty"`
|
|
ToolCalls []wireToolCall `json:"tool_calls,omitempty"`
|
|
ToolCallID string `json:"tool_call_id,omitempty"`
|
|
Name string `json:"name,omitempty"`
|
|
}
|
|
|
|
func (c *openAIChat) Complete(ctx context.Context, req ChatRequest) (ChatReply, error) {
|
|
model := c.cfg.ModelFor(req.Tier)
|
|
if model == "" {
|
|
return ChatReply{}, ErrChatNotConfigured
|
|
}
|
|
|
|
messages := make([]wireMessage, 0, len(req.Messages))
|
|
for _, m := range req.Messages {
|
|
wire := wireMessage{Role: m.Role, Content: m.Content, ToolCallID: m.ToolCallID, Name: m.Name}
|
|
for _, call := range m.ToolCalls {
|
|
raw, err := json.Marshal(call.Arguments)
|
|
if err != nil {
|
|
return ChatReply{}, fmt.Errorf("assistant: encoding a tool call: %w", err)
|
|
}
|
|
wt := wireToolCall{ID: call.ID, Type: "function"}
|
|
wt.Function.Name = call.Name
|
|
wt.Function.Arguments = string(raw)
|
|
wire.ToolCalls = append(wire.ToolCalls, wt)
|
|
}
|
|
messages = append(messages, wire)
|
|
}
|
|
|
|
body := map[string]any{
|
|
"model": model,
|
|
"messages": messages,
|
|
}
|
|
if req.MaxTokens > 0 {
|
|
body["max_tokens"] = req.MaxTokens
|
|
}
|
|
if len(req.Tools) > 0 {
|
|
body["tools"] = toolsForWire(req.Tools)
|
|
// "auto", never "required": some questions are answered from what is
|
|
// already in the conversation, and forcing a call makes the model
|
|
// invent one to satisfy the demand.
|
|
body["tool_choice"] = "auto"
|
|
}
|
|
|
|
var out struct {
|
|
Choices []struct {
|
|
Message wireMessage `json:"message"`
|
|
FinishReason string `json:"finish_reason"`
|
|
} `json:"choices"`
|
|
Model string `json:"model"`
|
|
Error *struct {
|
|
Message string `json:"message"`
|
|
} `json:"error"`
|
|
}
|
|
|
|
auth := ""
|
|
if c.cfg.APIKey != "" {
|
|
auth = "Bearer " + c.cfg.APIKey
|
|
}
|
|
if err := postJSON(ctx, c.client, "assistant", c.base+"/chat/completions", auth, body, &out); err != nil {
|
|
return ChatReply{}, err
|
|
}
|
|
if out.Error != nil {
|
|
return ChatReply{}, fmt.Errorf("assistant: %s", out.Error.Message)
|
|
}
|
|
if len(out.Choices) == 0 {
|
|
return ChatReply{}, errors.New("assistant: the model returned no choices")
|
|
}
|
|
|
|
choice := out.Choices[0]
|
|
reply := ChatReply{
|
|
Content: choice.Message.Content,
|
|
Model: firstNonEmpty(out.Model, model),
|
|
StopReason: choice.FinishReason,
|
|
}
|
|
|
|
for _, call := range choice.Message.ToolCalls {
|
|
args := map[string]any{}
|
|
// An empty argument string is a call with no arguments, which is
|
|
// ordinary — `{}` and `""` both mean the same thing here.
|
|
if trimmed := strings.TrimSpace(call.Function.Arguments); trimmed != "" {
|
|
if err := json.Unmarshal([]byte(trimmed), &args); err != nil {
|
|
// Caught at the edge rather than passed on. A model emitting
|
|
// malformed JSON is a fact about this round trip, and the loop
|
|
// can retry or give up — but nothing downstream should have to
|
|
// parse it a second time.
|
|
return ChatReply{}, fmt.Errorf("assistant: the model sent unreadable arguments for %s: %w", call.Function.Name, err)
|
|
}
|
|
}
|
|
reply.ToolCalls = append(reply.ToolCalls, ToolCall{
|
|
ID: call.ID,
|
|
Name: call.Function.Name,
|
|
Arguments: args,
|
|
})
|
|
}
|
|
|
|
return reply, nil
|
|
}
|
|
|
|
// toolsForWire converts the registry's description into the provider's shape.
|
|
//
|
|
// The registry speaks `{name, description, input_schema}` because that is what
|
|
// MCP uses and what reads clearly. This is the one place that knows OpenAI
|
|
// wants it wrapped in a `function` object — so a second provider with a
|
|
// different shape is a change here and nowhere else.
|
|
func toolsForWire(defs []map[string]any) []map[string]any {
|
|
out := make([]map[string]any, 0, len(defs))
|
|
for _, def := range defs {
|
|
out = append(out, map[string]any{
|
|
"type": "function",
|
|
"function": map[string]any{
|
|
"name": def["name"],
|
|
"description": def["description"],
|
|
"parameters": def["input_schema"],
|
|
},
|
|
})
|
|
}
|
|
return out
|
|
}
|
|
|
|
func firstNonEmpty(values ...string) string {
|
|
for _, value := range values {
|
|
if value != "" {
|
|
return value
|
|
}
|
|
}
|
|
return ""
|
|
}
|