250 lines
5.5 KiB
Go
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,
|
|
},
|
|
}
|
|
}
|