Add LLM interface

Signed-off-by: Bryan Frimin <bryan@getprobo.com>
This commit is contained in:
Bryan Frimin
2026-03-11 22:33:12 +01:00
parent 307ca48394
commit 158c36d9ab
22 changed files with 2822 additions and 91 deletions

14
go.mod
View File

@@ -5,9 +5,11 @@ go 1.26.1
require (
codeberg.org/miekg/dns v0.6.65
github.com/99designs/gqlgen v0.17.87
github.com/aws/aws-sdk-go-v2 v1.41.2
github.com/anthropics/anthropic-sdk-go v1.26.0
github.com/aws/aws-sdk-go-v2 v1.41.3
github.com/aws/aws-sdk-go-v2/credentials v1.19.10
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.18
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.50.1
github.com/aws/aws-sdk-go-v2/service/s3 v1.96.2
github.com/brianvoe/gofakeit/v7 v7.14.1
github.com/chromedp/cdproto v0.0.0-20250803210736-d308e07a266d
@@ -49,15 +51,15 @@ require (
cloud.google.com/go/auth/oauth2adapt v0.2.8 // indirect
cloud.google.com/go/compute/metadata v0.9.0 // indirect
github.com/agnivade/levenshtein v1.2.1 // indirect
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.5 // indirect
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.18 // indirect
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.18 // indirect
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.6 // indirect
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.19 // indirect
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.19 // indirect
github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.18 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.5 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.10 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.18 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.18 // indirect
github.com/aws/smithy-go v1.24.1 // indirect
github.com/aws/smithy-go v1.24.2
github.com/beevik/etree v1.6.0 // indirect
github.com/beorn7/perks v1.0.1 // indirect
github.com/bitfield/gotestdox v0.2.2 // indirect
@@ -133,7 +135,7 @@ require (
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.40.0 // indirect
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.40.0 // indirect
go.opentelemetry.io/otel/metric v1.40.0 // indirect
go.opentelemetry.io/otel/sdk v1.40.0 // indirect
go.opentelemetry.io/otel/sdk v1.40.0
go.opentelemetry.io/proto/otlp v1.9.0 // indirect
go.yaml.in/yaml/v2 v2.4.3 // indirect
golang.org/x/mod v0.34.0 // indirect

26
go.sum
View File

@@ -12,22 +12,26 @@ github.com/agnivade/levenshtein v1.2.1 h1:EHBY3UOn1gwdy/VbFwgo4cxecRznFk7fKWN1KO
github.com/agnivade/levenshtein v1.2.1/go.mod h1:QVVI16kDrtSuwcpd0p1+xMC6Z/VfhtCyDIjcwga4/DU=
github.com/andreyvit/diff v0.0.0-20170406064948-c7f18ee00883 h1:bvNMNQO63//z+xNgfBlViaCIJKLlCJ6/fmUseuG0wVQ=
github.com/andreyvit/diff v0.0.0-20170406064948-c7f18ee00883/go.mod h1:rCTlJbsFo29Kk6CurOXKm700vrz8f0KW0JNfpkRJY/8=
github.com/anthropics/anthropic-sdk-go v1.26.0 h1:oUTzFaUpAevfuELAP1sjL6CQJ9HHAfT7CoSYSac11PY=
github.com/anthropics/anthropic-sdk-go v1.26.0/go.mod h1:qUKmaW+uuPB64iy1l+4kOSvaLqPXnHTTBKH6RVZ7q5Q=
github.com/arbovm/levenshtein v0.0.0-20160628152529-48b4e1c0c4d0 h1:jfIu9sQUG6Ig+0+Ap1h4unLjW6YQJpKZVmUzxsD4E/Q=
github.com/arbovm/levenshtein v0.0.0-20160628152529-48b4e1c0c4d0/go.mod h1:t2tdKJDJF9BV14lnkjHmOQgcvEKgtqs5a1N3LNdJhGE=
github.com/aws/aws-sdk-go-v2 v1.41.2 h1:LuT2rzqNQsauaGkPK/7813XxcZ3o3yePY0Iy891T2ls=
github.com/aws/aws-sdk-go-v2 v1.41.2/go.mod h1:IvvlAZQXvTXznUPfRVfryiG1fbzE2NGK6m9u39YQ+S4=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.5 h1:zWFmPmgw4sveAYi1mRqG+E/g0461cJ5M4bJ8/nc6d3Q=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.5/go.mod h1:nVUlMLVV8ycXSb7mSkcNu9e3v/1TJq2RTlrPwhYWr5c=
github.com/aws/aws-sdk-go-v2 v1.41.3 h1:4kQ/fa22KjDt13QCy1+bYADvdgcxpfH18f0zP542kZA=
github.com/aws/aws-sdk-go-v2 v1.41.3/go.mod h1:mwsPRE8ceUUpiTgF7QmQIJ7lgsKUPQOUl3o72QBrE1o=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.6 h1:N4lRUXZpZ1KVEUn6hxtco/1d2lgYhNn1fHkkl8WhlyQ=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.6/go.mod h1:lyw7GFp3qENLh7kwzf7iMzAxDn+NzjXEAGjKS2UOKqI=
github.com/aws/aws-sdk-go-v2/credentials v1.19.10 h1:EEhmEUFCE1Yhl7vDhNOI5OCL/iKMdkkYFTRpZXNw7m8=
github.com/aws/aws-sdk-go-v2/credentials v1.19.10/go.mod h1:RnnlFCAlxQCkN2Q379B67USkBMu1PipEEiibzYN5UTE=
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.18 h1:Ii4s+Sq3yDfaMLpjrJsqD6SmG/Wq/P5L/hw2qa78UAY=
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.18/go.mod h1:6x81qnY++ovptLE6nWQeWrpXxbnlIex+4H4eYYGcqfc=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.18 h1:F43zk1vemYIqPAwhjTjYIz0irU2EY7sOb/F5eJ3HuyM=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.18/go.mod h1:w1jdlZXrGKaJcNoL+Nnrj+k5wlpGXqnNrKoP22HvAug=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.18 h1:xCeWVjj0ki0l3nruoyP2slHsGArMxeiiaoPN5QZH6YQ=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.18/go.mod h1:r/eLGuGCBw6l36ZRWiw6PaZwPXb6YOj+i/7MizNl5/k=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.19 h1:/sECfyq2JTifMI2JPyZ4bdRN77zJmr6SrS1eL3augIA=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.19/go.mod h1:dMf8A5oAqr9/oxOfLkC/c2LU/uMcALP0Rgn2BD5LWn0=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.19 h1:AWeJMk33GTBf6J20XJe6qZoRSJo0WfUhsMdUKhoODXE=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.19/go.mod h1:+GWrYoaAsV7/4pNHpwh1kiNLXkKaSoppxQq9lbH8Ejw=
github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.18 h1:eZioDaZGJ0tMM4gzmkNIO2aAoQd+je7Ug7TkvAzlmkU=
github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.18/go.mod h1:CCXwUKAJdoWr6/NcxZ+zsiPr6oH/Q5aTooRGYieAyj4=
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.50.1 h1:tnLUbtNW5c056BEbQ4xvlZaakvgdaEdiKF87R1fxuoo=
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.50.1/go.mod h1:DYDD64rVUpCvpLyuWCiTaaSfrW2O9GiDo8S6fNo8ZI0=
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.5 h1:CeY9LUdur+Dxoeldqoun6y4WtJ3RQtzk0JMP2gfUay0=
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.5/go.mod h1:AZLZf2fMaahW5s/wMRciu1sYbdsikT/UHwbUjOdEVTc=
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.10 h1:fJvQ5mIBVfKtiyx0AHY6HeWcRX5LGANLpq8SVR+Uazs=
@@ -38,8 +42,8 @@ github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.18 h1:/A/xDuZAVD2Bp
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.18/go.mod h1:hWe9b4f+djUQGmyiGEeOnZv69dtMSgpDRIvNMvuvzvY=
github.com/aws/aws-sdk-go-v2/service/s3 v1.96.2 h1:M1A9AjcFwlxTLuf0Faj88L8Iqw0n/AJHjpZTQzMMsSc=
github.com/aws/aws-sdk-go-v2/service/s3 v1.96.2/go.mod h1:KsdTV6Q9WKUZm2mNJnUFmIoXfZux91M3sr/a4REX8e0=
github.com/aws/smithy-go v1.24.1 h1:VbyeNfmYkWoxMVpGUAbQumkODcYmfMRfZ8yQiH30SK0=
github.com/aws/smithy-go v1.24.1/go.mod h1:LEj2LM3rBRQJxPZTB4KuzZkaZYnZPnvgIhb4pu07mx0=
github.com/aws/smithy-go v1.24.2 h1:FzA3bu/nt/vDvmnkg+R8Xl46gmzEDam6mZ1hzmwXFng=
github.com/aws/smithy-go v1.24.2/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc=
github.com/beevik/etree v1.6.0 h1:u8Kwy8pp9D9XeITj2Z0XtA5qqZEmtJtuXZRQi+j03eE=
github.com/beevik/etree v1.6.0/go.mod h1:bh4zJxiIr62SOf9pRzN7UUYaEDa9HEKafK25+sLc0Gc=
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
@@ -81,6 +85,8 @@ github.com/digitorus/pkcs7 v0.0.0-20230713084857-e76b763bdc49 h1:h+XMRXf+WLY0h/3
github.com/digitorus/pkcs7 v0.0.0-20230713084857-e76b763bdc49/go.mod h1:SKVExuS+vpu2l9IoOc0RwqE7NYnb0JlcFHFnEJkVDzc=
github.com/digitorus/timestamp v0.0.0-20250524132541-c45532741eea h1:ALRwvjsSP53QmnN3Bcj0NpR8SsFLnskny/EIMebAk1c=
github.com/digitorus/timestamp v0.0.0-20250524132541-c45532741eea/go.mod h1:GvWntX9qiTlOud0WkQ6ewFm0LPy5JUR1Xo0Ngbd1w6Y=
github.com/dnaeon/go-vcr v1.2.0 h1:zHCHvJYTMh1N7xnV7zf1m1GPBF9Ad0Jk/whtQ1663qI=
github.com/dnaeon/go-vcr v1.2.0/go.mod h1:R4UdLID7HZT3taECzJs4YgbbH6PIGXB6W/sc5OLb6RQ=
github.com/dnephin/pflag v1.0.7 h1:oxONGlWxhmUct0YzKTgrpQv9AUA1wtPBn7zuSjJqptk=
github.com/dnephin/pflag v1.0.7/go.mod h1:uxE91IoWURlOiTUIA8Mq5ZZkAv3dPUfZNaT80Zm7OQE=
github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM=

View File

@@ -15,27 +15,17 @@
package agents
import (
"github.com/openai/openai-go"
"github.com/openai/openai-go/option"
"go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/llm"
)
type (
Agent struct {
l *log.Logger
cfg Config
client *openai.Client
}
Config struct {
OpenAIAPIKey string
Temperature float64
ModelName string
}
)
func NewAgent(l *log.Logger, cfg Config) *Agent {
client := openai.NewClient(option.WithAPIKey(cfg.OpenAIAPIKey))
return &Agent{l: l, cfg: cfg, client: &client}
type Agent struct {
l *log.Logger
client *llm.Client
model string
temp float64
}
func NewAgent(l *log.Logger, client *llm.Client, model string, temp float64) *Agent {
return &Agent{l: l, client: client, model: model, temp: temp}
}

View File

@@ -18,8 +18,7 @@ import (
"context"
"fmt"
"github.com/openai/openai-go"
"github.com/openai/openai-go/packages/param"
"go.probo.inc/probo/pkg/llm"
)
const (
@@ -49,23 +48,19 @@ const (
)
func (a *Agent) GenerateChangelog(ctx context.Context, oldContent string, newContent string) (*string, error) {
model := openai.ChatModel(a.cfg.ModelName)
chatCompletion, err := a.client.Chat.Completions.New(ctx, openai.ChatCompletionNewParams{
Messages: []openai.ChatCompletionMessageParamUnion{
openai.SystemMessage(changelogGeneratorSystemPrompt),
openai.UserMessage(fmt.Sprintf(`Old content: %s`, oldContent)),
openai.UserMessage(fmt.Sprintf(`New content: %s`, newContent)),
resp, err := a.client.ChatCompletion(ctx, &llm.ChatCompletionRequest{
Model: a.model,
Messages: []llm.Message{
{Role: llm.RoleSystem, Parts: []llm.Part{llm.TextPart{Text: changelogGeneratorSystemPrompt}}},
{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: fmt.Sprintf(`Old content: %s`, oldContent)}}},
{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: fmt.Sprintf(`New content: %s`, newContent)}}},
},
Model: model,
Temperature: param.NewOpt(a.cfg.Temperature),
Temperature: &a.temp,
})
if err != nil {
return nil, fmt.Errorf("cannot parse vendor info: %w", err)
return nil, fmt.Errorf("cannot generate changelog: %w", err)
}
if len(chatCompletion.Choices) == 0 {
return nil, fmt.Errorf("no completion choices returned from API")
}
return &chatCompletion.Choices[0].Message.Content, nil
text := resp.Message.Text()
return &text, nil
}

View File

@@ -19,8 +19,7 @@ import (
"encoding/json"
"fmt"
"github.com/openai/openai-go"
"github.com/openai/openai-go/packages/param"
"go.probo.inc/probo/pkg/llm"
)
type (
@@ -123,28 +122,25 @@ const (
)
func (a *Agent) AssessVendor(ctx context.Context, websiteURL string) (*vendorInfo, error) {
model := openai.ChatModel(a.cfg.ModelName)
chatCompletion, err := a.client.Chat.Completions.New(ctx, openai.ChatCompletionNewParams{
Messages: []openai.ChatCompletionMessageParamUnion{
openai.SystemMessage(assessVendorSystemPrompt),
openai.UserMessage(websiteURL),
resp, err := a.client.ChatCompletion(
ctx,
&llm.ChatCompletionRequest{
Model: a.model,
Messages: []llm.Message{
{Role: llm.RoleSystem, Parts: []llm.Part{llm.TextPart{Text: assessVendorSystemPrompt}}},
{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: websiteURL}}},
},
Temperature: &a.temp,
},
Model: model,
Temperature: param.NewOpt(a.cfg.Temperature),
})
)
if err != nil {
return nil, fmt.Errorf("cannot assess vendor: %w", err)
}
var info vendorInfo
if err := json.Unmarshal([]byte(resp.Message.Text()), &info); err != nil {
return nil, fmt.Errorf("cannot parse vendor info: %w", err)
}
if len(chatCompletion.Choices) == 0 {
return nil, fmt.Errorf("no completion choices returned from API")
}
var vendorInfo vendorInfo
err = json.Unmarshal([]byte(chatCompletion.Choices[0].Message.Content), &vendorInfo)
if err != nil {
return nil, fmt.Errorf("cannot parse vendor info: %w", err)
}
return &vendorInfo, nil
return &info, nil
}

View File

@@ -0,0 +1,446 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package anthropic
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strconv"
"time"
"github.com/anthropics/anthropic-sdk-go"
"github.com/anthropics/anthropic-sdk-go/option"
"github.com/anthropics/anthropic-sdk-go/packages/param"
"github.com/anthropics/anthropic-sdk-go/packages/ssestream"
"go.probo.inc/probo/pkg/llm"
)
type (
Provider struct {
client *anthropic.Client
}
Option func(*config)
config struct {
httpClient *http.Client
baseURL string
requestTimeout time.Duration
maxRetries *int
}
)
func WithHTTPClient(c *http.Client) Option {
return func(cfg *config) { cfg.httpClient = c }
}
func WithBaseURL(url string) Option {
return func(cfg *config) { cfg.baseURL = url }
}
func WithRequestTimeout(d time.Duration) Option {
return func(cfg *config) { cfg.requestTimeout = d }
}
func WithMaxRetries(n int) Option {
return func(cfg *config) { cfg.maxRetries = &n }
}
func NewProvider(apiKey string, opts ...Option) *Provider {
var cfg config
for _, o := range opts {
o(&cfg)
}
reqOpts := []option.RequestOption{
option.WithAPIKey(apiKey),
}
if cfg.httpClient != nil {
reqOpts = append(reqOpts, option.WithHTTPClient(cfg.httpClient))
}
if cfg.baseURL != "" {
reqOpts = append(reqOpts, option.WithBaseURL(cfg.baseURL))
}
if cfg.requestTimeout > 0 {
reqOpts = append(reqOpts, option.WithRequestTimeout(cfg.requestTimeout))
}
if cfg.maxRetries != nil {
reqOpts = append(reqOpts, option.WithMaxRetries(*cfg.maxRetries))
}
client := anthropic.NewClient(reqOpts...)
return &Provider{client: &client}
}
func (p *Provider) ChatCompletion(ctx context.Context, req *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) {
params, err := buildParams(req)
if err != nil {
return nil, err
}
msg, err := p.client.Messages.New(ctx, params)
if err != nil {
return nil, mapError(err)
}
return mapResponse(msg), nil
}
func (p *Provider) ChatCompletionStream(ctx context.Context, req *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) {
params, err := buildParams(req)
if err != nil {
return nil, err
}
stream := p.client.Messages.NewStreaming(ctx, params)
return &anthropicStream{stream: stream}, nil
}
func buildParams(req *llm.ChatCompletionRequest) (anthropic.MessageNewParams, error) {
if req.MaxTokens == nil {
return anthropic.MessageNewParams{}, &llm.ErrContextLength{
Err: fmt.Errorf("MaxTokens is required for Anthropic"),
}
}
system, messages := extractSystem(req.Messages)
params := anthropic.MessageNewParams{
Model: anthropic.Model(req.Model),
MaxTokens: int64(*req.MaxTokens),
Messages: buildMessages(messages),
}
if len(system) > 0 {
blocks := make([]anthropic.TextBlockParam, len(system))
for i, s := range system {
blocks[i] = anthropic.TextBlockParam{Text: s}
}
params.System = blocks
}
if req.Temperature != nil {
params.Temperature = param.NewOpt(*req.Temperature)
}
if req.TopP != nil {
params.TopP = param.NewOpt(*req.TopP)
}
if len(req.StopSequences) > 0 {
params.StopSequences = req.StopSequences
}
if len(req.Tools) > 0 {
params.Tools = buildTools(req.Tools)
}
if req.ToolChoice != nil {
params.ToolChoice = buildToolChoice(req.ToolChoice)
}
return params, nil
}
func extractSystem(messages []llm.Message) (system []string, rest []llm.Message) {
for _, msg := range messages {
if msg.Role == llm.RoleSystem {
system = append(system, msg.Text())
} else {
rest = append(rest, msg)
}
}
return
}
func buildMessages(messages []llm.Message) []anthropic.MessageParam {
out := make([]anthropic.MessageParam, 0, len(messages))
for _, msg := range messages {
switch msg.Role {
case llm.RoleUser:
blocks := make([]anthropic.ContentBlockParamUnion, 0, len(msg.Parts))
for _, p := range msg.Parts {
switch p := p.(type) {
case llm.TextPart:
blocks = append(blocks, anthropic.NewTextBlock(p.Text))
case llm.ImagePart:
blocks = append(
blocks,
anthropic.NewImageBlock(
anthropic.URLImageSourceParam{
URL: p.URL,
},
),
)
}
}
out = append(out, anthropic.NewUserMessage(blocks...))
case llm.RoleAssistant:
blocks := []anthropic.ContentBlockParamUnion{
anthropic.NewTextBlock(msg.Text()),
}
for _, tc := range msg.ToolCalls {
var input any
_ = json.Unmarshal([]byte(tc.Function.Arguments), &input)
blocks = append(blocks, anthropic.NewToolUseBlock(tc.ID, input, tc.Function.Name))
}
out = append(out, anthropic.NewAssistantMessage(blocks...))
case llm.RoleTool:
out = append(
out,
anthropic.NewUserMessage(
anthropic.NewToolResultBlock(
msg.ToolCallID,
msg.Text(),
false,
),
),
)
}
}
return out
}
func buildTools(tools []llm.Tool) []anthropic.ToolUnionParam {
out := make([]anthropic.ToolUnionParam, len(tools))
for i, t := range tools {
tool := anthropic.ToolParam{
Name: t.Name,
Description: param.NewOpt(t.Description),
}
if t.Parameters != nil {
var schema map[string]any
if err := json.Unmarshal(t.Parameters, &schema); err == nil {
props := schema["properties"]
required, _ := schema["required"].([]any)
reqStrings := make([]string, 0, len(required))
for _, r := range required {
if s, ok := r.(string); ok {
reqStrings = append(reqStrings, s)
}
}
tool.InputSchema = anthropic.ToolInputSchemaParam{
Properties: props,
Required: reqStrings,
}
}
}
out[i] = anthropic.ToolUnionParam{OfTool: &tool}
}
return out
}
func buildToolChoice(tc *llm.ToolChoice) anthropic.ToolChoiceUnionParam {
switch tc.Type {
case llm.ToolChoiceAuto:
return anthropic.ToolChoiceUnionParam{OfAuto: &anthropic.ToolChoiceAutoParam{}}
case llm.ToolChoiceNone:
return anthropic.ToolChoiceUnionParam{OfNone: &anthropic.ToolChoiceNoneParam{}}
case llm.ToolChoiceRequired:
return anthropic.ToolChoiceUnionParam{OfAny: &anthropic.ToolChoiceAnyParam{}}
case llm.ToolChoiceFunction:
return anthropic.ToolChoiceUnionParam{OfTool: &anthropic.ToolChoiceToolParam{Name: tc.Function}}
default:
return anthropic.ToolChoiceUnionParam{}
}
}
func mapResponse(msg *anthropic.Message) *llm.ChatCompletionResponse {
resp := &llm.ChatCompletionResponse{
Model: string(msg.Model),
FinishReason: mapStopReason(msg.StopReason),
Usage: llm.Usage{
InputTokens: int(msg.Usage.InputTokens),
OutputTokens: int(msg.Usage.OutputTokens),
},
Message: llm.Message{
Role: llm.RoleAssistant,
},
}
for _, block := range msg.Content {
switch block.Type {
case "text":
resp.Message.Parts = append(resp.Message.Parts, llm.TextPart{Text: block.Text})
case "tool_use":
tu := block.AsToolUse()
resp.Message.ToolCalls = append(resp.Message.ToolCalls, llm.ToolCall{
ID: tu.ID,
Function: llm.FunctionCall{
Name: tu.Name,
Arguments: string(tu.Input),
},
})
}
}
return resp
}
func mapStopReason(reason anthropic.StopReason) llm.FinishReason {
switch reason {
case anthropic.StopReasonEndTurn, anthropic.StopReasonStopSequence:
return llm.FinishReasonStop
case anthropic.StopReasonMaxTokens:
return llm.FinishReasonLength
case anthropic.StopReasonToolUse:
return llm.FinishReasonToolCalls
default:
return llm.FinishReasonStop
}
}
func mapError(err error) error {
var apiErr *anthropic.Error
if !errors.As(err, &apiErr) {
return err
}
switch apiErr.StatusCode {
case http.StatusTooManyRequests:
retryAfter := parseRetryAfter(apiErr.Response)
return &llm.ErrRateLimit{RetryAfter: retryAfter, Err: err}
case http.StatusUnauthorized:
return &llm.ErrAuthentication{Err: err}
default:
return err
}
}
func parseRetryAfter(resp *http.Response) time.Duration {
if resp == nil {
return 0
}
h := resp.Header.Get("Retry-After")
if h == "" {
return 0
}
if secs, err := strconv.Atoi(h); err == nil {
return time.Duration(secs) * time.Second
}
return 0
}
// anthropicStream adapts an Anthropic SSE stream to our ChatCompletionStream interface.
type anthropicStream struct {
stream *ssestream.Stream[anthropic.MessageStreamEventUnion]
current llm.ChatCompletionStreamEvent
// Track tool call indices for mapping content_block_start events.
toolCallIndex int
}
func (s *anthropicStream) Next() bool {
for s.stream.Next() {
event := s.stream.Current()
mapped, ok := s.mapStreamEvent(&event)
if ok {
s.current = mapped
return true
}
}
return false
}
func (s *anthropicStream) Event() llm.ChatCompletionStreamEvent {
return s.current
}
func (s *anthropicStream) Err() error {
err := s.stream.Err()
if err != nil {
return mapError(err)
}
return nil
}
func (s *anthropicStream) Close() error {
return s.stream.Close()
}
func (s *anthropicStream) mapStreamEvent(event *anthropic.MessageStreamEventUnion) (llm.ChatCompletionStreamEvent, bool) {
switch event.Type {
case "content_block_start":
cb := event.ContentBlock
if cb.Type == "tool_use" {
tu := cb.AsToolUse()
return llm.ChatCompletionStreamEvent{
Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{{
Index: s.toolCallIndex,
ID: tu.ID,
Name: tu.Name,
}},
},
}, true
}
return llm.ChatCompletionStreamEvent{}, false
case "content_block_delta":
delta := event.Delta
switch delta.Type {
case "text_delta":
return llm.ChatCompletionStreamEvent{
Delta: llm.MessageDelta{Content: delta.Text},
}, true
case "input_json_delta":
return llm.ChatCompletionStreamEvent{
Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{{
Index: s.toolCallIndex,
Arguments: delta.PartialJSON,
}},
},
}, true
}
return llm.ChatCompletionStreamEvent{}, false
case "content_block_stop":
if event.ContentBlock.Type == "tool_use" {
s.toolCallIndex++
}
return llm.ChatCompletionStreamEvent{}, false
case "message_delta":
fr := mapStopReason(anthropic.StopReason(event.Delta.StopReason))
evt := llm.ChatCompletionStreamEvent{
FinishReason: &fr,
}
if event.Usage.OutputTokens > 0 || event.Usage.InputTokens > 0 {
evt.Usage = &llm.Usage{
InputTokens: int(event.Usage.InputTokens),
OutputTokens: int(event.Usage.OutputTokens),
}
}
return evt, true
case "message_start":
if event.Message.Usage.InputTokens > 0 {
return llm.ChatCompletionStreamEvent{
Usage: &llm.Usage{
InputTokens: int(event.Message.Usage.InputTokens),
OutputTokens: int(event.Message.Usage.OutputTokens),
},
}, true
}
return llm.ChatCompletionStreamEvent{}, false
default:
return llm.ChatCompletionStreamEvent{}, false
}
}

446
pkg/llm/bedrock/provider.go Normal file
View File

@@ -0,0 +1,446 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package bedrock
import (
"context"
"encoding/json"
"errors"
"strings"
"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime/document"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime/types"
smithyhttp "github.com/aws/smithy-go/transport/http"
"go.probo.inc/probo/pkg/llm"
)
type (
Provider struct {
client *bedrockruntime.Client
}
Option func(*bedrockruntime.Options)
)
// WithBaseEndpoint overrides the Bedrock service endpoint URL.
func WithBaseEndpoint(url string) Option {
return func(o *bedrockruntime.Options) { o.BaseEndpoint = &url }
}
func NewProvider(cfg aws.Config, opts ...Option) *Provider {
fns := make([]func(*bedrockruntime.Options), len(opts))
for i, o := range opts {
fns[i] = func(bo *bedrockruntime.Options) { o(bo) }
}
client := bedrockruntime.NewFromConfig(cfg, fns...)
return &Provider{client: client}
}
func (p *Provider) ChatCompletion(ctx context.Context, req *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) {
input := buildInput(req)
output, err := p.client.Converse(ctx, input)
if err != nil {
return nil, mapError(err)
}
return mapResponse(output, req.Model), nil
}
func (p *Provider) ChatCompletionStream(ctx context.Context, req *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) {
input := &bedrockruntime.ConverseStreamInput{
ModelId: aws.String(req.Model),
Messages: buildMessages(req.Messages),
InferenceConfig: buildInferenceConfig(req),
}
system := buildSystem(req.Messages)
if len(system) > 0 {
input.System = system
}
if len(req.Tools) > 0 {
toolConfig := buildToolConfig(req)
input.ToolConfig = toolConfig
}
output, err := p.client.ConverseStream(ctx, input)
if err != nil {
return nil, mapError(err)
}
return newBedrockStream(output.GetStream()), nil
}
func buildInput(req *llm.ChatCompletionRequest) *bedrockruntime.ConverseInput {
input := &bedrockruntime.ConverseInput{
ModelId: aws.String(req.Model),
Messages: buildMessages(req.Messages),
InferenceConfig: buildInferenceConfig(req),
}
system := buildSystem(req.Messages)
if len(system) > 0 {
input.System = system
}
if len(req.Tools) > 0 {
toolConfig := buildToolConfig(req)
input.ToolConfig = toolConfig
}
return input
}
func buildInferenceConfig(req *llm.ChatCompletionRequest) *types.InferenceConfiguration {
cfg := &types.InferenceConfiguration{}
if req.MaxTokens != nil {
v := int32(*req.MaxTokens)
cfg.MaxTokens = &v
}
if req.Temperature != nil {
v := float32(*req.Temperature)
cfg.Temperature = &v
}
if req.TopP != nil {
v := float32(*req.TopP)
cfg.TopP = &v
}
if len(req.StopSequences) > 0 {
cfg.StopSequences = req.StopSequences
}
return cfg
}
func buildSystem(messages []llm.Message) []types.SystemContentBlock {
var system []types.SystemContentBlock
for _, msg := range messages {
if msg.Role == llm.RoleSystem {
system = append(
system,
&types.SystemContentBlockMemberText{
Value: msg.Text(),
},
)
}
}
return system
}
func buildMessages(messages []llm.Message) []types.Message {
var out []types.Message
for _, msg := range messages {
switch msg.Role {
case llm.RoleSystem:
continue
case llm.RoleUser:
var content []types.ContentBlock
for _, p := range msg.Parts {
if tp, ok := p.(llm.TextPart); ok {
content = append(content, &types.ContentBlockMemberText{Value: tp.Text})
}
}
out = append(
out, types.Message{
Role: types.ConversationRoleUser,
Content: content,
},
)
case llm.RoleAssistant:
var content []types.ContentBlock
if text := msg.Text(); text != "" {
content = append(content, &types.ContentBlockMemberText{Value: text})
}
for _, tc := range msg.ToolCalls {
var input any
_ = json.Unmarshal([]byte(tc.Function.Arguments), &input)
content = append(
content,
&types.ContentBlockMemberToolUse{
Value: types.ToolUseBlock{
ToolUseId: aws.String(tc.ID),
Name: aws.String(tc.Function.Name),
Input: document.NewLazyDocument(input),
},
},
)
}
out = append(
out, types.Message{
Role: types.ConversationRoleAssistant,
Content: content,
},
)
case llm.RoleTool:
out = append(out, types.Message{
Role: types.ConversationRoleUser,
Content: []types.ContentBlock{
&types.ContentBlockMemberToolResult{
Value: types.ToolResultBlock{
ToolUseId: aws.String(msg.ToolCallID),
Content: []types.ToolResultContentBlock{
&types.ToolResultContentBlockMemberText{Value: msg.Text()},
},
},
},
},
})
}
}
return out
}
func buildToolConfig(req *llm.ChatCompletionRequest) *types.ToolConfiguration {
config := &types.ToolConfiguration{}
tools := make([]types.Tool, len(req.Tools))
for i, t := range req.Tools {
spec := types.ToolSpecification{
Name: aws.String(t.Name),
Description: aws.String(t.Description),
}
if t.Parameters != nil {
var schema any
_ = json.Unmarshal(t.Parameters, &schema)
spec.InputSchema = &types.ToolInputSchemaMemberJson{
Value: document.NewLazyDocument(schema),
}
}
tools[i] = &types.ToolMemberToolSpec{Value: spec}
}
config.Tools = tools
if req.ToolChoice != nil {
config.ToolChoice = buildToolChoice(req.ToolChoice)
}
return config
}
func buildToolChoice(tc *llm.ToolChoice) types.ToolChoice {
switch tc.Type {
case llm.ToolChoiceAuto:
return &types.ToolChoiceMemberAuto{Value: types.AutoToolChoice{}}
case llm.ToolChoiceRequired:
return &types.ToolChoiceMemberAny{Value: types.AnyToolChoice{}}
case llm.ToolChoiceFunction:
return &types.ToolChoiceMemberTool{
Value: types.SpecificToolChoice{
Name: aws.String(tc.Function),
},
}
case llm.ToolChoiceNone:
// Bedrock doesn't have a "none" tool choice; omit tools instead.
return nil
default:
return nil
}
}
func mapResponse(output *bedrockruntime.ConverseOutput, model string) *llm.ChatCompletionResponse {
resp := &llm.ChatCompletionResponse{
Model: model,
FinishReason: mapStopReason(output.StopReason),
Message: llm.Message{
Role: llm.RoleAssistant,
},
}
if output.Usage != nil {
resp.Usage = llm.Usage{
InputTokens: int(aws.ToInt32(output.Usage.InputTokens)),
OutputTokens: int(aws.ToInt32(output.Usage.OutputTokens)),
}
}
// Extract message content from the response output union.
if msgOutput, ok := output.Output.(*types.ConverseOutputMemberMessage); ok {
for _, block := range msgOutput.Value.Content {
switch b := block.(type) {
case *types.ContentBlockMemberText:
resp.Message.Parts = append(resp.Message.Parts, llm.TextPart{Text: b.Value})
case *types.ContentBlockMemberToolUse:
var args any
if b.Value.Input != nil {
_ = b.Value.Input.UnmarshalSmithyDocument(&args)
}
argsJSON, _ := json.Marshal(args)
resp.Message.ToolCalls = append(resp.Message.ToolCalls, llm.ToolCall{
ID: aws.ToString(b.Value.ToolUseId),
Function: llm.FunctionCall{
Name: aws.ToString(b.Value.Name),
Arguments: string(argsJSON),
},
})
}
}
}
return resp
}
func mapStopReason(reason types.StopReason) llm.FinishReason {
switch reason {
case types.StopReasonEndTurn, types.StopReasonStopSequence:
return llm.FinishReasonStop
case types.StopReasonMaxTokens:
return llm.FinishReasonLength
case types.StopReasonToolUse:
return llm.FinishReasonToolCalls
case types.StopReasonContentFiltered, types.StopReasonGuardrailIntervened:
return llm.FinishReasonContentFilter
default:
return llm.FinishReasonStop
}
}
func mapError(err error) error {
var respErr *smithyhttp.ResponseError
if !errors.As(err, &respErr) {
// Check for common error types by message content.
msg := err.Error()
if strings.Contains(msg, "throttling") || strings.Contains(msg, "ThrottlingException") {
return &llm.ErrRateLimit{Err: err}
}
return err
}
switch respErr.HTTPStatusCode() {
case 429:
return &llm.ErrRateLimit{Err: err}
case 401, 403:
return &llm.ErrAuthentication{Err: err}
case 400:
msg := err.Error()
if strings.Contains(msg, "context") || strings.Contains(msg, "token") {
return &llm.ErrContextLength{Err: err}
}
return err
default:
return err
}
}
// bedrockStream adapts a Bedrock ConverseStream to our ChatCompletionStream interface.
type bedrockStream struct {
eventStream *bedrockruntime.ConverseStreamEventStream
events <-chan types.ConverseStreamOutput
current llm.ChatCompletionStreamEvent
err error
toolIndex int
}
func newBedrockStream(eventStream *bedrockruntime.ConverseStreamEventStream) *bedrockStream {
return &bedrockStream{
eventStream: eventStream,
events: eventStream.Events(),
}
}
func (s *bedrockStream) Next() bool {
for event := range s.events {
mapped, ok := s.mapEvent(event)
if ok {
s.current = mapped
return true
}
}
if err := s.eventStream.Err(); err != nil {
s.err = mapError(err)
}
return false
}
func (s *bedrockStream) Event() llm.ChatCompletionStreamEvent {
return s.current
}
func (s *bedrockStream) Err() error {
return s.err
}
func (s *bedrockStream) Close() error {
return s.eventStream.Close()
}
func (s *bedrockStream) mapEvent(event types.ConverseStreamOutput) (llm.ChatCompletionStreamEvent, bool) {
switch e := event.(type) {
case *types.ConverseStreamOutputMemberContentBlockStart:
if start, ok := e.Value.Start.(*types.ContentBlockStartMemberToolUse); ok {
return llm.ChatCompletionStreamEvent{
Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{{
Index: s.toolIndex,
ID: aws.ToString(start.Value.ToolUseId),
Name: aws.ToString(start.Value.Name),
}},
},
}, true
}
return llm.ChatCompletionStreamEvent{}, false
case *types.ConverseStreamOutputMemberContentBlockDelta:
switch d := e.Value.Delta.(type) {
case *types.ContentBlockDeltaMemberText:
return llm.ChatCompletionStreamEvent{
Delta: llm.MessageDelta{Content: d.Value},
}, true
case *types.ContentBlockDeltaMemberToolUse:
return llm.ChatCompletionStreamEvent{
Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{{
Index: s.toolIndex,
Arguments: aws.ToString(d.Value.Input),
}},
},
}, true
}
return llm.ChatCompletionStreamEvent{}, false
case *types.ConverseStreamOutputMemberContentBlockStop:
s.toolIndex++
return llm.ChatCompletionStreamEvent{}, false
case *types.ConverseStreamOutputMemberMessageStop:
fr := mapStopReason(e.Value.StopReason)
return llm.ChatCompletionStreamEvent{
FinishReason: &fr,
}, true
case *types.ConverseStreamOutputMemberMetadata:
if e.Value.Usage != nil {
return llm.ChatCompletionStreamEvent{
Usage: &llm.Usage{
InputTokens: int(aws.ToInt32(e.Value.Usage.InputTokens)),
OutputTokens: int(aws.ToInt32(e.Value.Usage.OutputTokens)),
},
}, true
}
return llm.ChatCompletionStreamEvent{}, false
default:
return llm.ChatCompletionStreamEvent{}, false
}
}

