diff --git a/pkg/probo/document_service.go b/pkg/probo/document_service.go index bd62c2912..89756622e 100644 --- a/pkg/probo/document_service.go +++ b/pkg/probo/document_service.go @@ -506,9 +506,9 @@ func (s DocumentService) generateChangelog( "changelog_generator", s.svc.llmClient, agent.WithInstructions(changelogGeneratorSystemPrompt), - agent.WithModel(s.svc.llmModel), - agent.WithTemperature(s.svc.llmTemperature), - agent.WithMaxTokens(s.svc.llmMaxTokens), + agent.WithModel(s.svc.llmConfig.Model), + agent.WithTemperature(s.svc.llmConfig.Temperature), + agent.WithMaxTokens(s.svc.llmConfig.MaxTokens), ) result, err := ag.Run( diff --git a/pkg/probo/service.go b/pkg/probo/service.go index 42965c867..2a95f0856 100644 --- a/pkg/probo/service.go +++ b/pkg/probo/service.go @@ -59,6 +59,14 @@ type ExportService interface { } type ( + // LLMConfig holds the model parameters used by the agents the probo + // service runs. + LLMConfig struct { + Model string + Temperature float64 + MaxTokens int + } + Service struct { pg *pg.Client s3 *s3.Client @@ -67,9 +75,7 @@ type ( baseURL string tokenSecret string llmClient *llm.Client - llmModel string - llmTemperature float64 - llmMaxTokens int + llmConfig LLMConfig html2pdfConverter *html2pdf.Converter acmeService *certmanager.ACMEService fileManager *filemanager.Service @@ -129,9 +135,7 @@ func NewService( baseURL string, tokenSecret string, llmClient *llm.Client, - llmModel string, - llmTemperature float64, - llmMaxTokens int, + llmConfig LLMConfig, html2pdfConverter *html2pdf.Converter, acmeService *certmanager.ACMEService, fileManagerService *filemanager.Service, @@ -157,9 +161,7 @@ func NewService( baseURL: baseURL, tokenSecret: tokenSecret, llmClient: llmClient, - llmModel: llmModel, - llmTemperature: llmTemperature, - llmMaxTokens: llmMaxTokens, + llmConfig: llmConfig, html2pdfConverter: html2pdfConverter, acmeService: acmeService, fileManager: fileManagerService, diff --git a/pkg/probod/probod.go b/pkg/probod/probod.go index a18ff1a96..619b6db75 100644 --- a/pkg/probod/probod.go +++ b/pkg/probod/probod.go @@ -39,6 +39,7 @@ import ( "go.gearno.de/kit/pg" "go.gearno.de/kit/unit" "go.gearno.de/kit/worker" + "go.gearno.de/x/ref" "go.opentelemetry.io/otel/trace" "go.probo.inc/probo/packages/emails" "go.probo.inc/probo/pkg/accessreview" @@ -507,9 +508,11 @@ func (impl *Implm) Run( baseURL.String(), impl.cfg.Auth.Cookie.Secret, proboLLMClient, - proboAgentCfg.ModelName, - *proboAgentCfg.Temperature, - *proboAgentCfg.MaxTokens, + probo.LLMConfig{ + Model: proboAgentCfg.ModelName, + Temperature: ref.UnrefOrZero(proboAgentCfg.Temperature), + MaxTokens: ref.UnrefOrZero(proboAgentCfg.MaxTokens), + }, html2pdfConverter, acmeService, fileManagerService, @@ -790,8 +793,8 @@ func (impl *Implm) Run( evidenceDescriberLLMClient, evidencedescriber.Config{ Model: evidenceDescriberAgentCfg.ModelName, - Temp: *evidenceDescriberAgentCfg.Temperature, - MaxTokens: *evidenceDescriberAgentCfg.MaxTokens, + Temp: ref.UnrefOrZero(evidenceDescriberAgentCfg.Temperature), + MaxTokens: ref.UnrefOrZero(evidenceDescriberAgentCfg.MaxTokens), }, ) evidenceDescriptionWorker := probo.NewEvidenceDescriptionWorker( diff --git a/pkg/probodconfig/llm_config.go b/pkg/probodconfig/llm_config.go index b745bde31..3cda73fde 100644 --- a/pkg/probodconfig/llm_config.go +++ b/pkg/probodconfig/llm_config.go @@ -96,12 +96,8 @@ func (c *AgentsConfig) ResolveAgent(agent LLMAgentConfig) LLMAgentConfig { agent.ModelName = c.Default.ModelName } - if agent.Temperature == nil { - agent.Temperature = c.Default.Temperature - } - - if agent.MaxTokens == nil { - agent.MaxTokens = c.Default.MaxTokens + if agent.MaxTokens == nil && c.Default.MaxTokens != nil { + agent.MaxTokens = new(*c.Default.MaxTokens) } return agent