From 5760570aec9ab4c23aa176a347f6d0f8a779a251 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=A2=AB=E9=81=97=E5=BF=98=E7=9A=84=E8=AE=B0=E5=BF=86?= <32097237+YIWANG-sketch@users.noreply.github.com> Date: Tue, 26 Nov 2024 14:05:16 +0800 Subject: [PATCH] add image upload --- index.mjs | 102 ++++++++++++++++++++++++++++++++++++++++++++++++------ 1 file changed, 92 insertions(+), 10 deletions(-) diff --git a/index.mjs b/index.mjs index d5b7545..16db5d4 100644 --- a/index.mjs +++ b/index.mjs @@ -6,6 +6,10 @@ import ngrok from 'ngrok'; import {v4 as uuidv4} from "uuid"; import './proxyAgent.mjs'; import SessionManager from './sessionManager.mjs'; +import { storeImage } from './imageStorage.mjs'; +import fetch from 'node-fetch'; +import path from 'path'; +import { Buffer } from 'buffer'; const app = express(); const port = process.env.PORT || 8080; @@ -109,7 +113,7 @@ app.post("/v1/chat/completions", OpenAIApiKeyAuth, (req, res) => { let jsonBody = JSON.parse(req.rawBody); // 规范化消息 - jsonBody.messages = openaiNormalizeMessages(jsonBody.messages); + jsonBody.messages = await openaiNormalizeMessages(jsonBody.messages); console.log("message length:" + jsonBody.messages.length); @@ -290,7 +294,7 @@ app.post("/v1/chat/completions", OpenAIApiKeyAuth, (req, res) => { }); // Helper function: Normalize messages -function openaiNormalizeMessages(messages) { +async function openaiNormalizeMessages(messages) { let normalizedMessages = []; let currentSystemMessage = ""; @@ -303,20 +307,90 @@ function openaiNormalizeMessages(messages) { } } else { if (currentSystemMessage) { - normalizedMessages.push({role: 'system', content: currentSystemMessage}); + normalizedMessages.push({ role: 'system', content: currentSystemMessage }); currentSystemMessage = ""; } - normalizedMessages.push(message); + + // 检查消息内容 + if (Array.isArray(message.content)) { + const textContent = message.content + .filter(item => item.type === 'text') + .map(item => item.text) + .join('\n'); + + // 处理图片内容,存储图片 + for (const item of message.content) { + if (item.type === 'image_url' && item.image_url?.url) { + // 获取媒体类型 + const mediaType = await getMediaTypeFromUrl(item.image_url.url); + // 获取图片 base64 + const base64Data = await fetchImageAsBase64(item.image_url.url); + if (base64Data) { + const { imageId } = storeImage(base64Data, mediaType); + console.log(`Image stored with ID: ${imageId}, Media Type: ${mediaType}`); + } else { + console.warn('Failed to store image due to missing data.'); + } + } + } + + normalizedMessages.push({ role: message.role, content: textContent }); + } else if (typeof message.content === 'string') { + normalizedMessages.push(message); + } else { + console.warn('未知的消息内容格式:', message.content); + normalizedMessages.push(message); + } } } if (currentSystemMessage) { - normalizedMessages.push({role: 'system', content: currentSystemMessage}); + normalizedMessages.push({ role: 'system', content: currentSystemMessage }); } return normalizedMessages; } +// 图片 URL 获取媒体类型 +async function getMediaTypeFromUrl(url) { + try { + const response = await fetch(url, { method: 'HEAD' }); + const contentType = response.headers.get('content-type'); + return contentType || guessMediaTypeFromUrl(url); + } catch (error) { + console.warn('无法获取媒体类型,尝试根据 URL 推断', error); + return guessMediaTypeFromUrl(url); + } +} + +function guessMediaTypeFromUrl(url) { + const ext = path.extname(url).toLowerCase(); + switch (ext) { + case '.jpg': + case '.jpeg': + return 'image/jpeg'; + case '.png': + return 'image/png'; + case '.gif': + return 'image/gif'; + default: + return 'application/octet-stream'; + } +} + +// 图片 URL 获取 base64 +async function fetchImageAsBase64(url) { + try { + const response = await fetch(url); + const arrayBuffer = await response.arrayBuffer(); + const buffer = Buffer.from(arrayBuffer); + return buffer.toString('base64'); + } catch (error) { + console.error('Failed to fetch image data:', error); + return null; + } +} + // handle anthropic format model request app.post("/v1/messages", AnthropicApiKeyAuth, (req, res) => { @@ -466,16 +540,24 @@ function anthropicNormalizeMessages(messages) { if (typeof message.content === 'string') { return message; } else if (Array.isArray(message.content)) { - // 新版格式,提取文本内容 + // 提取文本内容 const textContent = message.content .filter(item => item.type === 'text') .map(item => item.text) .join('\n'); - return {...message, content: textContent}; + + // 处理图片内容,存储图片 + message.content.forEach(item => { + if (item.type === 'image' && item.source?.type === 'base64') { + const { imageId, mediaType } = storeImage(item.source.data, item.source.media_type); + console.log(`Image stored with ID: ${imageId}, Media Type: ${mediaType}`); + } + }); + + return { ...message, content: textContent }; } else { - // 未知格式,返回原始消息 - console.warn('未知的消息格式:', message); - return message; + console.warn('Unknown message format:', message); + return message; // 未知格式,返回原始消息 } }); }