229
pkg/llm/chat.go Normal file
View File

@@ -0,0 +1,229 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package llm
import (
"encoding/json"
"strings"
)
type (
ChatCompletionRequest struct {
Model string
Messages []Message
MaxTokens *int
Temperature *float64
TopP *float64
StopSequences []string
Tools []Tool
ToolChoice *ToolChoice
ResponseFormat *ResponseFormat
}
ToolChoiceType string
ToolChoice struct {
Type ToolChoiceType
Function string // required when Type is ToolChoiceFunction
}
ResponseFormatType string
ResponseFormat struct {
Type ResponseFormatType
JSONSchema *JSONSchema // required when Type is ResponseFormatJSONSchema
}
JSONSchema struct {
Name string
Description string
Schema json.RawMessage
}
FinishReason string
Usage struct {
InputTokens int
OutputTokens int
}
ChatCompletionResponse struct {
Model string
Message Message
Usage Usage
FinishReason FinishReason
}
// ChatCompletionStream is an iterator over streaming chat completion events.
// Callers must call Close when done, even if Next returns false.
// The typical usage pattern is:
//
// stream, err := client.ChatCompletionStream(ctx, req)
// if err != nil { ... }
// defer stream.Close()
// for stream.Next() {
// event := stream.Event()
// // process event
// }
// if err := stream.Err(); err != nil { ... }
ChatCompletionStream interface {
Next() bool
Event() ChatCompletionStreamEvent
Err() error
Close() error
}
ChatCompletionStreamEvent struct {
Delta MessageDelta
Usage *Usage // present on final event if provider supports it
FinishReason *FinishReason // present on final event
}
MessageDelta struct {
Content string
ToolCalls []ToolCallDelta
}
ToolCallDelta struct {
Index int
ID string // set on first chunk for this tool call
Name string // set on first chunk for this tool call
Arguments string // incremental JSON fragment
}
)
const (
ToolChoiceAuto ToolChoiceType = "auto"
ToolChoiceNone ToolChoiceType = "none"
ToolChoiceRequired ToolChoiceType = "required"
ToolChoiceFunction ToolChoiceType = "function"
)
const (
ResponseFormatText ResponseFormatType = "text"
ResponseFormatJSONObject ResponseFormatType = "json_object"
ResponseFormatJSONSchema ResponseFormatType = "json_schema"
)
const (
FinishReasonStop FinishReason = "stop"
FinishReasonToolCalls FinishReason = "tool_calls"
FinishReasonLength FinishReason = "length"
FinishReasonContentFilter FinishReason = "content_filter"
)
func (u Usage) Add(other Usage) Usage {
return Usage{
InputTokens: u.InputTokens + other.InputTokens,
OutputTokens: u.OutputTokens + other.OutputTokens,
}
}
// StreamAccumulator wraps a ChatCompletionStream and reassembles the
// streamed deltas into a full ChatCompletionResponse. It proxies
// Next/Event/Err/Close transparently so callers can still observe
// individual deltas while accumulating the final result.
//
// After the stream is exhausted (Next returns false), call Response
// to get the fully assembled ChatCompletionResponse.
type StreamAccumulator struct {
stream ChatCompletionStream
current ChatCompletionStreamEvent
content strings.Builder
toolCalls map[int]*ToolCall
usage Usage
finishReason FinishReason
model string
}
func NewStreamAccumulator(stream ChatCompletionStream) *StreamAccumulator {
return &StreamAccumulator{
stream: stream,
toolCalls: make(map[int]*ToolCall),
}
}
func (a *StreamAccumulator) Next() bool {
if !a.stream.Next() {
return false
}
a.current = a.stream.Event()
a.accumulate(a.current)
return true
}
func (a *StreamAccumulator) Event() ChatCompletionStreamEvent {
return a.current
}
func (a *StreamAccumulator) Err() error {
return a.stream.Err()
}
func (a *StreamAccumulator) Close() error {
return a.stream.Close()
}
// Response returns the fully assembled ChatCompletionResponse after
// the stream has been exhausted. Must only be called after Next
// returns false and Err returns nil.
func (a *StreamAccumulator) Response() *ChatCompletionResponse {
toolCalls := make([]ToolCall, 0, len(a.toolCalls))
for i := 0; i < len(a.toolCalls); i++ {
if tc, ok := a.toolCalls[i]; ok {
toolCalls = append(toolCalls, *tc)
}
}
return &ChatCompletionResponse{
Model: a.model,
Message: Message{
Role: RoleAssistant,
Parts: []Part{TextPart{Text: a.content.String()}},
ToolCalls: toolCalls,
},
Usage: a.usage,
FinishReason: a.finishReason,
}
}
func (a *StreamAccumulator) accumulate(event ChatCompletionStreamEvent) {
a.content.WriteString(event.Delta.Content)
for _, tcd := range event.Delta.ToolCalls {
tc, ok := a.toolCalls[tcd.Index]
if !ok {
tc = &ToolCall{}
a.toolCalls[tcd.Index] = tc
}
if tcd.ID != "" {
tc.ID = tcd.ID
}
if tcd.Name != "" {
tc.Function.Name = tcd.Name
}
tc.Function.Arguments += tcd.Arguments
}
if event.Usage != nil {
a.usage = *event.Usage
}
if event.FinishReason != nil {
a.finishReason = *event.FinishReason
}
}

