resolve-model.ts 2.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778
  1. /**
  2. * Shared model resolution utilities for API routes.
  3. *
  4. * Extracts the repeated parseModelString → resolveApiKey → resolveBaseUrl →
  5. * resolveProxy → getModel boilerplate into a single call.
  6. */
  7. import type { NextRequest } from 'next/server';
  8. import { getModel, parseModelString, type ModelWithInfo } from '@/lib/ai/providers';
  9. import { resolveApiKey, resolveBaseUrl, resolveProxy } from '@/lib/server/provider-config';
  10. import { validateUrlForSSRF } from '@/lib/server/ssrf-guard';
  11. export interface ResolvedModel extends ModelWithInfo {
  12. /** Original model string (e.g. "openai/gpt-4o-mini") */
  13. modelString: string;
  14. /** Resolved provider ID (e.g. "openai", "ollama") */
  15. providerId: string;
  16. /** Effective API key after server-side fallback resolution */
  17. apiKey: string;
  18. }
  19. /**
  20. * Resolve a language model from explicit parameters.
  21. *
  22. * Use this when model config comes from the request body.
  23. */
  24. export async function resolveModel(params: {
  25. modelString?: string;
  26. apiKey?: string;
  27. baseUrl?: string;
  28. providerType?: string;
  29. }): Promise<ResolvedModel> {
  30. const modelString = params.modelString || process.env.DEFAULT_MODEL || 'gpt-4o-mini';
  31. const { providerId, modelId } = parseModelString(modelString);
  32. // SSRF validation applies only to client-supplied base URLs.
  33. // Server-configured URLs (e.g. OLLAMA_BASE_URL from env/YAML) flow through
  34. // resolveBaseUrl() and bypass this check — they're trusted by the operator.
  35. const clientBaseUrl = params.baseUrl || undefined;
  36. if (clientBaseUrl && process.env.NODE_ENV === 'production') {
  37. const ssrfError = await validateUrlForSSRF(clientBaseUrl);
  38. if (ssrfError) {
  39. throw new Error(ssrfError);
  40. }
  41. }
  42. const apiKey = clientBaseUrl
  43. ? params.apiKey || ''
  44. : resolveApiKey(providerId, params.apiKey || '');
  45. const baseUrl = clientBaseUrl ? clientBaseUrl : resolveBaseUrl(providerId, params.baseUrl);
  46. const proxy = resolveProxy(providerId);
  47. const { model, modelInfo } = getModel({
  48. providerId,
  49. modelId,
  50. apiKey,
  51. baseUrl,
  52. proxy,
  53. providerType: params.providerType as 'openai' | 'anthropic' | 'google' | undefined,
  54. });
  55. return { model, modelInfo, modelString, providerId, apiKey };
  56. }
  57. /**
  58. * Resolve a language model from standard request headers.
  59. *
  60. * Reads: x-model, x-api-key, x-base-url, x-provider-type
  61. * Note: requiresApiKey is derived server-side from the provider registry,
  62. * never from client headers, to prevent auth bypass.
  63. */
  64. export async function resolveModelFromHeaders(req: NextRequest): Promise<ResolvedModel> {
  65. return resolveModel({
  66. modelString: req.headers.get('x-model') || undefined,
  67. apiKey: req.headers.get('x-api-key') || undefined,
  68. baseUrl: req.headers.get('x-base-url') || undefined,
  69. providerType: req.headers.get('x-provider-type') || undefined,
  70. });
  71. }