82 lines
2.8 KiB
Go
82 lines
2.8 KiB
Go
package proposal_generation
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"strings"
|
|
|
|
"gitea.maximumdirect.net/eric/audita/internal/core/schema"
|
|
"gitea.maximumdirect.net/eric/audita/internal/framework/contracts"
|
|
"gitea.maximumdirect.net/eric/audita/internal/framework/stagename"
|
|
"gitea.maximumdirect.net/eric/audita/internal/prompts"
|
|
)
|
|
|
|
type ProposalMessageBuilder func(transcript *schema.Transcript, glossary *schema.Glossary, sectionIndex int, transcriptDescription string) ([]contracts.LLMMessage, error)
|
|
|
|
type ModuleProposalRequest struct {
|
|
ProposalRequest contracts.ProposalRequest
|
|
PromptID string
|
|
BuildMessages ProposalMessageBuilder
|
|
}
|
|
|
|
// ExecuteModuleProposal runs shared proposal generation plumbing for one
|
|
// module, leaving only module-specific prompt message building at call sites.
|
|
func ExecuteModuleProposal(ctx context.Context, req ModuleProposalRequest) (contracts.ProposalResult, error) {
|
|
if req.BuildMessages == nil {
|
|
return contracts.ProposalResult{}, fmt.Errorf("proposal message builder is required")
|
|
}
|
|
if strings.TrimSpace(req.PromptID) == "" {
|
|
return contracts.ProposalResult{}, fmt.Errorf("prompt ID must not be empty")
|
|
}
|
|
|
|
sectionIndex := 0
|
|
if req.ProposalRequest.Section != nil {
|
|
sectionIndex = req.ProposalRequest.Section.Index
|
|
}
|
|
|
|
transcriptDescription := ""
|
|
if req.ProposalRequest.Config != nil {
|
|
transcriptDescription = req.ProposalRequest.Config.TranscriptDescription
|
|
}
|
|
|
|
messages, err := req.BuildMessages(
|
|
req.ProposalRequest.WorkingTranscript,
|
|
req.ProposalRequest.Glossary,
|
|
sectionIndex,
|
|
transcriptDescription,
|
|
)
|
|
if err != nil {
|
|
return contracts.ProposalResult{}, err
|
|
}
|
|
|
|
promptMetadata, ok := prompts.LookupMetadata(req.PromptID)
|
|
if !ok {
|
|
return contracts.ProposalResult{}, fmt.Errorf("unknown prompt ID %q", req.PromptID)
|
|
}
|
|
|
|
generated, err := GenerateCandidates(ctx, Request{
|
|
ModuleKey: req.ProposalRequest.RunSpec.ModuleKey,
|
|
ModuleInstance: req.ProposalRequest.RunSpec.InstanceName,
|
|
ReplacementPolicy: req.ProposalRequest.RunSpec.ReplacementPolicy,
|
|
WorkingTranscript: req.ProposalRequest.WorkingTranscript,
|
|
Section: req.ProposalRequest.Section,
|
|
Glossary: req.ProposalRequest.Glossary,
|
|
Config: req.ProposalRequest.Config,
|
|
Messages: messages,
|
|
PromptMetadata: promptMetadata.DiagnosticsMap(),
|
|
StageName: stagename.ModuleProposal(req.ProposalRequest.RunSpec.InstanceName, sectionIndexPtr(req.ProposalRequest.Section)),
|
|
StartIndex: 0,
|
|
LLMClient: req.ProposalRequest.LLMClient,
|
|
Scheduler: req.ProposalRequest.LLMScheduler,
|
|
DiagnosticsDir: req.ProposalRequest.DiagnosticsDir,
|
|
})
|
|
if err != nil {
|
|
return contracts.ProposalResult{}, err
|
|
}
|
|
|
|
return contracts.ProposalResult{
|
|
Proposals: generated.Corrections,
|
|
Warnings: generated.Warnings,
|
|
}, nil
|
|
}
|