70
pkg/llm/errors.go Normal file
View File

@@ -0,0 +1,70 @@
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package llm
import (
"fmt"
"time"
)
type (
ErrRateLimit struct {
RetryAfter time.Duration
Err error
}
ErrContextLength struct {
MaxTokens int
Err error
}
ErrContentFilter struct {
Err error
}
ErrAuthentication struct {
Err error
}
)
func (e *ErrRateLimit) Error() string {
if e.RetryAfter > 0 {
return fmt.Sprintf("rate limited (retry after %s): %v", e.RetryAfter, e.Err)
}
return fmt.Sprintf("rate limited: %v", e.Err)
}
func (e *ErrRateLimit) Unwrap() error { return e.Err }
func (e *ErrContextLength) Error() string {
if e.MaxTokens > 0 {
return fmt.Sprintf("context length exceeded (max %d tokens): %v", e.MaxTokens, e.Err)
}
return fmt.Sprintf("context length exceeded: %v", e.Err)
}
func (e *ErrContextLength) Unwrap() error { return e.Err }
func (e *ErrContentFilter) Error() string {
return fmt.Sprintf("content filtered: %v", e.Err)
}
func (e *ErrContentFilter) Unwrap() error { return e.Err }
func (e *ErrAuthentication) Error() string {
return fmt.Sprintf("authentication failed: %v", e.Err)
}
func (e *ErrAuthentication) Unwrap() error { return e.Err }

