14
go.mod
14
go.mod
@@ -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
26
go.sum
@@ -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=
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
446
pkg/llm/anthropic/provider.go
Normal file
446
pkg/llm/anthropic/provider.go
Normal 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
446
pkg/llm/bedrock/provider.go
Normal 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
229
pkg/llm/chat.go
Normal 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
70
pkg/llm/errors.go
Normal 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
137
pkg/llm/llm.go
Normal 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
642
pkg/llm/llm_test.go
Normal 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
52
pkg/llm/message.go
Normal 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
437
pkg/llm/openai/provider.go
Normal 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, ¶ms); 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
32
pkg/llm/part.go
Normal 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
22
pkg/llm/provider.go
Normal 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
24
pkg/llm/role.go
Normal 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
143
pkg/llm/trace.go
Normal 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()
|
||||
})
|
||||
}
|
||||
@@ -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
56
pkg/probod/llm.go
Normal 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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user