model.ts 4.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156
  1. import { z } from "zod"
  2. import { eq, and } from "drizzle-orm"
  3. import { Database } from "./drizzle"
  4. import { ModelTable } from "./schema/model.sql"
  5. import { Identifier } from "./identifier"
  6. import { fn } from "./util/fn"
  7. import { Actor } from "./actor"
  8. import { Resource } from "@opencode-ai/console-resource"
  9. export namespace ZenData {
  10. const FormatSchema = z.enum(["anthropic", "google", "openai", "oa-compat"])
  11. export type Format = z.infer<typeof FormatSchema>
  12. const ModelCostSchema = z.object({
  13. input: z.number(),
  14. output: z.number(),
  15. cacheRead: z.number().optional(),
  16. cacheWrite5m: z.number().optional(),
  17. cacheWrite1h: z.number().optional(),
  18. })
  19. const ModelSchema = z.object({
  20. name: z.string(),
  21. cost: ModelCostSchema,
  22. cost200K: ModelCostSchema.optional(),
  23. allowAnonymous: z.boolean().optional(),
  24. byokProvider: z.enum(["openai", "anthropic", "google"]).optional(),
  25. stickyProvider: z.enum(["strict", "prefer"]).optional(),
  26. trialProvider: z.string().optional(),
  27. fallbackProvider: z.string().optional(),
  28. rateLimit: z.number().optional(),
  29. providers: z.array(
  30. z.object({
  31. id: z.string(),
  32. model: z.string(),
  33. weight: z.number().optional(),
  34. disabled: z.boolean().optional(),
  35. storeModel: z.string().optional(),
  36. payloadModifier: z.record(z.string(), z.any()).optional(),
  37. }),
  38. ),
  39. })
  40. const ProviderSchema = z.object({
  41. api: z.string(),
  42. apiKey: z.string(),
  43. format: FormatSchema.optional(),
  44. headerMappings: z.record(z.string(), z.string()).optional(),
  45. payloadModifier: z.record(z.string(), z.any()).optional(),
  46. payloadMappings: z.record(z.string(), z.string()).optional(),
  47. })
  48. const ModelsSchema = z.object({
  49. models: z.record(z.string(), z.union([ModelSchema, z.array(ModelSchema.extend({ formatFilter: FormatSchema }))])),
  50. liteModels: z.record(z.string(), ModelSchema),
  51. providers: z.record(z.string(), ProviderSchema),
  52. })
  53. export const validate = fn(ModelsSchema, (input) => {
  54. return input
  55. })
  56. export const list = fn(z.enum(["lite", "full"]), (modelList) => {
  57. const json = JSON.parse(
  58. Resource.ZEN_MODELS1.value +
  59. Resource.ZEN_MODELS2.value +
  60. Resource.ZEN_MODELS3.value +
  61. Resource.ZEN_MODELS4.value +
  62. Resource.ZEN_MODELS5.value +
  63. Resource.ZEN_MODELS6.value +
  64. Resource.ZEN_MODELS7.value +
  65. Resource.ZEN_MODELS8.value +
  66. Resource.ZEN_MODELS9.value +
  67. Resource.ZEN_MODELS10.value +
  68. Resource.ZEN_MODELS11.value +
  69. Resource.ZEN_MODELS12.value +
  70. Resource.ZEN_MODELS13.value +
  71. Resource.ZEN_MODELS14.value +
  72. Resource.ZEN_MODELS15.value +
  73. Resource.ZEN_MODELS16.value +
  74. Resource.ZEN_MODELS17.value +
  75. Resource.ZEN_MODELS18.value +
  76. Resource.ZEN_MODELS19.value +
  77. Resource.ZEN_MODELS20.value +
  78. Resource.ZEN_MODELS21.value +
  79. Resource.ZEN_MODELS22.value +
  80. Resource.ZEN_MODELS23.value +
  81. Resource.ZEN_MODELS24.value +
  82. Resource.ZEN_MODELS25.value +
  83. Resource.ZEN_MODELS26.value +
  84. Resource.ZEN_MODELS27.value +
  85. Resource.ZEN_MODELS28.value +
  86. Resource.ZEN_MODELS29.value +
  87. Resource.ZEN_MODELS30.value,
  88. )
  89. const { models, liteModels, providers } = ModelsSchema.parse(json)
  90. return {
  91. models: modelList === "lite" ? liteModels : models,
  92. providers,
  93. }
  94. })
  95. }
  96. export namespace Model {
  97. export const enable = fn(z.object({ model: z.string() }), ({ model }) => {
  98. Actor.assertAdmin()
  99. return Database.use((db) =>
  100. db.delete(ModelTable).where(and(eq(ModelTable.workspaceID, Actor.workspace()), eq(ModelTable.model, model))),
  101. )
  102. })
  103. export const disable = fn(z.object({ model: z.string() }), ({ model }) => {
  104. Actor.assertAdmin()
  105. return Database.use((db) =>
  106. db
  107. .insert(ModelTable)
  108. .values({
  109. id: Identifier.create("model"),
  110. workspaceID: Actor.workspace(),
  111. model: model,
  112. })
  113. .onDuplicateKeyUpdate({
  114. set: {
  115. timeDeleted: null,
  116. },
  117. }),
  118. )
  119. })
  120. export const listDisabled = fn(z.void(), () => {
  121. return Database.use((db) =>
  122. db
  123. .select({ model: ModelTable.model })
  124. .from(ModelTable)
  125. .where(eq(ModelTable.workspaceID, Actor.workspace()))
  126. .then((rows) => rows.map((row) => row.model)),
  127. )
  128. })
  129. export const isDisabled = fn(
  130. z.object({
  131. model: z.string(),
  132. }),
  133. ({ model }) => {
  134. return Database.use(async (db) => {
  135. const result = await db
  136. .select()
  137. .from(ModelTable)
  138. .where(and(eq(ModelTable.workspaceID, Actor.workspace()), eq(ModelTable.model, model)))
  139. .limit(1)
  140. return result.length > 0
  141. })
  142. },
  143. )
  144. }