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:
Jin Hai
2026-04-27 20:35:47 +08:00
committed by GitHub
parent 343bda1119
commit 965717c4fb
15 changed files with 456 additions and 181 deletions

View File

@@ -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
}

View File

@@ -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:

View 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")
}

View File

@@ -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
}
}

View File

@@ -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
}

View File

@@ -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) {