137
pkg/llm/llm.go Normal file
View File

@@ -0,0 +1,137 @@
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package llm
import (
"context"
"io"
"time"
"go.gearno.de/kit/log"
"go.opentelemetry.io/otel"
"go.opentelemetry.io/otel/trace"
)
var tracerName = "go.probo.inc/probo/pkg/llm"
type (
Option func(*Client)
Client struct {
provider Provider
system string
logger *log.Logger
tracerProvider trace.TracerProvider
tracer trace.Tracer
}
)
func WithLogger(l *log.Logger) Option {
return func(c *Client) {
c.logger = l
}
}
func WithTracerProvider(tp trace.TracerProvider) Option {
return func(c *Client) {
c.tracerProvider = tp
}
}
// NewClient creates a new instrumented LLM client.
// The system parameter identifies the provider for the OTel gen_ai.provider.name
// attribute (e.g., "openai", "anthropic", "aws.bedrock").
func NewClient(provider Provider, system string, opts ...Option) *Client {
c := &Client{
provider: provider,
system: system,
logger: log.NewLogger(log.WithOutput(io.Discard)),
tracerProvider: otel.GetTracerProvider(),
}
for _, opt := range opts {
opt(c)
}
c.logger = c.logger.Named("llm").With(log.String("system", system))
c.tracer = c.tracerProvider.Tracer(tracerName)
return c
}
func (c *Client) ChatCompletion(ctx context.Context, req *ChatCompletionRequest) (*ChatCompletionResponse, error) {
ctx, span := startChatSpan(ctx, c.tracer, c.system, req)
c.logger.InfoCtx(
ctx,
"chat completion request",
log.String("model", req.Model),
log.Int("message_count", len(req.Messages)),
log.Int("tool_count", len(req.Tools)),
)
start := time.Now()
resp, err := c.provider.ChatCompletion(ctx, req)
duration := time.Since(start)
if err != nil {
c.logger.ErrorCtx(
ctx,
"chat completion failed",
log.String("model", req.Model),
log.Duration("duration", duration),
log.Error(err),
)
endChatSpan(span, nil, err)
return nil, err
}
c.logger.InfoCtx(
ctx,
"chat completion response",
log.String("model", resp.Model),
log.Int("input_tokens", resp.Usage.InputTokens),
log.Int("output_tokens", resp.Usage.OutputTokens),
log.String("finish_reason", string(resp.FinishReason)),
log.Duration("duration", duration),
)
endChatSpan(span, resp, nil)
return resp, nil
}
func (c *Client) ChatCompletionStream(ctx context.Context, req *ChatCompletionRequest) (ChatCompletionStream, error) {
ctx, span := startChatSpan(ctx, c.tracer, c.system, req)
c.logger.InfoCtx(ctx, "chat completion stream request",
log.String("model", req.Model),
log.Int("message_count", len(req.Messages)),
log.Int("tool_count", len(req.Tools)),
)
stream, err := c.provider.ChatCompletionStream(ctx, req)
if err != nil {
c.logger.ErrorCtx(
ctx,
"chat completion stream failed",
log.String("model", req.Model),
log.Error(err),
)
endChatSpan(span, nil, err)
return nil, err
}
return newTracedStream(stream, span), nil
}

