mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-20 06:31:02 +08:00
Go: add new provider: google (#14395)
### What problem does this PR solve? As title. ### Type of change - [x] New Feature (non-breaking change which adds functionality) --------- Signed-off-by: Jin Hai <haijin.chn@gmail.com>
This commit is contained in:
@@ -149,8 +149,8 @@ type Features struct {
|
||||
}
|
||||
|
||||
type ModelThinking struct {
|
||||
DefaultValue bool `json:"default_value"`
|
||||
ClearContent bool `json:"clear_content"`
|
||||
DefaultValue bool `json:"default_value"`
|
||||
ClearThinking bool `json:"clear_thinking"`
|
||||
}
|
||||
|
||||
// Model represents a single LLM model
|
||||
@@ -226,37 +226,8 @@ func NewProviderManager(dirPath string) (*ProviderManager, error) {
|
||||
return nil, fmt.Errorf("error parsing JSON from file %s: %w", filePath, err)
|
||||
}
|
||||
|
||||
// Get support thinking models
|
||||
modelSupportThinking := make(map[string]bool)
|
||||
if provider.Features.Thinking != nil {
|
||||
for _, modelName := range provider.Features.Thinking.SupportedModels {
|
||||
modelSupportThinking[modelName] = true
|
||||
}
|
||||
}
|
||||
|
||||
modelClearThinking := make(map[string]bool)
|
||||
if provider.Features.ClearThinking != nil {
|
||||
for _, modelName := range provider.Features.ClearThinking.SupportedModels {
|
||||
modelClearThinking[modelName] = true
|
||||
}
|
||||
}
|
||||
|
||||
for _, model := range provider.Models {
|
||||
// if the prefix of mode.Name is matched with keys of modelSupportThinking
|
||||
for modelPrefix, _ := range modelSupportThinking {
|
||||
if strings.HasPrefix(model.Name, modelPrefix) {
|
||||
model.Thinking = &ModelThinking{
|
||||
DefaultValue: provider.Features.Thinking.DefaultValue,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for modelPrefix, _ := range modelClearThinking {
|
||||
if strings.HasPrefix(model.Name, modelPrefix) {
|
||||
model.Thinking.ClearContent = true
|
||||
}
|
||||
}
|
||||
|
||||
if provider.Type == "" {
|
||||
pos := strings.Index(model.Name, "-")
|
||||
modelType := model.Name[0:pos]
|
||||
@@ -553,7 +524,7 @@ func ConvertToFeaturesMap(model *Model) map[string]interface{} {
|
||||
if model.Thinking != nil {
|
||||
thinkingMap := map[string]interface{}{
|
||||
"default_value": model.Thinking.DefaultValue,
|
||||
"clear_reasoning": model.Thinking.ClearContent,
|
||||
"clear_reasoning": model.Thinking.ClearThinking,
|
||||
}
|
||||
featuresMap["thinking"] = thinkingMap
|
||||
}
|
||||
|
||||
@@ -45,6 +45,8 @@ func (f *ModelFactory) CreateModelDriver(providerName string, baseURL map[string
|
||||
return NewGiteeModel(baseURL, urlSuffix), nil
|
||||
case "siliconflow":
|
||||
return NewSiliconflowModel(baseURL, urlSuffix), nil
|
||||
case "google":
|
||||
return NewGoogleModel(baseURL, urlSuffix), nil
|
||||
case "aliyun":
|
||||
return NewAliyunModel(baseURL, urlSuffix), nil
|
||||
default:
|
||||
|
||||
173
internal/entity/models/google.go
Normal file
173
internal/entity/models/google.go
Normal file
@@ -0,0 +1,173 @@
|
||||
//
|
||||
// Copyright 2026 The InfiniFlow Authors. All Rights Reserved.
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
//
|
||||
|
||||
package models
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"ragflow/internal/logger"
|
||||
|
||||
"google.golang.org/genai"
|
||||
)
|
||||
|
||||
// GoogleModel implements ModelDriver for Dummy AI
|
||||
type GoogleModel struct {
|
||||
BaseURL map[string]string
|
||||
URLSuffix URLSuffix
|
||||
}
|
||||
|
||||
// NewGoogleModel creates a new Google AI model instance
|
||||
func NewGoogleModel(baseURL map[string]string, urlSuffix URLSuffix) *GoogleModel {
|
||||
return &GoogleModel{
|
||||
BaseURL: baseURL,
|
||||
URLSuffix: urlSuffix,
|
||||
}
|
||||
}
|
||||
|
||||
func (z *GoogleModel) Name() string {
|
||||
return "google"
|
||||
}
|
||||
|
||||
// Chat sends a message and returns response
|
||||
func (z *GoogleModel) Chat(modelName, message *string, apiConfig *APIConfig, chatModelConfig *ChatConfig) (*ChatResponse, error) {
|
||||
ctx := context.Background()
|
||||
client, err := genai.NewClient(ctx, &genai.ClientConfig{
|
||||
APIKey: *apiConfig.ApiKey,
|
||||
Backend: genai.BackendGeminiAPI,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
contents := []*genai.Content{
|
||||
genai.NewContentFromText(*message, genai.RoleUser),
|
||||
}
|
||||
|
||||
generateContentConfig := &genai.GenerateContentConfig{}
|
||||
generateContentConfig.ThinkingConfig = &genai.ThinkingConfig{}
|
||||
if chatModelConfig.Thinking != nil && *chatModelConfig.Thinking {
|
||||
generateContentConfig.ThinkingConfig.IncludeThoughts = true
|
||||
} else {
|
||||
generateContentConfig.ThinkingConfig.IncludeThoughts = false
|
||||
}
|
||||
|
||||
response, err := client.Models.GenerateContent(ctx, *modelName, contents, generateContentConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
content := response.Text()
|
||||
|
||||
var responseContent string
|
||||
if chatModelConfig.Thinking != nil && *chatModelConfig.Thinking {
|
||||
responseContent = response.Candidates[0].Content.Parts[0].Text
|
||||
}
|
||||
|
||||
chatResponse := &ChatResponse{
|
||||
Answer: &content,
|
||||
ReasonContent: &responseContent,
|
||||
}
|
||||
return chatResponse, nil
|
||||
}
|
||||
|
||||
// ChatWithMessages sends multiple messages with roles and returns response
|
||||
func (z *GoogleModel) ChatWithMessages(modelName string, apiKey *string, messages []Message, modelConfig *ChatConfig) (string, error) {
|
||||
return "", fmt.Errorf("not implemented")
|
||||
}
|
||||
|
||||
// ChatStreamlyWithSender sends a message and streams response via sender function (best performance, no channel)
|
||||
func (z *GoogleModel) ChatStreamlyWithSender(modelName, message *string, apiConfig *APIConfig, chatModelConfig *ChatConfig, sender func(*string, *string) error) error {
|
||||
ctx := context.Background()
|
||||
client, err := genai.NewClient(ctx, &genai.ClientConfig{
|
||||
APIKey: *apiConfig.ApiKey,
|
||||
Backend: genai.BackendGeminiAPI,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
contents := []*genai.Content{
|
||||
genai.NewContentFromText(*message, genai.RoleUser),
|
||||
}
|
||||
for response, err := range client.Models.GenerateContentStream(
|
||||
ctx,
|
||||
*modelName,
|
||||
contents,
|
||||
nil,
|
||||
) {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
content := response.Text()
|
||||
|
||||
var responseContent string
|
||||
if chatModelConfig.Thinking != nil && *chatModelConfig.Thinking {
|
||||
responseContent = response.Candidates[0].Content.Parts[0].Text
|
||||
}
|
||||
|
||||
if responseContent != "" {
|
||||
logger.Info(fmt.Sprintf("Thinking: %s", responseContent))
|
||||
if err = sender(nil, &responseContent); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if content != "" {
|
||||
logger.Info(fmt.Sprintf("Answer: %s", responseContent))
|
||||
if err = sender(&content, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// EncodeToEmbedding encodes a list of texts into embeddings
|
||||
func (z *GoogleModel) EncodeToEmbedding(modelName *string, texts []string, apiConfig *APIConfig, embeddingConfig *EmbeddingConfig) ([][]float64, error) {
|
||||
return nil, fmt.Errorf("not implemented")
|
||||
}
|
||||
|
||||
func (z *GoogleModel) ListModels(apiConfig *APIConfig) ([]string, error) {
|
||||
ctx := context.Background()
|
||||
client, err := genai.NewClient(ctx, &genai.ClientConfig{
|
||||
APIKey: *apiConfig.ApiKey,
|
||||
Backend: genai.BackendGeminiAPI,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Retrieve the list of models.
|
||||
models, err := client.Models.List(ctx, &genai.ListModelsConfig{})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var modelNames []string
|
||||
for _, m := range models.Items {
|
||||
modelNames = append(modelNames, m.Name)
|
||||
}
|
||||
return modelNames, nil
|
||||
}
|
||||
|
||||
func (z *GoogleModel) Balance(apiConfig *APIConfig) (map[string]interface{}, error) {
|
||||
return nil, fmt.Errorf("no such method")
|
||||
}
|
||||
|
||||
func (z *GoogleModel) CheckConnection(apiConfig *APIConfig) error {
|
||||
return fmt.Errorf("no such method")
|
||||
}
|
||||
@@ -208,9 +208,9 @@ func (z *ZhipuAIModel) ChatWithMessages(modelName string, apiKey *string, messag
|
||||
|
||||
// Build request body
|
||||
reqBody := map[string]interface{}{
|
||||
"model": modelName,
|
||||
"messages": apiMessages,
|
||||
"stream": false,
|
||||
"model": modelName,
|
||||
"messages": apiMessages,
|
||||
"stream": false,
|
||||
"temperature": 1,
|
||||
}
|
||||
|
||||
@@ -404,16 +404,16 @@ func (z *ZhipuAIModel) ChatStreamlyWithSender(modelName, message *string, apiCon
|
||||
continue
|
||||
}
|
||||
|
||||
content, ok := delta["content"].(string)
|
||||
if ok && content != "" {
|
||||
if err := sender(&content, nil); err != nil {
|
||||
reasoningContent, ok := delta["reasoning_content"].(string)
|
||||
if ok && reasoningContent != "" {
|
||||
if err := sender(nil, &reasoningContent); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
reasoningContent, ok := delta["reasoning_content"].(string)
|
||||
if ok && reasoningContent != "" {
|
||||
if err := sender(nil, &reasoningContent); err != nil {
|
||||
content, ok := delta["content"].(string)
|
||||
if ok && content != "" {
|
||||
if err := sender(&content, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -737,10 +737,10 @@ func (h *ProviderHandler) ChatToModel(c *gin.Context) {
|
||||
}
|
||||
|
||||
// Stream response using sender function (best performance, no channel)
|
||||
errorCode := h.modelProviderService.ChatToModelStreamWithSender(providerName, instanceName, req.ModelName, userID, req.Message, &apiConfig, &chatConfig, sender)
|
||||
errorCode, err := h.modelProviderService.ChatToModelStreamWithSender(providerName, instanceName, req.ModelName, userID, req.Message, &apiConfig, &chatConfig, sender)
|
||||
|
||||
if errorCode != common.CodeSuccess {
|
||||
c.SSEvent("error", "stream failed")
|
||||
c.SSEvent("error", err.Error())
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -844,15 +844,15 @@ func (m *ModelProviderService) ChatWithMessagesToModelByApiKey(providerName, mod
|
||||
}
|
||||
|
||||
// ChatToModelStreamWithSender streams chat response directly via sender function (best performance, no channel)
|
||||
func (m *ModelProviderService) ChatToModelStreamWithSender(providerName, instanceName, modelName, userID, message string, apiConfig *modelModule.APIConfig, modelConfig *modelModule.ChatConfig, sender func(*string, *string) error) common.ErrorCode {
|
||||
func (m *ModelProviderService) ChatToModelStreamWithSender(providerName, instanceName, modelName, userID, message string, apiConfig *modelModule.APIConfig, modelConfig *modelModule.ChatConfig, sender func(*string, *string) error) (common.ErrorCode, error) {
|
||||
// Get tenant ID from user
|
||||
tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner")
|
||||
if err != nil {
|
||||
return common.CodeServerError
|
||||
return common.CodeServerError, err
|
||||
}
|
||||
|
||||
if len(tenants) == 0 {
|
||||
return common.CodeNotFound
|
||||
return common.CodeNotFound, errors.New("user has no tenants")
|
||||
}
|
||||
|
||||
tenantID := tenants[0].TenantID
|
||||
@@ -860,30 +860,30 @@ func (m *ModelProviderService) ChatToModelStreamWithSender(providerName, instanc
|
||||
// Check if provider exists
|
||||
provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName)
|
||||
if err != nil {
|
||||
return common.CodeServerError
|
||||
return common.CodeServerError, err
|
||||
}
|
||||
|
||||
instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName)
|
||||
if err != nil {
|
||||
return common.CodeServerError
|
||||
return common.CodeServerError, err
|
||||
}
|
||||
|
||||
_, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, modelName)
|
||||
if err != nil {
|
||||
providerInfo := dao.GetModelProviderManager().FindProvider(providerName)
|
||||
if providerInfo == nil {
|
||||
return common.CodeNotFound
|
||||
return common.CodeNotFound, err
|
||||
}
|
||||
|
||||
_, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return common.CodeNotFound
|
||||
return common.CodeNotFound, err
|
||||
}
|
||||
|
||||
var extra map[string]string
|
||||
err = json.Unmarshal([]byte(instance.Extra), &extra)
|
||||
if err != nil {
|
||||
return common.CodeServerError
|
||||
return common.CodeServerError, err
|
||||
}
|
||||
|
||||
region := extra["region"]
|
||||
@@ -893,13 +893,13 @@ func (m *ModelProviderService) ChatToModelStreamWithSender(providerName, instanc
|
||||
// Direct call with sender function
|
||||
err = providerInfo.ModelDriver.ChatStreamlyWithSender(&modelName, &message, apiConfig, modelConfig, sender)
|
||||
if err != nil {
|
||||
return common.CodeServerError
|
||||
return common.CodeServerError, err
|
||||
}
|
||||
|
||||
return common.CodeSuccess
|
||||
return common.CodeSuccess, nil
|
||||
}
|
||||
|
||||
return common.CodeServerError
|
||||
return common.CodeServerError, errors.New("model is disabled")
|
||||
}
|
||||
|
||||
func (m *ModelProviderService) GetDefaultModel(modelType entity.ModelType, tenantID string) (*entity.ModelCredentials, error) {
|
||||
|
||||
Reference in New Issue
Block a user