Files
lassistanoque/backend/internal/adapter/llm/anyllm.go
T

250 lines
5.5 KiB
Go

package llm
import (
"context"
"fmt"
"log/slog"
anyllm "github.com/mozilla-ai/any-llm-go"
"github.com/mozilla-ai/any-llm-go/providers"
"github.com/mozilla-ai/any-llm-go/providers/anthropic"
"github.com/mozilla-ai/any-llm-go/providers/ollama"
"github.com/mozilla-ai/any-llm-go/providers/openai"
"trankilou.fr/lassistanoque/backend/internal/domain"
)
var providerTypes = []domain.Item{
{Value: "anthropic", Label: "Anthropic"},
{Value: "openai", Label: "OpenAI"},
{Value: "ollama", Label: "Ollama"},
{Value: "openaicomp", Label: "OpenAI compatible"},
{Value: "openrouter", Label: "Openrouter"},
}
type AnyLLMEngine struct {
}
func NewAnyLLMEngine() *AnyLLMEngine {
return &AnyLLMEngine{}
}
func (e *AnyLLMEngine) ListProviderTypes() []domain.Item {
return providerTypes
}
func providerFactory(provider *domain.Provider) (anyllm.Provider, error) {
switch provider.Type {
case "ollama":
return ollama.New(
anyllm.WithBaseURL(provider.URL),
)
case "openai":
return openai.New(
anyllm.WithAPIKey(provider.APIKey),
)
case "openaicomp":
return openai.New(
anyllm.WithBaseURL(provider.URL),
anyllm.WithAPIKey(provider.APIKey),
)
case "anthropic":
return anthropic.New(
anyllm.WithAPIKey(provider.APIKey),
)
}
return nil, fmt.Errorf("unknown provider type: %s", provider.Type)
}
func (e *AnyLLMEngine) ListModelsFromProvider(ctx context.Context, provider *domain.Provider) ([]string, error) {
models := make([]string, 0)
prov, err := providerFactory(provider)
if err != nil {
return nil, err
}
if lister, ok := prov.(anyllm.ModelLister); ok {
response, err := lister.ListModels(ctx)
if err != nil {
return nil, err
}
for _, m := range response.Data {
models = append(models, m.ID)
}
} else {
models = append(models, "Default")
}
return models, nil
}
func (e *AnyLLMEngine) Stream(
ctx context.Context,
provider *domain.Provider,
modelID string,
params *domain.LLMParams,
messages []*domain.Message,
) *domain.Message {
p, err := providerFactory(provider)
if err != nil {
return &domain.Message{
Role: string(domain.RoleAssistant),
Content: err.Error(),
}
}
anyllmMessages := make([]anyllm.Message, 0)
for _, m := range messages {
anyllmMessages = append(anyllmMessages, anyllm.Message{
Role: m.Role,
Content: m.Content,
})
}
if params.ReasoningEffort == "" {
params.ReasoningEffort = "none"
}
tools := make([]providers.Tool, len(params.Tools))
for i, t := range params.Tools {
tools[i] = toolConvert(t)
}
toolCalling := false
toolCalls := make([]domain.ToolCall, 0)
chunkChan, errChan := p.CompletionStream(ctx, anyllm.CompletionParams{
Model: modelID,
Messages: anyllmMessages,
Stream: true,
ReasoningEffort: providers.ReasoningEffort(params.ReasoningEffort),
Tools: tools,
})
fullContent := ""
for chunk := range chunkChan {
if len(chunk.Choices) > 0 {
content := chunk.Choices[0].Delta.Content
reasoning := chunk.Choices[0].Delta.Reasoning
if params.OnChunk != nil {
if content != "" {
fullContent += content
params.OnChunk(&domain.Chunk{
Done: false,
Role: anyllm.RoleAssistant,
Content: content,
Reasoning: false,
})
} else if reasoning != nil {
params.OnChunk(&domain.Chunk{
Done: false,
Role: anyllm.RoleAssistant,
Content: reasoning.Content,
Reasoning: true,
})
}
}
for _, tc := range chunk.Choices[0].Delta.ToolCalls {
toolCalling = true
slog.Info("Model is calling tool", "tool", tc.Function.Name, "arguments", tc.Function.Arguments)
extra := make(map[string]map[string]any)
for k, v := range tc.Extra {
m := map[string]any(v)
extra[k] = m
}
toolCalls = append(toolCalls, domain.ToolCall{
ID: tc.ID,
Type: tc.Type,
Function: domain.ToolCallFunction{
Name: tc.Function.Name,
Arguments: tc.Function.Arguments,
},
Extra: extra,
})
params.OnChunk(&domain.Chunk{
Done: false,
ToolName: tc.Function.Name,
Role: anyllm.RoleAssistant,
Content: tc.Function.Arguments,
Reasoning: false,
})
}
}
}
if !toolCalling {
params.OnChunk(&domain.Chunk{
Done: true,
Role: anyllm.RoleAssistant,
Content: "",
})
}
if err := <-errChan; err != nil {
return &domain.Message{
Role: string(domain.RoleAssistant),
Content: err.Error(),
}
}
return &domain.Message{
Role: string(domain.RoleAssistant),
Content: fullContent,
ToolCalls: toolCalls,
}
}
func (e *AnyLLMEngine) Generate(
ctx context.Context,
provider *domain.Provider,
modelID string,
messages []*domain.Message,
) (string, error) {
p, err := providerFactory(provider)
if err != nil {
return "", err
}
anyllmMessages := make([]anyllm.Message, 0)
for _, m := range messages {
anyllmMessages = append(anyllmMessages, anyllm.Message{
Role: m.Role,
Content: m.Content,
})
}
response, err := p.Completion(ctx, anyllm.CompletionParams{
Model: modelID,
Messages: anyllmMessages,
Stream: false,
})
if err != nil {
return "", err
}
return response.Choices[0].Message.Content.(string), nil
}
func toolConvert(domainTool *domain.ToolDefinition) providers.Tool {
return providers.Tool{
Type: domainTool.Type,
Function: providers.Function{
Name: domainTool.Function.Name,
Description: domainTool.Function.Description,
Parameters: domainTool.Function.Parameters,
},
}
}