642
pkg/llm/llm_test.go Normal file
View File

@@ -0,0 +1,642 @@
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package llm_test
import (
"context"
"errors"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/codes"
sdktrace "go.opentelemetry.io/otel/sdk/trace"
"go.opentelemetry.io/otel/sdk/trace/tracetest"
"go.probo.inc/probo/pkg/llm"
)
// ---------------------------------------------------------------------------
// Helpers
// ---------------------------------------------------------------------------
type mockProvider struct {
chatResp *llm.ChatCompletionResponse
chatErr error
streamResp llm.ChatCompletionStream
streamErr error
}
func (m *mockProvider) ChatCompletion(_ context.Context, _ *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) {
return m.chatResp, m.chatErr
}
func (m *mockProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) {
return m.streamResp, m.streamErr
}
type mockStream struct {
events []llm.ChatCompletionStreamEvent
idx int
current llm.ChatCompletionStreamEvent
err error
closed bool
}
func (s *mockStream) Next() bool {
if s.idx >= len(s.events) {
return false
}
s.current = s.events[s.idx]
s.idx++
return true
}
func (s *mockStream) Event() llm.ChatCompletionStreamEvent { return s.current }
func (s *mockStream) Err() error { return s.err }
func (s *mockStream) Close() error { s.closed = true; return nil }
func newTestClient(provider llm.Provider) (*llm.Client, *tracetest.SpanRecorder) {
recorder := tracetest.NewSpanRecorder()
tp := sdktrace.NewTracerProvider(sdktrace.WithSpanProcessor(recorder))
client := llm.NewClient(provider, "test",
llm.WithTracerProvider(tp),
)
return client, recorder
}
func spanAttrMap(recorder *tracetest.SpanRecorder) map[string]any {
spans := recorder.Ended()
if len(spans) == 0 {
return nil
}
m := make(map[string]any)
for _, a := range spans[0].Attributes() {
m[string(a.Key)] = a.Value.AsInterface()
}
return m
}
func ptr[T any](v T) *T { return &v }
// ---------------------------------------------------------------------------
// Message.Text
// ---------------------------------------------------------------------------
func TestMessageText(t *testing.T) {
t.Parallel()
tests := []struct {
name string
msg llm.Message
want string
}{
{
name: "single text part",
msg: llm.Message{
Parts: []llm.Part{llm.TextPart{Text: "hello"}},
},
want: "hello",
},
{
name: "multiple text parts concatenated",
msg: llm.Message{
Parts: []llm.Part{
llm.TextPart{Text: "hello"},
llm.TextPart{Text: " world"},
},
},
want: "hello world",
},
{
name: "image parts skipped",
msg: llm.Message{
Parts: []llm.Part{
llm.TextPart{Text: "before"},
llm.ImagePart{URL: "http://example.com/img.png"},
llm.TextPart{Text: "after"},
},
},
want: "beforeafter",
},
{
name: "no parts returns empty string",
msg: llm.Message{},
want: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
assert.Equal(t, tt.want, tt.msg.Text())
})
}
}
// ---------------------------------------------------------------------------
// Usage.Add
// ---------------------------------------------------------------------------
func TestUsageAdd(t *testing.T) {
t.Parallel()
a := llm.Usage{InputTokens: 10, OutputTokens: 5}
b := llm.Usage{InputTokens: 20, OutputTokens: 15}
c := a.Add(b)
assert.Equal(t, 30, c.InputTokens)
assert.Equal(t, 20, c.OutputTokens)
}
// ---------------------------------------------------------------------------
// Error types
// ---------------------------------------------------------------------------
func TestErrors(t *testing.T) {
t.Parallel()
inner := errors.New("upstream")
t.Run("ErrRateLimit", func(t *testing.T) {
t.Parallel()
t.Run("with retry after", func(t *testing.T) {
t.Parallel()
e := &llm.ErrRateLimit{RetryAfter: 30 * time.Second, Err: inner}
assert.Contains(t, e.Error(), "retry after 30s")
assert.Contains(t, e.Error(), "upstream")
assert.ErrorIs(t, e, inner)
})
t.Run("without retry after", func(t *testing.T) {
t.Parallel()
e := &llm.ErrRateLimit{Err: inner}
assert.Contains(t, e.Error(), "rate limited")
assert.NotContains(t, e.Error(), "retry after")
assert.ErrorIs(t, e, inner)
})
t.Run("errors.As", func(t *testing.T) {
t.Parallel()
var target *llm.ErrRateLimit
e := &llm.ErrRateLimit{RetryAfter: 5 * time.Second, Err: inner}
require.ErrorAs(t, e, &target)
assert.Equal(t, 5*time.Second, target.RetryAfter)
})
})
t.Run("ErrContextLength", func(t *testing.T) {
t.Parallel()
t.Run("with max tokens", func(t *testing.T) {
t.Parallel()
e := &llm.ErrContextLength{MaxTokens: 4096, Err: inner}
assert.Contains(t, e.Error(), "4096")
assert.ErrorIs(t, e, inner)
})
t.Run("without max tokens", func(t *testing.T) {
t.Parallel()
e := &llm.ErrContextLength{Err: inner}
assert.Contains(t, e.Error(), "context length exceeded")
assert.NotContains(t, e.Error(), "max")
assert.ErrorIs(t, e, inner)
})
})
t.Run("ErrContentFilter", func(t *testing.T) {
t.Parallel()
e := &llm.ErrContentFilter{Err: inner}
assert.Contains(t, e.Error(), "content filtered")
assert.ErrorIs(t, e, inner)
})
t.Run("ErrAuthentication", func(t *testing.T) {
t.Parallel()
e := &llm.ErrAuthentication{Err: inner}
assert.Contains(t, e.Error(), "authentication failed")
assert.ErrorIs(t, e, inner)
})
}
// ---------------------------------------------------------------------------
// Client — ChatCompletion
// ---------------------------------------------------------------------------
func TestChatCompletion(t *testing.T) {
t.Parallel()
t.Run("success", func(t *testing.T) {
t.Parallel()
provider := &mockProvider{
chatResp: &llm.ChatCompletionResponse{
Model: "test-model",
Message: llm.Message{
Role: llm.RoleAssistant,
Parts: []llm.Part{llm.TextPart{Text: "Hello!"}},
},
Usage: llm.Usage{InputTokens: 10, OutputTokens: 5},
FinishReason: llm.FinishReasonStop,
},
}
client, recorder := newTestClient(provider)
temp := 0.7
resp, err := client.ChatCompletion(context.Background(), &llm.ChatCompletionRequest{
Model: "test-model",
Messages: []llm.Message{
{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Hi"}}},
},
Temperature: &temp,
})
require.NoError(t, err)
assert.Equal(t, "test-model", resp.Model)
assert.Equal(t, "Hello!", resp.Message.Text())
assert.Equal(t, 10, resp.Usage.InputTokens)
assert.Equal(t, 5, resp.Usage.OutputTokens)
assert.Equal(t, llm.FinishReasonStop, resp.FinishReason)
spans := recorder.Ended()
require.Len(t, spans, 1)
assert.Equal(t, "chat test-model", spans[0].Name())
attrs := spanAttrMap(recorder)
assert.Equal(t, "test", attrs["gen_ai.provider.name"])
assert.Equal(t, "test-model", attrs["gen_ai.request.model"])
assert.Equal(t, 0.7, attrs["gen_ai.request.temperature"])
assert.Equal(t, "test-model", attrs["gen_ai.response.model"])
assert.Equal(t, int64(10), attrs["gen_ai.usage.input_tokens"])
assert.Equal(t, int64(5), attrs["gen_ai.usage.output_tokens"])
})
t.Run("all span attributes", func(t *testing.T) {
t.Parallel()
provider := &mockProvider{
chatResp: &llm.ChatCompletionResponse{
Model: "gpt-4",
FinishReason: llm.FinishReasonStop,
Usage: llm.Usage{InputTokens: 50, OutputTokens: 25},
Message: llm.Message{
Role: llm.RoleAssistant,
Parts: []llm.Part{llm.TextPart{Text: "ok"}},
},
},
}
client, recorder := newTestClient(provider)
maxTokens := 1024
topP := 0.9
resp, err := client.ChatCompletion(context.Background(), &llm.ChatCompletionRequest{
Model: "gpt-4",
Messages: []llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "test"}}}},
MaxTokens: &maxTokens,
TopP: &topP,
StopSequences: []string{"END", "STOP"},
})
require.NoError(t, err)
require.NotNil(t, resp)
attrs := spanAttrMap(recorder)
assert.Equal(t, int64(1024), attrs["gen_ai.request.max_tokens"])
assert.Equal(t, 0.9, attrs["gen_ai.request.top_p"])
assert.Equal(t, []string{"END", "STOP"}, attrs["gen_ai.request.stop_sequences"])
})
t.Run("error sets span status", func(t *testing.T) {
t.Parallel()
provider := &mockProvider{
chatErr: errors.New("provider error"),
}
client, recorder := newTestClient(provider)
_, err := client.ChatCompletion(context.Background(), &llm.ChatCompletionRequest{
Model: "test-model",
Messages: []llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Hi"}}}},
})
require.Error(t, err)
assert.Contains(t, err.Error(), "provider error")
spans := recorder.Ended()
require.Len(t, spans, 1)
assert.Equal(t, codes.Error, spans[0].Status().Code)
})
}
// ---------------------------------------------------------------------------
// Client — ChatCompletionStream
// ---------------------------------------------------------------------------
func TestChatCompletionStream(t *testing.T) {
t.Parallel()
t.Run("success collects deltas", func(t *testing.T) {
t.Parallel()
events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "Hello"}},
{Delta: llm.MessageDelta{Content: " world"}},
{
FinishReason: ptr(llm.FinishReasonStop),
Usage: &llm.Usage{InputTokens: 8, OutputTokens: 4},
},
}
client, recorder := newTestClient(&mockProvider{
streamResp: &mockStream{events: events},
})
stream, err := client.ChatCompletionStream(context.Background(), &llm.ChatCompletionRequest{
Model: "test-model",
Messages: []llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Hi"}}}},
})
require.NoError(t, err)
var collected []string
for stream.Next() {
e := stream.Event()
if e.Delta.Content != "" {
collected = append(collected, e.Delta.Content)
}
}
require.NoError(t, stream.Err())
require.NoError(t, stream.Close())
assert.Equal(t, []string{"Hello", " world"}, collected)
spans := recorder.Ended()
require.Len(t, spans, 1)
assert.Equal(t, "chat test-model", spans[0].Name())
})
t.Run("provider error sets span status", func(t *testing.T) {
t.Parallel()
client, recorder := newTestClient(&mockProvider{
streamErr: errors.New("stream open failed"),
})
_, err := client.ChatCompletionStream(context.Background(), &llm.ChatCompletionRequest{
Model: "test-model",
Messages: []llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Hi"}}}},
})
require.Error(t, err)
assert.Contains(t, err.Error(), "stream open failed")
spans := recorder.Ended()
require.Len(t, spans, 1)
assert.Equal(t, codes.Error, spans[0].Status().Code)
})
t.Run("inner stream error records error on span", func(t *testing.T) {
t.Parallel()
ms := &mockStream{
events: []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "partial"}},
},
err: errors.New("connection reset"),
}
client, recorder := newTestClient(&mockProvider{streamResp: ms})
stream, err := client.ChatCompletionStream(context.Background(), &llm.ChatCompletionRequest{
Model: "test-model",
Messages: []llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Hi"}}}},
})
require.NoError(t, err)
for stream.Next() {
}
assert.ErrorContains(t, stream.Err(), "connection reset")
_ = stream.Close()
spans := recorder.Ended()
require.Len(t, spans, 1)
assert.Equal(t, codes.Error, spans[0].Status().Code)
})
t.Run("span finalized on close before exhausting", func(t *testing.T) {
t.Parallel()
events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "a"}},
{Delta: llm.MessageDelta{Content: "b"}},
{Delta: llm.MessageDelta{Content: "c"}},
}
client, recorder := newTestClient(&mockProvider{
streamResp: &mockStream{events: events},
})
stream, err := client.ChatCompletionStream(context.Background(), &llm.ChatCompletionRequest{
Model: "test-model",
Messages: []llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Hi"}}}},
})
require.NoError(t, err)
// Read only one event, then close early.
require.True(t, stream.Next())
require.NoError(t, stream.Close())
spans := recorder.Ended()
require.Len(t, spans, 1, "span should be ended by Close even without exhausting stream")
})
t.Run("stream span records usage and finish reason", func(t *testing.T) {
t.Parallel()
events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "done"}},
{
FinishReason: ptr(llm.FinishReasonLength),
Usage: &llm.Usage{InputTokens: 100, OutputTokens: 50},
},
}
client, recorder := newTestClient(&mockProvider{
streamResp: &mockStream{events: events},
})
stream, err := client.ChatCompletionStream(context.Background(), &llm.ChatCompletionRequest{
Model: "test-model",
Messages: []llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "Hi"}}}},
})
require.NoError(t, err)
for stream.Next() {
}
require.NoError(t, stream.Err())
require.NoError(t, stream.Close())
spans := recorder.Ended()
require.Len(t, spans, 1)
attrs := make(map[string]any)
for _, a := range spans[0].Attributes() {
attrs[string(a.Key)] = a.Value.AsInterface()
}
assert.Equal(t, int64(100), attrs["gen_ai.usage.input_tokens"])
assert.Equal(t, int64(50), attrs["gen_ai.usage.output_tokens"])
assert.Equal(t, []string{"length"}, attrs["gen_ai.response.finish_reasons"])
})
}
// ---------------------------------------------------------------------------
// StreamAccumulator
// ---------------------------------------------------------------------------
func TestStreamAccumulator(t *testing.T) {
t.Parallel()
t.Run("text and single tool call", func(t *testing.T) {
t.Parallel()
events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "Hello"}},
{Delta: llm.MessageDelta{Content: " world"}},
{Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{
{Index: 0, ID: "tc_1", Name: "get_weather"},
},
}},
{Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{
{Index: 0, Arguments: `{"city":`},
},
}},
{Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{
{Index: 0, Arguments: `"Paris"}`},
},
}},
{
FinishReason: ptr(llm.FinishReasonToolCalls),
Usage: &llm.Usage{InputTokens: 20, OutputTokens: 15},
},
}
acc := llm.NewStreamAccumulator(&mockStream{events: events})
for acc.Next() {
}
require.NoError(t, acc.Err())
resp := acc.Response()
assert.Equal(t, "Hello world", resp.Message.Text())
assert.Equal(t, llm.RoleAssistant, resp.Message.Role)
assert.Equal(t, llm.FinishReasonToolCalls, resp.FinishReason)
assert.Equal(t, 20, resp.Usage.InputTokens)
assert.Equal(t, 15, resp.Usage.OutputTokens)
require.Len(t, resp.Message.ToolCalls, 1)
tc := resp.Message.ToolCalls[0]
assert.Equal(t, "tc_1", tc.ID)
assert.Equal(t, "get_weather", tc.Function.Name)
assert.Equal(t, `{"city":"Paris"}`, tc.Function.Arguments)
})
t.Run("multiple tool calls at different indices", func(t *testing.T) {
t.Parallel()
events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{
{Index: 0, ID: "tc_a", Name: "search"},
},
}},
{Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{
{Index: 0, Arguments: `{"q":"go"}`},
},
}},
{Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{
{Index: 1, ID: "tc_b", Name: "fetch"},
},
}},
{Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{
{Index: 1, Arguments: `{"url":"https://example.com"}`},
},
}},
{
FinishReason: ptr(llm.FinishReasonToolCalls),
Usage: &llm.Usage{InputTokens: 30, OutputTokens: 10},
},
}
acc := llm.NewStreamAccumulator(&mockStream{events: events})
for acc.Next() {
}
require.NoError(t, acc.Err())
resp := acc.Response()
require.Len(t, resp.Message.ToolCalls, 2)
assert.Equal(t, "tc_a", resp.Message.ToolCalls[0].ID)
assert.Equal(t, "search", resp.Message.ToolCalls[0].Function.Name)
assert.Equal(t, `{"q":"go"}`, resp.Message.ToolCalls[0].Function.Arguments)
assert.Equal(t, "tc_b", resp.Message.ToolCalls[1].ID)
assert.Equal(t, "fetch", resp.Message.ToolCalls[1].Function.Name)
assert.Equal(t, `{"url":"https://example.com"}`, resp.Message.ToolCalls[1].Function.Arguments)
})
t.Run("text only without tool calls", func(t *testing.T) {
t.Parallel()
events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "Just text."}},
{
FinishReason: ptr(llm.FinishReasonStop),
Usage: &llm.Usage{InputTokens: 5, OutputTokens: 3},
},
}
acc := llm.NewStreamAccumulator(&mockStream{events: events})
for acc.Next() {
}
require.NoError(t, acc.Err())
resp := acc.Response()
assert.Equal(t, "Just text.", resp.Message.Text())
assert.Equal(t, llm.FinishReasonStop, resp.FinishReason)
assert.Empty(t, resp.Message.ToolCalls)
})
t.Run("proxies events transparently", func(t *testing.T) {
t.Parallel()
events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "a"}},
{Delta: llm.MessageDelta{Content: "b"}},
{FinishReason: ptr(llm.FinishReasonStop)},
}
acc := llm.NewStreamAccumulator(&mockStream{events: events})
var seen []string
for acc.Next() {
e := acc.Event()
if e.Delta.Content != "" {
seen = append(seen, e.Delta.Content)
}
}
assert.Equal(t, []string{"a", "b"}, seen)
})
}

