mirror of
				https://github.com/songquanpeng/one-api.git
				synced 2025-10-31 22:03:41 +08:00 
			
		
		
		
	Compare commits
	
		
			2 Commits
		
	
	
		
			v0.6.6-alp
			...
			dev
		
	
	| Author | SHA1 | Date | |
|---|---|---|---|
|  | 42569c83c0 | ||
|  | b373882814 | 
							
								
								
									
										127
									
								
								common/config/config.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										127
									
								
								common/config/config.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,127 @@ | |||||||
|  | package config | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"os" | ||||||
|  | 	"strconv" | ||||||
|  | 	"sync" | ||||||
|  | 	"time" | ||||||
|  |  | ||||||
|  | 	"github.com/google/uuid" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | var SystemName = "One API" | ||||||
|  | var ServerAddress = "http://localhost:3000" | ||||||
|  | var Footer = "" | ||||||
|  | var Logo = "" | ||||||
|  | var TopUpLink = "" | ||||||
|  | var ChatLink = "" | ||||||
|  | var QuotaPerUnit = 500 * 1000.0 // $0.002 / 1K tokens | ||||||
|  | var DisplayInCurrencyEnabled = true | ||||||
|  | var DisplayTokenStatEnabled = true | ||||||
|  |  | ||||||
|  | // Any options with "Secret", "Token" in its key won't be return by GetOptions | ||||||
|  |  | ||||||
|  | var SessionSecret = uuid.New().String() | ||||||
|  |  | ||||||
|  | var OptionMap map[string]string | ||||||
|  | var OptionMapRWMutex sync.RWMutex | ||||||
|  |  | ||||||
|  | var ItemsPerPage = 10 | ||||||
|  | var MaxRecentItems = 100 | ||||||
|  |  | ||||||
|  | var PasswordLoginEnabled = true | ||||||
|  | var PasswordRegisterEnabled = true | ||||||
|  | var EmailVerificationEnabled = false | ||||||
|  | var GitHubOAuthEnabled = false | ||||||
|  | var WeChatAuthEnabled = false | ||||||
|  | var TurnstileCheckEnabled = false | ||||||
|  | var RegisterEnabled = true | ||||||
|  |  | ||||||
|  | var EmailDomainRestrictionEnabled = false | ||||||
|  | var EmailDomainWhitelist = []string{ | ||||||
|  | 	"gmail.com", | ||||||
|  | 	"163.com", | ||||||
|  | 	"126.com", | ||||||
|  | 	"qq.com", | ||||||
|  | 	"outlook.com", | ||||||
|  | 	"hotmail.com", | ||||||
|  | 	"icloud.com", | ||||||
|  | 	"yahoo.com", | ||||||
|  | 	"foxmail.com", | ||||||
|  | } | ||||||
|  |  | ||||||
|  | var DebugEnabled = os.Getenv("DEBUG") == "true" | ||||||
|  | var MemoryCacheEnabled = os.Getenv("MEMORY_CACHE_ENABLED") == "true" | ||||||
|  |  | ||||||
|  | var LogConsumeEnabled = true | ||||||
|  |  | ||||||
|  | var SMTPServer = "" | ||||||
|  | var SMTPPort = 587 | ||||||
|  | var SMTPAccount = "" | ||||||
|  | var SMTPFrom = "" | ||||||
|  | var SMTPToken = "" | ||||||
|  |  | ||||||
|  | var GitHubClientId = "" | ||||||
|  | var GitHubClientSecret = "" | ||||||
|  |  | ||||||
|  | var WeChatServerAddress = "" | ||||||
|  | var WeChatServerToken = "" | ||||||
|  | var WeChatAccountQRCodeImageURL = "" | ||||||
|  |  | ||||||
|  | var TurnstileSiteKey = "" | ||||||
|  | var TurnstileSecretKey = "" | ||||||
|  |  | ||||||
|  | var QuotaForNewUser = 0 | ||||||
|  | var QuotaForInviter = 0 | ||||||
|  | var QuotaForInvitee = 0 | ||||||
|  | var ChannelDisableThreshold = 5.0 | ||||||
|  | var AutomaticDisableChannelEnabled = false | ||||||
|  | var AutomaticEnableChannelEnabled = false | ||||||
|  | var QuotaRemindThreshold = 1000 | ||||||
|  | var PreConsumedQuota = 500 | ||||||
|  | var ApproximateTokenEnabled = false | ||||||
|  | var RetryTimes = 0 | ||||||
|  |  | ||||||
|  | var RootUserEmail = "" | ||||||
|  |  | ||||||
|  | var IsMasterNode = os.Getenv("NODE_TYPE") != "slave" | ||||||
|  |  | ||||||
|  | var requestInterval, _ = strconv.Atoi(os.Getenv("POLLING_INTERVAL")) | ||||||
|  | var RequestInterval = time.Duration(requestInterval) * time.Second | ||||||
|  |  | ||||||
|  | var SyncFrequency = helper.GetOrDefaultEnvInt("SYNC_FREQUENCY", 10*60) // unit is second | ||||||
|  |  | ||||||
|  | var BatchUpdateEnabled = false | ||||||
|  | var BatchUpdateInterval = helper.GetOrDefaultEnvInt("BATCH_UPDATE_INTERVAL", 5) | ||||||
|  |  | ||||||
|  | var RelayTimeout = helper.GetOrDefaultEnvInt("RELAY_TIMEOUT", 0) // unit is second | ||||||
|  |  | ||||||
|  | var GeminiSafetySetting = helper.GetOrDefaultEnvString("GEMINI_SAFETY_SETTING", "BLOCK_NONE") | ||||||
|  |  | ||||||
|  | var Theme = helper.GetOrDefaultEnvString("THEME", "default") | ||||||
|  | var ValidThemes = map[string]bool{ | ||||||
|  | 	"default": true, | ||||||
|  | 	"berry":   true, | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // All duration's unit is seconds | ||||||
|  | // Shouldn't larger then RateLimitKeyExpirationDuration | ||||||
|  | var ( | ||||||
|  | 	GlobalApiRateLimitNum            = helper.GetOrDefaultEnvInt("GLOBAL_API_RATE_LIMIT", 180) | ||||||
|  | 	GlobalApiRateLimitDuration int64 = 3 * 60 | ||||||
|  |  | ||||||
|  | 	GlobalWebRateLimitNum            = helper.GetOrDefaultEnvInt("GLOBAL_WEB_RATE_LIMIT", 60) | ||||||
|  | 	GlobalWebRateLimitDuration int64 = 3 * 60 | ||||||
|  |  | ||||||
|  | 	UploadRateLimitNum            = 10 | ||||||
|  | 	UploadRateLimitDuration int64 = 60 | ||||||
|  |  | ||||||
|  | 	DownloadRateLimitNum            = 10 | ||||||
|  | 	DownloadRateLimitDuration int64 = 60 | ||||||
|  |  | ||||||
|  | 	CriticalRateLimitNum            = 20 | ||||||
|  | 	CriticalRateLimitDuration int64 = 20 * 60 | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | var RateLimitKeyExpirationDuration = 20 * time.Minute | ||||||
| @@ -1,114 +1,9 @@ | |||||||
| package common | package common | ||||||
|  |  | ||||||
| import ( | import "time" | ||||||
| 	"os" |  | ||||||
| 	"strconv" |  | ||||||
| 	"sync" |  | ||||||
| 	"time" |  | ||||||
|  |  | ||||||
| 	"github.com/google/uuid" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| var StartTime = time.Now().Unix() // unit: second | var StartTime = time.Now().Unix() // unit: second | ||||||
| var Version = "v0.0.0"            // this hard coding will be replaced automatically when building, no need to manually change | var Version = "v0.0.0"            // this hard coding will be replaced automatically when building, no need to manually change | ||||||
| var SystemName = "One API" |  | ||||||
| var ServerAddress = "http://localhost:3000" |  | ||||||
| var Footer = "" |  | ||||||
| var Logo = "" |  | ||||||
| var TopUpLink = "" |  | ||||||
| var ChatLink = "" |  | ||||||
| var QuotaPerUnit = 500 * 1000.0 // $0.002 / 1K tokens |  | ||||||
| var DisplayInCurrencyEnabled = true |  | ||||||
| var DisplayTokenStatEnabled = true |  | ||||||
|  |  | ||||||
| // Any options with "Secret", "Token" in its key won't be return by GetOptions |  | ||||||
|  |  | ||||||
| var SessionSecret = uuid.New().String() |  | ||||||
|  |  | ||||||
| var OptionMap map[string]string |  | ||||||
| var OptionMapRWMutex sync.RWMutex |  | ||||||
|  |  | ||||||
| var ItemsPerPage = 10 |  | ||||||
| var MaxRecentItems = 100 |  | ||||||
|  |  | ||||||
| var PasswordLoginEnabled = true |  | ||||||
| var PasswordRegisterEnabled = true |  | ||||||
| var EmailVerificationEnabled = false |  | ||||||
| var GitHubOAuthEnabled = false |  | ||||||
| var WeChatAuthEnabled = false |  | ||||||
| var TurnstileCheckEnabled = false |  | ||||||
| var RegisterEnabled = true |  | ||||||
|  |  | ||||||
| var EmailDomainRestrictionEnabled = false |  | ||||||
| var EmailDomainWhitelist = []string{ |  | ||||||
| 	"gmail.com", |  | ||||||
| 	"163.com", |  | ||||||
| 	"126.com", |  | ||||||
| 	"qq.com", |  | ||||||
| 	"outlook.com", |  | ||||||
| 	"hotmail.com", |  | ||||||
| 	"icloud.com", |  | ||||||
| 	"yahoo.com", |  | ||||||
| 	"foxmail.com", |  | ||||||
| } |  | ||||||
|  |  | ||||||
| var DebugEnabled = os.Getenv("DEBUG") == "true" |  | ||||||
| var MemoryCacheEnabled = os.Getenv("MEMORY_CACHE_ENABLED") == "true" |  | ||||||
|  |  | ||||||
| var LogConsumeEnabled = true |  | ||||||
|  |  | ||||||
| var SMTPServer = "" |  | ||||||
| var SMTPPort = 587 |  | ||||||
| var SMTPAccount = "" |  | ||||||
| var SMTPFrom = "" |  | ||||||
| var SMTPToken = "" |  | ||||||
|  |  | ||||||
| var GitHubClientId = "" |  | ||||||
| var GitHubClientSecret = "" |  | ||||||
|  |  | ||||||
| var WeChatServerAddress = "" |  | ||||||
| var WeChatServerToken = "" |  | ||||||
| var WeChatAccountQRCodeImageURL = "" |  | ||||||
|  |  | ||||||
| var TurnstileSiteKey = "" |  | ||||||
| var TurnstileSecretKey = "" |  | ||||||
|  |  | ||||||
| var QuotaForNewUser = 0 |  | ||||||
| var QuotaForInviter = 0 |  | ||||||
| var QuotaForInvitee = 0 |  | ||||||
| var ChannelDisableThreshold = 5.0 |  | ||||||
| var AutomaticDisableChannelEnabled = false |  | ||||||
| var AutomaticEnableChannelEnabled = false |  | ||||||
| var QuotaRemindThreshold = 1000 |  | ||||||
| var PreConsumedQuota = 500 |  | ||||||
| var ApproximateTokenEnabled = false |  | ||||||
| var RetryTimes = 0 |  | ||||||
|  |  | ||||||
| var RootUserEmail = "" |  | ||||||
|  |  | ||||||
| var IsMasterNode = os.Getenv("NODE_TYPE") != "slave" |  | ||||||
|  |  | ||||||
| var requestInterval, _ = strconv.Atoi(os.Getenv("POLLING_INTERVAL")) |  | ||||||
| var RequestInterval = time.Duration(requestInterval) * time.Second |  | ||||||
|  |  | ||||||
| var SyncFrequency = GetOrDefault("SYNC_FREQUENCY", 10*60) // unit is second |  | ||||||
|  |  | ||||||
| var BatchUpdateEnabled = false |  | ||||||
| var BatchUpdateInterval = GetOrDefault("BATCH_UPDATE_INTERVAL", 5) |  | ||||||
|  |  | ||||||
| var RelayTimeout = GetOrDefault("RELAY_TIMEOUT", 0) // unit is second |  | ||||||
|  |  | ||||||
| var GeminiSafetySetting = GetOrDefaultString("GEMINI_SAFETY_SETTING", "BLOCK_NONE") |  | ||||||
|  |  | ||||||
| var Theme = GetOrDefaultString("THEME", "default") |  | ||||||
| var ValidThemes = map[string]bool{ |  | ||||||
| 	"default": true, |  | ||||||
| 	"berry":   true, |  | ||||||
| } |  | ||||||
|  |  | ||||||
| const ( |  | ||||||
| 	RequestIdKey = "X-Oneapi-Request-Id" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| const ( | const ( | ||||||
| 	RoleGuestUser  = 0 | 	RoleGuestUser  = 0 | ||||||
| @@ -117,34 +12,6 @@ const ( | |||||||
| 	RoleRootUser   = 100 | 	RoleRootUser   = 100 | ||||||
| ) | ) | ||||||
|  |  | ||||||
| var ( |  | ||||||
| 	FileUploadPermission    = RoleGuestUser |  | ||||||
| 	FileDownloadPermission  = RoleGuestUser |  | ||||||
| 	ImageUploadPermission   = RoleGuestUser |  | ||||||
| 	ImageDownloadPermission = RoleGuestUser |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| // All duration's unit is seconds |  | ||||||
| // Shouldn't larger then RateLimitKeyExpirationDuration |  | ||||||
| var ( |  | ||||||
| 	GlobalApiRateLimitNum            = GetOrDefault("GLOBAL_API_RATE_LIMIT", 180) |  | ||||||
| 	GlobalApiRateLimitDuration int64 = 3 * 60 |  | ||||||
|  |  | ||||||
| 	GlobalWebRateLimitNum            = GetOrDefault("GLOBAL_WEB_RATE_LIMIT", 60) |  | ||||||
| 	GlobalWebRateLimitDuration int64 = 3 * 60 |  | ||||||
|  |  | ||||||
| 	UploadRateLimitNum            = 10 |  | ||||||
| 	UploadRateLimitDuration int64 = 60 |  | ||||||
|  |  | ||||||
| 	DownloadRateLimitNum            = 10 |  | ||||||
| 	DownloadRateLimitDuration int64 = 60 |  | ||||||
|  |  | ||||||
| 	CriticalRateLimitNum            = 20 |  | ||||||
| 	CriticalRateLimitDuration int64 = 20 * 60 |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| var RateLimitKeyExpirationDuration = 20 * time.Minute |  | ||||||
|  |  | ||||||
| const ( | const ( | ||||||
| 	UserStatusEnabled  = 1 // don't use 0, 0 is the default value! | 	UserStatusEnabled  = 1 // don't use 0, 0 is the default value! | ||||||
| 	UserStatusDisabled = 2 // also don't use 0 | 	UserStatusDisabled = 2 // also don't use 0 | ||||||
| @@ -199,29 +66,29 @@ const ( | |||||||
| ) | ) | ||||||
|  |  | ||||||
| var ChannelBaseURLs = []string{ | var ChannelBaseURLs = []string{ | ||||||
| 	"",                                  // 0 | 	"",                              // 0 | ||||||
| 	"https://api.openai.com",            // 1 | 	"https://api.openai.com",        // 1 | ||||||
| 	"https://oa.api2d.net",              // 2 | 	"https://oa.api2d.net",          // 2 | ||||||
| 	"",                                  // 3 | 	"",                              // 3 | ||||||
| 	"https://api.closeai-proxy.xyz",     // 4 | 	"https://api.closeai-proxy.xyz", // 4 | ||||||
| 	"https://api.openai-sb.com",         // 5 | 	"https://api.openai-sb.com",     // 5 | ||||||
| 	"https://api.openaimax.com",         // 6 | 	"https://api.openaimax.com",     // 6 | ||||||
| 	"https://api.ohmygpt.com",           // 7 | 	"https://api.ohmygpt.com",       // 7 | ||||||
| 	"",                                  // 8 | 	"",                              // 8 | ||||||
| 	"https://api.caipacity.com",         // 9 | 	"https://api.caipacity.com",     // 9 | ||||||
| 	"https://api.aiproxy.io",            // 10 | 	"https://api.aiproxy.io",        // 10 | ||||||
| 	"",                                  // 11 | 	"https://generativelanguage.googleapis.com", // 11 | ||||||
| 	"https://api.api2gpt.com",           // 12 | 	"https://api.api2gpt.com",                   // 12 | ||||||
| 	"https://api.aigc2d.com",            // 13 | 	"https://api.aigc2d.com",                    // 13 | ||||||
| 	"https://api.anthropic.com",         // 14 | 	"https://api.anthropic.com",                 // 14 | ||||||
| 	"https://aip.baidubce.com",          // 15 | 	"https://aip.baidubce.com",                  // 15 | ||||||
| 	"https://open.bigmodel.cn",          // 16 | 	"https://open.bigmodel.cn",                  // 16 | ||||||
| 	"https://dashscope.aliyuncs.com",    // 17 | 	"https://dashscope.aliyuncs.com",            // 17 | ||||||
| 	"",                                  // 18 | 	"",                                          // 18 | ||||||
| 	"https://ai.360.cn",                 // 19 | 	"https://ai.360.cn",                         // 19 | ||||||
| 	"https://openrouter.ai/api",         // 20 | 	"https://openrouter.ai/api",                 // 20 | ||||||
| 	"https://api.aiproxy.io",            // 21 | 	"https://api.aiproxy.io",                    // 21 | ||||||
| 	"https://fastgpt.run/api/openapi",   // 22 | 	"https://fastgpt.run/api/openapi",           // 22 | ||||||
| 	"https://hunyuan.cloud.tencent.com", //23 | 	"https://hunyuan.cloud.tencent.com",         // 23 | ||||||
| 	"",                                  //24 | 	"https://generativelanguage.googleapis.com", // 24 | ||||||
| } | } | ||||||
|   | |||||||
| @@ -1,7 +1,9 @@ | |||||||
| package common | package common | ||||||
|  |  | ||||||
|  | import "one-api/common/helper" | ||||||
|  |  | ||||||
| var UsingSQLite = false | var UsingSQLite = false | ||||||
| var UsingPostgreSQL = false | var UsingPostgreSQL = false | ||||||
|  |  | ||||||
| var SQLitePath = "one-api.db" | var SQLitePath = "one-api.db" | ||||||
| var SQLiteBusyTimeout = GetOrDefault("SQLITE_BUSY_TIMEOUT", 3000) | var SQLiteBusyTimeout = helper.GetOrDefaultEnvInt("SQLITE_BUSY_TIMEOUT", 3000) | ||||||
|   | |||||||
| @@ -6,18 +6,19 @@ import ( | |||||||
| 	"encoding/base64" | 	"encoding/base64" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"net/smtp" | 	"net/smtp" | ||||||
|  | 	"one-api/common/config" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func SendEmail(subject string, receiver string, content string) error { | func SendEmail(subject string, receiver string, content string) error { | ||||||
| 	if SMTPFrom == "" { // for compatibility | 	if config.SMTPFrom == "" { // for compatibility | ||||||
| 		SMTPFrom = SMTPAccount | 		config.SMTPFrom = config.SMTPAccount | ||||||
| 	} | 	} | ||||||
| 	encodedSubject := fmt.Sprintf("=?UTF-8?B?%s?=", base64.StdEncoding.EncodeToString([]byte(subject))) | 	encodedSubject := fmt.Sprintf("=?UTF-8?B?%s?=", base64.StdEncoding.EncodeToString([]byte(subject))) | ||||||
|  |  | ||||||
| 	// Extract domain from SMTPFrom | 	// Extract domain from SMTPFrom | ||||||
| 	parts := strings.Split(SMTPFrom, "@") | 	parts := strings.Split(config.SMTPFrom, "@") | ||||||
| 	var domain string | 	var domain string | ||||||
| 	if len(parts) > 1 { | 	if len(parts) > 1 { | ||||||
| 		domain = parts[1] | 		domain = parts[1] | ||||||
| @@ -36,21 +37,21 @@ func SendEmail(subject string, receiver string, content string) error { | |||||||
| 		"Message-ID: %s\r\n"+ // add Message-ID header to avoid being treated as spam, RFC 5322 | 		"Message-ID: %s\r\n"+ // add Message-ID header to avoid being treated as spam, RFC 5322 | ||||||
| 		"Date: %s\r\n"+ | 		"Date: %s\r\n"+ | ||||||
| 		"Content-Type: text/html; charset=UTF-8\r\n\r\n%s\r\n", | 		"Content-Type: text/html; charset=UTF-8\r\n\r\n%s\r\n", | ||||||
| 		receiver, SystemName, SMTPFrom, encodedSubject, messageId, time.Now().Format(time.RFC1123Z), content)) | 		receiver, config.SystemName, config.SMTPFrom, encodedSubject, messageId, time.Now().Format(time.RFC1123Z), content)) | ||||||
| 	auth := smtp.PlainAuth("", SMTPAccount, SMTPToken, SMTPServer) | 	auth := smtp.PlainAuth("", config.SMTPAccount, config.SMTPToken, config.SMTPServer) | ||||||
| 	addr := fmt.Sprintf("%s:%d", SMTPServer, SMTPPort) | 	addr := fmt.Sprintf("%s:%d", config.SMTPServer, config.SMTPPort) | ||||||
| 	to := strings.Split(receiver, ";") | 	to := strings.Split(receiver, ";") | ||||||
|  |  | ||||||
| 	if SMTPPort == 465 { | 	if config.SMTPPort == 465 { | ||||||
| 		tlsConfig := &tls.Config{ | 		tlsConfig := &tls.Config{ | ||||||
| 			InsecureSkipVerify: true, | 			InsecureSkipVerify: true, | ||||||
| 			ServerName:         SMTPServer, | 			ServerName:         config.SMTPServer, | ||||||
| 		} | 		} | ||||||
| 		conn, err := tls.Dial("tcp", fmt.Sprintf("%s:%d", SMTPServer, SMTPPort), tlsConfig) | 		conn, err := tls.Dial("tcp", fmt.Sprintf("%s:%d", config.SMTPServer, config.SMTPPort), tlsConfig) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		client, err := smtp.NewClient(conn, SMTPServer) | 		client, err := smtp.NewClient(conn, config.SMTPServer) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| @@ -58,7 +59,7 @@ func SendEmail(subject string, receiver string, content string) error { | |||||||
| 		if err = client.Auth(auth); err != nil { | 		if err = client.Auth(auth); err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		if err = client.Mail(SMTPFrom); err != nil { | 		if err = client.Mail(config.SMTPFrom); err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		receiverEmails := strings.Split(receiver, ";") | 		receiverEmails := strings.Split(receiver, ";") | ||||||
| @@ -80,7 +81,7 @@ func SendEmail(subject string, receiver string, content string) error { | |||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 	} else { | 	} else { | ||||||
| 		err = smtp.SendMail(addr, auth, SMTPAccount, to, mail) | 		err = smtp.SendMail(addr, auth, config.SMTPAccount, to, mail) | ||||||
| 	} | 	} | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
|   | |||||||
| @@ -1,6 +1,9 @@ | |||||||
| package common | package common | ||||||
|  |  | ||||||
| import "encoding/json" | import ( | ||||||
|  | 	"encoding/json" | ||||||
|  | 	"one-api/common/logger" | ||||||
|  | ) | ||||||
|  |  | ||||||
| var GroupRatio = map[string]float64{ | var GroupRatio = map[string]float64{ | ||||||
| 	"default": 1, | 	"default": 1, | ||||||
| @@ -11,7 +14,7 @@ var GroupRatio = map[string]float64{ | |||||||
| func GroupRatio2JSONString() string { | func GroupRatio2JSONString() string { | ||||||
| 	jsonBytes, err := json.Marshal(GroupRatio) | 	jsonBytes, err := json.Marshal(GroupRatio) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		SysError("error marshalling model ratio: " + err.Error()) | 		logger.SysError("error marshalling model ratio: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return string(jsonBytes) | 	return string(jsonBytes) | ||||||
| } | } | ||||||
| @@ -24,7 +27,7 @@ func UpdateGroupRatioByJSONString(jsonStr string) error { | |||||||
| func GetGroupRatio(name string) float64 { | func GetGroupRatio(name string) float64 { | ||||||
| 	ratio, ok := GroupRatio[name] | 	ratio, ok := GroupRatio[name] | ||||||
| 	if !ok { | 	if !ok { | ||||||
| 		SysError("group ratio not found: " + name) | 		logger.SysError("group ratio not found: " + name) | ||||||
| 		return 1 | 		return 1 | ||||||
| 	} | 	} | ||||||
| 	return ratio | 	return ratio | ||||||
|   | |||||||
							
								
								
									
										224
									
								
								common/helper/helper.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										224
									
								
								common/helper/helper.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,224 @@ | |||||||
|  | package helper | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"fmt" | ||||||
|  | 	"github.com/google/uuid" | ||||||
|  | 	"html/template" | ||||||
|  | 	"log" | ||||||
|  | 	"math/rand" | ||||||
|  | 	"net" | ||||||
|  | 	"one-api/common/logger" | ||||||
|  | 	"os" | ||||||
|  | 	"os/exec" | ||||||
|  | 	"runtime" | ||||||
|  | 	"strconv" | ||||||
|  | 	"strings" | ||||||
|  | 	"time" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func OpenBrowser(url string) { | ||||||
|  | 	var err error | ||||||
|  |  | ||||||
|  | 	switch runtime.GOOS { | ||||||
|  | 	case "linux": | ||||||
|  | 		err = exec.Command("xdg-open", url).Start() | ||||||
|  | 	case "windows": | ||||||
|  | 		err = exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start() | ||||||
|  | 	case "darwin": | ||||||
|  | 		err = exec.Command("open", url).Start() | ||||||
|  | 	} | ||||||
|  | 	if err != nil { | ||||||
|  | 		log.Println(err) | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GetIp() (ip string) { | ||||||
|  | 	ips, err := net.InterfaceAddrs() | ||||||
|  | 	if err != nil { | ||||||
|  | 		log.Println(err) | ||||||
|  | 		return ip | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	for _, a := range ips { | ||||||
|  | 		if ipNet, ok := a.(*net.IPNet); ok && !ipNet.IP.IsLoopback() { | ||||||
|  | 			if ipNet.IP.To4() != nil { | ||||||
|  | 				ip = ipNet.IP.String() | ||||||
|  | 				if strings.HasPrefix(ip, "10") { | ||||||
|  | 					return | ||||||
|  | 				} | ||||||
|  | 				if strings.HasPrefix(ip, "172") { | ||||||
|  | 					return | ||||||
|  | 				} | ||||||
|  | 				if strings.HasPrefix(ip, "192.168") { | ||||||
|  | 					return | ||||||
|  | 				} | ||||||
|  | 				ip = "" | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return | ||||||
|  | } | ||||||
|  |  | ||||||
|  | var sizeKB = 1024 | ||||||
|  | var sizeMB = sizeKB * 1024 | ||||||
|  | var sizeGB = sizeMB * 1024 | ||||||
|  |  | ||||||
|  | func Bytes2Size(num int64) string { | ||||||
|  | 	numStr := "" | ||||||
|  | 	unit := "B" | ||||||
|  | 	if num/int64(sizeGB) > 1 { | ||||||
|  | 		numStr = fmt.Sprintf("%.2f", float64(num)/float64(sizeGB)) | ||||||
|  | 		unit = "GB" | ||||||
|  | 	} else if num/int64(sizeMB) > 1 { | ||||||
|  | 		numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeMB))) | ||||||
|  | 		unit = "MB" | ||||||
|  | 	} else if num/int64(sizeKB) > 1 { | ||||||
|  | 		numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeKB))) | ||||||
|  | 		unit = "KB" | ||||||
|  | 	} else { | ||||||
|  | 		numStr = fmt.Sprintf("%d", num) | ||||||
|  | 	} | ||||||
|  | 	return numStr + " " + unit | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func Seconds2Time(num int) (time string) { | ||||||
|  | 	if num/31104000 > 0 { | ||||||
|  | 		time += strconv.Itoa(num/31104000) + " 年 " | ||||||
|  | 		num %= 31104000 | ||||||
|  | 	} | ||||||
|  | 	if num/2592000 > 0 { | ||||||
|  | 		time += strconv.Itoa(num/2592000) + " 个月 " | ||||||
|  | 		num %= 2592000 | ||||||
|  | 	} | ||||||
|  | 	if num/86400 > 0 { | ||||||
|  | 		time += strconv.Itoa(num/86400) + " 天 " | ||||||
|  | 		num %= 86400 | ||||||
|  | 	} | ||||||
|  | 	if num/3600 > 0 { | ||||||
|  | 		time += strconv.Itoa(num/3600) + " 小时 " | ||||||
|  | 		num %= 3600 | ||||||
|  | 	} | ||||||
|  | 	if num/60 > 0 { | ||||||
|  | 		time += strconv.Itoa(num/60) + " 分钟 " | ||||||
|  | 		num %= 60 | ||||||
|  | 	} | ||||||
|  | 	time += strconv.Itoa(num) + " 秒" | ||||||
|  | 	return | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func Interface2String(inter interface{}) string { | ||||||
|  | 	switch inter.(type) { | ||||||
|  | 	case string: | ||||||
|  | 		return inter.(string) | ||||||
|  | 	case int: | ||||||
|  | 		return fmt.Sprintf("%d", inter.(int)) | ||||||
|  | 	case float64: | ||||||
|  | 		return fmt.Sprintf("%f", inter.(float64)) | ||||||
|  | 	} | ||||||
|  | 	return "Not Implemented" | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func UnescapeHTML(x string) interface{} { | ||||||
|  | 	return template.HTML(x) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func IntMax(a int, b int) int { | ||||||
|  | 	if a >= b { | ||||||
|  | 		return a | ||||||
|  | 	} else { | ||||||
|  | 		return b | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GetUUID() string { | ||||||
|  | 	code := uuid.New().String() | ||||||
|  | 	code = strings.Replace(code, "-", "", -1) | ||||||
|  | 	return code | ||||||
|  | } | ||||||
|  |  | ||||||
|  | const keyChars = "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ" | ||||||
|  |  | ||||||
|  | func init() { | ||||||
|  | 	rand.Seed(time.Now().UnixNano()) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GenerateKey() string { | ||||||
|  | 	rand.Seed(time.Now().UnixNano()) | ||||||
|  | 	key := make([]byte, 48) | ||||||
|  | 	for i := 0; i < 16; i++ { | ||||||
|  | 		key[i] = keyChars[rand.Intn(len(keyChars))] | ||||||
|  | 	} | ||||||
|  | 	uuid_ := GetUUID() | ||||||
|  | 	for i := 0; i < 32; i++ { | ||||||
|  | 		c := uuid_[i] | ||||||
|  | 		if i%2 == 0 && c >= 'a' && c <= 'z' { | ||||||
|  | 			c = c - 'a' + 'A' | ||||||
|  | 		} | ||||||
|  | 		key[i+16] = c | ||||||
|  | 	} | ||||||
|  | 	return string(key) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GetRandomString(length int) string { | ||||||
|  | 	rand.Seed(time.Now().UnixNano()) | ||||||
|  | 	key := make([]byte, length) | ||||||
|  | 	for i := 0; i < length; i++ { | ||||||
|  | 		key[i] = keyChars[rand.Intn(len(keyChars))] | ||||||
|  | 	} | ||||||
|  | 	return string(key) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GetTimestamp() int64 { | ||||||
|  | 	return time.Now().Unix() | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GetTimeString() string { | ||||||
|  | 	now := time.Now() | ||||||
|  | 	return fmt.Sprintf("%s%d", now.Format("20060102150405"), now.UnixNano()%1e9) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func Max(a int, b int) int { | ||||||
|  | 	if a >= b { | ||||||
|  | 		return a | ||||||
|  | 	} else { | ||||||
|  | 		return b | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GetOrDefaultEnvInt(env string, defaultValue int) int { | ||||||
|  | 	if env == "" || os.Getenv(env) == "" { | ||||||
|  | 		return defaultValue | ||||||
|  | 	} | ||||||
|  | 	num, err := strconv.Atoi(os.Getenv(env)) | ||||||
|  | 	if err != nil { | ||||||
|  | 		logger.SysError(fmt.Sprintf("failed to parse %s: %s, using default value: %d", env, err.Error(), defaultValue)) | ||||||
|  | 		return defaultValue | ||||||
|  | 	} | ||||||
|  | 	return num | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GetOrDefaultEnvString(env string, defaultValue string) string { | ||||||
|  | 	if env == "" || os.Getenv(env) == "" { | ||||||
|  | 		return defaultValue | ||||||
|  | 	} | ||||||
|  | 	return os.Getenv(env) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func AssignOrDefault(value string, defaultValue string) string { | ||||||
|  | 	if len(value) != 0 { | ||||||
|  | 		return value | ||||||
|  | 	} | ||||||
|  | 	return defaultValue | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func MessageWithRequestId(message string, id string) string { | ||||||
|  | 	return fmt.Sprintf("%s (request id: %s)", message, id) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func String2Int(str string) int { | ||||||
|  | 	num, err := strconv.Atoi(str) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return 0 | ||||||
|  | 	} | ||||||
|  | 	return num | ||||||
|  | } | ||||||
| @@ -4,6 +4,8 @@ import ( | |||||||
| 	"flag" | 	"flag" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"log" | 	"log" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"os" | 	"os" | ||||||
| 	"path/filepath" | 	"path/filepath" | ||||||
| ) | ) | ||||||
| @@ -37,9 +39,9 @@ func init() { | |||||||
|  |  | ||||||
| 	if os.Getenv("SESSION_SECRET") != "" { | 	if os.Getenv("SESSION_SECRET") != "" { | ||||||
| 		if os.Getenv("SESSION_SECRET") == "random_string" { | 		if os.Getenv("SESSION_SECRET") == "random_string" { | ||||||
| 			SysError("SESSION_SECRET is set to an example value, please change it to a random string.") | 			logger.SysError("SESSION_SECRET is set to an example value, please change it to a random string.") | ||||||
| 		} else { | 		} else { | ||||||
| 			SessionSecret = os.Getenv("SESSION_SECRET") | 			config.SessionSecret = os.Getenv("SESSION_SECRET") | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	if os.Getenv("SQLITE_PATH") != "" { | 	if os.Getenv("SQLITE_PATH") != "" { | ||||||
| @@ -57,5 +59,6 @@ func init() { | |||||||
| 				log.Fatal(err) | 				log.Fatal(err) | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
|  | 		logger.LogDir = *LogDir | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
							
								
								
									
										7
									
								
								common/logger/constants.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										7
									
								
								common/logger/constants.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,7 @@ | |||||||
|  | package logger | ||||||
|  |  | ||||||
|  | const ( | ||||||
|  | 	RequestIdKey = "X-Oneapi-Request-Id" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | var LogDir string | ||||||
| @@ -1,4 +1,4 @@ | |||||||
| package common | package logger | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"context" | 	"context" | ||||||
| @@ -25,7 +25,7 @@ var setupLogLock sync.Mutex | |||||||
| var setupLogWorking bool | var setupLogWorking bool | ||||||
| 
 | 
 | ||||||
| func SetupLogger() { | func SetupLogger() { | ||||||
| 	if *LogDir != "" { | 	if LogDir != "" { | ||||||
| 		ok := setupLogLock.TryLock() | 		ok := setupLogLock.TryLock() | ||||||
| 		if !ok { | 		if !ok { | ||||||
| 			log.Println("setup log is already working") | 			log.Println("setup log is already working") | ||||||
| @@ -35,7 +35,7 @@ func SetupLogger() { | |||||||
| 			setupLogLock.Unlock() | 			setupLogLock.Unlock() | ||||||
| 			setupLogWorking = false | 			setupLogWorking = false | ||||||
| 		}() | 		}() | ||||||
| 		logPath := filepath.Join(*LogDir, fmt.Sprintf("oneapi-%s.log", time.Now().Format("20060102"))) | 		logPath := filepath.Join(LogDir, fmt.Sprintf("oneapi-%s.log", time.Now().Format("20060102"))) | ||||||
| 		fd, err := os.OpenFile(logPath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644) | 		fd, err := os.OpenFile(logPath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			log.Fatal("failed to open log file") | 			log.Fatal("failed to open log file") | ||||||
| @@ -55,18 +55,30 @@ func SysError(s string) { | |||||||
| 	_, _ = fmt.Fprintf(gin.DefaultErrorWriter, "[SYS] %v | %s \n", t.Format("2006/01/02 - 15:04:05"), s) | 	_, _ = fmt.Fprintf(gin.DefaultErrorWriter, "[SYS] %v | %s \n", t.Format("2006/01/02 - 15:04:05"), s) | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func LogInfo(ctx context.Context, msg string) { | func Info(ctx context.Context, msg string) { | ||||||
| 	logHelper(ctx, loggerINFO, msg) | 	logHelper(ctx, loggerINFO, msg) | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func LogWarn(ctx context.Context, msg string) { | func Warn(ctx context.Context, msg string) { | ||||||
| 	logHelper(ctx, loggerWarn, msg) | 	logHelper(ctx, loggerWarn, msg) | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func LogError(ctx context.Context, msg string) { | func Error(ctx context.Context, msg string) { | ||||||
| 	logHelper(ctx, loggerError, msg) | 	logHelper(ctx, loggerError, msg) | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
|  | func Infof(ctx context.Context, format string, a ...any) { | ||||||
|  | 	Info(ctx, fmt.Sprintf(format, a)) | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func Warnf(ctx context.Context, format string, a ...any) { | ||||||
|  | 	Warn(ctx, fmt.Sprintf(format, a)) | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func Errorf(ctx context.Context, format string, a ...any) { | ||||||
|  | 	Error(ctx, fmt.Sprintf(format, a)) | ||||||
|  | } | ||||||
|  | 
 | ||||||
| func logHelper(ctx context.Context, level string, msg string) { | func logHelper(ctx context.Context, level string, msg string) { | ||||||
| 	writer := gin.DefaultErrorWriter | 	writer := gin.DefaultErrorWriter | ||||||
| 	if level == loggerINFO { | 	if level == loggerINFO { | ||||||
| @@ -90,11 +102,3 @@ func FatalLog(v ...any) { | |||||||
| 	_, _ = fmt.Fprintf(gin.DefaultErrorWriter, "[FATAL] %v | %v \n", t.Format("2006/01/02 - 15:04:05"), v) | 	_, _ = fmt.Fprintf(gin.DefaultErrorWriter, "[FATAL] %v | %v \n", t.Format("2006/01/02 - 15:04:05"), v) | ||||||
| 	os.Exit(1) | 	os.Exit(1) | ||||||
| } | } | ||||||
| 
 |  | ||||||
| func LogQuota(quota int) string { |  | ||||||
| 	if DisplayInCurrencyEnabled { |  | ||||||
| 		return fmt.Sprintf("$%.6f 额度", float64(quota)/QuotaPerUnit) |  | ||||||
| 	} else { |  | ||||||
| 		return fmt.Sprintf("%d 点额度", quota) |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
| @@ -2,6 +2,7 @@ package common | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
| @@ -107,7 +108,7 @@ var ModelRatio = map[string]float64{ | |||||||
| func ModelRatio2JSONString() string { | func ModelRatio2JSONString() string { | ||||||
| 	jsonBytes, err := json.Marshal(ModelRatio) | 	jsonBytes, err := json.Marshal(ModelRatio) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		SysError("error marshalling model ratio: " + err.Error()) | 		logger.SysError("error marshalling model ratio: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return string(jsonBytes) | 	return string(jsonBytes) | ||||||
| } | } | ||||||
| @@ -123,7 +124,7 @@ func GetModelRatio(name string) float64 { | |||||||
| 	} | 	} | ||||||
| 	ratio, ok := ModelRatio[name] | 	ratio, ok := ModelRatio[name] | ||||||
| 	if !ok { | 	if !ok { | ||||||
| 		SysError("model ratio not found: " + name) | 		logger.SysError("model ratio not found: " + name) | ||||||
| 		return 30 | 		return 30 | ||||||
| 	} | 	} | ||||||
| 	return ratio | 	return ratio | ||||||
|   | |||||||
| @@ -3,6 +3,7 @@ package common | |||||||
| import ( | import ( | ||||||
| 	"context" | 	"context" | ||||||
| 	"github.com/go-redis/redis/v8" | 	"github.com/go-redis/redis/v8" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"os" | 	"os" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
| @@ -14,18 +15,18 @@ var RedisEnabled = true | |||||||
| func InitRedisClient() (err error) { | func InitRedisClient() (err error) { | ||||||
| 	if os.Getenv("REDIS_CONN_STRING") == "" { | 	if os.Getenv("REDIS_CONN_STRING") == "" { | ||||||
| 		RedisEnabled = false | 		RedisEnabled = false | ||||||
| 		SysLog("REDIS_CONN_STRING not set, Redis is not enabled") | 		logger.SysLog("REDIS_CONN_STRING not set, Redis is not enabled") | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| 	if os.Getenv("SYNC_FREQUENCY") == "" { | 	if os.Getenv("SYNC_FREQUENCY") == "" { | ||||||
| 		RedisEnabled = false | 		RedisEnabled = false | ||||||
| 		SysLog("SYNC_FREQUENCY not set, Redis is disabled") | 		logger.SysLog("SYNC_FREQUENCY not set, Redis is disabled") | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| 	SysLog("Redis is enabled") | 	logger.SysLog("Redis is enabled") | ||||||
| 	opt, err := redis.ParseURL(os.Getenv("REDIS_CONN_STRING")) | 	opt, err := redis.ParseURL(os.Getenv("REDIS_CONN_STRING")) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		FatalLog("failed to parse Redis connection string: " + err.Error()) | 		logger.FatalLog("failed to parse Redis connection string: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	RDB = redis.NewClient(opt) | 	RDB = redis.NewClient(opt) | ||||||
|  |  | ||||||
| @@ -34,7 +35,7 @@ func InitRedisClient() (err error) { | |||||||
|  |  | ||||||
| 	_, err = RDB.Ping(ctx).Result() | 	_, err = RDB.Ping(ctx).Result() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		FatalLog("Redis ping test failed: " + err.Error()) | 		logger.FatalLog("Redis ping test failed: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
| @@ -42,7 +43,7 @@ func InitRedisClient() (err error) { | |||||||
| func ParseRedisOption() *redis.Options { | func ParseRedisOption() *redis.Options { | ||||||
| 	opt, err := redis.ParseURL(os.Getenv("REDIS_CONN_STRING")) | 	opt, err := redis.ParseURL(os.Getenv("REDIS_CONN_STRING")) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		FatalLog("failed to parse Redis connection string: " + err.Error()) | 		logger.FatalLog("failed to parse Redis connection string: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return opt | 	return opt | ||||||
| } | } | ||||||
|   | |||||||
							
								
								
									
										212
									
								
								common/utils.go
									
									
									
									
									
								
							
							
						
						
									
										212
									
								
								common/utils.go
									
									
									
									
									
								
							| @@ -2,215 +2,13 @@ package common | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/google/uuid" | 	"one-api/common/config" | ||||||
| 	"html/template" |  | ||||||
| 	"log" |  | ||||||
| 	"math/rand" |  | ||||||
| 	"net" |  | ||||||
| 	"os" |  | ||||||
| 	"os/exec" |  | ||||||
| 	"runtime" |  | ||||||
| 	"strconv" |  | ||||||
| 	"strings" |  | ||||||
| 	"time" |  | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func OpenBrowser(url string) { | func LogQuota(quota int) string { | ||||||
| 	var err error | 	if config.DisplayInCurrencyEnabled { | ||||||
|  | 		return fmt.Sprintf("$%.6f 额度", float64(quota)/config.QuotaPerUnit) | ||||||
| 	switch runtime.GOOS { |  | ||||||
| 	case "linux": |  | ||||||
| 		err = exec.Command("xdg-open", url).Start() |  | ||||||
| 	case "windows": |  | ||||||
| 		err = exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start() |  | ||||||
| 	case "darwin": |  | ||||||
| 		err = exec.Command("open", url).Start() |  | ||||||
| 	} |  | ||||||
| 	if err != nil { |  | ||||||
| 		log.Println(err) |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func GetIp() (ip string) { |  | ||||||
| 	ips, err := net.InterfaceAddrs() |  | ||||||
| 	if err != nil { |  | ||||||
| 		log.Println(err) |  | ||||||
| 		return ip |  | ||||||
| 	} |  | ||||||
|  |  | ||||||
| 	for _, a := range ips { |  | ||||||
| 		if ipNet, ok := a.(*net.IPNet); ok && !ipNet.IP.IsLoopback() { |  | ||||||
| 			if ipNet.IP.To4() != nil { |  | ||||||
| 				ip = ipNet.IP.String() |  | ||||||
| 				if strings.HasPrefix(ip, "10") { |  | ||||||
| 					return |  | ||||||
| 				} |  | ||||||
| 				if strings.HasPrefix(ip, "172") { |  | ||||||
| 					return |  | ||||||
| 				} |  | ||||||
| 				if strings.HasPrefix(ip, "192.168") { |  | ||||||
| 					return |  | ||||||
| 				} |  | ||||||
| 				ip = "" |  | ||||||
| 			} |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	return |  | ||||||
| } |  | ||||||
|  |  | ||||||
| var sizeKB = 1024 |  | ||||||
| var sizeMB = sizeKB * 1024 |  | ||||||
| var sizeGB = sizeMB * 1024 |  | ||||||
|  |  | ||||||
| func Bytes2Size(num int64) string { |  | ||||||
| 	numStr := "" |  | ||||||
| 	unit := "B" |  | ||||||
| 	if num/int64(sizeGB) > 1 { |  | ||||||
| 		numStr = fmt.Sprintf("%.2f", float64(num)/float64(sizeGB)) |  | ||||||
| 		unit = "GB" |  | ||||||
| 	} else if num/int64(sizeMB) > 1 { |  | ||||||
| 		numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeMB))) |  | ||||||
| 		unit = "MB" |  | ||||||
| 	} else if num/int64(sizeKB) > 1 { |  | ||||||
| 		numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeKB))) |  | ||||||
| 		unit = "KB" |  | ||||||
| 	} else { | 	} else { | ||||||
| 		numStr = fmt.Sprintf("%d", num) | 		return fmt.Sprintf("%d 点额度", quota) | ||||||
| 	} |  | ||||||
| 	return numStr + " " + unit |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func Seconds2Time(num int) (time string) { |  | ||||||
| 	if num/31104000 > 0 { |  | ||||||
| 		time += strconv.Itoa(num/31104000) + " 年 " |  | ||||||
| 		num %= 31104000 |  | ||||||
| 	} |  | ||||||
| 	if num/2592000 > 0 { |  | ||||||
| 		time += strconv.Itoa(num/2592000) + " 个月 " |  | ||||||
| 		num %= 2592000 |  | ||||||
| 	} |  | ||||||
| 	if num/86400 > 0 { |  | ||||||
| 		time += strconv.Itoa(num/86400) + " 天 " |  | ||||||
| 		num %= 86400 |  | ||||||
| 	} |  | ||||||
| 	if num/3600 > 0 { |  | ||||||
| 		time += strconv.Itoa(num/3600) + " 小时 " |  | ||||||
| 		num %= 3600 |  | ||||||
| 	} |  | ||||||
| 	if num/60 > 0 { |  | ||||||
| 		time += strconv.Itoa(num/60) + " 分钟 " |  | ||||||
| 		num %= 60 |  | ||||||
| 	} |  | ||||||
| 	time += strconv.Itoa(num) + " 秒" |  | ||||||
| 	return |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func Interface2String(inter interface{}) string { |  | ||||||
| 	switch inter.(type) { |  | ||||||
| 	case string: |  | ||||||
| 		return inter.(string) |  | ||||||
| 	case int: |  | ||||||
| 		return fmt.Sprintf("%d", inter.(int)) |  | ||||||
| 	case float64: |  | ||||||
| 		return fmt.Sprintf("%f", inter.(float64)) |  | ||||||
| 	} |  | ||||||
| 	return "Not Implemented" |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func UnescapeHTML(x string) interface{} { |  | ||||||
| 	return template.HTML(x) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func IntMax(a int, b int) int { |  | ||||||
| 	if a >= b { |  | ||||||
| 		return a |  | ||||||
| 	} else { |  | ||||||
| 		return b |  | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetUUID() string { |  | ||||||
| 	code := uuid.New().String() |  | ||||||
| 	code = strings.Replace(code, "-", "", -1) |  | ||||||
| 	return code |  | ||||||
| } |  | ||||||
|  |  | ||||||
| const keyChars = "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ" |  | ||||||
|  |  | ||||||
| func init() { |  | ||||||
| 	rand.Seed(time.Now().UnixNano()) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func GenerateKey() string { |  | ||||||
| 	rand.Seed(time.Now().UnixNano()) |  | ||||||
| 	key := make([]byte, 48) |  | ||||||
| 	for i := 0; i < 16; i++ { |  | ||||||
| 		key[i] = keyChars[rand.Intn(len(keyChars))] |  | ||||||
| 	} |  | ||||||
| 	uuid_ := GetUUID() |  | ||||||
| 	for i := 0; i < 32; i++ { |  | ||||||
| 		c := uuid_[i] |  | ||||||
| 		if i%2 == 0 && c >= 'a' && c <= 'z' { |  | ||||||
| 			c = c - 'a' + 'A' |  | ||||||
| 		} |  | ||||||
| 		key[i+16] = c |  | ||||||
| 	} |  | ||||||
| 	return string(key) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func GetRandomString(length int) string { |  | ||||||
| 	rand.Seed(time.Now().UnixNano()) |  | ||||||
| 	key := make([]byte, length) |  | ||||||
| 	for i := 0; i < length; i++ { |  | ||||||
| 		key[i] = keyChars[rand.Intn(len(keyChars))] |  | ||||||
| 	} |  | ||||||
| 	return string(key) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func GetTimestamp() int64 { |  | ||||||
| 	return time.Now().Unix() |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func GetTimeString() string { |  | ||||||
| 	now := time.Now() |  | ||||||
| 	return fmt.Sprintf("%s%d", now.Format("20060102150405"), now.UnixNano()%1e9) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func Max(a int, b int) int { |  | ||||||
| 	if a >= b { |  | ||||||
| 		return a |  | ||||||
| 	} else { |  | ||||||
| 		return b |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func GetOrDefault(env string, defaultValue int) int { |  | ||||||
| 	if env == "" || os.Getenv(env) == "" { |  | ||||||
| 		return defaultValue |  | ||||||
| 	} |  | ||||||
| 	num, err := strconv.Atoi(os.Getenv(env)) |  | ||||||
| 	if err != nil { |  | ||||||
| 		SysError(fmt.Sprintf("failed to parse %s: %s, using default value: %d", env, err.Error(), defaultValue)) |  | ||||||
| 		return defaultValue |  | ||||||
| 	} |  | ||||||
| 	return num |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func GetOrDefaultString(env string, defaultValue string) string { |  | ||||||
| 	if env == "" || os.Getenv(env) == "" { |  | ||||||
| 		return defaultValue |  | ||||||
| 	} |  | ||||||
| 	return os.Getenv(env) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func MessageWithRequestId(message string, id string) string { |  | ||||||
| 	return fmt.Sprintf("%s (request id: %s)", message, id) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func String2Int(str string) int { |  | ||||||
| 	num, err := strconv.Atoi(str) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return 0 |  | ||||||
| 	} |  | ||||||
| 	return num |  | ||||||
| } |  | ||||||
|   | |||||||
| @@ -2,7 +2,7 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| ) | ) | ||||||
| @@ -13,7 +13,7 @@ func GetSubscription(c *gin.Context) { | |||||||
| 	var err error | 	var err error | ||||||
| 	var token *model.Token | 	var token *model.Token | ||||||
| 	var expiredTime int64 | 	var expiredTime int64 | ||||||
| 	if common.DisplayTokenStatEnabled { | 	if config.DisplayTokenStatEnabled { | ||||||
| 		tokenId := c.GetInt("token_id") | 		tokenId := c.GetInt("token_id") | ||||||
| 		token, err = model.GetTokenById(tokenId) | 		token, err = model.GetTokenById(tokenId) | ||||||
| 		expiredTime = token.ExpiredTime | 		expiredTime = token.ExpiredTime | ||||||
| @@ -39,8 +39,8 @@ func GetSubscription(c *gin.Context) { | |||||||
| 	} | 	} | ||||||
| 	quota := remainQuota + usedQuota | 	quota := remainQuota + usedQuota | ||||||
| 	amount := float64(quota) | 	amount := float64(quota) | ||||||
| 	if common.DisplayInCurrencyEnabled { | 	if config.DisplayInCurrencyEnabled { | ||||||
| 		amount /= common.QuotaPerUnit | 		amount /= config.QuotaPerUnit | ||||||
| 	} | 	} | ||||||
| 	if token != nil && token.UnlimitedQuota { | 	if token != nil && token.UnlimitedQuota { | ||||||
| 		amount = 100000000 | 		amount = 100000000 | ||||||
| @@ -61,7 +61,7 @@ func GetUsage(c *gin.Context) { | |||||||
| 	var quota int | 	var quota int | ||||||
| 	var err error | 	var err error | ||||||
| 	var token *model.Token | 	var token *model.Token | ||||||
| 	if common.DisplayTokenStatEnabled { | 	if config.DisplayTokenStatEnabled { | ||||||
| 		tokenId := c.GetInt("token_id") | 		tokenId := c.GetInt("token_id") | ||||||
| 		token, err = model.GetTokenById(tokenId) | 		token, err = model.GetTokenById(tokenId) | ||||||
| 		quota = token.UsedQuota | 		quota = token.UsedQuota | ||||||
| @@ -80,8 +80,8 @@ func GetUsage(c *gin.Context) { | |||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	amount := float64(quota) | 	amount := float64(quota) | ||||||
| 	if common.DisplayInCurrencyEnabled { | 	if config.DisplayInCurrencyEnabled { | ||||||
| 		amount /= common.QuotaPerUnit | 		amount /= config.QuotaPerUnit | ||||||
| 	} | 	} | ||||||
| 	usage := OpenAIUsageResponse{ | 	usage := OpenAIUsageResponse{ | ||||||
| 		Object:     "list", | 		Object:     "list", | ||||||
|   | |||||||
| @@ -7,6 +7,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"one-api/relay/util" | 	"one-api/relay/util" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| @@ -314,7 +316,7 @@ func updateAllChannelsBalance() error { | |||||||
| 				disableChannel(channel.Id, channel.Name, "余额不足") | 				disableChannel(channel.Id, channel.Name, "余额不足") | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 		time.Sleep(common.RequestInterval) | 		time.Sleep(config.RequestInterval) | ||||||
| 	} | 	} | ||||||
| 	return nil | 	return nil | ||||||
| } | } | ||||||
| @@ -339,8 +341,8 @@ func UpdateAllChannelsBalance(c *gin.Context) { | |||||||
| func AutomaticallyUpdateChannels(frequency int) { | func AutomaticallyUpdateChannels(frequency int) { | ||||||
| 	for { | 	for { | ||||||
| 		time.Sleep(time.Duration(frequency) * time.Minute) | 		time.Sleep(time.Duration(frequency) * time.Minute) | ||||||
| 		common.SysLog("updating all channels") | 		logger.SysLog("updating all channels") | ||||||
| 		_ = updateAllChannelsBalance() | 		_ = updateAllChannelsBalance() | ||||||
| 		common.SysLog("channels update done") | 		logger.SysLog("channels update done") | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -8,6 +8,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/util" | 	"one-api/relay/util" | ||||||
| @@ -150,12 +152,12 @@ var testAllChannelsLock sync.Mutex | |||||||
| var testAllChannelsRunning bool = false | var testAllChannelsRunning bool = false | ||||||
|  |  | ||||||
| func notifyRootUser(subject string, content string) { | func notifyRootUser(subject string, content string) { | ||||||
| 	if common.RootUserEmail == "" { | 	if config.RootUserEmail == "" { | ||||||
| 		common.RootUserEmail = model.GetRootUserEmail() | 		config.RootUserEmail = model.GetRootUserEmail() | ||||||
| 	} | 	} | ||||||
| 	err := common.SendEmail(subject, common.RootUserEmail, content) | 	err := common.SendEmail(subject, config.RootUserEmail, content) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError(fmt.Sprintf("failed to send email: %s", err.Error())) | 		logger.SysError(fmt.Sprintf("failed to send email: %s", err.Error())) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -176,8 +178,8 @@ func enableChannel(channelId int, channelName string) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func testAllChannels(notify bool) error { | func testAllChannels(notify bool) error { | ||||||
| 	if common.RootUserEmail == "" { | 	if config.RootUserEmail == "" { | ||||||
| 		common.RootUserEmail = model.GetRootUserEmail() | 		config.RootUserEmail = model.GetRootUserEmail() | ||||||
| 	} | 	} | ||||||
| 	testAllChannelsLock.Lock() | 	testAllChannelsLock.Lock() | ||||||
| 	if testAllChannelsRunning { | 	if testAllChannelsRunning { | ||||||
| @@ -191,7 +193,7 @@ func testAllChannels(notify bool) error { | |||||||
| 		return err | 		return err | ||||||
| 	} | 	} | ||||||
| 	testRequest := buildTestRequest() | 	testRequest := buildTestRequest() | ||||||
| 	var disableThreshold = int64(common.ChannelDisableThreshold * 1000) | 	var disableThreshold = int64(config.ChannelDisableThreshold * 1000) | ||||||
| 	if disableThreshold == 0 { | 	if disableThreshold == 0 { | ||||||
| 		disableThreshold = 10000000 // a impossible value | 		disableThreshold = 10000000 // a impossible value | ||||||
| 	} | 	} | ||||||
| @@ -213,15 +215,15 @@ func testAllChannels(notify bool) error { | |||||||
| 				enableChannel(channel.Id, channel.Name) | 				enableChannel(channel.Id, channel.Name) | ||||||
| 			} | 			} | ||||||
| 			channel.UpdateResponseTime(milliseconds) | 			channel.UpdateResponseTime(milliseconds) | ||||||
| 			time.Sleep(common.RequestInterval) | 			time.Sleep(config.RequestInterval) | ||||||
| 		} | 		} | ||||||
| 		testAllChannelsLock.Lock() | 		testAllChannelsLock.Lock() | ||||||
| 		testAllChannelsRunning = false | 		testAllChannelsRunning = false | ||||||
| 		testAllChannelsLock.Unlock() | 		testAllChannelsLock.Unlock() | ||||||
| 		if notify { | 		if notify { | ||||||
| 			err := common.SendEmail("通道测试完成", common.RootUserEmail, "通道测试完成,如果没有收到禁用通知,说明所有通道都正常") | 			err := common.SendEmail("通道测试完成", config.RootUserEmail, "通道测试完成,如果没有收到禁用通知,说明所有通道都正常") | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError(fmt.Sprintf("failed to send email: %s", err.Error())) | 				logger.SysError(fmt.Sprintf("failed to send email: %s", err.Error())) | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 	}() | 	}() | ||||||
| @@ -247,8 +249,8 @@ func TestAllChannels(c *gin.Context) { | |||||||
| func AutomaticallyTestChannels(frequency int) { | func AutomaticallyTestChannels(frequency int) { | ||||||
| 	for { | 	for { | ||||||
| 		time.Sleep(time.Duration(frequency) * time.Minute) | 		time.Sleep(time.Duration(frequency) * time.Minute) | ||||||
| 		common.SysLog("testing all channels") | 		logger.SysLog("testing all channels") | ||||||
| 		_ = testAllChannels(false) | 		_ = testAllChannels(false) | ||||||
| 		common.SysLog("channel test finished") | 		logger.SysLog("channel test finished") | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -3,7 +3,8 @@ package controller | |||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
| @@ -14,7 +15,7 @@ func GetAllChannels(c *gin.Context) { | |||||||
| 	if p < 0 { | 	if p < 0 { | ||||||
| 		p = 0 | 		p = 0 | ||||||
| 	} | 	} | ||||||
| 	channels, err := model.GetAllChannels(p*common.ItemsPerPage, common.ItemsPerPage, false) | 	channels, err := model.GetAllChannels(p*config.ItemsPerPage, config.ItemsPerPage, false) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -83,7 +84,7 @@ func AddChannel(c *gin.Context) { | |||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	channel.CreatedTime = common.GetTimestamp() | 	channel.CreatedTime = helper.GetTimestamp() | ||||||
| 	keys := strings.Split(channel.Key, "\n") | 	keys := strings.Split(channel.Key, "\n") | ||||||
| 	channels := make([]model.Channel, 0, len(keys)) | 	channels := make([]model.Channel, 0, len(keys)) | ||||||
| 	for _, key := range keys { | 	for _, key := range keys { | ||||||
|   | |||||||
| @@ -9,6 +9,9 @@ import ( | |||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"time" | 	"time" | ||||||
| @@ -30,7 +33,7 @@ func getGitHubUserInfoByCode(code string) (*GitHubUser, error) { | |||||||
| 	if code == "" { | 	if code == "" { | ||||||
| 		return nil, errors.New("无效的参数") | 		return nil, errors.New("无效的参数") | ||||||
| 	} | 	} | ||||||
| 	values := map[string]string{"client_id": common.GitHubClientId, "client_secret": common.GitHubClientSecret, "code": code} | 	values := map[string]string{"client_id": config.GitHubClientId, "client_secret": config.GitHubClientSecret, "code": code} | ||||||
| 	jsonData, err := json.Marshal(values) | 	jsonData, err := json.Marshal(values) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return nil, err | 		return nil, err | ||||||
| @@ -46,7 +49,7 @@ func getGitHubUserInfoByCode(code string) (*GitHubUser, error) { | |||||||
| 	} | 	} | ||||||
| 	res, err := client.Do(req) | 	res, err := client.Do(req) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysLog(err.Error()) | 		logger.SysLog(err.Error()) | ||||||
| 		return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") | 		return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") | ||||||
| 	} | 	} | ||||||
| 	defer res.Body.Close() | 	defer res.Body.Close() | ||||||
| @@ -62,7 +65,7 @@ func getGitHubUserInfoByCode(code string) (*GitHubUser, error) { | |||||||
| 	req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", oAuthResponse.AccessToken)) | 	req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", oAuthResponse.AccessToken)) | ||||||
| 	res2, err := client.Do(req) | 	res2, err := client.Do(req) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysLog(err.Error()) | 		logger.SysLog(err.Error()) | ||||||
| 		return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") | 		return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") | ||||||
| 	} | 	} | ||||||
| 	defer res2.Body.Close() | 	defer res2.Body.Close() | ||||||
| @@ -93,7 +96,7 @@ func GitHubOAuth(c *gin.Context) { | |||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	if !common.GitHubOAuthEnabled { | 	if !config.GitHubOAuthEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| 			"message": "管理员未开启通过 GitHub 登录以及注册", | 			"message": "管理员未开启通过 GitHub 登录以及注册", | ||||||
| @@ -122,7 +125,7 @@ func GitHubOAuth(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	} else { | 	} else { | ||||||
| 		if common.RegisterEnabled { | 		if config.RegisterEnabled { | ||||||
| 			user.Username = "github_" + strconv.Itoa(model.GetMaxUserId()+1) | 			user.Username = "github_" + strconv.Itoa(model.GetMaxUserId()+1) | ||||||
| 			if githubUser.Name != "" { | 			if githubUser.Name != "" { | ||||||
| 				user.DisplayName = githubUser.Name | 				user.DisplayName = githubUser.Name | ||||||
| @@ -160,7 +163,7 @@ func GitHubOAuth(c *gin.Context) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func GitHubBind(c *gin.Context) { | func GitHubBind(c *gin.Context) { | ||||||
| 	if !common.GitHubOAuthEnabled { | 	if !config.GitHubOAuthEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| 			"message": "管理员未开启通过 GitHub 登录以及注册", | 			"message": "管理员未开启通过 GitHub 登录以及注册", | ||||||
| @@ -216,7 +219,7 @@ func GitHubBind(c *gin.Context) { | |||||||
|  |  | ||||||
| func GenerateOAuthCode(c *gin.Context) { | func GenerateOAuthCode(c *gin.Context) { | ||||||
| 	session := sessions.Default(c) | 	session := sessions.Default(c) | ||||||
| 	state := common.GetRandomString(12) | 	state := helper.GetRandomString(12) | ||||||
| 	session.Set("oauth_state", state) | 	session.Set("oauth_state", state) | ||||||
| 	err := session.Save() | 	err := session.Save() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
|   | |||||||
| @@ -3,7 +3,7 @@ package controller | |||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| ) | ) | ||||||
| @@ -20,7 +20,7 @@ func GetAllLogs(c *gin.Context) { | |||||||
| 	tokenName := c.Query("token_name") | 	tokenName := c.Query("token_name") | ||||||
| 	modelName := c.Query("model_name") | 	modelName := c.Query("model_name") | ||||||
| 	channel, _ := strconv.Atoi(c.Query("channel")) | 	channel, _ := strconv.Atoi(c.Query("channel")) | ||||||
| 	logs, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, p*common.ItemsPerPage, common.ItemsPerPage, channel) | 	logs, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, p*config.ItemsPerPage, config.ItemsPerPage, channel) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -47,7 +47,7 @@ func GetUserLogs(c *gin.Context) { | |||||||
| 	endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64) | 	endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64) | ||||||
| 	tokenName := c.Query("token_name") | 	tokenName := c.Query("token_name") | ||||||
| 	modelName := c.Query("model_name") | 	modelName := c.Query("model_name") | ||||||
| 	logs, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, p*common.ItemsPerPage, common.ItemsPerPage) | 	logs, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, p*config.ItemsPerPage, config.ItemsPerPage) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
|   | |||||||
| @@ -5,6 +5,7 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strings" | 	"strings" | ||||||
|  |  | ||||||
| @@ -18,55 +19,55 @@ func GetStatus(c *gin.Context) { | |||||||
| 		"data": gin.H{ | 		"data": gin.H{ | ||||||
| 			"version":             common.Version, | 			"version":             common.Version, | ||||||
| 			"start_time":          common.StartTime, | 			"start_time":          common.StartTime, | ||||||
| 			"email_verification":  common.EmailVerificationEnabled, | 			"email_verification":  config.EmailVerificationEnabled, | ||||||
| 			"github_oauth":        common.GitHubOAuthEnabled, | 			"github_oauth":        config.GitHubOAuthEnabled, | ||||||
| 			"github_client_id":    common.GitHubClientId, | 			"github_client_id":    config.GitHubClientId, | ||||||
| 			"system_name":         common.SystemName, | 			"system_name":         config.SystemName, | ||||||
| 			"logo":                common.Logo, | 			"logo":                config.Logo, | ||||||
| 			"footer_html":         common.Footer, | 			"footer_html":         config.Footer, | ||||||
| 			"wechat_qrcode":       common.WeChatAccountQRCodeImageURL, | 			"wechat_qrcode":       config.WeChatAccountQRCodeImageURL, | ||||||
| 			"wechat_login":        common.WeChatAuthEnabled, | 			"wechat_login":        config.WeChatAuthEnabled, | ||||||
| 			"server_address":      common.ServerAddress, | 			"server_address":      config.ServerAddress, | ||||||
| 			"turnstile_check":     common.TurnstileCheckEnabled, | 			"turnstile_check":     config.TurnstileCheckEnabled, | ||||||
| 			"turnstile_site_key":  common.TurnstileSiteKey, | 			"turnstile_site_key":  config.TurnstileSiteKey, | ||||||
| 			"top_up_link":         common.TopUpLink, | 			"top_up_link":         config.TopUpLink, | ||||||
| 			"chat_link":           common.ChatLink, | 			"chat_link":           config.ChatLink, | ||||||
| 			"quota_per_unit":      common.QuotaPerUnit, | 			"quota_per_unit":      config.QuotaPerUnit, | ||||||
| 			"display_in_currency": common.DisplayInCurrencyEnabled, | 			"display_in_currency": config.DisplayInCurrencyEnabled, | ||||||
| 		}, | 		}, | ||||||
| 	}) | 	}) | ||||||
| 	return | 	return | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetNotice(c *gin.Context) { | func GetNotice(c *gin.Context) { | ||||||
| 	common.OptionMapRWMutex.RLock() | 	config.OptionMapRWMutex.RLock() | ||||||
| 	defer common.OptionMapRWMutex.RUnlock() | 	defer config.OptionMapRWMutex.RUnlock() | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
| 		"message": "", | 		"message": "", | ||||||
| 		"data":    common.OptionMap["Notice"], | 		"data":    config.OptionMap["Notice"], | ||||||
| 	}) | 	}) | ||||||
| 	return | 	return | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetAbout(c *gin.Context) { | func GetAbout(c *gin.Context) { | ||||||
| 	common.OptionMapRWMutex.RLock() | 	config.OptionMapRWMutex.RLock() | ||||||
| 	defer common.OptionMapRWMutex.RUnlock() | 	defer config.OptionMapRWMutex.RUnlock() | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
| 		"message": "", | 		"message": "", | ||||||
| 		"data":    common.OptionMap["About"], | 		"data":    config.OptionMap["About"], | ||||||
| 	}) | 	}) | ||||||
| 	return | 	return | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetHomePageContent(c *gin.Context) { | func GetHomePageContent(c *gin.Context) { | ||||||
| 	common.OptionMapRWMutex.RLock() | 	config.OptionMapRWMutex.RLock() | ||||||
| 	defer common.OptionMapRWMutex.RUnlock() | 	defer config.OptionMapRWMutex.RUnlock() | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
| 		"message": "", | 		"message": "", | ||||||
| 		"data":    common.OptionMap["HomePageContent"], | 		"data":    config.OptionMap["HomePageContent"], | ||||||
| 	}) | 	}) | ||||||
| 	return | 	return | ||||||
| } | } | ||||||
| @@ -80,9 +81,9 @@ func SendEmailVerification(c *gin.Context) { | |||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if common.EmailDomainRestrictionEnabled { | 	if config.EmailDomainRestrictionEnabled { | ||||||
| 		allowed := false | 		allowed := false | ||||||
| 		for _, domain := range common.EmailDomainWhitelist { | 		for _, domain := range config.EmailDomainWhitelist { | ||||||
| 			if strings.HasSuffix(email, "@"+domain) { | 			if strings.HasSuffix(email, "@"+domain) { | ||||||
| 				allowed = true | 				allowed = true | ||||||
| 				break | 				break | ||||||
| @@ -105,10 +106,10 @@ func SendEmailVerification(c *gin.Context) { | |||||||
| 	} | 	} | ||||||
| 	code := common.GenerateVerificationCode(6) | 	code := common.GenerateVerificationCode(6) | ||||||
| 	common.RegisterVerificationCodeWithKey(email, code, common.EmailVerificationPurpose) | 	common.RegisterVerificationCodeWithKey(email, code, common.EmailVerificationPurpose) | ||||||
| 	subject := fmt.Sprintf("%s邮箱验证邮件", common.SystemName) | 	subject := fmt.Sprintf("%s邮箱验证邮件", config.SystemName) | ||||||
| 	content := fmt.Sprintf("<p>您好,你正在进行%s邮箱验证。</p>"+ | 	content := fmt.Sprintf("<p>您好,你正在进行%s邮箱验证。</p>"+ | ||||||
| 		"<p>您的验证码为: <strong>%s</strong></p>"+ | 		"<p>您的验证码为: <strong>%s</strong></p>"+ | ||||||
| 		"<p>验证码 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, code, common.VerificationValidMinutes) | 		"<p>验证码 %d 分钟内有效,如果不是本人操作,请忽略。</p>", config.SystemName, code, common.VerificationValidMinutes) | ||||||
| 	err := common.SendEmail(subject, email, content) | 	err := common.SendEmail(subject, email, content) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| @@ -142,12 +143,12 @@ func SendPasswordResetEmail(c *gin.Context) { | |||||||
| 	} | 	} | ||||||
| 	code := common.GenerateVerificationCode(0) | 	code := common.GenerateVerificationCode(0) | ||||||
| 	common.RegisterVerificationCodeWithKey(email, code, common.PasswordResetPurpose) | 	common.RegisterVerificationCodeWithKey(email, code, common.PasswordResetPurpose) | ||||||
| 	link := fmt.Sprintf("%s/user/reset?email=%s&token=%s", common.ServerAddress, email, code) | 	link := fmt.Sprintf("%s/user/reset?email=%s&token=%s", config.ServerAddress, email, code) | ||||||
| 	subject := fmt.Sprintf("%s密码重置", common.SystemName) | 	subject := fmt.Sprintf("%s密码重置", config.SystemName) | ||||||
| 	content := fmt.Sprintf("<p>您好,你正在进行%s密码重置。</p>"+ | 	content := fmt.Sprintf("<p>您好,你正在进行%s密码重置。</p>"+ | ||||||
| 		"<p>点击 <a href='%s'>此处</a> 进行密码重置。</p>"+ | 		"<p>点击 <a href='%s'>此处</a> 进行密码重置。</p>"+ | ||||||
| 		"<p>如果链接无法点击,请尝试点击下面的链接或将其复制到浏览器中打开:<br> %s </p>"+ | 		"<p>如果链接无法点击,请尝试点击下面的链接或将其复制到浏览器中打开:<br> %s </p>"+ | ||||||
| 		"<p>重置链接 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, link, link, common.VerificationValidMinutes) | 		"<p>重置链接 %d 分钟内有效,如果不是本人操作,请忽略。</p>", config.SystemName, link, link, common.VerificationValidMinutes) | ||||||
| 	err := common.SendEmail(subject, email, content) | 	err := common.SendEmail(subject, email, content) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
|   | |||||||
| @@ -3,7 +3,8 @@ package controller | |||||||
| import ( | import ( | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strings" | 	"strings" | ||||||
|  |  | ||||||
| @@ -12,17 +13,17 @@ import ( | |||||||
|  |  | ||||||
| func GetOptions(c *gin.Context) { | func GetOptions(c *gin.Context) { | ||||||
| 	var options []*model.Option | 	var options []*model.Option | ||||||
| 	common.OptionMapRWMutex.Lock() | 	config.OptionMapRWMutex.Lock() | ||||||
| 	for k, v := range common.OptionMap { | 	for k, v := range config.OptionMap { | ||||||
| 		if strings.HasSuffix(k, "Token") || strings.HasSuffix(k, "Secret") { | 		if strings.HasSuffix(k, "Token") || strings.HasSuffix(k, "Secret") { | ||||||
| 			continue | 			continue | ||||||
| 		} | 		} | ||||||
| 		options = append(options, &model.Option{ | 		options = append(options, &model.Option{ | ||||||
| 			Key:   k, | 			Key:   k, | ||||||
| 			Value: common.Interface2String(v), | 			Value: helper.Interface2String(v), | ||||||
| 		}) | 		}) | ||||||
| 	} | 	} | ||||||
| 	common.OptionMapRWMutex.Unlock() | 	config.OptionMapRWMutex.Unlock() | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
| 		"message": "", | 		"message": "", | ||||||
| @@ -43,7 +44,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 	} | 	} | ||||||
| 	switch option.Key { | 	switch option.Key { | ||||||
| 	case "Theme": | 	case "Theme": | ||||||
| 		if !common.ValidThemes[option.Value] { | 		if !config.ValidThemes[option.Value] { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "无效的主题", | 				"message": "无效的主题", | ||||||
| @@ -51,7 +52,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	case "GitHubOAuthEnabled": | 	case "GitHubOAuthEnabled": | ||||||
| 		if option.Value == "true" && common.GitHubClientId == "" { | 		if option.Value == "true" && config.GitHubClientId == "" { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "无法启用 GitHub OAuth,请先填入 GitHub Client Id 以及 GitHub Client Secret!", | 				"message": "无法启用 GitHub OAuth,请先填入 GitHub Client Id 以及 GitHub Client Secret!", | ||||||
| @@ -59,7 +60,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	case "EmailDomainRestrictionEnabled": | 	case "EmailDomainRestrictionEnabled": | ||||||
| 		if option.Value == "true" && len(common.EmailDomainWhitelist) == 0 { | 		if option.Value == "true" && len(config.EmailDomainWhitelist) == 0 { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "无法启用邮箱域名限制,请先填入限制的邮箱域名!", | 				"message": "无法启用邮箱域名限制,请先填入限制的邮箱域名!", | ||||||
| @@ -67,7 +68,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	case "WeChatAuthEnabled": | 	case "WeChatAuthEnabled": | ||||||
| 		if option.Value == "true" && common.WeChatServerAddress == "" { | 		if option.Value == "true" && config.WeChatServerAddress == "" { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "无法启用微信登录,请先填入微信登录相关配置信息!", | 				"message": "无法启用微信登录,请先填入微信登录相关配置信息!", | ||||||
| @@ -75,7 +76,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	case "TurnstileCheckEnabled": | 	case "TurnstileCheckEnabled": | ||||||
| 		if option.Value == "true" && common.TurnstileSiteKey == "" { | 		if option.Value == "true" && config.TurnstileSiteKey == "" { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "无法启用 Turnstile 校验,请先填入 Turnstile 校验相关配置信息!", | 				"message": "无法启用 Turnstile 校验,请先填入 Turnstile 校验相关配置信息!", | ||||||
|   | |||||||
| @@ -3,7 +3,8 @@ package controller | |||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| ) | ) | ||||||
| @@ -13,7 +14,7 @@ func GetAllRedemptions(c *gin.Context) { | |||||||
| 	if p < 0 { | 	if p < 0 { | ||||||
| 		p = 0 | 		p = 0 | ||||||
| 	} | 	} | ||||||
| 	redemptions, err := model.GetAllRedemptions(p*common.ItemsPerPage, common.ItemsPerPage) | 	redemptions, err := model.GetAllRedemptions(p*config.ItemsPerPage, config.ItemsPerPage) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -105,12 +106,12 @@ func AddRedemption(c *gin.Context) { | |||||||
| 	} | 	} | ||||||
| 	var keys []string | 	var keys []string | ||||||
| 	for i := 0; i < redemption.Count; i++ { | 	for i := 0; i < redemption.Count; i++ { | ||||||
| 		key := common.GetUUID() | 		key := helper.GetUUID() | ||||||
| 		cleanRedemption := model.Redemption{ | 		cleanRedemption := model.Redemption{ | ||||||
| 			UserId:      c.GetInt("id"), | 			UserId:      c.GetInt("id"), | ||||||
| 			Name:        redemption.Name, | 			Name:        redemption.Name, | ||||||
| 			Key:         key, | 			Key:         key, | ||||||
| 			CreatedTime: common.GetTimestamp(), | 			CreatedTime: helper.GetTimestamp(), | ||||||
| 			Quota:       redemption.Quota, | 			Quota:       redemption.Quota, | ||||||
| 		} | 		} | ||||||
| 		err = cleanRedemption.Insert() | 		err = cleanRedemption.Insert() | ||||||
|   | |||||||
| @@ -2,43 +2,22 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"fmt" | 	"fmt" | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| 	"one-api/relay/controller" | 	"one-api/relay/controller" | ||||||
| 	"one-api/relay/util" | 	"one-api/relay/util" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" |  | ||||||
|  |  | ||||||
| 	"github.com/gin-gonic/gin" |  | ||||||
| ) | ) | ||||||
|  |  | ||||||
| // https://platform.openai.com/docs/api-reference/chat | // https://platform.openai.com/docs/api-reference/chat | ||||||
|  |  | ||||||
| func Relay(c *gin.Context) { | func Relay(c *gin.Context) { | ||||||
| 	relayMode := constant.RelayModeUnknown | 	relayMode := constant.Path2RelayMode(c.Request.URL.Path) | ||||||
| 	if strings.HasPrefix(c.Request.URL.Path, "/v1/chat/completions") { |  | ||||||
| 		relayMode = constant.RelayModeChatCompletions |  | ||||||
| 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/completions") { |  | ||||||
| 		relayMode = constant.RelayModeCompletions |  | ||||||
| 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/embeddings") { |  | ||||||
| 		relayMode = constant.RelayModeEmbeddings |  | ||||||
| 	} else if strings.HasSuffix(c.Request.URL.Path, "embeddings") { |  | ||||||
| 		relayMode = constant.RelayModeEmbeddings |  | ||||||
| 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/moderations") { |  | ||||||
| 		relayMode = constant.RelayModeModerations |  | ||||||
| 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/images/generations") { |  | ||||||
| 		relayMode = constant.RelayModeImagesGenerations |  | ||||||
| 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/edits") { |  | ||||||
| 		relayMode = constant.RelayModeEdits |  | ||||||
| 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/audio/speech") { |  | ||||||
| 		relayMode = constant.RelayModeAudioSpeech |  | ||||||
| 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/audio/transcriptions") { |  | ||||||
| 		relayMode = constant.RelayModeAudioTranscription |  | ||||||
| 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/audio/translations") { |  | ||||||
| 		relayMode = constant.RelayModeAudioTranslation |  | ||||||
| 	} |  | ||||||
| 	var err *openai.ErrorWithStatusCode | 	var err *openai.ErrorWithStatusCode | ||||||
| 	switch relayMode { | 	switch relayMode { | ||||||
| 	case constant.RelayModeImagesGenerations: | 	case constant.RelayModeImagesGenerations: | ||||||
| @@ -53,11 +32,11 @@ func Relay(c *gin.Context) { | |||||||
| 		err = controller.RelayTextHelper(c, relayMode) | 		err = controller.RelayTextHelper(c, relayMode) | ||||||
| 	} | 	} | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		requestId := c.GetString(common.RequestIdKey) | 		requestId := c.GetString(logger.RequestIdKey) | ||||||
| 		retryTimesStr := c.Query("retry") | 		retryTimesStr := c.Query("retry") | ||||||
| 		retryTimes, _ := strconv.Atoi(retryTimesStr) | 		retryTimes, _ := strconv.Atoi(retryTimesStr) | ||||||
| 		if retryTimesStr == "" { | 		if retryTimesStr == "" { | ||||||
| 			retryTimes = common.RetryTimes | 			retryTimes = config.RetryTimes | ||||||
| 		} | 		} | ||||||
| 		if retryTimes > 0 { | 		if retryTimes > 0 { | ||||||
| 			c.Redirect(http.StatusTemporaryRedirect, fmt.Sprintf("%s?retry=%d", c.Request.URL.Path, retryTimes-1)) | 			c.Redirect(http.StatusTemporaryRedirect, fmt.Sprintf("%s?retry=%d", c.Request.URL.Path, retryTimes-1)) | ||||||
| @@ -65,13 +44,13 @@ func Relay(c *gin.Context) { | |||||||
| 			if err.StatusCode == http.StatusTooManyRequests { | 			if err.StatusCode == http.StatusTooManyRequests { | ||||||
| 				err.Error.Message = "当前分组上游负载已饱和,请稍后再试" | 				err.Error.Message = "当前分组上游负载已饱和,请稍后再试" | ||||||
| 			} | 			} | ||||||
| 			err.Error.Message = common.MessageWithRequestId(err.Error.Message, requestId) | 			err.Error.Message = helper.MessageWithRequestId(err.Error.Message, requestId) | ||||||
| 			c.JSON(err.StatusCode, gin.H{ | 			c.JSON(err.StatusCode, gin.H{ | ||||||
| 				"error": err.Error, | 				"error": err.Error, | ||||||
| 			}) | 			}) | ||||||
| 		} | 		} | ||||||
| 		channelId := c.GetInt("channel_id") | 		channelId := c.GetInt("channel_id") | ||||||
| 		common.LogError(c.Request.Context(), fmt.Sprintf("relay error (channel #%d): %s", channelId, err.Message)) | 		logger.Error(c.Request.Context(), fmt.Sprintf("relay error (channel #%d): %s", channelId, err.Message)) | ||||||
| 		// https://platform.openai.com/docs/guides/error-codes/api-errors | 		// https://platform.openai.com/docs/guides/error-codes/api-errors | ||||||
| 		if util.ShouldDisableChannel(&err.Error, err.StatusCode) { | 		if util.ShouldDisableChannel(&err.Error, err.StatusCode) { | ||||||
| 			channelId := c.GetInt("channel_id") | 			channelId := c.GetInt("channel_id") | ||||||
|   | |||||||
| @@ -4,6 +4,8 @@ import ( | |||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| ) | ) | ||||||
| @@ -14,7 +16,7 @@ func GetAllTokens(c *gin.Context) { | |||||||
| 	if p < 0 { | 	if p < 0 { | ||||||
| 		p = 0 | 		p = 0 | ||||||
| 	} | 	} | ||||||
| 	tokens, err := model.GetAllUserTokens(userId, p*common.ItemsPerPage, common.ItemsPerPage) | 	tokens, err := model.GetAllUserTokens(userId, p*config.ItemsPerPage, config.ItemsPerPage) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -119,9 +121,9 @@ func AddToken(c *gin.Context) { | |||||||
| 	cleanToken := model.Token{ | 	cleanToken := model.Token{ | ||||||
| 		UserId:         c.GetInt("id"), | 		UserId:         c.GetInt("id"), | ||||||
| 		Name:           token.Name, | 		Name:           token.Name, | ||||||
| 		Key:            common.GenerateKey(), | 		Key:            helper.GenerateKey(), | ||||||
| 		CreatedTime:    common.GetTimestamp(), | 		CreatedTime:    helper.GetTimestamp(), | ||||||
| 		AccessedTime:   common.GetTimestamp(), | 		AccessedTime:   helper.GetTimestamp(), | ||||||
| 		ExpiredTime:    token.ExpiredTime, | 		ExpiredTime:    token.ExpiredTime, | ||||||
| 		RemainQuota:    token.RemainQuota, | 		RemainQuota:    token.RemainQuota, | ||||||
| 		UnlimitedQuota: token.UnlimitedQuota, | 		UnlimitedQuota: token.UnlimitedQuota, | ||||||
| @@ -187,7 +189,7 @@ func UpdateToken(c *gin.Context) { | |||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if token.Status == common.TokenStatusEnabled { | 	if token.Status == common.TokenStatusEnabled { | ||||||
| 		if cleanToken.Status == common.TokenStatusExpired && cleanToken.ExpiredTime <= common.GetTimestamp() && cleanToken.ExpiredTime != -1 { | 		if cleanToken.Status == common.TokenStatusExpired && cleanToken.ExpiredTime <= helper.GetTimestamp() && cleanToken.ExpiredTime != -1 { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "令牌已过期,无法启用,请先修改令牌过期时间,或者设置为永不过期", | 				"message": "令牌已过期,无法启用,请先修改令牌过期时间,或者设置为永不过期", | ||||||
|   | |||||||
| @@ -5,6 +5,8 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"time" | 	"time" | ||||||
| @@ -19,7 +21,7 @@ type LoginRequest struct { | |||||||
| } | } | ||||||
|  |  | ||||||
| func Login(c *gin.Context) { | func Login(c *gin.Context) { | ||||||
| 	if !common.PasswordLoginEnabled { | 	if !config.PasswordLoginEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员关闭了密码登录", | 			"message": "管理员关闭了密码登录", | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -106,14 +108,14 @@ func Logout(c *gin.Context) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func Register(c *gin.Context) { | func Register(c *gin.Context) { | ||||||
| 	if !common.RegisterEnabled { | 	if !config.RegisterEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员关闭了新用户注册", | 			"message": "管理员关闭了新用户注册", | ||||||
| 			"success": false, | 			"success": false, | ||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if !common.PasswordRegisterEnabled { | 	if !config.PasswordRegisterEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员关闭了通过密码进行注册,请使用第三方账户验证的形式进行注册", | 			"message": "管理员关闭了通过密码进行注册,请使用第三方账户验证的形式进行注册", | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -136,7 +138,7 @@ func Register(c *gin.Context) { | |||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if common.EmailVerificationEnabled { | 	if config.EmailVerificationEnabled { | ||||||
| 		if user.Email == "" || user.VerificationCode == "" { | 		if user.Email == "" || user.VerificationCode == "" { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| @@ -160,7 +162,7 @@ func Register(c *gin.Context) { | |||||||
| 		DisplayName: user.Username, | 		DisplayName: user.Username, | ||||||
| 		InviterId:   inviterId, | 		InviterId:   inviterId, | ||||||
| 	} | 	} | ||||||
| 	if common.EmailVerificationEnabled { | 	if config.EmailVerificationEnabled { | ||||||
| 		cleanUser.Email = user.Email | 		cleanUser.Email = user.Email | ||||||
| 	} | 	} | ||||||
| 	if err := cleanUser.Insert(inviterId); err != nil { | 	if err := cleanUser.Insert(inviterId); err != nil { | ||||||
| @@ -182,7 +184,7 @@ func GetAllUsers(c *gin.Context) { | |||||||
| 	if p < 0 { | 	if p < 0 { | ||||||
| 		p = 0 | 		p = 0 | ||||||
| 	} | 	} | ||||||
| 	users, err := model.GetAllUsers(p*common.ItemsPerPage, common.ItemsPerPage) | 	users, err := model.GetAllUsers(p*config.ItemsPerPage, config.ItemsPerPage) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -282,7 +284,7 @@ func GenerateAccessToken(c *gin.Context) { | |||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	user.AccessToken = common.GetUUID() | 	user.AccessToken = helper.GetUUID() | ||||||
|  |  | ||||||
| 	if model.DB.Where("access_token = ?", user.AccessToken).First(user).RowsAffected != 0 { | 	if model.DB.Where("access_token = ?", user.AccessToken).First(user).RowsAffected != 0 { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| @@ -319,7 +321,7 @@ func GetAffCode(c *gin.Context) { | |||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if user.AffCode == "" { | 	if user.AffCode == "" { | ||||||
| 		user.AffCode = common.GetRandomString(4) | 		user.AffCode = helper.GetRandomString(4) | ||||||
| 		if err := user.Update(false); err != nil { | 		if err := user.Update(false); err != nil { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| @@ -726,7 +728,7 @@ func EmailBind(c *gin.Context) { | |||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if user.Role == common.RoleRootUser { | 	if user.Role == common.RoleRootUser { | ||||||
| 		common.RootUserEmail = email | 		config.RootUserEmail = email | ||||||
| 	} | 	} | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
|   | |||||||
| @@ -7,6 +7,7 @@ import ( | |||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"time" | 	"time" | ||||||
| @@ -22,11 +23,11 @@ func getWeChatIdByCode(code string) (string, error) { | |||||||
| 	if code == "" { | 	if code == "" { | ||||||
| 		return "", errors.New("无效的参数") | 		return "", errors.New("无效的参数") | ||||||
| 	} | 	} | ||||||
| 	req, err := http.NewRequest("GET", fmt.Sprintf("%s/api/wechat/user?code=%s", common.WeChatServerAddress, code), nil) | 	req, err := http.NewRequest("GET", fmt.Sprintf("%s/api/wechat/user?code=%s", config.WeChatServerAddress, code), nil) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return "", err | 		return "", err | ||||||
| 	} | 	} | ||||||
| 	req.Header.Set("Authorization", common.WeChatServerToken) | 	req.Header.Set("Authorization", config.WeChatServerToken) | ||||||
| 	client := http.Client{ | 	client := http.Client{ | ||||||
| 		Timeout: 5 * time.Second, | 		Timeout: 5 * time.Second, | ||||||
| 	} | 	} | ||||||
| @@ -50,7 +51,7 @@ func getWeChatIdByCode(code string) (string, error) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func WeChatAuth(c *gin.Context) { | func WeChatAuth(c *gin.Context) { | ||||||
| 	if !common.WeChatAuthEnabled { | 	if !config.WeChatAuthEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员未开启通过微信登录以及注册", | 			"message": "管理员未开启通过微信登录以及注册", | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -79,7 +80,7 @@ func WeChatAuth(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	} else { | 	} else { | ||||||
| 		if common.RegisterEnabled { | 		if config.RegisterEnabled { | ||||||
| 			user.Username = "wechat_" + strconv.Itoa(model.GetMaxUserId()+1) | 			user.Username = "wechat_" + strconv.Itoa(model.GetMaxUserId()+1) | ||||||
| 			user.DisplayName = "WeChat User" | 			user.DisplayName = "WeChat User" | ||||||
| 			user.Role = common.RoleCommonUser | 			user.Role = common.RoleCommonUser | ||||||
| @@ -112,7 +113,7 @@ func WeChatAuth(c *gin.Context) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func WeChatBind(c *gin.Context) { | func WeChatBind(c *gin.Context) { | ||||||
| 	if !common.WeChatAuthEnabled { | 	if !config.WeChatAuthEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员未开启通过微信登录以及注册", | 			"message": "管理员未开启通过微信登录以及注册", | ||||||
| 			"success": false, | 			"success": false, | ||||||
|   | |||||||
							
								
								
									
										44
									
								
								main.go
									
									
									
									
									
								
							
							
						
						
									
										44
									
								
								main.go
									
									
									
									
									
								
							| @@ -7,6 +7,8 @@ import ( | |||||||
| 	"github.com/gin-contrib/sessions/cookie" | 	"github.com/gin-contrib/sessions/cookie" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/controller" | 	"one-api/controller" | ||||||
| 	"one-api/middleware" | 	"one-api/middleware" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| @@ -20,65 +22,65 @@ import ( | |||||||
| var buildFS embed.FS | var buildFS embed.FS | ||||||
|  |  | ||||||
| func main() { | func main() { | ||||||
| 	common.SetupLogger() | 	logger.SetupLogger() | ||||||
| 	common.SysLog(fmt.Sprintf("One API %s started", common.Version)) | 	logger.SysLog(fmt.Sprintf("One API %s started", common.Version)) | ||||||
| 	if os.Getenv("GIN_MODE") != "debug" { | 	if os.Getenv("GIN_MODE") != "debug" { | ||||||
| 		gin.SetMode(gin.ReleaseMode) | 		gin.SetMode(gin.ReleaseMode) | ||||||
| 	} | 	} | ||||||
| 	if common.DebugEnabled { | 	if config.DebugEnabled { | ||||||
| 		common.SysLog("running in debug mode") | 		logger.SysLog("running in debug mode") | ||||||
| 	} | 	} | ||||||
| 	// Initialize SQL Database | 	// Initialize SQL Database | ||||||
| 	err := model.InitDB() | 	err := model.InitDB() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.FatalLog("failed to initialize database: " + err.Error()) | 		logger.FatalLog("failed to initialize database: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	defer func() { | 	defer func() { | ||||||
| 		err := model.CloseDB() | 		err := model.CloseDB() | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.FatalLog("failed to close database: " + err.Error()) | 			logger.FatalLog("failed to close database: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 	}() | 	}() | ||||||
|  |  | ||||||
| 	// Initialize Redis | 	// Initialize Redis | ||||||
| 	err = common.InitRedisClient() | 	err = common.InitRedisClient() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.FatalLog("failed to initialize Redis: " + err.Error()) | 		logger.FatalLog("failed to initialize Redis: " + err.Error()) | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	// Initialize options | 	// Initialize options | ||||||
| 	model.InitOptionMap() | 	model.InitOptionMap() | ||||||
| 	common.SysLog(fmt.Sprintf("using theme %s", common.Theme)) | 	logger.SysLog(fmt.Sprintf("using theme %s", config.Theme)) | ||||||
| 	if common.RedisEnabled { | 	if common.RedisEnabled { | ||||||
| 		// for compatibility with old versions | 		// for compatibility with old versions | ||||||
| 		common.MemoryCacheEnabled = true | 		config.MemoryCacheEnabled = true | ||||||
| 	} | 	} | ||||||
| 	if common.MemoryCacheEnabled { | 	if config.MemoryCacheEnabled { | ||||||
| 		common.SysLog("memory cache enabled") | 		logger.SysLog("memory cache enabled") | ||||||
| 		common.SysError(fmt.Sprintf("sync frequency: %d seconds", common.SyncFrequency)) | 		logger.SysError(fmt.Sprintf("sync frequency: %d seconds", config.SyncFrequency)) | ||||||
| 		model.InitChannelCache() | 		model.InitChannelCache() | ||||||
| 	} | 	} | ||||||
| 	if common.MemoryCacheEnabled { | 	if config.MemoryCacheEnabled { | ||||||
| 		go model.SyncOptions(common.SyncFrequency) | 		go model.SyncOptions(config.SyncFrequency) | ||||||
| 		go model.SyncChannelCache(common.SyncFrequency) | 		go model.SyncChannelCache(config.SyncFrequency) | ||||||
| 	} | 	} | ||||||
| 	if os.Getenv("CHANNEL_UPDATE_FREQUENCY") != "" { | 	if os.Getenv("CHANNEL_UPDATE_FREQUENCY") != "" { | ||||||
| 		frequency, err := strconv.Atoi(os.Getenv("CHANNEL_UPDATE_FREQUENCY")) | 		frequency, err := strconv.Atoi(os.Getenv("CHANNEL_UPDATE_FREQUENCY")) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.FatalLog("failed to parse CHANNEL_UPDATE_FREQUENCY: " + err.Error()) | 			logger.FatalLog("failed to parse CHANNEL_UPDATE_FREQUENCY: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 		go controller.AutomaticallyUpdateChannels(frequency) | 		go controller.AutomaticallyUpdateChannels(frequency) | ||||||
| 	} | 	} | ||||||
| 	if os.Getenv("CHANNEL_TEST_FREQUENCY") != "" { | 	if os.Getenv("CHANNEL_TEST_FREQUENCY") != "" { | ||||||
| 		frequency, err := strconv.Atoi(os.Getenv("CHANNEL_TEST_FREQUENCY")) | 		frequency, err := strconv.Atoi(os.Getenv("CHANNEL_TEST_FREQUENCY")) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.FatalLog("failed to parse CHANNEL_TEST_FREQUENCY: " + err.Error()) | 			logger.FatalLog("failed to parse CHANNEL_TEST_FREQUENCY: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 		go controller.AutomaticallyTestChannels(frequency) | 		go controller.AutomaticallyTestChannels(frequency) | ||||||
| 	} | 	} | ||||||
| 	if os.Getenv("BATCH_UPDATE_ENABLED") == "true" { | 	if os.Getenv("BATCH_UPDATE_ENABLED") == "true" { | ||||||
| 		common.BatchUpdateEnabled = true | 		config.BatchUpdateEnabled = true | ||||||
| 		common.SysLog("batch update enabled with interval " + strconv.Itoa(common.BatchUpdateInterval) + "s") | 		logger.SysLog("batch update enabled with interval " + strconv.Itoa(config.BatchUpdateInterval) + "s") | ||||||
| 		model.InitBatchUpdater() | 		model.InitBatchUpdater() | ||||||
| 	} | 	} | ||||||
| 	openai.InitTokenEncoders() | 	openai.InitTokenEncoders() | ||||||
| @@ -91,7 +93,7 @@ func main() { | |||||||
| 	server.Use(middleware.RequestId()) | 	server.Use(middleware.RequestId()) | ||||||
| 	middleware.SetUpLogger(server) | 	middleware.SetUpLogger(server) | ||||||
| 	// Initialize session store | 	// Initialize session store | ||||||
| 	store := cookie.NewStore([]byte(common.SessionSecret)) | 	store := cookie.NewStore([]byte(config.SessionSecret)) | ||||||
| 	server.Use(sessions.Sessions("session", store)) | 	server.Use(sessions.Sessions("session", store)) | ||||||
|  |  | ||||||
| 	router.SetRouter(server, buildFS) | 	router.SetRouter(server, buildFS) | ||||||
| @@ -101,6 +103,6 @@ func main() { | |||||||
| 	} | 	} | ||||||
| 	err = server.Run(":" + port) | 	err = server.Run(":" + port) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.FatalLog("failed to start HTTP server: " + err.Error()) | 		logger.FatalLog("failed to start HTTP server: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -4,6 +4,7 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
| @@ -69,7 +70,7 @@ func Distribute() func(c *gin.Context) { | |||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				message := fmt.Sprintf("当前分组 %s 下对于模型 %s 无可用渠道", userGroup, modelRequest.Model) | 				message := fmt.Sprintf("当前分组 %s 下对于模型 %s 无可用渠道", userGroup, modelRequest.Model) | ||||||
| 				if channel != nil { | 				if channel != nil { | ||||||
| 					common.SysError(fmt.Sprintf("渠道不存在:%d", channel.Id)) | 					logger.SysError(fmt.Sprintf("渠道不存在:%d", channel.Id)) | ||||||
| 					message = "数据库一致性已被破坏,请联系管理员" | 					message = "数据库一致性已被破坏,请联系管理员" | ||||||
| 				} | 				} | ||||||
| 				abortWithMessage(c, http.StatusServiceUnavailable, message) | 				abortWithMessage(c, http.StatusServiceUnavailable, message) | ||||||
|   | |||||||
| @@ -3,14 +3,14 @@ package middleware | |||||||
| import ( | import ( | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"one-api/common" | 	"one-api/common/logger" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func SetUpLogger(server *gin.Engine) { | func SetUpLogger(server *gin.Engine) { | ||||||
| 	server.Use(gin.LoggerWithFormatter(func(param gin.LogFormatterParams) string { | 	server.Use(gin.LoggerWithFormatter(func(param gin.LogFormatterParams) string { | ||||||
| 		var requestID string | 		var requestID string | ||||||
| 		if param.Keys != nil { | 		if param.Keys != nil { | ||||||
| 			requestID = param.Keys[common.RequestIdKey].(string) | 			requestID = param.Keys[logger.RequestIdKey].(string) | ||||||
| 		} | 		} | ||||||
| 		return fmt.Sprintf("[GIN] %s | %s | %3d | %13v | %15s | %7s %s\n", | 		return fmt.Sprintf("[GIN] %s | %s | %3d | %13v | %15s | %7s %s\n", | ||||||
| 			param.TimeStamp.Format("2006/01/02 - 15:04:05"), | 			param.TimeStamp.Format("2006/01/02 - 15:04:05"), | ||||||
|   | |||||||
| @@ -6,6 +6,7 @@ import ( | |||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -26,7 +27,7 @@ func redisRateLimiter(c *gin.Context, maxRequestNum int, duration int64, mark st | |||||||
| 	} | 	} | ||||||
| 	if listLength < int64(maxRequestNum) { | 	if listLength < int64(maxRequestNum) { | ||||||
| 		rdb.LPush(ctx, key, time.Now().Format(timeFormat)) | 		rdb.LPush(ctx, key, time.Now().Format(timeFormat)) | ||||||
| 		rdb.Expire(ctx, key, common.RateLimitKeyExpirationDuration) | 		rdb.Expire(ctx, key, config.RateLimitKeyExpirationDuration) | ||||||
| 	} else { | 	} else { | ||||||
| 		oldTimeStr, _ := rdb.LIndex(ctx, key, -1).Result() | 		oldTimeStr, _ := rdb.LIndex(ctx, key, -1).Result() | ||||||
| 		oldTime, err := time.Parse(timeFormat, oldTimeStr) | 		oldTime, err := time.Parse(timeFormat, oldTimeStr) | ||||||
| @@ -47,14 +48,14 @@ func redisRateLimiter(c *gin.Context, maxRequestNum int, duration int64, mark st | |||||||
| 		// time.Since will return negative number! | 		// time.Since will return negative number! | ||||||
| 		// See: https://stackoverflow.com/questions/50970900/why-is-time-since-returning-negative-durations-on-windows | 		// See: https://stackoverflow.com/questions/50970900/why-is-time-since-returning-negative-durations-on-windows | ||||||
| 		if int64(nowTime.Sub(oldTime).Seconds()) < duration { | 		if int64(nowTime.Sub(oldTime).Seconds()) < duration { | ||||||
| 			rdb.Expire(ctx, key, common.RateLimitKeyExpirationDuration) | 			rdb.Expire(ctx, key, config.RateLimitKeyExpirationDuration) | ||||||
| 			c.Status(http.StatusTooManyRequests) | 			c.Status(http.StatusTooManyRequests) | ||||||
| 			c.Abort() | 			c.Abort() | ||||||
| 			return | 			return | ||||||
| 		} else { | 		} else { | ||||||
| 			rdb.LPush(ctx, key, time.Now().Format(timeFormat)) | 			rdb.LPush(ctx, key, time.Now().Format(timeFormat)) | ||||||
| 			rdb.LTrim(ctx, key, 0, int64(maxRequestNum-1)) | 			rdb.LTrim(ctx, key, 0, int64(maxRequestNum-1)) | ||||||
| 			rdb.Expire(ctx, key, common.RateLimitKeyExpirationDuration) | 			rdb.Expire(ctx, key, config.RateLimitKeyExpirationDuration) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| @@ -75,7 +76,7 @@ func rateLimitFactory(maxRequestNum int, duration int64, mark string) func(c *gi | |||||||
| 		} | 		} | ||||||
| 	} else { | 	} else { | ||||||
| 		// It's safe to call multi times. | 		// It's safe to call multi times. | ||||||
| 		inMemoryRateLimiter.Init(common.RateLimitKeyExpirationDuration) | 		inMemoryRateLimiter.Init(config.RateLimitKeyExpirationDuration) | ||||||
| 		return func(c *gin.Context) { | 		return func(c *gin.Context) { | ||||||
| 			memoryRateLimiter(c, maxRequestNum, duration, mark) | 			memoryRateLimiter(c, maxRequestNum, duration, mark) | ||||||
| 		} | 		} | ||||||
| @@ -83,21 +84,21 @@ func rateLimitFactory(maxRequestNum int, duration int64, mark string) func(c *gi | |||||||
| } | } | ||||||
|  |  | ||||||
| func GlobalWebRateLimit() func(c *gin.Context) { | func GlobalWebRateLimit() func(c *gin.Context) { | ||||||
| 	return rateLimitFactory(common.GlobalWebRateLimitNum, common.GlobalWebRateLimitDuration, "GW") | 	return rateLimitFactory(config.GlobalWebRateLimitNum, config.GlobalWebRateLimitDuration, "GW") | ||||||
| } | } | ||||||
|  |  | ||||||
| func GlobalAPIRateLimit() func(c *gin.Context) { | func GlobalAPIRateLimit() func(c *gin.Context) { | ||||||
| 	return rateLimitFactory(common.GlobalApiRateLimitNum, common.GlobalApiRateLimitDuration, "GA") | 	return rateLimitFactory(config.GlobalApiRateLimitNum, config.GlobalApiRateLimitDuration, "GA") | ||||||
| } | } | ||||||
|  |  | ||||||
| func CriticalRateLimit() func(c *gin.Context) { | func CriticalRateLimit() func(c *gin.Context) { | ||||||
| 	return rateLimitFactory(common.CriticalRateLimitNum, common.CriticalRateLimitDuration, "CT") | 	return rateLimitFactory(config.CriticalRateLimitNum, config.CriticalRateLimitDuration, "CT") | ||||||
| } | } | ||||||
|  |  | ||||||
| func DownloadRateLimit() func(c *gin.Context) { | func DownloadRateLimit() func(c *gin.Context) { | ||||||
| 	return rateLimitFactory(common.DownloadRateLimitNum, common.DownloadRateLimitDuration, "DW") | 	return rateLimitFactory(config.DownloadRateLimitNum, config.DownloadRateLimitDuration, "DW") | ||||||
| } | } | ||||||
|  |  | ||||||
| func UploadRateLimit() func(c *gin.Context) { | func UploadRateLimit() func(c *gin.Context) { | ||||||
| 	return rateLimitFactory(common.UploadRateLimitNum, common.UploadRateLimitDuration, "UP") | 	return rateLimitFactory(config.UploadRateLimitNum, config.UploadRateLimitDuration, "UP") | ||||||
| } | } | ||||||
|   | |||||||
| @@ -4,7 +4,7 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common/logger" | ||||||
| 	"runtime/debug" | 	"runtime/debug" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -12,8 +12,8 @@ func RelayPanicRecover() gin.HandlerFunc { | |||||||
| 	return func(c *gin.Context) { | 	return func(c *gin.Context) { | ||||||
| 		defer func() { | 		defer func() { | ||||||
| 			if err := recover(); err != nil { | 			if err := recover(); err != nil { | ||||||
| 				common.SysError(fmt.Sprintf("panic detected: %v", err)) | 				logger.SysError(fmt.Sprintf("panic detected: %v", err)) | ||||||
| 				common.SysError(fmt.Sprintf("stacktrace from panic: %s", string(debug.Stack()))) | 				logger.SysError(fmt.Sprintf("stacktrace from panic: %s", string(debug.Stack()))) | ||||||
| 				c.JSON(http.StatusInternalServerError, gin.H{ | 				c.JSON(http.StatusInternalServerError, gin.H{ | ||||||
| 					"error": gin.H{ | 					"error": gin.H{ | ||||||
| 						"message": fmt.Sprintf("Panic detected, error: %v. Please submit a issue here: https://github.com/songquanpeng/one-api", err), | 						"message": fmt.Sprintf("Panic detected, error: %v. Please submit a issue here: https://github.com/songquanpeng/one-api", err), | ||||||
|   | |||||||
| @@ -3,16 +3,17 @@ package middleware | |||||||
| import ( | import ( | ||||||
| 	"context" | 	"context" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"one-api/common" | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func RequestId() func(c *gin.Context) { | func RequestId() func(c *gin.Context) { | ||||||
| 	return func(c *gin.Context) { | 	return func(c *gin.Context) { | ||||||
| 		id := common.GetTimeString() + common.GetRandomString(8) | 		id := helper.GetTimeString() + helper.GetRandomString(8) | ||||||
| 		c.Set(common.RequestIdKey, id) | 		c.Set(logger.RequestIdKey, id) | ||||||
| 		ctx := context.WithValue(c.Request.Context(), common.RequestIdKey, id) | 		ctx := context.WithValue(c.Request.Context(), logger.RequestIdKey, id) | ||||||
| 		c.Request = c.Request.WithContext(ctx) | 		c.Request = c.Request.WithContext(ctx) | ||||||
| 		c.Header(common.RequestIdKey, id) | 		c.Header(logger.RequestIdKey, id) | ||||||
| 		c.Next() | 		c.Next() | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -6,7 +6,8 @@ import ( | |||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"net/url" | 	"net/url" | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| type turnstileCheckResponse struct { | type turnstileCheckResponse struct { | ||||||
| @@ -15,7 +16,7 @@ type turnstileCheckResponse struct { | |||||||
|  |  | ||||||
| func TurnstileCheck() gin.HandlerFunc { | func TurnstileCheck() gin.HandlerFunc { | ||||||
| 	return func(c *gin.Context) { | 	return func(c *gin.Context) { | ||||||
| 		if common.TurnstileCheckEnabled { | 		if config.TurnstileCheckEnabled { | ||||||
| 			session := sessions.Default(c) | 			session := sessions.Default(c) | ||||||
| 			turnstileChecked := session.Get("turnstile") | 			turnstileChecked := session.Get("turnstile") | ||||||
| 			if turnstileChecked != nil { | 			if turnstileChecked != nil { | ||||||
| @@ -32,12 +33,12 @@ func TurnstileCheck() gin.HandlerFunc { | |||||||
| 				return | 				return | ||||||
| 			} | 			} | ||||||
| 			rawRes, err := http.PostForm("https://challenges.cloudflare.com/turnstile/v0/siteverify", url.Values{ | 			rawRes, err := http.PostForm("https://challenges.cloudflare.com/turnstile/v0/siteverify", url.Values{ | ||||||
| 				"secret":   {common.TurnstileSecretKey}, | 				"secret":   {config.TurnstileSecretKey}, | ||||||
| 				"response": {response}, | 				"response": {response}, | ||||||
| 				"remoteip": {c.ClientIP()}, | 				"remoteip": {c.ClientIP()}, | ||||||
| 			}) | 			}) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError(err.Error()) | 				logger.SysError(err.Error()) | ||||||
| 				c.JSON(http.StatusOK, gin.H{ | 				c.JSON(http.StatusOK, gin.H{ | ||||||
| 					"success": false, | 					"success": false, | ||||||
| 					"message": err.Error(), | 					"message": err.Error(), | ||||||
| @@ -49,7 +50,7 @@ func TurnstileCheck() gin.HandlerFunc { | |||||||
| 			var res turnstileCheckResponse | 			var res turnstileCheckResponse | ||||||
| 			err = json.NewDecoder(rawRes.Body).Decode(&res) | 			err = json.NewDecoder(rawRes.Body).Decode(&res) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError(err.Error()) | 				logger.SysError(err.Error()) | ||||||
| 				c.JSON(http.StatusOK, gin.H{ | 				c.JSON(http.StatusOK, gin.H{ | ||||||
| 					"success": false, | 					"success": false, | ||||||
| 					"message": err.Error(), | 					"message": err.Error(), | ||||||
|   | |||||||
| @@ -2,16 +2,17 @@ package middleware | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"one-api/common" | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func abortWithMessage(c *gin.Context, statusCode int, message string) { | func abortWithMessage(c *gin.Context, statusCode int, message string) { | ||||||
| 	c.JSON(statusCode, gin.H{ | 	c.JSON(statusCode, gin.H{ | ||||||
| 		"error": gin.H{ | 		"error": gin.H{ | ||||||
| 			"message": common.MessageWithRequestId(message, c.GetString(common.RequestIdKey)), | 			"message": helper.MessageWithRequestId(message, c.GetString(logger.RequestIdKey)), | ||||||
| 			"type":    "one_api_error", | 			"type":    "one_api_error", | ||||||
| 		}, | 		}, | ||||||
| 	}) | 	}) | ||||||
| 	c.Abort() | 	c.Abort() | ||||||
| 	common.LogError(c.Request.Context(), message) | 	logger.Error(c.Request.Context(), message) | ||||||
| } | } | ||||||
|   | |||||||
| @@ -6,6 +6,8 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"math/rand" | 	"math/rand" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"sort" | 	"sort" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
| @@ -14,10 +16,10 @@ import ( | |||||||
| ) | ) | ||||||
|  |  | ||||||
| var ( | var ( | ||||||
| 	TokenCacheSeconds         = common.SyncFrequency | 	TokenCacheSeconds         = config.SyncFrequency | ||||||
| 	UserId2GroupCacheSeconds  = common.SyncFrequency | 	UserId2GroupCacheSeconds  = config.SyncFrequency | ||||||
| 	UserId2QuotaCacheSeconds  = common.SyncFrequency | 	UserId2QuotaCacheSeconds  = config.SyncFrequency | ||||||
| 	UserId2StatusCacheSeconds = common.SyncFrequency | 	UserId2StatusCacheSeconds = config.SyncFrequency | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func CacheGetTokenByKey(key string) (*Token, error) { | func CacheGetTokenByKey(key string) (*Token, error) { | ||||||
| @@ -42,7 +44,7 @@ func CacheGetTokenByKey(key string) (*Token, error) { | |||||||
| 		} | 		} | ||||||
| 		err = common.RedisSet(fmt.Sprintf("token:%s", key), string(jsonBytes), time.Duration(TokenCacheSeconds)*time.Second) | 		err = common.RedisSet(fmt.Sprintf("token:%s", key), string(jsonBytes), time.Duration(TokenCacheSeconds)*time.Second) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("Redis set token error: " + err.Error()) | 			logger.SysError("Redis set token error: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 		return &token, nil | 		return &token, nil | ||||||
| 	} | 	} | ||||||
| @@ -62,7 +64,7 @@ func CacheGetUserGroup(id int) (group string, err error) { | |||||||
| 		} | 		} | ||||||
| 		err = common.RedisSet(fmt.Sprintf("user_group:%d", id), group, time.Duration(UserId2GroupCacheSeconds)*time.Second) | 		err = common.RedisSet(fmt.Sprintf("user_group:%d", id), group, time.Duration(UserId2GroupCacheSeconds)*time.Second) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("Redis set user group error: " + err.Error()) | 			logger.SysError("Redis set user group error: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	return group, err | 	return group, err | ||||||
| @@ -80,7 +82,7 @@ func CacheGetUserQuota(id int) (quota int, err error) { | |||||||
| 		} | 		} | ||||||
| 		err = common.RedisSet(fmt.Sprintf("user_quota:%d", id), fmt.Sprintf("%d", quota), time.Duration(UserId2QuotaCacheSeconds)*time.Second) | 		err = common.RedisSet(fmt.Sprintf("user_quota:%d", id), fmt.Sprintf("%d", quota), time.Duration(UserId2QuotaCacheSeconds)*time.Second) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("Redis set user quota error: " + err.Error()) | 			logger.SysError("Redis set user quota error: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 		return quota, err | 		return quota, err | ||||||
| 	} | 	} | ||||||
| @@ -127,7 +129,7 @@ func CacheIsUserEnabled(userId int) (bool, error) { | |||||||
| 	} | 	} | ||||||
| 	err = common.RedisSet(fmt.Sprintf("user_enabled:%d", userId), enabled, time.Duration(UserId2StatusCacheSeconds)*time.Second) | 	err = common.RedisSet(fmt.Sprintf("user_enabled:%d", userId), enabled, time.Duration(UserId2StatusCacheSeconds)*time.Second) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("Redis set user enabled error: " + err.Error()) | 		logger.SysError("Redis set user enabled error: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return userEnabled, err | 	return userEnabled, err | ||||||
| } | } | ||||||
| @@ -178,19 +180,19 @@ func InitChannelCache() { | |||||||
| 	channelSyncLock.Lock() | 	channelSyncLock.Lock() | ||||||
| 	group2model2channels = newGroup2model2channels | 	group2model2channels = newGroup2model2channels | ||||||
| 	channelSyncLock.Unlock() | 	channelSyncLock.Unlock() | ||||||
| 	common.SysLog("channels synced from database") | 	logger.SysLog("channels synced from database") | ||||||
| } | } | ||||||
|  |  | ||||||
| func SyncChannelCache(frequency int) { | func SyncChannelCache(frequency int) { | ||||||
| 	for { | 	for { | ||||||
| 		time.Sleep(time.Duration(frequency) * time.Second) | 		time.Sleep(time.Duration(frequency) * time.Second) | ||||||
| 		common.SysLog("syncing channels from database") | 		logger.SysLog("syncing channels from database") | ||||||
| 		InitChannelCache() | 		InitChannelCache() | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func CacheGetRandomSatisfiedChannel(group string, model string) (*Channel, error) { | func CacheGetRandomSatisfiedChannel(group string, model string) (*Channel, error) { | ||||||
| 	if !common.MemoryCacheEnabled { | 	if !config.MemoryCacheEnabled { | ||||||
| 		return GetRandomSatisfiedChannel(group, model) | 		return GetRandomSatisfiedChannel(group, model) | ||||||
| 	} | 	} | ||||||
| 	channelSyncLock.RLock() | 	channelSyncLock.RLock() | ||||||
|   | |||||||
| @@ -1,8 +1,13 @@ | |||||||
| package model | package model | ||||||
|  |  | ||||||
| import ( | import ( | ||||||
|  | 	"encoding/json" | ||||||
|  | 	"fmt" | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| type Channel struct { | type Channel struct { | ||||||
| @@ -42,7 +47,7 @@ func SearchChannels(keyword string) (channels []*Channel, err error) { | |||||||
| 	if common.UsingPostgreSQL { | 	if common.UsingPostgreSQL { | ||||||
| 		keyCol = `"key"` | 		keyCol = `"key"` | ||||||
| 	} | 	} | ||||||
| 	err = DB.Omit("key").Where("id = ? or name LIKE ? or "+keyCol+" = ?", common.String2Int(keyword), keyword+"%", keyword).Find(&channels).Error | 	err = DB.Omit("key").Where("id = ? or name LIKE ? or "+keyCol+" = ?", helper.String2Int(keyword), keyword+"%", keyword).Find(&channels).Error | ||||||
| 	return channels, err | 	return channels, err | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -86,11 +91,17 @@ func (channel *Channel) GetBaseURL() string { | |||||||
| 	return *channel.BaseURL | 	return *channel.BaseURL | ||||||
| } | } | ||||||
|  |  | ||||||
| func (channel *Channel) GetModelMapping() string { | func (channel *Channel) GetModelMapping() map[string]string { | ||||||
| 	if channel.ModelMapping == nil { | 	if channel.ModelMapping == nil || *channel.ModelMapping == "" || *channel.ModelMapping == "{}" { | ||||||
| 		return "" | 		return nil | ||||||
| 	} | 	} | ||||||
| 	return *channel.ModelMapping | 	modelMapping := make(map[string]string) | ||||||
|  | 	err := json.Unmarshal([]byte(*channel.ModelMapping), &modelMapping) | ||||||
|  | 	if err != nil { | ||||||
|  | 		logger.SysError(fmt.Sprintf("failed to unmarshal model mapping for channel %d, error: %s", channel.Id, err.Error())) | ||||||
|  | 		return nil | ||||||
|  | 	} | ||||||
|  | 	return modelMapping | ||||||
| } | } | ||||||
|  |  | ||||||
| func (channel *Channel) Insert() error { | func (channel *Channel) Insert() error { | ||||||
| @@ -116,21 +127,21 @@ func (channel *Channel) Update() error { | |||||||
|  |  | ||||||
| func (channel *Channel) UpdateResponseTime(responseTime int64) { | func (channel *Channel) UpdateResponseTime(responseTime int64) { | ||||||
| 	err := DB.Model(channel).Select("response_time", "test_time").Updates(Channel{ | 	err := DB.Model(channel).Select("response_time", "test_time").Updates(Channel{ | ||||||
| 		TestTime:     common.GetTimestamp(), | 		TestTime:     helper.GetTimestamp(), | ||||||
| 		ResponseTime: int(responseTime), | 		ResponseTime: int(responseTime), | ||||||
| 	}).Error | 	}).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("failed to update response time: " + err.Error()) | 		logger.SysError("failed to update response time: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func (channel *Channel) UpdateBalance(balance float64) { | func (channel *Channel) UpdateBalance(balance float64) { | ||||||
| 	err := DB.Model(channel).Select("balance_updated_time", "balance").Updates(Channel{ | 	err := DB.Model(channel).Select("balance_updated_time", "balance").Updates(Channel{ | ||||||
| 		BalanceUpdatedTime: common.GetTimestamp(), | 		BalanceUpdatedTime: helper.GetTimestamp(), | ||||||
| 		Balance:            balance, | 		Balance:            balance, | ||||||
| 	}).Error | 	}).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("failed to update balance: " + err.Error()) | 		logger.SysError("failed to update balance: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -147,16 +158,16 @@ func (channel *Channel) Delete() error { | |||||||
| func UpdateChannelStatusById(id int, status int) { | func UpdateChannelStatusById(id int, status int) { | ||||||
| 	err := UpdateAbilityStatus(id, status == common.ChannelStatusEnabled) | 	err := UpdateAbilityStatus(id, status == common.ChannelStatusEnabled) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("failed to update ability status: " + err.Error()) | 		logger.SysError("failed to update ability status: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	err = DB.Model(&Channel{}).Where("id = ?", id).Update("status", status).Error | 	err = DB.Model(&Channel{}).Where("id = ?", id).Update("status", status).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("failed to update channel status: " + err.Error()) | 		logger.SysError("failed to update channel status: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func UpdateChannelUsedQuota(id int, quota int) { | func UpdateChannelUsedQuota(id int, quota int) { | ||||||
| 	if common.BatchUpdateEnabled { | 	if config.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeChannelUsedQuota, id, quota) | 		addNewRecord(BatchUpdateTypeChannelUsedQuota, id, quota) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| @@ -166,7 +177,7 @@ func UpdateChannelUsedQuota(id int, quota int) { | |||||||
| func updateChannelUsedQuota(id int, quota int) { | func updateChannelUsedQuota(id int, quota int) { | ||||||
| 	err := DB.Model(&Channel{}).Where("id = ?", id).Update("used_quota", gorm.Expr("used_quota + ?", quota)).Error | 	err := DB.Model(&Channel{}).Where("id = ?", id).Update("used_quota", gorm.Expr("used_quota + ?", quota)).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("failed to update channel used quota: " + err.Error()) | 		logger.SysError("failed to update channel used quota: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
|   | |||||||
							
								
								
									
										21
									
								
								model/log.go
									
									
									
									
									
								
							
							
						
						
									
										21
									
								
								model/log.go
									
									
									
									
									
								
							| @@ -4,6 +4,9 @@ import ( | |||||||
| 	"context" | 	"context" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
|  |  | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
| ) | ) | ||||||
| @@ -32,31 +35,31 @@ const ( | |||||||
| ) | ) | ||||||
|  |  | ||||||
| func RecordLog(userId int, logType int, content string) { | func RecordLog(userId int, logType int, content string) { | ||||||
| 	if logType == LogTypeConsume && !common.LogConsumeEnabled { | 	if logType == LogTypeConsume && !config.LogConsumeEnabled { | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	log := &Log{ | 	log := &Log{ | ||||||
| 		UserId:    userId, | 		UserId:    userId, | ||||||
| 		Username:  GetUsernameById(userId), | 		Username:  GetUsernameById(userId), | ||||||
| 		CreatedAt: common.GetTimestamp(), | 		CreatedAt: helper.GetTimestamp(), | ||||||
| 		Type:      logType, | 		Type:      logType, | ||||||
| 		Content:   content, | 		Content:   content, | ||||||
| 	} | 	} | ||||||
| 	err := DB.Create(log).Error | 	err := DB.Create(log).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("failed to record log: " + err.Error()) | 		logger.SysError("failed to record log: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func RecordConsumeLog(ctx context.Context, userId int, channelId int, promptTokens int, completionTokens int, modelName string, tokenName string, quota int, content string) { | func RecordConsumeLog(ctx context.Context, userId int, channelId int, promptTokens int, completionTokens int, modelName string, tokenName string, quota int, content string) { | ||||||
| 	common.LogInfo(ctx, fmt.Sprintf("record consume log: userId=%d, channelId=%d, promptTokens=%d, completionTokens=%d, modelName=%s, tokenName=%s, quota=%d, content=%s", userId, channelId, promptTokens, completionTokens, modelName, tokenName, quota, content)) | 	logger.Info(ctx, fmt.Sprintf("record consume log: userId=%d, channelId=%d, promptTokens=%d, completionTokens=%d, modelName=%s, tokenName=%s, quota=%d, content=%s", userId, channelId, promptTokens, completionTokens, modelName, tokenName, quota, content)) | ||||||
| 	if !common.LogConsumeEnabled { | 	if !config.LogConsumeEnabled { | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	log := &Log{ | 	log := &Log{ | ||||||
| 		UserId:           userId, | 		UserId:           userId, | ||||||
| 		Username:         GetUsernameById(userId), | 		Username:         GetUsernameById(userId), | ||||||
| 		CreatedAt:        common.GetTimestamp(), | 		CreatedAt:        helper.GetTimestamp(), | ||||||
| 		Type:             LogTypeConsume, | 		Type:             LogTypeConsume, | ||||||
| 		Content:          content, | 		Content:          content, | ||||||
| 		PromptTokens:     promptTokens, | 		PromptTokens:     promptTokens, | ||||||
| @@ -68,7 +71,7 @@ func RecordConsumeLog(ctx context.Context, userId int, channelId int, promptToke | |||||||
| 	} | 	} | ||||||
| 	err := DB.Create(log).Error | 	err := DB.Create(log).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.LogError(ctx, "failed to record log: "+err.Error()) | 		logger.Error(ctx, "failed to record log: "+err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -125,12 +128,12 @@ func GetUserLogs(userId int, logType int, startTimestamp int64, endTimestamp int | |||||||
| } | } | ||||||
|  |  | ||||||
| func SearchAllLogs(keyword string) (logs []*Log, err error) { | func SearchAllLogs(keyword string) (logs []*Log, err error) { | ||||||
| 	err = DB.Where("type = ? or content LIKE ?", keyword, keyword+"%").Order("id desc").Limit(common.MaxRecentItems).Find(&logs).Error | 	err = DB.Where("type = ? or content LIKE ?", keyword, keyword+"%").Order("id desc").Limit(config.MaxRecentItems).Find(&logs).Error | ||||||
| 	return logs, err | 	return logs, err | ||||||
| } | } | ||||||
|  |  | ||||||
| func SearchUserLogs(userId int, keyword string) (logs []*Log, err error) { | func SearchUserLogs(userId int, keyword string) (logs []*Log, err error) { | ||||||
| 	err = DB.Where("user_id = ? and type = ?", userId, keyword).Order("id desc").Limit(common.MaxRecentItems).Omit("id").Find(&logs).Error | 	err = DB.Where("user_id = ? and type = ?", userId, keyword).Order("id desc").Limit(config.MaxRecentItems).Omit("id").Find(&logs).Error | ||||||
| 	return logs, err | 	return logs, err | ||||||
| } | } | ||||||
|  |  | ||||||
|   | |||||||
| @@ -7,6 +7,9 @@ import ( | |||||||
| 	"gorm.io/driver/sqlite" | 	"gorm.io/driver/sqlite" | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"os" | 	"os" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"time" | 	"time" | ||||||
| @@ -18,7 +21,7 @@ func createRootAccountIfNeed() error { | |||||||
| 	var user User | 	var user User | ||||||
| 	//if user.Status != util.UserStatusEnabled { | 	//if user.Status != util.UserStatusEnabled { | ||||||
| 	if err := DB.First(&user).Error; err != nil { | 	if err := DB.First(&user).Error; err != nil { | ||||||
| 		common.SysLog("no user exists, create a root user for you: username is root, password is 123456") | 		logger.SysLog("no user exists, create a root user for you: username is root, password is 123456") | ||||||
| 		hashedPassword, err := common.Password2Hash("123456") | 		hashedPassword, err := common.Password2Hash("123456") | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| @@ -29,7 +32,7 @@ func createRootAccountIfNeed() error { | |||||||
| 			Role:        common.RoleRootUser, | 			Role:        common.RoleRootUser, | ||||||
| 			Status:      common.UserStatusEnabled, | 			Status:      common.UserStatusEnabled, | ||||||
| 			DisplayName: "Root User", | 			DisplayName: "Root User", | ||||||
| 			AccessToken: common.GetUUID(), | 			AccessToken: helper.GetUUID(), | ||||||
| 			Quota:       100000000, | 			Quota:       100000000, | ||||||
| 		} | 		} | ||||||
| 		DB.Create(&rootUser) | 		DB.Create(&rootUser) | ||||||
| @@ -42,7 +45,7 @@ func chooseDB() (*gorm.DB, error) { | |||||||
| 		dsn := os.Getenv("SQL_DSN") | 		dsn := os.Getenv("SQL_DSN") | ||||||
| 		if strings.HasPrefix(dsn, "postgres://") { | 		if strings.HasPrefix(dsn, "postgres://") { | ||||||
| 			// Use PostgreSQL | 			// Use PostgreSQL | ||||||
| 			common.SysLog("using PostgreSQL as database") | 			logger.SysLog("using PostgreSQL as database") | ||||||
| 			common.UsingPostgreSQL = true | 			common.UsingPostgreSQL = true | ||||||
| 			return gorm.Open(postgres.New(postgres.Config{ | 			return gorm.Open(postgres.New(postgres.Config{ | ||||||
| 				DSN:                  dsn, | 				DSN:                  dsn, | ||||||
| @@ -52,13 +55,13 @@ func chooseDB() (*gorm.DB, error) { | |||||||
| 			}) | 			}) | ||||||
| 		} | 		} | ||||||
| 		// Use MySQL | 		// Use MySQL | ||||||
| 		common.SysLog("using MySQL as database") | 		logger.SysLog("using MySQL as database") | ||||||
| 		return gorm.Open(mysql.Open(dsn), &gorm.Config{ | 		return gorm.Open(mysql.Open(dsn), &gorm.Config{ | ||||||
| 			PrepareStmt: true, // precompile SQL | 			PrepareStmt: true, // precompile SQL | ||||||
| 		}) | 		}) | ||||||
| 	} | 	} | ||||||
| 	// Use SQLite | 	// Use SQLite | ||||||
| 	common.SysLog("SQL_DSN not set, using SQLite as database") | 	logger.SysLog("SQL_DSN not set, using SQLite as database") | ||||||
| 	common.UsingSQLite = true | 	common.UsingSQLite = true | ||||||
| 	config := fmt.Sprintf("?_busy_timeout=%d", common.SQLiteBusyTimeout) | 	config := fmt.Sprintf("?_busy_timeout=%d", common.SQLiteBusyTimeout) | ||||||
| 	return gorm.Open(sqlite.Open(common.SQLitePath+config), &gorm.Config{ | 	return gorm.Open(sqlite.Open(common.SQLitePath+config), &gorm.Config{ | ||||||
| @@ -69,7 +72,7 @@ func chooseDB() (*gorm.DB, error) { | |||||||
| func InitDB() (err error) { | func InitDB() (err error) { | ||||||
| 	db, err := chooseDB() | 	db, err := chooseDB() | ||||||
| 	if err == nil { | 	if err == nil { | ||||||
| 		if common.DebugEnabled { | 		if config.DebugEnabled { | ||||||
| 			db = db.Debug() | 			db = db.Debug() | ||||||
| 		} | 		} | ||||||
| 		DB = db | 		DB = db | ||||||
| @@ -77,14 +80,14 @@ func InitDB() (err error) { | |||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		sqlDB.SetMaxIdleConns(common.GetOrDefault("SQL_MAX_IDLE_CONNS", 100)) | 		sqlDB.SetMaxIdleConns(helper.GetOrDefaultEnvInt("SQL_MAX_IDLE_CONNS", 100)) | ||||||
| 		sqlDB.SetMaxOpenConns(common.GetOrDefault("SQL_MAX_OPEN_CONNS", 1000)) | 		sqlDB.SetMaxOpenConns(helper.GetOrDefaultEnvInt("SQL_MAX_OPEN_CONNS", 1000)) | ||||||
| 		sqlDB.SetConnMaxLifetime(time.Second * time.Duration(common.GetOrDefault("SQL_MAX_LIFETIME", 60))) | 		sqlDB.SetConnMaxLifetime(time.Second * time.Duration(helper.GetOrDefaultEnvInt("SQL_MAX_LIFETIME", 60))) | ||||||
|  |  | ||||||
| 		if !common.IsMasterNode { | 		if !config.IsMasterNode { | ||||||
| 			return nil | 			return nil | ||||||
| 		} | 		} | ||||||
| 		common.SysLog("database migration started") | 		logger.SysLog("database migration started") | ||||||
| 		err = db.AutoMigrate(&Channel{}) | 		err = db.AutoMigrate(&Channel{}) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| @@ -113,11 +116,11 @@ func InitDB() (err error) { | |||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		common.SysLog("database migrated") | 		logger.SysLog("database migrated") | ||||||
| 		err = createRootAccountIfNeed() | 		err = createRootAccountIfNeed() | ||||||
| 		return err | 		return err | ||||||
| 	} else { | 	} else { | ||||||
| 		common.FatalLog(err) | 		logger.FatalLog(err) | ||||||
| 	} | 	} | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
|   | |||||||
							
								
								
									
										213
									
								
								model/option.go
									
									
									
									
									
								
							
							
						
						
									
										213
									
								
								model/option.go
									
									
									
									
									
								
							| @@ -2,6 +2,8 @@ package model | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"time" | 	"time" | ||||||
| @@ -20,60 +22,56 @@ func AllOption() ([]*Option, error) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func InitOptionMap() { | func InitOptionMap() { | ||||||
| 	common.OptionMapRWMutex.Lock() | 	config.OptionMapRWMutex.Lock() | ||||||
| 	common.OptionMap = make(map[string]string) | 	config.OptionMap = make(map[string]string) | ||||||
| 	common.OptionMap["FileUploadPermission"] = strconv.Itoa(common.FileUploadPermission) | 	config.OptionMap["PasswordLoginEnabled"] = strconv.FormatBool(config.PasswordLoginEnabled) | ||||||
| 	common.OptionMap["FileDownloadPermission"] = strconv.Itoa(common.FileDownloadPermission) | 	config.OptionMap["PasswordRegisterEnabled"] = strconv.FormatBool(config.PasswordRegisterEnabled) | ||||||
| 	common.OptionMap["ImageUploadPermission"] = strconv.Itoa(common.ImageUploadPermission) | 	config.OptionMap["EmailVerificationEnabled"] = strconv.FormatBool(config.EmailVerificationEnabled) | ||||||
| 	common.OptionMap["ImageDownloadPermission"] = strconv.Itoa(common.ImageDownloadPermission) | 	config.OptionMap["GitHubOAuthEnabled"] = strconv.FormatBool(config.GitHubOAuthEnabled) | ||||||
| 	common.OptionMap["PasswordLoginEnabled"] = strconv.FormatBool(common.PasswordLoginEnabled) | 	config.OptionMap["WeChatAuthEnabled"] = strconv.FormatBool(config.WeChatAuthEnabled) | ||||||
| 	common.OptionMap["PasswordRegisterEnabled"] = strconv.FormatBool(common.PasswordRegisterEnabled) | 	config.OptionMap["TurnstileCheckEnabled"] = strconv.FormatBool(config.TurnstileCheckEnabled) | ||||||
| 	common.OptionMap["EmailVerificationEnabled"] = strconv.FormatBool(common.EmailVerificationEnabled) | 	config.OptionMap["RegisterEnabled"] = strconv.FormatBool(config.RegisterEnabled) | ||||||
| 	common.OptionMap["GitHubOAuthEnabled"] = strconv.FormatBool(common.GitHubOAuthEnabled) | 	config.OptionMap["AutomaticDisableChannelEnabled"] = strconv.FormatBool(config.AutomaticDisableChannelEnabled) | ||||||
| 	common.OptionMap["WeChatAuthEnabled"] = strconv.FormatBool(common.WeChatAuthEnabled) | 	config.OptionMap["AutomaticEnableChannelEnabled"] = strconv.FormatBool(config.AutomaticEnableChannelEnabled) | ||||||
| 	common.OptionMap["TurnstileCheckEnabled"] = strconv.FormatBool(common.TurnstileCheckEnabled) | 	config.OptionMap["ApproximateTokenEnabled"] = strconv.FormatBool(config.ApproximateTokenEnabled) | ||||||
| 	common.OptionMap["RegisterEnabled"] = strconv.FormatBool(common.RegisterEnabled) | 	config.OptionMap["LogConsumeEnabled"] = strconv.FormatBool(config.LogConsumeEnabled) | ||||||
| 	common.OptionMap["AutomaticDisableChannelEnabled"] = strconv.FormatBool(common.AutomaticDisableChannelEnabled) | 	config.OptionMap["DisplayInCurrencyEnabled"] = strconv.FormatBool(config.DisplayInCurrencyEnabled) | ||||||
| 	common.OptionMap["AutomaticEnableChannelEnabled"] = strconv.FormatBool(common.AutomaticEnableChannelEnabled) | 	config.OptionMap["DisplayTokenStatEnabled"] = strconv.FormatBool(config.DisplayTokenStatEnabled) | ||||||
| 	common.OptionMap["ApproximateTokenEnabled"] = strconv.FormatBool(common.ApproximateTokenEnabled) | 	config.OptionMap["ChannelDisableThreshold"] = strconv.FormatFloat(config.ChannelDisableThreshold, 'f', -1, 64) | ||||||
| 	common.OptionMap["LogConsumeEnabled"] = strconv.FormatBool(common.LogConsumeEnabled) | 	config.OptionMap["EmailDomainRestrictionEnabled"] = strconv.FormatBool(config.EmailDomainRestrictionEnabled) | ||||||
| 	common.OptionMap["DisplayInCurrencyEnabled"] = strconv.FormatBool(common.DisplayInCurrencyEnabled) | 	config.OptionMap["EmailDomainWhitelist"] = strings.Join(config.EmailDomainWhitelist, ",") | ||||||
| 	common.OptionMap["DisplayTokenStatEnabled"] = strconv.FormatBool(common.DisplayTokenStatEnabled) | 	config.OptionMap["SMTPServer"] = "" | ||||||
| 	common.OptionMap["ChannelDisableThreshold"] = strconv.FormatFloat(common.ChannelDisableThreshold, 'f', -1, 64) | 	config.OptionMap["SMTPFrom"] = "" | ||||||
| 	common.OptionMap["EmailDomainRestrictionEnabled"] = strconv.FormatBool(common.EmailDomainRestrictionEnabled) | 	config.OptionMap["SMTPPort"] = strconv.Itoa(config.SMTPPort) | ||||||
| 	common.OptionMap["EmailDomainWhitelist"] = strings.Join(common.EmailDomainWhitelist, ",") | 	config.OptionMap["SMTPAccount"] = "" | ||||||
| 	common.OptionMap["SMTPServer"] = "" | 	config.OptionMap["SMTPToken"] = "" | ||||||
| 	common.OptionMap["SMTPFrom"] = "" | 	config.OptionMap["Notice"] = "" | ||||||
| 	common.OptionMap["SMTPPort"] = strconv.Itoa(common.SMTPPort) | 	config.OptionMap["About"] = "" | ||||||
| 	common.OptionMap["SMTPAccount"] = "" | 	config.OptionMap["HomePageContent"] = "" | ||||||
| 	common.OptionMap["SMTPToken"] = "" | 	config.OptionMap["Footer"] = config.Footer | ||||||
| 	common.OptionMap["Notice"] = "" | 	config.OptionMap["SystemName"] = config.SystemName | ||||||
| 	common.OptionMap["About"] = "" | 	config.OptionMap["Logo"] = config.Logo | ||||||
| 	common.OptionMap["HomePageContent"] = "" | 	config.OptionMap["ServerAddress"] = "" | ||||||
| 	common.OptionMap["Footer"] = common.Footer | 	config.OptionMap["GitHubClientId"] = "" | ||||||
| 	common.OptionMap["SystemName"] = common.SystemName | 	config.OptionMap["GitHubClientSecret"] = "" | ||||||
| 	common.OptionMap["Logo"] = common.Logo | 	config.OptionMap["WeChatServerAddress"] = "" | ||||||
| 	common.OptionMap["ServerAddress"] = "" | 	config.OptionMap["WeChatServerToken"] = "" | ||||||
| 	common.OptionMap["GitHubClientId"] = "" | 	config.OptionMap["WeChatAccountQRCodeImageURL"] = "" | ||||||
| 	common.OptionMap["GitHubClientSecret"] = "" | 	config.OptionMap["TurnstileSiteKey"] = "" | ||||||
| 	common.OptionMap["WeChatServerAddress"] = "" | 	config.OptionMap["TurnstileSecretKey"] = "" | ||||||
| 	common.OptionMap["WeChatServerToken"] = "" | 	config.OptionMap["QuotaForNewUser"] = strconv.Itoa(config.QuotaForNewUser) | ||||||
| 	common.OptionMap["WeChatAccountQRCodeImageURL"] = "" | 	config.OptionMap["QuotaForInviter"] = strconv.Itoa(config.QuotaForInviter) | ||||||
| 	common.OptionMap["TurnstileSiteKey"] = "" | 	config.OptionMap["QuotaForInvitee"] = strconv.Itoa(config.QuotaForInvitee) | ||||||
| 	common.OptionMap["TurnstileSecretKey"] = "" | 	config.OptionMap["QuotaRemindThreshold"] = strconv.Itoa(config.QuotaRemindThreshold) | ||||||
| 	common.OptionMap["QuotaForNewUser"] = strconv.Itoa(common.QuotaForNewUser) | 	config.OptionMap["PreConsumedQuota"] = strconv.Itoa(config.PreConsumedQuota) | ||||||
| 	common.OptionMap["QuotaForInviter"] = strconv.Itoa(common.QuotaForInviter) | 	config.OptionMap["ModelRatio"] = common.ModelRatio2JSONString() | ||||||
| 	common.OptionMap["QuotaForInvitee"] = strconv.Itoa(common.QuotaForInvitee) | 	config.OptionMap["GroupRatio"] = common.GroupRatio2JSONString() | ||||||
| 	common.OptionMap["QuotaRemindThreshold"] = strconv.Itoa(common.QuotaRemindThreshold) | 	config.OptionMap["TopUpLink"] = config.TopUpLink | ||||||
| 	common.OptionMap["PreConsumedQuota"] = strconv.Itoa(common.PreConsumedQuota) | 	config.OptionMap["ChatLink"] = config.ChatLink | ||||||
| 	common.OptionMap["ModelRatio"] = common.ModelRatio2JSONString() | 	config.OptionMap["QuotaPerUnit"] = strconv.FormatFloat(config.QuotaPerUnit, 'f', -1, 64) | ||||||
| 	common.OptionMap["GroupRatio"] = common.GroupRatio2JSONString() | 	config.OptionMap["RetryTimes"] = strconv.Itoa(config.RetryTimes) | ||||||
| 	common.OptionMap["TopUpLink"] = common.TopUpLink | 	config.OptionMap["Theme"] = config.Theme | ||||||
| 	common.OptionMap["ChatLink"] = common.ChatLink | 	config.OptionMapRWMutex.Unlock() | ||||||
| 	common.OptionMap["QuotaPerUnit"] = strconv.FormatFloat(common.QuotaPerUnit, 'f', -1, 64) |  | ||||||
| 	common.OptionMap["RetryTimes"] = strconv.Itoa(common.RetryTimes) |  | ||||||
| 	common.OptionMap["Theme"] = common.Theme |  | ||||||
| 	common.OptionMapRWMutex.Unlock() |  | ||||||
| 	loadOptionsFromDatabase() | 	loadOptionsFromDatabase() | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -82,7 +80,7 @@ func loadOptionsFromDatabase() { | |||||||
| 	for _, option := range options { | 	for _, option := range options { | ||||||
| 		err := updateOptionMap(option.Key, option.Value) | 		err := updateOptionMap(option.Key, option.Value) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("failed to update option map: " + err.Error()) | 			logger.SysError("failed to update option map: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| @@ -90,7 +88,7 @@ func loadOptionsFromDatabase() { | |||||||
| func SyncOptions(frequency int) { | func SyncOptions(frequency int) { | ||||||
| 	for { | 	for { | ||||||
| 		time.Sleep(time.Duration(frequency) * time.Second) | 		time.Sleep(time.Duration(frequency) * time.Second) | ||||||
| 		common.SysLog("syncing options from database") | 		logger.SysLog("syncing options from database") | ||||||
| 		loadOptionsFromDatabase() | 		loadOptionsFromDatabase() | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| @@ -112,117 +110,104 @@ func UpdateOption(key string, value string) error { | |||||||
| } | } | ||||||
|  |  | ||||||
| func updateOptionMap(key string, value string) (err error) { | func updateOptionMap(key string, value string) (err error) { | ||||||
| 	common.OptionMapRWMutex.Lock() | 	config.OptionMapRWMutex.Lock() | ||||||
| 	defer common.OptionMapRWMutex.Unlock() | 	defer config.OptionMapRWMutex.Unlock() | ||||||
| 	common.OptionMap[key] = value | 	config.OptionMap[key] = value | ||||||
| 	if strings.HasSuffix(key, "Permission") { |  | ||||||
| 		intValue, _ := strconv.Atoi(value) |  | ||||||
| 		switch key { |  | ||||||
| 		case "FileUploadPermission": |  | ||||||
| 			common.FileUploadPermission = intValue |  | ||||||
| 		case "FileDownloadPermission": |  | ||||||
| 			common.FileDownloadPermission = intValue |  | ||||||
| 		case "ImageUploadPermission": |  | ||||||
| 			common.ImageUploadPermission = intValue |  | ||||||
| 		case "ImageDownloadPermission": |  | ||||||
| 			common.ImageDownloadPermission = intValue |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	if strings.HasSuffix(key, "Enabled") { | 	if strings.HasSuffix(key, "Enabled") { | ||||||
| 		boolValue := value == "true" | 		boolValue := value == "true" | ||||||
| 		switch key { | 		switch key { | ||||||
| 		case "PasswordRegisterEnabled": | 		case "PasswordRegisterEnabled": | ||||||
| 			common.PasswordRegisterEnabled = boolValue | 			config.PasswordRegisterEnabled = boolValue | ||||||
| 		case "PasswordLoginEnabled": | 		case "PasswordLoginEnabled": | ||||||
| 			common.PasswordLoginEnabled = boolValue | 			config.PasswordLoginEnabled = boolValue | ||||||
| 		case "EmailVerificationEnabled": | 		case "EmailVerificationEnabled": | ||||||
| 			common.EmailVerificationEnabled = boolValue | 			config.EmailVerificationEnabled = boolValue | ||||||
| 		case "GitHubOAuthEnabled": | 		case "GitHubOAuthEnabled": | ||||||
| 			common.GitHubOAuthEnabled = boolValue | 			config.GitHubOAuthEnabled = boolValue | ||||||
| 		case "WeChatAuthEnabled": | 		case "WeChatAuthEnabled": | ||||||
| 			common.WeChatAuthEnabled = boolValue | 			config.WeChatAuthEnabled = boolValue | ||||||
| 		case "TurnstileCheckEnabled": | 		case "TurnstileCheckEnabled": | ||||||
| 			common.TurnstileCheckEnabled = boolValue | 			config.TurnstileCheckEnabled = boolValue | ||||||
| 		case "RegisterEnabled": | 		case "RegisterEnabled": | ||||||
| 			common.RegisterEnabled = boolValue | 			config.RegisterEnabled = boolValue | ||||||
| 		case "EmailDomainRestrictionEnabled": | 		case "EmailDomainRestrictionEnabled": | ||||||
| 			common.EmailDomainRestrictionEnabled = boolValue | 			config.EmailDomainRestrictionEnabled = boolValue | ||||||
| 		case "AutomaticDisableChannelEnabled": | 		case "AutomaticDisableChannelEnabled": | ||||||
| 			common.AutomaticDisableChannelEnabled = boolValue | 			config.AutomaticDisableChannelEnabled = boolValue | ||||||
| 		case "AutomaticEnableChannelEnabled": | 		case "AutomaticEnableChannelEnabled": | ||||||
| 			common.AutomaticEnableChannelEnabled = boolValue | 			config.AutomaticEnableChannelEnabled = boolValue | ||||||
| 		case "ApproximateTokenEnabled": | 		case "ApproximateTokenEnabled": | ||||||
| 			common.ApproximateTokenEnabled = boolValue | 			config.ApproximateTokenEnabled = boolValue | ||||||
| 		case "LogConsumeEnabled": | 		case "LogConsumeEnabled": | ||||||
| 			common.LogConsumeEnabled = boolValue | 			config.LogConsumeEnabled = boolValue | ||||||
| 		case "DisplayInCurrencyEnabled": | 		case "DisplayInCurrencyEnabled": | ||||||
| 			common.DisplayInCurrencyEnabled = boolValue | 			config.DisplayInCurrencyEnabled = boolValue | ||||||
| 		case "DisplayTokenStatEnabled": | 		case "DisplayTokenStatEnabled": | ||||||
| 			common.DisplayTokenStatEnabled = boolValue | 			config.DisplayTokenStatEnabled = boolValue | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	switch key { | 	switch key { | ||||||
| 	case "EmailDomainWhitelist": | 	case "EmailDomainWhitelist": | ||||||
| 		common.EmailDomainWhitelist = strings.Split(value, ",") | 		config.EmailDomainWhitelist = strings.Split(value, ",") | ||||||
| 	case "SMTPServer": | 	case "SMTPServer": | ||||||
| 		common.SMTPServer = value | 		config.SMTPServer = value | ||||||
| 	case "SMTPPort": | 	case "SMTPPort": | ||||||
| 		intValue, _ := strconv.Atoi(value) | 		intValue, _ := strconv.Atoi(value) | ||||||
| 		common.SMTPPort = intValue | 		config.SMTPPort = intValue | ||||||
| 	case "SMTPAccount": | 	case "SMTPAccount": | ||||||
| 		common.SMTPAccount = value | 		config.SMTPAccount = value | ||||||
| 	case "SMTPFrom": | 	case "SMTPFrom": | ||||||
| 		common.SMTPFrom = value | 		config.SMTPFrom = value | ||||||
| 	case "SMTPToken": | 	case "SMTPToken": | ||||||
| 		common.SMTPToken = value | 		config.SMTPToken = value | ||||||
| 	case "ServerAddress": | 	case "ServerAddress": | ||||||
| 		common.ServerAddress = value | 		config.ServerAddress = value | ||||||
| 	case "GitHubClientId": | 	case "GitHubClientId": | ||||||
| 		common.GitHubClientId = value | 		config.GitHubClientId = value | ||||||
| 	case "GitHubClientSecret": | 	case "GitHubClientSecret": | ||||||
| 		common.GitHubClientSecret = value | 		config.GitHubClientSecret = value | ||||||
| 	case "Footer": | 	case "Footer": | ||||||
| 		common.Footer = value | 		config.Footer = value | ||||||
| 	case "SystemName": | 	case "SystemName": | ||||||
| 		common.SystemName = value | 		config.SystemName = value | ||||||
| 	case "Logo": | 	case "Logo": | ||||||
| 		common.Logo = value | 		config.Logo = value | ||||||
| 	case "WeChatServerAddress": | 	case "WeChatServerAddress": | ||||||
| 		common.WeChatServerAddress = value | 		config.WeChatServerAddress = value | ||||||
| 	case "WeChatServerToken": | 	case "WeChatServerToken": | ||||||
| 		common.WeChatServerToken = value | 		config.WeChatServerToken = value | ||||||
| 	case "WeChatAccountQRCodeImageURL": | 	case "WeChatAccountQRCodeImageURL": | ||||||
| 		common.WeChatAccountQRCodeImageURL = value | 		config.WeChatAccountQRCodeImageURL = value | ||||||
| 	case "TurnstileSiteKey": | 	case "TurnstileSiteKey": | ||||||
| 		common.TurnstileSiteKey = value | 		config.TurnstileSiteKey = value | ||||||
| 	case "TurnstileSecretKey": | 	case "TurnstileSecretKey": | ||||||
| 		common.TurnstileSecretKey = value | 		config.TurnstileSecretKey = value | ||||||
| 	case "QuotaForNewUser": | 	case "QuotaForNewUser": | ||||||
| 		common.QuotaForNewUser, _ = strconv.Atoi(value) | 		config.QuotaForNewUser, _ = strconv.Atoi(value) | ||||||
| 	case "QuotaForInviter": | 	case "QuotaForInviter": | ||||||
| 		common.QuotaForInviter, _ = strconv.Atoi(value) | 		config.QuotaForInviter, _ = strconv.Atoi(value) | ||||||
| 	case "QuotaForInvitee": | 	case "QuotaForInvitee": | ||||||
| 		common.QuotaForInvitee, _ = strconv.Atoi(value) | 		config.QuotaForInvitee, _ = strconv.Atoi(value) | ||||||
| 	case "QuotaRemindThreshold": | 	case "QuotaRemindThreshold": | ||||||
| 		common.QuotaRemindThreshold, _ = strconv.Atoi(value) | 		config.QuotaRemindThreshold, _ = strconv.Atoi(value) | ||||||
| 	case "PreConsumedQuota": | 	case "PreConsumedQuota": | ||||||
| 		common.PreConsumedQuota, _ = strconv.Atoi(value) | 		config.PreConsumedQuota, _ = strconv.Atoi(value) | ||||||
| 	case "RetryTimes": | 	case "RetryTimes": | ||||||
| 		common.RetryTimes, _ = strconv.Atoi(value) | 		config.RetryTimes, _ = strconv.Atoi(value) | ||||||
| 	case "ModelRatio": | 	case "ModelRatio": | ||||||
| 		err = common.UpdateModelRatioByJSONString(value) | 		err = common.UpdateModelRatioByJSONString(value) | ||||||
| 	case "GroupRatio": | 	case "GroupRatio": | ||||||
| 		err = common.UpdateGroupRatioByJSONString(value) | 		err = common.UpdateGroupRatioByJSONString(value) | ||||||
| 	case "TopUpLink": | 	case "TopUpLink": | ||||||
| 		common.TopUpLink = value | 		config.TopUpLink = value | ||||||
| 	case "ChatLink": | 	case "ChatLink": | ||||||
| 		common.ChatLink = value | 		config.ChatLink = value | ||||||
| 	case "ChannelDisableThreshold": | 	case "ChannelDisableThreshold": | ||||||
| 		common.ChannelDisableThreshold, _ = strconv.ParseFloat(value, 64) | 		config.ChannelDisableThreshold, _ = strconv.ParseFloat(value, 64) | ||||||
| 	case "QuotaPerUnit": | 	case "QuotaPerUnit": | ||||||
| 		common.QuotaPerUnit, _ = strconv.ParseFloat(value, 64) | 		config.QuotaPerUnit, _ = strconv.ParseFloat(value, 64) | ||||||
| 	case "Theme": | 	case "Theme": | ||||||
| 		common.Theme = value | 		config.Theme = value | ||||||
| 	} | 	} | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
|   | |||||||
| @@ -5,6 +5,7 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/helper" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| type Redemption struct { | type Redemption struct { | ||||||
| @@ -67,7 +68,7 @@ func Redeem(key string, userId int) (quota int, err error) { | |||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		redemption.RedeemedTime = common.GetTimestamp() | 		redemption.RedeemedTime = helper.GetTimestamp() | ||||||
| 		redemption.Status = common.RedemptionCodeStatusUsed | 		redemption.Status = common.RedemptionCodeStatusUsed | ||||||
| 		err = tx.Save(redemption).Error | 		err = tx.Save(redemption).Error | ||||||
| 		return err | 		return err | ||||||
|   | |||||||
| @@ -5,6 +5,9 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| type Token struct { | type Token struct { | ||||||
| @@ -39,7 +42,7 @@ func ValidateUserToken(key string) (token *Token, err error) { | |||||||
| 	} | 	} | ||||||
| 	token, err = CacheGetTokenByKey(key) | 	token, err = CacheGetTokenByKey(key) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("CacheGetTokenByKey failed: " + err.Error()) | 		logger.SysError("CacheGetTokenByKey failed: " + err.Error()) | ||||||
| 		if errors.Is(err, gorm.ErrRecordNotFound) { | 		if errors.Is(err, gorm.ErrRecordNotFound) { | ||||||
| 			return nil, errors.New("无效的令牌") | 			return nil, errors.New("无效的令牌") | ||||||
| 		} | 		} | ||||||
| @@ -53,12 +56,12 @@ func ValidateUserToken(key string) (token *Token, err error) { | |||||||
| 	if token.Status != common.TokenStatusEnabled { | 	if token.Status != common.TokenStatusEnabled { | ||||||
| 		return nil, errors.New("该令牌状态不可用") | 		return nil, errors.New("该令牌状态不可用") | ||||||
| 	} | 	} | ||||||
| 	if token.ExpiredTime != -1 && token.ExpiredTime < common.GetTimestamp() { | 	if token.ExpiredTime != -1 && token.ExpiredTime < helper.GetTimestamp() { | ||||||
| 		if !common.RedisEnabled { | 		if !common.RedisEnabled { | ||||||
| 			token.Status = common.TokenStatusExpired | 			token.Status = common.TokenStatusExpired | ||||||
| 			err := token.SelectUpdate() | 			err := token.SelectUpdate() | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("failed to update token status" + err.Error()) | 				logger.SysError("failed to update token status" + err.Error()) | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 		return nil, errors.New("该令牌已过期") | 		return nil, errors.New("该令牌已过期") | ||||||
| @@ -69,7 +72,7 @@ func ValidateUserToken(key string) (token *Token, err error) { | |||||||
| 			token.Status = common.TokenStatusExhausted | 			token.Status = common.TokenStatusExhausted | ||||||
| 			err := token.SelectUpdate() | 			err := token.SelectUpdate() | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("failed to update token status" + err.Error()) | 				logger.SysError("failed to update token status" + err.Error()) | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 		return nil, errors.New("该令牌额度已用尽") | 		return nil, errors.New("该令牌额度已用尽") | ||||||
| @@ -138,7 +141,7 @@ func IncreaseTokenQuota(id int, quota int) (err error) { | |||||||
| 	if quota < 0 { | 	if quota < 0 { | ||||||
| 		return errors.New("quota 不能为负数!") | 		return errors.New("quota 不能为负数!") | ||||||
| 	} | 	} | ||||||
| 	if common.BatchUpdateEnabled { | 	if config.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeTokenQuota, id, quota) | 		addNewRecord(BatchUpdateTypeTokenQuota, id, quota) | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| @@ -150,7 +153,7 @@ func increaseTokenQuota(id int, quota int) (err error) { | |||||||
| 		map[string]interface{}{ | 		map[string]interface{}{ | ||||||
| 			"remain_quota":  gorm.Expr("remain_quota + ?", quota), | 			"remain_quota":  gorm.Expr("remain_quota + ?", quota), | ||||||
| 			"used_quota":    gorm.Expr("used_quota - ?", quota), | 			"used_quota":    gorm.Expr("used_quota - ?", quota), | ||||||
| 			"accessed_time": common.GetTimestamp(), | 			"accessed_time": helper.GetTimestamp(), | ||||||
| 		}, | 		}, | ||||||
| 	).Error | 	).Error | ||||||
| 	return err | 	return err | ||||||
| @@ -160,7 +163,7 @@ func DecreaseTokenQuota(id int, quota int) (err error) { | |||||||
| 	if quota < 0 { | 	if quota < 0 { | ||||||
| 		return errors.New("quota 不能为负数!") | 		return errors.New("quota 不能为负数!") | ||||||
| 	} | 	} | ||||||
| 	if common.BatchUpdateEnabled { | 	if config.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeTokenQuota, id, -quota) | 		addNewRecord(BatchUpdateTypeTokenQuota, id, -quota) | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| @@ -172,7 +175,7 @@ func decreaseTokenQuota(id int, quota int) (err error) { | |||||||
| 		map[string]interface{}{ | 		map[string]interface{}{ | ||||||
| 			"remain_quota":  gorm.Expr("remain_quota - ?", quota), | 			"remain_quota":  gorm.Expr("remain_quota - ?", quota), | ||||||
| 			"used_quota":    gorm.Expr("used_quota + ?", quota), | 			"used_quota":    gorm.Expr("used_quota + ?", quota), | ||||||
| 			"accessed_time": common.GetTimestamp(), | 			"accessed_time": helper.GetTimestamp(), | ||||||
| 		}, | 		}, | ||||||
| 	).Error | 	).Error | ||||||
| 	return err | 	return err | ||||||
| @@ -196,24 +199,24 @@ func PreConsumeTokenQuota(tokenId int, quota int) (err error) { | |||||||
| 	if userQuota < quota { | 	if userQuota < quota { | ||||||
| 		return errors.New("用户额度不足") | 		return errors.New("用户额度不足") | ||||||
| 	} | 	} | ||||||
| 	quotaTooLow := userQuota >= common.QuotaRemindThreshold && userQuota-quota < common.QuotaRemindThreshold | 	quotaTooLow := userQuota >= config.QuotaRemindThreshold && userQuota-quota < config.QuotaRemindThreshold | ||||||
| 	noMoreQuota := userQuota-quota <= 0 | 	noMoreQuota := userQuota-quota <= 0 | ||||||
| 	if quotaTooLow || noMoreQuota { | 	if quotaTooLow || noMoreQuota { | ||||||
| 		go func() { | 		go func() { | ||||||
| 			email, err := GetUserEmail(token.UserId) | 			email, err := GetUserEmail(token.UserId) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("failed to fetch user email: " + err.Error()) | 				logger.SysError("failed to fetch user email: " + err.Error()) | ||||||
| 			} | 			} | ||||||
| 			prompt := "您的额度即将用尽" | 			prompt := "您的额度即将用尽" | ||||||
| 			if noMoreQuota { | 			if noMoreQuota { | ||||||
| 				prompt = "您的额度已用尽" | 				prompt = "您的额度已用尽" | ||||||
| 			} | 			} | ||||||
| 			if email != "" { | 			if email != "" { | ||||||
| 				topUpLink := fmt.Sprintf("%s/topup", common.ServerAddress) | 				topUpLink := fmt.Sprintf("%s/topup", config.ServerAddress) | ||||||
| 				err = common.SendEmail(prompt, email, | 				err = common.SendEmail(prompt, email, | ||||||
| 					fmt.Sprintf("%s,当前剩余额度为 %d,为了不影响您的使用,请及时充值。<br/>充值链接:<a href='%s'>%s</a>", prompt, userQuota, topUpLink, topUpLink)) | 					fmt.Sprintf("%s,当前剩余额度为 %d,为了不影响您的使用,请及时充值。<br/>充值链接:<a href='%s'>%s</a>", prompt, userQuota, topUpLink, topUpLink)) | ||||||
| 				if err != nil { | 				if err != nil { | ||||||
| 					common.SysError("failed to send email" + err.Error()) | 					logger.SysError("failed to send email" + err.Error()) | ||||||
| 				} | 				} | ||||||
| 			} | 			} | ||||||
| 		}() | 		}() | ||||||
|   | |||||||
| @@ -5,6 +5,9 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -89,24 +92,24 @@ func (user *User) Insert(inviterId int) error { | |||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	user.Quota = common.QuotaForNewUser | 	user.Quota = config.QuotaForNewUser | ||||||
| 	user.AccessToken = common.GetUUID() | 	user.AccessToken = helper.GetUUID() | ||||||
| 	user.AffCode = common.GetRandomString(4) | 	user.AffCode = helper.GetRandomString(4) | ||||||
| 	result := DB.Create(user) | 	result := DB.Create(user) | ||||||
| 	if result.Error != nil { | 	if result.Error != nil { | ||||||
| 		return result.Error | 		return result.Error | ||||||
| 	} | 	} | ||||||
| 	if common.QuotaForNewUser > 0 { | 	if config.QuotaForNewUser > 0 { | ||||||
| 		RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", common.LogQuota(common.QuotaForNewUser))) | 		RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", common.LogQuota(config.QuotaForNewUser))) | ||||||
| 	} | 	} | ||||||
| 	if inviterId != 0 { | 	if inviterId != 0 { | ||||||
| 		if common.QuotaForInvitee > 0 { | 		if config.QuotaForInvitee > 0 { | ||||||
| 			_ = IncreaseUserQuota(user.Id, common.QuotaForInvitee) | 			_ = IncreaseUserQuota(user.Id, config.QuotaForInvitee) | ||||||
| 			RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("使用邀请码赠送 %s", common.LogQuota(common.QuotaForInvitee))) | 			RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("使用邀请码赠送 %s", common.LogQuota(config.QuotaForInvitee))) | ||||||
| 		} | 		} | ||||||
| 		if common.QuotaForInviter > 0 { | 		if config.QuotaForInviter > 0 { | ||||||
| 			_ = IncreaseUserQuota(inviterId, common.QuotaForInviter) | 			_ = IncreaseUserQuota(inviterId, config.QuotaForInviter) | ||||||
| 			RecordLog(inviterId, LogTypeSystem, fmt.Sprintf("邀请用户赠送 %s", common.LogQuota(common.QuotaForInviter))) | 			RecordLog(inviterId, LogTypeSystem, fmt.Sprintf("邀请用户赠送 %s", common.LogQuota(config.QuotaForInviter))) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	return nil | 	return nil | ||||||
| @@ -232,7 +235,7 @@ func IsAdmin(userId int) bool { | |||||||
| 	var user User | 	var user User | ||||||
| 	err := DB.Where("id = ?", userId).Select("role").Find(&user).Error | 	err := DB.Where("id = ?", userId).Select("role").Find(&user).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("no such user " + err.Error()) | 		logger.SysError("no such user " + err.Error()) | ||||||
| 		return false | 		return false | ||||||
| 	} | 	} | ||||||
| 	return user.Role >= common.RoleAdminUser | 	return user.Role >= common.RoleAdminUser | ||||||
| @@ -291,7 +294,7 @@ func IncreaseUserQuota(id int, quota int) (err error) { | |||||||
| 	if quota < 0 { | 	if quota < 0 { | ||||||
| 		return errors.New("quota 不能为负数!") | 		return errors.New("quota 不能为负数!") | ||||||
| 	} | 	} | ||||||
| 	if common.BatchUpdateEnabled { | 	if config.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeUserQuota, id, quota) | 		addNewRecord(BatchUpdateTypeUserQuota, id, quota) | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| @@ -307,7 +310,7 @@ func DecreaseUserQuota(id int, quota int) (err error) { | |||||||
| 	if quota < 0 { | 	if quota < 0 { | ||||||
| 		return errors.New("quota 不能为负数!") | 		return errors.New("quota 不能为负数!") | ||||||
| 	} | 	} | ||||||
| 	if common.BatchUpdateEnabled { | 	if config.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeUserQuota, id, -quota) | 		addNewRecord(BatchUpdateTypeUserQuota, id, -quota) | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| @@ -325,7 +328,7 @@ func GetRootUserEmail() (email string) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func UpdateUserUsedQuotaAndRequestCount(id int, quota int) { | func UpdateUserUsedQuotaAndRequestCount(id int, quota int) { | ||||||
| 	if common.BatchUpdateEnabled { | 	if config.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeUsedQuota, id, quota) | 		addNewRecord(BatchUpdateTypeUsedQuota, id, quota) | ||||||
| 		addNewRecord(BatchUpdateTypeRequestCount, id, 1) | 		addNewRecord(BatchUpdateTypeRequestCount, id, 1) | ||||||
| 		return | 		return | ||||||
| @@ -341,7 +344,7 @@ func updateUserUsedQuotaAndRequestCount(id int, quota int, count int) { | |||||||
| 		}, | 		}, | ||||||
| 	).Error | 	).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("failed to update user used quota and request count: " + err.Error()) | 		logger.SysError("failed to update user used quota and request count: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -352,14 +355,14 @@ func updateUserUsedQuota(id int, quota int) { | |||||||
| 		}, | 		}, | ||||||
| 	).Error | 	).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("failed to update user used quota: " + err.Error()) | 		logger.SysError("failed to update user used quota: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func updateUserRequestCount(id int, count int) { | func updateUserRequestCount(id int, count int) { | ||||||
| 	err := DB.Model(&User{}).Where("id = ?", id).Update("request_count", gorm.Expr("request_count + ?", count)).Error | 	err := DB.Model(&User{}).Where("id = ?", id).Update("request_count", gorm.Expr("request_count + ?", count)).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("failed to update user request count: " + err.Error()) | 		logger.SysError("failed to update user request count: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
|   | |||||||
| @@ -1,7 +1,8 @@ | |||||||
| package model | package model | ||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"sync" | 	"sync" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
| @@ -28,7 +29,7 @@ func init() { | |||||||
| func InitBatchUpdater() { | func InitBatchUpdater() { | ||||||
| 	go func() { | 	go func() { | ||||||
| 		for { | 		for { | ||||||
| 			time.Sleep(time.Duration(common.BatchUpdateInterval) * time.Second) | 			time.Sleep(time.Duration(config.BatchUpdateInterval) * time.Second) | ||||||
| 			batchUpdate() | 			batchUpdate() | ||||||
| 		} | 		} | ||||||
| 	}() | 	}() | ||||||
| @@ -45,7 +46,7 @@ func addNewRecord(type_ int, id int, value int) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func batchUpdate() { | func batchUpdate() { | ||||||
| 	common.SysLog("batch update started") | 	logger.SysLog("batch update started") | ||||||
| 	for i := 0; i < BatchUpdateTypeCount; i++ { | 	for i := 0; i < BatchUpdateTypeCount; i++ { | ||||||
| 		batchUpdateLocks[i].Lock() | 		batchUpdateLocks[i].Lock() | ||||||
| 		store := batchUpdateStores[i] | 		store := batchUpdateStores[i] | ||||||
| @@ -57,12 +58,12 @@ func batchUpdate() { | |||||||
| 			case BatchUpdateTypeUserQuota: | 			case BatchUpdateTypeUserQuota: | ||||||
| 				err := increaseUserQuota(key, value) | 				err := increaseUserQuota(key, value) | ||||||
| 				if err != nil { | 				if err != nil { | ||||||
| 					common.SysError("failed to batch update user quota: " + err.Error()) | 					logger.SysError("failed to batch update user quota: " + err.Error()) | ||||||
| 				} | 				} | ||||||
| 			case BatchUpdateTypeTokenQuota: | 			case BatchUpdateTypeTokenQuota: | ||||||
| 				err := increaseTokenQuota(key, value) | 				err := increaseTokenQuota(key, value) | ||||||
| 				if err != nil { | 				if err != nil { | ||||||
| 					common.SysError("failed to batch update token quota: " + err.Error()) | 					logger.SysError("failed to batch update token quota: " + err.Error()) | ||||||
| 				} | 				} | ||||||
| 			case BatchUpdateTypeUsedQuota: | 			case BatchUpdateTypeUsedQuota: | ||||||
| 				updateUserUsedQuota(key, value) | 				updateUserUsedQuota(key, value) | ||||||
| @@ -73,5 +74,5 @@ func batchUpdate() { | |||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	common.SysLog("batch update finished") | 	logger.SysLog("batch update finished") | ||||||
| } | } | ||||||
|   | |||||||
							
								
								
									
										22
									
								
								relay/channel/aiproxy/adaptor.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										22
									
								
								relay/channel/aiproxy/adaptor.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,22 @@ | |||||||
|  | package aiproxy | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor struct { | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) Auth(c *gin.Context) error { | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) ConvertRequest(request *openai.GeneralOpenAIRequest) (any, error) { | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, *openai.Usage, error) { | ||||||
|  | 	return nil, nil, nil | ||||||
|  | } | ||||||
| @@ -8,6 +8,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| @@ -50,9 +52,9 @@ func responseAIProxyLibrary2OpenAI(response *LibraryResponse) *openai.TextRespon | |||||||
| 		FinishReason: "stop", | 		FinishReason: "stop", | ||||||
| 	} | 	} | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := openai.TextResponse{ | ||||||
| 		Id:      common.GetUUID(), | 		Id:      helper.GetUUID(), | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Choices: []openai.TextResponseChoice{choice}, | 		Choices: []openai.TextResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| 	return &fullTextResponse | 	return &fullTextResponse | ||||||
| @@ -63,9 +65,9 @@ func documentsAIProxyLibrary(documents []LibraryDocument) *openai.ChatCompletion | |||||||
| 	choice.Delta.Content = aiProxyDocuments2Markdown(documents) | 	choice.Delta.Content = aiProxyDocuments2Markdown(documents) | ||||||
| 	choice.FinishReason = &constant.StopFinishReason | 	choice.FinishReason = &constant.StopFinishReason | ||||||
| 	return &openai.ChatCompletionsStreamResponse{ | 	return &openai.ChatCompletionsStreamResponse{ | ||||||
| 		Id:      common.GetUUID(), | 		Id:      helper.GetUUID(), | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Model:   "", | 		Model:   "", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| @@ -75,9 +77,9 @@ func streamResponseAIProxyLibrary2OpenAI(response *LibraryStreamResponse) *opena | |||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice openai.ChatCompletionsStreamResponseChoice | ||||||
| 	choice.Delta.Content = response.Content | 	choice.Delta.Content = response.Content | ||||||
| 	return &openai.ChatCompletionsStreamResponse{ | 	return &openai.ChatCompletionsStreamResponse{ | ||||||
| 		Id:      common.GetUUID(), | 		Id:      helper.GetUUID(), | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Model:   response.Model, | 		Model:   response.Model, | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| @@ -122,7 +124,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			var AIProxyLibraryResponse LibraryStreamResponse | 			var AIProxyLibraryResponse LibraryStreamResponse | ||||||
| 			err := json.Unmarshal([]byte(data), &AIProxyLibraryResponse) | 			err := json.Unmarshal([]byte(data), &AIProxyLibraryResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error unmarshalling stream response: " + err.Error()) | 				logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			if len(AIProxyLibraryResponse.Documents) != 0 { | 			if len(AIProxyLibraryResponse.Documents) != 0 { | ||||||
| @@ -131,7 +133,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			response := streamResponseAIProxyLibrary2OpenAI(&AIProxyLibraryResponse) | 			response := streamResponseAIProxyLibrary2OpenAI(&AIProxyLibraryResponse) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | ||||||
| @@ -140,7 +142,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			response := documentsAIProxyLibrary(documents) | 			response := documentsAIProxyLibrary(documents) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | ||||||
|   | |||||||
							
								
								
									
										22
									
								
								relay/channel/ali/adaptor.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										22
									
								
								relay/channel/ali/adaptor.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,22 @@ | |||||||
|  | package ali | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor struct { | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) Auth(c *gin.Context) error { | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) ConvertRequest(request *openai.GeneralOpenAIRequest) (any, error) { | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, *openai.Usage, error) { | ||||||
|  | 	return nil, nil, nil | ||||||
|  | } | ||||||
| @@ -7,6 +7,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
| @@ -118,7 +120,7 @@ func responseAli2OpenAI(response *ChatResponse) *openai.TextResponse { | |||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := openai.TextResponse{ | ||||||
| 		Id:      response.RequestId, | 		Id:      response.RequestId, | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Choices: []openai.TextResponseChoice{choice}, | 		Choices: []openai.TextResponseChoice{choice}, | ||||||
| 		Usage: openai.Usage{ | 		Usage: openai.Usage{ | ||||||
| 			PromptTokens:     response.Usage.InputTokens, | 			PromptTokens:     response.Usage.InputTokens, | ||||||
| @@ -139,7 +141,7 @@ func streamResponseAli2OpenAI(aliResponse *ChatResponse) *openai.ChatCompletions | |||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := openai.ChatCompletionsStreamResponse{ | ||||||
| 		Id:      aliResponse.RequestId, | 		Id:      aliResponse.RequestId, | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Model:   "qwen", | 		Model:   "qwen", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| @@ -185,7 +187,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			var aliResponse ChatResponse | 			var aliResponse ChatResponse | ||||||
| 			err := json.Unmarshal([]byte(data), &aliResponse) | 			err := json.Unmarshal([]byte(data), &aliResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error unmarshalling stream response: " + err.Error()) | 				logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			if aliResponse.Usage.OutputTokens != 0 { | 			if aliResponse.Usage.OutputTokens != 0 { | ||||||
| @@ -198,7 +200,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			//lastResponseText = aliResponse.Output.Text | 			//lastResponseText = aliResponse.Output.Text | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | ||||||
|   | |||||||
							
								
								
									
										22
									
								
								relay/channel/anthropic/adaptor.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										22
									
								
								relay/channel/anthropic/adaptor.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,22 @@ | |||||||
|  | package anthropic | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor struct { | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) Auth(c *gin.Context) error { | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) ConvertRequest(request *openai.GeneralOpenAIRequest) (any, error) { | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, *openai.Usage, error) { | ||||||
|  | 	return nil, nil, nil | ||||||
|  | } | ||||||
| @@ -8,6 +8,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
| @@ -78,9 +80,9 @@ func responseClaude2OpenAI(claudeResponse *Response) *openai.TextResponse { | |||||||
| 		FinishReason: stopReasonClaude2OpenAI(claudeResponse.StopReason), | 		FinishReason: stopReasonClaude2OpenAI(claudeResponse.StopReason), | ||||||
| 	} | 	} | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := openai.TextResponse{ | ||||||
| 		Id:      fmt.Sprintf("chatcmpl-%s", common.GetUUID()), | 		Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Choices: []openai.TextResponseChoice{choice}, | 		Choices: []openai.TextResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| 	return &fullTextResponse | 	return &fullTextResponse | ||||||
| @@ -88,8 +90,8 @@ func responseClaude2OpenAI(claudeResponse *Response) *openai.TextResponse { | |||||||
|  |  | ||||||
| func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, string) { | func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, string) { | ||||||
| 	responseText := "" | 	responseText := "" | ||||||
| 	responseId := fmt.Sprintf("chatcmpl-%s", common.GetUUID()) | 	responseId := fmt.Sprintf("chatcmpl-%s", helper.GetUUID()) | ||||||
| 	createdTime := common.GetTimestamp() | 	createdTime := helper.GetTimestamp() | ||||||
| 	scanner := bufio.NewScanner(resp.Body) | 	scanner := bufio.NewScanner(resp.Body) | ||||||
| 	scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) { | 	scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) { | ||||||
| 		if atEOF && len(data) == 0 { | 		if atEOF && len(data) == 0 { | ||||||
| @@ -125,7 +127,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			var claudeResponse Response | 			var claudeResponse Response | ||||||
| 			err := json.Unmarshal([]byte(data), &claudeResponse) | 			err := json.Unmarshal([]byte(data), &claudeResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error unmarshalling stream response: " + err.Error()) | 				logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			responseText += claudeResponse.Completion | 			responseText += claudeResponse.Completion | ||||||
| @@ -134,7 +136,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			response.Created = createdTime | 			response.Created = createdTime | ||||||
| 			jsonStr, err := json.Marshal(response) | 			jsonStr, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonStr)}) | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonStr)}) | ||||||
|   | |||||||
							
								
								
									
										22
									
								
								relay/channel/baidu/adaptor.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										22
									
								
								relay/channel/baidu/adaptor.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,22 @@ | |||||||
|  | package baidu | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor struct { | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) Auth(c *gin.Context) error { | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) ConvertRequest(request *openai.GeneralOpenAIRequest) (any, error) { | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, *openai.Usage, error) { | ||||||
|  | 	return nil, nil, nil | ||||||
|  | } | ||||||
| @@ -9,6 +9,7 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| 	"one-api/relay/util" | 	"one-api/relay/util" | ||||||
| @@ -19,49 +20,49 @@ import ( | |||||||
|  |  | ||||||
| // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/flfmc9do2 | // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/flfmc9do2 | ||||||
|  |  | ||||||
| type BaiduTokenResponse struct { | type TokenResponse struct { | ||||||
| 	ExpiresIn   int    `json:"expires_in"` | 	ExpiresIn   int    `json:"expires_in"` | ||||||
| 	AccessToken string `json:"access_token"` | 	AccessToken string `json:"access_token"` | ||||||
| } | } | ||||||
|  |  | ||||||
| type BaiduMessage struct { | type Message struct { | ||||||
| 	Role    string `json:"role"` | 	Role    string `json:"role"` | ||||||
| 	Content string `json:"content"` | 	Content string `json:"content"` | ||||||
| } | } | ||||||
|  |  | ||||||
| type BaiduChatRequest struct { | type ChatRequest struct { | ||||||
| 	Messages []BaiduMessage `json:"messages"` | 	Messages []Message `json:"messages"` | ||||||
| 	Stream   bool           `json:"stream"` | 	Stream   bool      `json:"stream"` | ||||||
| 	UserId   string         `json:"user_id,omitempty"` | 	UserId   string    `json:"user_id,omitempty"` | ||||||
| } | } | ||||||
|  |  | ||||||
| type BaiduError struct { | type Error struct { | ||||||
| 	ErrorCode int    `json:"error_code"` | 	ErrorCode int    `json:"error_code"` | ||||||
| 	ErrorMsg  string `json:"error_msg"` | 	ErrorMsg  string `json:"error_msg"` | ||||||
| } | } | ||||||
|  |  | ||||||
| var baiduTokenStore sync.Map | var baiduTokenStore sync.Map | ||||||
|  |  | ||||||
| func ConvertRequest(request openai.GeneralOpenAIRequest) *BaiduChatRequest { | func ConvertRequest(request openai.GeneralOpenAIRequest) *ChatRequest { | ||||||
| 	messages := make([]BaiduMessage, 0, len(request.Messages)) | 	messages := make([]Message, 0, len(request.Messages)) | ||||||
| 	for _, message := range request.Messages { | 	for _, message := range request.Messages { | ||||||
| 		if message.Role == "system" { | 		if message.Role == "system" { | ||||||
| 			messages = append(messages, BaiduMessage{ | 			messages = append(messages, Message{ | ||||||
| 				Role:    "user", | 				Role:    "user", | ||||||
| 				Content: message.StringContent(), | 				Content: message.StringContent(), | ||||||
| 			}) | 			}) | ||||||
| 			messages = append(messages, BaiduMessage{ | 			messages = append(messages, Message{ | ||||||
| 				Role:    "assistant", | 				Role:    "assistant", | ||||||
| 				Content: "Okay", | 				Content: "Okay", | ||||||
| 			}) | 			}) | ||||||
| 		} else { | 		} else { | ||||||
| 			messages = append(messages, BaiduMessage{ | 			messages = append(messages, Message{ | ||||||
| 				Role:    message.Role, | 				Role:    message.Role, | ||||||
| 				Content: message.StringContent(), | 				Content: message.StringContent(), | ||||||
| 			}) | 			}) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	return &BaiduChatRequest{ | 	return &ChatRequest{ | ||||||
| 		Messages: messages, | 		Messages: messages, | ||||||
| 		Stream:   request.Stream, | 		Stream:   request.Stream, | ||||||
| 	} | 	} | ||||||
| @@ -160,7 +161,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			var baiduResponse ChatStreamResponse | 			var baiduResponse ChatStreamResponse | ||||||
| 			err := json.Unmarshal([]byte(data), &baiduResponse) | 			err := json.Unmarshal([]byte(data), &baiduResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error unmarshalling stream response: " + err.Error()) | 				logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			if baiduResponse.Usage.TotalTokens != 0 { | 			if baiduResponse.Usage.TotalTokens != 0 { | ||||||
| @@ -171,7 +172,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			response := streamResponseBaidu2OpenAI(&baiduResponse) | 			response := streamResponseBaidu2OpenAI(&baiduResponse) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | ||||||
|   | |||||||
| @@ -13,7 +13,7 @@ type ChatResponse struct { | |||||||
| 	IsTruncated      bool         `json:"is_truncated"` | 	IsTruncated      bool         `json:"is_truncated"` | ||||||
| 	NeedClearHistory bool         `json:"need_clear_history"` | 	NeedClearHistory bool         `json:"need_clear_history"` | ||||||
| 	Usage            openai.Usage `json:"usage"` | 	Usage            openai.Usage `json:"usage"` | ||||||
| 	BaiduError | 	Error | ||||||
| } | } | ||||||
|  |  | ||||||
| type ChatStreamResponse struct { | type ChatStreamResponse struct { | ||||||
| @@ -38,7 +38,7 @@ type EmbeddingResponse struct { | |||||||
| 	Created int64           `json:"created"` | 	Created int64           `json:"created"` | ||||||
| 	Data    []EmbeddingData `json:"data"` | 	Data    []EmbeddingData `json:"data"` | ||||||
| 	Usage   openai.Usage    `json:"usage"` | 	Usage   openai.Usage    `json:"usage"` | ||||||
| 	BaiduError | 	Error | ||||||
| } | } | ||||||
|  |  | ||||||
| type AccessToken struct { | type AccessToken struct { | ||||||
|   | |||||||
							
								
								
									
										22
									
								
								relay/channel/google/adaptor.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										22
									
								
								relay/channel/google/adaptor.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,22 @@ | |||||||
|  | package google | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor struct { | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) Auth(c *gin.Context) error { | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) ConvertRequest(request *openai.GeneralOpenAIRequest) (any, error) { | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, *openai.Usage, error) { | ||||||
|  | 	return nil, nil, nil | ||||||
|  | } | ||||||
| @@ -7,7 +7,10 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/helper" | ||||||
| 	"one-api/common/image" | 	"one-api/common/image" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| 	"strings" | 	"strings" | ||||||
| @@ -28,19 +31,19 @@ func ConvertGeminiRequest(textRequest openai.GeneralOpenAIRequest) *GeminiChatRe | |||||||
| 		SafetySettings: []GeminiChatSafetySettings{ | 		SafetySettings: []GeminiChatSafetySettings{ | ||||||
| 			{ | 			{ | ||||||
| 				Category:  "HARM_CATEGORY_HARASSMENT", | 				Category:  "HARM_CATEGORY_HARASSMENT", | ||||||
| 				Threshold: common.GeminiSafetySetting, | 				Threshold: config.GeminiSafetySetting, | ||||||
| 			}, | 			}, | ||||||
| 			{ | 			{ | ||||||
| 				Category:  "HARM_CATEGORY_HATE_SPEECH", | 				Category:  "HARM_CATEGORY_HATE_SPEECH", | ||||||
| 				Threshold: common.GeminiSafetySetting, | 				Threshold: config.GeminiSafetySetting, | ||||||
| 			}, | 			}, | ||||||
| 			{ | 			{ | ||||||
| 				Category:  "HARM_CATEGORY_SEXUALLY_EXPLICIT", | 				Category:  "HARM_CATEGORY_SEXUALLY_EXPLICIT", | ||||||
| 				Threshold: common.GeminiSafetySetting, | 				Threshold: config.GeminiSafetySetting, | ||||||
| 			}, | 			}, | ||||||
| 			{ | 			{ | ||||||
| 				Category:  "HARM_CATEGORY_DANGEROUS_CONTENT", | 				Category:  "HARM_CATEGORY_DANGEROUS_CONTENT", | ||||||
| 				Threshold: common.GeminiSafetySetting, | 				Threshold: config.GeminiSafetySetting, | ||||||
| 			}, | 			}, | ||||||
| 		}, | 		}, | ||||||
| 		GenerationConfig: GeminiChatGenerationConfig{ | 		GenerationConfig: GeminiChatGenerationConfig{ | ||||||
| @@ -151,9 +154,9 @@ type GeminiChatPromptFeedback struct { | |||||||
|  |  | ||||||
| func responseGeminiChat2OpenAI(response *GeminiChatResponse) *openai.TextResponse { | func responseGeminiChat2OpenAI(response *GeminiChatResponse) *openai.TextResponse { | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := openai.TextResponse{ | ||||||
| 		Id:      fmt.Sprintf("chatcmpl-%s", common.GetUUID()), | 		Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Choices: make([]openai.TextResponseChoice, 0, len(response.Candidates)), | 		Choices: make([]openai.TextResponseChoice, 0, len(response.Candidates)), | ||||||
| 	} | 	} | ||||||
| 	for i, candidate := range response.Candidates { | 	for i, candidate := range response.Candidates { | ||||||
| @@ -229,15 +232,15 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			var choice openai.ChatCompletionsStreamResponseChoice | 			var choice openai.ChatCompletionsStreamResponseChoice | ||||||
| 			choice.Delta.Content = dummy.Content | 			choice.Delta.Content = dummy.Content | ||||||
| 			response := openai.ChatCompletionsStreamResponse{ | 			response := openai.ChatCompletionsStreamResponse{ | ||||||
| 				Id:      fmt.Sprintf("chatcmpl-%s", common.GetUUID()), | 				Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), | ||||||
| 				Object:  "chat.completion.chunk", | 				Object:  "chat.completion.chunk", | ||||||
| 				Created: common.GetTimestamp(), | 				Created: helper.GetTimestamp(), | ||||||
| 				Model:   "gemini-pro", | 				Model:   "gemini-pro", | ||||||
| 				Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 				Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 			} | 			} | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | ||||||
|   | |||||||
| @@ -7,6 +7,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| ) | ) | ||||||
| @@ -71,27 +73,27 @@ func streamResponsePaLM2OpenAI(palmResponse *PaLMChatResponse) *openai.ChatCompl | |||||||
|  |  | ||||||
| func PaLMStreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, string) { | func PaLMStreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, string) { | ||||||
| 	responseText := "" | 	responseText := "" | ||||||
| 	responseId := fmt.Sprintf("chatcmpl-%s", common.GetUUID()) | 	responseId := fmt.Sprintf("chatcmpl-%s", helper.GetUUID()) | ||||||
| 	createdTime := common.GetTimestamp() | 	createdTime := helper.GetTimestamp() | ||||||
| 	dataChan := make(chan string) | 	dataChan := make(chan string) | ||||||
| 	stopChan := make(chan bool) | 	stopChan := make(chan bool) | ||||||
| 	go func() { | 	go func() { | ||||||
| 		responseBody, err := io.ReadAll(resp.Body) | 		responseBody, err := io.ReadAll(resp.Body) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("error reading stream response: " + err.Error()) | 			logger.SysError("error reading stream response: " + err.Error()) | ||||||
| 			stopChan <- true | 			stopChan <- true | ||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 		err = resp.Body.Close() | 		err = resp.Body.Close() | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("error closing stream response: " + err.Error()) | 			logger.SysError("error closing stream response: " + err.Error()) | ||||||
| 			stopChan <- true | 			stopChan <- true | ||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 		var palmResponse PaLMChatResponse | 		var palmResponse PaLMChatResponse | ||||||
| 		err = json.Unmarshal(responseBody, &palmResponse) | 		err = json.Unmarshal(responseBody, &palmResponse) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("error unmarshalling stream response: " + err.Error()) | 			logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 			stopChan <- true | 			stopChan <- true | ||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| @@ -103,7 +105,7 @@ func PaLMStreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithSt | |||||||
| 		} | 		} | ||||||
| 		jsonResponse, err := json.Marshal(fullTextResponse) | 		jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("error marshalling stream response: " + err.Error()) | 			logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 			stopChan <- true | 			stopChan <- true | ||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
|   | |||||||
							
								
								
									
										15
									
								
								relay/channel/interface.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										15
									
								
								relay/channel/interface.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,15 @@ | |||||||
|  | package channel | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor interface { | ||||||
|  | 	GetRequestURL() string | ||||||
|  | 	Auth(c *gin.Context) error | ||||||
|  | 	ConvertRequest(request *openai.GeneralOpenAIRequest) (any, error) | ||||||
|  | 	DoRequest(request *openai.GeneralOpenAIRequest) error | ||||||
|  | 	DoResponse(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, *openai.Usage, error) | ||||||
|  | } | ||||||
							
								
								
									
										21
									
								
								relay/channel/openai/adaptor.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										21
									
								
								relay/channel/openai/adaptor.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,21 @@ | |||||||
|  | package openai | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor struct { | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) Auth(c *gin.Context) error { | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) ConvertRequest(request *GeneralOpenAIRequest) (any, error) { | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response) (*ErrorWithStatusCode, *Usage, error) { | ||||||
|  | 	return nil, nil, nil | ||||||
|  | } | ||||||
| @@ -8,6 +8,7 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
| @@ -46,7 +47,7 @@ func StreamHandler(c *gin.Context, resp *http.Response, relayMode int) (*ErrorWi | |||||||
| 					var streamResponse ChatCompletionsStreamResponse | 					var streamResponse ChatCompletionsStreamResponse | ||||||
| 					err := json.Unmarshal([]byte(data), &streamResponse) | 					err := json.Unmarshal([]byte(data), &streamResponse) | ||||||
| 					if err != nil { | 					if err != nil { | ||||||
| 						common.SysError("error unmarshalling stream response: " + err.Error()) | 						logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 						continue // just ignore the error | 						continue // just ignore the error | ||||||
| 					} | 					} | ||||||
| 					for _, choice := range streamResponse.Choices { | 					for _, choice := range streamResponse.Choices { | ||||||
| @@ -56,7 +57,7 @@ func StreamHandler(c *gin.Context, resp *http.Response, relayMode int) (*ErrorWi | |||||||
| 					var streamResponse CompletionsStreamResponse | 					var streamResponse CompletionsStreamResponse | ||||||
| 					err := json.Unmarshal([]byte(data), &streamResponse) | 					err := json.Unmarshal([]byte(data), &streamResponse) | ||||||
| 					if err != nil { | 					if err != nil { | ||||||
| 						common.SysError("error unmarshalling stream response: " + err.Error()) | 						logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 						continue | 						continue | ||||||
| 					} | 					} | ||||||
| 					for _, choice := range streamResponse.Choices { | 					for _, choice := range streamResponse.Choices { | ||||||
|   | |||||||
| @@ -207,6 +207,11 @@ type Usage struct { | |||||||
| 	TotalTokens      int `json:"total_tokens"` | 	TotalTokens      int `json:"total_tokens"` | ||||||
| } | } | ||||||
|  |  | ||||||
|  | type UsageOrResponseText struct { | ||||||
|  | 	*Usage | ||||||
|  | 	ResponseText string | ||||||
|  | } | ||||||
|  |  | ||||||
| type Error struct { | type Error struct { | ||||||
| 	Message string `json:"message"` | 	Message string `json:"message"` | ||||||
| 	Type    string `json:"type"` | 	Type    string `json:"type"` | ||||||
|   | |||||||
| @@ -6,7 +6,9 @@ import ( | |||||||
| 	"github.com/pkoukk/tiktoken-go" | 	"github.com/pkoukk/tiktoken-go" | ||||||
| 	"math" | 	"math" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
| 	"one-api/common/image" | 	"one-api/common/image" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -15,15 +17,15 @@ var tokenEncoderMap = map[string]*tiktoken.Tiktoken{} | |||||||
| var defaultTokenEncoder *tiktoken.Tiktoken | var defaultTokenEncoder *tiktoken.Tiktoken | ||||||
|  |  | ||||||
| func InitTokenEncoders() { | func InitTokenEncoders() { | ||||||
| 	common.SysLog("initializing token encoders") | 	logger.SysLog("initializing token encoders") | ||||||
| 	gpt35TokenEncoder, err := tiktoken.EncodingForModel("gpt-3.5-turbo") | 	gpt35TokenEncoder, err := tiktoken.EncodingForModel("gpt-3.5-turbo") | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.FatalLog(fmt.Sprintf("failed to get gpt-3.5-turbo token encoder: %s", err.Error())) | 		logger.FatalLog(fmt.Sprintf("failed to get gpt-3.5-turbo token encoder: %s", err.Error())) | ||||||
| 	} | 	} | ||||||
| 	defaultTokenEncoder = gpt35TokenEncoder | 	defaultTokenEncoder = gpt35TokenEncoder | ||||||
| 	gpt4TokenEncoder, err := tiktoken.EncodingForModel("gpt-4") | 	gpt4TokenEncoder, err := tiktoken.EncodingForModel("gpt-4") | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.FatalLog(fmt.Sprintf("failed to get gpt-4 token encoder: %s", err.Error())) | 		logger.FatalLog(fmt.Sprintf("failed to get gpt-4 token encoder: %s", err.Error())) | ||||||
| 	} | 	} | ||||||
| 	for model, _ := range common.ModelRatio { | 	for model, _ := range common.ModelRatio { | ||||||
| 		if strings.HasPrefix(model, "gpt-3.5") { | 		if strings.HasPrefix(model, "gpt-3.5") { | ||||||
| @@ -34,7 +36,7 @@ func InitTokenEncoders() { | |||||||
| 			tokenEncoderMap[model] = nil | 			tokenEncoderMap[model] = nil | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	common.SysLog("token encoders initialized") | 	logger.SysLog("token encoders initialized") | ||||||
| } | } | ||||||
|  |  | ||||||
| func getTokenEncoder(model string) *tiktoken.Tiktoken { | func getTokenEncoder(model string) *tiktoken.Tiktoken { | ||||||
| @@ -45,7 +47,7 @@ func getTokenEncoder(model string) *tiktoken.Tiktoken { | |||||||
| 	if ok { | 	if ok { | ||||||
| 		tokenEncoder, err := tiktoken.EncodingForModel(model) | 		tokenEncoder, err := tiktoken.EncodingForModel(model) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError(fmt.Sprintf("failed to get token encoder for model %s: %s, using encoder for gpt-3.5-turbo", model, err.Error())) | 			logger.SysError(fmt.Sprintf("failed to get token encoder for model %s: %s, using encoder for gpt-3.5-turbo", model, err.Error())) | ||||||
| 			tokenEncoder = defaultTokenEncoder | 			tokenEncoder = defaultTokenEncoder | ||||||
| 		} | 		} | ||||||
| 		tokenEncoderMap[model] = tokenEncoder | 		tokenEncoderMap[model] = tokenEncoder | ||||||
| @@ -55,7 +57,7 @@ func getTokenEncoder(model string) *tiktoken.Tiktoken { | |||||||
| } | } | ||||||
|  |  | ||||||
| func getTokenNum(tokenEncoder *tiktoken.Tiktoken, text string) int { | func getTokenNum(tokenEncoder *tiktoken.Tiktoken, text string) int { | ||||||
| 	if common.ApproximateTokenEnabled { | 	if config.ApproximateTokenEnabled { | ||||||
| 		return int(float64(len(text)) * 0.38) | 		return int(float64(len(text)) * 0.38) | ||||||
| 	} | 	} | ||||||
| 	return len(tokenEncoder.Encode(text, nil, nil)) | 	return len(tokenEncoder.Encode(text, nil, nil)) | ||||||
| @@ -99,7 +101,7 @@ func CountTokenMessages(messages []Message, model string) int { | |||||||
| 						} | 						} | ||||||
| 						imageTokens, err := countImageTokens(url, detail) | 						imageTokens, err := countImageTokens(url, detail) | ||||||
| 						if err != nil { | 						if err != nil { | ||||||
| 							common.SysError("error counting image tokens: " + err.Error()) | 							logger.SysError("error counting image tokens: " + err.Error()) | ||||||
| 						} else { | 						} else { | ||||||
| 							tokenNum += imageTokens | 							tokenNum += imageTokens | ||||||
| 						} | 						} | ||||||
|   | |||||||
							
								
								
									
										22
									
								
								relay/channel/tencent/adaptor.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										22
									
								
								relay/channel/tencent/adaptor.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,22 @@ | |||||||
|  | package tencent | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor struct { | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) Auth(c *gin.Context) error { | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) ConvertRequest(request *openai.GeneralOpenAIRequest) (any, error) { | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, *openai.Usage, error) { | ||||||
|  | 	return nil, nil, nil | ||||||
|  | } | ||||||
| @@ -12,6 +12,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| 	"sort" | 	"sort" | ||||||
| @@ -46,9 +48,9 @@ func ConvertRequest(request openai.GeneralOpenAIRequest) *ChatRequest { | |||||||
| 		stream = 1 | 		stream = 1 | ||||||
| 	} | 	} | ||||||
| 	return &ChatRequest{ | 	return &ChatRequest{ | ||||||
| 		Timestamp:   common.GetTimestamp(), | 		Timestamp:   helper.GetTimestamp(), | ||||||
| 		Expired:     common.GetTimestamp() + 24*60*60, | 		Expired:     helper.GetTimestamp() + 24*60*60, | ||||||
| 		QueryID:     common.GetUUID(), | 		QueryID:     helper.GetUUID(), | ||||||
| 		Temperature: request.Temperature, | 		Temperature: request.Temperature, | ||||||
| 		TopP:        request.TopP, | 		TopP:        request.TopP, | ||||||
| 		Stream:      stream, | 		Stream:      stream, | ||||||
| @@ -59,7 +61,7 @@ func ConvertRequest(request openai.GeneralOpenAIRequest) *ChatRequest { | |||||||
| func responseTencent2OpenAI(response *ChatResponse) *openai.TextResponse { | func responseTencent2OpenAI(response *ChatResponse) *openai.TextResponse { | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := openai.TextResponse{ | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Usage:   response.Usage, | 		Usage:   response.Usage, | ||||||
| 	} | 	} | ||||||
| 	if len(response.Choices) > 0 { | 	if len(response.Choices) > 0 { | ||||||
| @@ -79,7 +81,7 @@ func responseTencent2OpenAI(response *ChatResponse) *openai.TextResponse { | |||||||
| func streamResponseTencent2OpenAI(TencentResponse *ChatResponse) *openai.ChatCompletionsStreamResponse { | func streamResponseTencent2OpenAI(TencentResponse *ChatResponse) *openai.ChatCompletionsStreamResponse { | ||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := openai.ChatCompletionsStreamResponse{ | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Model:   "tencent-hunyuan", | 		Model:   "tencent-hunyuan", | ||||||
| 	} | 	} | ||||||
| 	if len(TencentResponse.Choices) > 0 { | 	if len(TencentResponse.Choices) > 0 { | ||||||
| @@ -131,7 +133,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			var TencentResponse ChatResponse | 			var TencentResponse ChatResponse | ||||||
| 			err := json.Unmarshal([]byte(data), &TencentResponse) | 			err := json.Unmarshal([]byte(data), &TencentResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error unmarshalling stream response: " + err.Error()) | 				logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			response := streamResponseTencent2OpenAI(&TencentResponse) | 			response := streamResponseTencent2OpenAI(&TencentResponse) | ||||||
| @@ -140,7 +142,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			} | 			} | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | ||||||
|   | |||||||
							
								
								
									
										22
									
								
								relay/channel/xunfei/adaptor.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										22
									
								
								relay/channel/xunfei/adaptor.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,22 @@ | |||||||
|  | package xunfei | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor struct { | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) Auth(c *gin.Context) error { | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) ConvertRequest(request *openai.GeneralOpenAIRequest) (any, error) { | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, *openai.Usage, error) { | ||||||
|  | 	return nil, nil, nil | ||||||
|  | } | ||||||
| @@ -12,6 +12,8 @@ import ( | |||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"net/url" | 	"net/url" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| 	"strings" | 	"strings" | ||||||
| @@ -68,7 +70,7 @@ func responseXunfei2OpenAI(response *ChatResponse) *openai.TextResponse { | |||||||
| 	} | 	} | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := openai.TextResponse{ | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Choices: []openai.TextResponseChoice{choice}, | 		Choices: []openai.TextResponseChoice{choice}, | ||||||
| 		Usage:   response.Payload.Usage.Text, | 		Usage:   response.Payload.Usage.Text, | ||||||
| 	} | 	} | ||||||
| @@ -90,7 +92,7 @@ func streamResponseXunfei2OpenAI(xunfeiResponse *ChatResponse) *openai.ChatCompl | |||||||
| 	} | 	} | ||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := openai.ChatCompletionsStreamResponse{ | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Model:   "SparkDesk", | 		Model:   "SparkDesk", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| @@ -140,7 +142,7 @@ func StreamHandler(c *gin.Context, textRequest openai.GeneralOpenAIRequest, appI | |||||||
| 			response := streamResponseXunfei2OpenAI(&xunfeiResponse) | 			response := streamResponseXunfei2OpenAI(&xunfeiResponse) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | ||||||
| @@ -215,20 +217,20 @@ func xunfeiMakeRequest(textRequest openai.GeneralOpenAIRequest, domain, authUrl, | |||||||
| 		for { | 		for { | ||||||
| 			_, msg, err := conn.ReadMessage() | 			_, msg, err := conn.ReadMessage() | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error reading stream response: " + err.Error()) | 				logger.SysError("error reading stream response: " + err.Error()) | ||||||
| 				break | 				break | ||||||
| 			} | 			} | ||||||
| 			var response ChatResponse | 			var response ChatResponse | ||||||
| 			err = json.Unmarshal(msg, &response) | 			err = json.Unmarshal(msg, &response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error unmarshalling stream response: " + err.Error()) | 				logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				break | 				break | ||||||
| 			} | 			} | ||||||
| 			dataChan <- response | 			dataChan <- response | ||||||
| 			if response.Payload.Choices.Status == 2 { | 			if response.Payload.Choices.Status == 2 { | ||||||
| 				err := conn.Close() | 				err := conn.Close() | ||||||
| 				if err != nil { | 				if err != nil { | ||||||
| 					common.SysError("error closing websocket connection: " + err.Error()) | 					logger.SysError("error closing websocket connection: " + err.Error()) | ||||||
| 				} | 				} | ||||||
| 				break | 				break | ||||||
| 			} | 			} | ||||||
| @@ -247,7 +249,7 @@ func getXunfeiAuthUrl(c *gin.Context, apiKey string, apiSecret string) (string, | |||||||
| 	} | 	} | ||||||
| 	if apiVersion == "" { | 	if apiVersion == "" { | ||||||
| 		apiVersion = "v1.1" | 		apiVersion = "v1.1" | ||||||
| 		common.SysLog("api_version not found, use default: " + apiVersion) | 		logger.SysLog("api_version not found, use default: " + apiVersion) | ||||||
| 	} | 	} | ||||||
| 	domain := "general" | 	domain := "general" | ||||||
| 	if apiVersion != "v1.1" { | 	if apiVersion != "v1.1" { | ||||||
|   | |||||||
							
								
								
									
										22
									
								
								relay/channel/zhipu/adaptor.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										22
									
								
								relay/channel/zhipu/adaptor.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,22 @@ | |||||||
|  | package zhipu | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Adaptor struct { | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) Auth(c *gin.Context) error { | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) ConvertRequest(request *openai.GeneralOpenAIRequest) (any, error) { | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatusCode, *openai.Usage, error) { | ||||||
|  | 	return nil, nil, nil | ||||||
|  | } | ||||||
| @@ -8,6 +8,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| 	"strings" | 	"strings" | ||||||
| @@ -34,7 +36,7 @@ func GetToken(apikey string) string { | |||||||
|  |  | ||||||
| 	split := strings.Split(apikey, ".") | 	split := strings.Split(apikey, ".") | ||||||
| 	if len(split) != 2 { | 	if len(split) != 2 { | ||||||
| 		common.SysError("invalid zhipu key: " + apikey) | 		logger.SysError("invalid zhipu key: " + apikey) | ||||||
| 		return "" | 		return "" | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| @@ -101,7 +103,7 @@ func responseZhipu2OpenAI(response *Response) *openai.TextResponse { | |||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := openai.TextResponse{ | ||||||
| 		Id:      response.Data.TaskId, | 		Id:      response.Data.TaskId, | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Choices: make([]openai.TextResponseChoice, 0, len(response.Data.Choices)), | 		Choices: make([]openai.TextResponseChoice, 0, len(response.Data.Choices)), | ||||||
| 		Usage:   response.Data.Usage, | 		Usage:   response.Data.Usage, | ||||||
| 	} | 	} | ||||||
| @@ -127,7 +129,7 @@ func streamResponseZhipu2OpenAI(zhipuResponse string) *openai.ChatCompletionsStr | |||||||
| 	choice.Delta.Content = zhipuResponse | 	choice.Delta.Content = zhipuResponse | ||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := openai.ChatCompletionsStreamResponse{ | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Model:   "chatglm", | 		Model:   "chatglm", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| @@ -141,7 +143,7 @@ func streamMetaResponseZhipu2OpenAI(zhipuResponse *StreamMetaResponse) (*openai. | |||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := openai.ChatCompletionsStreamResponse{ | ||||||
| 		Id:      zhipuResponse.RequestId, | 		Id:      zhipuResponse.RequestId, | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: common.GetTimestamp(), | 		Created: helper.GetTimestamp(), | ||||||
| 		Model:   "chatglm", | 		Model:   "chatglm", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| @@ -193,7 +195,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			response := streamResponseZhipu2OpenAI(data) | 			response := streamResponseZhipu2OpenAI(data) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | ||||||
| @@ -202,13 +204,13 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*openai.ErrorWithStatus | |||||||
| 			var zhipuResponse StreamMetaResponse | 			var zhipuResponse StreamMetaResponse | ||||||
| 			err := json.Unmarshal([]byte(data), &zhipuResponse) | 			err := json.Unmarshal([]byte(data), &zhipuResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error unmarshalling stream response: " + err.Error()) | 				logger.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			response, zhipuUsage := streamMetaResponseZhipu2OpenAI(&zhipuResponse) | 			response, zhipuUsage := streamMetaResponseZhipu2OpenAI(&zhipuResponse) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.SysError("error marshalling stream response: " + err.Error()) | 				logger.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			usage = zhipuUsage | 			usage = zhipuUsage | ||||||
|   | |||||||
							
								
								
									
										69
									
								
								relay/constant/api_type.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										69
									
								
								relay/constant/api_type.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,69 @@ | |||||||
|  | package constant | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"one-api/common" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | const ( | ||||||
|  | 	APITypeOpenAI = iota | ||||||
|  | 	APITypeClaude | ||||||
|  | 	APITypePaLM | ||||||
|  | 	APITypeBaidu | ||||||
|  | 	APITypeZhipu | ||||||
|  | 	APITypeAli | ||||||
|  | 	APITypeXunfei | ||||||
|  | 	APITypeAIProxyLibrary | ||||||
|  | 	APITypeTencent | ||||||
|  | 	APITypeGemini | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func ChannelType2APIType(channelType int) int { | ||||||
|  | 	apiType := APITypeOpenAI | ||||||
|  | 	switch channelType { | ||||||
|  | 	case common.ChannelTypeAnthropic: | ||||||
|  | 		apiType = APITypeClaude | ||||||
|  | 	case common.ChannelTypeBaidu: | ||||||
|  | 		apiType = APITypeBaidu | ||||||
|  | 	case common.ChannelTypePaLM: | ||||||
|  | 		apiType = APITypePaLM | ||||||
|  | 	case common.ChannelTypeZhipu: | ||||||
|  | 		apiType = APITypeZhipu | ||||||
|  | 	case common.ChannelTypeAli: | ||||||
|  | 		apiType = APITypeAli | ||||||
|  | 	case common.ChannelTypeXunfei: | ||||||
|  | 		apiType = APITypeXunfei | ||||||
|  | 	case common.ChannelTypeAIProxyLibrary: | ||||||
|  | 		apiType = APITypeAIProxyLibrary | ||||||
|  | 	case common.ChannelTypeTencent: | ||||||
|  | 		apiType = APITypeTencent | ||||||
|  | 	case common.ChannelTypeGemini: | ||||||
|  | 		apiType = APITypeGemini | ||||||
|  | 	} | ||||||
|  | 	return apiType | ||||||
|  | } | ||||||
|  |  | ||||||
|  | //func GetAdaptor(apiType int) channel.Adaptor { | ||||||
|  | //	switch apiType { | ||||||
|  | //	case APITypeOpenAI: | ||||||
|  | //		return &openai.Adaptor{} | ||||||
|  | //	case APITypeClaude: | ||||||
|  | //		return &anthropic.Adaptor{} | ||||||
|  | //	case APITypePaLM: | ||||||
|  | //		return &google.Adaptor{} | ||||||
|  | //	case APITypeZhipu: | ||||||
|  | //		return &baidu.Adaptor{} | ||||||
|  | //	case APITypeBaidu: | ||||||
|  | //		return &baidu.Adaptor{} | ||||||
|  | //	case APITypeAli: | ||||||
|  | //		return &ali.Adaptor{} | ||||||
|  | //	case APITypeXunfei: | ||||||
|  | //		return &xunfei.Adaptor{} | ||||||
|  | //	case APITypeAIProxyLibrary: | ||||||
|  | //		return &aiproxy.Adaptor{} | ||||||
|  | //	case APITypeTencent: | ||||||
|  | //		return &tencent.Adaptor{} | ||||||
|  | //	case APITypeGemini: | ||||||
|  | //		return &google.Adaptor{} | ||||||
|  | //	} | ||||||
|  | //	return nil | ||||||
|  | //} | ||||||
							
								
								
									
										3
									
								
								relay/constant/common.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										3
									
								
								relay/constant/common.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,3 @@ | |||||||
|  | package constant | ||||||
|  |  | ||||||
|  | var StopFinishReason = "stop" | ||||||
| @@ -1,16 +0,0 @@ | |||||||
| package constant |  | ||||||
|  |  | ||||||
| const ( |  | ||||||
| 	RelayModeUnknown = iota |  | ||||||
| 	RelayModeChatCompletions |  | ||||||
| 	RelayModeCompletions |  | ||||||
| 	RelayModeEmbeddings |  | ||||||
| 	RelayModeModerations |  | ||||||
| 	RelayModeImagesGenerations |  | ||||||
| 	RelayModeEdits |  | ||||||
| 	RelayModeAudioSpeech |  | ||||||
| 	RelayModeAudioTranscription |  | ||||||
| 	RelayModeAudioTranslation |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| var StopFinishReason = "stop" |  | ||||||
							
								
								
									
										42
									
								
								relay/constant/relay_mode.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										42
									
								
								relay/constant/relay_mode.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,42 @@ | |||||||
|  | package constant | ||||||
|  |  | ||||||
|  | import "strings" | ||||||
|  |  | ||||||
|  | const ( | ||||||
|  | 	RelayModeUnknown = iota | ||||||
|  | 	RelayModeChatCompletions | ||||||
|  | 	RelayModeCompletions | ||||||
|  | 	RelayModeEmbeddings | ||||||
|  | 	RelayModeModerations | ||||||
|  | 	RelayModeImagesGenerations | ||||||
|  | 	RelayModeEdits | ||||||
|  | 	RelayModeAudioSpeech | ||||||
|  | 	RelayModeAudioTranscription | ||||||
|  | 	RelayModeAudioTranslation | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func Path2RelayMode(path string) int { | ||||||
|  | 	relayMode := RelayModeUnknown | ||||||
|  | 	if strings.HasPrefix(path, "/v1/chat/completions") { | ||||||
|  | 		relayMode = RelayModeChatCompletions | ||||||
|  | 	} else if strings.HasPrefix(path, "/v1/completions") { | ||||||
|  | 		relayMode = RelayModeCompletions | ||||||
|  | 	} else if strings.HasPrefix(path, "/v1/embeddings") { | ||||||
|  | 		relayMode = RelayModeEmbeddings | ||||||
|  | 	} else if strings.HasSuffix(path, "embeddings") { | ||||||
|  | 		relayMode = RelayModeEmbeddings | ||||||
|  | 	} else if strings.HasPrefix(path, "/v1/moderations") { | ||||||
|  | 		relayMode = RelayModeModerations | ||||||
|  | 	} else if strings.HasPrefix(path, "/v1/images/generations") { | ||||||
|  | 		relayMode = RelayModeImagesGenerations | ||||||
|  | 	} else if strings.HasPrefix(path, "/v1/edits") { | ||||||
|  | 		relayMode = RelayModeEdits | ||||||
|  | 	} else if strings.HasPrefix(path, "/v1/audio/speech") { | ||||||
|  | 		relayMode = RelayModeAudioSpeech | ||||||
|  | 	} else if strings.HasPrefix(path, "/v1/audio/transcriptions") { | ||||||
|  | 		relayMode = RelayModeAudioTranscription | ||||||
|  | 	} else if strings.HasPrefix(path, "/v1/audio/translations") { | ||||||
|  | 		relayMode = RelayModeAudioTranslation | ||||||
|  | 	} | ||||||
|  | 	return relayMode | ||||||
|  | } | ||||||
| @@ -11,6 +11,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| @@ -53,7 +55,7 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode | |||||||
| 		preConsumedQuota = int(float64(len(ttsRequest.Input)) * ratio) | 		preConsumedQuota = int(float64(len(ttsRequest.Input)) * ratio) | ||||||
| 		quota = preConsumedQuota | 		quota = preConsumedQuota | ||||||
| 	default: | 	default: | ||||||
| 		preConsumedQuota = int(float64(common.PreConsumedQuota) * ratio) | 		preConsumedQuota = int(float64(config.PreConsumedQuota) * ratio) | ||||||
| 	} | 	} | ||||||
| 	userQuota, err := model.CacheGetUserQuota(userId) | 	userQuota, err := model.CacheGetUserQuota(userId) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| @@ -102,7 +104,7 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode | |||||||
| 	fullRequestURL := util.GetFullRequestURL(baseURL, requestURL, channelType) | 	fullRequestURL := util.GetFullRequestURL(baseURL, requestURL, channelType) | ||||||
| 	if relayMode == constant.RelayModeAudioTranscription && channelType == common.ChannelTypeAzure { | 	if relayMode == constant.RelayModeAudioTranscription && channelType == common.ChannelTypeAzure { | ||||||
| 		// https://learn.microsoft.com/en-us/azure/ai-services/openai/whisper-quickstart?tabs=command-line#rest-api | 		// https://learn.microsoft.com/en-us/azure/ai-services/openai/whisper-quickstart?tabs=command-line#rest-api | ||||||
| 		apiVersion := util.GetAPIVersion(c) | 		apiVersion := util.GetAzureAPIVersion(c) | ||||||
| 		fullRequestURL = fmt.Sprintf("%s/openai/deployments/%s/audio/transcriptions?api-version=%s", baseURL, audioModel, apiVersion) | 		fullRequestURL = fmt.Sprintf("%s/openai/deployments/%s/audio/transcriptions?api-version=%s", baseURL, audioModel, apiVersion) | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| @@ -191,7 +193,7 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode | |||||||
| 					// negative means add quota back for token & user | 					// negative means add quota back for token & user | ||||||
| 					err := model.PostConsumeTokenQuota(tokenId, -preConsumedQuota) | 					err := model.PostConsumeTokenQuota(tokenId, -preConsumedQuota) | ||||||
| 					if err != nil { | 					if err != nil { | ||||||
| 						common.LogError(ctx, fmt.Sprintf("error rollback pre-consumed quota: %s", err.Error())) | 						logger.Error(ctx, fmt.Sprintf("error rollback pre-consumed quota: %s", err.Error())) | ||||||
| 					} | 					} | ||||||
| 				}() | 				}() | ||||||
| 			}(c.Request.Context()) | 			}(c.Request.Context()) | ||||||
|   | |||||||
| @@ -9,6 +9,7 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/util" | 	"one-api/relay/util" | ||||||
| @@ -112,7 +113,7 @@ func RelayImageHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode | |||||||
| 	fullRequestURL := util.GetFullRequestURL(baseURL, requestURL, channelType) | 	fullRequestURL := util.GetFullRequestURL(baseURL, requestURL, channelType) | ||||||
| 	if channelType == common.ChannelTypeAzure { | 	if channelType == common.ChannelTypeAzure { | ||||||
| 		// https://learn.microsoft.com/en-us/azure/ai-services/openai/dall-e-quickstart?tabs=dalle3%2Ccommand-line&pivots=rest-api | 		// https://learn.microsoft.com/en-us/azure/ai-services/openai/dall-e-quickstart?tabs=dalle3%2Ccommand-line&pivots=rest-api | ||||||
| 		apiVersion := util.GetAPIVersion(c) | 		apiVersion := util.GetAzureAPIVersion(c) | ||||||
| 		// https://{resource_name}.openai.azure.com/openai/deployments/dall-e-3/images/generations?api-version=2023-06-01-preview | 		// https://{resource_name}.openai.azure.com/openai/deployments/dall-e-3/images/generations?api-version=2023-06-01-preview | ||||||
| 		fullRequestURL = fmt.Sprintf("%s/openai/deployments/%s/images/generations?api-version=%s", baseURL, imageModel, apiVersion) | 		fullRequestURL = fmt.Sprintf("%s/openai/deployments/%s/images/generations?api-version=%s", baseURL, imageModel, apiVersion) | ||||||
| 	} | 	} | ||||||
| @@ -175,11 +176,11 @@ func RelayImageHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode | |||||||
| 		} | 		} | ||||||
| 		err := model.PostConsumeTokenQuota(tokenId, quota) | 		err := model.PostConsumeTokenQuota(tokenId, quota) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("error consuming token remain quota: " + err.Error()) | 			logger.SysError("error consuming token remain quota: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 		err = model.CacheUpdateUserQuota(userId) | 		err = model.CacheUpdateUserQuota(userId) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			common.SysError("error update user quota cache: " + err.Error()) | 			logger.SysError("error update user quota cache: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 		if quota != 0 { | 		if quota != 0 { | ||||||
| 			tokenName := c.GetString("token_name") | 			tokenName := c.GetString("token_name") | ||||||
|   | |||||||
| @@ -1,206 +1,47 @@ | |||||||
| package controller | package controller | ||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"bytes" |  | ||||||
| 	"context" | 	"context" | ||||||
| 	"encoding/json" |  | ||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"io" |  | ||||||
| 	"math" | 	"math" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"one-api/relay/channel/aiproxy" |  | ||||||
| 	"one-api/relay/channel/ali" |  | ||||||
| 	"one-api/relay/channel/anthropic" |  | ||||||
| 	"one-api/relay/channel/baidu" |  | ||||||
| 	"one-api/relay/channel/google" |  | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"one-api/relay/channel/tencent" |  | ||||||
| 	"one-api/relay/channel/xunfei" |  | ||||||
| 	"one-api/relay/channel/zhipu" |  | ||||||
| 	"one-api/relay/constant" | 	"one-api/relay/constant" | ||||||
| 	"one-api/relay/util" | 	"one-api/relay/util" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| const ( |  | ||||||
| 	APITypeOpenAI = iota |  | ||||||
| 	APITypeClaude |  | ||||||
| 	APITypePaLM |  | ||||||
| 	APITypeBaidu |  | ||||||
| 	APITypeZhipu |  | ||||||
| 	APITypeAli |  | ||||||
| 	APITypeXunfei |  | ||||||
| 	APITypeAIProxyLibrary |  | ||||||
| 	APITypeTencent |  | ||||||
| 	APITypeGemini |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| func RelayTextHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode { | func RelayTextHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode { | ||||||
| 	channelType := c.GetInt("channel") | 	ctx := c.Request.Context() | ||||||
| 	channelId := c.GetInt("channel_id") | 	meta := util.GetRelayMeta(c) | ||||||
| 	tokenId := c.GetInt("token_id") |  | ||||||
| 	userId := c.GetInt("id") |  | ||||||
| 	group := c.GetString("group") |  | ||||||
| 	var textRequest openai.GeneralOpenAIRequest | 	var textRequest openai.GeneralOpenAIRequest | ||||||
| 	err := common.UnmarshalBodyReusable(c, &textRequest) | 	err := common.UnmarshalBodyReusable(c, &textRequest) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "bind_request_body_failed", http.StatusBadRequest) | 		return openai.ErrorWrapper(err, "bind_request_body_failed", http.StatusBadRequest) | ||||||
| 	} | 	} | ||||||
| 	if textRequest.MaxTokens < 0 || textRequest.MaxTokens > math.MaxInt32/2 { |  | ||||||
| 		return openai.ErrorWrapper(errors.New("max_tokens is invalid"), "invalid_max_tokens", http.StatusBadRequest) |  | ||||||
| 	} |  | ||||||
| 	if relayMode == constant.RelayModeModerations && textRequest.Model == "" { | 	if relayMode == constant.RelayModeModerations && textRequest.Model == "" { | ||||||
| 		textRequest.Model = "text-moderation-latest" | 		textRequest.Model = "text-moderation-latest" | ||||||
| 	} | 	} | ||||||
| 	if relayMode == constant.RelayModeEmbeddings && textRequest.Model == "" { | 	if relayMode == constant.RelayModeEmbeddings && textRequest.Model == "" { | ||||||
| 		textRequest.Model = c.Param("model") | 		textRequest.Model = c.Param("model") | ||||||
| 	} | 	} | ||||||
| 	// request validation | 	err = util.ValidateTextRequest(&textRequest, relayMode) | ||||||
| 	if textRequest.Model == "" { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(errors.New("model is required"), "required_field_missing", http.StatusBadRequest) | 		return openai.ErrorWrapper(err, "invalid_text_request", http.StatusBadRequest) | ||||||
| 	} | 	} | ||||||
| 	switch relayMode { | 	var isModelMapped bool | ||||||
| 	case constant.RelayModeCompletions: | 	textRequest.Model, isModelMapped = util.GetMappedModelName(textRequest.Model, meta.ModelMapping) | ||||||
| 		if textRequest.Prompt == "" { | 	apiType := constant.ChannelType2APIType(meta.ChannelType) | ||||||
| 			return openai.ErrorWrapper(errors.New("field prompt is required"), "required_field_missing", http.StatusBadRequest) | 	fullRequestURL, err := GetRequestURL(c.Request.URL.String(), apiType, relayMode, meta, &textRequest) | ||||||
| 		} | 	if err != nil { | ||||||
| 	case constant.RelayModeChatCompletions: | 		logger.Error(ctx, fmt.Sprintf("util.GetRequestURL failed: %s", err.Error())) | ||||||
| 		if textRequest.Messages == nil || len(textRequest.Messages) == 0 { | 		return openai.ErrorWrapper(fmt.Errorf("util.GetRequestURL failed"), "get_request_url_failed", http.StatusInternalServerError) | ||||||
| 			return openai.ErrorWrapper(errors.New("field messages is required"), "required_field_missing", http.StatusBadRequest) |  | ||||||
| 		} |  | ||||||
| 	case constant.RelayModeEmbeddings: |  | ||||||
| 	case constant.RelayModeModerations: |  | ||||||
| 		if textRequest.Input == "" { |  | ||||||
| 			return openai.ErrorWrapper(errors.New("field input is required"), "required_field_missing", http.StatusBadRequest) |  | ||||||
| 		} |  | ||||||
| 	case constant.RelayModeEdits: |  | ||||||
| 		if textRequest.Instruction == "" { |  | ||||||
| 			return openai.ErrorWrapper(errors.New("field instruction is required"), "required_field_missing", http.StatusBadRequest) |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	// map model name |  | ||||||
| 	modelMapping := c.GetString("model_mapping") |  | ||||||
| 	isModelMapped := false |  | ||||||
| 	if modelMapping != "" && modelMapping != "{}" { |  | ||||||
| 		modelMap := make(map[string]string) |  | ||||||
| 		err := json.Unmarshal([]byte(modelMapping), &modelMap) |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "unmarshal_model_mapping_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		if modelMap[textRequest.Model] != "" { |  | ||||||
| 			textRequest.Model = modelMap[textRequest.Model] |  | ||||||
| 			isModelMapped = true |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	apiType := APITypeOpenAI |  | ||||||
| 	switch channelType { |  | ||||||
| 	case common.ChannelTypeAnthropic: |  | ||||||
| 		apiType = APITypeClaude |  | ||||||
| 	case common.ChannelTypeBaidu: |  | ||||||
| 		apiType = APITypeBaidu |  | ||||||
| 	case common.ChannelTypePaLM: |  | ||||||
| 		apiType = APITypePaLM |  | ||||||
| 	case common.ChannelTypeZhipu: |  | ||||||
| 		apiType = APITypeZhipu |  | ||||||
| 	case common.ChannelTypeAli: |  | ||||||
| 		apiType = APITypeAli |  | ||||||
| 	case common.ChannelTypeXunfei: |  | ||||||
| 		apiType = APITypeXunfei |  | ||||||
| 	case common.ChannelTypeAIProxyLibrary: |  | ||||||
| 		apiType = APITypeAIProxyLibrary |  | ||||||
| 	case common.ChannelTypeTencent: |  | ||||||
| 		apiType = APITypeTencent |  | ||||||
| 	case common.ChannelTypeGemini: |  | ||||||
| 		apiType = APITypeGemini |  | ||||||
| 	} |  | ||||||
| 	baseURL := common.ChannelBaseURLs[channelType] |  | ||||||
| 	requestURL := c.Request.URL.String() |  | ||||||
| 	if c.GetString("base_url") != "" { |  | ||||||
| 		baseURL = c.GetString("base_url") |  | ||||||
| 	} |  | ||||||
| 	fullRequestURL := util.GetFullRequestURL(baseURL, requestURL, channelType) |  | ||||||
| 	switch apiType { |  | ||||||
| 	case APITypeOpenAI: |  | ||||||
| 		if channelType == common.ChannelTypeAzure { |  | ||||||
| 			// https://learn.microsoft.com/en-us/azure/cognitive-services/openai/chatgpt-quickstart?pivots=rest-api&tabs=command-line#rest-api |  | ||||||
| 			apiVersion := util.GetAPIVersion(c) |  | ||||||
| 			requestURL := strings.Split(requestURL, "?")[0] |  | ||||||
| 			requestURL = fmt.Sprintf("%s?api-version=%s", requestURL, apiVersion) |  | ||||||
| 			baseURL = c.GetString("base_url") |  | ||||||
| 			task := strings.TrimPrefix(requestURL, "/v1/") |  | ||||||
| 			model_ := textRequest.Model |  | ||||||
| 			model_ = strings.Replace(model_, ".", "", -1) |  | ||||||
| 			// https://github.com/songquanpeng/one-api/issues/67 |  | ||||||
| 			model_ = strings.TrimSuffix(model_, "-0301") |  | ||||||
| 			model_ = strings.TrimSuffix(model_, "-0314") |  | ||||||
| 			model_ = strings.TrimSuffix(model_, "-0613") |  | ||||||
|  |  | ||||||
| 			requestURL = fmt.Sprintf("/openai/deployments/%s/%s", model_, task) |  | ||||||
| 			fullRequestURL = util.GetFullRequestURL(baseURL, requestURL, channelType) |  | ||||||
| 		} |  | ||||||
| 	case APITypeClaude: |  | ||||||
| 		fullRequestURL = "https://api.anthropic.com/v1/complete" |  | ||||||
| 		if baseURL != "" { |  | ||||||
| 			fullRequestURL = fmt.Sprintf("%s/v1/complete", baseURL) |  | ||||||
| 		} |  | ||||||
| 	case APITypeBaidu: |  | ||||||
| 		switch textRequest.Model { |  | ||||||
| 		case "ERNIE-Bot": |  | ||||||
| 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/completions" |  | ||||||
| 		case "ERNIE-Bot-turbo": |  | ||||||
| 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/eb-instant" |  | ||||||
| 		case "ERNIE-Bot-4": |  | ||||||
| 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/completions_pro" |  | ||||||
| 		case "BLOOMZ-7B": |  | ||||||
| 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/bloomz_7b1" |  | ||||||
| 		case "Embedding-V1": |  | ||||||
| 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/embeddings/embedding-v1" |  | ||||||
| 		} |  | ||||||
| 		apiKey := c.Request.Header.Get("Authorization") |  | ||||||
| 		apiKey = strings.TrimPrefix(apiKey, "Bearer ") |  | ||||||
| 		var err error |  | ||||||
| 		if apiKey, err = baidu.GetAccessToken(apiKey); err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "invalid_baidu_config", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		fullRequestURL += "?access_token=" + apiKey |  | ||||||
| 	case APITypePaLM: |  | ||||||
| 		fullRequestURL = "https://generativelanguage.googleapis.com/v1beta2/models/chat-bison-001:generateMessage" |  | ||||||
| 		if baseURL != "" { |  | ||||||
| 			fullRequestURL = fmt.Sprintf("%s/v1beta2/models/chat-bison-001:generateMessage", baseURL) |  | ||||||
| 		} |  | ||||||
| 	case APITypeGemini: |  | ||||||
| 		requestBaseURL := "https://generativelanguage.googleapis.com" |  | ||||||
| 		if baseURL != "" { |  | ||||||
| 			requestBaseURL = baseURL |  | ||||||
| 		} |  | ||||||
| 		version := "v1" |  | ||||||
| 		if c.GetString("api_version") != "" { |  | ||||||
| 			version = c.GetString("api_version") |  | ||||||
| 		} |  | ||||||
| 		action := "generateContent" |  | ||||||
| 		if textRequest.Stream { |  | ||||||
| 			action = "streamGenerateContent" |  | ||||||
| 		} |  | ||||||
| 		fullRequestURL = fmt.Sprintf("%s/%s/models/%s:%s", requestBaseURL, version, textRequest.Model, action) |  | ||||||
| 	case APITypeZhipu: |  | ||||||
| 		method := "invoke" |  | ||||||
| 		if textRequest.Stream { |  | ||||||
| 			method = "sse-invoke" |  | ||||||
| 		} |  | ||||||
| 		fullRequestURL = fmt.Sprintf("https://open.bigmodel.cn/api/paas/v3/model-api/%s/%s", textRequest.Model, method) |  | ||||||
| 	case APITypeAli: |  | ||||||
| 		fullRequestURL = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation" |  | ||||||
| 		if relayMode == constant.RelayModeEmbeddings { |  | ||||||
| 			fullRequestURL = "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding" |  | ||||||
| 		} |  | ||||||
| 	case APITypeTencent: |  | ||||||
| 		fullRequestURL = "https://hunyuan.cloud.tencent.com/hyllm/v1/chat/completions" |  | ||||||
| 	case APITypeAIProxyLibrary: |  | ||||||
| 		fullRequestURL = fmt.Sprintf("%s/api/library/ask", baseURL) |  | ||||||
| 	} | 	} | ||||||
| 	var promptTokens int | 	var promptTokens int | ||||||
| 	var completionTokens int | 	var completionTokens int | ||||||
| @@ -212,22 +53,22 @@ func RelayTextHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode | |||||||
| 	case constant.RelayModeModerations: | 	case constant.RelayModeModerations: | ||||||
| 		promptTokens = openai.CountTokenInput(textRequest.Input, textRequest.Model) | 		promptTokens = openai.CountTokenInput(textRequest.Input, textRequest.Model) | ||||||
| 	} | 	} | ||||||
| 	preConsumedTokens := common.PreConsumedQuota | 	preConsumedTokens := config.PreConsumedQuota | ||||||
| 	if textRequest.MaxTokens != 0 { | 	if textRequest.MaxTokens != 0 { | ||||||
| 		preConsumedTokens = promptTokens + textRequest.MaxTokens | 		preConsumedTokens = promptTokens + textRequest.MaxTokens | ||||||
| 	} | 	} | ||||||
| 	modelRatio := common.GetModelRatio(textRequest.Model) | 	modelRatio := common.GetModelRatio(textRequest.Model) | ||||||
| 	groupRatio := common.GetGroupRatio(group) | 	groupRatio := common.GetGroupRatio(meta.Group) | ||||||
| 	ratio := modelRatio * groupRatio | 	ratio := modelRatio * groupRatio | ||||||
| 	preConsumedQuota := int(float64(preConsumedTokens) * ratio) | 	preConsumedQuota := int(float64(preConsumedTokens) * ratio) | ||||||
| 	userQuota, err := model.CacheGetUserQuota(userId) | 	userQuota, err := model.CacheGetUserQuota(meta.UserId) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "get_user_quota_failed", http.StatusInternalServerError) | 		return openai.ErrorWrapper(err, "get_user_quota_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 	if userQuota-preConsumedQuota < 0 { | 	if userQuota-preConsumedQuota < 0 { | ||||||
| 		return openai.ErrorWrapper(errors.New("user quota is not enough"), "insufficient_user_quota", http.StatusForbidden) | 		return openai.ErrorWrapper(errors.New("user quota is not enough"), "insufficient_user_quota", http.StatusForbidden) | ||||||
| 	} | 	} | ||||||
| 	err = model.CacheDecreaseUserQuota(userId, preConsumedQuota) | 	err = model.CacheDecreaseUserQuota(meta.UserId, preConsumedQuota) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "decrease_user_quota_failed", http.StatusInternalServerError) | 		return openai.ErrorWrapper(err, "decrease_user_quota_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| @@ -235,165 +76,28 @@ func RelayTextHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode | |||||||
| 		// in this case, we do not pre-consume quota | 		// in this case, we do not pre-consume quota | ||||||
| 		// because the user has enough quota | 		// because the user has enough quota | ||||||
| 		preConsumedQuota = 0 | 		preConsumedQuota = 0 | ||||||
| 		common.LogInfo(c.Request.Context(), fmt.Sprintf("user %d has enough quota %d, trusted and no need to pre-consume", userId, userQuota)) | 		logger.Info(c.Request.Context(), fmt.Sprintf("user %d has enough quota %d, trusted and no need to pre-consume", meta.UserId, userQuota)) | ||||||
| 	} | 	} | ||||||
| 	if preConsumedQuota > 0 { | 	if preConsumedQuota > 0 { | ||||||
| 		err := model.PreConsumeTokenQuota(tokenId, preConsumedQuota) | 		err := model.PreConsumeTokenQuota(meta.TokenId, preConsumedQuota) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return openai.ErrorWrapper(err, "pre_consume_token_quota_failed", http.StatusForbidden) | 			return openai.ErrorWrapper(err, "pre_consume_token_quota_failed", http.StatusForbidden) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	var requestBody io.Reader | 	requestBody, err := GetRequestBody(c, textRequest, isModelMapped, apiType, relayMode) | ||||||
| 	if isModelMapped { | 	if err != nil { | ||||||
| 		jsonStr, err := json.Marshal(textRequest) | 		return openai.ErrorWrapper(err, "get_request_body_failed", http.StatusInternalServerError) | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		requestBody = bytes.NewBuffer(jsonStr) |  | ||||||
| 	} else { |  | ||||||
| 		requestBody = c.Request.Body |  | ||||||
| 	} | 	} | ||||||
| 	switch apiType { |  | ||||||
| 	case APITypeClaude: |  | ||||||
| 		claudeRequest := anthropic.ConvertRequest(textRequest) |  | ||||||
| 		jsonStr, err := json.Marshal(claudeRequest) |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		requestBody = bytes.NewBuffer(jsonStr) |  | ||||||
| 	case APITypeBaidu: |  | ||||||
| 		var jsonData []byte |  | ||||||
| 		var err error |  | ||||||
| 		switch relayMode { |  | ||||||
| 		case constant.RelayModeEmbeddings: |  | ||||||
| 			baiduEmbeddingRequest := baidu.ConvertEmbeddingRequest(textRequest) |  | ||||||
| 			jsonData, err = json.Marshal(baiduEmbeddingRequest) |  | ||||||
| 		default: |  | ||||||
| 			baiduRequest := baidu.ConvertRequest(textRequest) |  | ||||||
| 			jsonData, err = json.Marshal(baiduRequest) |  | ||||||
| 		} |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		requestBody = bytes.NewBuffer(jsonData) |  | ||||||
| 	case APITypePaLM: |  | ||||||
| 		palmRequest := google.ConvertPaLMRequest(textRequest) |  | ||||||
| 		jsonStr, err := json.Marshal(palmRequest) |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		requestBody = bytes.NewBuffer(jsonStr) |  | ||||||
| 	case APITypeGemini: |  | ||||||
| 		geminiChatRequest := google.ConvertGeminiRequest(textRequest) |  | ||||||
| 		jsonStr, err := json.Marshal(geminiChatRequest) |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		requestBody = bytes.NewBuffer(jsonStr) |  | ||||||
| 	case APITypeZhipu: |  | ||||||
| 		zhipuRequest := zhipu.ConvertRequest(textRequest) |  | ||||||
| 		jsonStr, err := json.Marshal(zhipuRequest) |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		requestBody = bytes.NewBuffer(jsonStr) |  | ||||||
| 	case APITypeAli: |  | ||||||
| 		var jsonStr []byte |  | ||||||
| 		var err error |  | ||||||
| 		switch relayMode { |  | ||||||
| 		case constant.RelayModeEmbeddings: |  | ||||||
| 			aliEmbeddingRequest := ali.ConvertEmbeddingRequest(textRequest) |  | ||||||
| 			jsonStr, err = json.Marshal(aliEmbeddingRequest) |  | ||||||
| 		default: |  | ||||||
| 			aliRequest := ali.ConvertRequest(textRequest) |  | ||||||
| 			jsonStr, err = json.Marshal(aliRequest) |  | ||||||
| 		} |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		requestBody = bytes.NewBuffer(jsonStr) |  | ||||||
| 	case APITypeTencent: |  | ||||||
| 		apiKey := c.Request.Header.Get("Authorization") |  | ||||||
| 		apiKey = strings.TrimPrefix(apiKey, "Bearer ") |  | ||||||
| 		appId, secretId, secretKey, err := tencent.ParseConfig(apiKey) |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "invalid_tencent_config", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		tencentRequest := tencent.ConvertRequest(textRequest) |  | ||||||
| 		tencentRequest.AppId = appId |  | ||||||
| 		tencentRequest.SecretId = secretId |  | ||||||
| 		jsonStr, err := json.Marshal(tencentRequest) |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		sign := tencent.GetSign(*tencentRequest, secretKey) |  | ||||||
| 		c.Request.Header.Set("Authorization", sign) |  | ||||||
| 		requestBody = bytes.NewBuffer(jsonStr) |  | ||||||
| 	case APITypeAIProxyLibrary: |  | ||||||
| 		aiProxyLibraryRequest := aiproxy.ConvertRequest(textRequest) |  | ||||||
| 		aiProxyLibraryRequest.LibraryId = c.GetString("library_id") |  | ||||||
| 		jsonStr, err := json.Marshal(aiProxyLibraryRequest) |  | ||||||
| 		if err != nil { |  | ||||||
| 			return openai.ErrorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) |  | ||||||
| 		} |  | ||||||
| 		requestBody = bytes.NewBuffer(jsonStr) |  | ||||||
| 	} |  | ||||||
|  |  | ||||||
| 	var req *http.Request | 	var req *http.Request | ||||||
| 	var resp *http.Response | 	var resp *http.Response | ||||||
| 	isStream := textRequest.Stream | 	isStream := textRequest.Stream | ||||||
|  |  | ||||||
| 	if apiType != APITypeXunfei { // cause xunfei use websocket | 	if apiType != constant.APITypeXunfei { // cause xunfei use websocket | ||||||
| 		req, err = http.NewRequest(c.Request.Method, fullRequestURL, requestBody) | 		req, err = http.NewRequest(c.Request.Method, fullRequestURL, requestBody) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return openai.ErrorWrapper(err, "new_request_failed", http.StatusInternalServerError) | 			return openai.ErrorWrapper(err, "new_request_failed", http.StatusInternalServerError) | ||||||
| 		} | 		} | ||||||
| 		apiKey := c.Request.Header.Get("Authorization") | 		SetupRequestHeaders(c, req, apiType, meta, isStream) | ||||||
| 		apiKey = strings.TrimPrefix(apiKey, "Bearer ") |  | ||||||
| 		switch apiType { |  | ||||||
| 		case APITypeOpenAI: |  | ||||||
| 			if channelType == common.ChannelTypeAzure { |  | ||||||
| 				req.Header.Set("api-key", apiKey) |  | ||||||
| 			} else { |  | ||||||
| 				req.Header.Set("Authorization", c.Request.Header.Get("Authorization")) |  | ||||||
| 				if channelType == common.ChannelTypeOpenRouter { |  | ||||||
| 					req.Header.Set("HTTP-Referer", "https://github.com/songquanpeng/one-api") |  | ||||||
| 					req.Header.Set("X-Title", "One API") |  | ||||||
| 				} |  | ||||||
| 			} |  | ||||||
| 		case APITypeClaude: |  | ||||||
| 			req.Header.Set("x-api-key", apiKey) |  | ||||||
| 			anthropicVersion := c.Request.Header.Get("anthropic-version") |  | ||||||
| 			if anthropicVersion == "" { |  | ||||||
| 				anthropicVersion = "2023-06-01" |  | ||||||
| 			} |  | ||||||
| 			req.Header.Set("anthropic-version", anthropicVersion) |  | ||||||
| 		case APITypeZhipu: |  | ||||||
| 			token := zhipu.GetToken(apiKey) |  | ||||||
| 			req.Header.Set("Authorization", token) |  | ||||||
| 		case APITypeAli: |  | ||||||
| 			req.Header.Set("Authorization", "Bearer "+apiKey) |  | ||||||
| 			if textRequest.Stream { |  | ||||||
| 				req.Header.Set("X-DashScope-SSE", "enable") |  | ||||||
| 			} |  | ||||||
| 			if c.GetString("plugin") != "" { |  | ||||||
| 				req.Header.Set("X-DashScope-Plugin", c.GetString("plugin")) |  | ||||||
| 			} |  | ||||||
| 		case APITypeTencent: |  | ||||||
| 			req.Header.Set("Authorization", apiKey) |  | ||||||
| 		case APITypePaLM: |  | ||||||
| 			req.Header.Set("x-goog-api-key", apiKey) |  | ||||||
| 		case APITypeGemini: |  | ||||||
| 			req.Header.Set("x-goog-api-key", apiKey) |  | ||||||
| 		default: |  | ||||||
| 			req.Header.Set("Authorization", "Bearer "+apiKey) |  | ||||||
| 		} |  | ||||||
| 		req.Header.Set("Content-Type", c.Request.Header.Get("Content-Type")) |  | ||||||
| 		req.Header.Set("Accept", c.Request.Header.Get("Accept")) |  | ||||||
| 		if isStream && c.Request.Header.Get("Accept") == "" { |  | ||||||
| 			req.Header.Set("Accept", "text/event-stream") |  | ||||||
| 		} |  | ||||||
| 		//req.Header.Set("Connection", c.Request.Header.Get("Connection")) |  | ||||||
| 		resp, err = util.HTTPClient.Do(req) | 		resp, err = util.HTTPClient.Do(req) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return openai.ErrorWrapper(err, "do_request_failed", http.StatusInternalServerError) | 			return openai.ErrorWrapper(err, "do_request_failed", http.StatusInternalServerError) | ||||||
| @@ -409,29 +113,31 @@ func RelayTextHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode | |||||||
| 		isStream = isStream || strings.HasPrefix(resp.Header.Get("Content-Type"), "text/event-stream") | 		isStream = isStream || strings.HasPrefix(resp.Header.Get("Content-Type"), "text/event-stream") | ||||||
|  |  | ||||||
| 		if resp.StatusCode != http.StatusOK { | 		if resp.StatusCode != http.StatusOK { | ||||||
| 			if preConsumedQuota != 0 { | 			util.ReturnPreConsumedQuota(ctx, preConsumedQuota, meta.TokenId) | ||||||
| 				go func(ctx context.Context) { |  | ||||||
| 					// return pre-consumed quota |  | ||||||
| 					err := model.PostConsumeTokenQuota(tokenId, -preConsumedQuota) |  | ||||||
| 					if err != nil { |  | ||||||
| 						common.LogError(ctx, "error return pre-consumed quota: "+err.Error()) |  | ||||||
| 					} |  | ||||||
| 				}(c.Request.Context()) |  | ||||||
| 			} |  | ||||||
| 			return util.RelayErrorHandler(resp) | 			return util.RelayErrorHandler(resp) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	var textResponse openai.SlimTextResponse | 	var respErr *openai.ErrorWithStatusCode | ||||||
| 	tokenName := c.GetString("token_name") | 	var usage *openai.Usage | ||||||
|  |  | ||||||
| 	defer func(ctx context.Context) { | 	defer func(ctx context.Context) { | ||||||
| 		// c.Writer.Flush() | 		// Why we use defer here? Because if error happened, we will have to return the pre-consumed quota. | ||||||
|  | 		if respErr != nil { | ||||||
|  | 			logger.Errorf(ctx, "respErr is not nil: %+v", respErr) | ||||||
|  | 			util.ReturnPreConsumedQuota(ctx, preConsumedQuota, meta.TokenId) | ||||||
|  | 			return | ||||||
|  | 		} | ||||||
|  | 		if usage == nil { | ||||||
|  | 			logger.Error(ctx, "usage is nil, which is unexpected") | ||||||
|  | 			return | ||||||
|  | 		} | ||||||
|  |  | ||||||
| 		go func() { | 		go func() { | ||||||
| 			quota := 0 | 			quota := 0 | ||||||
| 			completionRatio := common.GetCompletionRatio(textRequest.Model) | 			completionRatio := common.GetCompletionRatio(textRequest.Model) | ||||||
| 			promptTokens = textResponse.Usage.PromptTokens | 			promptTokens = usage.PromptTokens | ||||||
| 			completionTokens = textResponse.Usage.CompletionTokens | 			completionTokens = usage.CompletionTokens | ||||||
| 			quota = int(math.Ceil((float64(promptTokens) + float64(completionTokens)*completionRatio) * ratio)) | 			quota = int(math.Ceil((float64(promptTokens) + float64(completionTokens)*completionRatio) * ratio)) | ||||||
| 			if ratio != 0 && quota <= 0 { | 			if ratio != 0 && quota <= 0 { | ||||||
| 				quota = 1 | 				quota = 1 | ||||||
| @@ -443,239 +149,25 @@ func RelayTextHelper(c *gin.Context, relayMode int) *openai.ErrorWithStatusCode | |||||||
| 				quota = 0 | 				quota = 0 | ||||||
| 			} | 			} | ||||||
| 			quotaDelta := quota - preConsumedQuota | 			quotaDelta := quota - preConsumedQuota | ||||||
| 			err := model.PostConsumeTokenQuota(tokenId, quotaDelta) | 			err := model.PostConsumeTokenQuota(meta.TokenId, quotaDelta) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.LogError(ctx, "error consuming token remain quota: "+err.Error()) | 				logger.Error(ctx, "error consuming token remain quota: "+err.Error()) | ||||||
| 			} | 			} | ||||||
| 			err = model.CacheUpdateUserQuota(userId) | 			err = model.CacheUpdateUserQuota(meta.UserId) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				common.LogError(ctx, "error update user quota cache: "+err.Error()) | 				logger.Error(ctx, "error update user quota cache: "+err.Error()) | ||||||
| 			} | 			} | ||||||
| 			if quota != 0 { | 			if quota != 0 { | ||||||
| 				logContent := fmt.Sprintf("模型倍率 %.2f,分组倍率 %.2f", modelRatio, groupRatio) | 				logContent := fmt.Sprintf("模型倍率 %.2f,分组倍率 %.2f", modelRatio, groupRatio) | ||||||
| 				model.RecordConsumeLog(ctx, userId, channelId, promptTokens, completionTokens, textRequest.Model, tokenName, quota, logContent) | 				model.RecordConsumeLog(ctx, meta.UserId, meta.ChannelId, promptTokens, completionTokens, textRequest.Model, meta.TokenName, quota, logContent) | ||||||
| 				model.UpdateUserUsedQuotaAndRequestCount(userId, quota) | 				model.UpdateUserUsedQuotaAndRequestCount(meta.UserId, quota) | ||||||
| 				model.UpdateChannelUsedQuota(channelId, quota) | 				model.UpdateChannelUsedQuota(meta.ChannelId, quota) | ||||||
| 			} | 			} | ||||||
|  |  | ||||||
| 		}() | 		}() | ||||||
| 	}(c.Request.Context()) | 	}(ctx) | ||||||
| 	switch apiType { | 	usage, respErr = DoResponse(c, &textRequest, resp, relayMode, apiType, isStream, promptTokens) | ||||||
| 	case APITypeOpenAI: | 	if respErr != nil { | ||||||
| 		if isStream { | 		return respErr | ||||||
| 			err, responseText := openai.StreamHandler(c, resp, relayMode) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			textResponse.Usage.PromptTokens = promptTokens |  | ||||||
| 			textResponse.Usage.CompletionTokens = openai.CountTokenText(responseText, textRequest.Model) |  | ||||||
| 			return nil |  | ||||||
| 		} else { |  | ||||||
| 			err, usage := openai.Handler(c, resp, promptTokens, textRequest.Model) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} |  | ||||||
| 	case APITypeClaude: |  | ||||||
| 		if isStream { |  | ||||||
| 			err, responseText := anthropic.StreamHandler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			textResponse.Usage.PromptTokens = promptTokens |  | ||||||
| 			textResponse.Usage.CompletionTokens = openai.CountTokenText(responseText, textRequest.Model) |  | ||||||
| 			return nil |  | ||||||
| 		} else { |  | ||||||
| 			err, usage := anthropic.Handler(c, resp, promptTokens, textRequest.Model) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} |  | ||||||
| 	case APITypeBaidu: |  | ||||||
| 		if isStream { |  | ||||||
| 			err, usage := baidu.StreamHandler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} else { |  | ||||||
| 			var err *openai.ErrorWithStatusCode |  | ||||||
| 			var usage *openai.Usage |  | ||||||
| 			switch relayMode { |  | ||||||
| 			case constant.RelayModeEmbeddings: |  | ||||||
| 				err, usage = baidu.EmbeddingHandler(c, resp) |  | ||||||
| 			default: |  | ||||||
| 				err, usage = baidu.Handler(c, resp) |  | ||||||
| 			} |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} |  | ||||||
| 	case APITypePaLM: |  | ||||||
| 		if textRequest.Stream { // PaLM2 API does not support stream |  | ||||||
| 			err, responseText := google.PaLMStreamHandler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			textResponse.Usage.PromptTokens = promptTokens |  | ||||||
| 			textResponse.Usage.CompletionTokens = openai.CountTokenText(responseText, textRequest.Model) |  | ||||||
| 			return nil |  | ||||||
| 		} else { |  | ||||||
| 			err, usage := google.PaLMHandler(c, resp, promptTokens, textRequest.Model) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} |  | ||||||
| 	case APITypeGemini: |  | ||||||
| 		if textRequest.Stream { |  | ||||||
| 			err, responseText := google.StreamHandler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			textResponse.Usage.PromptTokens = promptTokens |  | ||||||
| 			textResponse.Usage.CompletionTokens = openai.CountTokenText(responseText, textRequest.Model) |  | ||||||
| 			return nil |  | ||||||
| 		} else { |  | ||||||
| 			err, usage := google.GeminiHandler(c, resp, promptTokens, textRequest.Model) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} |  | ||||||
| 	case APITypeZhipu: |  | ||||||
| 		if isStream { |  | ||||||
| 			err, usage := zhipu.StreamHandler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			// zhipu's API does not return prompt tokens & completion tokens |  | ||||||
| 			textResponse.Usage.PromptTokens = textResponse.Usage.TotalTokens |  | ||||||
| 			return nil |  | ||||||
| 		} else { |  | ||||||
| 			err, usage := zhipu.Handler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			// zhipu's API does not return prompt tokens & completion tokens |  | ||||||
| 			textResponse.Usage.PromptTokens = textResponse.Usage.TotalTokens |  | ||||||
| 			return nil |  | ||||||
| 		} |  | ||||||
| 	case APITypeAli: |  | ||||||
| 		if isStream { |  | ||||||
| 			err, usage := ali.StreamHandler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} else { |  | ||||||
| 			var err *openai.ErrorWithStatusCode |  | ||||||
| 			var usage *openai.Usage |  | ||||||
| 			switch relayMode { |  | ||||||
| 			case constant.RelayModeEmbeddings: |  | ||||||
| 				err, usage = ali.EmbeddingHandler(c, resp) |  | ||||||
| 			default: |  | ||||||
| 				err, usage = ali.Handler(c, resp) |  | ||||||
| 			} |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} |  | ||||||
| 	case APITypeXunfei: |  | ||||||
| 		auth := c.Request.Header.Get("Authorization") |  | ||||||
| 		auth = strings.TrimPrefix(auth, "Bearer ") |  | ||||||
| 		splits := strings.Split(auth, "|") |  | ||||||
| 		if len(splits) != 3 { |  | ||||||
| 			return openai.ErrorWrapper(errors.New("invalid auth"), "invalid_auth", http.StatusBadRequest) |  | ||||||
| 		} |  | ||||||
| 		var err *openai.ErrorWithStatusCode |  | ||||||
| 		var usage *openai.Usage |  | ||||||
| 		if isStream { |  | ||||||
| 			err, usage = xunfei.StreamHandler(c, textRequest, splits[0], splits[1], splits[2]) |  | ||||||
| 		} else { |  | ||||||
| 			err, usage = xunfei.Handler(c, textRequest, splits[0], splits[1], splits[2]) |  | ||||||
| 		} |  | ||||||
| 		if err != nil { |  | ||||||
| 			return err |  | ||||||
| 		} |  | ||||||
| 		if usage != nil { |  | ||||||
| 			textResponse.Usage = *usage |  | ||||||
| 		} |  | ||||||
| 		return nil |  | ||||||
| 	case APITypeAIProxyLibrary: |  | ||||||
| 		if isStream { |  | ||||||
| 			err, usage := aiproxy.StreamHandler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} else { |  | ||||||
| 			err, usage := aiproxy.Handler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} |  | ||||||
| 	case APITypeTencent: |  | ||||||
| 		if isStream { |  | ||||||
| 			err, responseText := tencent.StreamHandler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			textResponse.Usage.PromptTokens = promptTokens |  | ||||||
| 			textResponse.Usage.CompletionTokens = openai.CountTokenText(responseText, textRequest.Model) |  | ||||||
| 			return nil |  | ||||||
| 		} else { |  | ||||||
| 			err, usage := tencent.Handler(c, resp) |  | ||||||
| 			if err != nil { |  | ||||||
| 				return err |  | ||||||
| 			} |  | ||||||
| 			if usage != nil { |  | ||||||
| 				textResponse.Usage = *usage |  | ||||||
| 			} |  | ||||||
| 			return nil |  | ||||||
| 		} |  | ||||||
| 	default: |  | ||||||
| 		return openai.ErrorWrapper(errors.New("unknown api type"), "unknown_api_type", http.StatusInternalServerError) |  | ||||||
| 	} | 	} | ||||||
|  | 	return nil | ||||||
| } | } | ||||||
|   | |||||||
							
								
								
									
										337
									
								
								relay/controller/util.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										337
									
								
								relay/controller/util.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,337 @@ | |||||||
|  | package controller | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"bytes" | ||||||
|  | 	"encoding/json" | ||||||
|  | 	"errors" | ||||||
|  | 	"fmt" | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"io" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/common/helper" | ||||||
|  | 	"one-api/relay/channel/aiproxy" | ||||||
|  | 	"one-api/relay/channel/ali" | ||||||
|  | 	"one-api/relay/channel/anthropic" | ||||||
|  | 	"one-api/relay/channel/baidu" | ||||||
|  | 	"one-api/relay/channel/google" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | 	"one-api/relay/channel/tencent" | ||||||
|  | 	"one-api/relay/channel/xunfei" | ||||||
|  | 	"one-api/relay/channel/zhipu" | ||||||
|  | 	"one-api/relay/constant" | ||||||
|  | 	"one-api/relay/util" | ||||||
|  | 	"strings" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func GetRequestURL(requestURL string, apiType int, relayMode int, meta *util.RelayMeta, textRequest *openai.GeneralOpenAIRequest) (string, error) { | ||||||
|  | 	fullRequestURL := util.GetFullRequestURL(meta.BaseURL, requestURL, meta.ChannelType) | ||||||
|  | 	switch apiType { | ||||||
|  | 	case constant.APITypeOpenAI: | ||||||
|  | 		if meta.ChannelType == common.ChannelTypeAzure { | ||||||
|  | 			// https://learn.microsoft.com/en-us/azure/cognitive-services/openai/chatgpt-quickstart?pivots=rest-api&tabs=command-line#rest-api | ||||||
|  | 			requestURL := strings.Split(requestURL, "?")[0] | ||||||
|  | 			requestURL = fmt.Sprintf("%s?api-version=%s", requestURL, meta.APIVersion) | ||||||
|  | 			task := strings.TrimPrefix(requestURL, "/v1/") | ||||||
|  | 			model_ := textRequest.Model | ||||||
|  | 			model_ = strings.Replace(model_, ".", "", -1) | ||||||
|  | 			// https://github.com/songquanpeng/one-api/issues/67 | ||||||
|  | 			model_ = strings.TrimSuffix(model_, "-0301") | ||||||
|  | 			model_ = strings.TrimSuffix(model_, "-0314") | ||||||
|  | 			model_ = strings.TrimSuffix(model_, "-0613") | ||||||
|  |  | ||||||
|  | 			requestURL = fmt.Sprintf("/openai/deployments/%s/%s", model_, task) | ||||||
|  | 			fullRequestURL = util.GetFullRequestURL(meta.BaseURL, requestURL, meta.ChannelType) | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeClaude: | ||||||
|  | 		fullRequestURL = fmt.Sprintf("%s/v1/complete", meta.BaseURL) | ||||||
|  | 	case constant.APITypeBaidu: | ||||||
|  | 		switch textRequest.Model { | ||||||
|  | 		case "ERNIE-Bot": | ||||||
|  | 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/completions" | ||||||
|  | 		case "ERNIE-Bot-turbo": | ||||||
|  | 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/eb-instant" | ||||||
|  | 		case "ERNIE-Bot-4": | ||||||
|  | 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/completions_pro" | ||||||
|  | 		case "BLOOMZ-7B": | ||||||
|  | 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/bloomz_7b1" | ||||||
|  | 		case "Embedding-V1": | ||||||
|  | 			fullRequestURL = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/embeddings/embedding-v1" | ||||||
|  | 		} | ||||||
|  | 		var accessToken string | ||||||
|  | 		var err error | ||||||
|  | 		if accessToken, err = baidu.GetAccessToken(meta.APIKey); err != nil { | ||||||
|  | 			return "", fmt.Errorf("failed to get baidu access token: %w", err) | ||||||
|  | 		} | ||||||
|  | 		fullRequestURL += "?access_token=" + accessToken | ||||||
|  | 	case constant.APITypePaLM: | ||||||
|  | 		fullRequestURL = fmt.Sprintf("%s/v1beta2/models/chat-bison-001:generateMessage", meta.BaseURL) | ||||||
|  | 	case constant.APITypeGemini: | ||||||
|  | 		version := helper.AssignOrDefault(meta.APIVersion, "v1") | ||||||
|  | 		action := "generateContent" | ||||||
|  | 		if textRequest.Stream { | ||||||
|  | 			action = "streamGenerateContent" | ||||||
|  | 		} | ||||||
|  | 		fullRequestURL = fmt.Sprintf("%s/%s/models/%s:%s", meta.BaseURL, version, textRequest.Model, action) | ||||||
|  | 	case constant.APITypeZhipu: | ||||||
|  | 		method := "invoke" | ||||||
|  | 		if textRequest.Stream { | ||||||
|  | 			method = "sse-invoke" | ||||||
|  | 		} | ||||||
|  | 		fullRequestURL = fmt.Sprintf("https://open.bigmodel.cn/api/paas/v3/model-api/%s/%s", textRequest.Model, method) | ||||||
|  | 	case constant.APITypeAli: | ||||||
|  | 		fullRequestURL = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation" | ||||||
|  | 		if relayMode == constant.RelayModeEmbeddings { | ||||||
|  | 			fullRequestURL = "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding" | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeTencent: | ||||||
|  | 		fullRequestURL = "https://hunyuan.cloud.tencent.com/hyllm/v1/chat/completions" | ||||||
|  | 	case constant.APITypeAIProxyLibrary: | ||||||
|  | 		fullRequestURL = fmt.Sprintf("%s/api/library/ask", meta.BaseURL) | ||||||
|  | 	} | ||||||
|  | 	return fullRequestURL, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GetRequestBody(c *gin.Context, textRequest openai.GeneralOpenAIRequest, isModelMapped bool, apiType int, relayMode int) (io.Reader, error) { | ||||||
|  | 	var requestBody io.Reader | ||||||
|  | 	if isModelMapped { | ||||||
|  | 		jsonStr, err := json.Marshal(textRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	} else { | ||||||
|  | 		requestBody = c.Request.Body | ||||||
|  | 	} | ||||||
|  | 	switch apiType { | ||||||
|  | 	case constant.APITypeClaude: | ||||||
|  | 		claudeRequest := anthropic.ConvertRequest(textRequest) | ||||||
|  | 		jsonStr, err := json.Marshal(claudeRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case constant.APITypeBaidu: | ||||||
|  | 		var jsonData []byte | ||||||
|  | 		var err error | ||||||
|  | 		switch relayMode { | ||||||
|  | 		case constant.RelayModeEmbeddings: | ||||||
|  | 			baiduEmbeddingRequest := baidu.ConvertEmbeddingRequest(textRequest) | ||||||
|  | 			jsonData, err = json.Marshal(baiduEmbeddingRequest) | ||||||
|  | 		default: | ||||||
|  | 			baiduRequest := baidu.ConvertRequest(textRequest) | ||||||
|  | 			jsonData, err = json.Marshal(baiduRequest) | ||||||
|  | 		} | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonData) | ||||||
|  | 	case constant.APITypePaLM: | ||||||
|  | 		palmRequest := google.ConvertPaLMRequest(textRequest) | ||||||
|  | 		jsonStr, err := json.Marshal(palmRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case constant.APITypeGemini: | ||||||
|  | 		geminiChatRequest := google.ConvertGeminiRequest(textRequest) | ||||||
|  | 		jsonStr, err := json.Marshal(geminiChatRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case constant.APITypeZhipu: | ||||||
|  | 		zhipuRequest := zhipu.ConvertRequest(textRequest) | ||||||
|  | 		jsonStr, err := json.Marshal(zhipuRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case constant.APITypeAli: | ||||||
|  | 		var jsonStr []byte | ||||||
|  | 		var err error | ||||||
|  | 		switch relayMode { | ||||||
|  | 		case constant.RelayModeEmbeddings: | ||||||
|  | 			aliEmbeddingRequest := ali.ConvertEmbeddingRequest(textRequest) | ||||||
|  | 			jsonStr, err = json.Marshal(aliEmbeddingRequest) | ||||||
|  | 		default: | ||||||
|  | 			aliRequest := ali.ConvertRequest(textRequest) | ||||||
|  | 			jsonStr, err = json.Marshal(aliRequest) | ||||||
|  | 		} | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case constant.APITypeTencent: | ||||||
|  | 		apiKey := c.Request.Header.Get("Authorization") | ||||||
|  | 		apiKey = strings.TrimPrefix(apiKey, "Bearer ") | ||||||
|  | 		appId, secretId, secretKey, err := tencent.ParseConfig(apiKey) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		tencentRequest := tencent.ConvertRequest(textRequest) | ||||||
|  | 		tencentRequest.AppId = appId | ||||||
|  | 		tencentRequest.SecretId = secretId | ||||||
|  | 		jsonStr, err := json.Marshal(tencentRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		sign := tencent.GetSign(*tencentRequest, secretKey) | ||||||
|  | 		c.Request.Header.Set("Authorization", sign) | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case constant.APITypeAIProxyLibrary: | ||||||
|  | 		aiProxyLibraryRequest := aiproxy.ConvertRequest(textRequest) | ||||||
|  | 		aiProxyLibraryRequest.LibraryId = c.GetString("library_id") | ||||||
|  | 		jsonStr, err := json.Marshal(aiProxyLibraryRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return nil, err | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	} | ||||||
|  | 	return requestBody, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func SetupRequestHeaders(c *gin.Context, req *http.Request, apiType int, meta *util.RelayMeta, isStream bool) { | ||||||
|  | 	SetupAuthHeaders(c, req, apiType, meta, isStream) | ||||||
|  | 	req.Header.Set("Content-Type", c.Request.Header.Get("Content-Type")) | ||||||
|  | 	req.Header.Set("Accept", c.Request.Header.Get("Accept")) | ||||||
|  | 	if isStream && c.Request.Header.Get("Accept") == "" { | ||||||
|  | 		req.Header.Set("Accept", "text/event-stream") | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func SetupAuthHeaders(c *gin.Context, req *http.Request, apiType int, meta *util.RelayMeta, isStream bool) { | ||||||
|  | 	apiKey := meta.APIKey | ||||||
|  | 	switch apiType { | ||||||
|  | 	case constant.APITypeOpenAI: | ||||||
|  | 		if meta.ChannelType == common.ChannelTypeAzure { | ||||||
|  | 			req.Header.Set("api-key", apiKey) | ||||||
|  | 		} else { | ||||||
|  | 			req.Header.Set("Authorization", c.Request.Header.Get("Authorization")) | ||||||
|  | 			if meta.ChannelType == common.ChannelTypeOpenRouter { | ||||||
|  | 				req.Header.Set("HTTP-Referer", "https://github.com/songquanpeng/one-api") | ||||||
|  | 				req.Header.Set("X-Title", "One API") | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeClaude: | ||||||
|  | 		req.Header.Set("x-api-key", apiKey) | ||||||
|  | 		anthropicVersion := c.Request.Header.Get("anthropic-version") | ||||||
|  | 		if anthropicVersion == "" { | ||||||
|  | 			anthropicVersion = "2023-06-01" | ||||||
|  | 		} | ||||||
|  | 		req.Header.Set("anthropic-version", anthropicVersion) | ||||||
|  | 	case constant.APITypeZhipu: | ||||||
|  | 		token := zhipu.GetToken(apiKey) | ||||||
|  | 		req.Header.Set("Authorization", token) | ||||||
|  | 	case constant.APITypeAli: | ||||||
|  | 		req.Header.Set("Authorization", "Bearer "+apiKey) | ||||||
|  | 		if isStream { | ||||||
|  | 			req.Header.Set("X-DashScope-SSE", "enable") | ||||||
|  | 		} | ||||||
|  | 		if c.GetString("plugin") != "" { | ||||||
|  | 			req.Header.Set("X-DashScope-Plugin", c.GetString("plugin")) | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeTencent: | ||||||
|  | 		req.Header.Set("Authorization", apiKey) | ||||||
|  | 	case constant.APITypePaLM: | ||||||
|  | 		req.Header.Set("x-goog-api-key", apiKey) | ||||||
|  | 	case constant.APITypeGemini: | ||||||
|  | 		req.Header.Set("x-goog-api-key", apiKey) | ||||||
|  | 	default: | ||||||
|  | 		req.Header.Set("Authorization", "Bearer "+apiKey) | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func DoResponse(c *gin.Context, textRequest *openai.GeneralOpenAIRequest, resp *http.Response, relayMode int, apiType int, isStream bool, promptTokens int) (usage *openai.Usage, err *openai.ErrorWithStatusCode) { | ||||||
|  | 	var responseText string | ||||||
|  | 	switch apiType { | ||||||
|  | 	case constant.APITypeOpenAI: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, responseText = openai.StreamHandler(c, resp, relayMode) | ||||||
|  | 		} else { | ||||||
|  | 			err, usage = openai.Handler(c, resp, promptTokens, textRequest.Model) | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeClaude: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, responseText = anthropic.StreamHandler(c, resp) | ||||||
|  | 		} else { | ||||||
|  | 			err, usage = anthropic.Handler(c, resp, promptTokens, textRequest.Model) | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeBaidu: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, usage = baidu.StreamHandler(c, resp) | ||||||
|  | 		} else { | ||||||
|  | 			switch relayMode { | ||||||
|  | 			case constant.RelayModeEmbeddings: | ||||||
|  | 				err, usage = baidu.EmbeddingHandler(c, resp) | ||||||
|  | 			default: | ||||||
|  | 				err, usage = baidu.Handler(c, resp) | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypePaLM: | ||||||
|  | 		if isStream { // PaLM2 API does not support stream | ||||||
|  | 			err, responseText = google.PaLMStreamHandler(c, resp) | ||||||
|  | 		} else { | ||||||
|  | 			err, usage = google.PaLMHandler(c, resp, promptTokens, textRequest.Model) | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeGemini: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, responseText = google.StreamHandler(c, resp) | ||||||
|  | 		} else { | ||||||
|  | 			err, usage = google.GeminiHandler(c, resp, promptTokens, textRequest.Model) | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeZhipu: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, usage = zhipu.StreamHandler(c, resp) | ||||||
|  | 		} else { | ||||||
|  | 			err, usage = zhipu.Handler(c, resp) | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeAli: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, usage = ali.StreamHandler(c, resp) | ||||||
|  | 		} else { | ||||||
|  | 			switch relayMode { | ||||||
|  | 			case constant.RelayModeEmbeddings: | ||||||
|  | 				err, usage = ali.EmbeddingHandler(c, resp) | ||||||
|  | 			default: | ||||||
|  | 				err, usage = ali.Handler(c, resp) | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeXunfei: | ||||||
|  | 		auth := c.Request.Header.Get("Authorization") | ||||||
|  | 		auth = strings.TrimPrefix(auth, "Bearer ") | ||||||
|  | 		splits := strings.Split(auth, "|") | ||||||
|  | 		if len(splits) != 3 { | ||||||
|  | 			return nil, openai.ErrorWrapper(errors.New("invalid auth"), "invalid_auth", http.StatusBadRequest) | ||||||
|  | 		} | ||||||
|  | 		if isStream { | ||||||
|  | 			err, usage = xunfei.StreamHandler(c, *textRequest, splits[0], splits[1], splits[2]) | ||||||
|  | 		} else { | ||||||
|  | 			err, usage = xunfei.Handler(c, *textRequest, splits[0], splits[1], splits[2]) | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeAIProxyLibrary: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, usage = aiproxy.StreamHandler(c, resp) | ||||||
|  | 		} else { | ||||||
|  | 			err, usage = aiproxy.Handler(c, resp) | ||||||
|  | 		} | ||||||
|  | 	case constant.APITypeTencent: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, responseText = tencent.StreamHandler(c, resp) | ||||||
|  | 		} else { | ||||||
|  | 			err, usage = tencent.Handler(c, resp) | ||||||
|  | 		} | ||||||
|  | 	default: | ||||||
|  | 		return nil, openai.ErrorWrapper(errors.New("unknown api type"), "unknown_api_type", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	if err != nil { | ||||||
|  | 		return nil, err | ||||||
|  | 	} | ||||||
|  | 	if usage == nil && responseText != "" { | ||||||
|  | 		usage = &openai.Usage{} | ||||||
|  | 		usage.PromptTokens = promptTokens | ||||||
|  | 		usage.CompletionTokens = openai.CountTokenText(responseText, textRequest.Model) | ||||||
|  | 		usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens | ||||||
|  | 	} | ||||||
|  | 	return usage, nil | ||||||
|  | } | ||||||
							
								
								
									
										19
									
								
								relay/util/billing.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										19
									
								
								relay/util/billing.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,19 @@ | |||||||
|  | package util | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"context" | ||||||
|  | 	"one-api/common/logger" | ||||||
|  | 	"one-api/model" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func ReturnPreConsumedQuota(ctx context.Context, preConsumedQuota int, tokenId int) { | ||||||
|  | 	if preConsumedQuota != 0 { | ||||||
|  | 		go func(ctx context.Context) { | ||||||
|  | 			// return pre-consumed quota | ||||||
|  | 			err := model.PostConsumeTokenQuota(tokenId, -preConsumedQuota) | ||||||
|  | 			if err != nil { | ||||||
|  | 				logger.Error(ctx, "error return pre-consumed quota: "+err.Error()) | ||||||
|  | 			} | ||||||
|  | 		}(ctx) | ||||||
|  | 	} | ||||||
|  | } | ||||||
| @@ -7,6 +7,8 @@ import ( | |||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"one-api/model" | 	"one-api/model" | ||||||
| 	"one-api/relay/channel/openai" | 	"one-api/relay/channel/openai" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| @@ -16,7 +18,7 @@ import ( | |||||||
| ) | ) | ||||||
|  |  | ||||||
| func ShouldDisableChannel(err *openai.Error, statusCode int) bool { | func ShouldDisableChannel(err *openai.Error, statusCode int) bool { | ||||||
| 	if !common.AutomaticDisableChannelEnabled { | 	if !config.AutomaticDisableChannelEnabled { | ||||||
| 		return false | 		return false | ||||||
| 	} | 	} | ||||||
| 	if err == nil { | 	if err == nil { | ||||||
| @@ -32,7 +34,7 @@ func ShouldDisableChannel(err *openai.Error, statusCode int) bool { | |||||||
| } | } | ||||||
|  |  | ||||||
| func ShouldEnableChannel(err error, openAIErr *openai.Error) bool { | func ShouldEnableChannel(err error, openAIErr *openai.Error) bool { | ||||||
| 	if !common.AutomaticEnableChannelEnabled { | 	if !config.AutomaticEnableChannelEnabled { | ||||||
| 		return false | 		return false | ||||||
| 	} | 	} | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| @@ -138,11 +140,11 @@ func PostConsumeQuota(ctx context.Context, tokenId int, quotaDelta int, totalQuo | |||||||
| 	// quotaDelta is remaining quota to be consumed | 	// quotaDelta is remaining quota to be consumed | ||||||
| 	err := model.PostConsumeTokenQuota(tokenId, quotaDelta) | 	err := model.PostConsumeTokenQuota(tokenId, quotaDelta) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("error consuming token remain quota: " + err.Error()) | 		logger.SysError("error consuming token remain quota: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	err = model.CacheUpdateUserQuota(userId) | 	err = model.CacheUpdateUserQuota(userId) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		common.SysError("error update user quota cache: " + err.Error()) | 		logger.SysError("error update user quota cache: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	// totalQuota is total quota consumed | 	// totalQuota is total quota consumed | ||||||
| 	if totalQuota != 0 { | 	if totalQuota != 0 { | ||||||
| @@ -152,11 +154,11 @@ func PostConsumeQuota(ctx context.Context, tokenId int, quotaDelta int, totalQuo | |||||||
| 		model.UpdateChannelUsedQuota(channelId, totalQuota) | 		model.UpdateChannelUsedQuota(channelId, totalQuota) | ||||||
| 	} | 	} | ||||||
| 	if totalQuota <= 0 { | 	if totalQuota <= 0 { | ||||||
| 		common.LogError(ctx, fmt.Sprintf("totalQuota consumed is %d, something is wrong", totalQuota)) | 		logger.Error(ctx, fmt.Sprintf("totalQuota consumed is %d, something is wrong", totalQuota)) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetAPIVersion(c *gin.Context) string { | func GetAzureAPIVersion(c *gin.Context) string { | ||||||
| 	query := c.Request.URL.Query() | 	query := c.Request.URL.Query() | ||||||
| 	apiVersion := query.Get("api-version") | 	apiVersion := query.Get("api-version") | ||||||
| 	if apiVersion == "" { | 	if apiVersion == "" { | ||||||
|   | |||||||
| @@ -2,7 +2,7 @@ package util | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -10,11 +10,11 @@ var HTTPClient *http.Client | |||||||
| var ImpatientHTTPClient *http.Client | var ImpatientHTTPClient *http.Client | ||||||
|  |  | ||||||
| func init() { | func init() { | ||||||
| 	if common.RelayTimeout == 0 { | 	if config.RelayTimeout == 0 { | ||||||
| 		HTTPClient = &http.Client{} | 		HTTPClient = &http.Client{} | ||||||
| 	} else { | 	} else { | ||||||
| 		HTTPClient = &http.Client{ | 		HTTPClient = &http.Client{ | ||||||
| 			Timeout: time.Duration(common.RelayTimeout) * time.Second, | 			Timeout: time.Duration(config.RelayTimeout) * time.Second, | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
|   | |||||||
							
								
								
									
										12
									
								
								relay/util/model_mapping.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										12
									
								
								relay/util/model_mapping.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,12 @@ | |||||||
|  | package util | ||||||
|  |  | ||||||
|  | func GetMappedModelName(modelName string, mapping map[string]string) (string, bool) { | ||||||
|  | 	if mapping == nil { | ||||||
|  | 		return modelName, false | ||||||
|  | 	} | ||||||
|  | 	mappedModelName := mapping[modelName] | ||||||
|  | 	if mappedModelName != "" { | ||||||
|  | 		return mappedModelName, true | ||||||
|  | 	} | ||||||
|  | 	return modelName, false | ||||||
|  | } | ||||||
							
								
								
									
										44
									
								
								relay/util/relay_meta.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										44
									
								
								relay/util/relay_meta.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,44 @@ | |||||||
|  | package util | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"strings" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type RelayMeta struct { | ||||||
|  | 	ChannelType  int | ||||||
|  | 	ChannelId    int | ||||||
|  | 	TokenId      int | ||||||
|  | 	TokenName    string | ||||||
|  | 	UserId       int | ||||||
|  | 	Group        string | ||||||
|  | 	ModelMapping map[string]string | ||||||
|  | 	BaseURL      string | ||||||
|  | 	APIVersion   string | ||||||
|  | 	APIKey       string | ||||||
|  | 	Config       map[string]string | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func GetRelayMeta(c *gin.Context) *RelayMeta { | ||||||
|  | 	meta := RelayMeta{ | ||||||
|  | 		ChannelType:  c.GetInt("channel"), | ||||||
|  | 		ChannelId:    c.GetInt("channel_id"), | ||||||
|  | 		TokenId:      c.GetInt("token_id"), | ||||||
|  | 		TokenName:    c.GetString("token_name"), | ||||||
|  | 		UserId:       c.GetInt("id"), | ||||||
|  | 		Group:        c.GetString("group"), | ||||||
|  | 		ModelMapping: c.GetStringMapString("model_mapping"), | ||||||
|  | 		BaseURL:      c.GetString("base_url"), | ||||||
|  | 		APIVersion:   c.GetString("api_version"), | ||||||
|  | 		APIKey:       strings.TrimPrefix(c.Request.Header.Get("Authorization"), "Bearer "), | ||||||
|  | 		Config:       nil, | ||||||
|  | 	} | ||||||
|  | 	if meta.ChannelType == common.ChannelTypeAzure { | ||||||
|  | 		meta.APIVersion = GetAzureAPIVersion(c) | ||||||
|  | 	} | ||||||
|  | 	if meta.BaseURL == "" { | ||||||
|  | 		meta.BaseURL = common.ChannelBaseURLs[meta.ChannelType] | ||||||
|  | 	} | ||||||
|  | 	return &meta | ||||||
|  | } | ||||||
							
								
								
									
										37
									
								
								relay/util/validation.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										37
									
								
								relay/util/validation.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,37 @@ | |||||||
|  | package util | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"errors" | ||||||
|  | 	"math" | ||||||
|  | 	"one-api/relay/channel/openai" | ||||||
|  | 	"one-api/relay/constant" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func ValidateTextRequest(textRequest *openai.GeneralOpenAIRequest, relayMode int) error { | ||||||
|  | 	if textRequest.MaxTokens < 0 || textRequest.MaxTokens > math.MaxInt32/2 { | ||||||
|  | 		return errors.New("max_tokens is invalid") | ||||||
|  | 	} | ||||||
|  | 	if textRequest.Model == "" { | ||||||
|  | 		return errors.New("model is required") | ||||||
|  | 	} | ||||||
|  | 	switch relayMode { | ||||||
|  | 	case constant.RelayModeCompletions: | ||||||
|  | 		if textRequest.Prompt == "" { | ||||||
|  | 			return errors.New("field prompt is required") | ||||||
|  | 		} | ||||||
|  | 	case constant.RelayModeChatCompletions: | ||||||
|  | 		if textRequest.Messages == nil || len(textRequest.Messages) == 0 { | ||||||
|  | 			return errors.New("field messages is required") | ||||||
|  | 		} | ||||||
|  | 	case constant.RelayModeEmbeddings: | ||||||
|  | 	case constant.RelayModeModerations: | ||||||
|  | 		if textRequest.Input == "" { | ||||||
|  | 			return errors.New("field input is required") | ||||||
|  | 		} | ||||||
|  | 	case constant.RelayModeEdits: | ||||||
|  | 		if textRequest.Instruction == "" { | ||||||
|  | 			return errors.New("field instruction is required") | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
| @@ -5,7 +5,8 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common/config" | ||||||
|  | 	"one-api/common/logger" | ||||||
| 	"os" | 	"os" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
| @@ -15,9 +16,9 @@ func SetRouter(router *gin.Engine, buildFS embed.FS) { | |||||||
| 	SetDashboardRouter(router) | 	SetDashboardRouter(router) | ||||||
| 	SetRelayRouter(router) | 	SetRelayRouter(router) | ||||||
| 	frontendBaseUrl := os.Getenv("FRONTEND_BASE_URL") | 	frontendBaseUrl := os.Getenv("FRONTEND_BASE_URL") | ||||||
| 	if common.IsMasterNode && frontendBaseUrl != "" { | 	if config.IsMasterNode && frontendBaseUrl != "" { | ||||||
| 		frontendBaseUrl = "" | 		frontendBaseUrl = "" | ||||||
| 		common.SysLog("FRONTEND_BASE_URL is ignored on master node") | 		logger.SysLog("FRONTEND_BASE_URL is ignored on master node") | ||||||
| 	} | 	} | ||||||
| 	if frontendBaseUrl == "" { | 	if frontendBaseUrl == "" { | ||||||
| 		SetWebRouter(router, buildFS) | 		SetWebRouter(router, buildFS) | ||||||
|   | |||||||
| @@ -8,17 +8,18 @@ import ( | |||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"one-api/common" | 	"one-api/common" | ||||||
|  | 	"one-api/common/config" | ||||||
| 	"one-api/controller" | 	"one-api/controller" | ||||||
| 	"one-api/middleware" | 	"one-api/middleware" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func SetWebRouter(router *gin.Engine, buildFS embed.FS) { | func SetWebRouter(router *gin.Engine, buildFS embed.FS) { | ||||||
| 	indexPageData, _ := buildFS.ReadFile(fmt.Sprintf("web/build/%s/index.html", common.Theme)) | 	indexPageData, _ := buildFS.ReadFile(fmt.Sprintf("web/build/%s/index.html", config.Theme)) | ||||||
| 	router.Use(gzip.Gzip(gzip.DefaultCompression)) | 	router.Use(gzip.Gzip(gzip.DefaultCompression)) | ||||||
| 	router.Use(middleware.GlobalWebRateLimit()) | 	router.Use(middleware.GlobalWebRateLimit()) | ||||||
| 	router.Use(middleware.Cache()) | 	router.Use(middleware.Cache()) | ||||||
| 	router.Use(static.Serve("/", common.EmbedFolder(buildFS, fmt.Sprintf("web/build/%s", common.Theme)))) | 	router.Use(static.Serve("/", common.EmbedFolder(buildFS, fmt.Sprintf("web/build/%s", config.Theme)))) | ||||||
| 	router.NoRoute(func(c *gin.Context) { | 	router.NoRoute(func(c *gin.Context) { | ||||||
| 		if strings.HasPrefix(c.Request.RequestURI, "/v1") || strings.HasPrefix(c.Request.RequestURI, "/api") { | 		if strings.HasPrefix(c.Request.RequestURI, "/v1") || strings.HasPrefix(c.Request.RequestURI, "/api") { | ||||||
| 			controller.RelayNotFound(c) | 			controller.RelayNotFound(c) | ||||||
|   | |||||||
		Reference in New Issue
	
	Block a user