models.go 5.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228
  1. package models
  2. import (
  3. "context"
  4. "errors"
  5. "log"
  6. "github.com/cloudwego/eino-ext/components/model/claude"
  7. "github.com/cloudwego/eino-ext/components/model/openai"
  8. "github.com/cloudwego/eino/components/model"
  9. "github.com/spf13/viper"
  10. )
  11. type (
  12. ModelID string
  13. ModelProvider string
  14. )
  15. type Model struct {
  16. ID ModelID `json:"id"`
  17. Name string `json:"name"`
  18. Provider ModelProvider `json:"provider"`
  19. APIModel string `json:"api_model"`
  20. CostPer1MIn float64 `json:"cost_per_1m_in"`
  21. CostPer1MOut float64 `json:"cost_per_1m_out"`
  22. }
  23. const (
  24. DefaultBigModel = GPT4oMini
  25. DefaultLittleModel = GPT4oMini
  26. )
  27. // Model IDs
  28. const (
  29. // OpenAI
  30. GPT4o ModelID = "gpt-4o"
  31. GPT4oMini ModelID = "gpt-4o-mini"
  32. GPT45 ModelID = "gpt-4.5"
  33. O1 ModelID = "o1"
  34. O1Mini ModelID = "o1-mini"
  35. // Anthropic
  36. Claude35Sonnet ModelID = "claude-3.5-sonnet"
  37. Claude3Haiku ModelID = "claude-3-haiku"
  38. Claude37Sonnet ModelID = "claude-3.7-sonnet"
  39. // Google
  40. Gemini20Pro ModelID = "gemini-2.0-pro"
  41. Gemini15Flash ModelID = "gemini-1.5-flash"
  42. Gemini20Flash ModelID = "gemini-2.0-flash"
  43. // xAI
  44. Grok3 ModelID = "grok-3"
  45. Grok2Mini ModelID = "grok-2-mini"
  46. // DeepSeek
  47. DeepSeekR1 ModelID = "deepseek-r1"
  48. DeepSeekCoder ModelID = "deepseek-coder"
  49. // Meta
  50. Llama3 ModelID = "llama-3"
  51. Llama270B ModelID = "llama-2-70b"
  52. // GROQ
  53. GroqLlama3SpecDec ModelID = "groq-llama-3-spec-dec"
  54. GroqQwen32BCoder ModelID = "qwen-2.5-coder-32b"
  55. )
  56. const (
  57. ProviderOpenAI ModelProvider = "openai"
  58. ProviderAnthropic ModelProvider = "anthropic"
  59. ProviderGoogle ModelProvider = "google"
  60. ProviderXAI ModelProvider = "xai"
  61. ProviderDeepSeek ModelProvider = "deepseek"
  62. ProviderMeta ModelProvider = "meta"
  63. ProviderGroq ModelProvider = "groq"
  64. )
  65. var SupportedModels = map[ModelID]Model{
  66. // OpenAI
  67. GPT4o: {
  68. ID: GPT4o,
  69. Name: "GPT-4o",
  70. Provider: ProviderOpenAI,
  71. APIModel: "gpt-4o",
  72. },
  73. GPT4oMini: {
  74. ID: GPT4oMini,
  75. Name: "GPT-4o Mini",
  76. Provider: ProviderOpenAI,
  77. APIModel: "gpt-4o-mini",
  78. CostPer1MIn: 0.150,
  79. CostPer1MOut: 0.600,
  80. },
  81. GPT45: {
  82. ID: GPT45,
  83. Name: "GPT-4.5",
  84. Provider: ProviderOpenAI,
  85. APIModel: "gpt-4.5",
  86. },
  87. O1: {
  88. ID: O1,
  89. Name: "o1",
  90. Provider: ProviderOpenAI,
  91. APIModel: "o1",
  92. },
  93. O1Mini: {
  94. ID: O1Mini,
  95. Name: "o1 Mini",
  96. Provider: ProviderOpenAI,
  97. APIModel: "o1-mini",
  98. },
  99. // Anthropic
  100. Claude35Sonnet: {
  101. ID: Claude35Sonnet,
  102. Name: "Claude 3.5 Sonnet",
  103. Provider: ProviderAnthropic,
  104. APIModel: "claude-3.5-sonnet",
  105. },
  106. Claude3Haiku: {
  107. ID: Claude3Haiku,
  108. Name: "Claude 3 Haiku",
  109. Provider: ProviderAnthropic,
  110. APIModel: "claude-3-haiku",
  111. },
  112. Claude37Sonnet: {
  113. ID: Claude37Sonnet,
  114. Name: "Claude 3.7 Sonnet",
  115. Provider: ProviderAnthropic,
  116. APIModel: "claude-3-7-sonnet-20250219",
  117. },
  118. // Google
  119. Gemini20Pro: {
  120. ID: Gemini20Pro,
  121. Name: "Gemini 2.0 Pro",
  122. Provider: ProviderGoogle,
  123. APIModel: "gemini-2.0-pro",
  124. },
  125. Gemini15Flash: {
  126. ID: Gemini15Flash,
  127. Name: "Gemini 1.5 Flash",
  128. Provider: ProviderGoogle,
  129. APIModel: "gemini-1.5-flash",
  130. },
  131. Gemini20Flash: {
  132. ID: Gemini20Flash,
  133. Name: "Gemini 2.0 Flash",
  134. Provider: ProviderGoogle,
  135. APIModel: "gemini-2.0-flash",
  136. },
  137. // xAI
  138. Grok3: {
  139. ID: Grok3,
  140. Name: "Grok 3",
  141. Provider: ProviderXAI,
  142. APIModel: "grok-3",
  143. },
  144. Grok2Mini: {
  145. ID: Grok2Mini,
  146. Name: "Grok 2 Mini",
  147. Provider: ProviderXAI,
  148. APIModel: "grok-2-mini",
  149. },
  150. // DeepSeek
  151. DeepSeekR1: {
  152. ID: DeepSeekR1,
  153. Name: "DeepSeek R1",
  154. Provider: ProviderDeepSeek,
  155. APIModel: "deepseek-r1",
  156. },
  157. DeepSeekCoder: {
  158. ID: DeepSeekCoder,
  159. Name: "DeepSeek Coder",
  160. Provider: ProviderDeepSeek,
  161. APIModel: "deepseek-coder",
  162. },
  163. // Meta
  164. Llama3: {
  165. ID: Llama3,
  166. Name: "LLaMA 3",
  167. Provider: ProviderMeta,
  168. APIModel: "llama-3",
  169. },
  170. Llama270B: {
  171. ID: Llama270B,
  172. Name: "LLaMA 2 70B",
  173. Provider: ProviderMeta,
  174. APIModel: "llama-2-70b",
  175. },
  176. // GROQ
  177. GroqLlama3SpecDec: {
  178. ID: GroqLlama3SpecDec,
  179. Name: "GROQ LLaMA 3 SpecDec",
  180. Provider: ProviderGroq,
  181. APIModel: "llama-3.3-70b-specdec",
  182. },
  183. GroqQwen32BCoder: {
  184. ID: GroqQwen32BCoder,
  185. Name: "GROQ Qwen 2.5 Coder 32B",
  186. Provider: ProviderGroq,
  187. APIModel: "qwen-2.5-coder-32b",
  188. },
  189. }
  190. func GetModel(ctx context.Context, model ModelID) (model.ChatModel, error) {
  191. provider := SupportedModels[model].Provider
  192. log.Printf("Provider: %s", provider)
  193. maxTokens := viper.GetInt("providers.common.max_tokens")
  194. switch provider {
  195. case ProviderOpenAI:
  196. return openai.NewChatModel(ctx, &openai.ChatModelConfig{
  197. APIKey: viper.GetString("providers.openai.key"),
  198. Model: string(SupportedModels[model].APIModel),
  199. MaxTokens: &maxTokens,
  200. })
  201. case ProviderAnthropic:
  202. return claude.NewChatModel(ctx, &claude.Config{
  203. APIKey: viper.GetString("providers.anthropic.key"),
  204. Model: string(SupportedModels[model].APIModel),
  205. MaxTokens: maxTokens,
  206. })
  207. case ProviderGroq:
  208. return openai.NewChatModel(ctx, &openai.ChatModelConfig{
  209. BaseURL: "https://api.groq.com/openai/v1",
  210. APIKey: viper.GetString("providers.groq.key"),
  211. Model: string(SupportedModels[model].APIModel),
  212. MaxTokens: &maxTokens,
  213. })
  214. }
  215. return nil, errors.New("unsupported provider")
  216. }