mirror of
				https://github.com/songquanpeng/one-api.git
				synced 2025-10-31 22:03:41 +08:00 
			
		
		
		
	Compare commits
	
		
			3 Commits
		
	
	
		
			v0.6.2-alp
			...
			v0.5.11-de
		
	
	| Author | SHA1 | Date | |
|---|---|---|---|
|  | 227e11c5ac | ||
|  | b0bf224bb1 | ||
|  | 5342af9222 | 
							
								
								
									
										49
									
								
								.github/workflows/docker-image-amd64-en.yml
									
									
									
									
										vendored
									
									
								
							
							
						
						
									
										49
									
								
								.github/workflows/docker-image-amd64-en.yml
									
									
									
									
										vendored
									
									
								
							| @@ -1,49 +0,0 @@ | |||||||
| name: Publish Docker image (amd64, English) |  | ||||||
|  |  | ||||||
| on: |  | ||||||
|   push: |  | ||||||
|     tags: |  | ||||||
|       - '*' |  | ||||||
|   workflow_dispatch: |  | ||||||
|     inputs: |  | ||||||
|       name: |  | ||||||
|         description: 'reason' |  | ||||||
|         required: false |  | ||||||
| jobs: |  | ||||||
|   push_to_registries: |  | ||||||
|     name: Push Docker image to multiple registries |  | ||||||
|     runs-on: ubuntu-latest |  | ||||||
|     permissions: |  | ||||||
|       packages: write |  | ||||||
|       contents: read |  | ||||||
|     steps: |  | ||||||
|       - name: Check out the repo |  | ||||||
|         uses: actions/checkout@v3 |  | ||||||
|  |  | ||||||
|       - name: Save version info |  | ||||||
|         run: | |  | ||||||
|           git describe --tags > VERSION  |  | ||||||
|  |  | ||||||
|       - name: Translate |  | ||||||
|         run: | |  | ||||||
|           python ./i18n/translate.py --repository_path . --json_file_path ./i18n/en.json |  | ||||||
|       - name: Log in to Docker Hub |  | ||||||
|         uses: docker/login-action@v2 |  | ||||||
|         with: |  | ||||||
|           username: ${{ secrets.DOCKERHUB_USERNAME }} |  | ||||||
|           password: ${{ secrets.DOCKERHUB_TOKEN }} |  | ||||||
|  |  | ||||||
|       - name: Extract metadata (tags, labels) for Docker |  | ||||||
|         id: meta |  | ||||||
|         uses: docker/metadata-action@v4 |  | ||||||
|         with: |  | ||||||
|           images: | |  | ||||||
|             justsong/one-api-en |  | ||||||
|  |  | ||||||
|       - name: Build and push Docker images |  | ||||||
|         uses: docker/build-push-action@v3 |  | ||||||
|         with: |  | ||||||
|           context: . |  | ||||||
|           push: true |  | ||||||
|           tags: ${{ steps.meta.outputs.tags }} |  | ||||||
|           labels: ${{ steps.meta.outputs.labels }} |  | ||||||
							
								
								
									
										22
									
								
								.github/workflows/docker-image-amd64.yml
									
									
									
									
										vendored
									
									
								
							
							
						
						
									
										22
									
								
								.github/workflows/docker-image-amd64.yml
									
									
									
									
										vendored
									
									
								
							| @@ -1,34 +1,28 @@ | |||||||
| name: Publish Docker image (amd64) | name: Publish Docker image (amd64) | ||||||
|  |  | ||||||
| on: | on: | ||||||
|   push: |  | ||||||
|     tags: |  | ||||||
|       - '*' |  | ||||||
|   workflow_dispatch: |   workflow_dispatch: | ||||||
|     inputs: |     inputs: | ||||||
|       name: |       name: | ||||||
|         description: 'reason' |         description: 'reason' | ||||||
|         required: false |         required: false | ||||||
|  |  | ||||||
| jobs: | jobs: | ||||||
|   push_to_registries: |   build-and-push-image: | ||||||
|     name: Push Docker image to multiple registries |     name: Push Docker image to GitHub registry | ||||||
|     runs-on: ubuntu-latest |     runs-on: ubuntu-latest | ||||||
|     permissions: |     permissions: | ||||||
|       packages: write |       packages: write | ||||||
|       contents: read |       contents: read | ||||||
|  |  | ||||||
|     steps: |     steps: | ||||||
|  |  | ||||||
|       - name: Check out the repo |       - name: Check out the repo | ||||||
|         uses: actions/checkout@v3 |         uses: actions/checkout@v3 | ||||||
|  |  | ||||||
|       - name: Save version info |       - name: Save version info | ||||||
|         run: | |         run: | | ||||||
|           git describe --tags > VERSION  |           git describe --tags > VERSION | ||||||
|  |  | ||||||
|       - name: Log in to Docker Hub |  | ||||||
|         uses: docker/login-action@v2 |  | ||||||
|         with: |  | ||||||
|           username: ${{ secrets.DOCKERHUB_USERNAME }} |  | ||||||
|           password: ${{ secrets.DOCKERHUB_TOKEN }} |  | ||||||
|  |  | ||||||
|       - name: Log in to the Container registry |       - name: Log in to the Container registry | ||||||
|         uses: docker/login-action@v2 |         uses: docker/login-action@v2 | ||||||
| @@ -41,9 +35,7 @@ jobs: | |||||||
|         id: meta |         id: meta | ||||||
|         uses: docker/metadata-action@v4 |         uses: docker/metadata-action@v4 | ||||||
|         with: |         with: | ||||||
|           images: | |           images: ghcr.io/${{ github.repository }} | ||||||
|             justsong/one-api |  | ||||||
|             ghcr.io/${{ github.repository }} |  | ||||||
|  |  | ||||||
|       - name: Build and push Docker images |       - name: Build and push Docker images | ||||||
|         uses: docker/build-push-action@v3 |         uses: docker/build-push-action@v3 | ||||||
|   | |||||||
							
								
								
									
										62
									
								
								.github/workflows/docker-image-arm64.yml
									
									
									
									
										vendored
									
									
								
							
							
						
						
									
										62
									
								
								.github/workflows/docker-image-arm64.yml
									
									
									
									
										vendored
									
									
								
							| @@ -1,62 +0,0 @@ | |||||||
| name: Publish Docker image (arm64) |  | ||||||
|  |  | ||||||
| on: |  | ||||||
|   push: |  | ||||||
|     tags: |  | ||||||
|       - '*' |  | ||||||
|       - '!*-alpha*' |  | ||||||
|   workflow_dispatch: |  | ||||||
|     inputs: |  | ||||||
|       name: |  | ||||||
|         description: 'reason' |  | ||||||
|         required: false |  | ||||||
| jobs: |  | ||||||
|   push_to_registries: |  | ||||||
|     name: Push Docker image to multiple registries |  | ||||||
|     runs-on: ubuntu-latest |  | ||||||
|     permissions: |  | ||||||
|       packages: write |  | ||||||
|       contents: read |  | ||||||
|     steps: |  | ||||||
|       - name: Check out the repo |  | ||||||
|         uses: actions/checkout@v3 |  | ||||||
|  |  | ||||||
|       - name: Save version info |  | ||||||
|         run: | |  | ||||||
|           git describe --tags > VERSION  |  | ||||||
|  |  | ||||||
|       - name: Set up QEMU |  | ||||||
|         uses: docker/setup-qemu-action@v2 |  | ||||||
|  |  | ||||||
|       - name: Set up Docker Buildx |  | ||||||
|         uses: docker/setup-buildx-action@v2 |  | ||||||
|  |  | ||||||
|       - name: Log in to Docker Hub |  | ||||||
|         uses: docker/login-action@v2 |  | ||||||
|         with: |  | ||||||
|           username: ${{ secrets.DOCKERHUB_USERNAME }} |  | ||||||
|           password: ${{ secrets.DOCKERHUB_TOKEN }} |  | ||||||
|  |  | ||||||
|       - name: Log in to the Container registry |  | ||||||
|         uses: docker/login-action@v2 |  | ||||||
|         with: |  | ||||||
|           registry: ghcr.io |  | ||||||
|           username: ${{ github.actor }} |  | ||||||
|           password: ${{ secrets.GITHUB_TOKEN }} |  | ||||||
|  |  | ||||||
|       - name: Extract metadata (tags, labels) for Docker |  | ||||||
|         id: meta |  | ||||||
|         uses: docker/metadata-action@v4 |  | ||||||
|         with: |  | ||||||
|           images: | |  | ||||||
|             justsong/one-api |  | ||||||
|             ghcr.io/${{ github.repository }} |  | ||||||
|  |  | ||||||
|       - name: Build and push Docker images |  | ||||||
|         uses: docker/build-push-action@v3 |  | ||||||
|         with: |  | ||||||
|           context: . |  | ||||||
|           platforms: linux/amd64,linux/arm64 |  | ||||||
|           push: true |  | ||||||
|           tags: ${{ steps.meta.outputs.tags }} |  | ||||||
|           labels: ${{ steps.meta.outputs.labels }} |  | ||||||
							
								
								
									
										59
									
								
								.github/workflows/linux-release.yml
									
									
									
									
										vendored
									
									
								
							
							
						
						
									
										59
									
								
								.github/workflows/linux-release.yml
									
									
									
									
										vendored
									
									
								
							| @@ -1,59 +0,0 @@ | |||||||
| name: Linux Release |  | ||||||
| permissions: |  | ||||||
|   contents: write |  | ||||||
|  |  | ||||||
| on: |  | ||||||
|   push: |  | ||||||
|     tags: |  | ||||||
|       - '*' |  | ||||||
|       - '!*-alpha*' |  | ||||||
|   workflow_dispatch: |  | ||||||
|     inputs: |  | ||||||
|       name: |  | ||||||
|         description: 'reason' |  | ||||||
|         required: false |  | ||||||
| jobs: |  | ||||||
|   release: |  | ||||||
|     runs-on: ubuntu-latest |  | ||||||
|     steps: |  | ||||||
|       - name: Checkout |  | ||||||
|         uses: actions/checkout@v3 |  | ||||||
|         with: |  | ||||||
|           fetch-depth: 0 |  | ||||||
|       - uses: actions/setup-node@v3 |  | ||||||
|         with: |  | ||||||
|           node-version: 16 |  | ||||||
|       - name: Build Frontend |  | ||||||
|         env: |  | ||||||
|           CI: "" |  | ||||||
|         run: | |  | ||||||
|           cd web |  | ||||||
|           git describe --tags > VERSION |  | ||||||
|           REACT_APP_VERSION=$(git describe --tags) chmod u+x ./build.sh && ./build.sh |  | ||||||
|           cd .. |  | ||||||
|       - name: Set up Go |  | ||||||
|         uses: actions/setup-go@v3 |  | ||||||
|         with: |  | ||||||
|           go-version: '>=1.18.0' |  | ||||||
|       - name: Build Backend (amd64) |  | ||||||
|         run: | |  | ||||||
|           go mod download |  | ||||||
|           go build -ldflags "-s -w -X 'github.com/songquanpeng/one-api/common.Version=$(git describe --tags)' -extldflags '-static'" -o one-api |  | ||||||
|  |  | ||||||
|       - name: Build Backend (arm64) |  | ||||||
|         run: | |  | ||||||
|           sudo apt-get update |  | ||||||
|           sudo apt-get install gcc-aarch64-linux-gnu |  | ||||||
|           CC=aarch64-linux-gnu-gcc CGO_ENABLED=1 GOOS=linux GOARCH=arm64 go build -ldflags "-s -w -X 'one-api/common.Version=$(git describe --tags)' -extldflags '-static'" -o one-api-arm64 |  | ||||||
|  |  | ||||||
|       - name: Release |  | ||||||
|         uses: softprops/action-gh-release@v1 |  | ||||||
|         if: startsWith(github.ref, 'refs/tags/') |  | ||||||
|         with: |  | ||||||
|           files: | |  | ||||||
|             one-api |  | ||||||
|             one-api-arm64 |  | ||||||
|           draft: true |  | ||||||
|           generate_release_notes: true |  | ||||||
|         env: |  | ||||||
|           GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} |  | ||||||
							
								
								
									
										50
									
								
								.github/workflows/macos-release.yml
									
									
									
									
										vendored
									
									
								
							
							
						
						
									
										50
									
								
								.github/workflows/macos-release.yml
									
									
									
									
										vendored
									
									
								
							| @@ -1,50 +0,0 @@ | |||||||
| name: macOS Release |  | ||||||
| permissions: |  | ||||||
|   contents: write |  | ||||||
|  |  | ||||||
| on: |  | ||||||
|   push: |  | ||||||
|     tags: |  | ||||||
|       - '*' |  | ||||||
|       - '!*-alpha*' |  | ||||||
|   workflow_dispatch: |  | ||||||
|     inputs: |  | ||||||
|       name: |  | ||||||
|         description: 'reason' |  | ||||||
|         required: false |  | ||||||
| jobs: |  | ||||||
|   release: |  | ||||||
|     runs-on: macos-latest |  | ||||||
|     steps: |  | ||||||
|       - name: Checkout |  | ||||||
|         uses: actions/checkout@v3 |  | ||||||
|         with: |  | ||||||
|           fetch-depth: 0 |  | ||||||
|       - uses: actions/setup-node@v3 |  | ||||||
|         with: |  | ||||||
|           node-version: 16 |  | ||||||
|       - name: Build Frontend |  | ||||||
|         env: |  | ||||||
|           CI: "" |  | ||||||
|         run: | |  | ||||||
|           cd web |  | ||||||
|           git describe --tags > VERSION |  | ||||||
|           REACT_APP_VERSION=$(git describe --tags) chmod u+x ./build.sh && ./build.sh |  | ||||||
|           cd .. |  | ||||||
|       - name: Set up Go |  | ||||||
|         uses: actions/setup-go@v3 |  | ||||||
|         with: |  | ||||||
|           go-version: '>=1.18.0' |  | ||||||
|       - name: Build Backend |  | ||||||
|         run: | |  | ||||||
|           go mod download |  | ||||||
|           go build -ldflags "-X 'github.com/songquanpeng/one-api/common.Version=$(git describe --tags)'" -o one-api-macos |  | ||||||
|       - name: Release |  | ||||||
|         uses: softprops/action-gh-release@v1 |  | ||||||
|         if: startsWith(github.ref, 'refs/tags/') |  | ||||||
|         with: |  | ||||||
|           files: one-api-macos |  | ||||||
|           draft: true |  | ||||||
|           generate_release_notes: true |  | ||||||
|         env: |  | ||||||
|           GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} |  | ||||||
							
								
								
									
										53
									
								
								.github/workflows/windows-release.yml
									
									
									
									
										vendored
									
									
								
							
							
						
						
									
										53
									
								
								.github/workflows/windows-release.yml
									
									
									
									
										vendored
									
									
								
							| @@ -1,53 +0,0 @@ | |||||||
| name: Windows Release |  | ||||||
| permissions: |  | ||||||
|   contents: write |  | ||||||
|  |  | ||||||
| on: |  | ||||||
|   push: |  | ||||||
|     tags: |  | ||||||
|       - '*' |  | ||||||
|       - '!*-alpha*' |  | ||||||
|   workflow_dispatch: |  | ||||||
|     inputs: |  | ||||||
|       name: |  | ||||||
|         description: 'reason' |  | ||||||
|         required: false |  | ||||||
| jobs: |  | ||||||
|   release: |  | ||||||
|     runs-on: windows-latest |  | ||||||
|     defaults: |  | ||||||
|       run: |  | ||||||
|         shell: bash |  | ||||||
|     steps: |  | ||||||
|       - name: Checkout |  | ||||||
|         uses: actions/checkout@v3 |  | ||||||
|         with: |  | ||||||
|           fetch-depth: 0 |  | ||||||
|       - uses: actions/setup-node@v3 |  | ||||||
|         with: |  | ||||||
|           node-version: 16 |  | ||||||
|       - name: Build Frontend |  | ||||||
|         env: |  | ||||||
|           CI: "" |  | ||||||
|         run: | |  | ||||||
|           cd web/default |  | ||||||
|           npm install |  | ||||||
|           REACT_APP_VERSION=$(git describe --tags) npm run build |  | ||||||
|           cd ../.. |  | ||||||
|       - name: Set up Go |  | ||||||
|         uses: actions/setup-go@v3 |  | ||||||
|         with: |  | ||||||
|           go-version: '>=1.18.0' |  | ||||||
|       - name: Build Backend |  | ||||||
|         run: | |  | ||||||
|           go mod download |  | ||||||
|           go build -ldflags "-s -w -X 'github.com/songquanpeng/one-api/common.Version=$(git describe --tags)'" -o one-api.exe |  | ||||||
|       - name: Release |  | ||||||
|         uses: softprops/action-gh-release@v1 |  | ||||||
|         if: startsWith(github.ref, 'refs/tags/') |  | ||||||
|         with: |  | ||||||
|           files: one-api.exe |  | ||||||
|           draft: true |  | ||||||
|           generate_release_notes: true |  | ||||||
|         env: |  | ||||||
|           GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} |  | ||||||
							
								
								
									
										3
									
								
								.gitignore
									
									
									
									
										vendored
									
									
								
							
							
						
						
									
										3
									
								
								.gitignore
									
									
									
									
										vendored
									
									
								
							| @@ -6,5 +6,4 @@ upload | |||||||
| build | build | ||||||
| *.db-journal | *.db-journal | ||||||
| logs | logs | ||||||
| data | data | ||||||
| /web/node_modules |  | ||||||
| @@ -23,7 +23,7 @@ ADD go.mod go.sum ./ | |||||||
| RUN go mod download | RUN go mod download | ||||||
| COPY . . | COPY . . | ||||||
| COPY --from=builder /web/build ./web/build | COPY --from=builder /web/build ./web/build | ||||||
| RUN go build -ldflags "-s -w -X 'github.com/songquanpeng/one-api/common.Version=$(cat VERSION)' -extldflags '-static'" -o one-api | RUN go build -ldflags "-s -w -X 'one-api/common.Version=$(cat VERSION)' -extldflags '-static'" -o one-api | ||||||
|  |  | ||||||
| FROM alpine | FROM alpine | ||||||
|  |  | ||||||
|   | |||||||
| @@ -134,12 +134,12 @@ The initial account username is `root` and password is `123456`. | |||||||
|    git clone https://github.com/songquanpeng/one-api.git |    git clone https://github.com/songquanpeng/one-api.git | ||||||
|     |     | ||||||
|    # Build the frontend |    # Build the frontend | ||||||
|    cd one-api/web/default |    cd one-api/web | ||||||
|    npm install |    npm install | ||||||
|    npm run build |    npm run build | ||||||
|     |     | ||||||
|    # Build the backend |    # Build the backend | ||||||
|    cd ../.. |    cd .. | ||||||
|    go mod download |    go mod download | ||||||
|    go build -ldflags "-s -w" -o one-api |    go build -ldflags "-s -w" -o one-api | ||||||
|    ``` |    ``` | ||||||
|   | |||||||
| @@ -135,12 +135,12 @@ sudo service nginx restart | |||||||
|    git clone https://github.com/songquanpeng/one-api.git |    git clone https://github.com/songquanpeng/one-api.git | ||||||
|  |  | ||||||
|    # フロントエンドのビルド |    # フロントエンドのビルド | ||||||
|    cd one-api/web/default |    cd one-api/web | ||||||
|    npm install |    npm install | ||||||
|    npm run build |    npm run build | ||||||
|  |  | ||||||
|    # バックエンドのビルド |    # バックエンドのビルド | ||||||
|    cd ../.. |    cd .. | ||||||
|    go mod download |    go mod download | ||||||
|    go build -ldflags "-s -w" -o one-api |    go build -ldflags "-s -w" -o one-api | ||||||
|    ``` |    ``` | ||||||
|   | |||||||
							
								
								
									
										17
									
								
								README.md
									
									
									
									
									
								
							
							
						
						
									
										17
									
								
								README.md
									
									
									
									
									
								
							| @@ -67,18 +67,12 @@ _✨ 通过标准的 OpenAI API 格式访问所有的大模型,开箱即用  | |||||||
|    + [x] [OpenAI ChatGPT 系列模型](https://platform.openai.com/docs/guides/gpt/chat-completions-api)(支持 [Azure OpenAI API](https://learn.microsoft.com/en-us/azure/ai-services/openai/reference)) |    + [x] [OpenAI ChatGPT 系列模型](https://platform.openai.com/docs/guides/gpt/chat-completions-api)(支持 [Azure OpenAI API](https://learn.microsoft.com/en-us/azure/ai-services/openai/reference)) | ||||||
|    + [x] [Anthropic Claude 系列模型](https://anthropic.com) |    + [x] [Anthropic Claude 系列模型](https://anthropic.com) | ||||||
|    + [x] [Google PaLM2/Gemini 系列模型](https://developers.generativeai.google) |    + [x] [Google PaLM2/Gemini 系列模型](https://developers.generativeai.google) | ||||||
|    + [x] [Mistral 系列模型](https://mistral.ai/) |  | ||||||
|    + [x] [百度文心一言系列模型](https://cloud.baidu.com/doc/WENXINWORKSHOP/index.html) |    + [x] [百度文心一言系列模型](https://cloud.baidu.com/doc/WENXINWORKSHOP/index.html) | ||||||
|    + [x] [阿里通义千问系列模型](https://help.aliyun.com/document_detail/2400395.html) |    + [x] [阿里通义千问系列模型](https://help.aliyun.com/document_detail/2400395.html) | ||||||
|    + [x] [讯飞星火认知大模型](https://www.xfyun.cn/doc/spark/Web.html) |    + [x] [讯飞星火认知大模型](https://www.xfyun.cn/doc/spark/Web.html) | ||||||
|    + [x] [智谱 ChatGLM 系列模型](https://bigmodel.cn) |    + [x] [智谱 ChatGLM 系列模型](https://bigmodel.cn) | ||||||
|    + [x] [360 智脑](https://ai.360.cn) |    + [x] [360 智脑](https://ai.360.cn) | ||||||
|    + [x] [腾讯混元大模型](https://cloud.tencent.com/document/product/1729) |    + [x] [腾讯混元大模型](https://cloud.tencent.com/document/product/1729) | ||||||
|    + [x] [Moonshot AI](https://platform.moonshot.cn/) |  | ||||||
|    + [x] [百川大模型](https://platform.baichuan-ai.com) |  | ||||||
|    + [ ] [字节云雀大模型](https://www.volcengine.com/product/ark) (WIP) |  | ||||||
|    + [x] [MINIMAX](https://api.minimax.chat/) |  | ||||||
|    + [x] [Groq](https://wow.groq.com/) |  | ||||||
| 2. 支持配置镜像以及众多[第三方代理服务](https://iamazing.cn/page/openai-api-third-party-services)。 | 2. 支持配置镜像以及众多[第三方代理服务](https://iamazing.cn/page/openai-api-third-party-services)。 | ||||||
| 3. 支持通过**负载均衡**的方式访问多个渠道。 | 3. 支持通过**负载均衡**的方式访问多个渠道。 | ||||||
| 4. 支持 **stream 模式**,可以通过流式传输实现打字机效果。 | 4. 支持 **stream 模式**,可以通过流式传输实现打字机效果。 | ||||||
| @@ -106,7 +100,6 @@ _✨ 通过标准的 OpenAI API 格式访问所有的大模型,开箱即用  | |||||||
|     + [GitHub 开放授权](https://github.com/settings/applications/new)。 |     + [GitHub 开放授权](https://github.com/settings/applications/new)。 | ||||||
|     + 微信公众号授权(需要额外部署 [WeChat Server](https://github.com/songquanpeng/wechat-server))。 |     + 微信公众号授权(需要额外部署 [WeChat Server](https://github.com/songquanpeng/wechat-server))。 | ||||||
| 23. 支持主题切换,设置环境变量 `THEME` 即可,默认为 `default`,欢迎 PR 更多主题,具体参考[此处](./web/README.md)。 | 23. 支持主题切换,设置环境变量 `THEME` 即可,默认为 `default`,欢迎 PR 更多主题,具体参考[此处](./web/README.md)。 | ||||||
| 24. 配合 [Message Pusher](https://github.com/songquanpeng/message-pusher) 可将报警信息推送到多种 App 上。 |  | ||||||
|  |  | ||||||
| ## 部署 | ## 部署 | ||||||
| ### 基于 Docker 进行部署 | ### 基于 Docker 进行部署 | ||||||
| @@ -181,12 +174,12 @@ docker-compose ps | |||||||
|    git clone https://github.com/songquanpeng/one-api.git |    git clone https://github.com/songquanpeng/one-api.git | ||||||
|     |     | ||||||
|    # 构建前端 |    # 构建前端 | ||||||
|    cd one-api/web/default |    cd one-api/web | ||||||
|    npm install |    npm install | ||||||
|    npm run build |    npm run build | ||||||
|     |     | ||||||
|    # 构建后端 |    # 构建后端 | ||||||
|    cd ../.. |    cd .. | ||||||
|    go mod download |    go mod download | ||||||
|    go build -ldflags "-s -w" -o one-api |    go build -ldflags "-s -w" -o one-api | ||||||
|    ```` |    ```` | ||||||
| @@ -376,9 +369,6 @@ graph LR | |||||||
| 16. `SQLITE_BUSY_TIMEOUT`:SQLite 锁等待超时设置,单位为毫秒,默认 `3000`。 | 16. `SQLITE_BUSY_TIMEOUT`:SQLite 锁等待超时设置,单位为毫秒,默认 `3000`。 | ||||||
| 17. `GEMINI_SAFETY_SETTING`:Gemini 的安全设置,默认 `BLOCK_NONE`。 | 17. `GEMINI_SAFETY_SETTING`:Gemini 的安全设置,默认 `BLOCK_NONE`。 | ||||||
| 18. `THEME`:系统的主题设置,默认为 `default`,具体可选值参考[此处](./web/README.md)。 | 18. `THEME`:系统的主题设置,默认为 `default`,具体可选值参考[此处](./web/README.md)。 | ||||||
| 19. `ENABLE_METRIC`:是否根据请求成功率禁用渠道,默认不开启,可选值为 `true` 和 `false`。 |  | ||||||
| 20. `METRIC_QUEUE_SIZE`:请求成功率统计队列大小,默认为 `10`。 |  | ||||||
| 21. `METRIC_SUCCESS_RATE_THRESHOLD`:请求成功率阈值,默认为 `0.8`。 |  | ||||||
|  |  | ||||||
| ### 命令行参数 | ### 命令行参数 | ||||||
| 1. `--port <port_number>`: 指定服务器监听的端口号,默认为 `3000`。 | 1. `--port <port_number>`: 指定服务器监听的端口号,默认为 `3000`。 | ||||||
| @@ -424,9 +414,6 @@ https://openai.justsong.cn | |||||||
| 8. 升级之前数据库需要做变更吗? | 8. 升级之前数据库需要做变更吗? | ||||||
|    + 一般情况下不需要,系统将在初始化的时候自动调整。 |    + 一般情况下不需要,系统将在初始化的时候自动调整。 | ||||||
|    + 如果需要的话,我会在更新日志中说明,并给出脚本。 |    + 如果需要的话,我会在更新日志中说明,并给出脚本。 | ||||||
| 9. 手动修改数据库后报错:`数据库一致性已被破坏,请联系管理员`? |  | ||||||
|    + 这是检测到 ability 表里有些记录的通道 id 是不存在的,这大概率是因为你删了 channel 表里的记录但是没有同步在 ability 表里清理无效的通道。 |  | ||||||
|    + 对于每一个通道,其所支持的模型都需要有一个专门的 ability 表的记录,表示该通道支持该模型。 |  | ||||||
|  |  | ||||||
| ## 相关项目 | ## 相关项目 | ||||||
| * [FastGPT](https://github.com/labring/FastGPT): 基于 LLM 大语言模型的知识库问答系统 | * [FastGPT](https://github.com/labring/FastGPT): 基于 LLM 大语言模型的知识库问答系统 | ||||||
|   | |||||||
| @@ -1,29 +0,0 @@ | |||||||
| package blacklist |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"fmt" |  | ||||||
| 	"sync" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| var blackList sync.Map |  | ||||||
|  |  | ||||||
| func init() { |  | ||||||
| 	blackList = sync.Map{} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func userId2Key(id int) string { |  | ||||||
| 	return fmt.Sprintf("userid_%d", id) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func BanUser(id int) { |  | ||||||
| 	blackList.Store(userId2Key(id), true) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func UnbanUser(id int) { |  | ||||||
| 	blackList.Delete(userId2Key(id)) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func IsUserBanned(id int) bool { |  | ||||||
| 	_, ok := blackList.Load(userId2Key(id)) |  | ||||||
| 	return ok |  | ||||||
| } |  | ||||||
| @@ -1,137 +0,0 @@ | |||||||
| package config |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"github.com/songquanpeng/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 DebugSQLEnabled = os.Getenv("DEBUG_SQL") == "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 MessagePusherAddress = "" |  | ||||||
| var MessagePusherToken = "" |  | ||||||
|  |  | ||||||
| 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 |  | ||||||
|  |  | ||||||
| var EnableMetric = helper.GetOrDefaultEnvBool("ENABLE_METRIC", false) |  | ||||||
| var MetricQueueSize = helper.GetOrDefaultEnvInt("METRIC_QUEUE_SIZE", 10) |  | ||||||
| var MetricSuccessRateThreshold = helper.GetOrDefaultEnvFloat64("METRIC_SUCCESS_RATE_THRESHOLD", 0.8) |  | ||||||
| var MetricSuccessChanSize = helper.GetOrDefaultEnvInt("METRIC_SUCCESS_CHAN_SIZE", 1024) |  | ||||||
| var MetricFailChanSize = helper.GetOrDefaultEnvInt("METRIC_FAIL_CHAN_SIZE", 128) |  | ||||||
| @@ -1,9 +1,115 @@ | |||||||
| package common | package common | ||||||
|  |  | ||||||
| import "time" | import ( | ||||||
|  | 	"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 SMTPAuthLoginEnabled = false | ||||||
|  |  | ||||||
|  | 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 | ||||||
| @@ -12,10 +118,37 @@ 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 | ||||||
| 	UserStatusDeleted  = 3 |  | ||||||
| ) | ) | ||||||
|  |  | ||||||
| const ( | const ( | ||||||
| @@ -39,77 +172,57 @@ const ( | |||||||
| ) | ) | ||||||
|  |  | ||||||
| const ( | const ( | ||||||
| 	ChannelTypeUnknown = iota | 	ChannelTypeUnknown        = 0 | ||||||
| 	ChannelTypeOpenAI | 	ChannelTypeOpenAI         = 1 | ||||||
| 	ChannelTypeAPI2D | 	ChannelTypeAPI2D          = 2 | ||||||
| 	ChannelTypeAzure | 	ChannelTypeAzure          = 3 | ||||||
| 	ChannelTypeCloseAI | 	ChannelTypeCloseAI        = 4 | ||||||
| 	ChannelTypeOpenAISB | 	ChannelTypeOpenAISB       = 5 | ||||||
| 	ChannelTypeOpenAIMax | 	ChannelTypeOpenAIMax      = 6 | ||||||
| 	ChannelTypeOhMyGPT | 	ChannelTypeOhMyGPT        = 7 | ||||||
| 	ChannelTypeCustom | 	ChannelTypeCustom         = 8 | ||||||
| 	ChannelTypeAILS | 	ChannelTypeAILS           = 9 | ||||||
| 	ChannelTypeAIProxy | 	ChannelTypeAIProxy        = 10 | ||||||
| 	ChannelTypePaLM | 	ChannelTypePaLM           = 11 | ||||||
| 	ChannelTypeAPI2GPT | 	ChannelTypeAPI2GPT        = 12 | ||||||
| 	ChannelTypeAIGC2D | 	ChannelTypeAIGC2D         = 13 | ||||||
| 	ChannelTypeAnthropic | 	ChannelTypeAnthropic      = 14 | ||||||
| 	ChannelTypeBaidu | 	ChannelTypeBaidu          = 15 | ||||||
| 	ChannelTypeZhipu | 	ChannelTypeZhipu          = 16 | ||||||
| 	ChannelTypeAli | 	ChannelTypeAli            = 17 | ||||||
| 	ChannelTypeXunfei | 	ChannelTypeXunfei         = 18 | ||||||
| 	ChannelType360 | 	ChannelType360            = 19 | ||||||
| 	ChannelTypeOpenRouter | 	ChannelTypeOpenRouter     = 20 | ||||||
| 	ChannelTypeAIProxyLibrary | 	ChannelTypeAIProxyLibrary = 21 | ||||||
| 	ChannelTypeFastGPT | 	ChannelTypeFastGPT        = 22 | ||||||
| 	ChannelTypeTencent | 	ChannelTypeTencent        = 23 | ||||||
| 	ChannelTypeGemini | 	ChannelTypeGemini         = 24 | ||||||
| 	ChannelTypeMoonshot |  | ||||||
| 	ChannelTypeBaichuan |  | ||||||
| 	ChannelTypeMinimax |  | ||||||
| 	ChannelTypeMistral |  | ||||||
| 	ChannelTypeGroq |  | ||||||
|  |  | ||||||
| 	ChannelTypeDummy |  | ||||||
| ) | ) | ||||||
|  |  | ||||||
| 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 | ||||||
| 	"https://generativelanguage.googleapis.com", // 11 | 	"",                                  // 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 | ||||||
| 	"https://generativelanguage.googleapis.com", // 24 | 	"",                                  //24 | ||||||
| 	"https://api.moonshot.cn",                   // 25 |  | ||||||
| 	"https://api.baichuan-ai.com",               // 26 |  | ||||||
| 	"https://api.minimax.chat",                  // 27 |  | ||||||
| 	"https://api.mistral.ai",                    // 28 |  | ||||||
| 	"https://api.groq.com/openai",               // 29 |  | ||||||
| } | } | ||||||
|  |  | ||||||
| const ( |  | ||||||
| 	ConfigKeyPrefix = "cfg_" |  | ||||||
|  |  | ||||||
| 	ConfigKeyAPIVersion = ConfigKeyPrefix + "api_version" |  | ||||||
| 	ConfigKeyLibraryID  = ConfigKeyPrefix + "library_id" |  | ||||||
| 	ConfigKeyPlugin     = ConfigKeyPrefix + "plugin" |  | ||||||
| ) |  | ||||||
|   | |||||||
| @@ -1,9 +1,7 @@ | |||||||
| package common | package common | ||||||
|  |  | ||||||
| import "github.com/songquanpeng/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 = helper.GetOrDefaultEnvInt("SQLITE_BUSY_TIMEOUT", 3000) | var SQLiteBusyTimeout = GetOrDefault("SQLITE_BUSY_TIMEOUT", 3000) | ||||||
|   | |||||||
| @@ -1,27 +1,50 @@ | |||||||
| package message | package common | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"crypto/rand" | 	"crypto/rand" | ||||||
| 	"crypto/tls" | 	"crypto/tls" | ||||||
| 	"encoding/base64" | 	"encoding/base64" | ||||||
|  | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"net/smtp" | 	"net/smtp" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
| 
 | 
 | ||||||
| func SendEmail(subject string, receiver string, content string) error { | type loginAuth struct { | ||||||
| 	if receiver == "" { | 	username, password string | ||||||
| 		return fmt.Errorf("receiver is empty") | } | ||||||
|  | 
 | ||||||
|  | func LoginAuth(username, password string) smtp.Auth { | ||||||
|  | 	return &loginAuth{username, password} | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func (a *loginAuth) Start(_ *smtp.ServerInfo) (string, []byte, error) { | ||||||
|  | 	return "LOGIN", []byte(a.username), nil | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func (a *loginAuth) Next(fromServer []byte, more bool) ([]byte, error) { | ||||||
|  | 	if more { | ||||||
|  | 		switch string(fromServer) { | ||||||
|  | 		case "Username:": | ||||||
|  | 			return []byte(a.username), nil | ||||||
|  | 		case "Password:": | ||||||
|  | 			return []byte(a.password), nil | ||||||
|  | 		default: | ||||||
|  | 			return nil, errors.New("unknown command from server during login auth") | ||||||
|  | 		} | ||||||
| 	} | 	} | ||||||
| 	if config.SMTPFrom == "" { // for compatibility | 	return nil, nil | ||||||
| 		config.SMTPFrom = config.SMTPAccount | } | ||||||
|  | 
 | ||||||
|  | func SendEmail(subject string, receiver string, content string) error { | ||||||
|  | 	if SMTPFrom == "" { // for compatibility | ||||||
|  | 		SMTPFrom = 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(config.SMTPFrom, "@") | 	parts := strings.Split(SMTPFrom, "@") | ||||||
| 	var domain string | 	var domain string | ||||||
| 	if len(parts) > 1 { | 	if len(parts) > 1 { | ||||||
| 		domain = parts[1] | 		domain = parts[1] | ||||||
| @@ -40,21 +63,27 @@ 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, config.SystemName, config.SMTPFrom, encodedSubject, messageId, time.Now().Format(time.RFC1123Z), content)) | 		receiver, SystemName, SMTPFrom, encodedSubject, messageId, time.Now().Format(time.RFC1123Z), content)) | ||||||
| 	auth := smtp.PlainAuth("", config.SMTPAccount, config.SMTPToken, config.SMTPServer) | 
 | ||||||
| 	addr := fmt.Sprintf("%s:%d", config.SMTPServer, config.SMTPPort) | 	var auth smtp.Auth | ||||||
|  | 	if SMTPAuthLoginEnabled { | ||||||
|  | 		auth = LoginAuth(SMTPAccount, SMTPToken) | ||||||
|  | 	} else { | ||||||
|  | 		auth = smtp.PlainAuth("", SMTPAccount, SMTPToken, SMTPServer) | ||||||
|  | 	} | ||||||
|  | 	addr := fmt.Sprintf("%s:%d", SMTPServer, SMTPPort) | ||||||
| 	to := strings.Split(receiver, ";") | 	to := strings.Split(receiver, ";") | ||||||
| 
 | 
 | ||||||
| 	if config.SMTPPort == 465 { | 	if SMTPPort == 465 { | ||||||
| 		tlsConfig := &tls.Config{ | 		tlsConfig := &tls.Config{ | ||||||
| 			InsecureSkipVerify: true, | 			InsecureSkipVerify: true, | ||||||
| 			ServerName:         config.SMTPServer, | 			ServerName:         SMTPServer, | ||||||
| 		} | 		} | ||||||
| 		conn, err := tls.Dial("tcp", fmt.Sprintf("%s:%d", config.SMTPServer, config.SMTPPort), tlsConfig) | 		conn, err := tls.Dial("tcp", fmt.Sprintf("%s:%d", SMTPServer, SMTPPort), tlsConfig) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		client, err := smtp.NewClient(conn, config.SMTPServer) | 		client, err := smtp.NewClient(conn, SMTPServer) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| @@ -62,7 +91,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(config.SMTPFrom); err != nil { | 		if err = client.Mail(SMTPFrom); err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		receiverEmails := strings.Split(receiver, ";") | 		receiverEmails := strings.Split(receiver, ";") | ||||||
| @@ -84,7 +113,7 @@ func SendEmail(subject string, receiver string, content string) error { | |||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 	} else { | 	} else { | ||||||
| 		err = smtp.SendMail(addr, auth, config.SMTPAccount, to, mail) | 		err = smtp.SendMail(addr, auth, SMTPAccount, to, mail) | ||||||
| 	} | 	} | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
| @@ -15,7 +15,10 @@ type embedFileSystem struct { | |||||||
|  |  | ||||||
| func (e embedFileSystem) Exists(prefix string, path string) bool { | func (e embedFileSystem) Exists(prefix string, path string) bool { | ||||||
| 	_, err := e.Open(path) | 	_, err := e.Open(path) | ||||||
| 	return err == nil | 	if err != nil { | ||||||
|  | 		return false | ||||||
|  | 	} | ||||||
|  | 	return true | ||||||
| } | } | ||||||
|  |  | ||||||
| func EmbedFolder(fsEmbed embed.FS, targetPath string) static.ServeFileSystem { | func EmbedFolder(fsEmbed embed.FS, targetPath string) static.ServeFileSystem { | ||||||
|   | |||||||
| @@ -8,24 +8,12 @@ import ( | |||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| const KeyRequestBody = "key_request_body" | func UnmarshalBodyReusable(c *gin.Context, v any) error { | ||||||
|  |  | ||||||
| func GetRequestBody(c *gin.Context) ([]byte, error) { |  | ||||||
| 	requestBody, _ := c.Get(KeyRequestBody) |  | ||||||
| 	if requestBody != nil { |  | ||||||
| 		return requestBody.([]byte), nil |  | ||||||
| 	} |  | ||||||
| 	requestBody, err := io.ReadAll(c.Request.Body) | 	requestBody, err := io.ReadAll(c.Request.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return nil, err | 		return err | ||||||
| 	} | 	} | ||||||
| 	_ = c.Request.Body.Close() | 	err = c.Request.Body.Close() | ||||||
| 	c.Set(KeyRequestBody, requestBody) |  | ||||||
| 	return requestBody.([]byte), nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func UnmarshalBodyReusable(c *gin.Context, v any) error { |  | ||||||
| 	requestBody, err := GetRequestBody(c) |  | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return err | 		return err | ||||||
| 	} | 	} | ||||||
| @@ -43,11 +31,3 @@ func UnmarshalBodyReusable(c *gin.Context, v any) error { | |||||||
| 	c.Request.Body = io.NopCloser(bytes.NewBuffer(requestBody)) | 	c.Request.Body = io.NopCloser(bytes.NewBuffer(requestBody)) | ||||||
| 	return nil | 	return nil | ||||||
| } | } | ||||||
|  |  | ||||||
| func SetEventStreamHeaders(c *gin.Context) { |  | ||||||
| 	c.Writer.Header().Set("Content-Type", "text/event-stream") |  | ||||||
| 	c.Writer.Header().Set("Cache-Control", "no-cache") |  | ||||||
| 	c.Writer.Header().Set("Connection", "keep-alive") |  | ||||||
| 	c.Writer.Header().Set("Transfer-Encoding", "chunked") |  | ||||||
| 	c.Writer.Header().Set("X-Accel-Buffering", "no") |  | ||||||
| } |  | ||||||
|   | |||||||
| @@ -1,9 +1,6 @@ | |||||||
| package common | package common | ||||||
|  |  | ||||||
| import ( | import "encoding/json" | ||||||
| 	"encoding/json" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| var GroupRatio = map[string]float64{ | var GroupRatio = map[string]float64{ | ||||||
| 	"default": 1, | 	"default": 1, | ||||||
| @@ -14,7 +11,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 { | ||||||
| 		logger.SysError("error marshalling model ratio: " + err.Error()) | 		SysError("error marshalling model ratio: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return string(jsonBytes) | 	return string(jsonBytes) | ||||||
| } | } | ||||||
| @@ -27,7 +24,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 { | ||||||
| 		logger.SysError("group ratio not found: " + name) | 		SysError("group ratio not found: " + name) | ||||||
| 		return 1 | 		return 1 | ||||||
| 	} | 	} | ||||||
| 	return ratio | 	return ratio | ||||||
|   | |||||||
| @@ -1,253 +0,0 @@ | |||||||
| package helper |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/google/uuid" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"html/template" |  | ||||||
| 	"log" |  | ||||||
| 	"math/rand" |  | ||||||
| 	"net" |  | ||||||
| 	"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 := inter.(type) { |  | ||||||
| 	case string: |  | ||||||
| 		return inter |  | ||||||
| 	case int: |  | ||||||
| 		return fmt.Sprintf("%d", inter) |  | ||||||
| 	case float64: |  | ||||||
| 		return fmt.Sprintf("%f", inter) |  | ||||||
| 	} |  | ||||||
| 	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" |  | ||||||
| const keyNumbers = "0123456789" |  | ||||||
|  |  | ||||||
| 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 GetRandomNumberString(length int) string { |  | ||||||
| 	rand.Seed(time.Now().UnixNano()) |  | ||||||
| 	key := make([]byte, length) |  | ||||||
| 	for i := 0; i < length; i++ { |  | ||||||
| 		key[i] = keyNumbers[rand.Intn(len(keyNumbers))] |  | ||||||
| 	} |  | ||||||
| 	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 GetOrDefaultEnvBool(env string, defaultValue bool) bool { |  | ||||||
| 	if env == "" || os.Getenv(env) == "" { |  | ||||||
| 		return defaultValue |  | ||||||
| 	} |  | ||||||
| 	return os.Getenv(env) == "true" |  | ||||||
| } |  | ||||||
|  |  | ||||||
| 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 GetOrDefaultEnvFloat64(env string, defaultValue float64) float64 { |  | ||||||
| 	if env == "" || os.Getenv(env) == "" { |  | ||||||
| 		return defaultValue |  | ||||||
| 	} |  | ||||||
| 	num, err := strconv.ParseFloat(os.Getenv(env), 64) |  | ||||||
| 	if err != nil { |  | ||||||
| 		logger.SysError(fmt.Sprintf("failed to parse %s: %s, using default value: %f", 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 |  | ||||||
| } |  | ||||||
| @@ -12,7 +12,7 @@ import ( | |||||||
| 	"strings" | 	"strings" | ||||||
| 	"testing" | 	"testing" | ||||||
|  |  | ||||||
| 	img "github.com/songquanpeng/one-api/common/image" | 	img "one-api/common/image" | ||||||
|  |  | ||||||
| 	"github.com/stretchr/testify/assert" | 	"github.com/stretchr/testify/assert" | ||||||
| 	_ "golang.org/x/image/webp" | 	_ "golang.org/x/image/webp" | ||||||
|   | |||||||
| @@ -3,8 +3,6 @@ package common | |||||||
| import ( | import ( | ||||||
| 	"flag" | 	"flag" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"log" | 	"log" | ||||||
| 	"os" | 	"os" | ||||||
| 	"path/filepath" | 	"path/filepath" | ||||||
| @@ -39,9 +37,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" { | ||||||
| 			logger.SysError("SESSION_SECRET is set to an example value, please change it to a random string.") | 			SysError("SESSION_SECRET is set to an example value, please change it to a random string.") | ||||||
| 		} else { | 		} else { | ||||||
| 			config.SessionSecret = os.Getenv("SESSION_SECRET") | 			SessionSecret = os.Getenv("SESSION_SECRET") | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	if os.Getenv("SQLITE_PATH") != "" { | 	if os.Getenv("SQLITE_PATH") != "" { | ||||||
| @@ -59,6 +57,5 @@ func init() { | |||||||
| 				log.Fatal(err) | 				log.Fatal(err) | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 		logger.LogDir = *LogDir |  | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -1,4 +1,4 @@ | |||||||
| package logger | package common | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"context" | 	"context" | ||||||
| @@ -13,7 +13,6 @@ import ( | |||||||
| ) | ) | ||||||
| 
 | 
 | ||||||
| const ( | const ( | ||||||
| 	loggerDEBUG = "DEBUG" |  | ||||||
| 	loggerINFO  = "INFO" | 	loggerINFO  = "INFO" | ||||||
| 	loggerWarn  = "WARN" | 	loggerWarn  = "WARN" | ||||||
| 	loggerError = "ERR" | 	loggerError = "ERR" | ||||||
| @@ -26,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") | ||||||
| @@ -36,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") | ||||||
| @@ -56,38 +55,18 @@ 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 Debug(ctx context.Context, msg string) { | func LogInfo(ctx context.Context, msg string) { | ||||||
| 	logHelper(ctx, loggerDEBUG, msg) |  | ||||||
| } |  | ||||||
| 
 |  | ||||||
| func Info(ctx context.Context, msg string) { |  | ||||||
| 	logHelper(ctx, loggerINFO, msg) | 	logHelper(ctx, loggerINFO, msg) | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Warn(ctx context.Context, msg string) { | func LogWarn(ctx context.Context, msg string) { | ||||||
| 	logHelper(ctx, loggerWarn, msg) | 	logHelper(ctx, loggerWarn, msg) | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Error(ctx context.Context, msg string) { | func LogError(ctx context.Context, msg string) { | ||||||
| 	logHelper(ctx, loggerError, msg) | 	logHelper(ctx, loggerError, msg) | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Debugf(ctx context.Context, format string, a ...any) { |  | ||||||
| 	Debug(ctx, fmt.Sprintf(format, a...)) |  | ||||||
| } |  | ||||||
| 
 |  | ||||||
| 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 { | ||||||
| @@ -111,3 +90,11 @@ 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) | ||||||
|  | 	} | ||||||
|  | } | ||||||
| @@ -1,7 +0,0 @@ | |||||||
| package logger |  | ||||||
|  |  | ||||||
| const ( |  | ||||||
| 	RequestIdKey = "X-Oneapi-Request-Id" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| var LogDir string |  | ||||||
| @@ -1,22 +0,0 @@ | |||||||
| package message |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| const ( |  | ||||||
| 	ByAll           = "all" |  | ||||||
| 	ByEmail         = "email" |  | ||||||
| 	ByMessagePusher = "message_pusher" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| func Notify(by string, title string, description string, content string) error { |  | ||||||
| 	if by == ByEmail { |  | ||||||
| 		return SendEmail(title, config.RootUserEmail, content) |  | ||||||
| 	} |  | ||||||
| 	if by == ByMessagePusher { |  | ||||||
| 		return SendMessage(title, description, content) |  | ||||||
| 	} |  | ||||||
| 	return fmt.Errorf("unknown notify method: %s", by) |  | ||||||
| } |  | ||||||
| @@ -1,53 +0,0 @@ | |||||||
| package message |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"bytes" |  | ||||||
| 	"encoding/json" |  | ||||||
| 	"errors" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"net/http" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| type request struct { |  | ||||||
| 	Title       string `json:"title"` |  | ||||||
| 	Description string `json:"description"` |  | ||||||
| 	Content     string `json:"content"` |  | ||||||
| 	URL         string `json:"url"` |  | ||||||
| 	Channel     string `json:"channel"` |  | ||||||
| 	Token       string `json:"token"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type response struct { |  | ||||||
| 	Success bool   `json:"success"` |  | ||||||
| 	Message string `json:"message"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func SendMessage(title string, description string, content string) error { |  | ||||||
| 	if config.MessagePusherAddress == "" { |  | ||||||
| 		return errors.New("message pusher address is not set") |  | ||||||
| 	} |  | ||||||
| 	req := request{ |  | ||||||
| 		Title:       title, |  | ||||||
| 		Description: description, |  | ||||||
| 		Content:     content, |  | ||||||
| 		Token:       config.MessagePusherToken, |  | ||||||
| 	} |  | ||||||
| 	data, err := json.Marshal(req) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return err |  | ||||||
| 	} |  | ||||||
| 	resp, err := http.Post(config.MessagePusherAddress, |  | ||||||
| 		"application/json", bytes.NewBuffer(data)) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return err |  | ||||||
| 	} |  | ||||||
| 	var res response |  | ||||||
| 	err = json.NewDecoder(resp.Body).Decode(&res) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return err |  | ||||||
| 	} |  | ||||||
| 	if !res.Success { |  | ||||||
| 		return errors.New(res.Message) |  | ||||||
| 	} |  | ||||||
| 	return nil |  | ||||||
| } |  | ||||||
| @@ -2,89 +2,91 @@ package common | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| const ( | var DalleSizeRatios = map[string]map[string]float64{ | ||||||
| 	USD2RMB = 7 | 	"dall-e-2": { | ||||||
| 	USD     = 500 // $0.002 = 1 -> $1 = 500 | 		"256x256":   1, | ||||||
| 	RMB     = USD / USD2RMB | 		"512x512":   1.125, | ||||||
| ) | 		"1024x1024": 1.25, | ||||||
|  | 	}, | ||||||
|  | 	"dall-e-3": { | ||||||
|  | 		"1024x1024": 1, | ||||||
|  | 		"1024x1792": 2, | ||||||
|  | 		"1792x1024": 2, | ||||||
|  | 	}, | ||||||
|  | } | ||||||
|  |  | ||||||
|  | var DalleGenerationImageAmounts = map[string][2]int{ | ||||||
|  | 	"dall-e-2": {1, 10}, | ||||||
|  | 	"dall-e-3": {1, 1}, // OpenAI allows n=1 currently. | ||||||
|  | } | ||||||
|  |  | ||||||
|  | var DalleImagePromptLengthLimitations = map[string]int{ | ||||||
|  | 	"dall-e-2": 1000, | ||||||
|  | 	"dall-e-3": 4000, | ||||||
|  | } | ||||||
|  |  | ||||||
| // ModelRatio | // ModelRatio | ||||||
| // https://platform.openai.com/docs/models/model-endpoint-compatibility | // https://platform.openai.com/docs/models/model-endpoint-compatibility | ||||||
| // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/Blfmc9dlf | // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/Blfmc9dlf | ||||||
| // https://openai.com/pricing | // https://openai.com/pricing | ||||||
|  | // TODO: when a new api is enabled, check the pricing here | ||||||
| // 1 === $0.002 / 1K tokens | // 1 === $0.002 / 1K tokens | ||||||
| // 1 === ¥0.014 / 1k tokens | // 1 === ¥0.014 / 1k tokens | ||||||
| var ModelRatio = map[string]float64{ | var ModelRatio = map[string]float64{ | ||||||
| 	// https://openai.com/pricing | 	"gpt-4":                     15, | ||||||
| 	"gpt-4":                   15, | 	"gpt-4-0314":                15, | ||||||
| 	"gpt-4-0314":              15, | 	"gpt-4-0613":                15, | ||||||
| 	"gpt-4-0613":              15, | 	"gpt-4-32k":                 30, | ||||||
| 	"gpt-4-32k":               30, | 	"gpt-4-32k-0314":            30, | ||||||
| 	"gpt-4-32k-0314":          30, | 	"gpt-4-32k-0613":            30, | ||||||
| 	"gpt-4-32k-0613":          30, | 	"gpt-4-1106-preview":        5,    // $0.01 / 1K tokens | ||||||
| 	"gpt-4-1106-preview":      5,    // $0.01 / 1K tokens | 	"gpt-4-vision-preview":      5,    // $0.01 / 1K tokens | ||||||
| 	"gpt-4-0125-preview":      5,    // $0.01 / 1K tokens | 	"gpt-3.5-turbo":             0.75, // $0.0015 / 1K tokens | ||||||
| 	"gpt-4-turbo-preview":     5,    // $0.01 / 1K tokens | 	"gpt-3.5-turbo-0301":        0.75, | ||||||
| 	"gpt-4-vision-preview":    5,    // $0.01 / 1K tokens | 	"gpt-3.5-turbo-0613":        0.75, | ||||||
| 	"gpt-3.5-turbo":           0.75, // $0.0015 / 1K tokens | 	"gpt-3.5-turbo-16k":         1.5, // $0.003 / 1K tokens | ||||||
| 	"gpt-3.5-turbo-0301":      0.75, | 	"gpt-3.5-turbo-16k-0613":    1.5, | ||||||
| 	"gpt-3.5-turbo-0613":      0.75, | 	"gpt-3.5-turbo-instruct":    0.75, // $0.0015 / 1K tokens | ||||||
| 	"gpt-3.5-turbo-16k":       1.5, // $0.003 / 1K tokens | 	"gpt-3.5-turbo-1106":        0.5,  // $0.001 / 1K tokens | ||||||
| 	"gpt-3.5-turbo-16k-0613":  1.5, | 	"davinci-002":               1,    // $0.002 / 1K tokens | ||||||
| 	"gpt-3.5-turbo-instruct":  0.75, // $0.0015 / 1K tokens | 	"babbage-002":               0.2,  // $0.0004 / 1K tokens | ||||||
| 	"gpt-3.5-turbo-1106":      0.5,  // $0.001 / 1K tokens | 	"text-ada-001":              0.2, | ||||||
| 	"gpt-3.5-turbo-0125":      0.25, // $0.0005 / 1K tokens | 	"text-babbage-001":          0.25, | ||||||
| 	"davinci-002":             1,    // $0.002 / 1K tokens | 	"text-curie-001":            1, | ||||||
| 	"babbage-002":             0.2,  // $0.0004 / 1K tokens | 	"text-davinci-002":          10, | ||||||
| 	"text-ada-001":            0.2, | 	"text-davinci-003":          10, | ||||||
| 	"text-babbage-001":        0.25, | 	"text-davinci-edit-001":     10, | ||||||
| 	"text-curie-001":          1, | 	"code-davinci-edit-001":     10, | ||||||
| 	"text-davinci-002":        10, | 	"whisper-1":                 15,  // $0.006 / minute -> $0.006 / 150 words -> $0.006 / 200 tokens -> $0.03 / 1k tokens | ||||||
| 	"text-davinci-003":        10, | 	"tts-1":                     7.5, // $0.015 / 1K characters | ||||||
| 	"text-davinci-edit-001":   10, | 	"tts-1-1106":                7.5, | ||||||
| 	"code-davinci-edit-001":   10, | 	"tts-1-hd":                  15, // $0.030 / 1K characters | ||||||
| 	"whisper-1":               15,  // $0.006 / minute -> $0.006 / 150 words -> $0.006 / 200 tokens -> $0.03 / 1k tokens | 	"tts-1-hd-1106":             15, | ||||||
| 	"tts-1":                   7.5, // $0.015 / 1K characters | 	"davinci":                   10, | ||||||
| 	"tts-1-1106":              7.5, | 	"curie":                     10, | ||||||
| 	"tts-1-hd":                15, // $0.030 / 1K characters | 	"babbage":                   10, | ||||||
| 	"tts-1-hd-1106":           15, | 	"ada":                       10, | ||||||
| 	"davinci":                 10, | 	"text-embedding-ada-002":    0.05, | ||||||
| 	"curie":                   10, | 	"text-search-ada-doc-001":   10, | ||||||
| 	"babbage":                 10, | 	"text-moderation-stable":    0.1, | ||||||
| 	"ada":                     10, | 	"text-moderation-latest":    0.1, | ||||||
| 	"text-embedding-ada-002":  0.05, | 	"dall-e-2":                  8,      // $0.016 - $0.020 / image | ||||||
| 	"text-embedding-3-small":  0.01, | 	"dall-e-3":                  20,     // $0.040 - $0.120 / image | ||||||
| 	"text-embedding-3-large":  0.065, | 	"claude-instant-1":          0.815,  // $1.63 / 1M tokens | ||||||
| 	"text-search-ada-doc-001": 10, | 	"claude-2":                  5.51,   // $11.02 / 1M tokens | ||||||
| 	"text-moderation-stable":  0.1, | 	"claude-2.0":                5.51,   // $11.02 / 1M tokens | ||||||
| 	"text-moderation-latest":  0.1, | 	"claude-2.1":                5.51,   // $11.02 / 1M tokens | ||||||
| 	"dall-e-2":                8,  // $0.016 - $0.020 / image | 	"ERNIE-Bot":                 0.8572, // ¥0.012 / 1k tokens | ||||||
| 	"dall-e-3":                20, // $0.040 - $0.120 / image | 	"ERNIE-Bot-turbo":           0.5715, // ¥0.008 / 1k tokens | ||||||
| 	// https://www.anthropic.com/api#pricing | 	"ERNIE-Bot-4":               8.572,  // ¥0.12 / 1k tokens | ||||||
| 	"claude-instant-1.2":       0.8 / 1000 * USD, | 	"Embedding-V1":              0.1429, // ¥0.002 / 1k tokens | ||||||
| 	"claude-2.0":               8.0 / 1000 * USD, | 	"PaLM-2":                    1, | ||||||
| 	"claude-2.1":               8.0 / 1000 * USD, | 	"gemini-pro":                1,      // $0.00025 / 1k characters -> $0.001 / 1k tokens | ||||||
| 	"claude-3-haiku-20240229":  0.25 / 1000 * USD, | 	"gemini-pro-vision":         1,      // $0.00025 / 1k characters -> $0.001 / 1k tokens | ||||||
| 	"claude-3-sonnet-20240229": 3.0 / 1000 * USD, |  | ||||||
| 	"claude-3-opus-20240229":   15.0 / 1000 * USD, |  | ||||||
| 	// https://cloud.baidu.com/doc/WENXINWORKSHOP/s/hlrk4akp7 |  | ||||||
| 	"ERNIE-Bot":         0.8572,     // ¥0.012 / 1k tokens |  | ||||||
| 	"ERNIE-Bot-turbo":   0.5715,     // ¥0.008 / 1k tokens |  | ||||||
| 	"ERNIE-Bot-4":       0.12 * RMB, // ¥0.12 / 1k tokens |  | ||||||
| 	"ERNIE-Bot-8k":      0.024 * RMB, |  | ||||||
| 	"Embedding-V1":      0.1429, // ¥0.002 / 1k tokens |  | ||||||
| 	"PaLM-2":            1, |  | ||||||
| 	"gemini-pro":        1, // $0.00025 / 1k characters -> $0.001 / 1k tokens |  | ||||||
| 	"gemini-pro-vision": 1, // $0.00025 / 1k characters -> $0.001 / 1k tokens |  | ||||||
| 	// https://open.bigmodel.cn/pricing |  | ||||||
| 	"glm-4":                     0.1 * RMB, |  | ||||||
| 	"glm-4v":                    0.1 * RMB, |  | ||||||
| 	"glm-3-turbo":               0.005 * RMB, |  | ||||||
| 	"chatglm_turbo":             0.3572, // ¥0.005 / 1k tokens | 	"chatglm_turbo":             0.3572, // ¥0.005 / 1k tokens | ||||||
| 	"chatglm_pro":               0.7143, // ¥0.01 / 1k tokens | 	"chatglm_pro":               0.7143, // ¥0.01 / 1k tokens | ||||||
| 	"chatglm_std":               0.3572, // ¥0.005 / 1k tokens | 	"chatglm_std":               0.3572, // ¥0.005 / 1k tokens | ||||||
| @@ -95,63 +97,17 @@ var ModelRatio = map[string]float64{ | |||||||
| 	"qwen-max-longcontext":      1.4286, // ¥0.02 / 1k tokens | 	"qwen-max-longcontext":      1.4286, // ¥0.02 / 1k tokens | ||||||
| 	"text-embedding-v1":         0.05,   // ¥0.0007 / 1k tokens | 	"text-embedding-v1":         0.05,   // ¥0.0007 / 1k tokens | ||||||
| 	"SparkDesk":                 1.2858, // ¥0.018 / 1k tokens | 	"SparkDesk":                 1.2858, // ¥0.018 / 1k tokens | ||||||
| 	"SparkDesk-v1.1":            1.2858, // ¥0.018 / 1k tokens |  | ||||||
| 	"SparkDesk-v2.1":            1.2858, // ¥0.018 / 1k tokens |  | ||||||
| 	"SparkDesk-v3.1":            1.2858, // ¥0.018 / 1k tokens |  | ||||||
| 	"SparkDesk-v3.5":            1.2858, // ¥0.018 / 1k tokens |  | ||||||
| 	"360GPT_S2_V9":              0.8572, // ¥0.012 / 1k tokens | 	"360GPT_S2_V9":              0.8572, // ¥0.012 / 1k tokens | ||||||
| 	"embedding-bert-512-v1":     0.0715, // ¥0.001 / 1k tokens | 	"embedding-bert-512-v1":     0.0715, // ¥0.001 / 1k tokens | ||||||
| 	"embedding_s1_v1":           0.0715, // ¥0.001 / 1k tokens | 	"embedding_s1_v1":           0.0715, // ¥0.001 / 1k tokens | ||||||
| 	"semantic_similarity_s1_v1": 0.0715, // ¥0.001 / 1k tokens | 	"semantic_similarity_s1_v1": 0.0715, // ¥0.001 / 1k tokens | ||||||
| 	"hunyuan":                   7.143,  // ¥0.1 / 1k tokens  // https://cloud.tencent.com/document/product/1729/97731#e0e6be58-60c8-469f-bdeb-6c264ce3b4d0 | 	"hunyuan":                   7.143,  // ¥0.1 / 1k tokens  // https://cloud.tencent.com/document/product/1729/97731#e0e6be58-60c8-469f-bdeb-6c264ce3b4d0 | ||||||
| 	"ChatStd":                   0.01 * RMB, |  | ||||||
| 	"ChatPro":                   0.1 * RMB, |  | ||||||
| 	// https://platform.moonshot.cn/pricing |  | ||||||
| 	"moonshot-v1-8k":   0.012 * RMB, |  | ||||||
| 	"moonshot-v1-32k":  0.024 * RMB, |  | ||||||
| 	"moonshot-v1-128k": 0.06 * RMB, |  | ||||||
| 	// https://platform.baichuan-ai.com/price |  | ||||||
| 	"Baichuan2-Turbo":      0.008 * RMB, |  | ||||||
| 	"Baichuan2-Turbo-192k": 0.016 * RMB, |  | ||||||
| 	"Baichuan2-53B":        0.02 * RMB, |  | ||||||
| 	// https://api.minimax.chat/document/price |  | ||||||
| 	"abab6-chat":    0.1 * RMB, |  | ||||||
| 	"abab5.5-chat":  0.015 * RMB, |  | ||||||
| 	"abab5.5s-chat": 0.005 * RMB, |  | ||||||
| 	// https://docs.mistral.ai/platform/pricing/ |  | ||||||
| 	"open-mistral-7b":       0.25 / 1000 * USD, |  | ||||||
| 	"open-mixtral-8x7b":     0.7 / 1000 * USD, |  | ||||||
| 	"mistral-small-latest":  2.0 / 1000 * USD, |  | ||||||
| 	"mistral-medium-latest": 2.7 / 1000 * USD, |  | ||||||
| 	"mistral-large-latest":  8.0 / 1000 * USD, |  | ||||||
| 	"mistral-embed":         0.1 / 1000 * USD, |  | ||||||
| 	// https://wow.groq.com/ |  | ||||||
| 	"llama2-70b-4096":    0.7 / 1000 * USD, |  | ||||||
| 	"llama2-7b-2048":     0.1 / 1000 * USD, |  | ||||||
| 	"mixtral-8x7b-32768": 0.27 / 1000 * USD, |  | ||||||
| 	"gemma-7b-it":        0.1 / 1000 * USD, |  | ||||||
| } |  | ||||||
|  |  | ||||||
| var CompletionRatio = map[string]float64{} |  | ||||||
|  |  | ||||||
| var DefaultModelRatio map[string]float64 |  | ||||||
| var DefaultCompletionRatio map[string]float64 |  | ||||||
|  |  | ||||||
| func init() { |  | ||||||
| 	DefaultModelRatio = make(map[string]float64) |  | ||||||
| 	for k, v := range ModelRatio { |  | ||||||
| 		DefaultModelRatio[k] = v |  | ||||||
| 	} |  | ||||||
| 	DefaultCompletionRatio = make(map[string]float64) |  | ||||||
| 	for k, v := range CompletionRatio { |  | ||||||
| 		DefaultCompletionRatio[k] = v |  | ||||||
| 	} |  | ||||||
| } | } | ||||||
|  |  | ||||||
| func ModelRatio2JSONString() string { | func ModelRatio2JSONString() string { | ||||||
| 	jsonBytes, err := json.Marshal(ModelRatio) | 	jsonBytes, err := json.Marshal(ModelRatio) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.SysError("error marshalling model ratio: " + err.Error()) | 		SysError("error marshalling model ratio: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return string(jsonBytes) | 	return string(jsonBytes) | ||||||
| } | } | ||||||
| @@ -167,41 +123,14 @@ func GetModelRatio(name string) float64 { | |||||||
| 	} | 	} | ||||||
| 	ratio, ok := ModelRatio[name] | 	ratio, ok := ModelRatio[name] | ||||||
| 	if !ok { | 	if !ok { | ||||||
| 		ratio, ok = DefaultModelRatio[name] | 		SysError("model ratio not found: " + name) | ||||||
| 	} |  | ||||||
| 	if !ok { |  | ||||||
| 		logger.SysError("model ratio not found: " + name) |  | ||||||
| 		return 30 | 		return 30 | ||||||
| 	} | 	} | ||||||
| 	return ratio | 	return ratio | ||||||
| } | } | ||||||
|  |  | ||||||
| func CompletionRatio2JSONString() string { |  | ||||||
| 	jsonBytes, err := json.Marshal(CompletionRatio) |  | ||||||
| 	if err != nil { |  | ||||||
| 		logger.SysError("error marshalling completion ratio: " + err.Error()) |  | ||||||
| 	} |  | ||||||
| 	return string(jsonBytes) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func UpdateCompletionRatioByJSONString(jsonStr string) error { |  | ||||||
| 	CompletionRatio = make(map[string]float64) |  | ||||||
| 	return json.Unmarshal([]byte(jsonStr), &CompletionRatio) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func GetCompletionRatio(name string) float64 { | func GetCompletionRatio(name string) float64 { | ||||||
| 	if ratio, ok := CompletionRatio[name]; ok { |  | ||||||
| 		return ratio |  | ||||||
| 	} |  | ||||||
| 	if ratio, ok := DefaultCompletionRatio[name]; ok { |  | ||||||
| 		return ratio |  | ||||||
| 	} |  | ||||||
| 	if strings.HasPrefix(name, "gpt-3.5") { | 	if strings.HasPrefix(name, "gpt-3.5") { | ||||||
| 		if strings.HasSuffix(name, "0125") { |  | ||||||
| 			// https://openai.com/blog/new-embedding-models-and-api-updates |  | ||||||
| 			// Updated GPT-3.5 Turbo model and lower pricing |  | ||||||
| 			return 3 |  | ||||||
| 		} |  | ||||||
| 		if strings.HasSuffix(name, "1106") { | 		if strings.HasSuffix(name, "1106") { | ||||||
| 			return 2 | 			return 2 | ||||||
| 		} | 		} | ||||||
| @@ -214,7 +143,7 @@ func GetCompletionRatio(name string) float64 { | |||||||
| 				return 2 | 				return 2 | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 		return 4.0 / 3.0 | 		return 1.333333 | ||||||
| 	} | 	} | ||||||
| 	if strings.HasPrefix(name, "gpt-4") { | 	if strings.HasPrefix(name, "gpt-4") { | ||||||
| 		if strings.HasSuffix(name, "preview") { | 		if strings.HasSuffix(name, "preview") { | ||||||
| @@ -222,18 +151,11 @@ func GetCompletionRatio(name string) float64 { | |||||||
| 		} | 		} | ||||||
| 		return 2 | 		return 2 | ||||||
| 	} | 	} | ||||||
| 	if strings.HasPrefix(name, "claude-3") { | 	if strings.HasPrefix(name, "claude-instant-1") { | ||||||
| 		return 5 | 		return 3.38 | ||||||
| 	} | 	} | ||||||
| 	if strings.HasPrefix(name, "claude-") { | 	if strings.HasPrefix(name, "claude-2") { | ||||||
| 		return 3 | 		return 2.965517 | ||||||
| 	} |  | ||||||
| 	if strings.HasPrefix(name, "mistral-") { |  | ||||||
| 		return 3 |  | ||||||
| 	} |  | ||||||
| 	switch name { |  | ||||||
| 	case "llama2-70b-4096": |  | ||||||
| 		return 0.8 / 0.7 |  | ||||||
| 	} | 	} | ||||||
| 	return 1 | 	return 1 | ||||||
| } | } | ||||||
|   | |||||||
| @@ -1,8 +0,0 @@ | |||||||
| package common |  | ||||||
|  |  | ||||||
| import "math/rand" |  | ||||||
|  |  | ||||||
| // RandRange returns a random number between min and max (max is not included) |  | ||||||
| func RandRange(min, max int) int { |  | ||||||
| 	return min + rand.Intn(max-min) |  | ||||||
| } |  | ||||||
| @@ -3,7 +3,6 @@ package common | |||||||
| import ( | import ( | ||||||
| 	"context" | 	"context" | ||||||
| 	"github.com/go-redis/redis/v8" | 	"github.com/go-redis/redis/v8" | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"os" | 	"os" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
| @@ -15,18 +14,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 | ||||||
| 		logger.SysLog("REDIS_CONN_STRING not set, Redis is not enabled") | 		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 | ||||||
| 		logger.SysLog("SYNC_FREQUENCY not set, Redis is disabled") | 		SysLog("SYNC_FREQUENCY not set, Redis is disabled") | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| 	logger.SysLog("Redis is enabled") | 	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 { | ||||||
| 		logger.FatalLog("failed to parse Redis connection string: " + err.Error()) | 		FatalLog("failed to parse Redis connection string: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	RDB = redis.NewClient(opt) | 	RDB = redis.NewClient(opt) | ||||||
|  |  | ||||||
| @@ -35,7 +34,7 @@ func InitRedisClient() (err error) { | |||||||
|  |  | ||||||
| 	_, err = RDB.Ping(ctx).Result() | 	_, err = RDB.Ping(ctx).Result() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.FatalLog("Redis ping test failed: " + err.Error()) | 		FatalLog("Redis ping test failed: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
| @@ -43,7 +42,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 { | ||||||
| 		logger.FatalLog("failed to parse Redis connection string: " + err.Error()) | 		FatalLog("failed to parse Redis connection string: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return opt | 	return opt | ||||||
| } | } | ||||||
|   | |||||||
							
								
								
									
										214
									
								
								common/utils.go
									
									
									
									
									
								
							
							
						
						
									
										214
									
								
								common/utils.go
									
									
									
									
									
								
							| @@ -2,13 +2,215 @@ package common | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" | 	"github.com/google/uuid" | ||||||
|  | 	"html/template" | ||||||
|  | 	"log" | ||||||
|  | 	"math/rand" | ||||||
|  | 	"net" | ||||||
|  | 	"os" | ||||||
|  | 	"os/exec" | ||||||
|  | 	"runtime" | ||||||
|  | 	"strconv" | ||||||
|  | 	"strings" | ||||||
|  | 	"time" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func LogQuota(quota int) string { | func OpenBrowser(url string) { | ||||||
| 	if config.DisplayInCurrencyEnabled { | 	var err error | ||||||
| 		return fmt.Sprintf("$%.6f 额度", float64(quota)/config.QuotaPerUnit) |  | ||||||
| 	} else { | 	switch runtime.GOOS { | ||||||
| 		return fmt.Sprintf("%d 点额度", quota) | 	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 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,9 +2,8 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" | 	"one-api/common" | ||||||
| 	"github.com/songquanpeng/one-api/model" | 	"one-api/model" | ||||||
| 	relaymodel "github.com/songquanpeng/one-api/relay/model" |  | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func GetSubscription(c *gin.Context) { | func GetSubscription(c *gin.Context) { | ||||||
| @@ -13,7 +12,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 config.DisplayTokenStatEnabled { | 	if common.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 | ||||||
| @@ -22,27 +21,25 @@ func GetSubscription(c *gin.Context) { | |||||||
| 	} else { | 	} else { | ||||||
| 		userId := c.GetInt("id") | 		userId := c.GetInt("id") | ||||||
| 		remainQuota, err = model.GetUserQuota(userId) | 		remainQuota, err = model.GetUserQuota(userId) | ||||||
| 		if err != nil { | 		usedQuota, err = model.GetUserUsedQuota(userId) | ||||||
| 			usedQuota, err = model.GetUserUsedQuota(userId) |  | ||||||
| 		} |  | ||||||
| 	} | 	} | ||||||
| 	if expiredTime <= 0 { | 	if expiredTime <= 0 { | ||||||
| 		expiredTime = 0 | 		expiredTime = 0 | ||||||
| 	} | 	} | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		Error := relaymodel.Error{ | 		openAIError := OpenAIError{ | ||||||
| 			Message: err.Error(), | 			Message: err.Error(), | ||||||
| 			Type:    "upstream_error", | 			Type:    "upstream_error", | ||||||
| 		} | 		} | ||||||
| 		c.JSON(200, gin.H{ | 		c.JSON(200, gin.H{ | ||||||
| 			"error": Error, | 			"error": openAIError, | ||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	quota := remainQuota + usedQuota | 	quota := remainQuota + usedQuota | ||||||
| 	amount := float64(quota) | 	amount := float64(quota) | ||||||
| 	if config.DisplayInCurrencyEnabled { | 	if common.DisplayInCurrencyEnabled { | ||||||
| 		amount /= config.QuotaPerUnit | 		amount /= common.QuotaPerUnit | ||||||
| 	} | 	} | ||||||
| 	if token != nil && token.UnlimitedQuota { | 	if token != nil && token.UnlimitedQuota { | ||||||
| 		amount = 100000000 | 		amount = 100000000 | ||||||
| @@ -63,7 +60,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 config.DisplayTokenStatEnabled { | 	if common.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 | ||||||
| @@ -72,18 +69,18 @@ func GetUsage(c *gin.Context) { | |||||||
| 		quota, err = model.GetUserUsedQuota(userId) | 		quota, err = model.GetUserUsedQuota(userId) | ||||||
| 	} | 	} | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		Error := relaymodel.Error{ | 		openAIError := OpenAIError{ | ||||||
| 			Message: err.Error(), | 			Message: err.Error(), | ||||||
| 			Type:    "one_api_error", | 			Type:    "one_api_error", | ||||||
| 		} | 		} | ||||||
| 		c.JSON(200, gin.H{ | 		c.JSON(200, gin.H{ | ||||||
| 			"error": Error, | 			"error": openAIError, | ||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	amount := float64(quota) | 	amount := float64(quota) | ||||||
| 	if config.DisplayInCurrencyEnabled { | 	if common.DisplayInCurrencyEnabled { | ||||||
| 		amount /= config.QuotaPerUnit | 		amount /= common.QuotaPerUnit | ||||||
| 	} | 	} | ||||||
| 	usage := OpenAIUsageResponse{ | 	usage := OpenAIUsageResponse{ | ||||||
| 		Object:     "list", | 		Object:     "list", | ||||||
|   | |||||||
| @@ -4,14 +4,10 @@ import ( | |||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/monitor" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"time" | 	"time" | ||||||
|  |  | ||||||
| @@ -96,7 +92,7 @@ func GetResponseBody(method, url string, channel *model.Channel, headers http.He | |||||||
| 	for k := range headers { | 	for k := range headers { | ||||||
| 		req.Header.Add(k, headers.Get(k)) | 		req.Header.Add(k, headers.Get(k)) | ||||||
| 	} | 	} | ||||||
| 	res, err := util.HTTPClient.Do(req) | 	res, err := httpClient.Do(req) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return nil, err | 		return nil, err | ||||||
| 	} | 	} | ||||||
| @@ -296,7 +292,7 @@ func UpdateChannelBalance(c *gin.Context) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func updateAllChannelsBalance() error { | func updateAllChannelsBalance() error { | ||||||
| 	channels, err := model.GetAllChannels(0, 0, "all") | 	channels, err := model.GetAllChannels(0, 0, true) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return err | 		return err | ||||||
| 	} | 	} | ||||||
| @@ -314,10 +310,10 @@ func updateAllChannelsBalance() error { | |||||||
| 		} else { | 		} else { | ||||||
| 			// err is nil & balance <= 0 means quota is used up | 			// err is nil & balance <= 0 means quota is used up | ||||||
| 			if balance <= 0 { | 			if balance <= 0 { | ||||||
| 				monitor.DisableChannel(channel.Id, channel.Name, "余额不足") | 				disableChannel(channel.Id, channel.Name, "余额不足") | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 		time.Sleep(config.RequestInterval) | 		time.Sleep(common.RequestInterval) | ||||||
| 	} | 	} | ||||||
| 	return nil | 	return nil | ||||||
| } | } | ||||||
| @@ -342,8 +338,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) | ||||||
| 		logger.SysLog("updating all channels") | 		common.SysLog("updating all channels") | ||||||
| 		_ = updateAllChannelsBalance() | 		_ = updateAllChannelsBalance() | ||||||
| 		logger.SysLog("channels update done") | 		common.SysLog("channels update done") | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -5,36 +5,98 @@ import ( | |||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/message" |  | ||||||
| 	"github.com/songquanpeng/one-api/middleware" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/monitor" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/helper" |  | ||||||
| 	relaymodel "github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"net/http/httptest" | 	"one-api/common" | ||||||
| 	"net/url" | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" |  | ||||||
| 	"sync" | 	"sync" | ||||||
| 	"time" | 	"time" | ||||||
|  |  | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func buildTestRequest() *relaymodel.GeneralOpenAIRequest { | func testChannel(channel *model.Channel, request ChatRequest) (err error, openaiErr *OpenAIError) { | ||||||
| 	testRequest := &relaymodel.GeneralOpenAIRequest{ | 	switch channel.Type { | ||||||
| 		MaxTokens: 1, | 	case common.ChannelTypePaLM: | ||||||
| 		Stream:    false, | 		fallthrough | ||||||
| 		Model:     "gpt-3.5-turbo", | 	case common.ChannelTypeGemini: | ||||||
|  | 		fallthrough | ||||||
|  | 	case common.ChannelTypeAnthropic: | ||||||
|  | 		fallthrough | ||||||
|  | 	case common.ChannelTypeBaidu: | ||||||
|  | 		fallthrough | ||||||
|  | 	case common.ChannelTypeZhipu: | ||||||
|  | 		fallthrough | ||||||
|  | 	case common.ChannelTypeAli: | ||||||
|  | 		fallthrough | ||||||
|  | 	case common.ChannelType360: | ||||||
|  | 		fallthrough | ||||||
|  | 	case common.ChannelTypeXunfei: | ||||||
|  | 		return errors.New("该渠道类型当前版本不支持测试,请手动测试"), nil | ||||||
|  | 	case common.ChannelTypeAzure: | ||||||
|  | 		request.Model = "gpt-35-turbo" | ||||||
|  | 		defer func() { | ||||||
|  | 			if err != nil { | ||||||
|  | 				err = errors.New("请确保已在 Azure 上创建了 gpt-35-turbo 模型,并且 apiVersion 已正确填写!") | ||||||
|  | 			} | ||||||
|  | 		}() | ||||||
|  | 	default: | ||||||
|  | 		request.Model = "gpt-3.5-turbo" | ||||||
| 	} | 	} | ||||||
| 	testMessage := relaymodel.Message{ | 	requestURL := common.ChannelBaseURLs[channel.Type] | ||||||
|  | 	if channel.Type == common.ChannelTypeAzure { | ||||||
|  | 		requestURL = getFullRequestURL(channel.GetBaseURL(), fmt.Sprintf("/openai/deployments/%s/chat/completions?api-version=2023-03-15-preview", request.Model), channel.Type) | ||||||
|  | 	} else { | ||||||
|  | 		if baseURL := channel.GetBaseURL(); len(baseURL) > 0 { | ||||||
|  | 			requestURL = baseURL | ||||||
|  | 		} | ||||||
|  |  | ||||||
|  | 		requestURL = getFullRequestURL(requestURL, "/v1/chat/completions", channel.Type) | ||||||
|  | 	} | ||||||
|  | 	jsonData, err := json.Marshal(request) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err, nil | ||||||
|  | 	} | ||||||
|  | 	req, err := http.NewRequest("POST", requestURL, bytes.NewBuffer(jsonData)) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err, nil | ||||||
|  | 	} | ||||||
|  | 	if channel.Type == common.ChannelTypeAzure { | ||||||
|  | 		req.Header.Set("api-key", channel.Key) | ||||||
|  | 	} else { | ||||||
|  | 		req.Header.Set("Authorization", "Bearer "+channel.Key) | ||||||
|  | 	} | ||||||
|  | 	req.Header.Set("Content-Type", "application/json") | ||||||
|  | 	resp, err := httpClient.Do(req) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err, nil | ||||||
|  | 	} | ||||||
|  | 	defer resp.Body.Close() | ||||||
|  | 	var response TextResponse | ||||||
|  | 	body, err := io.ReadAll(resp.Body) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return err, nil | ||||||
|  | 	} | ||||||
|  | 	err = json.Unmarshal(body, &response) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return fmt.Errorf("Error: %s\nResp body: %s", err, body), nil | ||||||
|  | 	} | ||||||
|  | 	if response.Usage.CompletionTokens == 0 { | ||||||
|  | 		if response.Error.Message == "" { | ||||||
|  | 			response.Error.Message = "补全 tokens 非预期返回 0" | ||||||
|  | 		} | ||||||
|  | 		return errors.New(fmt.Sprintf("type %s, code %v, message %s", response.Error.Type, response.Error.Code, response.Error.Message)), &response.Error | ||||||
|  | 	} | ||||||
|  | 	return nil, nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func buildTestRequest() *ChatRequest { | ||||||
|  | 	testRequest := &ChatRequest{ | ||||||
|  | 		Model:     "", // this will be set later | ||||||
|  | 		MaxTokens: 1, | ||||||
|  | 	} | ||||||
|  | 	testMessage := Message{ | ||||||
| 		Role:    "user", | 		Role:    "user", | ||||||
| 		Content: "hi", | 		Content: "hi", | ||||||
| 	} | 	} | ||||||
| @@ -42,72 +104,6 @@ func buildTestRequest() *relaymodel.GeneralOpenAIRequest { | |||||||
| 	return testRequest | 	return testRequest | ||||||
| } | } | ||||||
|  |  | ||||||
| func testChannel(channel *model.Channel) (err error, openaiErr *relaymodel.Error) { |  | ||||||
| 	w := httptest.NewRecorder() |  | ||||||
| 	c, _ := gin.CreateTestContext(w) |  | ||||||
| 	c.Request = &http.Request{ |  | ||||||
| 		Method: "POST", |  | ||||||
| 		URL:    &url.URL{Path: "/v1/chat/completions"}, |  | ||||||
| 		Body:   nil, |  | ||||||
| 		Header: make(http.Header), |  | ||||||
| 	} |  | ||||||
| 	c.Request.Header.Set("Authorization", "Bearer "+channel.Key) |  | ||||||
| 	c.Request.Header.Set("Content-Type", "application/json") |  | ||||||
| 	c.Set("channel", channel.Type) |  | ||||||
| 	c.Set("base_url", channel.GetBaseURL()) |  | ||||||
| 	middleware.SetupContextForSelectedChannel(c, channel, "") |  | ||||||
| 	meta := util.GetRelayMeta(c) |  | ||||||
| 	apiType := constant.ChannelType2APIType(channel.Type) |  | ||||||
| 	adaptor := helper.GetAdaptor(apiType) |  | ||||||
| 	if adaptor == nil { |  | ||||||
| 		return fmt.Errorf("invalid api type: %d, adaptor is nil", apiType), nil |  | ||||||
| 	} |  | ||||||
| 	adaptor.Init(meta) |  | ||||||
| 	modelName := adaptor.GetModelList()[0] |  | ||||||
| 	if !strings.Contains(channel.Models, modelName) { |  | ||||||
| 		modelNames := strings.Split(channel.Models, ",") |  | ||||||
| 		if len(modelNames) > 0 { |  | ||||||
| 			modelName = modelNames[0] |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	request := buildTestRequest() |  | ||||||
| 	request.Model = modelName |  | ||||||
| 	meta.OriginModelName, meta.ActualModelName = modelName, modelName |  | ||||||
| 	convertedRequest, err := adaptor.ConvertRequest(c, constant.RelayModeChatCompletions, request) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return err, nil |  | ||||||
| 	} |  | ||||||
| 	jsonData, err := json.Marshal(convertedRequest) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return err, nil |  | ||||||
| 	} |  | ||||||
| 	requestBody := bytes.NewBuffer(jsonData) |  | ||||||
| 	c.Request.Body = io.NopCloser(requestBody) |  | ||||||
| 	resp, err := adaptor.DoRequest(c, meta, requestBody) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return err, nil |  | ||||||
| 	} |  | ||||||
| 	if resp.StatusCode != http.StatusOK { |  | ||||||
| 		err := util.RelayErrorHandler(resp) |  | ||||||
| 		return fmt.Errorf("status code %d: %s", resp.StatusCode, err.Error.Message), &err.Error |  | ||||||
| 	} |  | ||||||
| 	usage, respErr := adaptor.DoResponse(c, resp, meta) |  | ||||||
| 	if respErr != nil { |  | ||||||
| 		return fmt.Errorf("%s", respErr.Error.Message), &respErr.Error |  | ||||||
| 	} |  | ||||||
| 	if usage == nil { |  | ||||||
| 		return errors.New("usage is nil"), nil |  | ||||||
| 	} |  | ||||||
| 	result := w.Result() |  | ||||||
| 	// print result.Body |  | ||||||
| 	respBody, err := io.ReadAll(result.Body) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return err, nil |  | ||||||
| 	} |  | ||||||
| 	logger.SysLog(fmt.Sprintf("testing channel #%d, response: \n%s", channel.Id, string(respBody))) |  | ||||||
| 	return nil, nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func TestChannel(c *gin.Context) { | func TestChannel(c *gin.Context) { | ||||||
| 	id, err := strconv.Atoi(c.Param("id")) | 	id, err := strconv.Atoi(c.Param("id")) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| @@ -125,8 +121,9 @@ func TestChannel(c *gin.Context) { | |||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
|  | 	testRequest := buildTestRequest() | ||||||
| 	tik := time.Now() | 	tik := time.Now() | ||||||
| 	err, _ = testChannel(channel) | 	err, _ = testChannel(channel, *testRequest) | ||||||
| 	tok := time.Now() | 	tok := time.Now() | ||||||
| 	milliseconds := tok.Sub(tik).Milliseconds() | 	milliseconds := tok.Sub(tik).Milliseconds() | ||||||
| 	go channel.UpdateResponseTime(milliseconds) | 	go channel.UpdateResponseTime(milliseconds) | ||||||
| @@ -150,9 +147,35 @@ func TestChannel(c *gin.Context) { | |||||||
| var testAllChannelsLock sync.Mutex | var testAllChannelsLock sync.Mutex | ||||||
| var testAllChannelsRunning bool = false | var testAllChannelsRunning bool = false | ||||||
|  |  | ||||||
| func testChannels(notify bool, scope string) error { | func notifyRootUser(subject string, content string) { | ||||||
| 	if config.RootUserEmail == "" { | 	if common.RootUserEmail == "" { | ||||||
| 		config.RootUserEmail = model.GetRootUserEmail() | 		common.RootUserEmail = model.GetRootUserEmail() | ||||||
|  | 	} | ||||||
|  | 	err := common.SendEmail(subject, common.RootUserEmail, content) | ||||||
|  | 	if err != nil { | ||||||
|  | 		common.SysError(fmt.Sprintf("failed to send email: %s", err.Error())) | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // disable & notify | ||||||
|  | func disableChannel(channelId int, channelName string, reason string) { | ||||||
|  | 	model.UpdateChannelStatusById(channelId, common.ChannelStatusAutoDisabled) | ||||||
|  | 	subject := fmt.Sprintf("通道「%s」(#%d)已被禁用", channelName, channelId) | ||||||
|  | 	content := fmt.Sprintf("通道「%s」(#%d)已被禁用,原因:%s", channelName, channelId, reason) | ||||||
|  | 	notifyRootUser(subject, content) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // enable & notify | ||||||
|  | func enableChannel(channelId int, channelName string) { | ||||||
|  | 	model.UpdateChannelStatusById(channelId, common.ChannelStatusEnabled) | ||||||
|  | 	subject := fmt.Sprintf("通道「%s」(#%d)已被启用", channelName, channelId) | ||||||
|  | 	content := fmt.Sprintf("通道「%s」(#%d)已被启用", channelName, channelId) | ||||||
|  | 	notifyRootUser(subject, content) | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func testAllChannels(notify bool) error { | ||||||
|  | 	if common.RootUserEmail == "" { | ||||||
|  | 		common.RootUserEmail = model.GetRootUserEmail() | ||||||
| 	} | 	} | ||||||
| 	testAllChannelsLock.Lock() | 	testAllChannelsLock.Lock() | ||||||
| 	if testAllChannelsRunning { | 	if testAllChannelsRunning { | ||||||
| @@ -161,11 +184,12 @@ func testChannels(notify bool, scope string) error { | |||||||
| 	} | 	} | ||||||
| 	testAllChannelsRunning = true | 	testAllChannelsRunning = true | ||||||
| 	testAllChannelsLock.Unlock() | 	testAllChannelsLock.Unlock() | ||||||
| 	channels, err := model.GetAllChannels(0, 0, scope) | 	channels, err := model.GetAllChannels(0, 0, true) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return err | 		return err | ||||||
| 	} | 	} | ||||||
| 	var disableThreshold = int64(config.ChannelDisableThreshold * 1000) | 	testRequest := buildTestRequest() | ||||||
|  | 	var disableThreshold = int64(common.ChannelDisableThreshold * 1000) | ||||||
| 	if disableThreshold == 0 { | 	if disableThreshold == 0 { | ||||||
| 		disableThreshold = 10000000 // a impossible value | 		disableThreshold = 10000000 // a impossible value | ||||||
| 	} | 	} | ||||||
| @@ -173,41 +197,37 @@ func testChannels(notify bool, scope string) error { | |||||||
| 		for _, channel := range channels { | 		for _, channel := range channels { | ||||||
| 			isChannelEnabled := channel.Status == common.ChannelStatusEnabled | 			isChannelEnabled := channel.Status == common.ChannelStatusEnabled | ||||||
| 			tik := time.Now() | 			tik := time.Now() | ||||||
| 			err, openaiErr := testChannel(channel) | 			err, openaiErr := testChannel(channel, *testRequest) | ||||||
| 			tok := time.Now() | 			tok := time.Now() | ||||||
| 			milliseconds := tok.Sub(tik).Milliseconds() | 			milliseconds := tok.Sub(tik).Milliseconds() | ||||||
| 			if isChannelEnabled && milliseconds > disableThreshold { | 			if isChannelEnabled && milliseconds > disableThreshold { | ||||||
| 				err = errors.New(fmt.Sprintf("响应时间 %.2fs 超过阈值 %.2fs", float64(milliseconds)/1000.0, float64(disableThreshold)/1000.0)) | 				err = errors.New(fmt.Sprintf("响应时间 %.2fs 超过阈值 %.2fs", float64(milliseconds)/1000.0, float64(disableThreshold)/1000.0)) | ||||||
| 				monitor.DisableChannel(channel.Id, channel.Name, err.Error()) | 				disableChannel(channel.Id, channel.Name, err.Error()) | ||||||
| 			} | 			} | ||||||
| 			if isChannelEnabled && util.ShouldDisableChannel(openaiErr, -1) { | 			if isChannelEnabled && shouldDisableChannel(openaiErr, -1) { | ||||||
| 				monitor.DisableChannel(channel.Id, channel.Name, err.Error()) | 				disableChannel(channel.Id, channel.Name, err.Error()) | ||||||
| 			} | 			} | ||||||
| 			if !isChannelEnabled && util.ShouldEnableChannel(err, openaiErr) { | 			if !isChannelEnabled && shouldEnableChannel(err, openaiErr) { | ||||||
| 				monitor.EnableChannel(channel.Id, channel.Name) | 				enableChannel(channel.Id, channel.Name) | ||||||
| 			} | 			} | ||||||
| 			channel.UpdateResponseTime(milliseconds) | 			channel.UpdateResponseTime(milliseconds) | ||||||
| 			time.Sleep(config.RequestInterval) | 			time.Sleep(common.RequestInterval) | ||||||
| 		} | 		} | ||||||
| 		testAllChannelsLock.Lock() | 		testAllChannelsLock.Lock() | ||||||
| 		testAllChannelsRunning = false | 		testAllChannelsRunning = false | ||||||
| 		testAllChannelsLock.Unlock() | 		testAllChannelsLock.Unlock() | ||||||
| 		if notify { | 		if notify { | ||||||
| 			err := message.Notify(message.ByAll, "通道测试完成", "", "通道测试完成,如果没有收到禁用通知,说明所有通道都正常") | 			err := common.SendEmail("通道测试完成", common.RootUserEmail, "通道测试完成,如果没有收到禁用通知,说明所有通道都正常") | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError(fmt.Sprintf("failed to send email: %s", err.Error())) | 				common.SysError(fmt.Sprintf("failed to send email: %s", err.Error())) | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 	}() | 	}() | ||||||
| 	return nil | 	return nil | ||||||
| } | } | ||||||
|  |  | ||||||
| func TestChannels(c *gin.Context) { | func TestAllChannels(c *gin.Context) { | ||||||
| 	scope := c.Query("scope") | 	err := testAllChannels(true) | ||||||
| 	if scope == "" { |  | ||||||
| 		scope = "all" |  | ||||||
| 	} |  | ||||||
| 	err := testChannels(true, scope) |  | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -225,8 +245,8 @@ func TestChannels(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) | ||||||
| 		logger.SysLog("testing all channels") | 		common.SysLog("testing all channels") | ||||||
| 		_ = testChannels(false, "all") | 		_ = testAllChannels(false) | ||||||
| 		logger.SysLog("channel test finished") | 		common.SysLog("channel test finished") | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -2,10 +2,9 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
| @@ -15,7 +14,7 @@ func GetAllChannels(c *gin.Context) { | |||||||
| 	if p < 0 { | 	if p < 0 { | ||||||
| 		p = 0 | 		p = 0 | ||||||
| 	} | 	} | ||||||
| 	channels, err := model.GetAllChannels(p*config.ItemsPerPage, config.ItemsPerPage, "limited") | 	channels, err := model.GetAllChannels(p*common.ItemsPerPage, common.ItemsPerPage, false) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -84,7 +83,7 @@ func AddChannel(c *gin.Context) { | |||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	channel.CreatedTime = helper.GetTimestamp() | 	channel.CreatedTime = common.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 { | ||||||
|   | |||||||
| @@ -7,12 +7,9 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-contrib/sessions" | 	"github.com/gin-contrib/sessions" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
| @@ -33,7 +30,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": config.GitHubClientId, "client_secret": config.GitHubClientSecret, "code": code} | 	values := map[string]string{"client_id": common.GitHubClientId, "client_secret": common.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 | ||||||
| @@ -49,7 +46,7 @@ func getGitHubUserInfoByCode(code string) (*GitHubUser, error) { | |||||||
| 	} | 	} | ||||||
| 	res, err := client.Do(req) | 	res, err := client.Do(req) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.SysLog(err.Error()) | 		common.SysLog(err.Error()) | ||||||
| 		return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") | 		return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") | ||||||
| 	} | 	} | ||||||
| 	defer res.Body.Close() | 	defer res.Body.Close() | ||||||
| @@ -65,7 +62,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 { | ||||||
| 		logger.SysLog(err.Error()) | 		common.SysLog(err.Error()) | ||||||
| 		return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") | 		return nil, errors.New("无法连接至 GitHub 服务器,请稍后重试!") | ||||||
| 	} | 	} | ||||||
| 	defer res2.Body.Close() | 	defer res2.Body.Close() | ||||||
| @@ -96,7 +93,7 @@ func GitHubOAuth(c *gin.Context) { | |||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	if !config.GitHubOAuthEnabled { | 	if !common.GitHubOAuthEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| 			"message": "管理员未开启通过 GitHub 登录以及注册", | 			"message": "管理员未开启通过 GitHub 登录以及注册", | ||||||
| @@ -125,7 +122,7 @@ func GitHubOAuth(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	} else { | 	} else { | ||||||
| 		if config.RegisterEnabled { | 		if common.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 | ||||||
| @@ -163,7 +160,7 @@ func GitHubOAuth(c *gin.Context) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func GitHubBind(c *gin.Context) { | func GitHubBind(c *gin.Context) { | ||||||
| 	if !config.GitHubOAuthEnabled { | 	if !common.GitHubOAuthEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| 			"message": "管理员未开启通过 GitHub 登录以及注册", | 			"message": "管理员未开启通过 GitHub 登录以及注册", | ||||||
| @@ -219,7 +216,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 := helper.GetRandomString(12) | 	state := common.GetRandomString(12) | ||||||
| 	session.Set("oauth_state", state) | 	session.Set("oauth_state", state) | ||||||
| 	err := session.Save() | 	err := session.Save() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
|   | |||||||
| @@ -2,13 +2,13 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func GetGroups(c *gin.Context) { | func GetGroups(c *gin.Context) { | ||||||
| 	groupNames := make([]string, 0) | 	groupNames := make([]string, 0) | ||||||
| 	for groupName := range common.GroupRatio { | 	for groupName, _ := range common.GroupRatio { | ||||||
| 		groupNames = append(groupNames, groupName) | 		groupNames = append(groupNames, groupName) | ||||||
| 	} | 	} | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
|   | |||||||
| @@ -2,12 +2,27 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
|  | 	"strings" | ||||||
| ) | ) | ||||||
|  |  | ||||||
|  | func parseIntArray(input string) []int { | ||||||
|  | 	values := strings.Split(input, ",") | ||||||
|  | 	result := make([]int, 0) | ||||||
|  |  | ||||||
|  | 	for _, value := range values { | ||||||
|  | 		num, err := strconv.Atoi(strings.TrimSpace(value)) | ||||||
|  | 		if err == nil { | ||||||
|  | 			result = append(result, num) | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	return result | ||||||
|  | } | ||||||
|  |  | ||||||
| func GetAllLogs(c *gin.Context) { | func GetAllLogs(c *gin.Context) { | ||||||
| 	p, _ := strconv.Atoi(c.Query("p")) | 	p, _ := strconv.Atoi(c.Query("p")) | ||||||
| 	if p < 0 { | 	if p < 0 { | ||||||
| @@ -19,8 +34,8 @@ func GetAllLogs(c *gin.Context) { | |||||||
| 	username := c.Query("username") | 	username := c.Query("username") | ||||||
| 	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")) | 	channels := parseIntArray(c.Query("channel")) | ||||||
| 	logs, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, p*config.ItemsPerPage, config.ItemsPerPage, channel) | 	logs, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, p*common.ItemsPerPage, common.ItemsPerPage, channels) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -47,7 +62,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*config.ItemsPerPage, config.ItemsPerPage) | 	logs, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, p*common.ItemsPerPage, common.ItemsPerPage) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -107,8 +122,8 @@ func GetLogsStat(c *gin.Context) { | |||||||
| 	tokenName := c.Query("token_name") | 	tokenName := c.Query("token_name") | ||||||
| 	username := c.Query("username") | 	username := c.Query("username") | ||||||
| 	modelName := c.Query("model_name") | 	modelName := c.Query("model_name") | ||||||
| 	channel, _ := strconv.Atoi(c.Query("channel")) | 	channels := parseIntArray(c.Query("channel")) | ||||||
| 	quotaNum := model.SumUsedQuota(logType, startTimestamp, endTimestamp, modelName, username, tokenName, channel) | 	quotaNum := model.SumUsedQuota(logType, startTimestamp, endTimestamp, modelName, username, tokenName, channels) | ||||||
| 	//tokenNum := model.SumUsedToken(logType, startTimestamp, endTimestamp, modelName, username, "") | 	//tokenNum := model.SumUsedToken(logType, startTimestamp, endTimestamp, modelName, username, "") | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
| @@ -128,8 +143,8 @@ func GetLogsSelfStat(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") | ||||||
| 	channel, _ := strconv.Atoi(c.Query("channel")) | 	channels := parseIntArray(c.Query("channel")) | ||||||
| 	quotaNum := model.SumUsedQuota(logType, startTimestamp, endTimestamp, modelName, username, tokenName, channel) | 	quotaNum := model.SumUsedQuota(logType, startTimestamp, endTimestamp, modelName, username, tokenName, channels) | ||||||
| 	//tokenNum := model.SumUsedToken(logType, startTimestamp, endTimestamp, modelName, username, tokenName) | 	//tokenNum := model.SumUsedToken(logType, startTimestamp, endTimestamp, modelName, username, tokenName) | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
|   | |||||||
| @@ -3,11 +3,9 @@ package controller | |||||||
| import ( | import ( | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/message" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strings" | 	"strings" | ||||||
|  |  | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| @@ -20,55 +18,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":  config.EmailVerificationEnabled, | 			"email_verification":  common.EmailVerificationEnabled, | ||||||
| 			"github_oauth":        config.GitHubOAuthEnabled, | 			"github_oauth":        common.GitHubOAuthEnabled, | ||||||
| 			"github_client_id":    config.GitHubClientId, | 			"github_client_id":    common.GitHubClientId, | ||||||
| 			"system_name":         config.SystemName, | 			"system_name":         common.SystemName, | ||||||
| 			"logo":                config.Logo, | 			"logo":                common.Logo, | ||||||
| 			"footer_html":         config.Footer, | 			"footer_html":         common.Footer, | ||||||
| 			"wechat_qrcode":       config.WeChatAccountQRCodeImageURL, | 			"wechat_qrcode":       common.WeChatAccountQRCodeImageURL, | ||||||
| 			"wechat_login":        config.WeChatAuthEnabled, | 			"wechat_login":        common.WeChatAuthEnabled, | ||||||
| 			"server_address":      config.ServerAddress, | 			"server_address":      common.ServerAddress, | ||||||
| 			"turnstile_check":     config.TurnstileCheckEnabled, | 			"turnstile_check":     common.TurnstileCheckEnabled, | ||||||
| 			"turnstile_site_key":  config.TurnstileSiteKey, | 			"turnstile_site_key":  common.TurnstileSiteKey, | ||||||
| 			"top_up_link":         config.TopUpLink, | 			"top_up_link":         common.TopUpLink, | ||||||
| 			"chat_link":           config.ChatLink, | 			"chat_link":           common.ChatLink, | ||||||
| 			"quota_per_unit":      config.QuotaPerUnit, | 			"quota_per_unit":      common.QuotaPerUnit, | ||||||
| 			"display_in_currency": config.DisplayInCurrencyEnabled, | 			"display_in_currency": common.DisplayInCurrencyEnabled, | ||||||
| 		}, | 		}, | ||||||
| 	}) | 	}) | ||||||
| 	return | 	return | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetNotice(c *gin.Context) { | func GetNotice(c *gin.Context) { | ||||||
| 	config.OptionMapRWMutex.RLock() | 	common.OptionMapRWMutex.RLock() | ||||||
| 	defer config.OptionMapRWMutex.RUnlock() | 	defer common.OptionMapRWMutex.RUnlock() | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
| 		"message": "", | 		"message": "", | ||||||
| 		"data":    config.OptionMap["Notice"], | 		"data":    common.OptionMap["Notice"], | ||||||
| 	}) | 	}) | ||||||
| 	return | 	return | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetAbout(c *gin.Context) { | func GetAbout(c *gin.Context) { | ||||||
| 	config.OptionMapRWMutex.RLock() | 	common.OptionMapRWMutex.RLock() | ||||||
| 	defer config.OptionMapRWMutex.RUnlock() | 	defer common.OptionMapRWMutex.RUnlock() | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
| 		"message": "", | 		"message": "", | ||||||
| 		"data":    config.OptionMap["About"], | 		"data":    common.OptionMap["About"], | ||||||
| 	}) | 	}) | ||||||
| 	return | 	return | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetHomePageContent(c *gin.Context) { | func GetHomePageContent(c *gin.Context) { | ||||||
| 	config.OptionMapRWMutex.RLock() | 	common.OptionMapRWMutex.RLock() | ||||||
| 	defer config.OptionMapRWMutex.RUnlock() | 	defer common.OptionMapRWMutex.RUnlock() | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
| 		"message": "", | 		"message": "", | ||||||
| 		"data":    config.OptionMap["HomePageContent"], | 		"data":    common.OptionMap["HomePageContent"], | ||||||
| 	}) | 	}) | ||||||
| 	return | 	return | ||||||
| } | } | ||||||
| @@ -82,9 +80,9 @@ func SendEmailVerification(c *gin.Context) { | |||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if config.EmailDomainRestrictionEnabled { | 	if common.EmailDomainRestrictionEnabled { | ||||||
| 		allowed := false | 		allowed := false | ||||||
| 		for _, domain := range config.EmailDomainWhitelist { | 		for _, domain := range common.EmailDomainWhitelist { | ||||||
| 			if strings.HasSuffix(email, "@"+domain) { | 			if strings.HasSuffix(email, "@"+domain) { | ||||||
| 				allowed = true | 				allowed = true | ||||||
| 				break | 				break | ||||||
| @@ -107,11 +105,11 @@ 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邮箱验证邮件", config.SystemName) | 	subject := fmt.Sprintf("%s邮箱验证邮件", common.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>", config.SystemName, code, common.VerificationValidMinutes) | 		"<p>验证码 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, code, common.VerificationValidMinutes) | ||||||
| 	err := message.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{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -144,13 +142,13 @@ 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", config.ServerAddress, email, code) | 	link := fmt.Sprintf("%s/user/reset?email=%s&token=%s", common.ServerAddress, email, code) | ||||||
| 	subject := fmt.Sprintf("%s密码重置", config.SystemName) | 	subject := fmt.Sprintf("%s密码重置", common.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>", config.SystemName, link, link, common.VerificationValidMinutes) | 		"<p>重置链接 %d 分钟内有效,如果不是本人操作,请忽略。</p>", common.SystemName, link, link, common.VerificationValidMinutes) | ||||||
| 	err := message.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{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
|   | |||||||
| @@ -2,14 +2,8 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"fmt" | 	"fmt" | ||||||
|  |  | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/helper" |  | ||||||
| 	relaymodel "github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"net/http" |  | ||||||
| ) | ) | ||||||
|  |  | ||||||
| // https://platform.openai.com/docs/api-reference/models/list | // https://platform.openai.com/docs/api-reference/models/list | ||||||
| @@ -41,7 +35,6 @@ type OpenAIModels struct { | |||||||
|  |  | ||||||
| var openAIModels []OpenAIModels | var openAIModels []OpenAIModels | ||||||
| var openAIModelsMap map[string]OpenAIModels | var openAIModelsMap map[string]OpenAIModels | ||||||
| var channelId2Models map[int][]string |  | ||||||
|  |  | ||||||
| func init() { | func init() { | ||||||
| 	var permission []OpenAIModelPermission | 	var permission []OpenAIModelPermission | ||||||
| @@ -60,63 +53,552 @@ func init() { | |||||||
| 		IsBlocking:         false, | 		IsBlocking:         false, | ||||||
| 	}) | 	}) | ||||||
| 	// https://platform.openai.com/docs/models/model-endpoint-compatibility | 	// https://platform.openai.com/docs/models/model-endpoint-compatibility | ||||||
| 	for i := 0; i < constant.APITypeDummy; i++ { | 	openAIModels = []OpenAIModels{ | ||||||
| 		if i == constant.APITypeAIProxyLibrary { | 		{ | ||||||
| 			continue | 			Id:         "dall-e-2", | ||||||
| 		} | 			Object:     "model", | ||||||
| 		adaptor := helper.GetAdaptor(i) | 			Created:    1677649963, | ||||||
| 		channelName := adaptor.GetChannelName() | 			OwnedBy:    "openai", | ||||||
| 		modelNames := adaptor.GetModelList() | 			Permission: permission, | ||||||
| 		for _, modelName := range modelNames { | 			Root:       "dall-e-2", | ||||||
| 			openAIModels = append(openAIModels, OpenAIModels{ | 			Parent:     nil, | ||||||
| 				Id:         modelName, | 		}, | ||||||
| 				Object:     "model", | 		{ | ||||||
| 				Created:    1626777600, | 			Id:         "dall-e-3", | ||||||
| 				OwnedBy:    channelName, | 			Object:     "model", | ||||||
| 				Permission: permission, | 			Created:    1677649963, | ||||||
| 				Root:       modelName, | 			OwnedBy:    "openai", | ||||||
| 				Parent:     nil, | 			Permission: permission, | ||||||
| 			}) | 			Root:       "dall-e-3", | ||||||
| 		} | 			Parent:     nil, | ||||||
| 	} | 		}, | ||||||
| 	for _, channelType := range openai.CompatibleChannels { | 		{ | ||||||
| 		if channelType == common.ChannelTypeAzure { | 			Id:         "whisper-1", | ||||||
| 			continue | 			Object:     "model", | ||||||
| 		} | 			Created:    1677649963, | ||||||
| 		channelName, channelModelList := openai.GetCompatibleChannelMeta(channelType) | 			OwnedBy:    "openai", | ||||||
| 		for _, modelName := range channelModelList { | 			Permission: permission, | ||||||
| 			openAIModels = append(openAIModels, OpenAIModels{ | 			Root:       "whisper-1", | ||||||
| 				Id:         modelName, | 			Parent:     nil, | ||||||
| 				Object:     "model", | 		}, | ||||||
| 				Created:    1626777600, | 		{ | ||||||
| 				OwnedBy:    channelName, | 			Id:         "tts-1", | ||||||
| 				Permission: permission, | 			Object:     "model", | ||||||
| 				Root:       modelName, | 			Created:    1677649963, | ||||||
| 				Parent:     nil, | 			OwnedBy:    "openai", | ||||||
| 			}) | 			Permission: permission, | ||||||
| 		} | 			Root:       "tts-1", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "tts-1-1106", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "tts-1-1106", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "tts-1-hd", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "tts-1-hd", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "tts-1-hd-1106", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "tts-1-hd-1106", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-3.5-turbo", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-3.5-turbo", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-3.5-turbo-0301", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-3.5-turbo-0301", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-3.5-turbo-0613", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-3.5-turbo-0613", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-3.5-turbo-16k", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-3.5-turbo-16k", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-3.5-turbo-16k-0613", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-3.5-turbo-16k-0613", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-3.5-turbo-1106", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1699593571, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-3.5-turbo-1106", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-3.5-turbo-instruct", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-3.5-turbo-instruct", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-4", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-4", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-4-0314", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-4-0314", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-4-0613", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-4-0613", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-4-32k", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-4-32k", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-4-32k-0314", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-4-32k-0314", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-4-32k-0613", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-4-32k-0613", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-4-1106-preview", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1699593571, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-4-1106-preview", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gpt-4-vision-preview", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1699593571, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gpt-4-vision-preview", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-embedding-ada-002", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-embedding-ada-002", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-davinci-003", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-davinci-003", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-davinci-002", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-davinci-002", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-curie-001", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-curie-001", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-babbage-001", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-babbage-001", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-ada-001", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-ada-001", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-moderation-latest", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-moderation-latest", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-moderation-stable", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-moderation-stable", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-davinci-edit-001", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-davinci-edit-001", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "code-davinci-edit-001", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "code-davinci-edit-001", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "davinci-002", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "davinci-002", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "babbage-002", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "openai", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "babbage-002", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "claude-instant-1", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "anthropic", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "claude-instant-1", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "claude-2", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "anthropic", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "claude-2", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "claude-2.1", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "anthropic", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "claude-2.1", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "claude-2.0", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "anthropic", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "claude-2.0", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "ERNIE-Bot", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "baidu", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "ERNIE-Bot", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "ERNIE-Bot-turbo", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "baidu", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "ERNIE-Bot-turbo", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "ERNIE-Bot-4", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "baidu", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "ERNIE-Bot-4", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "Embedding-V1", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "baidu", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "Embedding-V1", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "PaLM-2", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "google palm", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "PaLM-2", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gemini-pro", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "google gemini", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gemini-pro", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "gemini-pro-vision", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "google gemini", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "gemini-pro-vision", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "chatglm_turbo", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "zhipu", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "chatglm_turbo", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "chatglm_pro", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "zhipu", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "chatglm_pro", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "chatglm_std", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "zhipu", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "chatglm_std", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "chatglm_lite", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "zhipu", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "chatglm_lite", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "qwen-turbo", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "ali", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "qwen-turbo", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "qwen-plus", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "ali", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "qwen-plus", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "qwen-max", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "ali", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "qwen-max", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "qwen-max-longcontext", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "ali", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "qwen-max-longcontext", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "text-embedding-v1", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "ali", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "text-embedding-v1", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "SparkDesk", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "xunfei", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "SparkDesk", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "360GPT_S2_V9", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "360", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "360GPT_S2_V9", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "embedding-bert-512-v1", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "360", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "embedding-bert-512-v1", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "embedding_s1_v1", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "360", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "embedding_s1_v1", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "semantic_similarity_s1_v1", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "360", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "semantic_similarity_s1_v1", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
|  | 		{ | ||||||
|  | 			Id:         "hunyuan", | ||||||
|  | 			Object:     "model", | ||||||
|  | 			Created:    1677649963, | ||||||
|  | 			OwnedBy:    "tencent", | ||||||
|  | 			Permission: permission, | ||||||
|  | 			Root:       "hunyuan", | ||||||
|  | 			Parent:     nil, | ||||||
|  | 		}, | ||||||
| 	} | 	} | ||||||
| 	openAIModelsMap = make(map[string]OpenAIModels) | 	openAIModelsMap = make(map[string]OpenAIModels) | ||||||
| 	for _, model := range openAIModels { | 	for _, model := range openAIModels { | ||||||
| 		openAIModelsMap[model.Id] = model | 		openAIModelsMap[model.Id] = model | ||||||
| 	} | 	} | ||||||
| 	channelId2Models = make(map[int][]string) |  | ||||||
| 	for i := 1; i < common.ChannelTypeDummy; i++ { |  | ||||||
| 		adaptor := helper.GetAdaptor(constant.ChannelType2APIType(i)) |  | ||||||
| 		meta := &util.RelayMeta{ |  | ||||||
| 			ChannelType: i, |  | ||||||
| 		} |  | ||||||
| 		adaptor.Init(meta) |  | ||||||
| 		channelId2Models[i] = adaptor.GetModelList() |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func DashboardListModels(c *gin.Context) { |  | ||||||
| 	c.JSON(http.StatusOK, gin.H{ |  | ||||||
| 		"success": true, |  | ||||||
| 		"message": "", |  | ||||||
| 		"data":    channelId2Models, |  | ||||||
| 	}) |  | ||||||
| } | } | ||||||
|  |  | ||||||
| func ListModels(c *gin.Context) { | func ListModels(c *gin.Context) { | ||||||
| @@ -131,14 +613,14 @@ func RetrieveModel(c *gin.Context) { | |||||||
| 	if model, ok := openAIModelsMap[modelId]; ok { | 	if model, ok := openAIModelsMap[modelId]; ok { | ||||||
| 		c.JSON(200, model) | 		c.JSON(200, model) | ||||||
| 	} else { | 	} else { | ||||||
| 		Error := relaymodel.Error{ | 		openAIError := OpenAIError{ | ||||||
| 			Message: fmt.Sprintf("The model '%s' does not exist", modelId), | 			Message: fmt.Sprintf("The model '%s' does not exist", modelId), | ||||||
| 			Type:    "invalid_request_error", | 			Type:    "invalid_request_error", | ||||||
| 			Param:   "model", | 			Param:   "model", | ||||||
| 			Code:    "model_not_found", | 			Code:    "model_not_found", | ||||||
| 		} | 		} | ||||||
| 		c.JSON(200, gin.H{ | 		c.JSON(200, gin.H{ | ||||||
| 			"error": Error, | 			"error": openAIError, | ||||||
| 		}) | 		}) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -2,10 +2,9 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strings" | 	"strings" | ||||||
|  |  | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| @@ -13,17 +12,17 @@ import ( | |||||||
|  |  | ||||||
| func GetOptions(c *gin.Context) { | func GetOptions(c *gin.Context) { | ||||||
| 	var options []*model.Option | 	var options []*model.Option | ||||||
| 	config.OptionMapRWMutex.Lock() | 	common.OptionMapRWMutex.Lock() | ||||||
| 	for k, v := range config.OptionMap { | 	for k, v := range common.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: helper.Interface2String(v), | 			Value: common.Interface2String(v), | ||||||
| 		}) | 		}) | ||||||
| 	} | 	} | ||||||
| 	config.OptionMapRWMutex.Unlock() | 	common.OptionMapRWMutex.Unlock() | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
| 		"message": "", | 		"message": "", | ||||||
| @@ -44,7 +43,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 	} | 	} | ||||||
| 	switch option.Key { | 	switch option.Key { | ||||||
| 	case "Theme": | 	case "Theme": | ||||||
| 		if !config.ValidThemes[option.Value] { | 		if !common.ValidThemes[option.Value] { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "无效的主题", | 				"message": "无效的主题", | ||||||
| @@ -52,7 +51,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	case "GitHubOAuthEnabled": | 	case "GitHubOAuthEnabled": | ||||||
| 		if option.Value == "true" && config.GitHubClientId == "" { | 		if option.Value == "true" && common.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!", | ||||||
| @@ -60,7 +59,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	case "EmailDomainRestrictionEnabled": | 	case "EmailDomainRestrictionEnabled": | ||||||
| 		if option.Value == "true" && len(config.EmailDomainWhitelist) == 0 { | 		if option.Value == "true" && len(common.EmailDomainWhitelist) == 0 { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "无法启用邮箱域名限制,请先填入限制的邮箱域名!", | 				"message": "无法启用邮箱域名限制,请先填入限制的邮箱域名!", | ||||||
| @@ -68,7 +67,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	case "WeChatAuthEnabled": | 	case "WeChatAuthEnabled": | ||||||
| 		if option.Value == "true" && config.WeChatServerAddress == "" { | 		if option.Value == "true" && common.WeChatServerAddress == "" { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "无法启用微信登录,请先填入微信登录相关配置信息!", | 				"message": "无法启用微信登录,请先填入微信登录相关配置信息!", | ||||||
| @@ -76,7 +75,7 @@ func UpdateOption(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	case "TurnstileCheckEnabled": | 	case "TurnstileCheckEnabled": | ||||||
| 		if option.Value == "true" && config.TurnstileSiteKey == "" { | 		if option.Value == "true" && common.TurnstileSiteKey == "" { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "无法启用 Turnstile 校验,请先填入 Turnstile 校验相关配置信息!", | 				"message": "无法启用 Turnstile 校验,请先填入 Turnstile 校验相关配置信息!", | ||||||
|   | |||||||
| @@ -2,10 +2,9 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -14,7 +13,7 @@ func GetAllRedemptions(c *gin.Context) { | |||||||
| 	if p < 0 { | 	if p < 0 { | ||||||
| 		p = 0 | 		p = 0 | ||||||
| 	} | 	} | ||||||
| 	redemptions, err := model.GetAllRedemptions(p*config.ItemsPerPage, config.ItemsPerPage) | 	redemptions, err := model.GetAllRedemptions(p*common.ItemsPerPage, common.ItemsPerPage) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -106,12 +105,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 := helper.GetUUID() | 		key := common.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: helper.GetTimestamp(), | 			CreatedTime: common.GetTimestamp(), | ||||||
| 			Quota:       redemption.Quota, | 			Quota:       redemption.Quota, | ||||||
| 		} | 		} | ||||||
| 		err = cleanRedemption.Insert() | 		err = cleanRedemption.Insert() | ||||||
|   | |||||||
| @@ -1,37 +1,63 @@ | |||||||
| package aiproxy | package controller | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"bufio" | 	"bufio" | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
| 
 | 
 | ||||||
| // https://docs.aiproxy.io/dev/library#使用已经定制好的知识库进行对话问答 | // https://docs.aiproxy.io/dev/library#使用已经定制好的知识库进行对话问答 | ||||||
| 
 | 
 | ||||||
| func ConvertRequest(request model.GeneralOpenAIRequest) *LibraryRequest { | type AIProxyLibraryRequest struct { | ||||||
|  | 	Model     string `json:"model"` | ||||||
|  | 	Query     string `json:"query"` | ||||||
|  | 	LibraryId string `json:"libraryId"` | ||||||
|  | 	Stream    bool   `json:"stream"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AIProxyLibraryError struct { | ||||||
|  | 	ErrCode int    `json:"errCode"` | ||||||
|  | 	Message string `json:"message"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AIProxyLibraryDocument struct { | ||||||
|  | 	Title string `json:"title"` | ||||||
|  | 	URL   string `json:"url"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AIProxyLibraryResponse struct { | ||||||
|  | 	Success   bool                     `json:"success"` | ||||||
|  | 	Answer    string                   `json:"answer"` | ||||||
|  | 	Documents []AIProxyLibraryDocument `json:"documents"` | ||||||
|  | 	AIProxyLibraryError | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AIProxyLibraryStreamResponse struct { | ||||||
|  | 	Content   string                   `json:"content"` | ||||||
|  | 	Finish    bool                     `json:"finish"` | ||||||
|  | 	Model     string                   `json:"model"` | ||||||
|  | 	Documents []AIProxyLibraryDocument `json:"documents"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func requestOpenAI2AIProxyLibrary(request GeneralOpenAIRequest) *AIProxyLibraryRequest { | ||||||
| 	query := "" | 	query := "" | ||||||
| 	if len(request.Messages) != 0 { | 	if len(request.Messages) != 0 { | ||||||
| 		query = request.Messages[len(request.Messages)-1].StringContent() | 		query = request.Messages[len(request.Messages)-1].StringContent() | ||||||
| 	} | 	} | ||||||
| 	return &LibraryRequest{ | 	return &AIProxyLibraryRequest{ | ||||||
| 		Model:  request.Model, | 		Model:  request.Model, | ||||||
| 		Stream: request.Stream, | 		Stream: request.Stream, | ||||||
| 		Query:  query, | 		Query:  query, | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func aiProxyDocuments2Markdown(documents []LibraryDocument) string { | func aiProxyDocuments2Markdown(documents []AIProxyLibraryDocument) string { | ||||||
| 	if len(documents) == 0 { | 	if len(documents) == 0 { | ||||||
| 		return "" | 		return "" | ||||||
| 	} | 	} | ||||||
| @@ -42,52 +68,52 @@ func aiProxyDocuments2Markdown(documents []LibraryDocument) string { | |||||||
| 	return content | 	return content | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func responseAIProxyLibrary2OpenAI(response *LibraryResponse) *openai.TextResponse { | func responseAIProxyLibrary2OpenAI(response *AIProxyLibraryResponse) *OpenAITextResponse { | ||||||
| 	content := response.Answer + aiProxyDocuments2Markdown(response.Documents) | 	content := response.Answer + aiProxyDocuments2Markdown(response.Documents) | ||||||
| 	choice := openai.TextResponseChoice{ | 	choice := OpenAITextResponseChoice{ | ||||||
| 		Index: 0, | 		Index: 0, | ||||||
| 		Message: model.Message{ | 		Message: Message{ | ||||||
| 			Role:    "assistant", | 			Role:    "assistant", | ||||||
| 			Content: content, | 			Content: content, | ||||||
| 		}, | 		}, | ||||||
| 		FinishReason: "stop", | 		FinishReason: "stop", | ||||||
| 	} | 	} | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := OpenAITextResponse{ | ||||||
| 		Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), | 		Id:      common.GetUUID(), | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Choices: []openai.TextResponseChoice{choice}, | 		Choices: []OpenAITextResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| 	return &fullTextResponse | 	return &fullTextResponse | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func documentsAIProxyLibrary(documents []LibraryDocument) *openai.ChatCompletionsStreamResponse { | func documentsAIProxyLibrary(documents []AIProxyLibraryDocument) *ChatCompletionsStreamResponse { | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice ChatCompletionsStreamResponseChoice | ||||||
| 	choice.Delta.Content = aiProxyDocuments2Markdown(documents) | 	choice.Delta.Content = aiProxyDocuments2Markdown(documents) | ||||||
| 	choice.FinishReason = &constant.StopFinishReason | 	choice.FinishReason = &stopFinishReason | ||||||
| 	return &openai.ChatCompletionsStreamResponse{ | 	return &ChatCompletionsStreamResponse{ | ||||||
| 		Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), | 		Id:      common.GetUUID(), | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Model:   "", | 		Model:   "", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func streamResponseAIProxyLibrary2OpenAI(response *LibraryStreamResponse) *openai.ChatCompletionsStreamResponse { | func streamResponseAIProxyLibrary2OpenAI(response *AIProxyLibraryStreamResponse) *ChatCompletionsStreamResponse { | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice ChatCompletionsStreamResponseChoice | ||||||
| 	choice.Delta.Content = response.Content | 	choice.Delta.Content = response.Content | ||||||
| 	return &openai.ChatCompletionsStreamResponse{ | 	return &ChatCompletionsStreamResponse{ | ||||||
| 		Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), | 		Id:      common.GetUUID(), | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Model:   response.Model, | 		Model:   response.Model, | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func aiProxyLibraryStreamHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var usage model.Usage | 	var usage Usage | ||||||
| 	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 { | ||||||
| @@ -117,15 +143,15 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 		} | 		} | ||||||
| 		stopChan <- true | 		stopChan <- true | ||||||
| 	}() | 	}() | ||||||
| 	common.SetEventStreamHeaders(c) | 	setEventStreamHeaders(c) | ||||||
| 	var documents []LibraryDocument | 	var documents []AIProxyLibraryDocument | ||||||
| 	c.Stream(func(w io.Writer) bool { | 	c.Stream(func(w io.Writer) bool { | ||||||
| 		select { | 		select { | ||||||
| 		case data := <-dataChan: | 		case data := <-dataChan: | ||||||
| 			var AIProxyLibraryResponse LibraryStreamResponse | 			var AIProxyLibraryResponse AIProxyLibraryStreamResponse | ||||||
| 			err := json.Unmarshal([]byte(data), &AIProxyLibraryResponse) | 			err := json.Unmarshal([]byte(data), &AIProxyLibraryResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error unmarshalling stream response: " + err.Error()) | 				common.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			if len(AIProxyLibraryResponse.Documents) != 0 { | 			if len(AIProxyLibraryResponse.Documents) != 0 { | ||||||
| @@ -134,7 +160,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 			response := streamResponseAIProxyLibrary2OpenAI(&AIProxyLibraryResponse) | 			response := streamResponseAIProxyLibrary2OpenAI(&AIProxyLibraryResponse) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error marshalling stream response: " + err.Error()) | 				common.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)}) | ||||||
| @@ -143,7 +169,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 			response := documentsAIProxyLibrary(documents) | 			response := documentsAIProxyLibrary(documents) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error marshalling stream response: " + err.Error()) | 				common.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)}) | ||||||
| @@ -153,28 +179,28 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 	}) | 	}) | ||||||
| 	err := resp.Body.Close() | 	err := resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	return nil, &usage | 	return nil, &usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Handler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func aiProxyLibraryHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var AIProxyLibraryResponse LibraryResponse | 	var AIProxyLibraryResponse AIProxyLibraryResponse | ||||||
| 	responseBody, err := io.ReadAll(resp.Body) | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = json.Unmarshal(responseBody, &AIProxyLibraryResponse) | 	err = json.Unmarshal(responseBody, &AIProxyLibraryResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	if AIProxyLibraryResponse.ErrCode != 0 { | 	if AIProxyLibraryResponse.ErrCode != 0 { | ||||||
| 		return &model.ErrorWithStatusCode{ | 		return &OpenAIErrorWithStatusCode{ | ||||||
| 			Error: model.Error{ | 			OpenAIError: OpenAIError{ | ||||||
| 				Message: AIProxyLibraryResponse.Message, | 				Message: AIProxyLibraryResponse.Message, | ||||||
| 				Type:    strconv.Itoa(AIProxyLibraryResponse.ErrCode), | 				Type:    strconv.Itoa(AIProxyLibraryResponse.ErrCode), | ||||||
| 				Code:    AIProxyLibraryResponse.ErrCode, | 				Code:    AIProxyLibraryResponse.ErrCode, | ||||||
| @@ -185,13 +211,10 @@ func Handler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, * | |||||||
| 	fullTextResponse := responseAIProxyLibrary2OpenAI(&AIProxyLibraryResponse) | 	fullTextResponse := responseAIProxyLibrary2OpenAI(&AIProxyLibraryResponse) | ||||||
| 	jsonResponse, err := json.Marshal(fullTextResponse) | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
| 	c.Writer.WriteHeader(resp.StatusCode) | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
| 	_, err = c.Writer.Write(jsonResponse) | 	_, err = c.Writer.Write(jsonResponse) | ||||||
| 	if err != nil { |  | ||||||
| 		return openai.ErrorWrapper(err, "write_response_body_failed", http.StatusInternalServerError), nil |  | ||||||
| 	} |  | ||||||
| 	return nil, &fullTextResponse.Usage | 	return nil, &fullTextResponse.Usage | ||||||
| } | } | ||||||
| @@ -1,59 +1,118 @@ | |||||||
| package ali | package controller | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"bufio" | 	"bufio" | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
| 
 | 
 | ||||||
| // https://help.aliyun.com/document_detail/613695.html?spm=a2c4g.2399480.0.0.1adb778fAdzP9w#341800c0f8w0r | // https://help.aliyun.com/document_detail/613695.html?spm=a2c4g.2399480.0.0.1adb778fAdzP9w#341800c0f8w0r | ||||||
| 
 | 
 | ||||||
| const EnableSearchModelSuffix = "-internet" | type AliMessage struct { | ||||||
|  | 	Content string `json:"content"` | ||||||
|  | 	Role    string `json:"role"` | ||||||
|  | } | ||||||
| 
 | 
 | ||||||
| func ConvertRequest(request model.GeneralOpenAIRequest) *ChatRequest { | type AliInput struct { | ||||||
| 	messages := make([]Message, 0, len(request.Messages)) | 	//Prompt   string       `json:"prompt"` | ||||||
|  | 	Messages []AliMessage `json:"messages"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AliParameters struct { | ||||||
|  | 	TopP              float64 `json:"top_p,omitempty"` | ||||||
|  | 	TopK              int     `json:"top_k,omitempty"` | ||||||
|  | 	Seed              uint64  `json:"seed,omitempty"` | ||||||
|  | 	EnableSearch      bool    `json:"enable_search,omitempty"` | ||||||
|  | 	IncrementalOutput bool    `json:"incremental_output,omitempty"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AliChatRequest struct { | ||||||
|  | 	Model      string        `json:"model"` | ||||||
|  | 	Input      AliInput      `json:"input"` | ||||||
|  | 	Parameters AliParameters `json:"parameters,omitempty"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AliEmbeddingRequest struct { | ||||||
|  | 	Model string `json:"model"` | ||||||
|  | 	Input struct { | ||||||
|  | 		Texts []string `json:"texts"` | ||||||
|  | 	} `json:"input"` | ||||||
|  | 	Parameters *struct { | ||||||
|  | 		TextType string `json:"text_type,omitempty"` | ||||||
|  | 	} `json:"parameters,omitempty"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AliEmbedding struct { | ||||||
|  | 	Embedding []float64 `json:"embedding"` | ||||||
|  | 	TextIndex int       `json:"text_index"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AliEmbeddingResponse struct { | ||||||
|  | 	Output struct { | ||||||
|  | 		Embeddings []AliEmbedding `json:"embeddings"` | ||||||
|  | 	} `json:"output"` | ||||||
|  | 	Usage AliUsage `json:"usage"` | ||||||
|  | 	AliError | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AliError struct { | ||||||
|  | 	Code      string `json:"code"` | ||||||
|  | 	Message   string `json:"message"` | ||||||
|  | 	RequestId string `json:"request_id"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AliUsage struct { | ||||||
|  | 	InputTokens  int `json:"input_tokens"` | ||||||
|  | 	OutputTokens int `json:"output_tokens"` | ||||||
|  | 	TotalTokens  int `json:"total_tokens"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AliOutput struct { | ||||||
|  | 	Text         string `json:"text"` | ||||||
|  | 	FinishReason string `json:"finish_reason"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type AliChatResponse struct { | ||||||
|  | 	Output AliOutput `json:"output"` | ||||||
|  | 	Usage  AliUsage  `json:"usage"` | ||||||
|  | 	AliError | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | const AliEnableSearchModelSuffix = "-internet" | ||||||
|  | 
 | ||||||
|  | func requestOpenAI2Ali(request GeneralOpenAIRequest) *AliChatRequest { | ||||||
|  | 	messages := make([]AliMessage, 0, len(request.Messages)) | ||||||
| 	for i := 0; i < len(request.Messages); i++ { | 	for i := 0; i < len(request.Messages); i++ { | ||||||
| 		message := request.Messages[i] | 		message := request.Messages[i] | ||||||
| 		messages = append(messages, Message{ | 		messages = append(messages, AliMessage{ | ||||||
| 			Content: message.StringContent(), | 			Content: message.StringContent(), | ||||||
| 			Role:    strings.ToLower(message.Role), | 			Role:    strings.ToLower(message.Role), | ||||||
| 		}) | 		}) | ||||||
| 	} | 	} | ||||||
| 	enableSearch := false | 	enableSearch := false | ||||||
| 	aliModel := request.Model | 	aliModel := request.Model | ||||||
| 	if strings.HasSuffix(aliModel, EnableSearchModelSuffix) { | 	if strings.HasSuffix(aliModel, AliEnableSearchModelSuffix) { | ||||||
| 		enableSearch = true | 		enableSearch = true | ||||||
| 		aliModel = strings.TrimSuffix(aliModel, EnableSearchModelSuffix) | 		aliModel = strings.TrimSuffix(aliModel, AliEnableSearchModelSuffix) | ||||||
| 	} | 	} | ||||||
| 	if request.TopP >= 1 { | 	return &AliChatRequest{ | ||||||
| 		request.TopP = 0.9999 |  | ||||||
| 	} |  | ||||||
| 	return &ChatRequest{ |  | ||||||
| 		Model: aliModel, | 		Model: aliModel, | ||||||
| 		Input: Input{ | 		Input: AliInput{ | ||||||
| 			Messages: messages, | 			Messages: messages, | ||||||
| 		}, | 		}, | ||||||
| 		Parameters: Parameters{ | 		Parameters: AliParameters{ | ||||||
| 			EnableSearch:      enableSearch, | 			EnableSearch:      enableSearch, | ||||||
| 			IncrementalOutput: request.Stream, | 			IncrementalOutput: request.Stream, | ||||||
| 			Seed:              uint64(request.Seed), |  | ||||||
| 			MaxTokens:         request.MaxTokens, |  | ||||||
| 			Temperature:       request.Temperature, |  | ||||||
| 			TopP:              request.TopP, |  | ||||||
| 		}, | 		}, | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func ConvertEmbeddingRequest(request model.GeneralOpenAIRequest) *EmbeddingRequest { | func embeddingRequestOpenAI2Ali(request GeneralOpenAIRequest) *AliEmbeddingRequest { | ||||||
| 	return &EmbeddingRequest{ | 	return &AliEmbeddingRequest{ | ||||||
| 		Model: "text-embedding-v1", | 		Model: "text-embedding-v1", | ||||||
| 		Input: struct { | 		Input: struct { | ||||||
| 			Texts []string `json:"texts"` | 			Texts []string `json:"texts"` | ||||||
| @@ -63,21 +122,21 @@ func ConvertEmbeddingRequest(request model.GeneralOpenAIRequest) *EmbeddingReque | |||||||
| 	} | 	} | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func EmbeddingHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func aliEmbeddingHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var aliResponse EmbeddingResponse | 	var aliResponse AliEmbeddingResponse | ||||||
| 	err := json.NewDecoder(resp.Body).Decode(&aliResponse) | 	err := json.NewDecoder(resp.Body).Decode(&aliResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	if aliResponse.Code != "" { | 	if aliResponse.Code != "" { | ||||||
| 		return &model.ErrorWithStatusCode{ | 		return &OpenAIErrorWithStatusCode{ | ||||||
| 			Error: model.Error{ | 			OpenAIError: OpenAIError{ | ||||||
| 				Message: aliResponse.Message, | 				Message: aliResponse.Message, | ||||||
| 				Type:    aliResponse.Code, | 				Type:    aliResponse.Code, | ||||||
| 				Param:   aliResponse.RequestId, | 				Param:   aliResponse.RequestId, | ||||||
| @@ -90,7 +149,7 @@ func EmbeddingHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStat | |||||||
| 	fullTextResponse := embeddingResponseAli2OpenAI(&aliResponse) | 	fullTextResponse := embeddingResponseAli2OpenAI(&aliResponse) | ||||||
| 	jsonResponse, err := json.Marshal(fullTextResponse) | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
| 	c.Writer.WriteHeader(resp.StatusCode) | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
| @@ -98,16 +157,16 @@ func EmbeddingHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStat | |||||||
| 	return nil, &fullTextResponse.Usage | 	return nil, &fullTextResponse.Usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func embeddingResponseAli2OpenAI(response *EmbeddingResponse) *openai.EmbeddingResponse { | func embeddingResponseAli2OpenAI(response *AliEmbeddingResponse) *OpenAIEmbeddingResponse { | ||||||
| 	openAIEmbeddingResponse := openai.EmbeddingResponse{ | 	openAIEmbeddingResponse := OpenAIEmbeddingResponse{ | ||||||
| 		Object: "list", | 		Object: "list", | ||||||
| 		Data:   make([]openai.EmbeddingResponseItem, 0, len(response.Output.Embeddings)), | 		Data:   make([]OpenAIEmbeddingResponseItem, 0, len(response.Output.Embeddings)), | ||||||
| 		Model:  "text-embedding-v1", | 		Model:  "text-embedding-v1", | ||||||
| 		Usage:  model.Usage{TotalTokens: response.Usage.TotalTokens}, | 		Usage:  Usage{TotalTokens: response.Usage.TotalTokens}, | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	for _, item := range response.Output.Embeddings { | 	for _, item := range response.Output.Embeddings { | ||||||
| 		openAIEmbeddingResponse.Data = append(openAIEmbeddingResponse.Data, openai.EmbeddingResponseItem{ | 		openAIEmbeddingResponse.Data = append(openAIEmbeddingResponse.Data, OpenAIEmbeddingResponseItem{ | ||||||
| 			Object:    `embedding`, | 			Object:    `embedding`, | ||||||
| 			Index:     item.TextIndex, | 			Index:     item.TextIndex, | ||||||
| 			Embedding: item.Embedding, | 			Embedding: item.Embedding, | ||||||
| @@ -116,21 +175,21 @@ func embeddingResponseAli2OpenAI(response *EmbeddingResponse) *openai.EmbeddingR | |||||||
| 	return &openAIEmbeddingResponse | 	return &openAIEmbeddingResponse | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func responseAli2OpenAI(response *ChatResponse) *openai.TextResponse { | func responseAli2OpenAI(response *AliChatResponse) *OpenAITextResponse { | ||||||
| 	choice := openai.TextResponseChoice{ | 	choice := OpenAITextResponseChoice{ | ||||||
| 		Index: 0, | 		Index: 0, | ||||||
| 		Message: model.Message{ | 		Message: Message{ | ||||||
| 			Role:    "assistant", | 			Role:    "assistant", | ||||||
| 			Content: response.Output.Text, | 			Content: response.Output.Text, | ||||||
| 		}, | 		}, | ||||||
| 		FinishReason: response.Output.FinishReason, | 		FinishReason: response.Output.FinishReason, | ||||||
| 	} | 	} | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := OpenAITextResponse{ | ||||||
| 		Id:      response.RequestId, | 		Id:      response.RequestId, | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Choices: []openai.TextResponseChoice{choice}, | 		Choices: []OpenAITextResponseChoice{choice}, | ||||||
| 		Usage: model.Usage{ | 		Usage: Usage{ | ||||||
| 			PromptTokens:     response.Usage.InputTokens, | 			PromptTokens:     response.Usage.InputTokens, | ||||||
| 			CompletionTokens: response.Usage.OutputTokens, | 			CompletionTokens: response.Usage.OutputTokens, | ||||||
| 			TotalTokens:      response.Usage.InputTokens + response.Usage.OutputTokens, | 			TotalTokens:      response.Usage.InputTokens + response.Usage.OutputTokens, | ||||||
| @@ -139,25 +198,25 @@ func responseAli2OpenAI(response *ChatResponse) *openai.TextResponse { | |||||||
| 	return &fullTextResponse | 	return &fullTextResponse | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func streamResponseAli2OpenAI(aliResponse *ChatResponse) *openai.ChatCompletionsStreamResponse { | func streamResponseAli2OpenAI(aliResponse *AliChatResponse) *ChatCompletionsStreamResponse { | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice ChatCompletionsStreamResponseChoice | ||||||
| 	choice.Delta.Content = aliResponse.Output.Text | 	choice.Delta.Content = aliResponse.Output.Text | ||||||
| 	if aliResponse.Output.FinishReason != "null" { | 	if aliResponse.Output.FinishReason != "null" { | ||||||
| 		finishReason := aliResponse.Output.FinishReason | 		finishReason := aliResponse.Output.FinishReason | ||||||
| 		choice.FinishReason = &finishReason | 		choice.FinishReason = &finishReason | ||||||
| 	} | 	} | ||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := ChatCompletionsStreamResponse{ | ||||||
| 		Id:      aliResponse.RequestId, | 		Id:      aliResponse.RequestId, | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Model:   "qwen", | 		Model:   "qwen", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| 	return &response | 	return &response | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func aliStreamHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var usage model.Usage | 	var usage Usage | ||||||
| 	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 { | ||||||
| @@ -187,15 +246,15 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 		} | 		} | ||||||
| 		stopChan <- true | 		stopChan <- true | ||||||
| 	}() | 	}() | ||||||
| 	common.SetEventStreamHeaders(c) | 	setEventStreamHeaders(c) | ||||||
| 	//lastResponseText := "" | 	//lastResponseText := "" | ||||||
| 	c.Stream(func(w io.Writer) bool { | 	c.Stream(func(w io.Writer) bool { | ||||||
| 		select { | 		select { | ||||||
| 		case data := <-dataChan: | 		case data := <-dataChan: | ||||||
| 			var aliResponse ChatResponse | 			var aliResponse AliChatResponse | ||||||
| 			err := json.Unmarshal([]byte(data), &aliResponse) | 			err := json.Unmarshal([]byte(data), &aliResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error unmarshalling stream response: " + err.Error()) | 				common.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			if aliResponse.Usage.OutputTokens != 0 { | 			if aliResponse.Usage.OutputTokens != 0 { | ||||||
| @@ -208,7 +267,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 			//lastResponseText = aliResponse.Output.Text | 			//lastResponseText = aliResponse.Output.Text | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error marshalling stream response: " + err.Error()) | 				common.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)}) | ||||||
| @@ -220,28 +279,28 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 	}) | 	}) | ||||||
| 	err := resp.Body.Close() | 	err := resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	return nil, &usage | 	return nil, &usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Handler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func aliHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var aliResponse ChatResponse | 	var aliResponse AliChatResponse | ||||||
| 	responseBody, err := io.ReadAll(resp.Body) | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = json.Unmarshal(responseBody, &aliResponse) | 	err = json.Unmarshal(responseBody, &aliResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	if aliResponse.Code != "" { | 	if aliResponse.Code != "" { | ||||||
| 		return &model.ErrorWithStatusCode{ | 		return &OpenAIErrorWithStatusCode{ | ||||||
| 			Error: model.Error{ | 			OpenAIError: OpenAIError{ | ||||||
| 				Message: aliResponse.Message, | 				Message: aliResponse.Message, | ||||||
| 				Type:    aliResponse.Code, | 				Type:    aliResponse.Code, | ||||||
| 				Param:   aliResponse.RequestId, | 				Param:   aliResponse.RequestId, | ||||||
| @@ -254,7 +313,7 @@ func Handler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, * | |||||||
| 	fullTextResponse.Model = "qwen" | 	fullTextResponse.Model = "qwen" | ||||||
| 	jsonResponse, err := json.Marshal(fullTextResponse) | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
| 	c.Writer.WriteHeader(resp.StatusCode) | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
| @@ -8,20 +8,14 @@ import ( | |||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	relaymodel "github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
| 
 | 
 | ||||||
| func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatusCode { | func relayAudioHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode { | ||||||
| 	audioModel := "whisper-1" | 	audioModel := "whisper-1" | ||||||
| 
 | 
 | ||||||
| 	tokenId := c.GetInt("token_id") | 	tokenId := c.GetInt("token_id") | ||||||
| @@ -31,18 +25,18 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatus | |||||||
| 	group := c.GetString("group") | 	group := c.GetString("group") | ||||||
| 	tokenName := c.GetString("token_name") | 	tokenName := c.GetString("token_name") | ||||||
| 
 | 
 | ||||||
| 	var ttsRequest openai.TextToSpeechRequest | 	var ttsRequest TextToSpeechRequest | ||||||
| 	if relayMode == constant.RelayModeAudioSpeech { | 	if relayMode == RelayModeAudioSpeech { | ||||||
| 		// Read JSON | 		// Read JSON | ||||||
| 		err := common.UnmarshalBodyReusable(c, &ttsRequest) | 		err := common.UnmarshalBodyReusable(c, &ttsRequest) | ||||||
| 		// Check if JSON is valid | 		// Check if JSON is valid | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return openai.ErrorWrapper(err, "invalid_json", http.StatusBadRequest) | 			return errorWrapper(err, "invalid_json", http.StatusBadRequest) | ||||||
| 		} | 		} | ||||||
| 		audioModel = ttsRequest.Model | 		audioModel = ttsRequest.Model | ||||||
| 		// Check if text is too long 4096 | 		// Check if text is too long 4096 | ||||||
| 		if len(ttsRequest.Input) > 4096 { | 		if len(ttsRequest.Input) > 4096 { | ||||||
| 			return openai.ErrorWrapper(errors.New("input is too long (over 4096 characters)"), "text_too_long", http.StatusBadRequest) | 			return errorWrapper(errors.New("input is too long (over 4096 characters)"), "text_too_long", http.StatusBadRequest) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| @@ -52,24 +46,24 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatus | |||||||
| 	var quota int | 	var quota int | ||||||
| 	var preConsumedQuota int | 	var preConsumedQuota int | ||||||
| 	switch relayMode { | 	switch relayMode { | ||||||
| 	case constant.RelayModeAudioSpeech: | 	case RelayModeAudioSpeech: | ||||||
| 		preConsumedQuota = int(float64(len(ttsRequest.Input)) * ratio) | 		preConsumedQuota = int(float64(len(ttsRequest.Input)) * ratio) | ||||||
| 		quota = preConsumedQuota | 		quota = preConsumedQuota | ||||||
| 	default: | 	default: | ||||||
| 		preConsumedQuota = int(float64(config.PreConsumedQuota) * ratio) | 		preConsumedQuota = int(float64(common.PreConsumedQuota) * ratio) | ||||||
| 	} | 	} | ||||||
| 	userQuota, err := model.CacheGetUserQuota(userId) | 	userQuota, err := model.CacheGetUserQuota(userId) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "get_user_quota_failed", http.StatusInternalServerError) | 		return errorWrapper(err, "get_user_quota_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	// Check if user quota is enough | 	// Check if user quota is enough | ||||||
| 	if userQuota-preConsumedQuota < 0 { | 	if userQuota-preConsumedQuota < 0 { | ||||||
| 		return openai.ErrorWrapper(errors.New("user quota is not enough"), "insufficient_user_quota", http.StatusForbidden) | 		return errorWrapper(errors.New("user quota is not enough"), "insufficient_user_quota", http.StatusForbidden) | ||||||
| 	} | 	} | ||||||
| 	err = model.CacheDecreaseUserQuota(userId, preConsumedQuota) | 	err = model.CacheDecreaseUserQuota(userId, preConsumedQuota) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "decrease_user_quota_failed", http.StatusInternalServerError) | 		return errorWrapper(err, "decrease_user_quota_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 	if userQuota > 100*preConsumedQuota { | 	if userQuota > 100*preConsumedQuota { | ||||||
| 		// in this case, we do not pre-consume quota | 		// in this case, we do not pre-consume quota | ||||||
| @@ -79,7 +73,7 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatus | |||||||
| 	if preConsumedQuota > 0 { | 	if preConsumedQuota > 0 { | ||||||
| 		err := model.PreConsumeTokenQuota(tokenId, preConsumedQuota) | 		err := model.PreConsumeTokenQuota(tokenId, preConsumedQuota) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return openai.ErrorWrapper(err, "pre_consume_token_quota_failed", http.StatusForbidden) | 			return errorWrapper(err, "pre_consume_token_quota_failed", http.StatusForbidden) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| @@ -89,7 +83,7 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatus | |||||||
| 		modelMap := make(map[string]string) | 		modelMap := make(map[string]string) | ||||||
| 		err := json.Unmarshal([]byte(modelMapping), &modelMap) | 		err := json.Unmarshal([]byte(modelMapping), &modelMap) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return openai.ErrorWrapper(err, "unmarshal_model_mapping_failed", http.StatusInternalServerError) | 			return errorWrapper(err, "unmarshal_model_mapping_failed", http.StatusInternalServerError) | ||||||
| 		} | 		} | ||||||
| 		if modelMap[audioModel] != "" { | 		if modelMap[audioModel] != "" { | ||||||
| 			audioModel = modelMap[audioModel] | 			audioModel = modelMap[audioModel] | ||||||
| @@ -102,27 +96,27 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatus | |||||||
| 		baseURL = c.GetString("base_url") | 		baseURL = c.GetString("base_url") | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	fullRequestURL := util.GetFullRequestURL(baseURL, requestURL, channelType) | 	fullRequestURL := getFullRequestURL(baseURL, requestURL, channelType) | ||||||
| 	if relayMode == constant.RelayModeAudioTranscription && channelType == common.ChannelTypeAzure { | 	if relayMode == 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.GetAzureAPIVersion(c) | 		apiVersion := GetAPIVersion(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) | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	requestBody := &bytes.Buffer{} | 	requestBody := &bytes.Buffer{} | ||||||
| 	_, err = io.Copy(requestBody, c.Request.Body) | 	_, err = io.Copy(requestBody, c.Request.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "new_request_body_failed", http.StatusInternalServerError) | 		return errorWrapper(err, "new_request_body_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 	c.Request.Body = io.NopCloser(bytes.NewBuffer(requestBody.Bytes())) | 	c.Request.Body = io.NopCloser(bytes.NewBuffer(requestBody.Bytes())) | ||||||
| 	responseFormat := c.DefaultPostForm("response_format", "json") | 	responseFormat := c.DefaultPostForm("response_format", "json") | ||||||
| 
 | 
 | ||||||
| 	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 errorWrapper(err, "new_request_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	if relayMode == constant.RelayModeAudioTranscription && channelType == common.ChannelTypeAzure { | 	if relayMode == 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 | ||||||
| 		apiKey := c.Request.Header.Get("Authorization") | 		apiKey := c.Request.Header.Get("Authorization") | ||||||
| 		apiKey = strings.TrimPrefix(apiKey, "Bearer ") | 		apiKey = strings.TrimPrefix(apiKey, "Bearer ") | ||||||
| @@ -134,34 +128,34 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatus | |||||||
| 	req.Header.Set("Content-Type", c.Request.Header.Get("Content-Type")) | 	req.Header.Set("Content-Type", c.Request.Header.Get("Content-Type")) | ||||||
| 	req.Header.Set("Accept", c.Request.Header.Get("Accept")) | 	req.Header.Set("Accept", c.Request.Header.Get("Accept")) | ||||||
| 
 | 
 | ||||||
| 	resp, err := util.HTTPClient.Do(req) | 	resp, err := httpClient.Do(req) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "do_request_failed", http.StatusInternalServerError) | 		return errorWrapper(err, "do_request_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	err = req.Body.Close() | 	err = req.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_request_body_failed", http.StatusInternalServerError) | 		return errorWrapper(err, "close_request_body_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 	err = c.Request.Body.Close() | 	err = c.Request.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_request_body_failed", http.StatusInternalServerError) | 		return errorWrapper(err, "close_request_body_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	if relayMode != constant.RelayModeAudioSpeech { | 	if relayMode != RelayModeAudioSpeech { | ||||||
| 		responseBody, err := io.ReadAll(resp.Body) | 		responseBody, err := io.ReadAll(resp.Body) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return openai.ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError) | 			return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError) | ||||||
| 		} | 		} | ||||||
| 		err = resp.Body.Close() | 		err = resp.Body.Close() | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError) | 			return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError) | ||||||
| 		} | 		} | ||||||
| 
 | 
 | ||||||
| 		var openAIErr openai.SlimTextResponse | 		var openAIErr TextResponse | ||||||
| 		if err = json.Unmarshal(responseBody, &openAIErr); err == nil { | 		if err = json.Unmarshal(responseBody, &openAIErr); err == nil { | ||||||
| 			if openAIErr.Error.Message != "" { | 			if openAIErr.Error.Message != "" { | ||||||
| 				return openai.ErrorWrapper(fmt.Errorf("type %s, code %v, message %s", openAIErr.Error.Type, openAIErr.Error.Code, openAIErr.Error.Message), "request_error", http.StatusInternalServerError) | 				return errorWrapper(fmt.Errorf("type %s, code %v, message %s", openAIErr.Error.Type, openAIErr.Error.Code, openAIErr.Error.Message), "request_error", http.StatusInternalServerError) | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 
 | 
 | ||||||
| @@ -178,12 +172,12 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatus | |||||||
| 		case "vtt": | 		case "vtt": | ||||||
| 			text, err = getTextFromVTT(responseBody) | 			text, err = getTextFromVTT(responseBody) | ||||||
| 		default: | 		default: | ||||||
| 			return openai.ErrorWrapper(errors.New("unexpected_response_format"), "unexpected_response_format", http.StatusInternalServerError) | 			return errorWrapper(errors.New("unexpected_response_format"), "unexpected_response_format", http.StatusInternalServerError) | ||||||
| 		} | 		} | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return openai.ErrorWrapper(err, "get_text_from_body_err", http.StatusInternalServerError) | 			return errorWrapper(err, "get_text_from_body_err", http.StatusInternalServerError) | ||||||
| 		} | 		} | ||||||
| 		quota = openai.CountTokenText(text, audioModel) | 		quota = countTokenText(text, audioModel) | ||||||
| 		resp.Body = io.NopCloser(bytes.NewBuffer(responseBody)) | 		resp.Body = io.NopCloser(bytes.NewBuffer(responseBody)) | ||||||
| 	} | 	} | ||||||
| 	if resp.StatusCode != http.StatusOK { | 	if resp.StatusCode != http.StatusOK { | ||||||
| @@ -194,16 +188,16 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatus | |||||||
| 					// 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 { | ||||||
| 						logger.Error(ctx, fmt.Sprintf("error rollback pre-consumed quota: %s", err.Error())) | 						common.LogError(ctx, fmt.Sprintf("error rollback pre-consumed quota: %s", err.Error())) | ||||||
| 					} | 					} | ||||||
| 				}() | 				}() | ||||||
| 			}(c.Request.Context()) | 			}(c.Request.Context()) | ||||||
| 		} | 		} | ||||||
| 		return util.RelayErrorHandler(resp) | 		return relayErrorHandler(resp) | ||||||
| 	} | 	} | ||||||
| 	quotaDelta := quota - preConsumedQuota | 	quotaDelta := quota - preConsumedQuota | ||||||
| 	defer func(ctx context.Context) { | 	defer func(ctx context.Context) { | ||||||
| 		go util.PostConsumeQuota(ctx, tokenId, quotaDelta, quota, userId, channelId, modelRatio, groupRatio, audioModel, tokenName) | 		go postConsumeQuota(ctx, tokenId, quotaDelta, quota, userId, channelId, modelRatio, groupRatio, audioModel, tokenName) | ||||||
| 	}(c.Request.Context()) | 	}(c.Request.Context()) | ||||||
| 
 | 
 | ||||||
| 	for k, v := range resp.Header { | 	for k, v := range resp.Header { | ||||||
| @@ -213,11 +207,11 @@ func RelayAudioHelper(c *gin.Context, relayMode int) *relaymodel.ErrorWithStatus | |||||||
| 
 | 
 | ||||||
| 	_, err = io.Copy(c.Writer, resp.Body) | 	_, err = io.Copy(c.Writer, resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "copy_response_body_failed", http.StatusInternalServerError) | 		return errorWrapper(err, "copy_response_body_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError) | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError) | ||||||
| 	} | 	} | ||||||
| 	return nil | 	return nil | ||||||
| } | } | ||||||
| @@ -227,7 +221,7 @@ func getTextFromVTT(body []byte) (string, error) { | |||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func getTextFromVerboseJSON(body []byte) (string, error) { | func getTextFromVerboseJSON(body []byte) (string, error) { | ||||||
| 	var whisperResponse openai.WhisperVerboseJSONResponse | 	var whisperResponse WhisperVerboseJSONResponse | ||||||
| 	if err := json.Unmarshal(body, &whisperResponse); err != nil { | 	if err := json.Unmarshal(body, &whisperResponse); err != nil { | ||||||
| 		return "", fmt.Errorf("unmarshal_response_body_failed err :%w", err) | 		return "", fmt.Errorf("unmarshal_response_body_failed err :%w", err) | ||||||
| 	} | 	} | ||||||
| @@ -260,7 +254,7 @@ func getTextFromText(body []byte) (string, error) { | |||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func getTextFromJSON(body []byte) (string, error) { | func getTextFromJSON(body []byte) (string, error) { | ||||||
| 	var whisperResponse openai.WhisperJSONResponse | 	var whisperResponse WhisperJSONResponse | ||||||
| 	if err := json.Unmarshal(body, &whisperResponse); err != nil { | 	if err := json.Unmarshal(body, &whisperResponse); err != nil { | ||||||
| 		return "", fmt.Errorf("unmarshal_response_body_failed err :%w", err) | 		return "", fmt.Errorf("unmarshal_response_body_failed err :%w", err) | ||||||
| 	} | 	} | ||||||
| @@ -1,4 +1,4 @@ | |||||||
| package baidu | package controller | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"bufio" | 	"bufio" | ||||||
| @@ -6,14 +6,9 @@ import ( | |||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"sync" | 	"sync" | ||||||
| 	"time" | 	"time" | ||||||
| @@ -21,104 +16,148 @@ import ( | |||||||
| 
 | 
 | ||||||
| // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/flfmc9do2 | // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/flfmc9do2 | ||||||
| 
 | 
 | ||||||
| type TokenResponse struct { | type BaiduTokenResponse struct { | ||||||
| 	ExpiresIn   int    `json:"expires_in"` | 	ExpiresIn   int    `json:"expires_in"` | ||||||
| 	AccessToken string `json:"access_token"` | 	AccessToken string `json:"access_token"` | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| type Message struct { | type BaiduMessage struct { | ||||||
| 	Role    string `json:"role"` | 	Role    string `json:"role"` | ||||||
| 	Content string `json:"content"` | 	Content string `json:"content"` | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| type ChatRequest struct { | type BaiduChatRequest struct { | ||||||
| 	Messages []Message `json:"messages"` | 	Messages []BaiduMessage `json:"messages"` | ||||||
| 	Stream   bool      `json:"stream"` | 	Stream   bool           `json:"stream"` | ||||||
| 	UserId   string    `json:"user_id,omitempty"` | 	UserId   string         `json:"user_id,omitempty"` | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| type Error struct { | type BaiduError struct { | ||||||
| 	ErrorCode int    `json:"error_code"` | 	ErrorCode int    `json:"error_code"` | ||||||
| 	ErrorMsg  string `json:"error_msg"` | 	ErrorMsg  string `json:"error_msg"` | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
|  | type BaiduChatResponse struct { | ||||||
|  | 	Id               string `json:"id"` | ||||||
|  | 	Object           string `json:"object"` | ||||||
|  | 	Created          int64  `json:"created"` | ||||||
|  | 	Result           string `json:"result"` | ||||||
|  | 	IsTruncated      bool   `json:"is_truncated"` | ||||||
|  | 	NeedClearHistory bool   `json:"need_clear_history"` | ||||||
|  | 	Usage            Usage  `json:"usage"` | ||||||
|  | 	BaiduError | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type BaiduChatStreamResponse struct { | ||||||
|  | 	BaiduChatResponse | ||||||
|  | 	SentenceId int  `json:"sentence_id"` | ||||||
|  | 	IsEnd      bool `json:"is_end"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type BaiduEmbeddingRequest struct { | ||||||
|  | 	Input []string `json:"input"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type BaiduEmbeddingData struct { | ||||||
|  | 	Object    string    `json:"object"` | ||||||
|  | 	Embedding []float64 `json:"embedding"` | ||||||
|  | 	Index     int       `json:"index"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type BaiduEmbeddingResponse struct { | ||||||
|  | 	Id      string               `json:"id"` | ||||||
|  | 	Object  string               `json:"object"` | ||||||
|  | 	Created int64                `json:"created"` | ||||||
|  | 	Data    []BaiduEmbeddingData `json:"data"` | ||||||
|  | 	Usage   Usage                `json:"usage"` | ||||||
|  | 	BaiduError | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type BaiduAccessToken struct { | ||||||
|  | 	AccessToken      string    `json:"access_token"` | ||||||
|  | 	Error            string    `json:"error,omitempty"` | ||||||
|  | 	ErrorDescription string    `json:"error_description,omitempty"` | ||||||
|  | 	ExpiresIn        int64     `json:"expires_in,omitempty"` | ||||||
|  | 	ExpiresAt        time.Time `json:"-"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
| var baiduTokenStore sync.Map | var baiduTokenStore sync.Map | ||||||
| 
 | 
 | ||||||
| func ConvertRequest(request model.GeneralOpenAIRequest) *ChatRequest { | func requestOpenAI2Baidu(request GeneralOpenAIRequest) *BaiduChatRequest { | ||||||
| 	messages := make([]Message, 0, len(request.Messages)) | 	messages := make([]BaiduMessage, 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, Message{ | 			messages = append(messages, BaiduMessage{ | ||||||
| 				Role:    "user", | 				Role:    "user", | ||||||
| 				Content: message.StringContent(), | 				Content: message.StringContent(), | ||||||
| 			}) | 			}) | ||||||
| 			messages = append(messages, Message{ | 			messages = append(messages, BaiduMessage{ | ||||||
| 				Role:    "assistant", | 				Role:    "assistant", | ||||||
| 				Content: "Okay", | 				Content: "Okay", | ||||||
| 			}) | 			}) | ||||||
| 		} else { | 		} else { | ||||||
| 			messages = append(messages, Message{ | 			messages = append(messages, BaiduMessage{ | ||||||
| 				Role:    message.Role, | 				Role:    message.Role, | ||||||
| 				Content: message.StringContent(), | 				Content: message.StringContent(), | ||||||
| 			}) | 			}) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	return &ChatRequest{ | 	return &BaiduChatRequest{ | ||||||
| 		Messages: messages, | 		Messages: messages, | ||||||
| 		Stream:   request.Stream, | 		Stream:   request.Stream, | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func responseBaidu2OpenAI(response *ChatResponse) *openai.TextResponse { | func responseBaidu2OpenAI(response *BaiduChatResponse) *OpenAITextResponse { | ||||||
| 	choice := openai.TextResponseChoice{ | 	choice := OpenAITextResponseChoice{ | ||||||
| 		Index: 0, | 		Index: 0, | ||||||
| 		Message: model.Message{ | 		Message: Message{ | ||||||
| 			Role:    "assistant", | 			Role:    "assistant", | ||||||
| 			Content: response.Result, | 			Content: response.Result, | ||||||
| 		}, | 		}, | ||||||
| 		FinishReason: "stop", | 		FinishReason: "stop", | ||||||
| 	} | 	} | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := OpenAITextResponse{ | ||||||
| 		Id:      response.Id, | 		Id:      response.Id, | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: response.Created, | 		Created: response.Created, | ||||||
| 		Choices: []openai.TextResponseChoice{choice}, | 		Choices: []OpenAITextResponseChoice{choice}, | ||||||
| 		Usage:   response.Usage, | 		Usage:   response.Usage, | ||||||
| 	} | 	} | ||||||
| 	return &fullTextResponse | 	return &fullTextResponse | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func streamResponseBaidu2OpenAI(baiduResponse *ChatStreamResponse) *openai.ChatCompletionsStreamResponse { | func streamResponseBaidu2OpenAI(baiduResponse *BaiduChatStreamResponse) *ChatCompletionsStreamResponse { | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice ChatCompletionsStreamResponseChoice | ||||||
| 	choice.Delta.Content = baiduResponse.Result | 	choice.Delta.Content = baiduResponse.Result | ||||||
| 	if baiduResponse.IsEnd { | 	if baiduResponse.IsEnd { | ||||||
| 		choice.FinishReason = &constant.StopFinishReason | 		choice.FinishReason = &stopFinishReason | ||||||
| 	} | 	} | ||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := ChatCompletionsStreamResponse{ | ||||||
| 		Id:      baiduResponse.Id, | 		Id:      baiduResponse.Id, | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: baiduResponse.Created, | 		Created: baiduResponse.Created, | ||||||
| 		Model:   "ernie-bot", | 		Model:   "ernie-bot", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| 	return &response | 	return &response | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func ConvertEmbeddingRequest(request model.GeneralOpenAIRequest) *EmbeddingRequest { | func embeddingRequestOpenAI2Baidu(request GeneralOpenAIRequest) *BaiduEmbeddingRequest { | ||||||
| 	return &EmbeddingRequest{ | 	return &BaiduEmbeddingRequest{ | ||||||
| 		Input: request.ParseInput(), | 		Input: request.ParseInput(), | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func embeddingResponseBaidu2OpenAI(response *EmbeddingResponse) *openai.EmbeddingResponse { | func embeddingResponseBaidu2OpenAI(response *BaiduEmbeddingResponse) *OpenAIEmbeddingResponse { | ||||||
| 	openAIEmbeddingResponse := openai.EmbeddingResponse{ | 	openAIEmbeddingResponse := OpenAIEmbeddingResponse{ | ||||||
| 		Object: "list", | 		Object: "list", | ||||||
| 		Data:   make([]openai.EmbeddingResponseItem, 0, len(response.Data)), | 		Data:   make([]OpenAIEmbeddingResponseItem, 0, len(response.Data)), | ||||||
| 		Model:  "baidu-embedding", | 		Model:  "baidu-embedding", | ||||||
| 		Usage:  response.Usage, | 		Usage:  response.Usage, | ||||||
| 	} | 	} | ||||||
| 	for _, item := range response.Data { | 	for _, item := range response.Data { | ||||||
| 		openAIEmbeddingResponse.Data = append(openAIEmbeddingResponse.Data, openai.EmbeddingResponseItem{ | 		openAIEmbeddingResponse.Data = append(openAIEmbeddingResponse.Data, OpenAIEmbeddingResponseItem{ | ||||||
| 			Object:    item.Object, | 			Object:    item.Object, | ||||||
| 			Index:     item.Index, | 			Index:     item.Index, | ||||||
| 			Embedding: item.Embedding, | 			Embedding: item.Embedding, | ||||||
| @@ -127,8 +166,8 @@ func embeddingResponseBaidu2OpenAI(response *EmbeddingResponse) *openai.Embeddin | |||||||
| 	return &openAIEmbeddingResponse | 	return &openAIEmbeddingResponse | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func baiduStreamHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var usage model.Usage | 	var usage Usage | ||||||
| 	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 { | ||||||
| @@ -155,14 +194,14 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 		} | 		} | ||||||
| 		stopChan <- true | 		stopChan <- true | ||||||
| 	}() | 	}() | ||||||
| 	common.SetEventStreamHeaders(c) | 	setEventStreamHeaders(c) | ||||||
| 	c.Stream(func(w io.Writer) bool { | 	c.Stream(func(w io.Writer) bool { | ||||||
| 		select { | 		select { | ||||||
| 		case data := <-dataChan: | 		case data := <-dataChan: | ||||||
| 			var baiduResponse ChatStreamResponse | 			var baiduResponse BaiduChatStreamResponse | ||||||
| 			err := json.Unmarshal([]byte(data), &baiduResponse) | 			err := json.Unmarshal([]byte(data), &baiduResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error unmarshalling stream response: " + err.Error()) | 				common.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			if baiduResponse.Usage.TotalTokens != 0 { | 			if baiduResponse.Usage.TotalTokens != 0 { | ||||||
| @@ -173,7 +212,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 			response := streamResponseBaidu2OpenAI(&baiduResponse) | 			response := streamResponseBaidu2OpenAI(&baiduResponse) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error marshalling stream response: " + err.Error()) | 				common.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)}) | ||||||
| @@ -185,28 +224,28 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 	}) | 	}) | ||||||
| 	err := resp.Body.Close() | 	err := resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	return nil, &usage | 	return nil, &usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Handler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func baiduHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var baiduResponse ChatResponse | 	var baiduResponse BaiduChatResponse | ||||||
| 	responseBody, err := io.ReadAll(resp.Body) | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = json.Unmarshal(responseBody, &baiduResponse) | 	err = json.Unmarshal(responseBody, &baiduResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	if baiduResponse.ErrorMsg != "" { | 	if baiduResponse.ErrorMsg != "" { | ||||||
| 		return &model.ErrorWithStatusCode{ | 		return &OpenAIErrorWithStatusCode{ | ||||||
| 			Error: model.Error{ | 			OpenAIError: OpenAIError{ | ||||||
| 				Message: baiduResponse.ErrorMsg, | 				Message: baiduResponse.ErrorMsg, | ||||||
| 				Type:    "baidu_error", | 				Type:    "baidu_error", | ||||||
| 				Param:   "", | 				Param:   "", | ||||||
| @@ -219,7 +258,7 @@ func Handler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, * | |||||||
| 	fullTextResponse.Model = "ernie-bot" | 	fullTextResponse.Model = "ernie-bot" | ||||||
| 	jsonResponse, err := json.Marshal(fullTextResponse) | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
| 	c.Writer.WriteHeader(resp.StatusCode) | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
| @@ -227,23 +266,23 @@ func Handler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, * | |||||||
| 	return nil, &fullTextResponse.Usage | 	return nil, &fullTextResponse.Usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func EmbeddingHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func baiduEmbeddingHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var baiduResponse EmbeddingResponse | 	var baiduResponse BaiduEmbeddingResponse | ||||||
| 	responseBody, err := io.ReadAll(resp.Body) | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = json.Unmarshal(responseBody, &baiduResponse) | 	err = json.Unmarshal(responseBody, &baiduResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	if baiduResponse.ErrorMsg != "" { | 	if baiduResponse.ErrorMsg != "" { | ||||||
| 		return &model.ErrorWithStatusCode{ | 		return &OpenAIErrorWithStatusCode{ | ||||||
| 			Error: model.Error{ | 			OpenAIError: OpenAIError{ | ||||||
| 				Message: baiduResponse.ErrorMsg, | 				Message: baiduResponse.ErrorMsg, | ||||||
| 				Type:    "baidu_error", | 				Type:    "baidu_error", | ||||||
| 				Param:   "", | 				Param:   "", | ||||||
| @@ -255,7 +294,7 @@ func EmbeddingHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStat | |||||||
| 	fullTextResponse := embeddingResponseBaidu2OpenAI(&baiduResponse) | 	fullTextResponse := embeddingResponseBaidu2OpenAI(&baiduResponse) | ||||||
| 	jsonResponse, err := json.Marshal(fullTextResponse) | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
| 	c.Writer.WriteHeader(resp.StatusCode) | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
| @@ -263,10 +302,10 @@ func EmbeddingHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStat | |||||||
| 	return nil, &fullTextResponse.Usage | 	return nil, &fullTextResponse.Usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func GetAccessToken(apiKey string) (string, error) { | func getBaiduAccessToken(apiKey string) (string, error) { | ||||||
| 	if val, ok := baiduTokenStore.Load(apiKey); ok { | 	if val, ok := baiduTokenStore.Load(apiKey); ok { | ||||||
| 		var accessToken AccessToken | 		var accessToken BaiduAccessToken | ||||||
| 		if accessToken, ok = val.(AccessToken); ok { | 		if accessToken, ok = val.(BaiduAccessToken); ok { | ||||||
| 			// soon this will expire | 			// soon this will expire | ||||||
| 			if time.Now().Add(time.Hour).After(accessToken.ExpiresAt) { | 			if time.Now().Add(time.Hour).After(accessToken.ExpiresAt) { | ||||||
| 				go func() { | 				go func() { | ||||||
| @@ -281,12 +320,12 @@ func GetAccessToken(apiKey string) (string, error) { | |||||||
| 		return "", err | 		return "", err | ||||||
| 	} | 	} | ||||||
| 	if accessToken == nil { | 	if accessToken == nil { | ||||||
| 		return "", errors.New("GetAccessToken return a nil token") | 		return "", errors.New("getBaiduAccessToken return a nil token") | ||||||
| 	} | 	} | ||||||
| 	return (*accessToken).AccessToken, nil | 	return (*accessToken).AccessToken, nil | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func getBaiduAccessTokenHelper(apiKey string) (*AccessToken, error) { | func getBaiduAccessTokenHelper(apiKey string) (*BaiduAccessToken, error) { | ||||||
| 	parts := strings.Split(apiKey, "|") | 	parts := strings.Split(apiKey, "|") | ||||||
| 	if len(parts) != 2 { | 	if len(parts) != 2 { | ||||||
| 		return nil, errors.New("invalid baidu apikey") | 		return nil, errors.New("invalid baidu apikey") | ||||||
| @@ -298,13 +337,13 @@ func getBaiduAccessTokenHelper(apiKey string) (*AccessToken, error) { | |||||||
| 	} | 	} | ||||||
| 	req.Header.Add("Content-Type", "application/json") | 	req.Header.Add("Content-Type", "application/json") | ||||||
| 	req.Header.Add("Accept", "application/json") | 	req.Header.Add("Accept", "application/json") | ||||||
| 	res, err := util.ImpatientHTTPClient.Do(req) | 	res, err := impatientHTTPClient.Do(req) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return nil, err | 		return nil, err | ||||||
| 	} | 	} | ||||||
| 	defer res.Body.Close() | 	defer res.Body.Close() | ||||||
| 
 | 
 | ||||||
| 	var accessToken AccessToken | 	var accessToken BaiduAccessToken | ||||||
| 	err = json.NewDecoder(res.Body).Decode(&accessToken) | 	err = json.NewDecoder(res.Body).Decode(&accessToken) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return nil, err | 		return nil, err | ||||||
							
								
								
									
										223
									
								
								controller/relay-claude.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										223
									
								
								controller/relay-claude.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,223 @@ | |||||||
|  | package controller | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"bufio" | ||||||
|  | 	"encoding/json" | ||||||
|  | 	"fmt" | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"io" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"strings" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type ClaudeMetadata struct { | ||||||
|  | 	UserId string `json:"user_id"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type ClaudeRequest struct { | ||||||
|  | 	Model             string   `json:"model"` | ||||||
|  | 	Prompt            string   `json:"prompt"` | ||||||
|  | 	MaxTokensToSample int      `json:"max_tokens_to_sample"` | ||||||
|  | 	StopSequences     []string `json:"stop_sequences,omitempty"` | ||||||
|  | 	Temperature       float64  `json:"temperature,omitempty"` | ||||||
|  | 	TopP              float64  `json:"top_p,omitempty"` | ||||||
|  | 	TopK              int      `json:"top_k,omitempty"` | ||||||
|  | 	//ClaudeMetadata    `json:"metadata,omitempty"` | ||||||
|  | 	Stream bool `json:"stream,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type ClaudeError struct { | ||||||
|  | 	Type    string `json:"type"` | ||||||
|  | 	Message string `json:"message"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type ClaudeResponse struct { | ||||||
|  | 	Completion string      `json:"completion"` | ||||||
|  | 	StopReason string      `json:"stop_reason"` | ||||||
|  | 	Model      string      `json:"model"` | ||||||
|  | 	Error      ClaudeError `json:"error"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func stopReasonClaude2OpenAI(reason string) string { | ||||||
|  | 	switch reason { | ||||||
|  | 	case "stop_sequence": | ||||||
|  | 		return "stop" | ||||||
|  | 	case "max_tokens": | ||||||
|  | 		return "length" | ||||||
|  | 	default: | ||||||
|  | 		return reason | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func requestOpenAI2Claude(textRequest GeneralOpenAIRequest) *ClaudeRequest { | ||||||
|  | 	claudeRequest := ClaudeRequest{ | ||||||
|  | 		Model:             textRequest.Model, | ||||||
|  | 		Prompt:            "", | ||||||
|  | 		MaxTokensToSample: textRequest.MaxTokens, | ||||||
|  | 		StopSequences:     nil, | ||||||
|  | 		Temperature:       textRequest.Temperature, | ||||||
|  | 		TopP:              textRequest.TopP, | ||||||
|  | 		Stream:            textRequest.Stream, | ||||||
|  | 	} | ||||||
|  | 	if claudeRequest.MaxTokensToSample == 0 { | ||||||
|  | 		claudeRequest.MaxTokensToSample = 1000000 | ||||||
|  | 	} | ||||||
|  | 	prompt := "" | ||||||
|  | 	for _, message := range textRequest.Messages { | ||||||
|  | 		if message.Role == "user" { | ||||||
|  | 			prompt += fmt.Sprintf("\n\nHuman: %s", message.Content) | ||||||
|  | 		} else if message.Role == "assistant" { | ||||||
|  | 			prompt += fmt.Sprintf("\n\nAssistant: %s", message.Content) | ||||||
|  | 		} else if message.Role == "system" { | ||||||
|  | 			if prompt == "" { | ||||||
|  | 				prompt = message.StringContent() | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	prompt += "\n\nAssistant:" | ||||||
|  | 	claudeRequest.Prompt = prompt | ||||||
|  | 	return &claudeRequest | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func streamResponseClaude2OpenAI(claudeResponse *ClaudeResponse) *ChatCompletionsStreamResponse { | ||||||
|  | 	var choice ChatCompletionsStreamResponseChoice | ||||||
|  | 	choice.Delta.Content = claudeResponse.Completion | ||||||
|  | 	finishReason := stopReasonClaude2OpenAI(claudeResponse.StopReason) | ||||||
|  | 	if finishReason != "null" { | ||||||
|  | 		choice.FinishReason = &finishReason | ||||||
|  | 	} | ||||||
|  | 	var response ChatCompletionsStreamResponse | ||||||
|  | 	response.Object = "chat.completion.chunk" | ||||||
|  | 	response.Model = claudeResponse.Model | ||||||
|  | 	response.Choices = []ChatCompletionsStreamResponseChoice{choice} | ||||||
|  | 	return &response | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func responseClaude2OpenAI(claudeResponse *ClaudeResponse) *OpenAITextResponse { | ||||||
|  | 	choice := OpenAITextResponseChoice{ | ||||||
|  | 		Index: 0, | ||||||
|  | 		Message: Message{ | ||||||
|  | 			Role:    "assistant", | ||||||
|  | 			Content: strings.TrimPrefix(claudeResponse.Completion, " "), | ||||||
|  | 			Name:    nil, | ||||||
|  | 		}, | ||||||
|  | 		FinishReason: stopReasonClaude2OpenAI(claudeResponse.StopReason), | ||||||
|  | 	} | ||||||
|  | 	fullTextResponse := OpenAITextResponse{ | ||||||
|  | 		Id:      fmt.Sprintf("chatcmpl-%s", common.GetUUID()), | ||||||
|  | 		Object:  "chat.completion", | ||||||
|  | 		Created: common.GetTimestamp(), | ||||||
|  | 		Choices: []OpenAITextResponseChoice{choice}, | ||||||
|  | 	} | ||||||
|  | 	return &fullTextResponse | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func claudeStreamHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, string) { | ||||||
|  | 	responseText := "" | ||||||
|  | 	responseId := fmt.Sprintf("chatcmpl-%s", common.GetUUID()) | ||||||
|  | 	createdTime := common.GetTimestamp() | ||||||
|  | 	scanner := bufio.NewScanner(resp.Body) | ||||||
|  | 	scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) { | ||||||
|  | 		if atEOF && len(data) == 0 { | ||||||
|  | 			return 0, nil, nil | ||||||
|  | 		} | ||||||
|  | 		if i := strings.Index(string(data), "\r\n\r\n"); i >= 0 { | ||||||
|  | 			return i + 4, data[0:i], nil | ||||||
|  | 		} | ||||||
|  | 		if atEOF { | ||||||
|  | 			return len(data), data, nil | ||||||
|  | 		} | ||||||
|  | 		return 0, nil, nil | ||||||
|  | 	}) | ||||||
|  | 	dataChan := make(chan string) | ||||||
|  | 	stopChan := make(chan bool) | ||||||
|  | 	go func() { | ||||||
|  | 		for scanner.Scan() { | ||||||
|  | 			data := scanner.Text() | ||||||
|  | 			if !strings.HasPrefix(data, "event: completion") { | ||||||
|  | 				continue | ||||||
|  | 			} | ||||||
|  | 			data = strings.TrimPrefix(data, "event: completion\r\ndata: ") | ||||||
|  | 			dataChan <- data | ||||||
|  | 		} | ||||||
|  | 		stopChan <- true | ||||||
|  | 	}() | ||||||
|  | 	setEventStreamHeaders(c) | ||||||
|  | 	c.Stream(func(w io.Writer) bool { | ||||||
|  | 		select { | ||||||
|  | 		case data := <-dataChan: | ||||||
|  | 			// some implementations may add \r at the end of data | ||||||
|  | 			data = strings.TrimSuffix(data, "\r") | ||||||
|  | 			var claudeResponse ClaudeResponse | ||||||
|  | 			err := json.Unmarshal([]byte(data), &claudeResponse) | ||||||
|  | 			if err != nil { | ||||||
|  | 				common.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
|  | 				return true | ||||||
|  | 			} | ||||||
|  | 			responseText += claudeResponse.Completion | ||||||
|  | 			response := streamResponseClaude2OpenAI(&claudeResponse) | ||||||
|  | 			response.Id = responseId | ||||||
|  | 			response.Created = createdTime | ||||||
|  | 			jsonStr, err := json.Marshal(response) | ||||||
|  | 			if err != nil { | ||||||
|  | 				common.SysError("error marshalling stream response: " + err.Error()) | ||||||
|  | 				return true | ||||||
|  | 			} | ||||||
|  | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonStr)}) | ||||||
|  | 			return true | ||||||
|  | 		case <-stopChan: | ||||||
|  | 			c.Render(-1, common.CustomEvent{Data: "data: [DONE]"}) | ||||||
|  | 			return false | ||||||
|  | 		} | ||||||
|  | 	}) | ||||||
|  | 	err := resp.Body.Close() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), "" | ||||||
|  | 	} | ||||||
|  | 	return nil, responseText | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func claudeHandler(c *gin.Context, resp *http.Response, promptTokens int, model string) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
|  | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
|  | 	} | ||||||
|  | 	err = resp.Body.Close() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
|  | 	} | ||||||
|  | 	var claudeResponse ClaudeResponse | ||||||
|  | 	err = json.Unmarshal(responseBody, &claudeResponse) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
|  | 	} | ||||||
|  | 	if claudeResponse.Error.Type != "" { | ||||||
|  | 		return &OpenAIErrorWithStatusCode{ | ||||||
|  | 			OpenAIError: OpenAIError{ | ||||||
|  | 				Message: claudeResponse.Error.Message, | ||||||
|  | 				Type:    claudeResponse.Error.Type, | ||||||
|  | 				Param:   "", | ||||||
|  | 				Code:    claudeResponse.Error.Type, | ||||||
|  | 			}, | ||||||
|  | 			StatusCode: resp.StatusCode, | ||||||
|  | 		}, nil | ||||||
|  | 	} | ||||||
|  | 	fullTextResponse := responseClaude2OpenAI(&claudeResponse) | ||||||
|  | 	fullTextResponse.Model = model | ||||||
|  | 	completionTokens := countTokenText(claudeResponse.Completion, model) | ||||||
|  | 	usage := Usage{ | ||||||
|  | 		PromptTokens:     promptTokens, | ||||||
|  | 		CompletionTokens: completionTokens, | ||||||
|  | 		TotalTokens:      promptTokens + completionTokens, | ||||||
|  | 	} | ||||||
|  | 	fullTextResponse.Usage = usage | ||||||
|  | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
|  | 	} | ||||||
|  | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
|  | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
|  | 	_, err = c.Writer.Write(jsonResponse) | ||||||
|  | 	return nil, &usage | ||||||
|  | } | ||||||
| @@ -1,19 +1,13 @@ | |||||||
| package gemini | package controller | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"bufio" | 	"bufio" | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/image" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/common/image" | ||||||
| 	"strings" | 	"strings" | ||||||
| 
 | 
 | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| @@ -22,39 +16,79 @@ import ( | |||||||
| // https://ai.google.dev/docs/gemini_api_overview?hl=zh-cn | // https://ai.google.dev/docs/gemini_api_overview?hl=zh-cn | ||||||
| 
 | 
 | ||||||
| const ( | const ( | ||||||
| 	VisionMaxImageNum = 16 | 	GeminiVisionMaxImageNum = 16 | ||||||
| ) | ) | ||||||
| 
 | 
 | ||||||
|  | type GeminiChatRequest struct { | ||||||
|  | 	Contents         []GeminiChatContent        `json:"contents"` | ||||||
|  | 	SafetySettings   []GeminiChatSafetySettings `json:"safety_settings,omitempty"` | ||||||
|  | 	GenerationConfig GeminiChatGenerationConfig `json:"generation_config,omitempty"` | ||||||
|  | 	Tools            []GeminiChatTools          `json:"tools,omitempty"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type GeminiInlineData struct { | ||||||
|  | 	MimeType string `json:"mimeType"` | ||||||
|  | 	Data     string `json:"data"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type GeminiPart struct { | ||||||
|  | 	Text       string            `json:"text,omitempty"` | ||||||
|  | 	InlineData *GeminiInlineData `json:"inlineData,omitempty"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type GeminiChatContent struct { | ||||||
|  | 	Role  string       `json:"role,omitempty"` | ||||||
|  | 	Parts []GeminiPart `json:"parts"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type GeminiChatSafetySettings struct { | ||||||
|  | 	Category  string `json:"category"` | ||||||
|  | 	Threshold string `json:"threshold"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type GeminiChatTools struct { | ||||||
|  | 	FunctionDeclarations any `json:"functionDeclarations,omitempty"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type GeminiChatGenerationConfig struct { | ||||||
|  | 	Temperature     float64  `json:"temperature,omitempty"` | ||||||
|  | 	TopP            float64  `json:"topP,omitempty"` | ||||||
|  | 	TopK            float64  `json:"topK,omitempty"` | ||||||
|  | 	MaxOutputTokens int      `json:"maxOutputTokens,omitempty"` | ||||||
|  | 	CandidateCount  int      `json:"candidateCount,omitempty"` | ||||||
|  | 	StopSequences   []string `json:"stopSequences,omitempty"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
| // Setting safety to the lowest possible values since Gemini is already powerless enough | // Setting safety to the lowest possible values since Gemini is already powerless enough | ||||||
| func ConvertRequest(textRequest model.GeneralOpenAIRequest) *ChatRequest { | func requestOpenAI2Gemini(textRequest GeneralOpenAIRequest) *GeminiChatRequest { | ||||||
| 	geminiRequest := ChatRequest{ | 	geminiRequest := GeminiChatRequest{ | ||||||
| 		Contents: make([]ChatContent, 0, len(textRequest.Messages)), | 		Contents: make([]GeminiChatContent, 0, len(textRequest.Messages)), | ||||||
| 		SafetySettings: []ChatSafetySettings{ | 		SafetySettings: []GeminiChatSafetySettings{ | ||||||
| 			{ | 			{ | ||||||
| 				Category:  "HARM_CATEGORY_HARASSMENT", | 				Category:  "HARM_CATEGORY_HARASSMENT", | ||||||
| 				Threshold: config.GeminiSafetySetting, | 				Threshold: common.GeminiSafetySetting, | ||||||
| 			}, | 			}, | ||||||
| 			{ | 			{ | ||||||
| 				Category:  "HARM_CATEGORY_HATE_SPEECH", | 				Category:  "HARM_CATEGORY_HATE_SPEECH", | ||||||
| 				Threshold: config.GeminiSafetySetting, | 				Threshold: common.GeminiSafetySetting, | ||||||
| 			}, | 			}, | ||||||
| 			{ | 			{ | ||||||
| 				Category:  "HARM_CATEGORY_SEXUALLY_EXPLICIT", | 				Category:  "HARM_CATEGORY_SEXUALLY_EXPLICIT", | ||||||
| 				Threshold: config.GeminiSafetySetting, | 				Threshold: common.GeminiSafetySetting, | ||||||
| 			}, | 			}, | ||||||
| 			{ | 			{ | ||||||
| 				Category:  "HARM_CATEGORY_DANGEROUS_CONTENT", | 				Category:  "HARM_CATEGORY_DANGEROUS_CONTENT", | ||||||
| 				Threshold: config.GeminiSafetySetting, | 				Threshold: common.GeminiSafetySetting, | ||||||
| 			}, | 			}, | ||||||
| 		}, | 		}, | ||||||
| 		GenerationConfig: ChatGenerationConfig{ | 		GenerationConfig: GeminiChatGenerationConfig{ | ||||||
| 			Temperature:     textRequest.Temperature, | 			Temperature:     textRequest.Temperature, | ||||||
| 			TopP:            textRequest.TopP, | 			TopP:            textRequest.TopP, | ||||||
| 			MaxOutputTokens: textRequest.MaxTokens, | 			MaxOutputTokens: textRequest.MaxTokens, | ||||||
| 		}, | 		}, | ||||||
| 	} | 	} | ||||||
| 	if textRequest.Functions != nil { | 	if textRequest.Functions != nil { | ||||||
| 		geminiRequest.Tools = []ChatTools{ | 		geminiRequest.Tools = []GeminiChatTools{ | ||||||
| 			{ | 			{ | ||||||
| 				FunctionDeclarations: textRequest.Functions, | 				FunctionDeclarations: textRequest.Functions, | ||||||
| 			}, | 			}, | ||||||
| @@ -62,30 +96,30 @@ func ConvertRequest(textRequest model.GeneralOpenAIRequest) *ChatRequest { | |||||||
| 	} | 	} | ||||||
| 	shouldAddDummyModelMessage := false | 	shouldAddDummyModelMessage := false | ||||||
| 	for _, message := range textRequest.Messages { | 	for _, message := range textRequest.Messages { | ||||||
| 		content := ChatContent{ | 		content := GeminiChatContent{ | ||||||
| 			Role: message.Role, | 			Role: message.Role, | ||||||
| 			Parts: []Part{ | 			Parts: []GeminiPart{ | ||||||
| 				{ | 				{ | ||||||
| 					Text: message.StringContent(), | 					Text: message.StringContent(), | ||||||
| 				}, | 				}, | ||||||
| 			}, | 			}, | ||||||
| 		} | 		} | ||||||
| 		openaiContent := message.ParseContent() | 		openaiContent := message.ParseContent() | ||||||
| 		var parts []Part | 		var parts []GeminiPart | ||||||
| 		imageNum := 0 | 		imageNum := 0 | ||||||
| 		for _, part := range openaiContent { | 		for _, part := range openaiContent { | ||||||
| 			if part.Type == model.ContentTypeText { | 			if part.Type == ContentTypeText { | ||||||
| 				parts = append(parts, Part{ | 				parts = append(parts, GeminiPart{ | ||||||
| 					Text: part.Text, | 					Text: part.Text, | ||||||
| 				}) | 				}) | ||||||
| 			} else if part.Type == model.ContentTypeImageURL { | 			} else if part.Type == ContentTypeImageURL { | ||||||
| 				imageNum += 1 | 				imageNum += 1 | ||||||
| 				if imageNum > VisionMaxImageNum { | 				if imageNum > GeminiVisionMaxImageNum { | ||||||
| 					continue | 					continue | ||||||
| 				} | 				} | ||||||
| 				mimeType, data, _ := image.GetImageFromUrl(part.ImageURL.Url) | 				mimeType, data, _ := image.GetImageFromUrl(part.ImageURL.Url) | ||||||
| 				parts = append(parts, Part{ | 				parts = append(parts, GeminiPart{ | ||||||
| 					InlineData: &InlineData{ | 					InlineData: &GeminiInlineData{ | ||||||
| 						MimeType: mimeType, | 						MimeType: mimeType, | ||||||
| 						Data:     data, | 						Data:     data, | ||||||
| 					}, | 					}, | ||||||
| @@ -107,9 +141,9 @@ func ConvertRequest(textRequest model.GeneralOpenAIRequest) *ChatRequest { | |||||||
| 
 | 
 | ||||||
| 		// If a system message is the last message, we need to add a dummy model message to make gemini happy | 		// If a system message is the last message, we need to add a dummy model message to make gemini happy | ||||||
| 		if shouldAddDummyModelMessage { | 		if shouldAddDummyModelMessage { | ||||||
| 			geminiRequest.Contents = append(geminiRequest.Contents, ChatContent{ | 			geminiRequest.Contents = append(geminiRequest.Contents, GeminiChatContent{ | ||||||
| 				Role: "model", | 				Role: "model", | ||||||
| 				Parts: []Part{ | 				Parts: []GeminiPart{ | ||||||
| 					{ | 					{ | ||||||
| 						Text: "Okay", | 						Text: "Okay", | ||||||
| 					}, | 					}, | ||||||
| @@ -122,12 +156,12 @@ func ConvertRequest(textRequest model.GeneralOpenAIRequest) *ChatRequest { | |||||||
| 	return &geminiRequest | 	return &geminiRequest | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| type ChatResponse struct { | type GeminiChatResponse struct { | ||||||
| 	Candidates     []ChatCandidate    `json:"candidates"` | 	Candidates     []GeminiChatCandidate    `json:"candidates"` | ||||||
| 	PromptFeedback ChatPromptFeedback `json:"promptFeedback"` | 	PromptFeedback GeminiChatPromptFeedback `json:"promptFeedback"` | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func (g *ChatResponse) GetResponseText() string { | func (g *GeminiChatResponse) GetResponseText() string { | ||||||
| 	if g == nil { | 	if g == nil { | ||||||
| 		return "" | 		return "" | ||||||
| 	} | 	} | ||||||
| @@ -137,37 +171,37 @@ func (g *ChatResponse) GetResponseText() string { | |||||||
| 	return "" | 	return "" | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| type ChatCandidate struct { | type GeminiChatCandidate struct { | ||||||
| 	Content       ChatContent        `json:"content"` | 	Content       GeminiChatContent        `json:"content"` | ||||||
| 	FinishReason  string             `json:"finishReason"` | 	FinishReason  string                   `json:"finishReason"` | ||||||
| 	Index         int64              `json:"index"` | 	Index         int64                    `json:"index"` | ||||||
| 	SafetyRatings []ChatSafetyRating `json:"safetyRatings"` | 	SafetyRatings []GeminiChatSafetyRating `json:"safetyRatings"` | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| type ChatSafetyRating struct { | type GeminiChatSafetyRating struct { | ||||||
| 	Category    string `json:"category"` | 	Category    string `json:"category"` | ||||||
| 	Probability string `json:"probability"` | 	Probability string `json:"probability"` | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| type ChatPromptFeedback struct { | type GeminiChatPromptFeedback struct { | ||||||
| 	SafetyRatings []ChatSafetyRating `json:"safetyRatings"` | 	SafetyRatings []GeminiChatSafetyRating `json:"safetyRatings"` | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func responseGeminiChat2OpenAI(response *ChatResponse) *openai.TextResponse { | func responseGeminiChat2OpenAI(response *GeminiChatResponse) *OpenAITextResponse { | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := OpenAITextResponse{ | ||||||
| 		Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), | 		Id:      fmt.Sprintf("chatcmpl-%s", common.GetUUID()), | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Choices: make([]openai.TextResponseChoice, 0, len(response.Candidates)), | 		Choices: make([]OpenAITextResponseChoice, 0, len(response.Candidates)), | ||||||
| 	} | 	} | ||||||
| 	for i, candidate := range response.Candidates { | 	for i, candidate := range response.Candidates { | ||||||
| 		choice := openai.TextResponseChoice{ | 		choice := OpenAITextResponseChoice{ | ||||||
| 			Index: i, | 			Index: i, | ||||||
| 			Message: model.Message{ | 			Message: Message{ | ||||||
| 				Role:    "assistant", | 				Role:    "assistant", | ||||||
| 				Content: "", | 				Content: "", | ||||||
| 			}, | 			}, | ||||||
| 			FinishReason: constant.StopFinishReason, | 			FinishReason: stopFinishReason, | ||||||
| 		} | 		} | ||||||
| 		if len(candidate.Content.Parts) > 0 { | 		if len(candidate.Content.Parts) > 0 { | ||||||
| 			choice.Message.Content = candidate.Content.Parts[0].Text | 			choice.Message.Content = candidate.Content.Parts[0].Text | ||||||
| @@ -177,18 +211,18 @@ func responseGeminiChat2OpenAI(response *ChatResponse) *openai.TextResponse { | |||||||
| 	return &fullTextResponse | 	return &fullTextResponse | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func streamResponseGeminiChat2OpenAI(geminiResponse *ChatResponse) *openai.ChatCompletionsStreamResponse { | func streamResponseGeminiChat2OpenAI(geminiResponse *GeminiChatResponse) *ChatCompletionsStreamResponse { | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice ChatCompletionsStreamResponseChoice | ||||||
| 	choice.Delta.Content = geminiResponse.GetResponseText() | 	choice.Delta.Content = geminiResponse.GetResponseText() | ||||||
| 	choice.FinishReason = &constant.StopFinishReason | 	choice.FinishReason = &stopFinishReason | ||||||
| 	var response openai.ChatCompletionsStreamResponse | 	var response ChatCompletionsStreamResponse | ||||||
| 	response.Object = "chat.completion.chunk" | 	response.Object = "chat.completion.chunk" | ||||||
| 	response.Model = "gemini" | 	response.Model = "gemini" | ||||||
| 	response.Choices = []openai.ChatCompletionsStreamResponseChoice{choice} | 	response.Choices = []ChatCompletionsStreamResponseChoice{choice} | ||||||
| 	return &response | 	return &response | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, string) { | func geminiChatStreamHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, string) { | ||||||
| 	responseText := "" | 	responseText := "" | ||||||
| 	dataChan := make(chan string) | 	dataChan := make(chan string) | ||||||
| 	stopChan := make(chan bool) | 	stopChan := make(chan bool) | ||||||
| @@ -218,7 +252,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 		} | 		} | ||||||
| 		stopChan <- true | 		stopChan <- true | ||||||
| 	}() | 	}() | ||||||
| 	common.SetEventStreamHeaders(c) | 	setEventStreamHeaders(c) | ||||||
| 	c.Stream(func(w io.Writer) bool { | 	c.Stream(func(w io.Writer) bool { | ||||||
| 		select { | 		select { | ||||||
| 		case data := <-dataChan: | 		case data := <-dataChan: | ||||||
| @@ -230,18 +264,18 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 			var dummy dummyStruct | 			var dummy dummyStruct | ||||||
| 			err := json.Unmarshal([]byte(data), &dummy) | 			err := json.Unmarshal([]byte(data), &dummy) | ||||||
| 			responseText += dummy.Content | 			responseText += dummy.Content | ||||||
| 			var choice openai.ChatCompletionsStreamResponseChoice | 			var choice ChatCompletionsStreamResponseChoice | ||||||
| 			choice.Delta.Content = dummy.Content | 			choice.Delta.Content = dummy.Content | ||||||
| 			response := openai.ChatCompletionsStreamResponse{ | 			response := ChatCompletionsStreamResponse{ | ||||||
| 				Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), | 				Id:      fmt.Sprintf("chatcmpl-%s", common.GetUUID()), | ||||||
| 				Object:  "chat.completion.chunk", | 				Object:  "chat.completion.chunk", | ||||||
| 				Created: helper.GetTimestamp(), | 				Created: common.GetTimestamp(), | ||||||
| 				Model:   "gemini-pro", | 				Model:   "gemini-pro", | ||||||
| 				Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 				Choices: []ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 			} | 			} | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error marshalling stream response: " + err.Error()) | 				common.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)}) | ||||||
| @@ -253,28 +287,28 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 	}) | 	}) | ||||||
| 	err := resp.Body.Close() | 	err := resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), "" | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), "" | ||||||
| 	} | 	} | ||||||
| 	return nil, responseText | 	return nil, responseText | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName string) (*model.ErrorWithStatusCode, *model.Usage) { | func geminiChatHandler(c *gin.Context, resp *http.Response, promptTokens int, model string) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	responseBody, err := io.ReadAll(resp.Body) | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	var geminiResponse ChatResponse | 	var geminiResponse GeminiChatResponse | ||||||
| 	err = json.Unmarshal(responseBody, &geminiResponse) | 	err = json.Unmarshal(responseBody, &geminiResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	if len(geminiResponse.Candidates) == 0 { | 	if len(geminiResponse.Candidates) == 0 { | ||||||
| 		return &model.ErrorWithStatusCode{ | 		return &OpenAIErrorWithStatusCode{ | ||||||
| 			Error: model.Error{ | 			OpenAIError: OpenAIError{ | ||||||
| 				Message: "No candidates returned", | 				Message: "No candidates returned", | ||||||
| 				Type:    "server_error", | 				Type:    "server_error", | ||||||
| 				Param:   "", | 				Param:   "", | ||||||
| @@ -284,9 +318,9 @@ func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName st | |||||||
| 		}, nil | 		}, nil | ||||||
| 	} | 	} | ||||||
| 	fullTextResponse := responseGeminiChat2OpenAI(&geminiResponse) | 	fullTextResponse := responseGeminiChat2OpenAI(&geminiResponse) | ||||||
| 	fullTextResponse.Model = modelName | 	fullTextResponse.Model = model | ||||||
| 	completionTokens := openai.CountTokenText(geminiResponse.GetResponseText(), modelName) | 	completionTokens := countTokenText(geminiResponse.GetResponseText(), model) | ||||||
| 	usage := model.Usage{ | 	usage := Usage{ | ||||||
| 		PromptTokens:     promptTokens, | 		PromptTokens:     promptTokens, | ||||||
| 		CompletionTokens: completionTokens, | 		CompletionTokens: completionTokens, | ||||||
| 		TotalTokens:      promptTokens + completionTokens, | 		TotalTokens:      promptTokens + completionTokens, | ||||||
| @@ -294,7 +328,7 @@ func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName st | |||||||
| 	fullTextResponse.Usage = usage | 	fullTextResponse.Usage = usage | ||||||
| 	jsonResponse, err := json.Marshal(fullTextResponse) | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
| 	c.Writer.WriteHeader(resp.StatusCode) | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
							
								
								
									
										222
									
								
								controller/relay-image.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										222
									
								
								controller/relay-image.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,222 @@ | |||||||
|  | package controller | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"bytes" | ||||||
|  | 	"context" | ||||||
|  | 	"encoding/json" | ||||||
|  | 	"errors" | ||||||
|  | 	"fmt" | ||||||
|  | 	"io" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
|  | 	"strings" | ||||||
|  |  | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | func isWithinRange(element string, value int) bool { | ||||||
|  | 	if _, ok := common.DalleGenerationImageAmounts[element]; !ok { | ||||||
|  | 		return false | ||||||
|  | 	} | ||||||
|  | 	min := common.DalleGenerationImageAmounts[element][0] | ||||||
|  | 	max := common.DalleGenerationImageAmounts[element][1] | ||||||
|  |  | ||||||
|  | 	return value >= min && value <= max | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func relayImageHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode { | ||||||
|  | 	imageModel := "dall-e-2" | ||||||
|  | 	imageSize := "1024x1024" | ||||||
|  |  | ||||||
|  | 	tokenId := c.GetInt("token_id") | ||||||
|  | 	channelType := c.GetInt("channel") | ||||||
|  | 	channelId := c.GetInt("channel_id") | ||||||
|  | 	userId := c.GetInt("id") | ||||||
|  | 	group := c.GetString("group") | ||||||
|  |  | ||||||
|  | 	var imageRequest ImageRequest | ||||||
|  | 	err := common.UnmarshalBodyReusable(c, &imageRequest) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "bind_request_body_failed", http.StatusBadRequest) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	if imageRequest.N == 0 { | ||||||
|  | 		imageRequest.N = 1 | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	// Size validation | ||||||
|  | 	if imageRequest.Size != "" { | ||||||
|  | 		imageSize = imageRequest.Size | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	// Model validation | ||||||
|  | 	if imageRequest.Model != "" { | ||||||
|  | 		imageModel = imageRequest.Model | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	imageCostRatio, hasValidSize := common.DalleSizeRatios[imageModel][imageSize] | ||||||
|  |  | ||||||
|  | 	// Check if model is supported | ||||||
|  | 	if hasValidSize { | ||||||
|  | 		if imageRequest.Quality == "hd" && imageModel == "dall-e-3" { | ||||||
|  | 			if imageSize == "1024x1024" { | ||||||
|  | 				imageCostRatio *= 2 | ||||||
|  | 			} else { | ||||||
|  | 				imageCostRatio *= 1.5 | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 	} else { | ||||||
|  | 		return errorWrapper(errors.New("size not supported for this image model"), "size_not_supported", http.StatusBadRequest) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	// Prompt validation | ||||||
|  | 	if imageRequest.Prompt == "" { | ||||||
|  | 		return errorWrapper(errors.New("prompt is required"), "prompt_missing", http.StatusBadRequest) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	// Check prompt length | ||||||
|  | 	if len(imageRequest.Prompt) > common.DalleImagePromptLengthLimitations[imageModel] { | ||||||
|  | 		return errorWrapper(errors.New("prompt is too long"), "prompt_too_long", http.StatusBadRequest) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	// Number of generated images validation | ||||||
|  | 	if isWithinRange(imageModel, imageRequest.N) == false { | ||||||
|  | 		// channel not azure | ||||||
|  | 		if channelType != common.ChannelTypeAzure { | ||||||
|  | 			return errorWrapper(errors.New("invalid value of n"), "n_not_within_range", http.StatusBadRequest) | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	// map model name | ||||||
|  | 	modelMapping := c.GetString("model_mapping") | ||||||
|  | 	isModelMapped := false | ||||||
|  | 	if modelMapping != "" { | ||||||
|  | 		modelMap := make(map[string]string) | ||||||
|  | 		err := json.Unmarshal([]byte(modelMapping), &modelMap) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "unmarshal_model_mapping_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		if modelMap[imageModel] != "" { | ||||||
|  | 			imageModel = modelMap[imageModel] | ||||||
|  | 			isModelMapped = true | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	baseURL := common.ChannelBaseURLs[channelType] | ||||||
|  | 	requestURL := c.Request.URL.String() | ||||||
|  | 	if c.GetString("base_url") != "" { | ||||||
|  | 		baseURL = c.GetString("base_url") | ||||||
|  | 	} | ||||||
|  | 	fullRequestURL := getFullRequestURL(baseURL, requestURL, channelType) | ||||||
|  | 	if channelType == common.ChannelTypeAzure { | ||||||
|  | 		// https://learn.microsoft.com/en-us/azure/ai-services/openai/dall-e-quickstart?tabs=dalle3%2Ccommand-line&pivots=rest-api | ||||||
|  | 		apiVersion := GetAPIVersion(c) | ||||||
|  | 		// 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) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	var requestBody io.Reader | ||||||
|  | 	if isModelMapped || channelType == common.ChannelTypeAzure { // make Azure channel request body | ||||||
|  | 		jsonStr, err := json.Marshal(imageRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	} else { | ||||||
|  | 		requestBody = c.Request.Body | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	modelRatio := common.GetModelRatio(imageModel) | ||||||
|  | 	groupRatio := common.GetGroupRatio(group) | ||||||
|  | 	ratio := modelRatio * groupRatio | ||||||
|  | 	userQuota, err := model.CacheGetUserQuota(userId) | ||||||
|  |  | ||||||
|  | 	quota := int(ratio*imageCostRatio*1000) * imageRequest.N | ||||||
|  |  | ||||||
|  | 	if userQuota-quota < 0 { | ||||||
|  | 		return errorWrapper(errors.New("user quota is not enough"), "insufficient_user_quota", http.StatusForbidden) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "new_request_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	token := c.Request.Header.Get("Authorization") | ||||||
|  | 	if channelType == common.ChannelTypeAzure { // Azure authentication | ||||||
|  | 		token = strings.TrimPrefix(token, "Bearer ") | ||||||
|  | 		req.Header.Set("api-key", token) | ||||||
|  | 	} else { | ||||||
|  | 		req.Header.Set("Authorization", token) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	req.Header.Set("Content-Type", c.Request.Header.Get("Content-Type")) | ||||||
|  | 	req.Header.Set("Accept", c.Request.Header.Get("Accept")) | ||||||
|  |  | ||||||
|  | 	resp, err := httpClient.Do(req) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "do_request_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	err = req.Body.Close() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "close_request_body_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	err = c.Request.Body.Close() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "close_request_body_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	var textResponse ImageResponse | ||||||
|  |  | ||||||
|  | 	defer func(ctx context.Context) { | ||||||
|  | 		if resp.StatusCode != http.StatusOK { | ||||||
|  | 			return | ||||||
|  | 		} | ||||||
|  | 		err := model.PostConsumeTokenQuota(tokenId, quota) | ||||||
|  | 		if err != nil { | ||||||
|  | 			common.SysError("error consuming token remain quota: " + err.Error()) | ||||||
|  | 		} | ||||||
|  | 		err = model.CacheUpdateUserQuota(userId) | ||||||
|  | 		if err != nil { | ||||||
|  | 			common.SysError("error update user quota cache: " + err.Error()) | ||||||
|  | 		} | ||||||
|  | 		if quota != 0 { | ||||||
|  | 			tokenName := c.GetString("token_name") | ||||||
|  | 			logContent := fmt.Sprintf("模型倍率 %.2f,分组倍率 %.2f", modelRatio, groupRatio) | ||||||
|  | 			model.RecordConsumeLog(ctx, userId, channelId, 0, 0, imageModel, tokenName, quota, logContent) | ||||||
|  | 			model.UpdateUserUsedQuotaAndRequestCount(userId, quota) | ||||||
|  | 			channelId := c.GetInt("channel_id") | ||||||
|  | 			model.UpdateChannelUsedQuota(channelId, quota) | ||||||
|  | 		} | ||||||
|  | 	}(c.Request.Context()) | ||||||
|  |  | ||||||
|  | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
|  |  | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	err = resp.Body.Close() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	err = json.Unmarshal(responseBody, &textResponse) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	resp.Body = io.NopCloser(bytes.NewBuffer(responseBody)) | ||||||
|  |  | ||||||
|  | 	for k, v := range resp.Header { | ||||||
|  | 		c.Writer.Header().Set(k, v[0]) | ||||||
|  | 	} | ||||||
|  | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
|  |  | ||||||
|  | 	_, err = io.Copy(c.Writer, resp.Body) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "copy_response_body_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	err = resp.Body.Close() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
| @@ -1,20 +1,17 @@ | |||||||
| package openai | package controller | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"bufio" | 	"bufio" | ||||||
| 	"bytes" | 	"bytes" | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
| 
 | 
 | ||||||
| func StreamHandler(c *gin.Context, resp *http.Response, relayMode int) (*model.ErrorWithStatusCode, string, *model.Usage) { | func openaiStreamHandler(c *gin.Context, resp *http.Response, relayMode int) (*OpenAIErrorWithStatusCode, string) { | ||||||
| 	responseText := "" | 	responseText := "" | ||||||
| 	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) { | ||||||
| @@ -31,7 +28,6 @@ func StreamHandler(c *gin.Context, resp *http.Response, relayMode int) (*model.E | |||||||
| 	}) | 	}) | ||||||
| 	dataChan := make(chan string) | 	dataChan := make(chan string) | ||||||
| 	stopChan := make(chan bool) | 	stopChan := make(chan bool) | ||||||
| 	var usage *model.Usage |  | ||||||
| 	go func() { | 	go func() { | ||||||
| 		for scanner.Scan() { | 		for scanner.Scan() { | ||||||
| 			data := scanner.Text() | 			data := scanner.Text() | ||||||
| @@ -45,24 +41,21 @@ func StreamHandler(c *gin.Context, resp *http.Response, relayMode int) (*model.E | |||||||
| 			data = data[6:] | 			data = data[6:] | ||||||
| 			if !strings.HasPrefix(data, "[DONE]") { | 			if !strings.HasPrefix(data, "[DONE]") { | ||||||
| 				switch relayMode { | 				switch relayMode { | ||||||
| 				case constant.RelayModeChatCompletions: | 				case RelayModeChatCompletions: | ||||||
| 					var streamResponse ChatCompletionsStreamResponse | 					var streamResponse ChatCompletionsStreamResponse | ||||||
| 					err := json.Unmarshal([]byte(data), &streamResponse) | 					err := json.Unmarshal([]byte(data), &streamResponse) | ||||||
| 					if err != nil { | 					if err != nil { | ||||||
| 						logger.SysError("error unmarshalling stream response: " + err.Error()) | 						common.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 { | ||||||
| 						responseText += choice.Delta.Content | 						responseText += choice.Delta.Content | ||||||
| 					} | 					} | ||||||
| 					if streamResponse.Usage != nil { | 				case RelayModeCompletions: | ||||||
| 						usage = streamResponse.Usage |  | ||||||
| 					} |  | ||||||
| 				case constant.RelayModeCompletions: |  | ||||||
| 					var streamResponse CompletionsStreamResponse | 					var streamResponse CompletionsStreamResponse | ||||||
| 					err := json.Unmarshal([]byte(data), &streamResponse) | 					err := json.Unmarshal([]byte(data), &streamResponse) | ||||||
| 					if err != nil { | 					if err != nil { | ||||||
| 						logger.SysError("error unmarshalling stream response: " + err.Error()) | 						common.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 						continue | 						continue | ||||||
| 					} | 					} | ||||||
| 					for _, choice := range streamResponse.Choices { | 					for _, choice := range streamResponse.Choices { | ||||||
| @@ -73,7 +66,7 @@ func StreamHandler(c *gin.Context, resp *http.Response, relayMode int) (*model.E | |||||||
| 		} | 		} | ||||||
| 		stopChan <- true | 		stopChan <- true | ||||||
| 	}() | 	}() | ||||||
| 	common.SetEventStreamHeaders(c) | 	setEventStreamHeaders(c) | ||||||
| 	c.Stream(func(w io.Writer) bool { | 	c.Stream(func(w io.Writer) bool { | ||||||
| 		select { | 		select { | ||||||
| 		case data := <-dataChan: | 		case data := <-dataChan: | ||||||
| @@ -90,29 +83,29 @@ func StreamHandler(c *gin.Context, resp *http.Response, relayMode int) (*model.E | |||||||
| 	}) | 	}) | ||||||
| 	err := resp.Body.Close() | 	err := resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), "", nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), "" | ||||||
| 	} | 	} | ||||||
| 	return nil, responseText, usage | 	return nil, responseText | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName string) (*model.ErrorWithStatusCode, *model.Usage) { | func openaiHandler(c *gin.Context, resp *http.Response, promptTokens int, model string) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var textResponse SlimTextResponse | 	var textResponse TextResponse | ||||||
| 	responseBody, err := io.ReadAll(resp.Body) | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = json.Unmarshal(responseBody, &textResponse) | 	err = json.Unmarshal(responseBody, &textResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	if textResponse.Error.Type != "" { | 	if textResponse.Error.Type != "" { | ||||||
| 		return &model.ErrorWithStatusCode{ | 		return &OpenAIErrorWithStatusCode{ | ||||||
| 			Error:      textResponse.Error, | 			OpenAIError: textResponse.Error, | ||||||
| 			StatusCode: resp.StatusCode, | 			StatusCode:  resp.StatusCode, | ||||||
| 		}, nil | 		}, nil | ||||||
| 	} | 	} | ||||||
| 	// Reset response body | 	// Reset response body | ||||||
| @@ -120,7 +113,7 @@ func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName st | |||||||
| 
 | 
 | ||||||
| 	// We shouldn't set the header before we parse the response body, because the parse part may fail. | 	// We shouldn't set the header before we parse the response body, because the parse part may fail. | ||||||
| 	// And then we will have to send an error response, but in this case, the header has already been set. | 	// And then we will have to send an error response, but in this case, the header has already been set. | ||||||
| 	// So the HTTPClient will be confused by the response. | 	// So the httpClient will be confused by the response. | ||||||
| 	// For example, Postman will report error, and we cannot check the response at all. | 	// For example, Postman will report error, and we cannot check the response at all. | ||||||
| 	for k, v := range resp.Header { | 	for k, v := range resp.Header { | ||||||
| 		c.Writer.Header().Set(k, v[0]) | 		c.Writer.Header().Set(k, v[0]) | ||||||
| @@ -128,19 +121,19 @@ func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName st | |||||||
| 	c.Writer.WriteHeader(resp.StatusCode) | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
| 	_, err = io.Copy(c.Writer, resp.Body) | 	_, err = io.Copy(c.Writer, resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return ErrorWrapper(err, "copy_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "copy_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	if textResponse.Usage.TotalTokens == 0 { | 	if textResponse.Usage.TotalTokens == 0 { | ||||||
| 		completionTokens := 0 | 		completionTokens := 0 | ||||||
| 		for _, choice := range textResponse.Choices { | 		for _, choice := range textResponse.Choices { | ||||||
| 			completionTokens += CountTokenText(choice.Message.StringContent(), modelName) | 			completionTokens += countTokenText(choice.Message.StringContent(), model) | ||||||
| 		} | 		} | ||||||
| 		textResponse.Usage = model.Usage{ | 		textResponse.Usage = Usage{ | ||||||
| 			PromptTokens:     promptTokens, | 			PromptTokens:     promptTokens, | ||||||
| 			CompletionTokens: completionTokens, | 			CompletionTokens: completionTokens, | ||||||
| 			TotalTokens:      promptTokens + completionTokens, | 			TotalTokens:      promptTokens + completionTokens, | ||||||
| @@ -1,26 +1,56 @@ | |||||||
| package palm | package controller | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
| ) | ) | ||||||
| 
 | 
 | ||||||
| // https://developers.generativeai.google/api/rest/generativelanguage/models/generateMessage#request-body | // https://developers.generativeai.google/api/rest/generativelanguage/models/generateMessage#request-body | ||||||
| // https://developers.generativeai.google/api/rest/generativelanguage/models/generateMessage#response-body | // https://developers.generativeai.google/api/rest/generativelanguage/models/generateMessage#response-body | ||||||
| 
 | 
 | ||||||
| func ConvertRequest(textRequest model.GeneralOpenAIRequest) *ChatRequest { | type PaLMChatMessage struct { | ||||||
| 	palmRequest := ChatRequest{ | 	Author  string `json:"author"` | ||||||
| 		Prompt: Prompt{ | 	Content string `json:"content"` | ||||||
| 			Messages: make([]ChatMessage, 0, len(textRequest.Messages)), | } | ||||||
|  | 
 | ||||||
|  | type PaLMFilter struct { | ||||||
|  | 	Reason  string `json:"reason"` | ||||||
|  | 	Message string `json:"message"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type PaLMPrompt struct { | ||||||
|  | 	Messages []PaLMChatMessage `json:"messages"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type PaLMChatRequest struct { | ||||||
|  | 	Prompt         PaLMPrompt `json:"prompt"` | ||||||
|  | 	Temperature    float64    `json:"temperature,omitempty"` | ||||||
|  | 	CandidateCount int        `json:"candidateCount,omitempty"` | ||||||
|  | 	TopP           float64    `json:"topP,omitempty"` | ||||||
|  | 	TopK           int        `json:"topK,omitempty"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type PaLMError struct { | ||||||
|  | 	Code    int    `json:"code"` | ||||||
|  | 	Message string `json:"message"` | ||||||
|  | 	Status  string `json:"status"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type PaLMChatResponse struct { | ||||||
|  | 	Candidates []PaLMChatMessage `json:"candidates"` | ||||||
|  | 	Messages   []Message         `json:"messages"` | ||||||
|  | 	Filters    []PaLMFilter      `json:"filters"` | ||||||
|  | 	Error      PaLMError         `json:"error"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func requestOpenAI2PaLM(textRequest GeneralOpenAIRequest) *PaLMChatRequest { | ||||||
|  | 	palmRequest := PaLMChatRequest{ | ||||||
|  | 		Prompt: PaLMPrompt{ | ||||||
|  | 			Messages: make([]PaLMChatMessage, 0, len(textRequest.Messages)), | ||||||
| 		}, | 		}, | ||||||
| 		Temperature:    textRequest.Temperature, | 		Temperature:    textRequest.Temperature, | ||||||
| 		CandidateCount: textRequest.N, | 		CandidateCount: textRequest.N, | ||||||
| @@ -28,7 +58,7 @@ func ConvertRequest(textRequest model.GeneralOpenAIRequest) *ChatRequest { | |||||||
| 		TopK:           textRequest.MaxTokens, | 		TopK:           textRequest.MaxTokens, | ||||||
| 	} | 	} | ||||||
| 	for _, message := range textRequest.Messages { | 	for _, message := range textRequest.Messages { | ||||||
| 		palmMessage := ChatMessage{ | 		palmMessage := PaLMChatMessage{ | ||||||
| 			Content: message.StringContent(), | 			Content: message.StringContent(), | ||||||
| 		} | 		} | ||||||
| 		if message.Role == "user" { | 		if message.Role == "user" { | ||||||
| @@ -41,14 +71,14 @@ func ConvertRequest(textRequest model.GeneralOpenAIRequest) *ChatRequest { | |||||||
| 	return &palmRequest | 	return &palmRequest | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func responsePaLM2OpenAI(response *ChatResponse) *openai.TextResponse { | func responsePaLM2OpenAI(response *PaLMChatResponse) *OpenAITextResponse { | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := OpenAITextResponse{ | ||||||
| 		Choices: make([]openai.TextResponseChoice, 0, len(response.Candidates)), | 		Choices: make([]OpenAITextResponseChoice, 0, len(response.Candidates)), | ||||||
| 	} | 	} | ||||||
| 	for i, candidate := range response.Candidates { | 	for i, candidate := range response.Candidates { | ||||||
| 		choice := openai.TextResponseChoice{ | 		choice := OpenAITextResponseChoice{ | ||||||
| 			Index: i, | 			Index: i, | ||||||
| 			Message: model.Message{ | 			Message: Message{ | ||||||
| 				Role:    "assistant", | 				Role:    "assistant", | ||||||
| 				Content: candidate.Content, | 				Content: candidate.Content, | ||||||
| 			}, | 			}, | ||||||
| @@ -59,42 +89,42 @@ func responsePaLM2OpenAI(response *ChatResponse) *openai.TextResponse { | |||||||
| 	return &fullTextResponse | 	return &fullTextResponse | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func streamResponsePaLM2OpenAI(palmResponse *ChatResponse) *openai.ChatCompletionsStreamResponse { | func streamResponsePaLM2OpenAI(palmResponse *PaLMChatResponse) *ChatCompletionsStreamResponse { | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice ChatCompletionsStreamResponseChoice | ||||||
| 	if len(palmResponse.Candidates) > 0 { | 	if len(palmResponse.Candidates) > 0 { | ||||||
| 		choice.Delta.Content = palmResponse.Candidates[0].Content | 		choice.Delta.Content = palmResponse.Candidates[0].Content | ||||||
| 	} | 	} | ||||||
| 	choice.FinishReason = &constant.StopFinishReason | 	choice.FinishReason = &stopFinishReason | ||||||
| 	var response openai.ChatCompletionsStreamResponse | 	var response ChatCompletionsStreamResponse | ||||||
| 	response.Object = "chat.completion.chunk" | 	response.Object = "chat.completion.chunk" | ||||||
| 	response.Model = "palm2" | 	response.Model = "palm2" | ||||||
| 	response.Choices = []openai.ChatCompletionsStreamResponseChoice{choice} | 	response.Choices = []ChatCompletionsStreamResponseChoice{choice} | ||||||
| 	return &response | 	return &response | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, string) { | func palmStreamHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, string) { | ||||||
| 	responseText := "" | 	responseText := "" | ||||||
| 	responseId := fmt.Sprintf("chatcmpl-%s", helper.GetUUID()) | 	responseId := fmt.Sprintf("chatcmpl-%s", common.GetUUID()) | ||||||
| 	createdTime := helper.GetTimestamp() | 	createdTime := common.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 { | ||||||
| 			logger.SysError("error reading stream response: " + err.Error()) | 			common.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 { | ||||||
| 			logger.SysError("error closing stream response: " + err.Error()) | 			common.SysError("error closing stream response: " + err.Error()) | ||||||
| 			stopChan <- true | 			stopChan <- true | ||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 		var palmResponse ChatResponse | 		var palmResponse PaLMChatResponse | ||||||
| 		err = json.Unmarshal(responseBody, &palmResponse) | 		err = json.Unmarshal(responseBody, &palmResponse) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			logger.SysError("error unmarshalling stream response: " + err.Error()) | 			common.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
| 			stopChan <- true | 			stopChan <- true | ||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| @@ -106,14 +136,14 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 		} | 		} | ||||||
| 		jsonResponse, err := json.Marshal(fullTextResponse) | 		jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			logger.SysError("error marshalling stream response: " + err.Error()) | 			common.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 			stopChan <- true | 			stopChan <- true | ||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 		dataChan <- string(jsonResponse) | 		dataChan <- string(jsonResponse) | ||||||
| 		stopChan <- true | 		stopChan <- true | ||||||
| 	}() | 	}() | ||||||
| 	common.SetEventStreamHeaders(c) | 	setEventStreamHeaders(c) | ||||||
| 	c.Stream(func(w io.Writer) bool { | 	c.Stream(func(w io.Writer) bool { | ||||||
| 		select { | 		select { | ||||||
| 		case data := <-dataChan: | 		case data := <-dataChan: | ||||||
| @@ -126,28 +156,28 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 	}) | 	}) | ||||||
| 	err := resp.Body.Close() | 	err := resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), "" | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), "" | ||||||
| 	} | 	} | ||||||
| 	return nil, responseText | 	return nil, responseText | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName string) (*model.ErrorWithStatusCode, *model.Usage) { | func palmHandler(c *gin.Context, resp *http.Response, promptTokens int, model string) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	responseBody, err := io.ReadAll(resp.Body) | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	var palmResponse ChatResponse | 	var palmResponse PaLMChatResponse | ||||||
| 	err = json.Unmarshal(responseBody, &palmResponse) | 	err = json.Unmarshal(responseBody, &palmResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	if palmResponse.Error.Code != 0 || len(palmResponse.Candidates) == 0 { | 	if palmResponse.Error.Code != 0 || len(palmResponse.Candidates) == 0 { | ||||||
| 		return &model.ErrorWithStatusCode{ | 		return &OpenAIErrorWithStatusCode{ | ||||||
| 			Error: model.Error{ | 			OpenAIError: OpenAIError{ | ||||||
| 				Message: palmResponse.Error.Message, | 				Message: palmResponse.Error.Message, | ||||||
| 				Type:    palmResponse.Error.Status, | 				Type:    palmResponse.Error.Status, | ||||||
| 				Param:   "", | 				Param:   "", | ||||||
| @@ -157,9 +187,9 @@ func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName st | |||||||
| 		}, nil | 		}, nil | ||||||
| 	} | 	} | ||||||
| 	fullTextResponse := responsePaLM2OpenAI(&palmResponse) | 	fullTextResponse := responsePaLM2OpenAI(&palmResponse) | ||||||
| 	fullTextResponse.Model = modelName | 	fullTextResponse.Model = model | ||||||
| 	completionTokens := openai.CountTokenText(palmResponse.Candidates[0].Content, modelName) | 	completionTokens := countTokenText(palmResponse.Candidates[0].Content, model) | ||||||
| 	usage := model.Usage{ | 	usage := Usage{ | ||||||
| 		PromptTokens:     promptTokens, | 		PromptTokens:     promptTokens, | ||||||
| 		CompletionTokens: completionTokens, | 		CompletionTokens: completionTokens, | ||||||
| 		TotalTokens:      promptTokens + completionTokens, | 		TotalTokens:      promptTokens + completionTokens, | ||||||
| @@ -167,7 +197,7 @@ func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName st | |||||||
| 	fullTextResponse.Usage = usage | 	fullTextResponse.Usage = usage | ||||||
| 	jsonResponse, err := json.Marshal(fullTextResponse) | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
| 	c.Writer.WriteHeader(resp.StatusCode) | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
							
								
								
									
										288
									
								
								controller/relay-tencent.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										288
									
								
								controller/relay-tencent.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,288 @@ | |||||||
|  | package controller | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"bufio" | ||||||
|  | 	"crypto/hmac" | ||||||
|  | 	"crypto/sha1" | ||||||
|  | 	"encoding/base64" | ||||||
|  | 	"encoding/json" | ||||||
|  | 	"errors" | ||||||
|  | 	"fmt" | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"io" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"sort" | ||||||
|  | 	"strconv" | ||||||
|  | 	"strings" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | // https://cloud.tencent.com/document/product/1729/97732 | ||||||
|  |  | ||||||
|  | type TencentMessage struct { | ||||||
|  | 	Role    string `json:"role"` | ||||||
|  | 	Content string `json:"content"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type TencentChatRequest struct { | ||||||
|  | 	AppId    int64  `json:"app_id"`    // 腾讯云账号的 APPID | ||||||
|  | 	SecretId string `json:"secret_id"` // 官网 SecretId | ||||||
|  | 	// Timestamp当前 UNIX 时间戳,单位为秒,可记录发起 API 请求的时间。 | ||||||
|  | 	// 例如1529223702,如果与当前时间相差过大,会引起签名过期错误 | ||||||
|  | 	Timestamp int64 `json:"timestamp"` | ||||||
|  | 	// Expired 签名的有效期,是一个符合 UNIX Epoch 时间戳规范的数值, | ||||||
|  | 	// 单位为秒;Expired 必须大于 Timestamp 且 Expired-Timestamp 小于90天 | ||||||
|  | 	Expired int64  `json:"expired"` | ||||||
|  | 	QueryID string `json:"query_id"` //请求 Id,用于问题排查 | ||||||
|  | 	// Temperature 较高的数值会使输出更加随机,而较低的数值会使其更加集中和确定 | ||||||
|  | 	// 默认 1.0,取值区间为[0.0,2.0],非必要不建议使用,不合理的取值会影响效果 | ||||||
|  | 	// 建议该参数和 top_p 只设置1个,不要同时更改 top_p | ||||||
|  | 	Temperature float64 `json:"temperature"` | ||||||
|  | 	// TopP 影响输出文本的多样性,取值越大,生成文本的多样性越强 | ||||||
|  | 	// 默认1.0,取值区间为[0.0, 1.0],非必要不建议使用, 不合理的取值会影响效果 | ||||||
|  | 	// 建议该参数和 temperature 只设置1个,不要同时更改 | ||||||
|  | 	TopP float64 `json:"top_p"` | ||||||
|  | 	// Stream 0:同步,1:流式 (默认,协议:SSE) | ||||||
|  | 	// 同步请求超时:60s,如果内容较长建议使用流式 | ||||||
|  | 	Stream int `json:"stream"` | ||||||
|  | 	// Messages 会话内容, 长度最多为40, 按对话时间从旧到新在数组中排列 | ||||||
|  | 	// 输入 content 总数最大支持 3000 token。 | ||||||
|  | 	Messages []TencentMessage `json:"messages"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type TencentError struct { | ||||||
|  | 	Code    int    `json:"code"` | ||||||
|  | 	Message string `json:"message"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type TencentUsage struct { | ||||||
|  | 	InputTokens  int `json:"input_tokens"` | ||||||
|  | 	OutputTokens int `json:"output_tokens"` | ||||||
|  | 	TotalTokens  int `json:"total_tokens"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type TencentResponseChoices struct { | ||||||
|  | 	FinishReason string         `json:"finish_reason,omitempty"` // 流式结束标志位,为 stop 则表示尾包 | ||||||
|  | 	Messages     TencentMessage `json:"messages,omitempty"`      // 内容,同步模式返回内容,流模式为 null 输出 content 内容总数最多支持 1024token。 | ||||||
|  | 	Delta        TencentMessage `json:"delta,omitempty"`         // 内容,流模式返回内容,同步模式为 null 输出 content 内容总数最多支持 1024token。 | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type TencentChatResponse struct { | ||||||
|  | 	Choices []TencentResponseChoices `json:"choices,omitempty"` // 结果 | ||||||
|  | 	Created string                   `json:"created,omitempty"` // unix 时间戳的字符串 | ||||||
|  | 	Id      string                   `json:"id,omitempty"`      // 会话 id | ||||||
|  | 	Usage   Usage                    `json:"usage,omitempty"`   // token 数量 | ||||||
|  | 	Error   TencentError             `json:"error,omitempty"`   // 错误信息 注意:此字段可能返回 null,表示取不到有效值 | ||||||
|  | 	Note    string                   `json:"note,omitempty"`    // 注释 | ||||||
|  | 	ReqID   string                   `json:"req_id,omitempty"`  // 唯一请求 Id,每次请求都会返回。用于反馈接口入参 | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func requestOpenAI2Tencent(request GeneralOpenAIRequest) *TencentChatRequest { | ||||||
|  | 	messages := make([]TencentMessage, 0, len(request.Messages)) | ||||||
|  | 	for i := 0; i < len(request.Messages); i++ { | ||||||
|  | 		message := request.Messages[i] | ||||||
|  | 		if message.Role == "system" { | ||||||
|  | 			messages = append(messages, TencentMessage{ | ||||||
|  | 				Role:    "user", | ||||||
|  | 				Content: message.StringContent(), | ||||||
|  | 			}) | ||||||
|  | 			messages = append(messages, TencentMessage{ | ||||||
|  | 				Role:    "assistant", | ||||||
|  | 				Content: "Okay", | ||||||
|  | 			}) | ||||||
|  | 			continue | ||||||
|  | 		} | ||||||
|  | 		messages = append(messages, TencentMessage{ | ||||||
|  | 			Content: message.StringContent(), | ||||||
|  | 			Role:    message.Role, | ||||||
|  | 		}) | ||||||
|  | 	} | ||||||
|  | 	stream := 0 | ||||||
|  | 	if request.Stream { | ||||||
|  | 		stream = 1 | ||||||
|  | 	} | ||||||
|  | 	return &TencentChatRequest{ | ||||||
|  | 		Timestamp:   common.GetTimestamp(), | ||||||
|  | 		Expired:     common.GetTimestamp() + 24*60*60, | ||||||
|  | 		QueryID:     common.GetUUID(), | ||||||
|  | 		Temperature: request.Temperature, | ||||||
|  | 		TopP:        request.TopP, | ||||||
|  | 		Stream:      stream, | ||||||
|  | 		Messages:    messages, | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func responseTencent2OpenAI(response *TencentChatResponse) *OpenAITextResponse { | ||||||
|  | 	fullTextResponse := OpenAITextResponse{ | ||||||
|  | 		Object:  "chat.completion", | ||||||
|  | 		Created: common.GetTimestamp(), | ||||||
|  | 		Usage:   response.Usage, | ||||||
|  | 	} | ||||||
|  | 	if len(response.Choices) > 0 { | ||||||
|  | 		choice := OpenAITextResponseChoice{ | ||||||
|  | 			Index: 0, | ||||||
|  | 			Message: Message{ | ||||||
|  | 				Role:    "assistant", | ||||||
|  | 				Content: response.Choices[0].Messages.Content, | ||||||
|  | 			}, | ||||||
|  | 			FinishReason: response.Choices[0].FinishReason, | ||||||
|  | 		} | ||||||
|  | 		fullTextResponse.Choices = append(fullTextResponse.Choices, choice) | ||||||
|  | 	} | ||||||
|  | 	return &fullTextResponse | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func streamResponseTencent2OpenAI(TencentResponse *TencentChatResponse) *ChatCompletionsStreamResponse { | ||||||
|  | 	response := ChatCompletionsStreamResponse{ | ||||||
|  | 		Object:  "chat.completion.chunk", | ||||||
|  | 		Created: common.GetTimestamp(), | ||||||
|  | 		Model:   "tencent-hunyuan", | ||||||
|  | 	} | ||||||
|  | 	if len(TencentResponse.Choices) > 0 { | ||||||
|  | 		var choice ChatCompletionsStreamResponseChoice | ||||||
|  | 		choice.Delta.Content = TencentResponse.Choices[0].Delta.Content | ||||||
|  | 		if TencentResponse.Choices[0].FinishReason == "stop" { | ||||||
|  | 			choice.FinishReason = &stopFinishReason | ||||||
|  | 		} | ||||||
|  | 		response.Choices = append(response.Choices, choice) | ||||||
|  | 	} | ||||||
|  | 	return &response | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func tencentStreamHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, string) { | ||||||
|  | 	var responseText string | ||||||
|  | 	scanner := bufio.NewScanner(resp.Body) | ||||||
|  | 	scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) { | ||||||
|  | 		if atEOF && len(data) == 0 { | ||||||
|  | 			return 0, nil, nil | ||||||
|  | 		} | ||||||
|  | 		if i := strings.Index(string(data), "\n"); i >= 0 { | ||||||
|  | 			return i + 1, data[0:i], nil | ||||||
|  | 		} | ||||||
|  | 		if atEOF { | ||||||
|  | 			return len(data), data, nil | ||||||
|  | 		} | ||||||
|  | 		return 0, nil, nil | ||||||
|  | 	}) | ||||||
|  | 	dataChan := make(chan string) | ||||||
|  | 	stopChan := make(chan bool) | ||||||
|  | 	go func() { | ||||||
|  | 		for scanner.Scan() { | ||||||
|  | 			data := scanner.Text() | ||||||
|  | 			if len(data) < 5 { // ignore blank line or wrong format | ||||||
|  | 				continue | ||||||
|  | 			} | ||||||
|  | 			if data[:5] != "data:" { | ||||||
|  | 				continue | ||||||
|  | 			} | ||||||
|  | 			data = data[5:] | ||||||
|  | 			dataChan <- data | ||||||
|  | 		} | ||||||
|  | 		stopChan <- true | ||||||
|  | 	}() | ||||||
|  | 	setEventStreamHeaders(c) | ||||||
|  | 	c.Stream(func(w io.Writer) bool { | ||||||
|  | 		select { | ||||||
|  | 		case data := <-dataChan: | ||||||
|  | 			var TencentResponse TencentChatResponse | ||||||
|  | 			err := json.Unmarshal([]byte(data), &TencentResponse) | ||||||
|  | 			if err != nil { | ||||||
|  | 				common.SysError("error unmarshalling stream response: " + err.Error()) | ||||||
|  | 				return true | ||||||
|  | 			} | ||||||
|  | 			response := streamResponseTencent2OpenAI(&TencentResponse) | ||||||
|  | 			if len(response.Choices) != 0 { | ||||||
|  | 				responseText += response.Choices[0].Delta.Content | ||||||
|  | 			} | ||||||
|  | 			jsonResponse, err := json.Marshal(response) | ||||||
|  | 			if err != nil { | ||||||
|  | 				common.SysError("error marshalling stream response: " + err.Error()) | ||||||
|  | 				return true | ||||||
|  | 			} | ||||||
|  | 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonResponse)}) | ||||||
|  | 			return true | ||||||
|  | 		case <-stopChan: | ||||||
|  | 			c.Render(-1, common.CustomEvent{Data: "data: [DONE]"}) | ||||||
|  | 			return false | ||||||
|  | 		} | ||||||
|  | 	}) | ||||||
|  | 	err := resp.Body.Close() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), "" | ||||||
|  | 	} | ||||||
|  | 	return nil, responseText | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func tencentHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
|  | 	var TencentResponse TencentChatResponse | ||||||
|  | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
|  | 	} | ||||||
|  | 	err = resp.Body.Close() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
|  | 	} | ||||||
|  | 	err = json.Unmarshal(responseBody, &TencentResponse) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
|  | 	} | ||||||
|  | 	if TencentResponse.Error.Code != 0 { | ||||||
|  | 		return &OpenAIErrorWithStatusCode{ | ||||||
|  | 			OpenAIError: OpenAIError{ | ||||||
|  | 				Message: TencentResponse.Error.Message, | ||||||
|  | 				Code:    TencentResponse.Error.Code, | ||||||
|  | 			}, | ||||||
|  | 			StatusCode: resp.StatusCode, | ||||||
|  | 		}, nil | ||||||
|  | 	} | ||||||
|  | 	fullTextResponse := responseTencent2OpenAI(&TencentResponse) | ||||||
|  | 	fullTextResponse.Model = "hunyuan" | ||||||
|  | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
|  | 	} | ||||||
|  | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
|  | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
|  | 	_, err = c.Writer.Write(jsonResponse) | ||||||
|  | 	return nil, &fullTextResponse.Usage | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func parseTencentConfig(config string) (appId int64, secretId string, secretKey string, err error) { | ||||||
|  | 	parts := strings.Split(config, "|") | ||||||
|  | 	if len(parts) != 3 { | ||||||
|  | 		err = errors.New("invalid tencent config") | ||||||
|  | 		return | ||||||
|  | 	} | ||||||
|  | 	appId, err = strconv.ParseInt(parts[0], 10, 64) | ||||||
|  | 	secretId = parts[1] | ||||||
|  | 	secretKey = parts[2] | ||||||
|  | 	return | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func getTencentSign(req TencentChatRequest, secretKey string) string { | ||||||
|  | 	params := make([]string, 0) | ||||||
|  | 	params = append(params, "app_id="+strconv.FormatInt(req.AppId, 10)) | ||||||
|  | 	params = append(params, "secret_id="+req.SecretId) | ||||||
|  | 	params = append(params, "timestamp="+strconv.FormatInt(req.Timestamp, 10)) | ||||||
|  | 	params = append(params, "query_id="+req.QueryID) | ||||||
|  | 	params = append(params, "temperature="+strconv.FormatFloat(req.Temperature, 'f', -1, 64)) | ||||||
|  | 	params = append(params, "top_p="+strconv.FormatFloat(req.TopP, 'f', -1, 64)) | ||||||
|  | 	params = append(params, "stream="+strconv.Itoa(req.Stream)) | ||||||
|  | 	params = append(params, "expired="+strconv.FormatInt(req.Expired, 10)) | ||||||
|  |  | ||||||
|  | 	var messageStr string | ||||||
|  | 	for _, msg := range req.Messages { | ||||||
|  | 		messageStr += fmt.Sprintf(`{"role":"%s","content":"%s"},`, msg.Role, msg.Content) | ||||||
|  | 	} | ||||||
|  | 	messageStr = strings.TrimSuffix(messageStr, ",") | ||||||
|  | 	params = append(params, "messages=["+messageStr+"]") | ||||||
|  |  | ||||||
|  | 	sort.Sort(sort.StringSlice(params)) | ||||||
|  | 	url := "hunyuan.cloud.tencent.com/hyllm/v1/chat/completions?" + strings.Join(params, "&") | ||||||
|  | 	mac := hmac.New(sha1.New, []byte(secretKey)) | ||||||
|  | 	signURL := url | ||||||
|  | 	mac.Write([]byte(signURL)) | ||||||
|  | 	sign := mac.Sum([]byte(nil)) | ||||||
|  | 	return base64.StdEncoding.EncodeToString(sign) | ||||||
|  | } | ||||||
							
								
								
									
										689
									
								
								controller/relay-text.go
									
									
									
									
									
										Normal file
									
								
							
							
						
						
									
										689
									
								
								controller/relay-text.go
									
									
									
									
									
										Normal file
									
								
							| @@ -0,0 +1,689 @@ | |||||||
|  | package controller | ||||||
|  |  | ||||||
|  | import ( | ||||||
|  | 	"bytes" | ||||||
|  | 	"context" | ||||||
|  | 	"encoding/json" | ||||||
|  | 	"errors" | ||||||
|  | 	"fmt" | ||||||
|  | 	"io" | ||||||
|  | 	"math" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
|  | 	"strings" | ||||||
|  | 	"time" | ||||||
|  |  | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | const ( | ||||||
|  | 	APITypeOpenAI = iota | ||||||
|  | 	APITypeClaude | ||||||
|  | 	APITypePaLM | ||||||
|  | 	APITypeBaidu | ||||||
|  | 	APITypeZhipu | ||||||
|  | 	APITypeAli | ||||||
|  | 	APITypeXunfei | ||||||
|  | 	APITypeAIProxyLibrary | ||||||
|  | 	APITypeTencent | ||||||
|  | 	APITypeGemini | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | var httpClient *http.Client | ||||||
|  | var impatientHTTPClient *http.Client | ||||||
|  |  | ||||||
|  | func init() { | ||||||
|  | 	if common.RelayTimeout == 0 { | ||||||
|  | 		httpClient = &http.Client{} | ||||||
|  | 	} else { | ||||||
|  | 		httpClient = &http.Client{ | ||||||
|  | 			Timeout: time.Duration(common.RelayTimeout) * time.Second, | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	impatientHTTPClient = &http.Client{ | ||||||
|  | 		Timeout: 5 * time.Second, | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func relayTextHelper(c *gin.Context, relayMode int) *OpenAIErrorWithStatusCode { | ||||||
|  | 	channelType := c.GetInt("channel") | ||||||
|  | 	channelId := c.GetInt("channel_id") | ||||||
|  | 	tokenId := c.GetInt("token_id") | ||||||
|  | 	userId := c.GetInt("id") | ||||||
|  | 	group := c.GetString("group") | ||||||
|  | 	var textRequest GeneralOpenAIRequest | ||||||
|  | 	err := common.UnmarshalBodyReusable(c, &textRequest) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "bind_request_body_failed", http.StatusBadRequest) | ||||||
|  | 	} | ||||||
|  | 	if textRequest.MaxTokens < 0 || textRequest.MaxTokens > math.MaxInt32/2 { | ||||||
|  | 		return errorWrapper(errors.New("max_tokens is invalid"), "invalid_max_tokens", http.StatusBadRequest) | ||||||
|  | 	} | ||||||
|  | 	if relayMode == RelayModeModerations && textRequest.Model == "" { | ||||||
|  | 		textRequest.Model = "text-moderation-latest" | ||||||
|  | 	} | ||||||
|  | 	if relayMode == RelayModeEmbeddings && textRequest.Model == "" { | ||||||
|  | 		textRequest.Model = c.Param("model") | ||||||
|  | 	} | ||||||
|  | 	// request validation | ||||||
|  | 	if textRequest.Model == "" { | ||||||
|  | 		return errorWrapper(errors.New("model is required"), "required_field_missing", http.StatusBadRequest) | ||||||
|  | 	} | ||||||
|  | 	switch relayMode { | ||||||
|  | 	case RelayModeCompletions: | ||||||
|  | 		if textRequest.Prompt == "" { | ||||||
|  | 			return errorWrapper(errors.New("field prompt is required"), "required_field_missing", http.StatusBadRequest) | ||||||
|  | 		} | ||||||
|  | 	case RelayModeChatCompletions: | ||||||
|  | 		if textRequest.Messages == nil || len(textRequest.Messages) == 0 { | ||||||
|  | 			return errorWrapper(errors.New("field messages is required"), "required_field_missing", http.StatusBadRequest) | ||||||
|  | 		} | ||||||
|  | 	case RelayModeEmbeddings: | ||||||
|  | 	case RelayModeModerations: | ||||||
|  | 		if textRequest.Input == "" { | ||||||
|  | 			return errorWrapper(errors.New("field input is required"), "required_field_missing", http.StatusBadRequest) | ||||||
|  | 		} | ||||||
|  | 	case RelayModeEdits: | ||||||
|  | 		if textRequest.Instruction == "" { | ||||||
|  | 			return 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 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 := 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 := 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 = 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 = getBaiduAccessToken(apiKey); err != nil { | ||||||
|  | 			return 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 == 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 completionTokens int | ||||||
|  | 	switch relayMode { | ||||||
|  | 	case RelayModeChatCompletions: | ||||||
|  | 		promptTokens = countTokenMessages(textRequest.Messages, textRequest.Model) | ||||||
|  | 	case RelayModeCompletions: | ||||||
|  | 		promptTokens = countTokenInput(textRequest.Prompt, textRequest.Model) | ||||||
|  | 	case RelayModeModerations: | ||||||
|  | 		promptTokens = countTokenInput(textRequest.Input, textRequest.Model) | ||||||
|  | 	} | ||||||
|  | 	preConsumedTokens := common.PreConsumedQuota | ||||||
|  | 	if textRequest.MaxTokens != 0 { | ||||||
|  | 		preConsumedTokens = promptTokens + textRequest.MaxTokens | ||||||
|  | 	} | ||||||
|  | 	modelRatio := common.GetModelRatio(textRequest.Model) | ||||||
|  | 	groupRatio := common.GetGroupRatio(group) | ||||||
|  | 	ratio := modelRatio * groupRatio | ||||||
|  | 	preConsumedQuota := int(float64(preConsumedTokens) * ratio) | ||||||
|  | 	userQuota, err := model.CacheGetUserQuota(userId) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "get_user_quota_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	if userQuota-preConsumedQuota < 0 { | ||||||
|  | 		return errorWrapper(errors.New("user quota is not enough"), "insufficient_user_quota", http.StatusForbidden) | ||||||
|  | 	} | ||||||
|  | 	err = model.CacheDecreaseUserQuota(userId, preConsumedQuota) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return errorWrapper(err, "decrease_user_quota_failed", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | 	if userQuota > 100*preConsumedQuota { | ||||||
|  | 		// in this case, we do not pre-consume quota | ||||||
|  | 		// because the user has enough quota | ||||||
|  | 		preConsumedQuota = 0 | ||||||
|  | 		common.LogInfo(c.Request.Context(), fmt.Sprintf("user %d has enough quota %d, trusted and no need to pre-consume", userId, userQuota)) | ||||||
|  | 	} | ||||||
|  | 	if preConsumedQuota > 0 { | ||||||
|  | 		err := model.PreConsumeTokenQuota(tokenId, preConsumedQuota) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "pre_consume_token_quota_failed", http.StatusForbidden) | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	var requestBody io.Reader | ||||||
|  | 	if isModelMapped { | ||||||
|  | 		jsonStr, err := json.Marshal(textRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	} else { | ||||||
|  | 		requestBody = c.Request.Body | ||||||
|  | 	} | ||||||
|  | 	switch apiType { | ||||||
|  | 	case APITypeClaude: | ||||||
|  | 		claudeRequest := requestOpenAI2Claude(textRequest) | ||||||
|  | 		jsonStr, err := json.Marshal(claudeRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case APITypeBaidu: | ||||||
|  | 		var jsonData []byte | ||||||
|  | 		var err error | ||||||
|  | 		switch relayMode { | ||||||
|  | 		case RelayModeEmbeddings: | ||||||
|  | 			baiduEmbeddingRequest := embeddingRequestOpenAI2Baidu(textRequest) | ||||||
|  | 			jsonData, err = json.Marshal(baiduEmbeddingRequest) | ||||||
|  | 		default: | ||||||
|  | 			baiduRequest := requestOpenAI2Baidu(textRequest) | ||||||
|  | 			jsonData, err = json.Marshal(baiduRequest) | ||||||
|  | 		} | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonData) | ||||||
|  | 	case APITypePaLM: | ||||||
|  | 		palmRequest := requestOpenAI2PaLM(textRequest) | ||||||
|  | 		jsonStr, err := json.Marshal(palmRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case APITypeGemini: | ||||||
|  | 		geminiChatRequest := requestOpenAI2Gemini(textRequest) | ||||||
|  | 		jsonStr, err := json.Marshal(geminiChatRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case APITypeZhipu: | ||||||
|  | 		zhipuRequest := requestOpenAI2Zhipu(textRequest) | ||||||
|  | 		jsonStr, err := json.Marshal(zhipuRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case APITypeAli: | ||||||
|  | 		var jsonStr []byte | ||||||
|  | 		var err error | ||||||
|  | 		switch relayMode { | ||||||
|  | 		case RelayModeEmbeddings: | ||||||
|  | 			aliEmbeddingRequest := embeddingRequestOpenAI2Ali(textRequest) | ||||||
|  | 			jsonStr, err = json.Marshal(aliEmbeddingRequest) | ||||||
|  | 		default: | ||||||
|  | 			aliRequest := requestOpenAI2Ali(textRequest) | ||||||
|  | 			jsonStr, err = json.Marshal(aliRequest) | ||||||
|  | 		} | ||||||
|  | 		if err != nil { | ||||||
|  | 			return 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 := parseTencentConfig(apiKey) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "invalid_tencent_config", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		tencentRequest := requestOpenAI2Tencent(textRequest) | ||||||
|  | 		tencentRequest.AppId = appId | ||||||
|  | 		tencentRequest.SecretId = secretId | ||||||
|  | 		jsonStr, err := json.Marshal(tencentRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		sign := getTencentSign(*tencentRequest, secretKey) | ||||||
|  | 		c.Request.Header.Set("Authorization", sign) | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	case APITypeAIProxyLibrary: | ||||||
|  | 		aiProxyLibraryRequest := requestOpenAI2AIProxyLibrary(textRequest) | ||||||
|  | 		aiProxyLibraryRequest.LibraryId = c.GetString("library_id") | ||||||
|  | 		jsonStr, err := json.Marshal(aiProxyLibraryRequest) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "marshal_text_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		requestBody = bytes.NewBuffer(jsonStr) | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	var req *http.Request | ||||||
|  | 	var resp *http.Response | ||||||
|  | 	isStream := textRequest.Stream | ||||||
|  |  | ||||||
|  | 	if apiType != APITypeXunfei { // cause xunfei use websocket | ||||||
|  | 		req, err = http.NewRequest(c.Request.Method, fullRequestURL, requestBody) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "new_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		apiKey := c.Request.Header.Get("Authorization") | ||||||
|  | 		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 := getZhipuToken(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 = httpClient.Do(req) | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "do_request_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		err = req.Body.Close() | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "close_request_body_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		err = c.Request.Body.Close() | ||||||
|  | 		if err != nil { | ||||||
|  | 			return errorWrapper(err, "close_request_body_failed", http.StatusInternalServerError) | ||||||
|  | 		} | ||||||
|  | 		isStream = isStream || strings.HasPrefix(resp.Header.Get("Content-Type"), "text/event-stream") | ||||||
|  |  | ||||||
|  | 		if resp.StatusCode != http.StatusOK { | ||||||
|  | 			if preConsumedQuota != 0 { | ||||||
|  | 				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 relayErrorHandler(resp) | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  |  | ||||||
|  | 	var textResponse TextResponse | ||||||
|  | 	tokenName := c.GetString("token_name") | ||||||
|  |  | ||||||
|  | 	defer func(ctx context.Context) { | ||||||
|  | 		// c.Writer.Flush() | ||||||
|  | 		go func() { | ||||||
|  | 			quota := 0 | ||||||
|  | 			completionRatio := common.GetCompletionRatio(textRequest.Model) | ||||||
|  | 			promptTokens = textResponse.Usage.PromptTokens | ||||||
|  | 			completionTokens = textResponse.Usage.CompletionTokens | ||||||
|  | 			quota = int(math.Ceil((float64(promptTokens) + float64(completionTokens)*completionRatio) * ratio)) | ||||||
|  | 			if ratio != 0 && quota <= 0 { | ||||||
|  | 				quota = 1 | ||||||
|  | 			} | ||||||
|  | 			totalTokens := promptTokens + completionTokens | ||||||
|  | 			if totalTokens == 0 { | ||||||
|  | 				// in this case, must be some error happened | ||||||
|  | 				// we cannot just return, because we may have to return the pre-consumed quota | ||||||
|  | 				quota = 0 | ||||||
|  | 			} | ||||||
|  | 			quotaDelta := quota - preConsumedQuota | ||||||
|  | 			err := model.PostConsumeTokenQuota(tokenId, quotaDelta) | ||||||
|  | 			if err != nil { | ||||||
|  | 				common.LogError(ctx, "error consuming token remain quota: "+err.Error()) | ||||||
|  | 			} | ||||||
|  | 			err = model.CacheUpdateUserQuota(userId) | ||||||
|  | 			if err != nil { | ||||||
|  | 				common.LogError(ctx, "error update user quota cache: "+err.Error()) | ||||||
|  | 			} | ||||||
|  | 			if quota != 0 { | ||||||
|  | 				logContent := fmt.Sprintf("模型倍率 %.2f,分组倍率 %.2f", modelRatio, groupRatio) | ||||||
|  | 				model.RecordConsumeLog(ctx, userId, channelId, promptTokens, completionTokens, textRequest.Model, tokenName, quota, logContent) | ||||||
|  | 				model.UpdateUserUsedQuotaAndRequestCount(userId, quota) | ||||||
|  | 				model.UpdateChannelUsedQuota(channelId, quota) | ||||||
|  | 			} | ||||||
|  |  | ||||||
|  | 		}() | ||||||
|  | 	}(c.Request.Context()) | ||||||
|  | 	switch apiType { | ||||||
|  | 	case APITypeOpenAI: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, responseText := openaiStreamHandler(c, resp, relayMode) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			textResponse.Usage.PromptTokens = promptTokens | ||||||
|  | 			textResponse.Usage.CompletionTokens = countTokenText(responseText, textRequest.Model) | ||||||
|  | 			return nil | ||||||
|  | 		} else { | ||||||
|  | 			err, usage := openaiHandler(c, resp, promptTokens, textRequest.Model) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			if usage != nil { | ||||||
|  | 				textResponse.Usage = *usage | ||||||
|  | 			} | ||||||
|  | 			return nil | ||||||
|  | 		} | ||||||
|  | 	case APITypeClaude: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, responseText := claudeStreamHandler(c, resp) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			textResponse.Usage.PromptTokens = promptTokens | ||||||
|  | 			textResponse.Usage.CompletionTokens = countTokenText(responseText, textRequest.Model) | ||||||
|  | 			return nil | ||||||
|  | 		} else { | ||||||
|  | 			err, usage := claudeHandler(c, resp, promptTokens, textRequest.Model) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			if usage != nil { | ||||||
|  | 				textResponse.Usage = *usage | ||||||
|  | 			} | ||||||
|  | 			return nil | ||||||
|  | 		} | ||||||
|  | 	case APITypeBaidu: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, usage := baiduStreamHandler(c, resp) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			if usage != nil { | ||||||
|  | 				textResponse.Usage = *usage | ||||||
|  | 			} | ||||||
|  | 			return nil | ||||||
|  | 		} else { | ||||||
|  | 			var err *OpenAIErrorWithStatusCode | ||||||
|  | 			var usage *Usage | ||||||
|  | 			switch relayMode { | ||||||
|  | 			case RelayModeEmbeddings: | ||||||
|  | 				err, usage = baiduEmbeddingHandler(c, resp) | ||||||
|  | 			default: | ||||||
|  | 				err, usage = baiduHandler(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 := palmStreamHandler(c, resp) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			textResponse.Usage.PromptTokens = promptTokens | ||||||
|  | 			textResponse.Usage.CompletionTokens = countTokenText(responseText, textRequest.Model) | ||||||
|  | 			return nil | ||||||
|  | 		} else { | ||||||
|  | 			err, usage := 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 := geminiChatStreamHandler(c, resp) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			textResponse.Usage.PromptTokens = promptTokens | ||||||
|  | 			textResponse.Usage.CompletionTokens = countTokenText(responseText, textRequest.Model) | ||||||
|  | 			return nil | ||||||
|  | 		} else { | ||||||
|  | 			err, usage := geminiChatHandler(c, resp, promptTokens, textRequest.Model) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			if usage != nil { | ||||||
|  | 				textResponse.Usage = *usage | ||||||
|  | 			} | ||||||
|  | 			return nil | ||||||
|  | 		} | ||||||
|  | 	case APITypeZhipu: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, usage := zhipuStreamHandler(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 := zhipuHandler(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 := aliStreamHandler(c, resp) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			if usage != nil { | ||||||
|  | 				textResponse.Usage = *usage | ||||||
|  | 			} | ||||||
|  | 			return nil | ||||||
|  | 		} else { | ||||||
|  | 			var err *OpenAIErrorWithStatusCode | ||||||
|  | 			var usage *Usage | ||||||
|  | 			switch relayMode { | ||||||
|  | 			case RelayModeEmbeddings: | ||||||
|  | 				err, usage = aliEmbeddingHandler(c, resp) | ||||||
|  | 			default: | ||||||
|  | 				err, usage = aliHandler(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 errorWrapper(errors.New("invalid auth"), "invalid_auth", http.StatusBadRequest) | ||||||
|  | 		} | ||||||
|  | 		var err *OpenAIErrorWithStatusCode | ||||||
|  | 		var usage *Usage | ||||||
|  | 		if isStream { | ||||||
|  | 			err, usage = xunfeiStreamHandler(c, textRequest, splits[0], splits[1], splits[2]) | ||||||
|  | 		} else { | ||||||
|  | 			err, usage = xunfeiHandler(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 := aiProxyLibraryStreamHandler(c, resp) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			if usage != nil { | ||||||
|  | 				textResponse.Usage = *usage | ||||||
|  | 			} | ||||||
|  | 			return nil | ||||||
|  | 		} else { | ||||||
|  | 			err, usage := aiProxyLibraryHandler(c, resp) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			if usage != nil { | ||||||
|  | 				textResponse.Usage = *usage | ||||||
|  | 			} | ||||||
|  | 			return nil | ||||||
|  | 		} | ||||||
|  | 	case APITypeTencent: | ||||||
|  | 		if isStream { | ||||||
|  | 			err, responseText := tencentStreamHandler(c, resp) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			textResponse.Usage.PromptTokens = promptTokens | ||||||
|  | 			textResponse.Usage.CompletionTokens = countTokenText(responseText, textRequest.Model) | ||||||
|  | 			return nil | ||||||
|  | 		} else { | ||||||
|  | 			err, usage := tencentHandler(c, resp) | ||||||
|  | 			if err != nil { | ||||||
|  | 				return err | ||||||
|  | 			} | ||||||
|  | 			if usage != nil { | ||||||
|  | 				textResponse.Usage = *usage | ||||||
|  | 			} | ||||||
|  | 			return nil | ||||||
|  | 		} | ||||||
|  | 	default: | ||||||
|  | 		return errorWrapper(errors.New("unknown api type"), "unknown_api_type", http.StatusInternalServerError) | ||||||
|  | 	} | ||||||
|  | } | ||||||
| @@ -1,34 +1,41 @@ | |||||||
| package openai | package controller | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
|  | 	"context" | ||||||
|  | 	"encoding/json" | ||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/pkoukk/tiktoken-go" | 	"io" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/image" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"math" | 	"math" | ||||||
|  | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/common/image" | ||||||
|  | 	"one-api/model" | ||||||
|  | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
|  | 
 | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | 	"github.com/pkoukk/tiktoken-go" | ||||||
| ) | ) | ||||||
| 
 | 
 | ||||||
|  | var stopFinishReason = "stop" | ||||||
|  | 
 | ||||||
| // tokenEncoderMap won't grow after initialization | // tokenEncoderMap won't grow after initialization | ||||||
| var tokenEncoderMap = map[string]*tiktoken.Tiktoken{} | var tokenEncoderMap = map[string]*tiktoken.Tiktoken{} | ||||||
| var defaultTokenEncoder *tiktoken.Tiktoken | var defaultTokenEncoder *tiktoken.Tiktoken | ||||||
| 
 | 
 | ||||||
| func InitTokenEncoders() { | func InitTokenEncoders() { | ||||||
| 	logger.SysLog("initializing token encoders") | 	common.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 { | ||||||
| 		logger.FatalLog(fmt.Sprintf("failed to get gpt-3.5-turbo token encoder: %s", err.Error())) | 		common.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 { | ||||||
| 		logger.FatalLog(fmt.Sprintf("failed to get gpt-4 token encoder: %s", err.Error())) | 		common.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") { | ||||||
| 			tokenEncoderMap[model] = gpt35TokenEncoder | 			tokenEncoderMap[model] = gpt35TokenEncoder | ||||||
| 		} else if strings.HasPrefix(model, "gpt-4") { | 		} else if strings.HasPrefix(model, "gpt-4") { | ||||||
| @@ -37,7 +44,7 @@ func InitTokenEncoders() { | |||||||
| 			tokenEncoderMap[model] = nil | 			tokenEncoderMap[model] = nil | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	logger.SysLog("token encoders initialized") | 	common.SysLog("token encoders initialized") | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func getTokenEncoder(model string) *tiktoken.Tiktoken { | func getTokenEncoder(model string) *tiktoken.Tiktoken { | ||||||
| @@ -48,7 +55,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 { | ||||||
| 			logger.SysError(fmt.Sprintf("failed to get token encoder for model %s: %s, using encoder for gpt-3.5-turbo", model, err.Error())) | 			common.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 | ||||||
| @@ -58,13 +65,13 @@ func getTokenEncoder(model string) *tiktoken.Tiktoken { | |||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func getTokenNum(tokenEncoder *tiktoken.Tiktoken, text string) int { | func getTokenNum(tokenEncoder *tiktoken.Tiktoken, text string) int { | ||||||
| 	if config.ApproximateTokenEnabled { | 	if common.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)) | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func CountTokenMessages(messages []model.Message, model string) int { | func countTokenMessages(messages []Message, model string) int { | ||||||
| 	tokenEncoder := getTokenEncoder(model) | 	tokenEncoder := getTokenEncoder(model) | ||||||
| 	// Reference: | 	// Reference: | ||||||
| 	// https://github.com/openai/openai-cookbook/blob/main/examples/How_to_count_tokens_with_tiktoken.ipynb | 	// https://github.com/openai/openai-cookbook/blob/main/examples/How_to_count_tokens_with_tiktoken.ipynb | ||||||
| @@ -102,7 +109,7 @@ func CountTokenMessages(messages []model.Message, model string) int { | |||||||
| 						} | 						} | ||||||
| 						imageTokens, err := countImageTokens(url, detail) | 						imageTokens, err := countImageTokens(url, detail) | ||||||
| 						if err != nil { | 						if err != nil { | ||||||
| 							logger.SysError("error counting image tokens: " + err.Error()) | 							common.SysError("error counting image tokens: " + err.Error()) | ||||||
| 						} else { | 						} else { | ||||||
| 							tokenNum += imageTokens | 							tokenNum += imageTokens | ||||||
| 						} | 						} | ||||||
| @@ -188,21 +195,191 @@ func countImageTokens(url string, detail string) (_ int, err error) { | |||||||
| 	} | 	} | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func CountTokenInput(input any, model string) int { | func countTokenInput(input any, model string) int { | ||||||
| 	switch v := input.(type) { | 	switch v := input.(type) { | ||||||
| 	case string: | 	case string: | ||||||
| 		return CountTokenText(v, model) | 		return countTokenText(v, model) | ||||||
| 	case []string: | 	case []string: | ||||||
| 		text := "" | 		text := "" | ||||||
| 		for _, s := range v { | 		for _, s := range v { | ||||||
| 			text += s | 			text += s | ||||||
| 		} | 		} | ||||||
| 		return CountTokenText(text, model) | 		return countTokenText(text, model) | ||||||
| 	} | 	} | ||||||
| 	return 0 | 	return 0 | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func CountTokenText(text string, model string) int { | func countTokenText(text string, model string) int { | ||||||
| 	tokenEncoder := getTokenEncoder(model) | 	tokenEncoder := getTokenEncoder(model) | ||||||
| 	return getTokenNum(tokenEncoder, text) | 	return getTokenNum(tokenEncoder, text) | ||||||
| } | } | ||||||
|  | 
 | ||||||
|  | func errorWrapper(err error, code string, statusCode int) *OpenAIErrorWithStatusCode { | ||||||
|  | 	openAIError := OpenAIError{ | ||||||
|  | 		Message: err.Error(), | ||||||
|  | 		Type:    "one_api_error", | ||||||
|  | 		Code:    code, | ||||||
|  | 	} | ||||||
|  | 	return &OpenAIErrorWithStatusCode{ | ||||||
|  | 		OpenAIError: openAIError, | ||||||
|  | 		StatusCode:  statusCode, | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func shouldDisableChannel(err *OpenAIError, statusCode int) bool { | ||||||
|  | 	if !common.AutomaticDisableChannelEnabled { | ||||||
|  | 		return false | ||||||
|  | 	} | ||||||
|  | 	if err == nil { | ||||||
|  | 		return false | ||||||
|  | 	} | ||||||
|  | 	if statusCode == http.StatusUnauthorized { | ||||||
|  | 		return true | ||||||
|  | 	} | ||||||
|  | 	if err.Type == "insufficient_quota" || err.Code == "invalid_api_key" || err.Code == "account_deactivated" { | ||||||
|  | 		return true | ||||||
|  | 	} | ||||||
|  | 	return false | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func shouldEnableChannel(err error, openAIErr *OpenAIError) bool { | ||||||
|  | 	if !common.AutomaticEnableChannelEnabled { | ||||||
|  | 		return false | ||||||
|  | 	} | ||||||
|  | 	if err != nil { | ||||||
|  | 		return false | ||||||
|  | 	} | ||||||
|  | 	if openAIErr != nil { | ||||||
|  | 		return false | ||||||
|  | 	} | ||||||
|  | 	return true | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func setEventStreamHeaders(c *gin.Context) { | ||||||
|  | 	c.Writer.Header().Set("Content-Type", "text/event-stream") | ||||||
|  | 	c.Writer.Header().Set("Cache-Control", "no-cache") | ||||||
|  | 	c.Writer.Header().Set("Connection", "keep-alive") | ||||||
|  | 	c.Writer.Header().Set("Transfer-Encoding", "chunked") | ||||||
|  | 	c.Writer.Header().Set("X-Accel-Buffering", "no") | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type GeneralErrorResponse struct { | ||||||
|  | 	Error    OpenAIError `json:"error"` | ||||||
|  | 	Message  string      `json:"message"` | ||||||
|  | 	Msg      string      `json:"msg"` | ||||||
|  | 	Err      string      `json:"err"` | ||||||
|  | 	ErrorMsg string      `json:"error_msg"` | ||||||
|  | 	Header   struct { | ||||||
|  | 		Message string `json:"message"` | ||||||
|  | 	} `json:"header"` | ||||||
|  | 	Response struct { | ||||||
|  | 		Error struct { | ||||||
|  | 			Message string `json:"message"` | ||||||
|  | 		} `json:"error"` | ||||||
|  | 	} `json:"response"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func (e GeneralErrorResponse) ToMessage() string { | ||||||
|  | 	if e.Error.Message != "" { | ||||||
|  | 		return e.Error.Message | ||||||
|  | 	} | ||||||
|  | 	if e.Message != "" { | ||||||
|  | 		return e.Message | ||||||
|  | 	} | ||||||
|  | 	if e.Msg != "" { | ||||||
|  | 		return e.Msg | ||||||
|  | 	} | ||||||
|  | 	if e.Err != "" { | ||||||
|  | 		return e.Err | ||||||
|  | 	} | ||||||
|  | 	if e.ErrorMsg != "" { | ||||||
|  | 		return e.ErrorMsg | ||||||
|  | 	} | ||||||
|  | 	if e.Header.Message != "" { | ||||||
|  | 		return e.Header.Message | ||||||
|  | 	} | ||||||
|  | 	if e.Response.Error.Message != "" { | ||||||
|  | 		return e.Response.Error.Message | ||||||
|  | 	} | ||||||
|  | 	return "" | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func relayErrorHandler(resp *http.Response) (openAIErrorWithStatusCode *OpenAIErrorWithStatusCode) { | ||||||
|  | 	openAIErrorWithStatusCode = &OpenAIErrorWithStatusCode{ | ||||||
|  | 		StatusCode: resp.StatusCode, | ||||||
|  | 		OpenAIError: OpenAIError{ | ||||||
|  | 			Message: "", | ||||||
|  | 			Type:    "upstream_error", | ||||||
|  | 			Code:    "bad_response_status_code", | ||||||
|  | 			Param:   strconv.Itoa(resp.StatusCode), | ||||||
|  | 		}, | ||||||
|  | 	} | ||||||
|  | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return | ||||||
|  | 	} | ||||||
|  | 	err = resp.Body.Close() | ||||||
|  | 	if err != nil { | ||||||
|  | 		return | ||||||
|  | 	} | ||||||
|  | 	var errResponse GeneralErrorResponse | ||||||
|  | 	err = json.Unmarshal(responseBody, &errResponse) | ||||||
|  | 	if err != nil { | ||||||
|  | 		return | ||||||
|  | 	} | ||||||
|  | 	if errResponse.Error.Message != "" { | ||||||
|  | 		// OpenAI format error, so we override the default one | ||||||
|  | 		openAIErrorWithStatusCode.OpenAIError = errResponse.Error | ||||||
|  | 	} else { | ||||||
|  | 		openAIErrorWithStatusCode.OpenAIError.Message = errResponse.ToMessage() | ||||||
|  | 	} | ||||||
|  | 	if openAIErrorWithStatusCode.OpenAIError.Message == "" { | ||||||
|  | 		openAIErrorWithStatusCode.OpenAIError.Message = fmt.Sprintf("bad response status code %d", resp.StatusCode) | ||||||
|  | 	} | ||||||
|  | 	return | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func getFullRequestURL(baseURL string, requestURL string, channelType int) string { | ||||||
|  | 	fullRequestURL := fmt.Sprintf("%s%s", baseURL, requestURL) | ||||||
|  | 
 | ||||||
|  | 	if strings.HasPrefix(baseURL, "https://gateway.ai.cloudflare.com") { | ||||||
|  | 		switch channelType { | ||||||
|  | 		case common.ChannelTypeOpenAI: | ||||||
|  | 			fullRequestURL = fmt.Sprintf("%s%s", baseURL, strings.TrimPrefix(requestURL, "/v1")) | ||||||
|  | 		case common.ChannelTypeAzure: | ||||||
|  | 			fullRequestURL = fmt.Sprintf("%s%s", baseURL, strings.TrimPrefix(requestURL, "/openai/deployments")) | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return fullRequestURL | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func postConsumeQuota(ctx context.Context, tokenId int, quotaDelta int, totalQuota int, userId int, channelId int, modelRatio float64, groupRatio float64, modelName string, tokenName string) { | ||||||
|  | 	// quotaDelta is remaining quota to be consumed | ||||||
|  | 	err := model.PostConsumeTokenQuota(tokenId, quotaDelta) | ||||||
|  | 	if err != nil { | ||||||
|  | 		common.SysError("error consuming token remain quota: " + err.Error()) | ||||||
|  | 	} | ||||||
|  | 	err = model.CacheUpdateUserQuota(userId) | ||||||
|  | 	if err != nil { | ||||||
|  | 		common.SysError("error update user quota cache: " + err.Error()) | ||||||
|  | 	} | ||||||
|  | 	// totalQuota is total quota consumed | ||||||
|  | 	if totalQuota != 0 { | ||||||
|  | 		logContent := fmt.Sprintf("模型倍率 %.2f,分组倍率 %.2f", modelRatio, groupRatio) | ||||||
|  | 		model.RecordConsumeLog(ctx, userId, channelId, totalQuota, 0, modelName, tokenName, totalQuota, logContent) | ||||||
|  | 		model.UpdateUserUsedQuotaAndRequestCount(userId, totalQuota) | ||||||
|  | 		model.UpdateChannelUsedQuota(channelId, totalQuota) | ||||||
|  | 	} | ||||||
|  | 	if totalQuota <= 0 { | ||||||
|  | 		common.LogError(ctx, fmt.Sprintf("totalQuota consumed is %d, something is wrong", totalQuota)) | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func GetAPIVersion(c *gin.Context) string { | ||||||
|  | 	query := c.Request.URL.Query() | ||||||
|  | 	apiVersion := query.Get("api-version") | ||||||
|  | 	if apiVersion == "" { | ||||||
|  | 		apiVersion = c.GetString("api_version") | ||||||
|  | 	} | ||||||
|  | 	return apiVersion | ||||||
|  | } | ||||||
| @@ -1,4 +1,4 @@ | |||||||
| package xunfei | package controller | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"crypto/hmac" | 	"crypto/hmac" | ||||||
| @@ -8,15 +8,10 @@ import ( | |||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/gorilla/websocket" | 	"github.com/gorilla/websocket" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"net/url" | 	"net/url" | ||||||
|  | 	"one-api/common" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
| @@ -24,15 +19,82 @@ import ( | |||||||
| // https://console.xfyun.cn/services/cbm | // https://console.xfyun.cn/services/cbm | ||||||
| // https://www.xfyun.cn/doc/spark/Web.html | // https://www.xfyun.cn/doc/spark/Web.html | ||||||
| 
 | 
 | ||||||
| func requestOpenAI2Xunfei(request model.GeneralOpenAIRequest, xunfeiAppId string, domain string) *ChatRequest { | type XunfeiMessage struct { | ||||||
| 	messages := make([]Message, 0, len(request.Messages)) | 	Role    string `json:"role"` | ||||||
|  | 	Content string `json:"content"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type XunfeiChatRequest struct { | ||||||
|  | 	Header struct { | ||||||
|  | 		AppId string `json:"app_id"` | ||||||
|  | 	} `json:"header"` | ||||||
|  | 	Parameter struct { | ||||||
|  | 		Chat struct { | ||||||
|  | 			Domain      string  `json:"domain,omitempty"` | ||||||
|  | 			Temperature float64 `json:"temperature,omitempty"` | ||||||
|  | 			TopK        int     `json:"top_k,omitempty"` | ||||||
|  | 			MaxTokens   int     `json:"max_tokens,omitempty"` | ||||||
|  | 			Auditing    bool    `json:"auditing,omitempty"` | ||||||
|  | 		} `json:"chat"` | ||||||
|  | 	} `json:"parameter"` | ||||||
|  | 	Payload struct { | ||||||
|  | 		Message struct { | ||||||
|  | 			Text []XunfeiMessage `json:"text"` | ||||||
|  | 		} `json:"message"` | ||||||
|  | 	} `json:"payload"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type XunfeiChatResponseTextItem struct { | ||||||
|  | 	Content string `json:"content"` | ||||||
|  | 	Role    string `json:"role"` | ||||||
|  | 	Index   int    `json:"index"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type XunfeiChatResponse struct { | ||||||
|  | 	Header struct { | ||||||
|  | 		Code    int    `json:"code"` | ||||||
|  | 		Message string `json:"message"` | ||||||
|  | 		Sid     string `json:"sid"` | ||||||
|  | 		Status  int    `json:"status"` | ||||||
|  | 	} `json:"header"` | ||||||
|  | 	Payload struct { | ||||||
|  | 		Choices struct { | ||||||
|  | 			Status int                          `json:"status"` | ||||||
|  | 			Seq    int                          `json:"seq"` | ||||||
|  | 			Text   []XunfeiChatResponseTextItem `json:"text"` | ||||||
|  | 		} `json:"choices"` | ||||||
|  | 		Usage struct { | ||||||
|  | 			//Text struct { | ||||||
|  | 			//	QuestionTokens   string `json:"question_tokens"` | ||||||
|  | 			//	PromptTokens     string `json:"prompt_tokens"` | ||||||
|  | 			//	CompletionTokens string `json:"completion_tokens"` | ||||||
|  | 			//	TotalTokens      string `json:"total_tokens"` | ||||||
|  | 			//} `json:"text"` | ||||||
|  | 			Text Usage `json:"text"` | ||||||
|  | 		} `json:"usage"` | ||||||
|  | 	} `json:"payload"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | func requestOpenAI2Xunfei(request GeneralOpenAIRequest, xunfeiAppId string, domain string) *XunfeiChatRequest { | ||||||
|  | 	messages := make([]XunfeiMessage, 0, len(request.Messages)) | ||||||
| 	for _, message := range request.Messages { | 	for _, message := range request.Messages { | ||||||
| 		messages = append(messages, Message{ | 		if message.Role == "system" { | ||||||
| 			Role:    message.Role, | 			messages = append(messages, XunfeiMessage{ | ||||||
| 			Content: message.StringContent(), | 				Role:    "user", | ||||||
| 		}) | 				Content: message.StringContent(), | ||||||
|  | 			}) | ||||||
|  | 			messages = append(messages, XunfeiMessage{ | ||||||
|  | 				Role:    "assistant", | ||||||
|  | 				Content: "Okay", | ||||||
|  | 			}) | ||||||
|  | 		} else { | ||||||
|  | 			messages = append(messages, XunfeiMessage{ | ||||||
|  | 				Role:    message.Role, | ||||||
|  | 				Content: message.StringContent(), | ||||||
|  | 			}) | ||||||
|  | 		} | ||||||
| 	} | 	} | ||||||
| 	xunfeiRequest := ChatRequest{} | 	xunfeiRequest := XunfeiChatRequest{} | ||||||
| 	xunfeiRequest.Header.AppId = xunfeiAppId | 	xunfeiRequest.Header.AppId = xunfeiAppId | ||||||
| 	xunfeiRequest.Parameter.Chat.Domain = domain | 	xunfeiRequest.Parameter.Chat.Domain = domain | ||||||
| 	xunfeiRequest.Parameter.Chat.Temperature = request.Temperature | 	xunfeiRequest.Parameter.Chat.Temperature = request.Temperature | ||||||
| @@ -42,51 +104,49 @@ func requestOpenAI2Xunfei(request model.GeneralOpenAIRequest, xunfeiAppId string | |||||||
| 	return &xunfeiRequest | 	return &xunfeiRequest | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func responseXunfei2OpenAI(response *ChatResponse) *openai.TextResponse { | func responseXunfei2OpenAI(response *XunfeiChatResponse) *OpenAITextResponse { | ||||||
| 	if len(response.Payload.Choices.Text) == 0 { | 	if len(response.Payload.Choices.Text) == 0 { | ||||||
| 		response.Payload.Choices.Text = []ChatResponseTextItem{ | 		response.Payload.Choices.Text = []XunfeiChatResponseTextItem{ | ||||||
| 			{ | 			{ | ||||||
| 				Content: "", | 				Content: "", | ||||||
| 			}, | 			}, | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	choice := openai.TextResponseChoice{ | 	choice := OpenAITextResponseChoice{ | ||||||
| 		Index: 0, | 		Index: 0, | ||||||
| 		Message: model.Message{ | 		Message: Message{ | ||||||
| 			Role:    "assistant", | 			Role:    "assistant", | ||||||
| 			Content: response.Payload.Choices.Text[0].Content, | 			Content: response.Payload.Choices.Text[0].Content, | ||||||
| 		}, | 		}, | ||||||
| 		FinishReason: constant.StopFinishReason, | 		FinishReason: stopFinishReason, | ||||||
| 	} | 	} | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := OpenAITextResponse{ | ||||||
| 		Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), |  | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Choices: []openai.TextResponseChoice{choice}, | 		Choices: []OpenAITextResponseChoice{choice}, | ||||||
| 		Usage:   response.Payload.Usage.Text, | 		Usage:   response.Payload.Usage.Text, | ||||||
| 	} | 	} | ||||||
| 	return &fullTextResponse | 	return &fullTextResponse | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func streamResponseXunfei2OpenAI(xunfeiResponse *ChatResponse) *openai.ChatCompletionsStreamResponse { | func streamResponseXunfei2OpenAI(xunfeiResponse *XunfeiChatResponse) *ChatCompletionsStreamResponse { | ||||||
| 	if len(xunfeiResponse.Payload.Choices.Text) == 0 { | 	if len(xunfeiResponse.Payload.Choices.Text) == 0 { | ||||||
| 		xunfeiResponse.Payload.Choices.Text = []ChatResponseTextItem{ | 		xunfeiResponse.Payload.Choices.Text = []XunfeiChatResponseTextItem{ | ||||||
| 			{ | 			{ | ||||||
| 				Content: "", | 				Content: "", | ||||||
| 			}, | 			}, | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice ChatCompletionsStreamResponseChoice | ||||||
| 	choice.Delta.Content = xunfeiResponse.Payload.Choices.Text[0].Content | 	choice.Delta.Content = xunfeiResponse.Payload.Choices.Text[0].Content | ||||||
| 	if xunfeiResponse.Payload.Choices.Status == 2 { | 	if xunfeiResponse.Payload.Choices.Status == 2 { | ||||||
| 		choice.FinishReason = &constant.StopFinishReason | 		choice.FinishReason = &stopFinishReason | ||||||
| 	} | 	} | ||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := ChatCompletionsStreamResponse{ | ||||||
| 		Id:      fmt.Sprintf("chatcmpl-%s", helper.GetUUID()), |  | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Model:   "SparkDesk", | 		Model:   "SparkDesk", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| 	return &response | 	return &response | ||||||
| } | } | ||||||
| @@ -117,14 +177,14 @@ func buildXunfeiAuthUrl(hostUrl string, apiKey, apiSecret string) string { | |||||||
| 	return callUrl | 	return callUrl | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func StreamHandler(c *gin.Context, textRequest model.GeneralOpenAIRequest, appId string, apiSecret string, apiKey string) (*model.ErrorWithStatusCode, *model.Usage) { | func xunfeiStreamHandler(c *gin.Context, textRequest GeneralOpenAIRequest, appId string, apiSecret string, apiKey string) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	domain, authUrl := getXunfeiAuthUrl(c, apiKey, apiSecret, textRequest.Model) | 	domain, authUrl := getXunfeiAuthUrl(c, apiKey, apiSecret) | ||||||
| 	dataChan, stopChan, err := xunfeiMakeRequest(textRequest, domain, authUrl, appId) | 	dataChan, stopChan, err := xunfeiMakeRequest(textRequest, domain, authUrl, appId) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "make xunfei request err", http.StatusInternalServerError), nil | 		return errorWrapper(err, "make xunfei request err", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	common.SetEventStreamHeaders(c) | 	setEventStreamHeaders(c) | ||||||
| 	var usage model.Usage | 	var usage Usage | ||||||
| 	c.Stream(func(w io.Writer) bool { | 	c.Stream(func(w io.Writer) bool { | ||||||
| 		select { | 		select { | ||||||
| 		case xunfeiResponse := <-dataChan: | 		case xunfeiResponse := <-dataChan: | ||||||
| @@ -134,7 +194,7 @@ func StreamHandler(c *gin.Context, textRequest model.GeneralOpenAIRequest, appId | |||||||
| 			response := streamResponseXunfei2OpenAI(&xunfeiResponse) | 			response := streamResponseXunfei2OpenAI(&xunfeiResponse) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error marshalling stream response: " + err.Error()) | 				common.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)}) | ||||||
| @@ -147,15 +207,15 @@ func StreamHandler(c *gin.Context, textRequest model.GeneralOpenAIRequest, appId | |||||||
| 	return nil, &usage | 	return nil, &usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Handler(c *gin.Context, textRequest model.GeneralOpenAIRequest, appId string, apiSecret string, apiKey string) (*model.ErrorWithStatusCode, *model.Usage) { | func xunfeiHandler(c *gin.Context, textRequest GeneralOpenAIRequest, appId string, apiSecret string, apiKey string) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	domain, authUrl := getXunfeiAuthUrl(c, apiKey, apiSecret, textRequest.Model) | 	domain, authUrl := getXunfeiAuthUrl(c, apiKey, apiSecret) | ||||||
| 	dataChan, stopChan, err := xunfeiMakeRequest(textRequest, domain, authUrl, appId) | 	dataChan, stopChan, err := xunfeiMakeRequest(textRequest, domain, authUrl, appId) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "make xunfei request err", http.StatusInternalServerError), nil | 		return errorWrapper(err, "make xunfei request err", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	var usage model.Usage | 	var usage Usage | ||||||
| 	var content string | 	var content string | ||||||
| 	var xunfeiResponse ChatResponse | 	var xunfeiResponse XunfeiChatResponse | ||||||
| 	stop := false | 	stop := false | ||||||
| 	for !stop { | 	for !stop { | ||||||
| 		select { | 		select { | ||||||
| @@ -171,7 +231,7 @@ func Handler(c *gin.Context, textRequest model.GeneralOpenAIRequest, appId strin | |||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	if len(xunfeiResponse.Payload.Choices.Text) == 0 { | 	if len(xunfeiResponse.Payload.Choices.Text) == 0 { | ||||||
| 		xunfeiResponse.Payload.Choices.Text = []ChatResponseTextItem{ | 		xunfeiResponse.Payload.Choices.Text = []XunfeiChatResponseTextItem{ | ||||||
| 			{ | 			{ | ||||||
| 				Content: "", | 				Content: "", | ||||||
| 			}, | 			}, | ||||||
| @@ -182,14 +242,14 @@ func Handler(c *gin.Context, textRequest model.GeneralOpenAIRequest, appId strin | |||||||
| 	response := responseXunfei2OpenAI(&xunfeiResponse) | 	response := responseXunfei2OpenAI(&xunfeiResponse) | ||||||
| 	jsonResponse, err := json.Marshal(response) | 	jsonResponse, err := json.Marshal(response) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
| 	_, _ = c.Writer.Write(jsonResponse) | 	_, _ = c.Writer.Write(jsonResponse) | ||||||
| 	return nil, &usage | 	return nil, &usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func xunfeiMakeRequest(textRequest model.GeneralOpenAIRequest, domain, authUrl, appId string) (chan ChatResponse, chan bool, error) { | func xunfeiMakeRequest(textRequest GeneralOpenAIRequest, domain, authUrl, appId string) (chan XunfeiChatResponse, chan bool, error) { | ||||||
| 	d := websocket.Dialer{ | 	d := websocket.Dialer{ | ||||||
| 		HandshakeTimeout: 5 * time.Second, | 		HandshakeTimeout: 5 * time.Second, | ||||||
| 	} | 	} | ||||||
| @@ -203,26 +263,26 @@ func xunfeiMakeRequest(textRequest model.GeneralOpenAIRequest, domain, authUrl, | |||||||
| 		return nil, nil, err | 		return nil, nil, err | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	dataChan := make(chan ChatResponse) | 	dataChan := make(chan XunfeiChatResponse) | ||||||
| 	stopChan := make(chan bool) | 	stopChan := make(chan bool) | ||||||
| 	go func() { | 	go func() { | ||||||
| 		for { | 		for { | ||||||
| 			_, msg, err := conn.ReadMessage() | 			_, msg, err := conn.ReadMessage() | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error reading stream response: " + err.Error()) | 				common.SysError("error reading stream response: " + err.Error()) | ||||||
| 				break | 				break | ||||||
| 			} | 			} | ||||||
| 			var response ChatResponse | 			var response XunfeiChatResponse | ||||||
| 			err = json.Unmarshal(msg, &response) | 			err = json.Unmarshal(msg, &response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error unmarshalling stream response: " + err.Error()) | 				common.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 { | ||||||
| 					logger.SysError("error closing websocket connection: " + err.Error()) | 					common.SysError("error closing websocket connection: " + err.Error()) | ||||||
| 				} | 				} | ||||||
| 				break | 				break | ||||||
| 			} | 			} | ||||||
| @@ -233,45 +293,20 @@ func xunfeiMakeRequest(textRequest model.GeneralOpenAIRequest, domain, authUrl, | |||||||
| 	return dataChan, stopChan, nil | 	return dataChan, stopChan, nil | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func getAPIVersion(c *gin.Context, modelName string) string { | func getXunfeiAuthUrl(c *gin.Context, apiKey string, apiSecret string) (string, 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 == "" { | ||||||
| 		return apiVersion | 		apiVersion = c.GetString("api_version") | ||||||
| 	} | 	} | ||||||
| 	parts := strings.Split(modelName, "-") | 	if apiVersion == "" { | ||||||
| 	if len(parts) == 2 { | 		apiVersion = "v1.1" | ||||||
| 		apiVersion = parts[1] | 		common.SysLog("api_version not found, use default: " + apiVersion) | ||||||
| 		return apiVersion |  | ||||||
| 
 |  | ||||||
| 	} | 	} | ||||||
| 	apiVersion = c.GetString(common.ConfigKeyAPIVersion) | 	domain := "general" | ||||||
| 	if apiVersion != "" { | 	if apiVersion != "v1.1" { | ||||||
| 		return apiVersion | 		domain += strings.Split(apiVersion, ".")[0] | ||||||
| 	} | 	} | ||||||
| 	apiVersion = "v1.1" |  | ||||||
| 	logger.SysLog("api_version not found, using default: " + apiVersion) |  | ||||||
| 	return apiVersion |  | ||||||
| } |  | ||||||
| 
 |  | ||||||
| // https://www.xfyun.cn/doc/spark/Web.html#_1-%E6%8E%A5%E5%8F%A3%E8%AF%B4%E6%98%8E |  | ||||||
| func apiVersion2domain(apiVersion string) string { |  | ||||||
| 	switch apiVersion { |  | ||||||
| 	case "v1.1": |  | ||||||
| 		return "general" |  | ||||||
| 	case "v2.1": |  | ||||||
| 		return "generalv2" |  | ||||||
| 	case "v3.1": |  | ||||||
| 		return "generalv3" |  | ||||||
| 	case "v3.5": |  | ||||||
| 		return "generalv3.5" |  | ||||||
| 	} |  | ||||||
| 	return "general" + apiVersion |  | ||||||
| } |  | ||||||
| 
 |  | ||||||
| func getXunfeiAuthUrl(c *gin.Context, apiKey string, apiSecret string, modelName string) (string, string) { |  | ||||||
| 	apiVersion := getAPIVersion(c, modelName) |  | ||||||
| 	domain := apiVersion2domain(apiVersion) |  | ||||||
| 	authUrl := buildXunfeiAuthUrl(fmt.Sprintf("wss://spark-api.xf-yun.com/%s/chat", apiVersion), apiKey, apiSecret) | 	authUrl := buildXunfeiAuthUrl(fmt.Sprintf("wss://spark-api.xf-yun.com/%s/chat", apiVersion), apiKey, apiSecret) | ||||||
| 	return domain, authUrl | 	return domain, authUrl | ||||||
| } | } | ||||||
| @@ -1,18 +1,13 @@ | |||||||
| package zhipu | package controller | ||||||
| 
 | 
 | ||||||
| import ( | import ( | ||||||
| 	"bufio" | 	"bufio" | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/golang-jwt/jwt" | 	"github.com/golang-jwt/jwt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"io" | 	"io" | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"sync" | 	"sync" | ||||||
| 	"time" | 	"time" | ||||||
| @@ -23,13 +18,53 @@ import ( | |||||||
| // https://open.bigmodel.cn/api/paas/v3/model-api/chatglm_std/invoke | // https://open.bigmodel.cn/api/paas/v3/model-api/chatglm_std/invoke | ||||||
| // https://open.bigmodel.cn/api/paas/v3/model-api/chatglm_std/sse-invoke | // https://open.bigmodel.cn/api/paas/v3/model-api/chatglm_std/sse-invoke | ||||||
| 
 | 
 | ||||||
|  | type ZhipuMessage struct { | ||||||
|  | 	Role    string `json:"role"` | ||||||
|  | 	Content string `json:"content"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type ZhipuRequest struct { | ||||||
|  | 	Prompt      []ZhipuMessage `json:"prompt"` | ||||||
|  | 	Temperature float64        `json:"temperature,omitempty"` | ||||||
|  | 	TopP        float64        `json:"top_p,omitempty"` | ||||||
|  | 	RequestId   string         `json:"request_id,omitempty"` | ||||||
|  | 	Incremental bool           `json:"incremental,omitempty"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type ZhipuResponseData struct { | ||||||
|  | 	TaskId     string         `json:"task_id"` | ||||||
|  | 	RequestId  string         `json:"request_id"` | ||||||
|  | 	TaskStatus string         `json:"task_status"` | ||||||
|  | 	Choices    []ZhipuMessage `json:"choices"` | ||||||
|  | 	Usage      `json:"usage"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type ZhipuResponse struct { | ||||||
|  | 	Code    int               `json:"code"` | ||||||
|  | 	Msg     string            `json:"msg"` | ||||||
|  | 	Success bool              `json:"success"` | ||||||
|  | 	Data    ZhipuResponseData `json:"data"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type ZhipuStreamMetaResponse struct { | ||||||
|  | 	RequestId  string `json:"request_id"` | ||||||
|  | 	TaskId     string `json:"task_id"` | ||||||
|  | 	TaskStatus string `json:"task_status"` | ||||||
|  | 	Usage      `json:"usage"` | ||||||
|  | } | ||||||
|  | 
 | ||||||
|  | type zhipuTokenData struct { | ||||||
|  | 	Token      string | ||||||
|  | 	ExpiryTime time.Time | ||||||
|  | } | ||||||
|  | 
 | ||||||
| var zhipuTokens sync.Map | var zhipuTokens sync.Map | ||||||
| var expSeconds int64 = 24 * 3600 | var expSeconds int64 = 24 * 3600 | ||||||
| 
 | 
 | ||||||
| func GetToken(apikey string) string { | func getZhipuToken(apikey string) string { | ||||||
| 	data, ok := zhipuTokens.Load(apikey) | 	data, ok := zhipuTokens.Load(apikey) | ||||||
| 	if ok { | 	if ok { | ||||||
| 		tokenData := data.(tokenData) | 		tokenData := data.(zhipuTokenData) | ||||||
| 		if time.Now().Before(tokenData.ExpiryTime) { | 		if time.Now().Before(tokenData.ExpiryTime) { | ||||||
| 			return tokenData.Token | 			return tokenData.Token | ||||||
| 		} | 		} | ||||||
| @@ -37,7 +72,7 @@ func GetToken(apikey string) string { | |||||||
| 
 | 
 | ||||||
| 	split := strings.Split(apikey, ".") | 	split := strings.Split(apikey, ".") | ||||||
| 	if len(split) != 2 { | 	if len(split) != 2 { | ||||||
| 		logger.SysError("invalid zhipu key: " + apikey) | 		common.SysError("invalid zhipu key: " + apikey) | ||||||
| 		return "" | 		return "" | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| @@ -65,7 +100,7 @@ func GetToken(apikey string) string { | |||||||
| 		return "" | 		return "" | ||||||
| 	} | 	} | ||||||
| 
 | 
 | ||||||
| 	zhipuTokens.Store(apikey, tokenData{ | 	zhipuTokens.Store(apikey, zhipuTokenData{ | ||||||
| 		Token:      tokenString, | 		Token:      tokenString, | ||||||
| 		ExpiryTime: expiryTime, | 		ExpiryTime: expiryTime, | ||||||
| 	}) | 	}) | ||||||
| @@ -73,15 +108,26 @@ func GetToken(apikey string) string { | |||||||
| 	return tokenString | 	return tokenString | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func ConvertRequest(request model.GeneralOpenAIRequest) *Request { | func requestOpenAI2Zhipu(request GeneralOpenAIRequest) *ZhipuRequest { | ||||||
| 	messages := make([]Message, 0, len(request.Messages)) | 	messages := make([]ZhipuMessage, 0, len(request.Messages)) | ||||||
| 	for _, message := range request.Messages { | 	for _, message := range request.Messages { | ||||||
| 		messages = append(messages, Message{ | 		if message.Role == "system" { | ||||||
| 			Role:    message.Role, | 			messages = append(messages, ZhipuMessage{ | ||||||
| 			Content: message.StringContent(), | 				Role:    "system", | ||||||
| 		}) | 				Content: message.StringContent(), | ||||||
|  | 			}) | ||||||
|  | 			messages = append(messages, ZhipuMessage{ | ||||||
|  | 				Role:    "user", | ||||||
|  | 				Content: "Okay", | ||||||
|  | 			}) | ||||||
|  | 		} else { | ||||||
|  | 			messages = append(messages, ZhipuMessage{ | ||||||
|  | 				Role:    message.Role, | ||||||
|  | 				Content: message.StringContent(), | ||||||
|  | 			}) | ||||||
|  | 		} | ||||||
| 	} | 	} | ||||||
| 	return &Request{ | 	return &ZhipuRequest{ | ||||||
| 		Prompt:      messages, | 		Prompt:      messages, | ||||||
| 		Temperature: request.Temperature, | 		Temperature: request.Temperature, | ||||||
| 		TopP:        request.TopP, | 		TopP:        request.TopP, | ||||||
| @@ -89,18 +135,18 @@ func ConvertRequest(request model.GeneralOpenAIRequest) *Request { | |||||||
| 	} | 	} | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func responseZhipu2OpenAI(response *Response) *openai.TextResponse { | func responseZhipu2OpenAI(response *ZhipuResponse) *OpenAITextResponse { | ||||||
| 	fullTextResponse := openai.TextResponse{ | 	fullTextResponse := OpenAITextResponse{ | ||||||
| 		Id:      response.Data.TaskId, | 		Id:      response.Data.TaskId, | ||||||
| 		Object:  "chat.completion", | 		Object:  "chat.completion", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Choices: make([]openai.TextResponseChoice, 0, len(response.Data.Choices)), | 		Choices: make([]OpenAITextResponseChoice, 0, len(response.Data.Choices)), | ||||||
| 		Usage:   response.Data.Usage, | 		Usage:   response.Data.Usage, | ||||||
| 	} | 	} | ||||||
| 	for i, choice := range response.Data.Choices { | 	for i, choice := range response.Data.Choices { | ||||||
| 		openaiChoice := openai.TextResponseChoice{ | 		openaiChoice := OpenAITextResponseChoice{ | ||||||
| 			Index: i, | 			Index: i, | ||||||
| 			Message: model.Message{ | 			Message: Message{ | ||||||
| 				Role:    choice.Role, | 				Role:    choice.Role, | ||||||
| 				Content: strings.Trim(choice.Content, "\""), | 				Content: strings.Trim(choice.Content, "\""), | ||||||
| 			}, | 			}, | ||||||
| @@ -114,34 +160,34 @@ func responseZhipu2OpenAI(response *Response) *openai.TextResponse { | |||||||
| 	return &fullTextResponse | 	return &fullTextResponse | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func streamResponseZhipu2OpenAI(zhipuResponse string) *openai.ChatCompletionsStreamResponse { | func streamResponseZhipu2OpenAI(zhipuResponse string) *ChatCompletionsStreamResponse { | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice ChatCompletionsStreamResponseChoice | ||||||
| 	choice.Delta.Content = zhipuResponse | 	choice.Delta.Content = zhipuResponse | ||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := ChatCompletionsStreamResponse{ | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Model:   "chatglm", | 		Model:   "chatglm", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| 	return &response | 	return &response | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func streamMetaResponseZhipu2OpenAI(zhipuResponse *StreamMetaResponse) (*openai.ChatCompletionsStreamResponse, *model.Usage) { | func streamMetaResponseZhipu2OpenAI(zhipuResponse *ZhipuStreamMetaResponse) (*ChatCompletionsStreamResponse, *Usage) { | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice | 	var choice ChatCompletionsStreamResponseChoice | ||||||
| 	choice.Delta.Content = "" | 	choice.Delta.Content = "" | ||||||
| 	choice.FinishReason = &constant.StopFinishReason | 	choice.FinishReason = &stopFinishReason | ||||||
| 	response := openai.ChatCompletionsStreamResponse{ | 	response := ChatCompletionsStreamResponse{ | ||||||
| 		Id:      zhipuResponse.RequestId, | 		Id:      zhipuResponse.RequestId, | ||||||
| 		Object:  "chat.completion.chunk", | 		Object:  "chat.completion.chunk", | ||||||
| 		Created: helper.GetTimestamp(), | 		Created: common.GetTimestamp(), | ||||||
| 		Model:   "chatglm", | 		Model:   "chatglm", | ||||||
| 		Choices: []openai.ChatCompletionsStreamResponseChoice{choice}, | 		Choices: []ChatCompletionsStreamResponseChoice{choice}, | ||||||
| 	} | 	} | ||||||
| 	return &response, &zhipuResponse.Usage | 	return &response, &zhipuResponse.Usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func zhipuStreamHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var usage *model.Usage | 	var usage *Usage | ||||||
| 	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 { | ||||||
| @@ -178,29 +224,29 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 		} | 		} | ||||||
| 		stopChan <- true | 		stopChan <- true | ||||||
| 	}() | 	}() | ||||||
| 	common.SetEventStreamHeaders(c) | 	setEventStreamHeaders(c) | ||||||
| 	c.Stream(func(w io.Writer) bool { | 	c.Stream(func(w io.Writer) bool { | ||||||
| 		select { | 		select { | ||||||
| 		case data := <-dataChan: | 		case data := <-dataChan: | ||||||
| 			response := streamResponseZhipu2OpenAI(data) | 			response := streamResponseZhipu2OpenAI(data) | ||||||
| 			jsonResponse, err := json.Marshal(response) | 			jsonResponse, err := json.Marshal(response) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error marshalling stream response: " + err.Error()) | 				common.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)}) | ||||||
| 			return true | 			return true | ||||||
| 		case data := <-metaChan: | 		case data := <-metaChan: | ||||||
| 			var zhipuResponse StreamMetaResponse | 			var zhipuResponse ZhipuStreamMetaResponse | ||||||
| 			err := json.Unmarshal([]byte(data), &zhipuResponse) | 			err := json.Unmarshal([]byte(data), &zhipuResponse) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError("error unmarshalling stream response: " + err.Error()) | 				common.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 { | ||||||
| 				logger.SysError("error marshalling stream response: " + err.Error()) | 				common.SysError("error marshalling stream response: " + err.Error()) | ||||||
| 				return true | 				return true | ||||||
| 			} | 			} | ||||||
| 			usage = zhipuUsage | 			usage = zhipuUsage | ||||||
| @@ -213,28 +259,28 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC | |||||||
| 	}) | 	}) | ||||||
| 	err := resp.Body.Close() | 	err := resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	return nil, usage | 	return nil, usage | ||||||
| } | } | ||||||
| 
 | 
 | ||||||
| func Handler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { | func zhipuHandler(c *gin.Context, resp *http.Response) (*OpenAIErrorWithStatusCode, *Usage) { | ||||||
| 	var zhipuResponse Response | 	var zhipuResponse ZhipuResponse | ||||||
| 	responseBody, err := io.ReadAll(resp.Body) | 	responseBody, err := io.ReadAll(resp.Body) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = resp.Body.Close() | 	err = resp.Body.Close() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	err = json.Unmarshal(responseBody, &zhipuResponse) | 	err = json.Unmarshal(responseBody, &zhipuResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	if !zhipuResponse.Success { | 	if !zhipuResponse.Success { | ||||||
| 		return &model.ErrorWithStatusCode{ | 		return &OpenAIErrorWithStatusCode{ | ||||||
| 			Error: model.Error{ | 			OpenAIError: OpenAIError{ | ||||||
| 				Message: zhipuResponse.Msg, | 				Message: zhipuResponse.Msg, | ||||||
| 				Type:    "zhipu_error", | 				Type:    "zhipu_error", | ||||||
| 				Param:   "", | 				Param:   "", | ||||||
| @@ -247,7 +293,7 @@ func Handler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, * | |||||||
| 	fullTextResponse.Model = "chatglm" | 	fullTextResponse.Model = "chatglm" | ||||||
| 	jsonResponse, err := json.Marshal(fullTextResponse) | 	jsonResponse, err := json.Marshal(fullTextResponse) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | 		return errorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil | ||||||
| 	} | 	} | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") | 	c.Writer.Header().Set("Content-Type", "application/json") | ||||||
| 	c.Writer.WriteHeader(resp.StatusCode) | 	c.Writer.WriteHeader(resp.StatusCode) | ||||||
| @@ -1,132 +1,384 @@ | |||||||
| package controller | package controller | ||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"bytes" |  | ||||||
| 	"context" |  | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" |  | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/middleware" |  | ||||||
| 	dbmodel "github.com/songquanpeng/one-api/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/monitor" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/controller" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"strconv" | ||||||
|  | 	"strings" | ||||||
|  |  | ||||||
|  | 	"github.com/gin-gonic/gin" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type Message struct { | ||||||
|  | 	Role    string  `json:"role"` | ||||||
|  | 	Content any     `json:"content"` | ||||||
|  | 	Name    *string `json:"name,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type ImageURL struct { | ||||||
|  | 	Url    string `json:"url,omitempty"` | ||||||
|  | 	Detail string `json:"detail,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type TextContent struct { | ||||||
|  | 	Type string `json:"type,omitempty"` | ||||||
|  | 	Text string `json:"text,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type ImageContent struct { | ||||||
|  | 	Type     string    `json:"type,omitempty"` | ||||||
|  | 	ImageURL *ImageURL `json:"image_url,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | const ( | ||||||
|  | 	ContentTypeText     = "text" | ||||||
|  | 	ContentTypeImageURL = "image_url" | ||||||
|  | ) | ||||||
|  |  | ||||||
|  | type OpenAIMessageContent struct { | ||||||
|  | 	Type     string    `json:"type,omitempty"` | ||||||
|  | 	Text     string    `json:"text"` | ||||||
|  | 	ImageURL *ImageURL `json:"image_url,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (m Message) IsStringContent() bool { | ||||||
|  | 	_, ok := m.Content.(string) | ||||||
|  | 	return ok | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (m Message) StringContent() string { | ||||||
|  | 	content, ok := m.Content.(string) | ||||||
|  | 	if ok { | ||||||
|  | 		return content | ||||||
|  | 	} | ||||||
|  | 	contentList, ok := m.Content.([]any) | ||||||
|  | 	if ok { | ||||||
|  | 		var contentStr string | ||||||
|  | 		for _, contentItem := range contentList { | ||||||
|  | 			contentMap, ok := contentItem.(map[string]any) | ||||||
|  | 			if !ok { | ||||||
|  | 				continue | ||||||
|  | 			} | ||||||
|  | 			if contentMap["type"] == ContentTypeText { | ||||||
|  | 				if subStr, ok := contentMap["text"].(string); ok { | ||||||
|  | 					contentStr += subStr | ||||||
|  | 				} | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 		return contentStr | ||||||
|  | 	} | ||||||
|  | 	return "" | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (m Message) ParseContent() []OpenAIMessageContent { | ||||||
|  | 	var contentList []OpenAIMessageContent | ||||||
|  | 	content, ok := m.Content.(string) | ||||||
|  | 	if ok { | ||||||
|  | 		contentList = append(contentList, OpenAIMessageContent{ | ||||||
|  | 			Type: ContentTypeText, | ||||||
|  | 			Text: content, | ||||||
|  | 		}) | ||||||
|  | 		return contentList | ||||||
|  | 	} | ||||||
|  | 	anyList, ok := m.Content.([]any) | ||||||
|  | 	if ok { | ||||||
|  | 		for _, contentItem := range anyList { | ||||||
|  | 			contentMap, ok := contentItem.(map[string]any) | ||||||
|  | 			if !ok { | ||||||
|  | 				continue | ||||||
|  | 			} | ||||||
|  | 			switch contentMap["type"] { | ||||||
|  | 			case ContentTypeText: | ||||||
|  | 				if subStr, ok := contentMap["text"].(string); ok { | ||||||
|  | 					contentList = append(contentList, OpenAIMessageContent{ | ||||||
|  | 						Type: ContentTypeText, | ||||||
|  | 						Text: subStr, | ||||||
|  | 					}) | ||||||
|  | 				} | ||||||
|  | 			case ContentTypeImageURL: | ||||||
|  | 				if subObj, ok := contentMap["image_url"].(map[string]any); ok { | ||||||
|  | 					contentList = append(contentList, OpenAIMessageContent{ | ||||||
|  | 						Type: ContentTypeImageURL, | ||||||
|  | 						ImageURL: &ImageURL{ | ||||||
|  | 							Url: subObj["url"].(string), | ||||||
|  | 						}, | ||||||
|  | 					}) | ||||||
|  | 				} | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 		return contentList | ||||||
|  | 	} | ||||||
|  | 	return nil | ||||||
|  | } | ||||||
|  |  | ||||||
|  | const ( | ||||||
|  | 	RelayModeUnknown = iota | ||||||
|  | 	RelayModeChatCompletions | ||||||
|  | 	RelayModeCompletions | ||||||
|  | 	RelayModeEmbeddings | ||||||
|  | 	RelayModeModerations | ||||||
|  | 	RelayModeImagesGenerations | ||||||
|  | 	RelayModeEdits | ||||||
|  | 	RelayModeAudioSpeech | ||||||
|  | 	RelayModeAudioTranscription | ||||||
|  | 	RelayModeAudioTranslation | ||||||
| ) | ) | ||||||
|  |  | ||||||
| // https://platform.openai.com/docs/api-reference/chat | // https://platform.openai.com/docs/api-reference/chat | ||||||
|  |  | ||||||
| func relay(c *gin.Context, relayMode int) *model.ErrorWithStatusCode { | type ResponseFormat struct { | ||||||
| 	var err *model.ErrorWithStatusCode | 	Type string `json:"type,omitempty"` | ||||||
| 	switch relayMode { | } | ||||||
| 	case constant.RelayModeImagesGenerations: |  | ||||||
| 		err = controller.RelayImageHelper(c, relayMode) | type GeneralOpenAIRequest struct { | ||||||
| 	case constant.RelayModeAudioSpeech: | 	Model            string          `json:"model,omitempty"` | ||||||
| 		fallthrough | 	Messages         []Message       `json:"messages,omitempty"` | ||||||
| 	case constant.RelayModeAudioTranslation: | 	Prompt           any             `json:"prompt,omitempty"` | ||||||
| 		fallthrough | 	Stream           bool            `json:"stream,omitempty"` | ||||||
| 	case constant.RelayModeAudioTranscription: | 	MaxTokens        int             `json:"max_tokens,omitempty"` | ||||||
| 		err = controller.RelayAudioHelper(c, relayMode) | 	Temperature      float64         `json:"temperature,omitempty"` | ||||||
| 	default: | 	TopP             float64         `json:"top_p,omitempty"` | ||||||
| 		err = controller.RelayTextHelper(c) | 	N                int             `json:"n,omitempty"` | ||||||
|  | 	Input            any             `json:"input,omitempty"` | ||||||
|  | 	Instruction      string          `json:"instruction,omitempty"` | ||||||
|  | 	Size             string          `json:"size,omitempty"` | ||||||
|  | 	Functions        any             `json:"functions,omitempty"` | ||||||
|  | 	FrequencyPenalty float64         `json:"frequency_penalty,omitempty"` | ||||||
|  | 	PresencePenalty  float64         `json:"presence_penalty,omitempty"` | ||||||
|  | 	ResponseFormat   *ResponseFormat `json:"response_format,omitempty"` | ||||||
|  | 	Seed             float64         `json:"seed,omitempty"` | ||||||
|  | 	Tools            any             `json:"tools,omitempty"` | ||||||
|  | 	ToolChoice       any             `json:"tool_choice,omitempty"` | ||||||
|  | 	User             string          `json:"user,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | func (r GeneralOpenAIRequest) ParseInput() []string { | ||||||
|  | 	if r.Input == nil { | ||||||
|  | 		return nil | ||||||
| 	} | 	} | ||||||
| 	return err | 	var input []string | ||||||
|  | 	switch r.Input.(type) { | ||||||
|  | 	case string: | ||||||
|  | 		input = []string{r.Input.(string)} | ||||||
|  | 	case []any: | ||||||
|  | 		input = make([]string, 0, len(r.Input.([]any))) | ||||||
|  | 		for _, item := range r.Input.([]any) { | ||||||
|  | 			if str, ok := item.(string); ok { | ||||||
|  | 				input = append(input, str) | ||||||
|  | 			} | ||||||
|  | 		} | ||||||
|  | 	} | ||||||
|  | 	return input | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type ChatRequest struct { | ||||||
|  | 	Model     string    `json:"model"` | ||||||
|  | 	Messages  []Message `json:"messages"` | ||||||
|  | 	MaxTokens int       `json:"max_tokens"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type TextRequest struct { | ||||||
|  | 	Model     string    `json:"model"` | ||||||
|  | 	Messages  []Message `json:"messages"` | ||||||
|  | 	Prompt    string    `json:"prompt"` | ||||||
|  | 	MaxTokens int       `json:"max_tokens"` | ||||||
|  | 	//Stream   bool      `json:"stream"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | // ImageRequest docs: https://platform.openai.com/docs/api-reference/images/create | ||||||
|  | type ImageRequest struct { | ||||||
|  | 	Model          string `json:"model"` | ||||||
|  | 	Prompt         string `json:"prompt" binding:"required"` | ||||||
|  | 	N              int    `json:"n,omitempty"` | ||||||
|  | 	Size           string `json:"size,omitempty"` | ||||||
|  | 	Quality        string `json:"quality,omitempty"` | ||||||
|  | 	ResponseFormat string `json:"response_format,omitempty"` | ||||||
|  | 	Style          string `json:"style,omitempty"` | ||||||
|  | 	User           string `json:"user,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type WhisperJSONResponse struct { | ||||||
|  | 	Text string `json:"text,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type WhisperVerboseJSONResponse struct { | ||||||
|  | 	Task     string    `json:"task,omitempty"` | ||||||
|  | 	Language string    `json:"language,omitempty"` | ||||||
|  | 	Duration float64   `json:"duration,omitempty"` | ||||||
|  | 	Text     string    `json:"text,omitempty"` | ||||||
|  | 	Segments []Segment `json:"segments,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type Segment struct { | ||||||
|  | 	Id               int     `json:"id"` | ||||||
|  | 	Seek             int     `json:"seek"` | ||||||
|  | 	Start            float64 `json:"start"` | ||||||
|  | 	End              float64 `json:"end"` | ||||||
|  | 	Text             string  `json:"text"` | ||||||
|  | 	Tokens           []int   `json:"tokens"` | ||||||
|  | 	Temperature      float64 `json:"temperature"` | ||||||
|  | 	AvgLogprob       float64 `json:"avg_logprob"` | ||||||
|  | 	CompressionRatio float64 `json:"compression_ratio"` | ||||||
|  | 	NoSpeechProb     float64 `json:"no_speech_prob"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type TextToSpeechRequest struct { | ||||||
|  | 	Model          string  `json:"model" binding:"required"` | ||||||
|  | 	Input          string  `json:"input" binding:"required"` | ||||||
|  | 	Voice          string  `json:"voice" binding:"required"` | ||||||
|  | 	Speed          float64 `json:"speed"` | ||||||
|  | 	ResponseFormat string  `json:"response_format"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type Usage struct { | ||||||
|  | 	PromptTokens     int `json:"prompt_tokens"` | ||||||
|  | 	CompletionTokens int `json:"completion_tokens"` | ||||||
|  | 	TotalTokens      int `json:"total_tokens"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type OpenAIError struct { | ||||||
|  | 	Message string `json:"message"` | ||||||
|  | 	Type    string `json:"type"` | ||||||
|  | 	Param   string `json:"param"` | ||||||
|  | 	Code    any    `json:"code"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type OpenAIErrorWithStatusCode struct { | ||||||
|  | 	OpenAIError | ||||||
|  | 	StatusCode int `json:"status_code"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type TextResponse struct { | ||||||
|  | 	Choices []OpenAITextResponseChoice `json:"choices"` | ||||||
|  | 	Usage   `json:"usage"` | ||||||
|  | 	Error   OpenAIError `json:"error"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type OpenAITextResponseChoice struct { | ||||||
|  | 	Index        int `json:"index"` | ||||||
|  | 	Message      `json:"message"` | ||||||
|  | 	FinishReason string `json:"finish_reason"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type OpenAITextResponse struct { | ||||||
|  | 	Id      string                     `json:"id"` | ||||||
|  | 	Model   string                     `json:"model,omitempty"` | ||||||
|  | 	Object  string                     `json:"object"` | ||||||
|  | 	Created int64                      `json:"created"` | ||||||
|  | 	Choices []OpenAITextResponseChoice `json:"choices"` | ||||||
|  | 	Usage   `json:"usage"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type OpenAIEmbeddingResponseItem struct { | ||||||
|  | 	Object    string    `json:"object"` | ||||||
|  | 	Index     int       `json:"index"` | ||||||
|  | 	Embedding []float64 `json:"embedding"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type OpenAIEmbeddingResponse struct { | ||||||
|  | 	Object string                        `json:"object"` | ||||||
|  | 	Data   []OpenAIEmbeddingResponseItem `json:"data"` | ||||||
|  | 	Model  string                        `json:"model"` | ||||||
|  | 	Usage  `json:"usage"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type ImageResponse struct { | ||||||
|  | 	Created int `json:"created"` | ||||||
|  | 	Data    []struct { | ||||||
|  | 		Url string `json:"url"` | ||||||
|  | 	} | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type ChatCompletionsStreamResponseChoice struct { | ||||||
|  | 	Delta struct { | ||||||
|  | 		Content string `json:"content"` | ||||||
|  | 	} `json:"delta"` | ||||||
|  | 	FinishReason *string `json:"finish_reason,omitempty"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type ChatCompletionsStreamResponse struct { | ||||||
|  | 	Id      string                                `json:"id"` | ||||||
|  | 	Object  string                                `json:"object"` | ||||||
|  | 	Created int64                                 `json:"created"` | ||||||
|  | 	Model   string                                `json:"model"` | ||||||
|  | 	Choices []ChatCompletionsStreamResponseChoice `json:"choices"` | ||||||
|  | } | ||||||
|  |  | ||||||
|  | type CompletionsStreamResponse struct { | ||||||
|  | 	Choices []struct { | ||||||
|  | 		Text         string `json:"text"` | ||||||
|  | 		FinishReason string `json:"finish_reason"` | ||||||
|  | 	} `json:"choices"` | ||||||
| } | } | ||||||
|  |  | ||||||
| func Relay(c *gin.Context) { | func Relay(c *gin.Context) { | ||||||
| 	ctx := c.Request.Context() | 	relayMode := RelayModeUnknown | ||||||
| 	relayMode := constant.Path2RelayMode(c.Request.URL.Path) | 	if strings.HasPrefix(c.Request.URL.Path, "/v1/chat/completions") { | ||||||
| 	if config.DebugEnabled { | 		relayMode = RelayModeChatCompletions | ||||||
| 		requestBody, _ := common.GetRequestBody(c) | 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/completions") { | ||||||
| 		logger.Debugf(ctx, "request body: %s", string(requestBody)) | 		relayMode = RelayModeCompletions | ||||||
|  | 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/embeddings") { | ||||||
|  | 		relayMode = RelayModeEmbeddings | ||||||
|  | 	} else if strings.HasSuffix(c.Request.URL.Path, "embeddings") { | ||||||
|  | 		relayMode = RelayModeEmbeddings | ||||||
|  | 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/moderations") { | ||||||
|  | 		relayMode = RelayModeModerations | ||||||
|  | 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/images/generations") { | ||||||
|  | 		relayMode = RelayModeImagesGenerations | ||||||
|  | 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/edits") { | ||||||
|  | 		relayMode = RelayModeEdits | ||||||
|  | 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/audio/speech") { | ||||||
|  | 		relayMode = RelayModeAudioSpeech | ||||||
|  | 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/audio/transcriptions") { | ||||||
|  | 		relayMode = RelayModeAudioTranscription | ||||||
|  | 	} else if strings.HasPrefix(c.Request.URL.Path, "/v1/audio/translations") { | ||||||
|  | 		relayMode = RelayModeAudioTranslation | ||||||
| 	} | 	} | ||||||
| 	channelId := c.GetInt("channel_id") | 	var err *OpenAIErrorWithStatusCode | ||||||
| 	bizErr := relay(c, relayMode) | 	switch relayMode { | ||||||
| 	if bizErr == nil { | 	case RelayModeImagesGenerations: | ||||||
| 		monitor.Emit(channelId, true) | 		err = relayImageHelper(c, relayMode) | ||||||
| 		return | 	case RelayModeAudioSpeech: | ||||||
|  | 		fallthrough | ||||||
|  | 	case RelayModeAudioTranslation: | ||||||
|  | 		fallthrough | ||||||
|  | 	case RelayModeAudioTranscription: | ||||||
|  | 		err = relayAudioHelper(c, relayMode) | ||||||
|  | 	default: | ||||||
|  | 		err = relayTextHelper(c, relayMode) | ||||||
| 	} | 	} | ||||||
| 	lastFailedChannelId := channelId | 	if err != nil { | ||||||
| 	channelName := c.GetString("channel_name") | 		requestId := c.GetString(common.RequestIdKey) | ||||||
| 	group := c.GetString("group") | 		retryTimesStr := c.Query("retry") | ||||||
| 	originalModel := c.GetString("original_model") | 		retryTimes, _ := strconv.Atoi(retryTimesStr) | ||||||
| 	go processChannelRelayError(ctx, channelId, channelName, bizErr) | 		if retryTimesStr == "" { | ||||||
| 	requestId := c.GetString(logger.RequestIdKey) | 			retryTimes = common.RetryTimes | ||||||
| 	retryTimes := config.RetryTimes |  | ||||||
| 	if !shouldRetry(c, bizErr.StatusCode) { |  | ||||||
| 		logger.Errorf(ctx, "relay error happen, status code is %d, won't retry in this case", bizErr.StatusCode) |  | ||||||
| 		retryTimes = 0 |  | ||||||
| 	} |  | ||||||
| 	for i := retryTimes; i > 0; i-- { |  | ||||||
| 		channel, err := dbmodel.CacheGetRandomSatisfiedChannel(group, originalModel, i != retryTimes) |  | ||||||
| 		if err != nil { |  | ||||||
| 			logger.Errorf(ctx, "CacheGetRandomSatisfiedChannel failed: %w", err) |  | ||||||
| 			break |  | ||||||
| 		} | 		} | ||||||
| 		logger.Infof(ctx, "using channel #%d to retry (remain times %d)", channel.Id, i) | 		if retryTimes > 0 { | ||||||
| 		if channel.Id == lastFailedChannelId { | 			c.Redirect(http.StatusTemporaryRedirect, fmt.Sprintf("%s?retry=%d", c.Request.URL.Path, retryTimes-1)) | ||||||
| 			continue | 		} else { | ||||||
| 		} | 			if err.StatusCode == http.StatusTooManyRequests { | ||||||
| 		middleware.SetupContextForSelectedChannel(c, channel, originalModel) | 				err.OpenAIError.Message = "当前分组上游负载已饱和,请稍后再试" | ||||||
| 		requestBody, err := common.GetRequestBody(c) | 			} | ||||||
| 		c.Request.Body = io.NopCloser(bytes.NewBuffer(requestBody)) | 			err.OpenAIError.Message = common.MessageWithRequestId(err.OpenAIError.Message, requestId) | ||||||
| 		bizErr = relay(c, relayMode) | 			c.JSON(err.StatusCode, gin.H{ | ||||||
| 		if bizErr == nil { | 				"error": err.OpenAIError, | ||||||
| 			return | 			}) | ||||||
| 		} | 		} | ||||||
| 		channelId := c.GetInt("channel_id") | 		channelId := c.GetInt("channel_id") | ||||||
| 		lastFailedChannelId = channelId | 		common.LogError(c.Request.Context(), fmt.Sprintf("relay error (channel #%d): %s", channelId, err.Message)) | ||||||
| 		channelName := c.GetString("channel_name") | 		// https://platform.openai.com/docs/guides/error-codes/api-errors | ||||||
| 		go processChannelRelayError(ctx, channelId, channelName, bizErr) | 		if shouldDisableChannel(&err.OpenAIError, err.StatusCode) { | ||||||
| 	} | 			channelId := c.GetInt("channel_id") | ||||||
| 	if bizErr != nil { | 			channelName := c.GetString("channel_name") | ||||||
| 		if bizErr.StatusCode == http.StatusTooManyRequests { | 			disableChannel(channelId, channelName, err.Message) | ||||||
| 			bizErr.Error.Message = "当前分组上游负载已饱和,请稍后再试" |  | ||||||
| 		} | 		} | ||||||
| 		bizErr.Error.Message = helper.MessageWithRequestId(bizErr.Error.Message, requestId) |  | ||||||
| 		c.JSON(bizErr.StatusCode, gin.H{ |  | ||||||
| 			"error": bizErr.Error, |  | ||||||
| 		}) |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func shouldRetry(c *gin.Context, statusCode int) bool { |  | ||||||
| 	if _, ok := c.Get("specific_channel_id"); ok { |  | ||||||
| 		return false |  | ||||||
| 	} |  | ||||||
| 	if statusCode == http.StatusTooManyRequests { |  | ||||||
| 		return true |  | ||||||
| 	} |  | ||||||
| 	if statusCode/100 == 5 { |  | ||||||
| 		return true |  | ||||||
| 	} |  | ||||||
| 	if statusCode == http.StatusBadRequest { |  | ||||||
| 		return false |  | ||||||
| 	} |  | ||||||
| 	if statusCode/100 == 2 { |  | ||||||
| 		return false |  | ||||||
| 	} |  | ||||||
| 	return true |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func processChannelRelayError(ctx context.Context, channelId int, channelName string, err *model.ErrorWithStatusCode) { |  | ||||||
| 	logger.Errorf(ctx, "relay error (channel #%d): %s", channelId, err.Message) |  | ||||||
| 	// https://platform.openai.com/docs/guides/error-codes/api-errors |  | ||||||
| 	if util.ShouldDisableChannel(&err.Error, err.StatusCode) { |  | ||||||
| 		monitor.DisableChannel(channelId, channelName, err.Message) |  | ||||||
| 	} else { |  | ||||||
| 		monitor.Emit(channelId, false) |  | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func RelayNotImplemented(c *gin.Context) { | func RelayNotImplemented(c *gin.Context) { | ||||||
| 	err := model.Error{ | 	err := OpenAIError{ | ||||||
| 		Message: "API not implemented", | 		Message: "API not implemented", | ||||||
| 		Type:    "one_api_error", | 		Type:    "one_api_error", | ||||||
| 		Param:   "", | 		Param:   "", | ||||||
| @@ -138,7 +390,7 @@ func RelayNotImplemented(c *gin.Context) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func RelayNotFound(c *gin.Context) { | func RelayNotFound(c *gin.Context) { | ||||||
| 	err := model.Error{ | 	err := OpenAIError{ | ||||||
| 		Message: fmt.Sprintf("Invalid URL (%s %s)", c.Request.Method, c.Request.URL.Path), | 		Message: fmt.Sprintf("Invalid URL (%s %s)", c.Request.Method, c.Request.URL.Path), | ||||||
| 		Type:    "invalid_request_error", | 		Type:    "invalid_request_error", | ||||||
| 		Param:   "", | 		Param:   "", | ||||||
|   | |||||||
| @@ -2,11 +2,9 @@ package controller | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -16,7 +14,7 @@ func GetAllTokens(c *gin.Context) { | |||||||
| 	if p < 0 { | 	if p < 0 { | ||||||
| 		p = 0 | 		p = 0 | ||||||
| 	} | 	} | ||||||
| 	tokens, err := model.GetAllUserTokens(userId, p*config.ItemsPerPage, config.ItemsPerPage) | 	tokens, err := model.GetAllUserTokens(userId, p*common.ItemsPerPage, common.ItemsPerPage) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -121,9 +119,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:            helper.GenerateKey(), | 		Key:            common.GenerateKey(), | ||||||
| 		CreatedTime:    helper.GetTimestamp(), | 		CreatedTime:    common.GetTimestamp(), | ||||||
| 		AccessedTime:   helper.GetTimestamp(), | 		AccessedTime:   common.GetTimestamp(), | ||||||
| 		ExpiredTime:    token.ExpiredTime, | 		ExpiredTime:    token.ExpiredTime, | ||||||
| 		RemainQuota:    token.RemainQuota, | 		RemainQuota:    token.RemainQuota, | ||||||
| 		UnlimitedQuota: token.UnlimitedQuota, | 		UnlimitedQuota: token.UnlimitedQuota, | ||||||
| @@ -189,7 +187,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 <= helper.GetTimestamp() && cleanToken.ExpiredTime != -1 { | 		if cleanToken.Status == common.TokenStatusExpired && cleanToken.ExpiredTime <= common.GetTimestamp() && cleanToken.ExpiredTime != -1 { | ||||||
| 			c.JSON(http.StatusOK, gin.H{ | 			c.JSON(http.StatusOK, gin.H{ | ||||||
| 				"success": false, | 				"success": false, | ||||||
| 				"message": "令牌已过期,无法启用,请先修改令牌过期时间,或者设置为永不过期", | 				"message": "令牌已过期,无法启用,请先修改令牌过期时间,或者设置为永不过期", | ||||||
|   | |||||||
| @@ -3,11 +3,9 @@ package controller | |||||||
| import ( | import ( | ||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"time" | 	"time" | ||||||
|  |  | ||||||
| @@ -21,7 +19,7 @@ type LoginRequest struct { | |||||||
| } | } | ||||||
|  |  | ||||||
| func Login(c *gin.Context) { | func Login(c *gin.Context) { | ||||||
| 	if !config.PasswordLoginEnabled { | 	if !common.PasswordLoginEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员关闭了密码登录", | 			"message": "管理员关闭了密码登录", | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -108,14 +106,14 @@ func Logout(c *gin.Context) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func Register(c *gin.Context) { | func Register(c *gin.Context) { | ||||||
| 	if !config.RegisterEnabled { | 	if !common.RegisterEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员关闭了新用户注册", | 			"message": "管理员关闭了新用户注册", | ||||||
| 			"success": false, | 			"success": false, | ||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if !config.PasswordRegisterEnabled { | 	if !common.PasswordRegisterEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员关闭了通过密码进行注册,请使用第三方账户验证的形式进行注册", | 			"message": "管理员关闭了通过密码进行注册,请使用第三方账户验证的形式进行注册", | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -138,7 +136,7 @@ func Register(c *gin.Context) { | |||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if config.EmailVerificationEnabled { | 	if common.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, | ||||||
| @@ -162,7 +160,7 @@ func Register(c *gin.Context) { | |||||||
| 		DisplayName: user.Username, | 		DisplayName: user.Username, | ||||||
| 		InviterId:   inviterId, | 		InviterId:   inviterId, | ||||||
| 	} | 	} | ||||||
| 	if config.EmailVerificationEnabled { | 	if common.EmailVerificationEnabled { | ||||||
| 		cleanUser.Email = user.Email | 		cleanUser.Email = user.Email | ||||||
| 	} | 	} | ||||||
| 	if err := cleanUser.Insert(inviterId); err != nil { | 	if err := cleanUser.Insert(inviterId); err != nil { | ||||||
| @@ -184,7 +182,7 @@ func GetAllUsers(c *gin.Context) { | |||||||
| 	if p < 0 { | 	if p < 0 { | ||||||
| 		p = 0 | 		p = 0 | ||||||
| 	} | 	} | ||||||
| 	users, err := model.GetAllUsers(p*config.ItemsPerPage, config.ItemsPerPage) | 	users, err := model.GetAllUsers(p*common.ItemsPerPage, common.ItemsPerPage) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -284,7 +282,7 @@ func GenerateAccessToken(c *gin.Context) { | |||||||
| 		}) | 		}) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	user.AccessToken = helper.GetUUID() | 	user.AccessToken = common.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{ | ||||||
| @@ -321,7 +319,7 @@ func GetAffCode(c *gin.Context) { | |||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if user.AffCode == "" { | 	if user.AffCode == "" { | ||||||
| 		user.AffCode = helper.GetRandomString(4) | 		user.AffCode = common.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, | ||||||
| @@ -728,7 +726,7 @@ func EmailBind(c *gin.Context) { | |||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	if user.Role == common.RoleRootUser { | 	if user.Role == common.RoleRootUser { | ||||||
| 		config.RootUserEmail = email | 		common.RootUserEmail = email | ||||||
| 	} | 	} | ||||||
| 	c.JSON(http.StatusOK, gin.H{ | 	c.JSON(http.StatusOK, gin.H{ | ||||||
| 		"success": true, | 		"success": true, | ||||||
|   | |||||||
| @@ -5,10 +5,9 @@ import ( | |||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
| @@ -23,11 +22,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", config.WeChatServerAddress, code), nil) | 	req, err := http.NewRequest("GET", fmt.Sprintf("%s/api/wechat/user?code=%s", common.WeChatServerAddress, code), nil) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return "", err | 		return "", err | ||||||
| 	} | 	} | ||||||
| 	req.Header.Set("Authorization", config.WeChatServerToken) | 	req.Header.Set("Authorization", common.WeChatServerToken) | ||||||
| 	client := http.Client{ | 	client := http.Client{ | ||||||
| 		Timeout: 5 * time.Second, | 		Timeout: 5 * time.Second, | ||||||
| 	} | 	} | ||||||
| @@ -51,7 +50,7 @@ func getWeChatIdByCode(code string) (string, error) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func WeChatAuth(c *gin.Context) { | func WeChatAuth(c *gin.Context) { | ||||||
| 	if !config.WeChatAuthEnabled { | 	if !common.WeChatAuthEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员未开启通过微信登录以及注册", | 			"message": "管理员未开启通过微信登录以及注册", | ||||||
| 			"success": false, | 			"success": false, | ||||||
| @@ -80,7 +79,7 @@ func WeChatAuth(c *gin.Context) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	} else { | 	} else { | ||||||
| 		if config.RegisterEnabled { | 		if common.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 | ||||||
| @@ -113,7 +112,7 @@ func WeChatAuth(c *gin.Context) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func WeChatBind(c *gin.Context) { | func WeChatBind(c *gin.Context) { | ||||||
| 	if !config.WeChatAuthEnabled { | 	if !common.WeChatAuthEnabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"message": "管理员未开启通过微信登录以及注册", | 			"message": "管理员未开启通过微信登录以及注册", | ||||||
| 			"success": false, | 			"success": false, | ||||||
|   | |||||||
							
								
								
									
										2
									
								
								go.mod
									
									
									
									
									
								
							
							
						
						
									
										2
									
								
								go.mod
									
									
									
									
									
								
							| @@ -1,4 +1,4 @@ | |||||||
| module github.com/songquanpeng/one-api | module one-api | ||||||
|  |  | ||||||
| // +heroku goVersion go1.18 | // +heroku goVersion go1.18 | ||||||
| go 1.18 | go 1.18 | ||||||
|   | |||||||
| @@ -330,6 +330,7 @@ | |||||||
|   "通常和邮箱地址保持一致": "Usually consistent with the email address", |   "通常和邮箱地址保持一致": "Usually consistent with the email address", | ||||||
|   "SMTP 访问凭证": "SMTP Access Credential", |   "SMTP 访问凭证": "SMTP Access Credential", | ||||||
|   "敏感信息不会发送到前端显示": "Sensitive information will not be displayed in the frontend", |   "敏感信息不会发送到前端显示": "Sensitive information will not be displayed in the frontend", | ||||||
|  |   "使用 SMTP LOGIN 认证方式": "Use LOGIN as SMTP authentication method", | ||||||
|   "保存 SMTP 设置": "Save SMTP Settings", |   "保存 SMTP 设置": "Save SMTP Settings", | ||||||
|   "配置 GitHub OAuth App": "Configure GitHub OAuth App", |   "配置 GitHub OAuth App": "Configure GitHub OAuth App", | ||||||
|   "用以支持通过 GitHub 进行登录注册": "To support login & registration via GitHub", |   "用以支持通过 GitHub 进行登录注册": "To support login & registration via GitHub", | ||||||
| @@ -456,7 +457,6 @@ | |||||||
|   "已绑定的邮箱账户": "Email Account Bound", |   "已绑定的邮箱账户": "Email Account Bound", | ||||||
|   "用户信息更新成功!": "User information updated successfully!", |   "用户信息更新成功!": "User information updated successfully!", | ||||||
|   "模型倍率 %.2f,分组倍率 %.2f": "model rate %.2f, group rate %.2f", |   "模型倍率 %.2f,分组倍率 %.2f": "model rate %.2f, group rate %.2f", | ||||||
|   "模型倍率 %.2f,分组倍率 %.2f,补全倍率 %.2f": "model rate %.2f, group rate %.2f, completion rate %.2f", |  | ||||||
|   "使用明细(总消耗额度:{renderQuota(stat.quota)})": "Usage Details (Total Consumption Quota: {renderQuota(stat.quota)})", |   "使用明细(总消耗额度:{renderQuota(stat.quota)})": "Usage Details (Total Consumption Quota: {renderQuota(stat.quota)})", | ||||||
|   "用户名称": "User Name", |   "用户名称": "User Name", | ||||||
|   "令牌名称": "Token Name", |   "令牌名称": "Token Name", | ||||||
|   | |||||||
							
								
								
									
										62
									
								
								main.go
									
									
									
									
									
								
							
							
						
						
									
										62
									
								
								main.go
									
									
									
									
									
								
							| @@ -6,15 +6,11 @@ import ( | |||||||
| 	"github.com/gin-contrib/sessions" | 	"github.com/gin-contrib/sessions" | ||||||
| 	"github.com/gin-contrib/sessions/cookie" | 	"github.com/gin-contrib/sessions/cookie" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" | 	"one-api/common" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" | 	"one-api/controller" | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" | 	"one-api/middleware" | ||||||
| 	"github.com/songquanpeng/one-api/common/message" | 	"one-api/model" | ||||||
| 	"github.com/songquanpeng/one-api/controller" | 	"one-api/router" | ||||||
| 	"github.com/songquanpeng/one-api/middleware" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/router" |  | ||||||
| 	"os" | 	"os" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| ) | ) | ||||||
| @@ -23,72 +19,68 @@ import ( | |||||||
| var buildFS embed.FS | var buildFS embed.FS | ||||||
|  |  | ||||||
| func main() { | func main() { | ||||||
| 	logger.SetupLogger() | 	common.SetupLogger() | ||||||
| 	logger.SysLog(fmt.Sprintf("One API %s started", common.Version)) | 	common.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 config.DebugEnabled { | 	if common.DebugEnabled { | ||||||
| 		logger.SysLog("running in debug mode") | 		common.SysLog("running in debug mode") | ||||||
| 	} | 	} | ||||||
| 	// Initialize SQL Database | 	// Initialize SQL Database | ||||||
| 	err := model.InitDB() | 	err := model.InitDB() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.FatalLog("failed to initialize database: " + err.Error()) | 		common.FatalLog("failed to initialize database: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	defer func() { | 	defer func() { | ||||||
| 		err := model.CloseDB() | 		err := model.CloseDB() | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			logger.FatalLog("failed to close database: " + err.Error()) | 			common.FatalLog("failed to close database: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 	}() | 	}() | ||||||
|  |  | ||||||
| 	// Initialize Redis | 	// Initialize Redis | ||||||
| 	err = common.InitRedisClient() | 	err = common.InitRedisClient() | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.FatalLog("failed to initialize Redis: " + err.Error()) | 		common.FatalLog("failed to initialize Redis: " + err.Error()) | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	// Initialize options | 	// Initialize options | ||||||
| 	model.InitOptionMap() | 	model.InitOptionMap() | ||||||
| 	logger.SysLog(fmt.Sprintf("using theme %s", config.Theme)) | 	common.SysLog(fmt.Sprintf("using theme %s", common.Theme)) | ||||||
| 	if common.RedisEnabled { | 	if common.RedisEnabled { | ||||||
| 		// for compatibility with old versions | 		// for compatibility with old versions | ||||||
| 		config.MemoryCacheEnabled = true | 		common.MemoryCacheEnabled = true | ||||||
| 	} | 	} | ||||||
| 	if config.MemoryCacheEnabled { | 	if common.MemoryCacheEnabled { | ||||||
| 		logger.SysLog("memory cache enabled") | 		common.SysLog("memory cache enabled") | ||||||
| 		logger.SysError(fmt.Sprintf("sync frequency: %d seconds", config.SyncFrequency)) | 		common.SysError(fmt.Sprintf("sync frequency: %d seconds", common.SyncFrequency)) | ||||||
| 		model.InitChannelCache() | 		model.InitChannelCache() | ||||||
| 	} | 	} | ||||||
| 	if config.MemoryCacheEnabled { | 	if common.MemoryCacheEnabled { | ||||||
| 		go model.SyncOptions(config.SyncFrequency) | 		go model.SyncOptions(common.SyncFrequency) | ||||||
| 		go model.SyncChannelCache(config.SyncFrequency) | 		go model.SyncChannelCache(common.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 { | ||||||
| 			logger.FatalLog("failed to parse CHANNEL_UPDATE_FREQUENCY: " + err.Error()) | 			common.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 { | ||||||
| 			logger.FatalLog("failed to parse CHANNEL_TEST_FREQUENCY: " + err.Error()) | 			common.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" { | ||||||
| 		config.BatchUpdateEnabled = true | 		common.BatchUpdateEnabled = true | ||||||
| 		logger.SysLog("batch update enabled with interval " + strconv.Itoa(config.BatchUpdateInterval) + "s") | 		common.SysLog("batch update enabled with interval " + strconv.Itoa(common.BatchUpdateInterval) + "s") | ||||||
| 		model.InitBatchUpdater() | 		model.InitBatchUpdater() | ||||||
| 	} | 	} | ||||||
| 	if config.EnableMetric { | 	controller.InitTokenEncoders() | ||||||
| 		logger.SysLog("metric enabled, will disable channel if too much request failed") |  | ||||||
| 	} |  | ||||||
| 	openai.InitTokenEncoders() |  | ||||||
| 	_ = message.SendMessage("One API", "", fmt.Sprintf("One API %s started", common.Version)) |  | ||||||
|  |  | ||||||
| 	// Initialize HTTP server | 	// Initialize HTTP server | ||||||
| 	server := gin.New() | 	server := gin.New() | ||||||
| @@ -98,7 +90,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(config.SessionSecret)) | 	store := cookie.NewStore([]byte(common.SessionSecret)) | ||||||
| 	server.Use(sessions.Sessions("session", store)) | 	server.Use(sessions.Sessions("session", store)) | ||||||
|  |  | ||||||
| 	router.SetRouter(server, buildFS) | 	router.SetRouter(server, buildFS) | ||||||
| @@ -108,6 +100,6 @@ func main() { | |||||||
| 	} | 	} | ||||||
| 	err = server.Run(":" + port) | 	err = server.Run(":" + port) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.FatalLog("failed to start HTTP server: " + err.Error()) | 		common.FatalLog("failed to start HTTP server: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -3,10 +3,9 @@ package middleware | |||||||
| import ( | import ( | ||||||
| 	"github.com/gin-contrib/sessions" | 	"github.com/gin-contrib/sessions" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/blacklist" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -43,14 +42,11 @@ func authHelper(c *gin.Context, minRole int) { | |||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	if status.(int) == common.UserStatusDisabled || blacklist.IsUserBanned(id.(int)) { | 	if status.(int) == common.UserStatusDisabled { | ||||||
| 		c.JSON(http.StatusOK, gin.H{ | 		c.JSON(http.StatusOK, gin.H{ | ||||||
| 			"success": false, | 			"success": false, | ||||||
| 			"message": "用户已被封禁", | 			"message": "用户已被封禁", | ||||||
| 		}) | 		}) | ||||||
| 		session := sessions.Default(c) |  | ||||||
| 		session.Clear() |  | ||||||
| 		_ = session.Save() |  | ||||||
| 		c.Abort() | 		c.Abort() | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| @@ -103,7 +99,7 @@ func TokenAuth() func(c *gin.Context) { | |||||||
| 			abortWithMessage(c, http.StatusInternalServerError, err.Error()) | 			abortWithMessage(c, http.StatusInternalServerError, err.Error()) | ||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| 		if !userEnabled || blacklist.IsUserBanned(token.UserId) { | 		if !userEnabled { | ||||||
| 			abortWithMessage(c, http.StatusForbidden, "用户已被封禁") | 			abortWithMessage(c, http.StatusForbidden, "用户已被封禁") | ||||||
| 			return | 			return | ||||||
| 		} | 		} | ||||||
| @@ -112,7 +108,7 @@ func TokenAuth() func(c *gin.Context) { | |||||||
| 		c.Set("token_name", token.Name) | 		c.Set("token_name", token.Name) | ||||||
| 		if len(parts) > 1 { | 		if len(parts) > 1 { | ||||||
| 			if model.IsAdmin(token.UserId) { | 			if model.IsAdmin(token.UserId) { | ||||||
| 				c.Set("specific_channel_id", parts[1]) | 				c.Set("channelId", parts[1]) | ||||||
| 			} else { | 			} else { | ||||||
| 				abortWithMessage(c, http.StatusForbidden, "普通用户不支持指定渠道") | 				abortWithMessage(c, http.StatusForbidden, "普通用户不支持指定渠道") | ||||||
| 				return | 				return | ||||||
|   | |||||||
| @@ -2,10 +2,9 @@ package middleware | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
|  | 	"one-api/model" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
|  |  | ||||||
| @@ -21,9 +20,8 @@ func Distribute() func(c *gin.Context) { | |||||||
| 		userId := c.GetInt("id") | 		userId := c.GetInt("id") | ||||||
| 		userGroup, _ := model.CacheGetUserGroup(userId) | 		userGroup, _ := model.CacheGetUserGroup(userId) | ||||||
| 		c.Set("group", userGroup) | 		c.Set("group", userGroup) | ||||||
| 		var requestModel string |  | ||||||
| 		var channel *model.Channel | 		var channel *model.Channel | ||||||
| 		channelId, ok := c.Get("specific_channel_id") | 		channelId, ok := c.Get("channelId") | ||||||
| 		if ok { | 		if ok { | ||||||
| 			id, err := strconv.Atoi(channelId.(string)) | 			id, err := strconv.Atoi(channelId.(string)) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| @@ -67,46 +65,35 @@ func Distribute() func(c *gin.Context) { | |||||||
| 					modelRequest.Model = "whisper-1" | 					modelRequest.Model = "whisper-1" | ||||||
| 				} | 				} | ||||||
| 			} | 			} | ||||||
| 			requestModel = modelRequest.Model | 			channel, err = model.CacheGetRandomSatisfiedChannel(userGroup, modelRequest.Model) | ||||||
| 			channel, err = model.CacheGetRandomSatisfiedChannel(userGroup, modelRequest.Model, false) |  | ||||||
| 			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 { | ||||||
| 					logger.SysError(fmt.Sprintf("渠道不存在:%d", channel.Id)) | 					common.SysError(fmt.Sprintf("渠道不存在:%d", channel.Id)) | ||||||
| 					message = "数据库一致性已被破坏,请联系管理员" | 					message = "数据库一致性已被破坏,请联系管理员" | ||||||
| 				} | 				} | ||||||
| 				abortWithMessage(c, http.StatusServiceUnavailable, message) | 				abortWithMessage(c, http.StatusServiceUnavailable, message) | ||||||
| 				return | 				return | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 		SetupContextForSelectedChannel(c, channel, requestModel) | 		c.Set("channel", channel.Type) | ||||||
|  | 		c.Set("channel_id", channel.Id) | ||||||
|  | 		c.Set("channel_name", channel.Name) | ||||||
|  | 		c.Set("model_mapping", channel.GetModelMapping()) | ||||||
|  | 		c.Request.Header.Set("Authorization", fmt.Sprintf("Bearer %s", channel.Key)) | ||||||
|  | 		c.Set("base_url", channel.GetBaseURL()) | ||||||
|  | 		switch channel.Type { | ||||||
|  | 		case common.ChannelTypeAzure: | ||||||
|  | 			c.Set("api_version", channel.Other) | ||||||
|  | 		case common.ChannelTypeXunfei: | ||||||
|  | 			c.Set("api_version", channel.Other) | ||||||
|  | 		case common.ChannelTypeGemini: | ||||||
|  | 			c.Set("api_version", channel.Other) | ||||||
|  | 		case common.ChannelTypeAIProxyLibrary: | ||||||
|  | 			c.Set("library_id", channel.Other) | ||||||
|  | 		case common.ChannelTypeAli: | ||||||
|  | 			c.Set("plugin", channel.Other) | ||||||
|  | 		} | ||||||
| 		c.Next() | 		c.Next() | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func SetupContextForSelectedChannel(c *gin.Context, channel *model.Channel, modelName string) { |  | ||||||
| 	c.Set("channel", channel.Type) |  | ||||||
| 	c.Set("channel_id", channel.Id) |  | ||||||
| 	c.Set("channel_name", channel.Name) |  | ||||||
| 	c.Set("model_mapping", channel.GetModelMapping()) |  | ||||||
| 	c.Set("original_model", modelName) // for retry |  | ||||||
| 	c.Request.Header.Set("Authorization", fmt.Sprintf("Bearer %s", channel.Key)) |  | ||||||
| 	c.Set("base_url", channel.GetBaseURL()) |  | ||||||
| 	// this is for backward compatibility |  | ||||||
| 	switch channel.Type { |  | ||||||
| 	case common.ChannelTypeAzure: |  | ||||||
| 		c.Set(common.ConfigKeyAPIVersion, channel.Other) |  | ||||||
| 	case common.ChannelTypeXunfei: |  | ||||||
| 		c.Set(common.ConfigKeyAPIVersion, channel.Other) |  | ||||||
| 	case common.ChannelTypeGemini: |  | ||||||
| 		c.Set(common.ConfigKeyAPIVersion, channel.Other) |  | ||||||
| 	case common.ChannelTypeAIProxyLibrary: |  | ||||||
| 		c.Set(common.ConfigKeyLibraryID, channel.Other) |  | ||||||
| 	case common.ChannelTypeAli: |  | ||||||
| 		c.Set(common.ConfigKeyPlugin, channel.Other) |  | ||||||
| 	} |  | ||||||
| 	cfg, _ := channel.LoadConfig() |  | ||||||
| 	for k, v := range cfg { |  | ||||||
| 		c.Set(common.ConfigKeyPrefix+k, v) |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|   | |||||||
| @@ -3,14 +3,14 @@ package middleware | |||||||
| import ( | import ( | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" | 	"one-api/common" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| 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[logger.RequestIdKey].(string) | 			requestID = param.Keys[common.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"), | ||||||
|   | |||||||
| @@ -4,9 +4,8 @@ import ( | |||||||
| 	"context" | 	"context" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -27,7 +26,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, config.RateLimitKeyExpirationDuration) | 		rdb.Expire(ctx, key, common.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) | ||||||
| @@ -48,14 +47,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, config.RateLimitKeyExpirationDuration) | 			rdb.Expire(ctx, key, common.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, config.RateLimitKeyExpirationDuration) | 			rdb.Expire(ctx, key, common.RateLimitKeyExpirationDuration) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| @@ -76,7 +75,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(config.RateLimitKeyExpirationDuration) | 		inMemoryRateLimiter.Init(common.RateLimitKeyExpirationDuration) | ||||||
| 		return func(c *gin.Context) { | 		return func(c *gin.Context) { | ||||||
| 			memoryRateLimiter(c, maxRequestNum, duration, mark) | 			memoryRateLimiter(c, maxRequestNum, duration, mark) | ||||||
| 		} | 		} | ||||||
| @@ -84,21 +83,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(config.GlobalWebRateLimitNum, config.GlobalWebRateLimitDuration, "GW") | 	return rateLimitFactory(common.GlobalWebRateLimitNum, common.GlobalWebRateLimitDuration, "GW") | ||||||
| } | } | ||||||
|  |  | ||||||
| func GlobalAPIRateLimit() func(c *gin.Context) { | func GlobalAPIRateLimit() func(c *gin.Context) { | ||||||
| 	return rateLimitFactory(config.GlobalApiRateLimitNum, config.GlobalApiRateLimitDuration, "GA") | 	return rateLimitFactory(common.GlobalApiRateLimitNum, common.GlobalApiRateLimitDuration, "GA") | ||||||
| } | } | ||||||
|  |  | ||||||
| func CriticalRateLimit() func(c *gin.Context) { | func CriticalRateLimit() func(c *gin.Context) { | ||||||
| 	return rateLimitFactory(config.CriticalRateLimitNum, config.CriticalRateLimitDuration, "CT") | 	return rateLimitFactory(common.CriticalRateLimitNum, common.CriticalRateLimitDuration, "CT") | ||||||
| } | } | ||||||
|  |  | ||||||
| func DownloadRateLimit() func(c *gin.Context) { | func DownloadRateLimit() func(c *gin.Context) { | ||||||
| 	return rateLimitFactory(config.DownloadRateLimitNum, config.DownloadRateLimitDuration, "DW") | 	return rateLimitFactory(common.DownloadRateLimitNum, common.DownloadRateLimitDuration, "DW") | ||||||
| } | } | ||||||
|  |  | ||||||
| func UploadRateLimit() func(c *gin.Context) { | func UploadRateLimit() func(c *gin.Context) { | ||||||
| 	return rateLimitFactory(config.UploadRateLimitNum, config.UploadRateLimitDuration, "UP") | 	return rateLimitFactory(common.UploadRateLimitNum, common.UploadRateLimitDuration, "UP") | ||||||
| } | } | ||||||
|   | |||||||
| @@ -3,8 +3,8 @@ package middleware | |||||||
| import ( | import ( | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
|  | 	"one-api/common" | ||||||
| 	"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 { | ||||||
| 				logger.SysError(fmt.Sprintf("panic detected: %v", err)) | 				common.SysError(fmt.Sprintf("panic detected: %v", err)) | ||||||
| 				logger.SysError(fmt.Sprintf("stacktrace from panic: %s", string(debug.Stack()))) | 				common.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,17 +3,16 @@ package middleware | |||||||
| import ( | import ( | ||||||
| 	"context" | 	"context" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" | 	"one-api/common" | ||||||
| 	"github.com/songquanpeng/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 := helper.GetTimeString() + helper.GetRandomNumberString(8) | 		id := common.GetTimeString() + common.GetRandomString(8) | ||||||
| 		c.Set(logger.RequestIdKey, id) | 		c.Set(common.RequestIdKey, id) | ||||||
| 		ctx := context.WithValue(c.Request.Context(), logger.RequestIdKey, id) | 		ctx := context.WithValue(c.Request.Context(), common.RequestIdKey, id) | ||||||
| 		c.Request = c.Request.WithContext(ctx) | 		c.Request = c.Request.WithContext(ctx) | ||||||
| 		c.Header(logger.RequestIdKey, id) | 		c.Header(common.RequestIdKey, id) | ||||||
| 		c.Next() | 		c.Next() | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|   | |||||||
| @@ -4,10 +4,9 @@ import ( | |||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"github.com/gin-contrib/sessions" | 	"github.com/gin-contrib/sessions" | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"net/http" | 	"net/http" | ||||||
| 	"net/url" | 	"net/url" | ||||||
|  | 	"one-api/common" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| type turnstileCheckResponse struct { | type turnstileCheckResponse struct { | ||||||
| @@ -16,7 +15,7 @@ type turnstileCheckResponse struct { | |||||||
|  |  | ||||||
| func TurnstileCheck() gin.HandlerFunc { | func TurnstileCheck() gin.HandlerFunc { | ||||||
| 	return func(c *gin.Context) { | 	return func(c *gin.Context) { | ||||||
| 		if config.TurnstileCheckEnabled { | 		if common.TurnstileCheckEnabled { | ||||||
| 			session := sessions.Default(c) | 			session := sessions.Default(c) | ||||||
| 			turnstileChecked := session.Get("turnstile") | 			turnstileChecked := session.Get("turnstile") | ||||||
| 			if turnstileChecked != nil { | 			if turnstileChecked != nil { | ||||||
| @@ -33,12 +32,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":   {config.TurnstileSecretKey}, | 				"secret":   {common.TurnstileSecretKey}, | ||||||
| 				"response": {response}, | 				"response": {response}, | ||||||
| 				"remoteip": {c.ClientIP()}, | 				"remoteip": {c.ClientIP()}, | ||||||
| 			}) | 			}) | ||||||
| 			if err != nil { | 			if err != nil { | ||||||
| 				logger.SysError(err.Error()) | 				common.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(), | ||||||
| @@ -50,7 +49,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 { | ||||||
| 				logger.SysError(err.Error()) | 				common.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,17 +2,16 @@ package middleware | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/gin-gonic/gin" | 	"github.com/gin-gonic/gin" | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" | 	"one-api/common" | ||||||
| 	"github.com/songquanpeng/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": helper.MessageWithRequestId(message, c.GetString(logger.RequestIdKey)), | 			"message": common.MessageWithRequestId(message, c.GetString(common.RequestIdKey)), | ||||||
| 			"type":    "one_api_error", | 			"type":    "one_api_error", | ||||||
| 		}, | 		}, | ||||||
| 	}) | 	}) | ||||||
| 	c.Abort() | 	c.Abort() | ||||||
| 	logger.Error(c.Request.Context(), message) | 	common.LogError(c.Request.Context(), message) | ||||||
| } | } | ||||||
|   | |||||||
| @@ -1,7 +1,7 @@ | |||||||
| package model | package model | ||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/songquanpeng/one-api/common" | 	"one-api/common" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
|  |  | ||||||
|   | |||||||
| @@ -4,10 +4,8 @@ import ( | |||||||
| 	"encoding/json" | 	"encoding/json" | ||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"math/rand" | 	"math/rand" | ||||||
|  | 	"one-api/common" | ||||||
| 	"sort" | 	"sort" | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
| @@ -16,10 +14,10 @@ import ( | |||||||
| ) | ) | ||||||
|  |  | ||||||
| var ( | var ( | ||||||
| 	TokenCacheSeconds         = config.SyncFrequency | 	TokenCacheSeconds         = common.SyncFrequency | ||||||
| 	UserId2GroupCacheSeconds  = config.SyncFrequency | 	UserId2GroupCacheSeconds  = common.SyncFrequency | ||||||
| 	UserId2QuotaCacheSeconds  = config.SyncFrequency | 	UserId2QuotaCacheSeconds  = common.SyncFrequency | ||||||
| 	UserId2StatusCacheSeconds = config.SyncFrequency | 	UserId2StatusCacheSeconds = common.SyncFrequency | ||||||
| ) | ) | ||||||
|  |  | ||||||
| func CacheGetTokenByKey(key string) (*Token, error) { | func CacheGetTokenByKey(key string) (*Token, error) { | ||||||
| @@ -44,7 +42,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 { | ||||||
| 			logger.SysError("Redis set token error: " + err.Error()) | 			common.SysError("Redis set token error: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 		return &token, nil | 		return &token, nil | ||||||
| 	} | 	} | ||||||
| @@ -64,7 +62,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 { | ||||||
| 			logger.SysError("Redis set user group error: " + err.Error()) | 			common.SysError("Redis set user group error: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	return group, err | 	return group, err | ||||||
| @@ -82,7 +80,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 { | ||||||
| 			logger.SysError("Redis set user quota error: " + err.Error()) | 			common.SysError("Redis set user quota error: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 		return quota, err | 		return quota, err | ||||||
| 	} | 	} | ||||||
| @@ -94,7 +92,7 @@ func CacheUpdateUserQuota(id int) error { | |||||||
| 	if !common.RedisEnabled { | 	if !common.RedisEnabled { | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| 	quota, err := CacheGetUserQuota(id) | 	quota, err := GetUserQuota(id) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		return err | 		return err | ||||||
| 	} | 	} | ||||||
| @@ -129,7 +127,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 { | ||||||
| 		logger.SysError("Redis set user enabled error: " + err.Error()) | 		common.SysError("Redis set user enabled error: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| 	return userEnabled, err | 	return userEnabled, err | ||||||
| } | } | ||||||
| @@ -180,19 +178,19 @@ func InitChannelCache() { | |||||||
| 	channelSyncLock.Lock() | 	channelSyncLock.Lock() | ||||||
| 	group2model2channels = newGroup2model2channels | 	group2model2channels = newGroup2model2channels | ||||||
| 	channelSyncLock.Unlock() | 	channelSyncLock.Unlock() | ||||||
| 	logger.SysLog("channels synced from database") | 	common.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) | ||||||
| 		logger.SysLog("syncing channels from database") | 		common.SysLog("syncing channels from database") | ||||||
| 		InitChannelCache() | 		InitChannelCache() | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func CacheGetRandomSatisfiedChannel(group string, model string, ignoreFirstPriority bool) (*Channel, error) { | func CacheGetRandomSatisfiedChannel(group string, model string) (*Channel, error) { | ||||||
| 	if !config.MemoryCacheEnabled { | 	if !common.MemoryCacheEnabled { | ||||||
| 		return GetRandomSatisfiedChannel(group, model) | 		return GetRandomSatisfiedChannel(group, model) | ||||||
| 	} | 	} | ||||||
| 	channelSyncLock.RLock() | 	channelSyncLock.RLock() | ||||||
| @@ -213,10 +211,5 @@ func CacheGetRandomSatisfiedChannel(group string, model string, ignoreFirstPrior | |||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	idx := rand.Intn(endIdx) | 	idx := rand.Intn(endIdx) | ||||||
| 	if ignoreFirstPriority { |  | ||||||
| 		if endIdx < len(channels) { // which means there are more than one priority |  | ||||||
| 			idx = common.RandRange(endIdx, len(channels)) |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	return channels[idx], nil | 	return channels[idx], nil | ||||||
| } | } | ||||||
|   | |||||||
| @@ -1,13 +1,8 @@ | |||||||
| package model | package model | ||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"encoding/json" |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
|  | 	"one-api/common" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| type Channel struct { | type Channel struct { | ||||||
| @@ -21,7 +16,7 @@ type Channel struct { | |||||||
| 	TestTime           int64   `json:"test_time" gorm:"bigint"` | 	TestTime           int64   `json:"test_time" gorm:"bigint"` | ||||||
| 	ResponseTime       int     `json:"response_time"` // in milliseconds | 	ResponseTime       int     `json:"response_time"` // in milliseconds | ||||||
| 	BaseURL            *string `json:"base_url" gorm:"column:base_url;default:''"` | 	BaseURL            *string `json:"base_url" gorm:"column:base_url;default:''"` | ||||||
| 	Other              string  `json:"other"`   // DEPRECATED: please save config to field Config | 	Other              string  `json:"other"` | ||||||
| 	Balance            float64 `json:"balance"` // in USD | 	Balance            float64 `json:"balance"` // in USD | ||||||
| 	BalanceUpdatedTime int64   `json:"balance_updated_time" gorm:"bigint"` | 	BalanceUpdatedTime int64   `json:"balance_updated_time" gorm:"bigint"` | ||||||
| 	Models             string  `json:"models"` | 	Models             string  `json:"models"` | ||||||
| @@ -29,18 +24,14 @@ type Channel struct { | |||||||
| 	UsedQuota          int64   `json:"used_quota" gorm:"bigint;default:0"` | 	UsedQuota          int64   `json:"used_quota" gorm:"bigint;default:0"` | ||||||
| 	ModelMapping       *string `json:"model_mapping" gorm:"type:varchar(1024);default:''"` | 	ModelMapping       *string `json:"model_mapping" gorm:"type:varchar(1024);default:''"` | ||||||
| 	Priority           *int64  `json:"priority" gorm:"bigint;default:0"` | 	Priority           *int64  `json:"priority" gorm:"bigint;default:0"` | ||||||
| 	Config             string  `json:"config"` |  | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetAllChannels(startIdx int, num int, scope string) ([]*Channel, error) { | func GetAllChannels(startIdx int, num int, selectAll bool) ([]*Channel, error) { | ||||||
| 	var channels []*Channel | 	var channels []*Channel | ||||||
| 	var err error | 	var err error | ||||||
| 	switch scope { | 	if selectAll { | ||||||
| 	case "all": |  | ||||||
| 		err = DB.Order("id desc").Find(&channels).Error | 		err = DB.Order("id desc").Find(&channels).Error | ||||||
| 	case "disabled": | 	} else { | ||||||
| 		err = DB.Order("id desc").Where("status = ? or status = ?", common.ChannelStatusAutoDisabled, common.ChannelStatusManuallyDisabled).Find(&channels).Error |  | ||||||
| 	default: |  | ||||||
| 		err = DB.Order("id desc").Limit(num).Offset(startIdx).Omit("key").Find(&channels).Error | 		err = DB.Order("id desc").Limit(num).Offset(startIdx).Omit("key").Find(&channels).Error | ||||||
| 	} | 	} | ||||||
| 	return channels, err | 	return channels, err | ||||||
| @@ -51,7 +42,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+" = ?", helper.String2Int(keyword), keyword+"%", keyword).Find(&channels).Error | 	err = DB.Omit("key").Where("id = ? or name LIKE ? or "+keyCol+" = ?", common.String2Int(keyword), keyword+"%", keyword).Find(&channels).Error | ||||||
| 	return channels, err | 	return channels, err | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -95,17 +86,11 @@ func (channel *Channel) GetBaseURL() string { | |||||||
| 	return *channel.BaseURL | 	return *channel.BaseURL | ||||||
| } | } | ||||||
|  |  | ||||||
| func (channel *Channel) GetModelMapping() map[string]string { | func (channel *Channel) GetModelMapping() string { | ||||||
| 	if channel.ModelMapping == nil || *channel.ModelMapping == "" || *channel.ModelMapping == "{}" { | 	if channel.ModelMapping == nil { | ||||||
| 		return nil | 		return "" | ||||||
| 	} | 	} | ||||||
| 	modelMapping := make(map[string]string) | 	return *channel.ModelMapping | ||||||
| 	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 { | ||||||
| @@ -131,21 +116,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:     helper.GetTimestamp(), | 		TestTime:     common.GetTimestamp(), | ||||||
| 		ResponseTime: int(responseTime), | 		ResponseTime: int(responseTime), | ||||||
| 	}).Error | 	}).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.SysError("failed to update response time: " + err.Error()) | 		common.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: helper.GetTimestamp(), | 		BalanceUpdatedTime: common.GetTimestamp(), | ||||||
| 		Balance:            balance, | 		Balance:            balance, | ||||||
| 	}).Error | 	}).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.SysError("failed to update balance: " + err.Error()) | 		common.SysError("failed to update balance: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -159,31 +144,19 @@ func (channel *Channel) Delete() error { | |||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
|  |  | ||||||
| func (channel *Channel) LoadConfig() (map[string]string, error) { |  | ||||||
| 	if channel.Config == "" { |  | ||||||
| 		return nil, nil |  | ||||||
| 	} |  | ||||||
| 	cfg := make(map[string]string) |  | ||||||
| 	err := json.Unmarshal([]byte(channel.Config), &cfg) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return nil, err |  | ||||||
| 	} |  | ||||||
| 	return cfg, nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| 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 { | ||||||
| 		logger.SysError("failed to update ability status: " + err.Error()) | 		common.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 { | ||||||
| 		logger.SysError("failed to update channel status: " + err.Error()) | 		common.SysError("failed to update channel status: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func UpdateChannelUsedQuota(id int, quota int) { | func UpdateChannelUsedQuota(id int, quota int) { | ||||||
| 	if config.BatchUpdateEnabled { | 	if common.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeChannelUsedQuota, id, quota) | 		addNewRecord(BatchUpdateTypeChannelUsedQuota, id, quota) | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| @@ -193,7 +166,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 { | ||||||
| 		logger.SysError("failed to update channel used quota: " + err.Error()) | 		common.SysError("failed to update channel used quota: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
|   | |||||||
							
								
								
									
										44
									
								
								model/log.go
									
									
									
									
									
								
							
							
						
						
									
										44
									
								
								model/log.go
									
									
									
									
									
								
							| @@ -3,18 +3,15 @@ package model | |||||||
| import ( | import ( | ||||||
| 	"context" | 	"context" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" | 	"one-api/common" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
|  |  | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| type Log struct { | type Log struct { | ||||||
| 	Id               int    `json:"id"` | 	Id               int    `json:"id;index:idx_created_at_id,priority:1"` | ||||||
| 	UserId           int    `json:"user_id" gorm:"index"` | 	UserId           int    `json:"user_id" gorm:"index"` | ||||||
| 	CreatedAt        int64  `json:"created_at" gorm:"bigint;index:idx_created_at_type"` | 	CreatedAt        int64  `json:"created_at" gorm:"bigint;index:idx_created_at_id,priority:2;index:idx_created_at_type"` | ||||||
| 	Type             int    `json:"type" gorm:"index:idx_created_at_type"` | 	Type             int    `json:"type" gorm:"index:idx_created_at_type"` | ||||||
| 	Content          string `json:"content"` | 	Content          string `json:"content"` | ||||||
| 	Username         string `json:"username" gorm:"index:index_username_model_name,priority:2;default:''"` | 	Username         string `json:"username" gorm:"index:index_username_model_name,priority:2;default:''"` | ||||||
| @@ -35,31 +32,31 @@ const ( | |||||||
| ) | ) | ||||||
|  |  | ||||||
| func RecordLog(userId int, logType int, content string) { | func RecordLog(userId int, logType int, content string) { | ||||||
| 	if logType == LogTypeConsume && !config.LogConsumeEnabled { | 	if logType == LogTypeConsume && !common.LogConsumeEnabled { | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	log := &Log{ | 	log := &Log{ | ||||||
| 		UserId:    userId, | 		UserId:    userId, | ||||||
| 		Username:  GetUsernameById(userId), | 		Username:  GetUsernameById(userId), | ||||||
| 		CreatedAt: helper.GetTimestamp(), | 		CreatedAt: common.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 { | ||||||
| 		logger.SysError("failed to record log: " + err.Error()) | 		common.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) { | ||||||
| 	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)) | 	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)) | ||||||
| 	if !config.LogConsumeEnabled { | 	if !common.LogConsumeEnabled { | ||||||
| 		return | 		return | ||||||
| 	} | 	} | ||||||
| 	log := &Log{ | 	log := &Log{ | ||||||
| 		UserId:           userId, | 		UserId:           userId, | ||||||
| 		Username:         GetUsernameById(userId), | 		Username:         GetUsernameById(userId), | ||||||
| 		CreatedAt:        helper.GetTimestamp(), | 		CreatedAt:        common.GetTimestamp(), | ||||||
| 		Type:             LogTypeConsume, | 		Type:             LogTypeConsume, | ||||||
| 		Content:          content, | 		Content:          content, | ||||||
| 		PromptTokens:     promptTokens, | 		PromptTokens:     promptTokens, | ||||||
| @@ -71,11 +68,11 @@ 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 { | ||||||
| 		logger.Error(ctx, "failed to record log: "+err.Error()) | 		common.LogError(ctx, "failed to record log: "+err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| func GetAllLogs(logType int, startTimestamp int64, endTimestamp int64, modelName string, username string, tokenName string, startIdx int, num int, channel int) (logs []*Log, err error) { | func GetAllLogs(logType int, startTimestamp int64, endTimestamp int64, modelName string, username string, tokenName string, startIdx int, num int, channels []int) (logs []*Log, err error) { | ||||||
| 	var tx *gorm.DB | 	var tx *gorm.DB | ||||||
| 	if logType == LogTypeUnknown { | 	if logType == LogTypeUnknown { | ||||||
| 		tx = DB | 		tx = DB | ||||||
| @@ -97,9 +94,12 @@ func GetAllLogs(logType int, startTimestamp int64, endTimestamp int64, modelName | |||||||
| 	if endTimestamp != 0 { | 	if endTimestamp != 0 { | ||||||
| 		tx = tx.Where("created_at <= ?", endTimestamp) | 		tx = tx.Where("created_at <= ?", endTimestamp) | ||||||
| 	} | 	} | ||||||
| 	if channel != 0 { | 	if len(channels) > 1 { | ||||||
| 		tx = tx.Where("channel_id = ?", channel) | 		tx = tx.Where("channel_id IN ?", channels) | ||||||
|  | 	} else if len(channels) == 1 && channels[0] != 0 { | ||||||
|  | 		tx = tx.Where("channel_id = ?", channels[0]) | ||||||
| 	} | 	} | ||||||
|  |  | ||||||
| 	err = tx.Order("id desc").Limit(num).Offset(startIdx).Find(&logs).Error | 	err = tx.Order("id desc").Limit(num).Offset(startIdx).Find(&logs).Error | ||||||
| 	return logs, err | 	return logs, err | ||||||
| } | } | ||||||
| @@ -128,16 +128,16 @@ 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(config.MaxRecentItems).Find(&logs).Error | 	err = DB.Where("type = ? or content LIKE ?", keyword, keyword+"%").Order("id desc").Limit(common.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(config.MaxRecentItems).Omit("id").Find(&logs).Error | 	err = DB.Where("user_id = ? and type = ?", userId, keyword).Order("id desc").Limit(common.MaxRecentItems).Omit("id").Find(&logs).Error | ||||||
| 	return logs, err | 	return logs, err | ||||||
| } | } | ||||||
|  |  | ||||||
| func SumUsedQuota(logType int, startTimestamp int64, endTimestamp int64, modelName string, username string, tokenName string, channel int) (quota int) { | func SumUsedQuota(logType int, startTimestamp int64, endTimestamp int64, modelName string, username string, tokenName string, channels []int) (quota int) { | ||||||
| 	tx := DB.Table("logs").Select("ifnull(sum(quota),0)") | 	tx := DB.Table("logs").Select("ifnull(sum(quota),0)") | ||||||
| 	if username != "" { | 	if username != "" { | ||||||
| 		tx = tx.Where("username = ?", username) | 		tx = tx.Where("username = ?", username) | ||||||
| @@ -154,8 +154,10 @@ func SumUsedQuota(logType int, startTimestamp int64, endTimestamp int64, modelNa | |||||||
| 	if modelName != "" { | 	if modelName != "" { | ||||||
| 		tx = tx.Where("model_name = ?", modelName) | 		tx = tx.Where("model_name = ?", modelName) | ||||||
| 	} | 	} | ||||||
| 	if channel != 0 { | 	if len(channels) > 1 { | ||||||
| 		tx = tx.Where("channel_id = ?", channel) | 		tx = tx.Where("channel_id IN ?", channels) | ||||||
|  | 	} else if len(channels) == 1 && channels[0] != 0 { | ||||||
|  | 		tx = tx.Where("channel_id = ?", channels[0]) | ||||||
| 	} | 	} | ||||||
| 	tx.Where("type = ?", LogTypeConsume).Scan("a) | 	tx.Where("type = ?", LogTypeConsume).Scan("a) | ||||||
| 	return quota | 	return quota | ||||||
|   | |||||||
| @@ -2,14 +2,11 @@ package model | |||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"gorm.io/driver/mysql" | 	"gorm.io/driver/mysql" | ||||||
| 	"gorm.io/driver/postgres" | 	"gorm.io/driver/postgres" | ||||||
| 	"gorm.io/driver/sqlite" | 	"gorm.io/driver/sqlite" | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
|  | 	"one-api/common" | ||||||
| 	"os" | 	"os" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"time" | 	"time" | ||||||
| @@ -19,9 +16,9 @@ var DB *gorm.DB | |||||||
|  |  | ||||||
| func createRootAccountIfNeed() error { | func createRootAccountIfNeed() error { | ||||||
| 	var user User | 	var user User | ||||||
| 	//if user.Status != util.UserStatusEnabled { | 	//if user.Status != common.UserStatusEnabled { | ||||||
| 	if err := DB.First(&user).Error; err != nil { | 	if err := DB.First(&user).Error; err != nil { | ||||||
| 		logger.SysLog("no user exists, create a root user for you: username is root, password is 123456") | 		common.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 | ||||||
| @@ -32,7 +29,7 @@ func createRootAccountIfNeed() error { | |||||||
| 			Role:        common.RoleRootUser, | 			Role:        common.RoleRootUser, | ||||||
| 			Status:      common.UserStatusEnabled, | 			Status:      common.UserStatusEnabled, | ||||||
| 			DisplayName: "Root User", | 			DisplayName: "Root User", | ||||||
| 			AccessToken: helper.GetUUID(), | 			AccessToken: common.GetUUID(), | ||||||
| 			Quota:       100000000, | 			Quota:       100000000, | ||||||
| 		} | 		} | ||||||
| 		DB.Create(&rootUser) | 		DB.Create(&rootUser) | ||||||
| @@ -45,7 +42,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 | ||||||
| 			logger.SysLog("using PostgreSQL as database") | 			common.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, | ||||||
| @@ -55,13 +52,13 @@ func chooseDB() (*gorm.DB, error) { | |||||||
| 			}) | 			}) | ||||||
| 		} | 		} | ||||||
| 		// Use MySQL | 		// Use MySQL | ||||||
| 		logger.SysLog("using MySQL as database") | 		common.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 | ||||||
| 	logger.SysLog("SQL_DSN not set, using SQLite as database") | 	common.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{ | ||||||
| @@ -72,7 +69,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 config.DebugSQLEnabled { | 		if common.DebugEnabled { | ||||||
| 			db = db.Debug() | 			db = db.Debug() | ||||||
| 		} | 		} | ||||||
| 		DB = db | 		DB = db | ||||||
| @@ -80,14 +77,14 @@ func InitDB() (err error) { | |||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		sqlDB.SetMaxIdleConns(helper.GetOrDefaultEnvInt("SQL_MAX_IDLE_CONNS", 100)) | 		sqlDB.SetMaxIdleConns(common.GetOrDefault("SQL_MAX_IDLE_CONNS", 100)) | ||||||
| 		sqlDB.SetMaxOpenConns(helper.GetOrDefaultEnvInt("SQL_MAX_OPEN_CONNS", 1000)) | 		sqlDB.SetMaxOpenConns(common.GetOrDefault("SQL_MAX_OPEN_CONNS", 1000)) | ||||||
| 		sqlDB.SetConnMaxLifetime(time.Second * time.Duration(helper.GetOrDefaultEnvInt("SQL_MAX_LIFETIME", 60))) | 		sqlDB.SetConnMaxLifetime(time.Second * time.Duration(common.GetOrDefault("SQL_MAX_LIFETIME", 60))) | ||||||
|  |  | ||||||
| 		if !config.IsMasterNode { | 		if !common.IsMasterNode { | ||||||
| 			return nil | 			return nil | ||||||
| 		} | 		} | ||||||
| 		logger.SysLog("database migration started") | 		common.SysLog("database migration started") | ||||||
| 		err = db.AutoMigrate(&Channel{}) | 		err = db.AutoMigrate(&Channel{}) | ||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| @@ -116,11 +113,11 @@ func InitDB() (err error) { | |||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		logger.SysLog("database migrated") | 		common.SysLog("database migrated") | ||||||
| 		err = createRootAccountIfNeed() | 		err = createRootAccountIfNeed() | ||||||
| 		return err | 		return err | ||||||
| 	} else { | 	} else { | ||||||
| 		logger.FatalLog(err) | 		common.FatalLog(err) | ||||||
| 	} | 	} | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
|   | |||||||
							
								
								
									
										227
									
								
								model/option.go
									
									
									
									
									
								
							
							
						
						
									
										227
									
								
								model/option.go
									
									
									
									
									
								
							| @@ -1,9 +1,7 @@ | |||||||
| package model | package model | ||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/songquanpeng/one-api/common" | 	"one-api/common" | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"strconv" | 	"strconv" | ||||||
| 	"strings" | 	"strings" | ||||||
| 	"time" | 	"time" | ||||||
| @@ -22,59 +20,61 @@ func AllOption() ([]*Option, error) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func InitOptionMap() { | func InitOptionMap() { | ||||||
| 	config.OptionMapRWMutex.Lock() | 	common.OptionMapRWMutex.Lock() | ||||||
| 	config.OptionMap = make(map[string]string) | 	common.OptionMap = make(map[string]string) | ||||||
| 	config.OptionMap["PasswordLoginEnabled"] = strconv.FormatBool(config.PasswordLoginEnabled) | 	common.OptionMap["FileUploadPermission"] = strconv.Itoa(common.FileUploadPermission) | ||||||
| 	config.OptionMap["PasswordRegisterEnabled"] = strconv.FormatBool(config.PasswordRegisterEnabled) | 	common.OptionMap["FileDownloadPermission"] = strconv.Itoa(common.FileDownloadPermission) | ||||||
| 	config.OptionMap["EmailVerificationEnabled"] = strconv.FormatBool(config.EmailVerificationEnabled) | 	common.OptionMap["ImageUploadPermission"] = strconv.Itoa(common.ImageUploadPermission) | ||||||
| 	config.OptionMap["GitHubOAuthEnabled"] = strconv.FormatBool(config.GitHubOAuthEnabled) | 	common.OptionMap["ImageDownloadPermission"] = strconv.Itoa(common.ImageDownloadPermission) | ||||||
| 	config.OptionMap["WeChatAuthEnabled"] = strconv.FormatBool(config.WeChatAuthEnabled) | 	common.OptionMap["PasswordLoginEnabled"] = strconv.FormatBool(common.PasswordLoginEnabled) | ||||||
| 	config.OptionMap["TurnstileCheckEnabled"] = strconv.FormatBool(config.TurnstileCheckEnabled) | 	common.OptionMap["PasswordRegisterEnabled"] = strconv.FormatBool(common.PasswordRegisterEnabled) | ||||||
| 	config.OptionMap["RegisterEnabled"] = strconv.FormatBool(config.RegisterEnabled) | 	common.OptionMap["EmailVerificationEnabled"] = strconv.FormatBool(common.EmailVerificationEnabled) | ||||||
| 	config.OptionMap["AutomaticDisableChannelEnabled"] = strconv.FormatBool(config.AutomaticDisableChannelEnabled) | 	common.OptionMap["GitHubOAuthEnabled"] = strconv.FormatBool(common.GitHubOAuthEnabled) | ||||||
| 	config.OptionMap["AutomaticEnableChannelEnabled"] = strconv.FormatBool(config.AutomaticEnableChannelEnabled) | 	common.OptionMap["WeChatAuthEnabled"] = strconv.FormatBool(common.WeChatAuthEnabled) | ||||||
| 	config.OptionMap["ApproximateTokenEnabled"] = strconv.FormatBool(config.ApproximateTokenEnabled) | 	common.OptionMap["TurnstileCheckEnabled"] = strconv.FormatBool(common.TurnstileCheckEnabled) | ||||||
| 	config.OptionMap["LogConsumeEnabled"] = strconv.FormatBool(config.LogConsumeEnabled) | 	common.OptionMap["RegisterEnabled"] = strconv.FormatBool(common.RegisterEnabled) | ||||||
| 	config.OptionMap["DisplayInCurrencyEnabled"] = strconv.FormatBool(config.DisplayInCurrencyEnabled) | 	common.OptionMap["AutomaticDisableChannelEnabled"] = strconv.FormatBool(common.AutomaticDisableChannelEnabled) | ||||||
| 	config.OptionMap["DisplayTokenStatEnabled"] = strconv.FormatBool(config.DisplayTokenStatEnabled) | 	common.OptionMap["AutomaticEnableChannelEnabled"] = strconv.FormatBool(common.AutomaticEnableChannelEnabled) | ||||||
| 	config.OptionMap["ChannelDisableThreshold"] = strconv.FormatFloat(config.ChannelDisableThreshold, 'f', -1, 64) | 	common.OptionMap["ApproximateTokenEnabled"] = strconv.FormatBool(common.ApproximateTokenEnabled) | ||||||
| 	config.OptionMap["EmailDomainRestrictionEnabled"] = strconv.FormatBool(config.EmailDomainRestrictionEnabled) | 	common.OptionMap["LogConsumeEnabled"] = strconv.FormatBool(common.LogConsumeEnabled) | ||||||
| 	config.OptionMap["EmailDomainWhitelist"] = strings.Join(config.EmailDomainWhitelist, ",") | 	common.OptionMap["DisplayInCurrencyEnabled"] = strconv.FormatBool(common.DisplayInCurrencyEnabled) | ||||||
| 	config.OptionMap["SMTPServer"] = "" | 	common.OptionMap["DisplayTokenStatEnabled"] = strconv.FormatBool(common.DisplayTokenStatEnabled) | ||||||
| 	config.OptionMap["SMTPFrom"] = "" | 	common.OptionMap["ChannelDisableThreshold"] = strconv.FormatFloat(common.ChannelDisableThreshold, 'f', -1, 64) | ||||||
| 	config.OptionMap["SMTPPort"] = strconv.Itoa(config.SMTPPort) | 	common.OptionMap["EmailDomainRestrictionEnabled"] = strconv.FormatBool(common.EmailDomainRestrictionEnabled) | ||||||
| 	config.OptionMap["SMTPAccount"] = "" | 	common.OptionMap["EmailDomainWhitelist"] = strings.Join(common.EmailDomainWhitelist, ",") | ||||||
| 	config.OptionMap["SMTPToken"] = "" | 	common.OptionMap["SMTPServer"] = "" | ||||||
| 	config.OptionMap["Notice"] = "" | 	common.OptionMap["SMTPFrom"] = "" | ||||||
| 	config.OptionMap["About"] = "" | 	common.OptionMap["SMTPPort"] = strconv.Itoa(common.SMTPPort) | ||||||
| 	config.OptionMap["HomePageContent"] = "" | 	common.OptionMap["SMTPAccount"] = "" | ||||||
| 	config.OptionMap["Footer"] = config.Footer | 	common.OptionMap["SMTPToken"] = "" | ||||||
| 	config.OptionMap["SystemName"] = config.SystemName | 	common.OptionMap["SMTPAuthLoginEnabled"] = strconv.FormatBool(common.SMTPAuthLoginEnabled) | ||||||
| 	config.OptionMap["Logo"] = config.Logo | 	common.OptionMap["Notice"] = "" | ||||||
| 	config.OptionMap["ServerAddress"] = "" | 	common.OptionMap["About"] = "" | ||||||
| 	config.OptionMap["GitHubClientId"] = "" | 	common.OptionMap["HomePageContent"] = "" | ||||||
| 	config.OptionMap["GitHubClientSecret"] = "" | 	common.OptionMap["Footer"] = common.Footer | ||||||
| 	config.OptionMap["WeChatServerAddress"] = "" | 	common.OptionMap["SystemName"] = common.SystemName | ||||||
| 	config.OptionMap["WeChatServerToken"] = "" | 	common.OptionMap["Logo"] = common.Logo | ||||||
| 	config.OptionMap["WeChatAccountQRCodeImageURL"] = "" | 	common.OptionMap["ServerAddress"] = "" | ||||||
| 	config.OptionMap["MessagePusherAddress"] = "" | 	common.OptionMap["GitHubClientId"] = "" | ||||||
| 	config.OptionMap["MessagePusherToken"] = "" | 	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["CompletionRatio"] = common.CompletionRatio2JSONString() | 	common.OptionMap["PreConsumedQuota"] = strconv.Itoa(common.PreConsumedQuota) | ||||||
| 	config.OptionMap["TopUpLink"] = config.TopUpLink | 	common.OptionMap["ModelRatio"] = common.ModelRatio2JSONString() | ||||||
| 	config.OptionMap["ChatLink"] = config.ChatLink | 	common.OptionMap["GroupRatio"] = common.GroupRatio2JSONString() | ||||||
| 	config.OptionMap["QuotaPerUnit"] = strconv.FormatFloat(config.QuotaPerUnit, 'f', -1, 64) | 	common.OptionMap["TopUpLink"] = common.TopUpLink | ||||||
| 	config.OptionMap["RetryTimes"] = strconv.Itoa(config.RetryTimes) | 	common.OptionMap["ChatLink"] = common.ChatLink | ||||||
| 	config.OptionMap["Theme"] = config.Theme | 	common.OptionMap["QuotaPerUnit"] = strconv.FormatFloat(common.QuotaPerUnit, 'f', -1, 64) | ||||||
| 	config.OptionMapRWMutex.Unlock() | 	common.OptionMap["RetryTimes"] = strconv.Itoa(common.RetryTimes) | ||||||
|  | 	common.OptionMap["Theme"] = common.Theme | ||||||
|  | 	common.OptionMapRWMutex.Unlock() | ||||||
| 	loadOptionsFromDatabase() | 	loadOptionsFromDatabase() | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -83,7 +83,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 { | ||||||
| 			logger.SysError("failed to update option map: " + err.Error()) | 			common.SysError("failed to update option map: " + err.Error()) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| @@ -91,7 +91,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) | ||||||
| 		logger.SysLog("syncing options from database") | 		common.SysLog("syncing options from database") | ||||||
| 		loadOptionsFromDatabase() | 		loadOptionsFromDatabase() | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
| @@ -113,110 +113,119 @@ func UpdateOption(key string, value string) error { | |||||||
| } | } | ||||||
|  |  | ||||||
| func updateOptionMap(key string, value string) (err error) { | func updateOptionMap(key string, value string) (err error) { | ||||||
| 	config.OptionMapRWMutex.Lock() | 	common.OptionMapRWMutex.Lock() | ||||||
| 	defer config.OptionMapRWMutex.Unlock() | 	defer common.OptionMapRWMutex.Unlock() | ||||||
| 	config.OptionMap[key] = value | 	common.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": | ||||||
| 			config.PasswordRegisterEnabled = boolValue | 			common.PasswordRegisterEnabled = boolValue | ||||||
| 		case "PasswordLoginEnabled": | 		case "PasswordLoginEnabled": | ||||||
| 			config.PasswordLoginEnabled = boolValue | 			common.PasswordLoginEnabled = boolValue | ||||||
| 		case "EmailVerificationEnabled": | 		case "EmailVerificationEnabled": | ||||||
| 			config.EmailVerificationEnabled = boolValue | 			common.EmailVerificationEnabled = boolValue | ||||||
| 		case "GitHubOAuthEnabled": | 		case "GitHubOAuthEnabled": | ||||||
| 			config.GitHubOAuthEnabled = boolValue | 			common.GitHubOAuthEnabled = boolValue | ||||||
| 		case "WeChatAuthEnabled": | 		case "WeChatAuthEnabled": | ||||||
| 			config.WeChatAuthEnabled = boolValue | 			common.WeChatAuthEnabled = boolValue | ||||||
| 		case "TurnstileCheckEnabled": | 		case "TurnstileCheckEnabled": | ||||||
| 			config.TurnstileCheckEnabled = boolValue | 			common.TurnstileCheckEnabled = boolValue | ||||||
| 		case "RegisterEnabled": | 		case "RegisterEnabled": | ||||||
| 			config.RegisterEnabled = boolValue | 			common.RegisterEnabled = boolValue | ||||||
| 		case "EmailDomainRestrictionEnabled": | 		case "EmailDomainRestrictionEnabled": | ||||||
| 			config.EmailDomainRestrictionEnabled = boolValue | 			common.EmailDomainRestrictionEnabled = boolValue | ||||||
| 		case "AutomaticDisableChannelEnabled": | 		case "AutomaticDisableChannelEnabled": | ||||||
| 			config.AutomaticDisableChannelEnabled = boolValue | 			common.AutomaticDisableChannelEnabled = boolValue | ||||||
| 		case "AutomaticEnableChannelEnabled": | 		case "AutomaticEnableChannelEnabled": | ||||||
| 			config.AutomaticEnableChannelEnabled = boolValue | 			common.AutomaticEnableChannelEnabled = boolValue | ||||||
| 		case "ApproximateTokenEnabled": | 		case "ApproximateTokenEnabled": | ||||||
| 			config.ApproximateTokenEnabled = boolValue | 			common.ApproximateTokenEnabled = boolValue | ||||||
| 		case "LogConsumeEnabled": | 		case "LogConsumeEnabled": | ||||||
| 			config.LogConsumeEnabled = boolValue | 			common.LogConsumeEnabled = boolValue | ||||||
| 		case "DisplayInCurrencyEnabled": | 		case "DisplayInCurrencyEnabled": | ||||||
| 			config.DisplayInCurrencyEnabled = boolValue | 			common.DisplayInCurrencyEnabled = boolValue | ||||||
| 		case "DisplayTokenStatEnabled": | 		case "DisplayTokenStatEnabled": | ||||||
| 			config.DisplayTokenStatEnabled = boolValue | 			common.DisplayTokenStatEnabled = boolValue | ||||||
|  | 		case "SMTPAuthLoginEnabled": | ||||||
|  | 			common.SMTPAuthLoginEnabled = boolValue | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	switch key { | 	switch key { | ||||||
| 	case "EmailDomainWhitelist": | 	case "EmailDomainWhitelist": | ||||||
| 		config.EmailDomainWhitelist = strings.Split(value, ",") | 		common.EmailDomainWhitelist = strings.Split(value, ",") | ||||||
| 	case "SMTPServer": | 	case "SMTPServer": | ||||||
| 		config.SMTPServer = value | 		common.SMTPServer = value | ||||||
| 	case "SMTPPort": | 	case "SMTPPort": | ||||||
| 		intValue, _ := strconv.Atoi(value) | 		intValue, _ := strconv.Atoi(value) | ||||||
| 		config.SMTPPort = intValue | 		common.SMTPPort = intValue | ||||||
| 	case "SMTPAccount": | 	case "SMTPAccount": | ||||||
| 		config.SMTPAccount = value | 		common.SMTPAccount = value | ||||||
| 	case "SMTPFrom": | 	case "SMTPFrom": | ||||||
| 		config.SMTPFrom = value | 		common.SMTPFrom = value | ||||||
| 	case "SMTPToken": | 	case "SMTPToken": | ||||||
| 		config.SMTPToken = value | 		common.SMTPToken = value | ||||||
| 	case "ServerAddress": | 	case "ServerAddress": | ||||||
| 		config.ServerAddress = value | 		common.ServerAddress = value | ||||||
| 	case "GitHubClientId": | 	case "GitHubClientId": | ||||||
| 		config.GitHubClientId = value | 		common.GitHubClientId = value | ||||||
| 	case "GitHubClientSecret": | 	case "GitHubClientSecret": | ||||||
| 		config.GitHubClientSecret = value | 		common.GitHubClientSecret = value | ||||||
| 	case "Footer": | 	case "Footer": | ||||||
| 		config.Footer = value | 		common.Footer = value | ||||||
| 	case "SystemName": | 	case "SystemName": | ||||||
| 		config.SystemName = value | 		common.SystemName = value | ||||||
| 	case "Logo": | 	case "Logo": | ||||||
| 		config.Logo = value | 		common.Logo = value | ||||||
| 	case "WeChatServerAddress": | 	case "WeChatServerAddress": | ||||||
| 		config.WeChatServerAddress = value | 		common.WeChatServerAddress = value | ||||||
| 	case "WeChatServerToken": | 	case "WeChatServerToken": | ||||||
| 		config.WeChatServerToken = value | 		common.WeChatServerToken = value | ||||||
| 	case "WeChatAccountQRCodeImageURL": | 	case "WeChatAccountQRCodeImageURL": | ||||||
| 		config.WeChatAccountQRCodeImageURL = value | 		common.WeChatAccountQRCodeImageURL = value | ||||||
| 	case "MessagePusherAddress": |  | ||||||
| 		config.MessagePusherAddress = value |  | ||||||
| 	case "MessagePusherToken": |  | ||||||
| 		config.MessagePusherToken = value |  | ||||||
| 	case "TurnstileSiteKey": | 	case "TurnstileSiteKey": | ||||||
| 		config.TurnstileSiteKey = value | 		common.TurnstileSiteKey = value | ||||||
| 	case "TurnstileSecretKey": | 	case "TurnstileSecretKey": | ||||||
| 		config.TurnstileSecretKey = value | 		common.TurnstileSecretKey = value | ||||||
| 	case "QuotaForNewUser": | 	case "QuotaForNewUser": | ||||||
| 		config.QuotaForNewUser, _ = strconv.Atoi(value) | 		common.QuotaForNewUser, _ = strconv.Atoi(value) | ||||||
| 	case "QuotaForInviter": | 	case "QuotaForInviter": | ||||||
| 		config.QuotaForInviter, _ = strconv.Atoi(value) | 		common.QuotaForInviter, _ = strconv.Atoi(value) | ||||||
| 	case "QuotaForInvitee": | 	case "QuotaForInvitee": | ||||||
| 		config.QuotaForInvitee, _ = strconv.Atoi(value) | 		common.QuotaForInvitee, _ = strconv.Atoi(value) | ||||||
| 	case "QuotaRemindThreshold": | 	case "QuotaRemindThreshold": | ||||||
| 		config.QuotaRemindThreshold, _ = strconv.Atoi(value) | 		common.QuotaRemindThreshold, _ = strconv.Atoi(value) | ||||||
| 	case "PreConsumedQuota": | 	case "PreConsumedQuota": | ||||||
| 		config.PreConsumedQuota, _ = strconv.Atoi(value) | 		common.PreConsumedQuota, _ = strconv.Atoi(value) | ||||||
| 	case "RetryTimes": | 	case "RetryTimes": | ||||||
| 		config.RetryTimes, _ = strconv.Atoi(value) | 		common.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 "CompletionRatio": |  | ||||||
| 		err = common.UpdateCompletionRatioByJSONString(value) |  | ||||||
| 	case "TopUpLink": | 	case "TopUpLink": | ||||||
| 		config.TopUpLink = value | 		common.TopUpLink = value | ||||||
| 	case "ChatLink": | 	case "ChatLink": | ||||||
| 		config.ChatLink = value | 		common.ChatLink = value | ||||||
| 	case "ChannelDisableThreshold": | 	case "ChannelDisableThreshold": | ||||||
| 		config.ChannelDisableThreshold, _ = strconv.ParseFloat(value, 64) | 		common.ChannelDisableThreshold, _ = strconv.ParseFloat(value, 64) | ||||||
| 	case "QuotaPerUnit": | 	case "QuotaPerUnit": | ||||||
| 		config.QuotaPerUnit, _ = strconv.ParseFloat(value, 64) | 		common.QuotaPerUnit, _ = strconv.ParseFloat(value, 64) | ||||||
| 	case "Theme": | 	case "Theme": | ||||||
| 		config.Theme = value | 		common.Theme = value | ||||||
| 	} | 	} | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
|   | |||||||
| @@ -3,9 +3,8 @@ package model | |||||||
| import ( | import ( | ||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
|  | 	"one-api/common" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| type Redemption struct { | type Redemption struct { | ||||||
| @@ -68,7 +67,7 @@ func Redeem(key string, userId int) (quota int, err error) { | |||||||
| 		if err != nil { | 		if err != nil { | ||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 		redemption.RedeemedTime = helper.GetTimestamp() | 		redemption.RedeemedTime = common.GetTimestamp() | ||||||
| 		redemption.Status = common.RedemptionCodeStatusUsed | 		redemption.Status = common.RedemptionCodeStatusUsed | ||||||
| 		err = tx.Save(redemption).Error | 		err = tx.Save(redemption).Error | ||||||
| 		return err | 		return err | ||||||
|   | |||||||
| @@ -3,12 +3,8 @@ package model | |||||||
| import ( | import ( | ||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/message" |  | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
|  | 	"one-api/common" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| type Token struct { | type Token struct { | ||||||
| @@ -43,7 +39,7 @@ func ValidateUserToken(key string) (token *Token, err error) { | |||||||
| 	} | 	} | ||||||
| 	token, err = CacheGetTokenByKey(key) | 	token, err = CacheGetTokenByKey(key) | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.SysError("CacheGetTokenByKey failed: " + err.Error()) | 		common.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("无效的令牌") | ||||||
| 		} | 		} | ||||||
| @@ -57,12 +53,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 < helper.GetTimestamp() { | 	if token.ExpiredTime != -1 && token.ExpiredTime < common.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 { | ||||||
| 				logger.SysError("failed to update token status" + err.Error()) | 				common.SysError("failed to update token status" + err.Error()) | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 		return nil, errors.New("该令牌已过期") | 		return nil, errors.New("该令牌已过期") | ||||||
| @@ -73,7 +69,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 { | ||||||
| 				logger.SysError("failed to update token status" + err.Error()) | 				common.SysError("failed to update token status" + err.Error()) | ||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 		return nil, errors.New("该令牌额度已用尽") | 		return nil, errors.New("该令牌额度已用尽") | ||||||
| @@ -142,7 +138,7 @@ func IncreaseTokenQuota(id int, quota int) (err error) { | |||||||
| 	if quota < 0 { | 	if quota < 0 { | ||||||
| 		return errors.New("quota 不能为负数!") | 		return errors.New("quota 不能为负数!") | ||||||
| 	} | 	} | ||||||
| 	if config.BatchUpdateEnabled { | 	if common.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeTokenQuota, id, quota) | 		addNewRecord(BatchUpdateTypeTokenQuota, id, quota) | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| @@ -154,7 +150,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": helper.GetTimestamp(), | 			"accessed_time": common.GetTimestamp(), | ||||||
| 		}, | 		}, | ||||||
| 	).Error | 	).Error | ||||||
| 	return err | 	return err | ||||||
| @@ -164,7 +160,7 @@ func DecreaseTokenQuota(id int, quota int) (err error) { | |||||||
| 	if quota < 0 { | 	if quota < 0 { | ||||||
| 		return errors.New("quota 不能为负数!") | 		return errors.New("quota 不能为负数!") | ||||||
| 	} | 	} | ||||||
| 	if config.BatchUpdateEnabled { | 	if common.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeTokenQuota, id, -quota) | 		addNewRecord(BatchUpdateTypeTokenQuota, id, -quota) | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| @@ -176,7 +172,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": helper.GetTimestamp(), | 			"accessed_time": common.GetTimestamp(), | ||||||
| 		}, | 		}, | ||||||
| 	).Error | 	).Error | ||||||
| 	return err | 	return err | ||||||
| @@ -200,24 +196,24 @@ func PreConsumeTokenQuota(tokenId int, quota int) (err error) { | |||||||
| 	if userQuota < quota { | 	if userQuota < quota { | ||||||
| 		return errors.New("用户额度不足") | 		return errors.New("用户额度不足") | ||||||
| 	} | 	} | ||||||
| 	quotaTooLow := userQuota >= config.QuotaRemindThreshold && userQuota-quota < config.QuotaRemindThreshold | 	quotaTooLow := userQuota >= common.QuotaRemindThreshold && userQuota-quota < common.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 { | ||||||
| 				logger.SysError("failed to fetch user email: " + err.Error()) | 				common.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", config.ServerAddress) | 				topUpLink := fmt.Sprintf("%s/topup", common.ServerAddress) | ||||||
| 				err = message.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 { | ||||||
| 					logger.SysError("failed to send email" + err.Error()) | 					common.SysError("failed to send email" + err.Error()) | ||||||
| 				} | 				} | ||||||
| 			} | 			} | ||||||
| 		}() | 		}() | ||||||
|   | |||||||
| @@ -3,12 +3,8 @@ package model | |||||||
| import ( | import ( | ||||||
| 	"errors" | 	"errors" | ||||||
| 	"fmt" | 	"fmt" | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/blacklist" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"gorm.io/gorm" | 	"gorm.io/gorm" | ||||||
|  | 	"one-api/common" | ||||||
| 	"strings" | 	"strings" | ||||||
| ) | ) | ||||||
|  |  | ||||||
| @@ -19,7 +15,7 @@ type User struct { | |||||||
| 	Username         string `json:"username" gorm:"unique;index" validate:"max=12"` | 	Username         string `json:"username" gorm:"unique;index" validate:"max=12"` | ||||||
| 	Password         string `json:"password" gorm:"not null;" validate:"min=8,max=20"` | 	Password         string `json:"password" gorm:"not null;" validate:"min=8,max=20"` | ||||||
| 	DisplayName      string `json:"display_name" gorm:"index" validate:"max=20"` | 	DisplayName      string `json:"display_name" gorm:"index" validate:"max=20"` | ||||||
| 	Role             int    `json:"role" gorm:"type:int;default:1"`   // admin, util | 	Role             int    `json:"role" gorm:"type:int;default:1"`   // admin, common | ||||||
| 	Status           int    `json:"status" gorm:"type:int;default:1"` // enabled, disabled | 	Status           int    `json:"status" gorm:"type:int;default:1"` // enabled, disabled | ||||||
| 	Email            string `json:"email" gorm:"index" validate:"max=50"` | 	Email            string `json:"email" gorm:"index" validate:"max=50"` | ||||||
| 	GitHubId         string `json:"github_id" gorm:"column:github_id;index"` | 	GitHubId         string `json:"github_id" gorm:"column:github_id;index"` | ||||||
| @@ -41,7 +37,7 @@ func GetMaxUserId() int { | |||||||
| } | } | ||||||
|  |  | ||||||
| func GetAllUsers(startIdx int, num int) (users []*User, err error) { | func GetAllUsers(startIdx int, num int) (users []*User, err error) { | ||||||
| 	err = DB.Order("id desc").Limit(num).Offset(startIdx).Omit("password").Where("status != ?", common.UserStatusDeleted).Find(&users).Error | 	err = DB.Order("id desc").Limit(num).Offset(startIdx).Omit("password").Find(&users).Error | ||||||
| 	return users, err | 	return users, err | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -93,24 +89,24 @@ func (user *User) Insert(inviterId int) error { | |||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	user.Quota = config.QuotaForNewUser | 	user.Quota = common.QuotaForNewUser | ||||||
| 	user.AccessToken = helper.GetUUID() | 	user.AccessToken = common.GetUUID() | ||||||
| 	user.AffCode = helper.GetRandomString(4) | 	user.AffCode = common.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 config.QuotaForNewUser > 0 { | 	if common.QuotaForNewUser > 0 { | ||||||
| 		RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", common.LogQuota(config.QuotaForNewUser))) | 		RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("新用户注册赠送 %s", common.LogQuota(common.QuotaForNewUser))) | ||||||
| 	} | 	} | ||||||
| 	if inviterId != 0 { | 	if inviterId != 0 { | ||||||
| 		if config.QuotaForInvitee > 0 { | 		if common.QuotaForInvitee > 0 { | ||||||
| 			_ = IncreaseUserQuota(user.Id, config.QuotaForInvitee) | 			_ = IncreaseUserQuota(user.Id, common.QuotaForInvitee) | ||||||
| 			RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("使用邀请码赠送 %s", common.LogQuota(config.QuotaForInvitee))) | 			RecordLog(user.Id, LogTypeSystem, fmt.Sprintf("使用邀请码赠送 %s", common.LogQuota(common.QuotaForInvitee))) | ||||||
| 		} | 		} | ||||||
| 		if config.QuotaForInviter > 0 { | 		if common.QuotaForInviter > 0 { | ||||||
| 			_ = IncreaseUserQuota(inviterId, config.QuotaForInviter) | 			_ = IncreaseUserQuota(inviterId, common.QuotaForInviter) | ||||||
| 			RecordLog(inviterId, LogTypeSystem, fmt.Sprintf("邀请用户赠送 %s", common.LogQuota(config.QuotaForInviter))) | 			RecordLog(inviterId, LogTypeSystem, fmt.Sprintf("邀请用户赠送 %s", common.LogQuota(common.QuotaForInviter))) | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	return nil | 	return nil | ||||||
| @@ -124,11 +120,6 @@ func (user *User) Update(updatePassword bool) error { | |||||||
| 			return err | 			return err | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	if user.Status == common.UserStatusDisabled { |  | ||||||
| 		blacklist.BanUser(user.Id) |  | ||||||
| 	} else if user.Status == common.UserStatusEnabled { |  | ||||||
| 		blacklist.UnbanUser(user.Id) |  | ||||||
| 	} |  | ||||||
| 	err = DB.Model(user).Updates(user).Error | 	err = DB.Model(user).Updates(user).Error | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
| @@ -137,10 +128,7 @@ func (user *User) Delete() error { | |||||||
| 	if user.Id == 0 { | 	if user.Id == 0 { | ||||||
| 		return errors.New("id 为空!") | 		return errors.New("id 为空!") | ||||||
| 	} | 	} | ||||||
| 	blacklist.BanUser(user.Id) | 	err := DB.Delete(user).Error | ||||||
| 	user.Username = fmt.Sprintf("deleted_%s", helper.GetUUID()) |  | ||||||
| 	user.Status = common.UserStatusDeleted |  | ||||||
| 	err := DB.Model(user).Updates(user).Error |  | ||||||
| 	return err | 	return err | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -153,15 +141,7 @@ func (user *User) ValidateAndFill() (err error) { | |||||||
| 	if user.Username == "" || password == "" { | 	if user.Username == "" || password == "" { | ||||||
| 		return errors.New("用户名或密码为空") | 		return errors.New("用户名或密码为空") | ||||||
| 	} | 	} | ||||||
| 	err = DB.Where("username = ?", user.Username).First(user).Error | 	DB.Where(User{Username: user.Username}).First(user) | ||||||
| 	if err != nil { |  | ||||||
| 		// we must make sure check username firstly |  | ||||||
| 		// consider this case: a malicious user set his username as other's email |  | ||||||
| 		err := DB.Where("email = ?", user.Username).First(user).Error |  | ||||||
| 		if err != nil { |  | ||||||
| 			return errors.New("用户名或密码错误,或用户已被封禁") |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	okay := common.ValidatePasswordAndHash(password, user.Password) | 	okay := common.ValidatePasswordAndHash(password, user.Password) | ||||||
| 	if !okay || user.Status != common.UserStatusEnabled { | 	if !okay || user.Status != common.UserStatusEnabled { | ||||||
| 		return errors.New("用户名或密码错误,或用户已被封禁") | 		return errors.New("用户名或密码错误,或用户已被封禁") | ||||||
| @@ -244,7 +224,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 { | ||||||
| 		logger.SysError("no such user " + err.Error()) | 		common.SysError("no such user " + err.Error()) | ||||||
| 		return false | 		return false | ||||||
| 	} | 	} | ||||||
| 	return user.Role >= common.RoleAdminUser | 	return user.Role >= common.RoleAdminUser | ||||||
| @@ -303,7 +283,7 @@ func IncreaseUserQuota(id int, quota int) (err error) { | |||||||
| 	if quota < 0 { | 	if quota < 0 { | ||||||
| 		return errors.New("quota 不能为负数!") | 		return errors.New("quota 不能为负数!") | ||||||
| 	} | 	} | ||||||
| 	if config.BatchUpdateEnabled { | 	if common.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeUserQuota, id, quota) | 		addNewRecord(BatchUpdateTypeUserQuota, id, quota) | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| @@ -319,7 +299,7 @@ func DecreaseUserQuota(id int, quota int) (err error) { | |||||||
| 	if quota < 0 { | 	if quota < 0 { | ||||||
| 		return errors.New("quota 不能为负数!") | 		return errors.New("quota 不能为负数!") | ||||||
| 	} | 	} | ||||||
| 	if config.BatchUpdateEnabled { | 	if common.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeUserQuota, id, -quota) | 		addNewRecord(BatchUpdateTypeUserQuota, id, -quota) | ||||||
| 		return nil | 		return nil | ||||||
| 	} | 	} | ||||||
| @@ -337,7 +317,7 @@ func GetRootUserEmail() (email string) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func UpdateUserUsedQuotaAndRequestCount(id int, quota int) { | func UpdateUserUsedQuotaAndRequestCount(id int, quota int) { | ||||||
| 	if config.BatchUpdateEnabled { | 	if common.BatchUpdateEnabled { | ||||||
| 		addNewRecord(BatchUpdateTypeUsedQuota, id, quota) | 		addNewRecord(BatchUpdateTypeUsedQuota, id, quota) | ||||||
| 		addNewRecord(BatchUpdateTypeRequestCount, id, 1) | 		addNewRecord(BatchUpdateTypeRequestCount, id, 1) | ||||||
| 		return | 		return | ||||||
| @@ -353,7 +333,7 @@ func updateUserUsedQuotaAndRequestCount(id int, quota int, count int) { | |||||||
| 		}, | 		}, | ||||||
| 	).Error | 	).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.SysError("failed to update user used quota and request count: " + err.Error()) | 		common.SysError("failed to update user used quota and request count: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
| @@ -364,14 +344,14 @@ func updateUserUsedQuota(id int, quota int) { | |||||||
| 		}, | 		}, | ||||||
| 	).Error | 	).Error | ||||||
| 	if err != nil { | 	if err != nil { | ||||||
| 		logger.SysError("failed to update user used quota: " + err.Error()) | 		common.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 { | ||||||
| 		logger.SysError("failed to update user request count: " + err.Error()) | 		common.SysError("failed to update user request count: " + err.Error()) | ||||||
| 	} | 	} | ||||||
| } | } | ||||||
|  |  | ||||||
|   | |||||||
| @@ -1,8 +1,7 @@ | |||||||
| package model | package model | ||||||
|  |  | ||||||
| import ( | import ( | ||||||
| 	"github.com/songquanpeng/one-api/common/config" | 	"one-api/common" | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"sync" | 	"sync" | ||||||
| 	"time" | 	"time" | ||||||
| ) | ) | ||||||
| @@ -29,7 +28,7 @@ func init() { | |||||||
| func InitBatchUpdater() { | func InitBatchUpdater() { | ||||||
| 	go func() { | 	go func() { | ||||||
| 		for { | 		for { | ||||||
| 			time.Sleep(time.Duration(config.BatchUpdateInterval) * time.Second) | 			time.Sleep(time.Duration(common.BatchUpdateInterval) * time.Second) | ||||||
| 			batchUpdate() | 			batchUpdate() | ||||||
| 		} | 		} | ||||||
| 	}() | 	}() | ||||||
| @@ -46,7 +45,7 @@ func addNewRecord(type_ int, id int, value int) { | |||||||
| } | } | ||||||
|  |  | ||||||
| func batchUpdate() { | func batchUpdate() { | ||||||
| 	logger.SysLog("batch update started") | 	common.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] | ||||||
| @@ -58,12 +57,12 @@ func batchUpdate() { | |||||||
| 			case BatchUpdateTypeUserQuota: | 			case BatchUpdateTypeUserQuota: | ||||||
| 				err := increaseUserQuota(key, value) | 				err := increaseUserQuota(key, value) | ||||||
| 				if err != nil { | 				if err != nil { | ||||||
| 					logger.SysError("failed to batch update user quota: " + err.Error()) | 					common.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 { | ||||||
| 					logger.SysError("failed to batch update token quota: " + err.Error()) | 					common.SysError("failed to batch update token quota: " + err.Error()) | ||||||
| 				} | 				} | ||||||
| 			case BatchUpdateTypeUsedQuota: | 			case BatchUpdateTypeUsedQuota: | ||||||
| 				updateUserUsedQuota(key, value) | 				updateUserUsedQuota(key, value) | ||||||
| @@ -74,5 +73,5 @@ func batchUpdate() { | |||||||
| 			} | 			} | ||||||
| 		} | 		} | ||||||
| 	} | 	} | ||||||
| 	logger.SysLog("batch update finished") | 	common.SysLog("batch update finished") | ||||||
| } | } | ||||||
|   | |||||||
| @@ -1,55 +0,0 @@ | |||||||
| package monitor |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/message" |  | ||||||
| 	"github.com/songquanpeng/one-api/model" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| func notifyRootUser(subject string, content string) { |  | ||||||
| 	if config.MessagePusherAddress != "" { |  | ||||||
| 		err := message.SendMessage(subject, content, content) |  | ||||||
| 		if err != nil { |  | ||||||
| 			logger.SysError(fmt.Sprintf("failed to send message: %s", err.Error())) |  | ||||||
| 		} else { |  | ||||||
| 			return |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	if config.RootUserEmail == "" { |  | ||||||
| 		config.RootUserEmail = model.GetRootUserEmail() |  | ||||||
| 	} |  | ||||||
| 	err := message.SendEmail(subject, config.RootUserEmail, content) |  | ||||||
| 	if err != nil { |  | ||||||
| 		logger.SysError(fmt.Sprintf("failed to send email: %s", err.Error())) |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| // DisableChannel disable & notify |  | ||||||
| func DisableChannel(channelId int, channelName string, reason string) { |  | ||||||
| 	model.UpdateChannelStatusById(channelId, common.ChannelStatusAutoDisabled) |  | ||||||
| 	logger.SysLog(fmt.Sprintf("channel #%d has been disabled: %s", channelId, reason)) |  | ||||||
| 	subject := fmt.Sprintf("通道「%s」(#%d)已被禁用", channelName, channelId) |  | ||||||
| 	content := fmt.Sprintf("通道「%s」(#%d)已被禁用,原因:%s", channelName, channelId, reason) |  | ||||||
| 	notifyRootUser(subject, content) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func MetricDisableChannel(channelId int, successRate float64) { |  | ||||||
| 	model.UpdateChannelStatusById(channelId, common.ChannelStatusAutoDisabled) |  | ||||||
| 	logger.SysLog(fmt.Sprintf("channel #%d has been disabled due to low success rate: %.2f", channelId, successRate*100)) |  | ||||||
| 	subject := fmt.Sprintf("通道 #%d 已被禁用", channelId) |  | ||||||
| 	content := fmt.Sprintf("该渠道在最近 %d 次调用中成功率为 %.2f%%,低于阈值 %.2f%%,因此被系统自动禁用。", |  | ||||||
| 		config.MetricQueueSize, successRate*100, config.MetricSuccessRateThreshold*100) |  | ||||||
| 	notifyRootUser(subject, content) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| // EnableChannel enable & notify |  | ||||||
| func EnableChannel(channelId int, channelName string) { |  | ||||||
| 	model.UpdateChannelStatusById(channelId, common.ChannelStatusEnabled) |  | ||||||
| 	logger.SysLog(fmt.Sprintf("channel #%d has been enabled", channelId)) |  | ||||||
| 	subject := fmt.Sprintf("通道「%s」(#%d)已被启用", channelName, channelId) |  | ||||||
| 	content := fmt.Sprintf("通道「%s」(#%d)已被启用", channelName, channelId) |  | ||||||
| 	notifyRootUser(subject, content) |  | ||||||
| } |  | ||||||
| @@ -1,79 +0,0 @@ | |||||||
| package monitor |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"github.com/songquanpeng/one-api/common/config" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| var store = make(map[int][]bool) |  | ||||||
| var metricSuccessChan = make(chan int, config.MetricSuccessChanSize) |  | ||||||
| var metricFailChan = make(chan int, config.MetricFailChanSize) |  | ||||||
|  |  | ||||||
| func consumeSuccess(channelId int) { |  | ||||||
| 	if len(store[channelId]) > config.MetricQueueSize { |  | ||||||
| 		store[channelId] = store[channelId][1:] |  | ||||||
| 	} |  | ||||||
| 	store[channelId] = append(store[channelId], true) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func consumeFail(channelId int) (bool, float64) { |  | ||||||
| 	if len(store[channelId]) > config.MetricQueueSize { |  | ||||||
| 		store[channelId] = store[channelId][1:] |  | ||||||
| 	} |  | ||||||
| 	store[channelId] = append(store[channelId], false) |  | ||||||
| 	successCount := 0 |  | ||||||
| 	for _, success := range store[channelId] { |  | ||||||
| 		if success { |  | ||||||
| 			successCount++ |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	successRate := float64(successCount) / float64(len(store[channelId])) |  | ||||||
| 	if len(store[channelId]) < config.MetricQueueSize { |  | ||||||
| 		return false, successRate |  | ||||||
| 	} |  | ||||||
| 	if successRate < config.MetricSuccessRateThreshold { |  | ||||||
| 		store[channelId] = make([]bool, 0) |  | ||||||
| 		return true, successRate |  | ||||||
| 	} |  | ||||||
| 	return false, successRate |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func metricSuccessConsumer() { |  | ||||||
| 	for { |  | ||||||
| 		select { |  | ||||||
| 		case channelId := <-metricSuccessChan: |  | ||||||
| 			consumeSuccess(channelId) |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func metricFailConsumer() { |  | ||||||
| 	for { |  | ||||||
| 		select { |  | ||||||
| 		case channelId := <-metricFailChan: |  | ||||||
| 			disable, successRate := consumeFail(channelId) |  | ||||||
| 			if disable { |  | ||||||
| 				go MetricDisableChannel(channelId, successRate) |  | ||||||
| 			} |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func init() { |  | ||||||
| 	if config.EnableMetric { |  | ||||||
| 		go metricSuccessConsumer() |  | ||||||
| 		go metricFailConsumer() |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func Emit(channelId int, success bool) { |  | ||||||
| 	if !config.EnableMetric { |  | ||||||
| 		return |  | ||||||
| 	} |  | ||||||
| 	go func() { |  | ||||||
| 		if success { |  | ||||||
| 			metricSuccessChan <- channelId |  | ||||||
| 		} else { |  | ||||||
| 			metricFailChan <- channelId |  | ||||||
| 		} |  | ||||||
| 	}() |  | ||||||
| } |  | ||||||
| @@ -1,8 +0,0 @@ | |||||||
| package ai360 |  | ||||||
|  |  | ||||||
| var ModelList = []string{ |  | ||||||
| 	"360GPT_S2_V9", |  | ||||||
| 	"embedding-bert-512-v1", |  | ||||||
| 	"embedding_s1_v1", |  | ||||||
| 	"semantic_similarity_s1_v1", |  | ||||||
| } |  | ||||||
| @@ -1,60 +0,0 @@ | |||||||
| package aiproxy |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"errors" |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/gin-gonic/gin" |  | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" |  | ||||||
| 	"net/http" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| type Adaptor struct { |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) Init(meta *util.RelayMeta) { |  | ||||||
|  |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetRequestURL(meta *util.RelayMeta) (string, error) { |  | ||||||
| 	return fmt.Sprintf("%s/api/library/ask", meta.BaseURL), nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Request, meta *util.RelayMeta) error { |  | ||||||
| 	channel.SetupCommonRequestHeader(c, req, meta) |  | ||||||
| 	req.Header.Set("Authorization", "Bearer "+meta.APIKey) |  | ||||||
| 	return nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) ConvertRequest(c *gin.Context, relayMode int, request *model.GeneralOpenAIRequest) (any, error) { |  | ||||||
| 	if request == nil { |  | ||||||
| 		return nil, errors.New("request is nil") |  | ||||||
| 	} |  | ||||||
| 	aiProxyLibraryRequest := ConvertRequest(*request) |  | ||||||
| 	aiProxyLibraryRequest.LibraryId = c.GetString(common.ConfigKeyLibraryID) |  | ||||||
| 	return aiProxyLibraryRequest, nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoRequest(c *gin.Context, meta *util.RelayMeta, requestBody io.Reader) (*http.Response, error) { |  | ||||||
| 	return channel.DoRequestHelper(a, c, meta, requestBody) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, meta *util.RelayMeta) (usage *model.Usage, err *model.ErrorWithStatusCode) { |  | ||||||
| 	if meta.IsStream { |  | ||||||
| 		err, usage = StreamHandler(c, resp) |  | ||||||
| 	} else { |  | ||||||
| 		err, usage = Handler(c, resp) |  | ||||||
| 	} |  | ||||||
| 	return |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetModelList() []string { |  | ||||||
| 	return ModelList |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetChannelName() string { |  | ||||||
| 	return "aiproxy" |  | ||||||
| } |  | ||||||
| @@ -1,9 +0,0 @@ | |||||||
| package aiproxy |  | ||||||
|  |  | ||||||
| import "github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
|  |  | ||||||
| var ModelList = []string{""} |  | ||||||
|  |  | ||||||
| func init() { |  | ||||||
| 	ModelList = openai.ModelList |  | ||||||
| } |  | ||||||
| @@ -1,32 +0,0 @@ | |||||||
| package aiproxy |  | ||||||
|  |  | ||||||
| type LibraryRequest struct { |  | ||||||
| 	Model     string `json:"model"` |  | ||||||
| 	Query     string `json:"query"` |  | ||||||
| 	LibraryId string `json:"libraryId"` |  | ||||||
| 	Stream    bool   `json:"stream"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type LibraryError struct { |  | ||||||
| 	ErrCode int    `json:"errCode"` |  | ||||||
| 	Message string `json:"message"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type LibraryDocument struct { |  | ||||||
| 	Title string `json:"title"` |  | ||||||
| 	URL   string `json:"url"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type LibraryResponse struct { |  | ||||||
| 	Success   bool              `json:"success"` |  | ||||||
| 	Answer    string            `json:"answer"` |  | ||||||
| 	Documents []LibraryDocument `json:"documents"` |  | ||||||
| 	LibraryError |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type LibraryStreamResponse struct { |  | ||||||
| 	Content   string            `json:"content"` |  | ||||||
| 	Finish    bool              `json:"finish"` |  | ||||||
| 	Model     string            `json:"model"` |  | ||||||
| 	Documents []LibraryDocument `json:"documents"` |  | ||||||
| } |  | ||||||
| @@ -1,83 +0,0 @@ | |||||||
| package ali |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"errors" |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/gin-gonic/gin" |  | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" |  | ||||||
| 	"net/http" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| // https://help.aliyun.com/zh/dashscope/developer-reference/api-details |  | ||||||
|  |  | ||||||
| type Adaptor struct { |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) Init(meta *util.RelayMeta) { |  | ||||||
|  |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetRequestURL(meta *util.RelayMeta) (string, error) { |  | ||||||
| 	fullRequestURL := fmt.Sprintf("%s/api/v1/services/aigc/text-generation/generation", meta.BaseURL) |  | ||||||
| 	if meta.Mode == constant.RelayModeEmbeddings { |  | ||||||
| 		fullRequestURL = fmt.Sprintf("%s/api/v1/services/embeddings/text-embedding/text-embedding", meta.BaseURL) |  | ||||||
| 	} |  | ||||||
| 	return fullRequestURL, nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Request, meta *util.RelayMeta) error { |  | ||||||
| 	channel.SetupCommonRequestHeader(c, req, meta) |  | ||||||
| 	req.Header.Set("Authorization", "Bearer "+meta.APIKey) |  | ||||||
| 	if meta.IsStream { |  | ||||||
| 		req.Header.Set("X-DashScope-SSE", "enable") |  | ||||||
| 	} |  | ||||||
| 	if c.GetString(common.ConfigKeyPlugin) != "" { |  | ||||||
| 		req.Header.Set("X-DashScope-Plugin", c.GetString(common.ConfigKeyPlugin)) |  | ||||||
| 	} |  | ||||||
| 	return nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) ConvertRequest(c *gin.Context, relayMode int, request *model.GeneralOpenAIRequest) (any, error) { |  | ||||||
| 	if request == nil { |  | ||||||
| 		return nil, errors.New("request is nil") |  | ||||||
| 	} |  | ||||||
| 	switch relayMode { |  | ||||||
| 	case constant.RelayModeEmbeddings: |  | ||||||
| 		baiduEmbeddingRequest := ConvertEmbeddingRequest(*request) |  | ||||||
| 		return baiduEmbeddingRequest, nil |  | ||||||
| 	default: |  | ||||||
| 		baiduRequest := ConvertRequest(*request) |  | ||||||
| 		return baiduRequest, nil |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoRequest(c *gin.Context, meta *util.RelayMeta, requestBody io.Reader) (*http.Response, error) { |  | ||||||
| 	return channel.DoRequestHelper(a, c, meta, requestBody) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, meta *util.RelayMeta) (usage *model.Usage, err *model.ErrorWithStatusCode) { |  | ||||||
| 	if meta.IsStream { |  | ||||||
| 		err, usage = StreamHandler(c, resp) |  | ||||||
| 	} else { |  | ||||||
| 		switch meta.Mode { |  | ||||||
| 		case constant.RelayModeEmbeddings: |  | ||||||
| 			err, usage = EmbeddingHandler(c, resp) |  | ||||||
| 		default: |  | ||||||
| 			err, usage = Handler(c, resp) |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	return |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetModelList() []string { |  | ||||||
| 	return ModelList |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetChannelName() string { |  | ||||||
| 	return "ali" |  | ||||||
| } |  | ||||||
| @@ -1,6 +0,0 @@ | |||||||
| package ali |  | ||||||
|  |  | ||||||
| var ModelList = []string{ |  | ||||||
| 	"qwen-turbo", "qwen-plus", "qwen-max", "qwen-max-longcontext", |  | ||||||
| 	"text-embedding-v1", |  | ||||||
| } |  | ||||||
| @@ -1,73 +0,0 @@ | |||||||
| package ali |  | ||||||
|  |  | ||||||
| type Message struct { |  | ||||||
| 	Content string `json:"content"` |  | ||||||
| 	Role    string `json:"role"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Input struct { |  | ||||||
| 	//Prompt   string       `json:"prompt"` |  | ||||||
| 	Messages []Message `json:"messages"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Parameters struct { |  | ||||||
| 	TopP              float64 `json:"top_p,omitempty"` |  | ||||||
| 	TopK              int     `json:"top_k,omitempty"` |  | ||||||
| 	Seed              uint64  `json:"seed,omitempty"` |  | ||||||
| 	EnableSearch      bool    `json:"enable_search,omitempty"` |  | ||||||
| 	IncrementalOutput bool    `json:"incremental_output,omitempty"` |  | ||||||
| 	MaxTokens         int     `json:"max_tokens,omitempty"` |  | ||||||
| 	Temperature       float64 `json:"temperature,omitempty"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type ChatRequest struct { |  | ||||||
| 	Model      string     `json:"model"` |  | ||||||
| 	Input      Input      `json:"input"` |  | ||||||
| 	Parameters Parameters `json:"parameters,omitempty"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type EmbeddingRequest struct { |  | ||||||
| 	Model string `json:"model"` |  | ||||||
| 	Input struct { |  | ||||||
| 		Texts []string `json:"texts"` |  | ||||||
| 	} `json:"input"` |  | ||||||
| 	Parameters *struct { |  | ||||||
| 		TextType string `json:"text_type,omitempty"` |  | ||||||
| 	} `json:"parameters,omitempty"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Embedding struct { |  | ||||||
| 	Embedding []float64 `json:"embedding"` |  | ||||||
| 	TextIndex int       `json:"text_index"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type EmbeddingResponse struct { |  | ||||||
| 	Output struct { |  | ||||||
| 		Embeddings []Embedding `json:"embeddings"` |  | ||||||
| 	} `json:"output"` |  | ||||||
| 	Usage Usage `json:"usage"` |  | ||||||
| 	Error |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Error struct { |  | ||||||
| 	Code      string `json:"code"` |  | ||||||
| 	Message   string `json:"message"` |  | ||||||
| 	RequestId string `json:"request_id"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Usage struct { |  | ||||||
| 	InputTokens  int `json:"input_tokens"` |  | ||||||
| 	OutputTokens int `json:"output_tokens"` |  | ||||||
| 	TotalTokens  int `json:"total_tokens"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Output struct { |  | ||||||
| 	Text         string `json:"text"` |  | ||||||
| 	FinishReason string `json:"finish_reason"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type ChatResponse struct { |  | ||||||
| 	Output Output `json:"output"` |  | ||||||
| 	Usage  Usage  `json:"usage"` |  | ||||||
| 	Error |  | ||||||
| } |  | ||||||
| @@ -1,63 +0,0 @@ | |||||||
| package anthropic |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"errors" |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/gin-gonic/gin" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" |  | ||||||
| 	"net/http" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| type Adaptor struct { |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) Init(meta *util.RelayMeta) { |  | ||||||
|  |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetRequestURL(meta *util.RelayMeta) (string, error) { |  | ||||||
| 	return fmt.Sprintf("%s/v1/messages", meta.BaseURL), nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Request, meta *util.RelayMeta) error { |  | ||||||
| 	channel.SetupCommonRequestHeader(c, req, meta) |  | ||||||
| 	req.Header.Set("x-api-key", meta.APIKey) |  | ||||||
| 	anthropicVersion := c.Request.Header.Get("anthropic-version") |  | ||||||
| 	if anthropicVersion == "" { |  | ||||||
| 		anthropicVersion = "2023-06-01" |  | ||||||
| 	} |  | ||||||
| 	req.Header.Set("anthropic-version", anthropicVersion) |  | ||||||
| 	req.Header.Set("anthropic-beta", "messages-2023-12-15") |  | ||||||
| 	return nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) ConvertRequest(c *gin.Context, relayMode int, request *model.GeneralOpenAIRequest) (any, error) { |  | ||||||
| 	if request == nil { |  | ||||||
| 		return nil, errors.New("request is nil") |  | ||||||
| 	} |  | ||||||
| 	return ConvertRequest(*request), nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoRequest(c *gin.Context, meta *util.RelayMeta, requestBody io.Reader) (*http.Response, error) { |  | ||||||
| 	return channel.DoRequestHelper(a, c, meta, requestBody) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, meta *util.RelayMeta) (usage *model.Usage, err *model.ErrorWithStatusCode) { |  | ||||||
| 	if meta.IsStream { |  | ||||||
| 		err, usage = StreamHandler(c, resp) |  | ||||||
| 	} else { |  | ||||||
| 		err, usage = Handler(c, resp, meta.PromptTokens, meta.ActualModelName) |  | ||||||
| 	} |  | ||||||
| 	return |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetModelList() []string { |  | ||||||
| 	return ModelList |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetChannelName() string { |  | ||||||
| 	return "authropic" |  | ||||||
| } |  | ||||||
| @@ -1,8 +0,0 @@ | |||||||
| package anthropic |  | ||||||
|  |  | ||||||
| var ModelList = []string{ |  | ||||||
| 	"claude-instant-1.2", "claude-2.0", "claude-2.1", |  | ||||||
| 	"claude-3-haiku-20240229", |  | ||||||
| 	"claude-3-sonnet-20240229", |  | ||||||
| 	"claude-3-opus-20240229", |  | ||||||
| } |  | ||||||
| @@ -1,272 +0,0 @@ | |||||||
| package anthropic |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"bufio" |  | ||||||
| 	"encoding/json" |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/gin-gonic/gin" |  | ||||||
| 	"github.com/songquanpeng/one-api/common" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/image" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/logger" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"io" |  | ||||||
| 	"net/http" |  | ||||||
| 	"strings" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| func stopReasonClaude2OpenAI(reason *string) string { |  | ||||||
| 	if reason == nil { |  | ||||||
| 		return "" |  | ||||||
| 	} |  | ||||||
| 	switch *reason { |  | ||||||
| 	case "end_turn": |  | ||||||
| 		return "stop" |  | ||||||
| 	case "stop_sequence": |  | ||||||
| 		return "stop" |  | ||||||
| 	case "max_tokens": |  | ||||||
| 		return "length" |  | ||||||
| 	default: |  | ||||||
| 		return *reason |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func ConvertRequest(textRequest model.GeneralOpenAIRequest) *Request { |  | ||||||
| 	claudeRequest := Request{ |  | ||||||
| 		Model:       textRequest.Model, |  | ||||||
| 		MaxTokens:   textRequest.MaxTokens, |  | ||||||
| 		Temperature: textRequest.Temperature, |  | ||||||
| 		TopP:        textRequest.TopP, |  | ||||||
| 		Stream:      textRequest.Stream, |  | ||||||
| 	} |  | ||||||
| 	if claudeRequest.MaxTokens == 0 { |  | ||||||
| 		claudeRequest.MaxTokens = 4096 |  | ||||||
| 	} |  | ||||||
| 	// legacy model name mapping |  | ||||||
| 	if claudeRequest.Model == "claude-instant-1" { |  | ||||||
| 		claudeRequest.Model = "claude-instant-1.1" |  | ||||||
| 	} else if claudeRequest.Model == "claude-2" { |  | ||||||
| 		claudeRequest.Model = "claude-2.1" |  | ||||||
| 	} |  | ||||||
| 	for _, message := range textRequest.Messages { |  | ||||||
| 		if message.Role == "system" && claudeRequest.System == "" { |  | ||||||
| 			claudeRequest.System = message.StringContent() |  | ||||||
| 			continue |  | ||||||
| 		} |  | ||||||
| 		claudeMessage := Message{ |  | ||||||
| 			Role: message.Role, |  | ||||||
| 		} |  | ||||||
| 		var content Content |  | ||||||
| 		if message.IsStringContent() { |  | ||||||
| 			content.Type = "text" |  | ||||||
| 			content.Text = message.StringContent() |  | ||||||
| 			claudeMessage.Content = append(claudeMessage.Content, content) |  | ||||||
| 			claudeRequest.Messages = append(claudeRequest.Messages, claudeMessage) |  | ||||||
| 			continue |  | ||||||
| 		} |  | ||||||
| 		var contents []Content |  | ||||||
| 		openaiContent := message.ParseContent() |  | ||||||
| 		for _, part := range openaiContent { |  | ||||||
| 			var content Content |  | ||||||
| 			if part.Type == model.ContentTypeText { |  | ||||||
| 				content.Type = "text" |  | ||||||
| 				content.Text = part.Text |  | ||||||
| 			} else if part.Type == model.ContentTypeImageURL { |  | ||||||
| 				content.Type = "image" |  | ||||||
| 				content.Source = &ImageSource{ |  | ||||||
| 					Type: "base64", |  | ||||||
| 				} |  | ||||||
| 				mimeType, data, _ := image.GetImageFromUrl(part.ImageURL.Url) |  | ||||||
| 				content.Source.MediaType = mimeType |  | ||||||
| 				content.Source.Data = data |  | ||||||
| 			} |  | ||||||
| 			contents = append(contents, content) |  | ||||||
| 		} |  | ||||||
| 		claudeMessage.Content = contents |  | ||||||
| 		claudeRequest.Messages = append(claudeRequest.Messages, claudeMessage) |  | ||||||
| 	} |  | ||||||
| 	return &claudeRequest |  | ||||||
| } |  | ||||||
|  |  | ||||||
| // https://docs.anthropic.com/claude/reference/messages-streaming |  | ||||||
| func streamResponseClaude2OpenAI(claudeResponse *StreamResponse) (*openai.ChatCompletionsStreamResponse, *Response) { |  | ||||||
| 	var response *Response |  | ||||||
| 	var responseText string |  | ||||||
| 	var stopReason string |  | ||||||
| 	switch claudeResponse.Type { |  | ||||||
| 	case "message_start": |  | ||||||
| 		return nil, claudeResponse.Message |  | ||||||
| 	case "content_block_start": |  | ||||||
| 		if claudeResponse.ContentBlock != nil { |  | ||||||
| 			responseText = claudeResponse.ContentBlock.Text |  | ||||||
| 		} |  | ||||||
| 	case "content_block_delta": |  | ||||||
| 		if claudeResponse.Delta != nil { |  | ||||||
| 			responseText = claudeResponse.Delta.Text |  | ||||||
| 		} |  | ||||||
| 	case "message_delta": |  | ||||||
| 		if claudeResponse.Usage != nil { |  | ||||||
| 			response = &Response{ |  | ||||||
| 				Usage: *claudeResponse.Usage, |  | ||||||
| 			} |  | ||||||
| 		} |  | ||||||
| 		if claudeResponse.Delta != nil && claudeResponse.Delta.StopReason != nil { |  | ||||||
| 			stopReason = *claudeResponse.Delta.StopReason |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	var choice openai.ChatCompletionsStreamResponseChoice |  | ||||||
| 	choice.Delta.Content = responseText |  | ||||||
| 	choice.Delta.Role = "assistant" |  | ||||||
| 	finishReason := stopReasonClaude2OpenAI(&stopReason) |  | ||||||
| 	if finishReason != "null" { |  | ||||||
| 		choice.FinishReason = &finishReason |  | ||||||
| 	} |  | ||||||
| 	var openaiResponse openai.ChatCompletionsStreamResponse |  | ||||||
| 	openaiResponse.Object = "chat.completion.chunk" |  | ||||||
| 	openaiResponse.Choices = []openai.ChatCompletionsStreamResponseChoice{choice} |  | ||||||
| 	return &openaiResponse, response |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func responseClaude2OpenAI(claudeResponse *Response) *openai.TextResponse { |  | ||||||
| 	var responseText string |  | ||||||
| 	if len(claudeResponse.Content) > 0 { |  | ||||||
| 		responseText = claudeResponse.Content[0].Text |  | ||||||
| 	} |  | ||||||
| 	choice := openai.TextResponseChoice{ |  | ||||||
| 		Index: 0, |  | ||||||
| 		Message: model.Message{ |  | ||||||
| 			Role:    "assistant", |  | ||||||
| 			Content: responseText, |  | ||||||
| 			Name:    nil, |  | ||||||
| 		}, |  | ||||||
| 		FinishReason: stopReasonClaude2OpenAI(claudeResponse.StopReason), |  | ||||||
| 	} |  | ||||||
| 	fullTextResponse := openai.TextResponse{ |  | ||||||
| 		Id:      fmt.Sprintf("chatcmpl-%s", claudeResponse.Id), |  | ||||||
| 		Model:   claudeResponse.Model, |  | ||||||
| 		Object:  "chat.completion", |  | ||||||
| 		Created: helper.GetTimestamp(), |  | ||||||
| 		Choices: []openai.TextResponseChoice{choice}, |  | ||||||
| 	} |  | ||||||
| 	return &fullTextResponse |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) { |  | ||||||
| 	createdTime := helper.GetTimestamp() |  | ||||||
| 	scanner := bufio.NewScanner(resp.Body) |  | ||||||
| 	scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) { |  | ||||||
| 		if atEOF && len(data) == 0 { |  | ||||||
| 			return 0, nil, nil |  | ||||||
| 		} |  | ||||||
| 		if i := strings.Index(string(data), "\n"); i >= 0 { |  | ||||||
| 			return i + 1, data[0:i], nil |  | ||||||
| 		} |  | ||||||
| 		if atEOF { |  | ||||||
| 			return len(data), data, nil |  | ||||||
| 		} |  | ||||||
| 		return 0, nil, nil |  | ||||||
| 	}) |  | ||||||
| 	dataChan := make(chan string) |  | ||||||
| 	stopChan := make(chan bool) |  | ||||||
| 	go func() { |  | ||||||
| 		for scanner.Scan() { |  | ||||||
| 			data := scanner.Text() |  | ||||||
| 			if len(data) < 6 { |  | ||||||
| 				continue |  | ||||||
| 			} |  | ||||||
| 			if !strings.HasPrefix(data, "data: ") { |  | ||||||
| 				continue |  | ||||||
| 			} |  | ||||||
| 			data = strings.TrimPrefix(data, "data: ") |  | ||||||
| 			dataChan <- data |  | ||||||
| 		} |  | ||||||
| 		stopChan <- true |  | ||||||
| 	}() |  | ||||||
| 	common.SetEventStreamHeaders(c) |  | ||||||
| 	var usage model.Usage |  | ||||||
| 	var modelName string |  | ||||||
| 	var id string |  | ||||||
| 	c.Stream(func(w io.Writer) bool { |  | ||||||
| 		select { |  | ||||||
| 		case data := <-dataChan: |  | ||||||
| 			// some implementations may add \r at the end of data |  | ||||||
| 			data = strings.TrimSuffix(data, "\r") |  | ||||||
| 			var claudeResponse StreamResponse |  | ||||||
| 			err := json.Unmarshal([]byte(data), &claudeResponse) |  | ||||||
| 			if err != nil { |  | ||||||
| 				logger.SysError("error unmarshalling stream response: " + err.Error()) |  | ||||||
| 				return true |  | ||||||
| 			} |  | ||||||
| 			response, meta := streamResponseClaude2OpenAI(&claudeResponse) |  | ||||||
| 			if meta != nil { |  | ||||||
| 				usage.PromptTokens += meta.Usage.InputTokens |  | ||||||
| 				usage.CompletionTokens += meta.Usage.OutputTokens |  | ||||||
| 				modelName = meta.Model |  | ||||||
| 				id = fmt.Sprintf("chatcmpl-%s", meta.Id) |  | ||||||
| 				return true |  | ||||||
| 			} |  | ||||||
| 			if response == nil { |  | ||||||
| 				return true |  | ||||||
| 			} |  | ||||||
| 			response.Id = id |  | ||||||
| 			response.Model = modelName |  | ||||||
| 			response.Created = createdTime |  | ||||||
| 			jsonStr, err := json.Marshal(response) |  | ||||||
| 			if err != nil { |  | ||||||
| 				logger.SysError("error marshalling stream response: " + err.Error()) |  | ||||||
| 				return true |  | ||||||
| 			} |  | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonStr)}) |  | ||||||
| 			return true |  | ||||||
| 		case <-stopChan: |  | ||||||
| 			c.Render(-1, common.CustomEvent{Data: "data: [DONE]"}) |  | ||||||
| 			return false |  | ||||||
| 		} |  | ||||||
| 	}) |  | ||||||
| 	_ = resp.Body.Close() |  | ||||||
| 	return nil, &usage |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func Handler(c *gin.Context, resp *http.Response, promptTokens int, modelName string) (*model.ErrorWithStatusCode, *model.Usage) { |  | ||||||
| 	responseBody, err := io.ReadAll(resp.Body) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return openai.ErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil |  | ||||||
| 	} |  | ||||||
| 	err = resp.Body.Close() |  | ||||||
| 	if err != nil { |  | ||||||
| 		return openai.ErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil |  | ||||||
| 	} |  | ||||||
| 	var claudeResponse Response |  | ||||||
| 	err = json.Unmarshal(responseBody, &claudeResponse) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil |  | ||||||
| 	} |  | ||||||
| 	if claudeResponse.Error.Type != "" { |  | ||||||
| 		return &model.ErrorWithStatusCode{ |  | ||||||
| 			Error: model.Error{ |  | ||||||
| 				Message: claudeResponse.Error.Message, |  | ||||||
| 				Type:    claudeResponse.Error.Type, |  | ||||||
| 				Param:   "", |  | ||||||
| 				Code:    claudeResponse.Error.Type, |  | ||||||
| 			}, |  | ||||||
| 			StatusCode: resp.StatusCode, |  | ||||||
| 		}, nil |  | ||||||
| 	} |  | ||||||
| 	fullTextResponse := responseClaude2OpenAI(&claudeResponse) |  | ||||||
| 	fullTextResponse.Model = modelName |  | ||||||
| 	usage := model.Usage{ |  | ||||||
| 		PromptTokens:     claudeResponse.Usage.InputTokens, |  | ||||||
| 		CompletionTokens: claudeResponse.Usage.OutputTokens, |  | ||||||
| 		TotalTokens:      claudeResponse.Usage.InputTokens + claudeResponse.Usage.OutputTokens, |  | ||||||
| 	} |  | ||||||
| 	fullTextResponse.Usage = usage |  | ||||||
| 	jsonResponse, err := json.Marshal(fullTextResponse) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil |  | ||||||
| 	} |  | ||||||
| 	c.Writer.Header().Set("Content-Type", "application/json") |  | ||||||
| 	c.Writer.WriteHeader(resp.StatusCode) |  | ||||||
| 	_, err = c.Writer.Write(jsonResponse) |  | ||||||
| 	return nil, &usage |  | ||||||
| } |  | ||||||
| @@ -1,75 +0,0 @@ | |||||||
| package anthropic |  | ||||||
|  |  | ||||||
| // https://docs.anthropic.com/claude/reference/messages_post |  | ||||||
|  |  | ||||||
| type Metadata struct { |  | ||||||
| 	UserId string `json:"user_id"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type ImageSource struct { |  | ||||||
| 	Type      string `json:"type"` |  | ||||||
| 	MediaType string `json:"media_type"` |  | ||||||
| 	Data      string `json:"data"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Content struct { |  | ||||||
| 	Type   string       `json:"type"` |  | ||||||
| 	Text   string       `json:"text,omitempty"` |  | ||||||
| 	Source *ImageSource `json:"source,omitempty"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Message struct { |  | ||||||
| 	Role    string    `json:"role"` |  | ||||||
| 	Content []Content `json:"content"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Request struct { |  | ||||||
| 	Model         string    `json:"model"` |  | ||||||
| 	Messages      []Message `json:"messages"` |  | ||||||
| 	System        string    `json:"system,omitempty"` |  | ||||||
| 	MaxTokens     int       `json:"max_tokens,omitempty"` |  | ||||||
| 	StopSequences []string  `json:"stop_sequences,omitempty"` |  | ||||||
| 	Stream        bool      `json:"stream,omitempty"` |  | ||||||
| 	Temperature   float64   `json:"temperature,omitempty"` |  | ||||||
| 	TopP          float64   `json:"top_p,omitempty"` |  | ||||||
| 	TopK          int       `json:"top_k,omitempty"` |  | ||||||
| 	//Metadata    `json:"metadata,omitempty"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Usage struct { |  | ||||||
| 	InputTokens  int `json:"input_tokens"` |  | ||||||
| 	OutputTokens int `json:"output_tokens"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Error struct { |  | ||||||
| 	Type    string `json:"type"` |  | ||||||
| 	Message string `json:"message"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Response struct { |  | ||||||
| 	Id           string    `json:"id"` |  | ||||||
| 	Type         string    `json:"type"` |  | ||||||
| 	Role         string    `json:"role"` |  | ||||||
| 	Content      []Content `json:"content"` |  | ||||||
| 	Model        string    `json:"model"` |  | ||||||
| 	StopReason   *string   `json:"stop_reason"` |  | ||||||
| 	StopSequence *string   `json:"stop_sequence"` |  | ||||||
| 	Usage        Usage     `json:"usage"` |  | ||||||
| 	Error        Error     `json:"error"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type Delta struct { |  | ||||||
| 	Type         string  `json:"type"` |  | ||||||
| 	Text         string  `json:"text"` |  | ||||||
| 	StopReason   *string `json:"stop_reason"` |  | ||||||
| 	StopSequence *string `json:"stop_sequence"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type StreamResponse struct { |  | ||||||
| 	Type         string    `json:"type"` |  | ||||||
| 	Message      *Response `json:"message"` |  | ||||||
| 	Index        int       `json:"index"` |  | ||||||
| 	ContentBlock *Content  `json:"content_block"` |  | ||||||
| 	Delta        *Delta    `json:"delta"` |  | ||||||
| 	Usage        *Usage    `json:"usage"` |  | ||||||
| } |  | ||||||
| @@ -1,7 +0,0 @@ | |||||||
| package baichuan |  | ||||||
|  |  | ||||||
| var ModelList = []string{ |  | ||||||
| 	"Baichuan2-Turbo", |  | ||||||
| 	"Baichuan2-Turbo-192k", |  | ||||||
| 	"Baichuan-Text-Embedding", |  | ||||||
| } |  | ||||||
| @@ -1,105 +0,0 @@ | |||||||
| package baidu |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"errors" |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/gin-gonic/gin" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/constant" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" |  | ||||||
| 	"net/http" |  | ||||||
| 	"strings" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| type Adaptor struct { |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) Init(meta *util.RelayMeta) { |  | ||||||
|  |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetRequestURL(meta *util.RelayMeta) (string, error) { |  | ||||||
| 	// https://cloud.baidu.com/doc/WENXINWORKSHOP/s/clntwmv7t |  | ||||||
| 	suffix := "chat/" |  | ||||||
| 	if strings.HasPrefix("Embedding", meta.ActualModelName) { |  | ||||||
| 		suffix = "embeddings/" |  | ||||||
| 	} |  | ||||||
| 	switch meta.ActualModelName { |  | ||||||
| 	case "ERNIE-4.0": |  | ||||||
| 		suffix += "completions_pro" |  | ||||||
| 	case "ERNIE-Bot-4": |  | ||||||
| 		suffix += "completions_pro" |  | ||||||
| 	case "ERNIE-3.5-8K": |  | ||||||
| 		suffix += "completions" |  | ||||||
| 	case "ERNIE-Bot-8K": |  | ||||||
| 		suffix += "ernie_bot_8k" |  | ||||||
| 	case "ERNIE-Bot": |  | ||||||
| 		suffix += "completions" |  | ||||||
| 	case "ERNIE-Speed": |  | ||||||
| 		suffix += "ernie_speed" |  | ||||||
| 	case "ERNIE-Bot-turbo": |  | ||||||
| 		suffix += "eb-instant" |  | ||||||
| 	case "BLOOMZ-7B": |  | ||||||
| 		suffix += "bloomz_7b1" |  | ||||||
| 	case "Embedding-V1": |  | ||||||
| 		suffix += "embedding-v1" |  | ||||||
| 	default: |  | ||||||
| 		suffix += meta.ActualModelName |  | ||||||
| 	} |  | ||||||
| 	fullRequestURL := fmt.Sprintf("%s/rpc/2.0/ai_custom/v1/wenxinworkshop/%s", meta.BaseURL, suffix) |  | ||||||
| 	var accessToken string |  | ||||||
| 	var err error |  | ||||||
| 	if accessToken, err = GetAccessToken(meta.APIKey); err != nil { |  | ||||||
| 		return "", err |  | ||||||
| 	} |  | ||||||
| 	fullRequestURL += "?access_token=" + accessToken |  | ||||||
| 	return fullRequestURL, nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Request, meta *util.RelayMeta) error { |  | ||||||
| 	channel.SetupCommonRequestHeader(c, req, meta) |  | ||||||
| 	req.Header.Set("Authorization", "Bearer "+meta.APIKey) |  | ||||||
| 	return nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) ConvertRequest(c *gin.Context, relayMode int, request *model.GeneralOpenAIRequest) (any, error) { |  | ||||||
| 	if request == nil { |  | ||||||
| 		return nil, errors.New("request is nil") |  | ||||||
| 	} |  | ||||||
| 	switch relayMode { |  | ||||||
| 	case constant.RelayModeEmbeddings: |  | ||||||
| 		baiduEmbeddingRequest := ConvertEmbeddingRequest(*request) |  | ||||||
| 		return baiduEmbeddingRequest, nil |  | ||||||
| 	default: |  | ||||||
| 		baiduRequest := ConvertRequest(*request) |  | ||||||
| 		return baiduRequest, nil |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoRequest(c *gin.Context, meta *util.RelayMeta, requestBody io.Reader) (*http.Response, error) { |  | ||||||
| 	return channel.DoRequestHelper(a, c, meta, requestBody) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, meta *util.RelayMeta) (usage *model.Usage, err *model.ErrorWithStatusCode) { |  | ||||||
| 	if meta.IsStream { |  | ||||||
| 		err, usage = StreamHandler(c, resp) |  | ||||||
| 	} else { |  | ||||||
| 		switch meta.Mode { |  | ||||||
| 		case constant.RelayModeEmbeddings: |  | ||||||
| 			err, usage = EmbeddingHandler(c, resp) |  | ||||||
| 		default: |  | ||||||
| 			err, usage = Handler(c, resp) |  | ||||||
| 		} |  | ||||||
| 	} |  | ||||||
| 	return |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetModelList() []string { |  | ||||||
| 	return ModelList |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetChannelName() string { |  | ||||||
| 	return "baidu" |  | ||||||
| } |  | ||||||
| @@ -1,10 +0,0 @@ | |||||||
| package baidu |  | ||||||
|  |  | ||||||
| var ModelList = []string{ |  | ||||||
| 	"ERNIE-Bot-4", |  | ||||||
| 	"ERNIE-Bot-8K", |  | ||||||
| 	"ERNIE-Bot", |  | ||||||
| 	"ERNIE-Speed", |  | ||||||
| 	"ERNIE-Bot-turbo", |  | ||||||
| 	"Embedding-V1", |  | ||||||
| } |  | ||||||
| @@ -1,50 +0,0 @@ | |||||||
| package baidu |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"time" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| type ChatResponse struct { |  | ||||||
| 	Id               string      `json:"id"` |  | ||||||
| 	Object           string      `json:"object"` |  | ||||||
| 	Created          int64       `json:"created"` |  | ||||||
| 	Result           string      `json:"result"` |  | ||||||
| 	IsTruncated      bool        `json:"is_truncated"` |  | ||||||
| 	NeedClearHistory bool        `json:"need_clear_history"` |  | ||||||
| 	Usage            model.Usage `json:"usage"` |  | ||||||
| 	Error |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type ChatStreamResponse struct { |  | ||||||
| 	ChatResponse |  | ||||||
| 	SentenceId int  `json:"sentence_id"` |  | ||||||
| 	IsEnd      bool `json:"is_end"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type EmbeddingRequest struct { |  | ||||||
| 	Input []string `json:"input"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type EmbeddingData struct { |  | ||||||
| 	Object    string    `json:"object"` |  | ||||||
| 	Embedding []float64 `json:"embedding"` |  | ||||||
| 	Index     int       `json:"index"` |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type EmbeddingResponse struct { |  | ||||||
| 	Id      string          `json:"id"` |  | ||||||
| 	Object  string          `json:"object"` |  | ||||||
| 	Created int64           `json:"created"` |  | ||||||
| 	Data    []EmbeddingData `json:"data"` |  | ||||||
| 	Usage   model.Usage     `json:"usage"` |  | ||||||
| 	Error |  | ||||||
| } |  | ||||||
|  |  | ||||||
| type AccessToken struct { |  | ||||||
| 	AccessToken      string    `json:"access_token"` |  | ||||||
| 	Error            string    `json:"error,omitempty"` |  | ||||||
| 	ErrorDescription string    `json:"error_description,omitempty"` |  | ||||||
| 	ExpiresIn        int64     `json:"expires_in,omitempty"` |  | ||||||
| 	ExpiresAt        time.Time `json:"-"` |  | ||||||
| } |  | ||||||
| @@ -1,51 +0,0 @@ | |||||||
| package channel |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"errors" |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/gin-gonic/gin" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" |  | ||||||
| 	"net/http" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| func SetupCommonRequestHeader(c *gin.Context, req *http.Request, meta *util.RelayMeta) { |  | ||||||
| 	req.Header.Set("Content-Type", c.Request.Header.Get("Content-Type")) |  | ||||||
| 	req.Header.Set("Accept", c.Request.Header.Get("Accept")) |  | ||||||
| 	if meta.IsStream && c.Request.Header.Get("Accept") == "" { |  | ||||||
| 		req.Header.Set("Accept", "text/event-stream") |  | ||||||
| 	} |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func DoRequestHelper(a Adaptor, c *gin.Context, meta *util.RelayMeta, requestBody io.Reader) (*http.Response, error) { |  | ||||||
| 	fullRequestURL, err := a.GetRequestURL(meta) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return nil, fmt.Errorf("get request url failed: %w", err) |  | ||||||
| 	} |  | ||||||
| 	req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return nil, fmt.Errorf("new request failed: %w", err) |  | ||||||
| 	} |  | ||||||
| 	err = a.SetupRequestHeader(c, req, meta) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return nil, fmt.Errorf("setup request header failed: %w", err) |  | ||||||
| 	} |  | ||||||
| 	resp, err := DoRequest(c, req) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return nil, fmt.Errorf("do request failed: %w", err) |  | ||||||
| 	} |  | ||||||
| 	return resp, nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func DoRequest(c *gin.Context, req *http.Request) (*http.Response, error) { |  | ||||||
| 	resp, err := util.HTTPClient.Do(req) |  | ||||||
| 	if err != nil { |  | ||||||
| 		return nil, err |  | ||||||
| 	} |  | ||||||
| 	if resp == nil { |  | ||||||
| 		return nil, errors.New("resp is nil") |  | ||||||
| 	} |  | ||||||
| 	_ = req.Body.Close() |  | ||||||
| 	_ = c.Request.Body.Close() |  | ||||||
| 	return resp, nil |  | ||||||
| } |  | ||||||
| @@ -1,66 +0,0 @@ | |||||||
| package gemini |  | ||||||
|  |  | ||||||
| import ( |  | ||||||
| 	"errors" |  | ||||||
| 	"fmt" |  | ||||||
| 	"github.com/gin-gonic/gin" |  | ||||||
| 	"github.com/songquanpeng/one-api/common/helper" |  | ||||||
| 	channelhelper "github.com/songquanpeng/one-api/relay/channel" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/channel/openai" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/model" |  | ||||||
| 	"github.com/songquanpeng/one-api/relay/util" |  | ||||||
| 	"io" |  | ||||||
| 	"net/http" |  | ||||||
| ) |  | ||||||
|  |  | ||||||
| type Adaptor struct { |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) Init(meta *util.RelayMeta) { |  | ||||||
|  |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetRequestURL(meta *util.RelayMeta) (string, error) { |  | ||||||
| 	version := helper.AssignOrDefault(meta.APIVersion, "v1") |  | ||||||
| 	action := "generateContent" |  | ||||||
| 	if meta.IsStream { |  | ||||||
| 		action = "streamGenerateContent" |  | ||||||
| 	} |  | ||||||
| 	return fmt.Sprintf("%s/%s/models/%s:%s", meta.BaseURL, version, meta.ActualModelName, action), nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Request, meta *util.RelayMeta) error { |  | ||||||
| 	channelhelper.SetupCommonRequestHeader(c, req, meta) |  | ||||||
| 	req.Header.Set("x-goog-api-key", meta.APIKey) |  | ||||||
| 	return nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) ConvertRequest(c *gin.Context, relayMode int, request *model.GeneralOpenAIRequest) (any, error) { |  | ||||||
| 	if request == nil { |  | ||||||
| 		return nil, errors.New("request is nil") |  | ||||||
| 	} |  | ||||||
| 	return ConvertRequest(*request), nil |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoRequest(c *gin.Context, meta *util.RelayMeta, requestBody io.Reader) (*http.Response, error) { |  | ||||||
| 	return channelhelper.DoRequestHelper(a, c, meta, requestBody) |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, meta *util.RelayMeta) (usage *model.Usage, err *model.ErrorWithStatusCode) { |  | ||||||
| 	if meta.IsStream { |  | ||||||
| 		var responseText string |  | ||||||
| 		err, responseText = StreamHandler(c, resp) |  | ||||||
| 		usage = openai.ResponseText2Usage(responseText, meta.ActualModelName, meta.PromptTokens) |  | ||||||
| 	} else { |  | ||||||
| 		err, usage = Handler(c, resp, meta.PromptTokens, meta.ActualModelName) |  | ||||||
| 	} |  | ||||||
| 	return |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetModelList() []string { |  | ||||||
| 	return ModelList |  | ||||||
| } |  | ||||||
|  |  | ||||||
| func (a *Adaptor) GetChannelName() string { |  | ||||||
| 	return "google gemini" |  | ||||||
| } |  | ||||||
| @@ -1,6 +0,0 @@ | |||||||
| package gemini |  | ||||||
|  |  | ||||||
| var ModelList = []string{ |  | ||||||
| 	"gemini-pro", "gemini-1.0-pro-001", |  | ||||||
| 	"gemini-pro-vision", "gemini-1.0-pro-vision-001", |  | ||||||
| } |  | ||||||
Some files were not shown because too many files have changed in this diff Show More
		Reference in New Issue
	
	Block a user