52
pkg/llm/message.go Normal file
View File

@@ -0,0 +1,52 @@
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package llm
import "encoding/json"
type (
Message struct {
Role Role
Parts []Part
ToolCalls []ToolCall
ToolCallID string // set when Role is RoleTool
}
ToolCall struct {
ID string
Function FunctionCall
}
FunctionCall struct {
Name string
Arguments string // JSON-encoded arguments
}
Tool struct {
Name string
Description string
Parameters json.RawMessage // JSON Schema
}
)
func (m Message) Text() string {
var s string
for _, p := range m.Parts {
if tp, ok := p.(TextPart); ok {
s += tp.Text
}
}
return s
}

437
pkg/llm/openai/provider.go Normal file
View File

@@ -0,0 +1,437 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package openai
import (
"context"
"encoding/json"
"errors"
"net/http"
"strconv"
"time"
"github.com/openai/openai-go"
"github.com/openai/openai-go/option"
"github.com/openai/openai-go/packages/param"
"github.com/openai/openai-go/packages/ssestream"
"github.com/openai/openai-go/shared"
"go.probo.inc/probo/pkg/llm"
)
type (
Provider struct {
client *openai.Client
}
Option func(*config)
config struct {
httpClient *http.Client
baseURL string
organization string
project string
requestTimeout time.Duration
maxRetries *int
}
)
func WithHTTPClient(c *http.Client) Option {
return func(cfg *config) { cfg.httpClient = c }
}
func WithBaseURL(url string) Option {
return func(cfg *config) { cfg.baseURL = url }
}
func WithOrganization(org string) Option {
return func(cfg *config) { cfg.organization = org }
}
func WithProject(project string) Option {
return func(cfg *config) { cfg.project = project }
}
func WithRequestTimeout(d time.Duration) Option {
return func(cfg *config) { cfg.requestTimeout = d }
}
func WithMaxRetries(n int) Option {
return func(cfg *config) { cfg.maxRetries = &n }
}
func NewProvider(apiKey string, opts ...Option) *Provider {
var cfg config
for _, o := range opts {
o(&cfg)
}
reqOpts := []option.RequestOption{
option.WithAPIKey(apiKey),
}
if cfg.httpClient != nil {
reqOpts = append(reqOpts, option.WithHTTPClient(cfg.httpClient))
}
if cfg.baseURL != "" {
reqOpts = append(reqOpts, option.WithBaseURL(cfg.baseURL))
}
if cfg.organization != "" {
reqOpts = append(reqOpts, option.WithOrganization(cfg.organization))
}
if cfg.project != "" {
reqOpts = append(reqOpts, option.WithProject(cfg.project))
}
if cfg.requestTimeout > 0 {
reqOpts = append(reqOpts, option.WithRequestTimeout(cfg.requestTimeout))
}
if cfg.maxRetries != nil {
reqOpts = append(reqOpts, option.WithMaxRetries(*cfg.maxRetries))
}
client := openai.NewClient(reqOpts...)
return &Provider{client: &client}
}
func (p *Provider) ChatCompletion(ctx context.Context, req *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) {
params := buildParams(req)
completion, err := p.client.Chat.Completions.New(ctx, params)
if err != nil {
return nil, mapError(err)
}
return mapResponse(completion), nil
}
func (p *Provider) ChatCompletionStream(ctx context.Context, req *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) {
params := buildParams(req)
params.StreamOptions = openai.ChatCompletionStreamOptionsParam{
IncludeUsage: param.NewOpt(true),
}
stream := p.client.Chat.Completions.NewStreaming(ctx, params)
return &openaiStream{stream: stream}, nil
}
func buildParams(req *llm.ChatCompletionRequest) openai.ChatCompletionNewParams {
params := openai.ChatCompletionNewParams{
Model: openai.ChatModel(req.Model),
Messages: buildMessages(req.Messages),
}
if req.MaxTokens != nil {
params.MaxCompletionTokens = param.NewOpt(int64(*req.MaxTokens))
}
if req.Temperature != nil {
params.Temperature = param.NewOpt(*req.Temperature)
}
if req.TopP != nil {
params.TopP = param.NewOpt(*req.TopP)
}
if len(req.StopSequences) > 0 {
params.Stop = openai.ChatCompletionNewParamsStopUnion{
OfStringArray: req.StopSequences,
}
}
if len(req.Tools) > 0 {
params.Tools = buildTools(req.Tools)
}
if req.ToolChoice != nil {
params.ToolChoice = buildToolChoice(req.ToolChoice)
}
if req.ResponseFormat != nil {
params.ResponseFormat = buildResponseFormat(req.ResponseFormat)
}
return params
}
func buildMessages(messages []llm.Message) []openai.ChatCompletionMessageParamUnion {
out := make([]openai.ChatCompletionMessageParamUnion, 0, len(messages))
for _, msg := range messages {
switch msg.Role {
case llm.RoleSystem:
out = append(out, openai.SystemMessage(msg.Text()))
case llm.RoleUser:
parts := make([]openai.ChatCompletionContentPartUnionParam, 0, len(msg.Parts))
for _, p := range msg.Parts {
switch p := p.(type) {
case llm.TextPart:
parts = append(parts, openai.TextContentPart(p.Text))
case llm.ImagePart:
parts = append(parts, openai.ImageContentPart(openai.ChatCompletionContentPartImageImageURLParam{
URL: p.URL,
}))
}
}
out = append(out, openai.UserMessage(parts))
case llm.RoleAssistant:
m := openai.ChatCompletionAssistantMessageParam{
Content: openai.ChatCompletionAssistantMessageParamContentUnion{
OfString: param.NewOpt(msg.Text()),
},
}
if len(msg.ToolCalls) > 0 {
m.ToolCalls = make([]openai.ChatCompletionMessageToolCallParam, len(msg.ToolCalls))
for i, tc := range msg.ToolCalls {
m.ToolCalls[i] = openai.ChatCompletionMessageToolCallParam{
ID: tc.ID,
Function: openai.ChatCompletionMessageToolCallFunctionParam{
Name: tc.Function.Name,
Arguments: tc.Function.Arguments,
},
}
}
}
out = append(out, openai.ChatCompletionMessageParamUnion{OfAssistant: &m})
case llm.RoleTool:
out = append(out, openai.ToolMessage(msg.Text(), msg.ToolCallID))
}
}
return out
}
func buildTools(tools []llm.Tool) []openai.ChatCompletionToolParam {
out := make([]openai.ChatCompletionToolParam, len(tools))
for i, t := range tools {
fn := shared.FunctionDefinitionParam{
Name: t.Name,
Description: param.NewOpt(t.Description),
Strict: param.NewOpt(true),
}
if t.Parameters != nil {
var params shared.FunctionParameters
if err := json.Unmarshal(t.Parameters, &params); err == nil {
fn.Parameters = params
}
}
out[i] = openai.ChatCompletionToolParam{Function: fn}
}
return out
}
func buildToolChoice(tc *llm.ToolChoice) openai.ChatCompletionToolChoiceOptionUnionParam {
switch tc.Type {
case llm.ToolChoiceAuto:
return openai.ChatCompletionToolChoiceOptionUnionParam{
OfAuto: param.NewOpt(string(openai.ChatCompletionToolChoiceOptionAutoAuto)),
}
case llm.ToolChoiceNone:
return openai.ChatCompletionToolChoiceOptionUnionParam{
OfAuto: param.NewOpt(string(openai.ChatCompletionToolChoiceOptionAutoNone)),
}
case llm.ToolChoiceRequired:
return openai.ChatCompletionToolChoiceOptionUnionParam{
OfAuto: param.NewOpt(string(openai.ChatCompletionToolChoiceOptionAutoRequired)),
}
case llm.ToolChoiceFunction:
return openai.ChatCompletionToolChoiceOptionParamOfChatCompletionNamedToolChoice(
openai.ChatCompletionNamedToolChoiceFunctionParam{Name: tc.Function},
)
default:
return openai.ChatCompletionToolChoiceOptionUnionParam{}
}
}
func buildResponseFormat(rf *llm.ResponseFormat) openai.ChatCompletionNewParamsResponseFormatUnion {
switch rf.Type {
case llm.ResponseFormatText:
return openai.ChatCompletionNewParamsResponseFormatUnion{
OfText: &shared.ResponseFormatTextParam{},
}
case llm.ResponseFormatJSONObject:
return openai.ChatCompletionNewParamsResponseFormatUnion{
OfJSONObject: &shared.ResponseFormatJSONObjectParam{},
}
case llm.ResponseFormatJSONSchema:
if rf.JSONSchema != nil {
schema := shared.ResponseFormatJSONSchemaJSONSchemaParam{
Name: rf.JSONSchema.Name,
Strict: param.NewOpt(true),
}
if rf.JSONSchema.Description != "" {
schema.Description = param.NewOpt(rf.JSONSchema.Description)
}
if rf.JSONSchema.Schema != nil {
schema.Schema = rf.JSONSchema.Schema
}
return openai.ChatCompletionNewParamsResponseFormatUnion{
OfJSONSchema: &shared.ResponseFormatJSONSchemaParam{JSONSchema: schema},
}
}
return openai.ChatCompletionNewParamsResponseFormatUnion{}
default:
return openai.ChatCompletionNewParamsResponseFormatUnion{}
}
}
func mapResponse(c *openai.ChatCompletion) *llm.ChatCompletionResponse {
resp := &llm.ChatCompletionResponse{
Model: c.Model,
Usage: llm.Usage{
InputTokens: int(c.Usage.PromptTokens),
OutputTokens: int(c.Usage.CompletionTokens),
},
}
if len(c.Choices) > 0 {
choice := c.Choices[0]
resp.FinishReason = mapFinishReason(choice.FinishReason)
resp.Message = llm.Message{
Role: llm.RoleAssistant,
Parts: []llm.Part{llm.TextPart{Text: choice.Message.Content}},
}
if len(choice.Message.ToolCalls) > 0 {
resp.Message.ToolCalls = make([]llm.ToolCall, len(choice.Message.ToolCalls))
for i, tc := range choice.Message.ToolCalls {
resp.Message.ToolCalls[i] = llm.ToolCall{
ID: tc.ID,
Function: llm.FunctionCall{
Name: tc.Function.Name,
Arguments: tc.Function.Arguments,
},
}
}
}
}
return resp
}
func mapFinishReason(reason string) llm.FinishReason {
switch reason {
case "stop":
return llm.FinishReasonStop
case "tool_calls":
return llm.FinishReasonToolCalls
case "length":
return llm.FinishReasonLength
case "content_filter":
return llm.FinishReasonContentFilter
default:
return llm.FinishReasonStop
}
}
func mapError(err error) error {
var apiErr *openai.Error
if !errors.As(err, &apiErr) {
return err
}
switch apiErr.StatusCode {
case http.StatusTooManyRequests:
retryAfter := parseRetryAfter(apiErr.Response)
return &llm.ErrRateLimit{RetryAfter: retryAfter, Err: err}
case http.StatusUnauthorized:
return &llm.ErrAuthentication{Err: err}
case http.StatusBadRequest:
if apiErr.Code == "context_length_exceeded" {
return &llm.ErrContextLength{Err: err}
}
if apiErr.Code == "content_filter" {
return &llm.ErrContentFilter{Err: err}
}
return err
default:
return err
}
}
func parseRetryAfter(resp *http.Response) time.Duration {
if resp == nil {
return 0
}
h := resp.Header.Get("Retry-After")
if h == "" {
return 0
}
if secs, err := strconv.Atoi(h); err == nil {
return time.Duration(secs) * time.Second
}
return 0
}
// openaiStream adapts an OpenAI SSE stream to our ChatCompletionStream interface.
type openaiStream struct {
stream *ssestream.Stream[openai.ChatCompletionChunk]
current llm.ChatCompletionStreamEvent
}
func (s *openaiStream) Next() bool {
if !s.stream.Next() {
return false
}
chunk := s.stream.Current()
s.current = mapChunkToEvent(&chunk)
return true
}
func (s *openaiStream) Event() llm.ChatCompletionStreamEvent {
return s.current
}
func (s *openaiStream) Err() error {
err := s.stream.Err()
if err != nil {
return mapError(err)
}
return nil
}
func (s *openaiStream) Close() error {
return s.stream.Close()
}
func mapChunkToEvent(chunk *openai.ChatCompletionChunk) llm.ChatCompletionStreamEvent {
event := llm.ChatCompletionStreamEvent{}
if chunk.Usage.PromptTokens > 0 || chunk.Usage.CompletionTokens > 0 {
usage := llm.Usage{
InputTokens: int(chunk.Usage.PromptTokens),
OutputTokens: int(chunk.Usage.CompletionTokens),
}
event.Usage = &usage
}
if len(chunk.Choices) > 0 {
choice := chunk.Choices[0]
delta := choice.Delta
event.Delta.Content = delta.Content
if len(delta.ToolCalls) > 0 {
event.Delta.ToolCalls = make([]llm.ToolCallDelta, len(delta.ToolCalls))
for i, tc := range delta.ToolCalls {
event.Delta.ToolCalls[i] = llm.ToolCallDelta{
Index: int(tc.Index),
ID: tc.ID,
Name: tc.Function.Name,
Arguments: tc.Function.Arguments,
}
}
}
if choice.FinishReason != "" {
fr := mapFinishReason(choice.FinishReason)
event.FinishReason = &fr
}
}
return event
}

