diff --git a/go.mod b/go.mod index b7552b034..6b797af1c 100644 --- a/go.mod +++ b/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 diff --git a/go.sum b/go.sum index 0f0b998e9..9d008c521 100644 --- a/go.sum +++ b/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= diff --git a/pkg/agents/agents.go b/pkg/agents/agents.go index 900d06879..ff5af3f21 100644 --- a/pkg/agents/agents.go +++ b/pkg/agents/agents.go @@ -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} } diff --git a/pkg/agents/changelog_generator.go b/pkg/agents/changelog_generator.go index 582ab8280..190073c4b 100644 --- a/pkg/agents/changelog_generator.go +++ b/pkg/agents/changelog_generator.go @@ -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 } diff --git a/pkg/agents/vendor_assessment.go b/pkg/agents/vendor_assessment.go index 8c8d0a9dc..c89fd1b7e 100644 --- a/pkg/agents/vendor_assessment.go +++ b/pkg/agents/vendor_assessment.go @@ -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 } diff --git a/pkg/llm/anthropic/provider.go b/pkg/llm/anthropic/provider.go new file mode 100644 index 000000000..529f321a3 --- /dev/null +++ b/pkg/llm/anthropic/provider.go @@ -0,0 +1,446 @@ +// Copyright (c) 2025 Probo Inc . +// +// 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 + } +} diff --git a/pkg/llm/bedrock/provider.go b/pkg/llm/bedrock/provider.go new file mode 100644 index 000000000..726b20137 --- /dev/null +++ b/pkg/llm/bedrock/provider.go @@ -0,0 +1,446 @@ +// Copyright (c) 2025 Probo Inc . +// +// 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 + } +} diff --git a/pkg/llm/chat.go b/pkg/llm/chat.go new file mode 100644 index 000000000..9b963b809 --- /dev/null +++ b/pkg/llm/chat.go @@ -0,0 +1,229 @@ +// Copyright (c) 2025 Probo Inc . +// +// 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 + } +} diff --git a/pkg/llm/errors.go b/pkg/llm/errors.go new file mode 100644 index 000000000..c1dc97990 --- /dev/null +++ b/pkg/llm/errors.go @@ -0,0 +1,70 @@ +// Copyright (c) 2026 Probo Inc . +// +// 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 } diff --git a/pkg/llm/llm.go b/pkg/llm/llm.go new file mode 100644 index 000000000..012c82978 --- /dev/null +++ b/pkg/llm/llm.go @@ -0,0 +1,137 @@ +// Copyright (c) 2026 Probo Inc . +// +// 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 +} diff --git a/pkg/llm/llm_test.go b/pkg/llm/llm_test.go new file mode 100644 index 000000000..5ab961aca --- /dev/null +++ b/pkg/llm/llm_test.go @@ -0,0 +1,642 @@ +// Copyright (c) 2026 Probo Inc . +// +// 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) + }) +} diff --git a/pkg/llm/message.go b/pkg/llm/message.go new file mode 100644 index 000000000..1b8e06b0a --- /dev/null +++ b/pkg/llm/message.go @@ -0,0 +1,52 @@ +// Copyright (c) 2026 Probo Inc . +// +// 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 +} diff --git a/pkg/llm/openai/provider.go b/pkg/llm/openai/provider.go new file mode 100644 index 000000000..e8bfbe12c --- /dev/null +++ b/pkg/llm/openai/provider.go @@ -0,0 +1,437 @@ +// Copyright (c) 2025 Probo Inc . +// +// 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 +} diff --git a/pkg/llm/part.go b/pkg/llm/part.go new file mode 100644 index 000000000..8383e0c6e --- /dev/null +++ b/pkg/llm/part.go @@ -0,0 +1,32 @@ +// Copyright (c) 2026 Probo Inc . +// +// 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() {} diff --git a/pkg/llm/provider.go b/pkg/llm/provider.go new file mode 100644 index 000000000..76fbaef98 --- /dev/null +++ b/pkg/llm/provider.go @@ -0,0 +1,22 @@ +// Copyright (c) 2026 Probo Inc . +// +// 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) +} diff --git a/pkg/llm/role.go b/pkg/llm/role.go new file mode 100644 index 000000000..f4440eea4 --- /dev/null +++ b/pkg/llm/role.go @@ -0,0 +1,24 @@ +// Copyright (c) 2026 Probo Inc . +// +// 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" +) diff --git a/pkg/llm/trace.go b/pkg/llm/trace.go new file mode 100644 index 000000000..a1b6bd373 --- /dev/null +++ b/pkg/llm/trace.go @@ -0,0 +1,143 @@ +// Copyright (c) 2026 Probo Inc . +// +// 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() + }) +} diff --git a/pkg/probo/service.go b/pkg/probo/service.go index 2793fa726..03087c0ef 100644 --- a/pkg/probo/service.go +++ b/pkg/probo/service.go @@ -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, } diff --git a/pkg/probod/llm.go b/pkg/probod/llm.go new file mode 100644 index 000000000..0ec1abdf1 --- /dev/null +++ b/pkg/probod/llm.go @@ -0,0 +1,56 @@ +// Copyright (c) 2025 Probo Inc . +// +// 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) + } +} diff --git a/pkg/probod/openai_config.go b/pkg/probod/openai_config.go index b61b9e5a9..af6e6bf90 100644 --- a/pkg/probod/openai_config.go +++ b/pkg/probod/openai_config.go @@ -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 diff --git a/pkg/probod/probod.go b/pkg/probod/probod.go index ff4b47ec7..f1fb66465 100644 --- a/pkg/probod/probod.go +++ b/pkg/probod/probod.go @@ -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"), diff --git a/pkg/server/server.go b/pkg/server/server.go index 59145a5e5..2d0ab3851 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -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 }