32
pkg/llm/part.go Normal file
View File

@@ -0,0 +1,32 @@
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package llm
type (
Part interface {
part()
}
TextPart struct {
Text string
}
ImagePart struct {
URL string
}
)
func (TextPart) part() {}
func (ImagePart) part() {}

22
pkg/llm/provider.go Normal file
View File

@@ -0,0 +1,22 @@
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package llm
import "context"
type Provider interface {
ChatCompletion(ctx context.Context, req *ChatCompletionRequest) (*ChatCompletionResponse, error)
ChatCompletionStream(ctx context.Context, req *ChatCompletionRequest) (ChatCompletionStream, error)
}

24
pkg/llm/role.go Normal file
View File

@@ -0,0 +1,24 @@
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package llm
type Role string
const (
RoleSystem Role = "system"
RoleUser Role = "user"
RoleAssistant Role = "assistant"
RoleTool Role = "tool"
)

143
pkg/llm/trace.go Normal file
View File

@@ -0,0 +1,143 @@
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package llm
import (
"context"
"fmt"
"sync"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/codes"
semconv "go.opentelemetry.io/otel/semconv/v1.37.0"
"go.opentelemetry.io/otel/trace"
)
func startChatSpan(ctx context.Context, tracer trace.Tracer, system string, req *ChatCompletionRequest) (context.Context, trace.Span) {
spanName := fmt.Sprintf("chat %s", req.Model)
attrs := []attribute.KeyValue{
semconv.GenAIOperationNameChat,
semconv.GenAIProviderNameKey.String(system),
semconv.GenAIRequestModel(req.Model),
}
if req.Temperature != nil {
attrs = append(attrs, semconv.GenAIRequestTemperature(*req.Temperature))
}
if req.MaxTokens != nil {
attrs = append(attrs, semconv.GenAIRequestMaxTokens(*req.MaxTokens))
}
if req.TopP != nil {
attrs = append(attrs, semconv.GenAIRequestTopP(*req.TopP))
}
if len(req.StopSequences) > 0 {
attrs = append(attrs, semconv.GenAIRequestStopSequences(req.StopSequences...))
}
return tracer.Start(ctx, spanName,
trace.WithSpanKind(trace.SpanKindClient),
trace.WithAttributes(attrs...),
)
}
func endChatSpan(span trace.Span, resp *ChatCompletionResponse, err error) {
if err != nil {
span.RecordError(err)
span.SetStatus(codes.Error, err.Error())
span.End()
return
}
span.SetAttributes(
semconv.GenAIResponseModel(resp.Model),
semconv.GenAIUsageInputTokens(resp.Usage.InputTokens),
semconv.GenAIUsageOutputTokens(resp.Usage.OutputTokens),
semconv.GenAIResponseFinishReasons(string(resp.FinishReason)),
)
span.End()
}
// tracedStream wraps a ChatCompletionStream and manages the OTel span
// lifecycle for streaming calls. The span is ended when Close is called
// or when Next returns false (whichever comes first).
type tracedStream struct {
inner ChatCompletionStream
span trace.Span
lastEvent ChatCompletionStreamEvent
closeOnce sync.Once
finishReason *FinishReason
usage *Usage
}
func newTracedStream(inner ChatCompletionStream, span trace.Span) *tracedStream {
return &tracedStream{
inner: inner,
span: span,
}
}
func (s *tracedStream) Next() bool {
if !s.inner.Next() {
s.finalizeSpan()
return false
}
s.lastEvent = s.inner.Event()
if s.lastEvent.FinishReason != nil {
s.finishReason = s.lastEvent.FinishReason
}
if s.lastEvent.Usage != nil {
s.usage = s.lastEvent.Usage
}
return true
}
func (s *tracedStream) Event() ChatCompletionStreamEvent {
return s.lastEvent
}
func (s *tracedStream) Err() error {
return s.inner.Err()
}
func (s *tracedStream) Close() error {
s.finalizeSpan()
return s.inner.Close()
}
func (s *tracedStream) finalizeSpan() {
s.closeOnce.Do(func() {
if err := s.inner.Err(); err != nil {
s.span.RecordError(err)
s.span.SetStatus(codes.Error, err.Error())
s.span.End()
return
}
var attrs []attribute.KeyValue
if s.usage != nil {
attrs = append(attrs,
semconv.GenAIUsageInputTokens(s.usage.InputTokens),
semconv.GenAIUsageOutputTokens(s.usage.OutputTokens),
)
}
if s.finishReason != nil {
attrs = append(attrs, semconv.GenAIResponseFinishReasons(string(*s.finishReason)))
}
if len(attrs) > 0 {
s.span.SetAttributes(attrs...)
}
s.span.End()
})
}

View File

@@ -24,6 +24,7 @@ import (
"go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/agents"
"go.probo.inc/probo/pkg/certmanager"
"go.probo.inc/probo/pkg/llm"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/crypto/cipher"
"go.probo.inc/probo/pkg/esign"
@@ -55,7 +56,9 @@ type (
encryptionKey cipher.EncryptionKey
baseURL string
tokenSecret string
agentConfig agents.Config
llmClient *llm.Client
llmModel string
llmTemperature float64
html2pdfConverter *html2pdf.Converter
acmeService *certmanager.ACMEService
fileManager *filemanager.Service
@@ -126,7 +129,9 @@ func NewService(
bucket string,
baseURL string,
tokenSecret string,
agentConfig agents.Config,
llmClient *llm.Client,
llmModel string,
llmTemperature float64,
html2pdfConverter *html2pdf.Converter,
acmeService *certmanager.ACMEService,
fileManagerService *filemanager.Service,
@@ -149,7 +154,9 @@ func NewService(
encryptionKey: encryptionKey,
baseURL: baseURL,
tokenSecret: tokenSecret,
agentConfig: agentConfig,
llmClient: llmClient,
llmModel: llmModel,
llmTemperature: llmTemperature,
html2pdfConverter: html2pdfConverter,
acmeService: acmeService,
fileManager: fileManagerService,
@@ -171,7 +178,7 @@ func (s *Service) WithTenant(tenantID gid.TenantID) *TenantService {
baseURL: s.baseURL,
scope: coredata.NewScope(tenantID),
tokenSecret: s.tokenSecret,
agent: agents.NewAgent(nil, s.agentConfig),
agent: agents.NewAgent(nil, s.llmClient, s.llmModel, s.llmTemperature),
fileManager: s.fileManager,
esign: s.esign,
}

56
pkg/probod/llm.go Normal file
View File

@@ -0,0 +1,56 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package probod
import (
"fmt"
"github.com/prometheus/client_golang/prometheus"
"go.gearno.de/kit/httpclient"
"go.gearno.de/kit/log"
"go.opentelemetry.io/otel/trace"
"go.probo.inc/probo/pkg/llm"
llmopenai "go.probo.inc/probo/pkg/llm/openai"
)
func buildLLMClient(cfg LLMConfig, l *log.Logger, tp trace.TracerProvider, r prometheus.Registerer) (*llm.Client, error) {
provider := cfg.Provider
if provider == "" {
provider = "openai"
}
httpClient := httpclient.DefaultPooledClient(
httpclient.WithLogger(l),
httpclient.WithTracerProvider(tp),
httpclient.WithRegisterer(r),
)
switch provider {
case "openai":
p := llmopenai.NewProvider(cfg.APIKey,
llmopenai.WithHTTPClient(httpClient),
)
return llm.NewClient(p, "openai",
llm.WithLogger(l),
llm.WithTracerProvider(tp),
), nil
case "anthropic":
return nil, fmt.Errorf("anthropic provider not yet wired; add import and construct here")
case "bedrock":
return nil, fmt.Errorf("bedrock provider not yet wired; requires aws.Config")
default:
return nil, fmt.Errorf("unsupported LLM provider: %q", provider)
}
}

View File

@@ -14,8 +14,11 @@
package probod
type OpenAIConfig struct {
APIKey string `json:"api-key"`
Temperature float64 `json:"temperature"`
type LLMConfig struct {
Provider string `json:"provider"` // "openai", "anthropic", "bedrock"
APIKey string `json:"api-key"` // for OpenAI and Anthropic
ModelName string `json:"model-name"`
Temperature float64 `json:"temperature"`
}
type OpenAIConfig = LLMConfig

View File

@@ -42,7 +42,6 @@ import (
"go.gearno.de/kit/pg"
"go.gearno.de/kit/unit"
"go.opentelemetry.io/otel/trace"
"go.probo.inc/probo/pkg/agents"
"go.probo.inc/probo/pkg/awsconfig"
"go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/certmanager"
@@ -315,14 +314,11 @@ func (impl *Implm) Run(
}
}
agentConfig := agents.Config{
OpenAIAPIKey: impl.cfg.OpenAI.APIKey,
Temperature: impl.cfg.OpenAI.Temperature,
ModelName: impl.cfg.OpenAI.ModelName,
llmClient, err := buildLLMClient(impl.cfg.OpenAI, l.Named("llm"), tp, r)
if err != nil {
return fmt.Errorf("cannot create LLM client: %w", err)
}
agent := agents.NewAgent(l.Named("agent"), agentConfig)
fileManagerService := filemanager.NewService(s3Client)
var samlCert *x509.Certificate
@@ -446,7 +442,9 @@ func (impl *Implm) Run(
impl.cfg.AWS.Bucket,
baseURL.String(),
impl.cfg.Auth.Cookie.Secret,
agentConfig,
llmClient,
impl.cfg.OpenAI.ModelName,
impl.cfg.OpenAI.Temperature,
html2pdfConverter,
acmeService,
fileManagerService,
@@ -486,7 +484,7 @@ func (impl *Implm) Run(
Slack: slackService,
ConnectorRegistry: defaultConnectorRegistry,
BaseURL: baseURL,
Agent: agent,
CustomDomainCname: impl.cfg.CustomDomains.CnameTarget,
TokenSecret: impl.cfg.Auth.Cookie.Secret,
Logger: l.Named("http.server"),

View File

@@ -22,7 +22,6 @@ import (
"github.com/go-chi/chi/v5"
"go.gearno.de/kit/httpserver"
"go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/agents"
"go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/connector"
"go.probo.inc/probo/pkg/esign"
@@ -52,7 +51,6 @@ type Config struct {
Cookie securecookie.Config
TokenSecret string
ConnectorRegistry *connector.ConnectorRegistry
Agent *agents.Agent
CustomDomainCname string
Logger *log.Logger
}