This commit is contained in:
kazusa 2025-09-16 15:12:44 +08:00
commit 3642aaa941
42 changed files with 2826 additions and 0 deletions

228
pom.xml Normal file
View File

@ -0,0 +1,228 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion>
<groupId>cn.kazusa</groupId>
<artifactId>chatgpt-sdk</artifactId>
<version>1.0-SNAPSHOT</version>
<name>chatgpt-sdk-java</name>
<description>OpenAI Java SDK, ChatGPT Java SDK .</description>
<properties>
<java.version>1.8</java.version>
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
<maven.compiler.source>1.8</maven.compiler.source>
<maven.compiler.target>1.8</maven.compiler.target>
<retrofit2.version>2.9.0</retrofit2.version>
<slf4j.version>2.0.6</slf4j.version>
</properties>
<dependencies>
<dependency>
<groupId>org.slf4j</groupId>
<artifactId>slf4j-api</artifactId>
<version>${slf4j.version}</version>
</dependency>
<dependency>
<groupId>org.slf4j</groupId>
<artifactId>slf4j-simple</artifactId>
<version>${slf4j.version}</version>
</dependency>
<dependency>
<groupId>com.fasterxml.jackson.core</groupId>
<artifactId>jackson-databind</artifactId>
<version>2.13.3</version>
</dependency>
<dependency>
<groupId>cn.hutool</groupId>
<artifactId>hutool-all</artifactId>
<version>5.8.12</version>
</dependency>
<dependency>
<groupId>com.squareup.okhttp3</groupId>
<artifactId>okhttp-sse</artifactId>
<version>3.14.9</version>
</dependency>
<dependency>
<groupId>com.squareup.okhttp3</groupId>
<artifactId>logging-interceptor</artifactId>
<version>3.14.9</version>
</dependency>
<dependency>
<groupId>com.squareup.retrofit2</groupId>
<artifactId>retrofit</artifactId>
<version>${retrofit2.version}</version>
</dependency>
<dependency>
<groupId>com.squareup.retrofit2</groupId>
<artifactId>converter-jackson</artifactId>
<version>${retrofit2.version}</version>
</dependency>
<dependency>
<groupId>com.squareup.retrofit2</groupId>
<artifactId>adapter-rxjava2</artifactId>
<version>${retrofit2.version}</version>
</dependency>
<dependency>
<groupId>junit</groupId>
<artifactId>junit</artifactId>
<version>4.13.2</version>
<scope>test</scope>
</dependency>
<dependency>
<groupId>org.jetbrains</groupId>
<artifactId>annotations</artifactId>
<version>RELEASE</version>
<scope>compile</scope>
</dependency>
<dependency>
<groupId>com.knuddels</groupId>
<artifactId>jtokkit</artifactId>
<version>0.2.0</version>
</dependency>
<dependency>
<groupId>org.projectlombok</groupId>
<artifactId>lombok</artifactId>
<version>1.18.24</version>
<scope>compile</scope>
</dependency>
<dependency>
<groupId>com.alibaba</groupId>
<artifactId>fastjson</artifactId>
<version>2.0.28</version>
</dependency>
</dependencies>
<build>
<finalName>chatgpt-sdk-java</finalName>
<plugins>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-surefire-plugin</artifactId>
<version>2.12.4</version>
<configuration>
<skipTests>true</skipTests>
</configuration>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-compiler-plugin</artifactId>
<configuration>
<source>8</source>
<target>8</target>
</configuration>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-jar-plugin</artifactId>
<version>2.3.1</version>
<configuration>
<archive>
<manifest>
<addDefaultImplementationEntries>true</addDefaultImplementationEntries>
</manifest>
</archive>
</configuration>
</plugin>
</plugins>
</build>
<!-- 发布使用
<profiles>
<profile>
&lt;!&ndash;注意,此id必须与setting.xml中指定的一致,不要自作聪明改它名字&ndash;&gt;
<id>release</id>
<activation>
<activeByDefault>true</activeByDefault>
</activation>
<build>
&lt;!&ndash;发布到中央SNAPSHOT仓库插件&ndash;&gt;
<plugins>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-release-plugin</artifactId>
<version>2.5.3</version>
<configuration>
<autoVersionSubmodules>true</autoVersionSubmodules>
<useReleaseProfile>false</useReleaseProfile>
<releaseProfiles>release</releaseProfiles>
<goals>deploy</goals>
</configuration>
</plugin>
&lt;!&ndash;生成源码插件&ndash;&gt;
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-source-plugin</artifactId>
<version>2.2.1</version>
<executions>
<execution>
<id>attach-sources</id>
<goals>
<goal>jar-no-fork</goal>
</goals>
</execution>
</executions>
</plugin>
&lt;!&ndash;生成API文档插件&ndash;&gt;
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-javadoc-plugin</artifactId>
<version>2.9.1</version>
<executions>
<execution>
<id>attach-javadocs</id>
<goals>
<goal>jar</goal>
</goals>
<configuration>
<additionalparam>-Xdoclint:none</additionalparam>
&lt;!&ndash; /usr/libexec/java_home -V &ndash;&gt;
<javadocExecutable>/Library/Java/JavaVirtualMachines/jdk1.8.0_311.jdk/Contents/Home/bin/javadoc</javadocExecutable>
</configuration>
</execution>
</executions>
</plugin>
&lt;!&ndash;gpg插件&ndash;&gt;
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-gpg-plugin</artifactId>
<version>1.5</version>
<executions>
<execution>
<id>sign-artifacts</id>
<phase>verify</phase>
<goals>
<goal>sign</goal>
</goals>
<configuration>
<keyname>fuzhengwei</keyname>
<passphrase>123456</passphrase>
</configuration>
</execution>
</executions>
</plugin>
</plugins>
</build>
<distributionManagement>
<snapshotRepository>
&lt;!&ndash;注意,此id必须与setting.xml中指定的一致&ndash;&gt;
<id>ossrh</id>
<url>https://s01.oss.sonatype.org/content/repositories/snapshots</url>
</snapshotRepository>
<repository>
<id>release</id>
<url>https://s01.oss.sonatype.org/content/repositories/releases/</url>
</repository>
</distributionManagement>
</profile>
</profiles>-->
</project>

View File

@ -0,0 +1,261 @@
package cn.kazusa.chatgpt;
import cn.kazusa.chatgpt.domain.billing.BillingUsage;
import cn.kazusa.chatgpt.domain.billing.Subscription;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionRequest;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionResponse;
import cn.kazusa.chatgpt.domain.edits.EditRequest;
import cn.kazusa.chatgpt.domain.edits.EditResponse;
import cn.kazusa.chatgpt.domain.embedd.EmbeddingRequest;
import cn.kazusa.chatgpt.domain.embedd.EmbeddingResponse;
import cn.kazusa.chatgpt.domain.files.DeleteFileResponse;
import cn.kazusa.chatgpt.domain.files.UploadFileResponse;
import cn.kazusa.chatgpt.domain.images.ImageRequest;
import cn.kazusa.chatgpt.domain.images.ImageResponse;
import cn.kazusa.chatgpt.domain.other.OpenAiResponse;
import cn.kazusa.chatgpt.domain.qa.QACompletionRequest;
import cn.kazusa.chatgpt.domain.qa.QACompletionResponse;
import cn.kazusa.chatgpt.domain.whisper.WhisperResponse;
import io.reactivex.Single;
import okhttp3.MultipartBody;
import okhttp3.RequestBody;
import okhttp3.ResponseBody;
import retrofit2.http.Body;
import retrofit2.http.DELETE;
import retrofit2.http.GET;
import retrofit2.http.Multipart;
import retrofit2.http.POST;
import retrofit2.http.Part;
import retrofit2.http.PartMap;
import retrofit2.http.Path;
import retrofit2.http.Query;
import retrofit2.http.Streaming;
import java.io.File;
import java.time.LocalDate;
import java.util.Map;
/**
* @author kazusa
* @description ChatGPT 官网 API 模型定义接口官网https://platform.openai.com/playground
*/
public interface IOpenAiApi {
String v1_completions = "v1/completions";
/**
* 文本问答
*
* @param qaCompletionRequest 请求信息
* @return 应答结果
*/
@POST(v1_completions)
Single<QACompletionResponse> completions(@Body QACompletionRequest qaCompletionRequest);
String v1_chat_completions = "v1/chat/completions";
/**
* 问答模型默认 GPT-3.5
*
* @param chatCompletionRequest 请求信息
* @return 应答结果
*/
@POST(v1_chat_completions)
Single<ChatCompletionResponse> completions(@Body ChatCompletionRequest chatCompletionRequest);
/**
* 文本修复
*
* @param editRequest 请求信息编辑文本的参数
* @return 应答结果
*/
@POST("v1/edits")
Single<EditResponse> edits(@Body EditRequest editRequest);
/**
* 生成图片
* curl https://api.openai.com/v1/images/generations \
* -H "Content-Type: application/json" \
* -H "Authorization: Bearer $OPENAI_API_KEY" \
* -d '{
* "prompt": "A cute baby sea otter",
* "n": 2,
* "size": "1024x1024"
* }'
* <p>
* {
* "created": 1589478378,
* "data": [
* {
* "url": "https://..."
* },
* {
* "url": "https://..."
* }
* ]
* }
*
* @param imageRequest 图片对象
* @return 应答结果
*/
@POST("v1/images/generations")
Single<ImageResponse> genImages(@Body ImageRequest imageRequest);
/**
* 修改图片
* <p>
* curl https://api.openai.com/v1/images/edits \
* -H "Authorization: Bearer $OPENAI_API_KEY" \
* -F image="@otter.png" \
* -F mask="@mask.png" \
* -F prompt="A cute baby sea otter wearing a beret" \
* -F n=2 \
* -F size="1024x1024"
* <p>
* {
* "created": 1589478378,
* "data": [
* {
* "url": "https://..."
* },
* {
* "url": "https://..."
* }
* ]
* }
*
* @param image 图片对象
* @param mask 图片对象
* @param requestBodyMap 请求参数
* @return 应答结果
*/
@Multipart
@POST("v1/images/edits")
Single<ImageResponse> editImages(@Part MultipartBody.Part image, @Part MultipartBody.Part mask, @PartMap Map<String, RequestBody> requestBodyMap);
/**
* 向量计算
* curl https://api.openai.com/v1/images/variations \
* -H "Authorization: Bearer $OPENAI_API_KEY" \
* -F image="@otter.png" \
* -F n=2 \
* -F size="1024x1024"
*
* @param embeddingRequest 请求对象
* @return 应答结果
*/
@POST("v1/embeddings")
Single<EmbeddingResponse> embeddings(@Body EmbeddingRequest embeddingRequest);
/**
* 文件列表在你上传文件到服务端后可以获取列表信息
* curl https://api.openai.com/v1/files \
* -H "Authorization: Bearer $OPENAI_API_KEY"
*
* @return 应答结果
*/
@GET("v1/files")
Single<OpenAiResponse<File>> files();
/**
* 上传文件上载一个文件该文件包含要在各种端点/功能中使用的文档目前一个组织上传的所有文件的大小最多可达1GB如果您需要增加存储限制请与官网联系
* curl https://api.openai.com/v1/files \
* -H "Authorization: Bearer $OPENAI_API_KEY" \
* -F purpose="fine-tune" \
* -F file="@mydata.jsonl"
*
* @param file 文件
* @param purpose "fine-tune"
* @return 应答结果
*/
@Multipart
@POST("v1/files")
Single<UploadFileResponse> uploadFile(@Part MultipartBody.Part file, @Part("purpose") RequestBody purpose);
/**
* 删除文件
* curl https://api.openai.com/v1/files/file-XjGxS3KTG0uNmNOK362iJua3 \
* -X DELETE \
* -H "Authorization: Bearer $OPENAI_API_KEY"
*
* @param fileId 文件ID
* @return 应答结果
*/
@DELETE("v1/files/{file_id}")
Single<DeleteFileResponse> deleteFile(@Path("file_id") String fileId);
/**
* 检索文件
* curl https://api.openai.com/v1/files/file-XjGxS3KTG0uNmNOK362iJua3 \
* -H "Authorization: Bearer $OPENAI_API_KEY"
*
* @param fileId 文件ID
* @return 应答结果
*/
@GET("v1/files/{file_id}")
Single<File> retrieveFile(@Path("file_id") String fileId);
/**
* 检索文件内容信息
* curl https://api.openai.com/v1/files/file-XjGxS3KTG0uNmNOK362iJua3/content \
* -H "Authorization: Bearer $OPENAI_API_KEY" > file.jsonl
*
* @param fileId 文件ID
* @return 应答结果
*/
@Streaming
@GET("v1/files/{file_id}/content")
Single<ResponseBody> retrieveFileContent(@Path("file_id") String fileId);
/**
* 语音转文字
* curl https://api.openai.com/v1/audio/transcriptions \
* -H "Authorization: Bearer $OPENAI_API_KEY" \
* -H "Content-Type: multipart/form-data" \
* -F file="@/path/to/file/audio.mp3" \
* -F model="whisper-1"
*
* @param file 语音文件
* @param requestBodyMap 请求信息
* @return 应答结果
*/
@Multipart
@POST("v1/audio/transcriptions")
Single<WhisperResponse> speed2TextTranscriptions(@Part MultipartBody.Part file, @PartMap() Map<String, RequestBody> requestBodyMap);
/**
* 语音翻译
* curl https://api.openai.com/v1/audio/translations \
* -H "Authorization: Bearer $OPENAI_API_KEY" \
* -H "Content-Type: multipart/form-data" \
* -F file="@/path/to/file/german.m4a" \
* -F model="whisper-1"
*
* @param file 语音文件
* @param requestBodyMap 请求信息
* @return 应答结果
*/
@Multipart
@POST("v1/audio/translations")
Single<WhisperResponse> speed2TextTranslations(@Part MultipartBody.Part file, @PartMap() Map<String, RequestBody> requestBodyMap);
/**
* 账单查询
*
* @return 应答结果
*/
@GET("v1/dashboard/billing/subscription")
Single<Subscription> subscription();
/**
* 消耗查询
*
* @param starDate 开始时间
* @param endDate 结束时间
* @return 应答数据
*/
@GET("v1/dashboard/billing/usage")
Single<BillingUsage> billingUsage(@Query("start_date") LocalDate starDate, @Query("end_date") LocalDate endDate);
}

View File

@ -0,0 +1,34 @@
package cn.kazusa.chatgpt.common;
/**
* @author kazusa
* @description 通用类
*/
public class Constants {
public final static String NULL = "NULL";
/**
* 官网支持的请求角色类型systemuserassistant
* https://platform.openai.com/docs/guides/chat/introduction
*/
public enum Role {
SYSTEM("system"),
USER("user"),
ASSISTANT("assistant"),
;
private String code;
Role(String code) {
this.code = code;
}
public String getCode() {
return code;
}
}
}

View File

@ -0,0 +1,25 @@
package cn.kazusa.chatgpt.domain.billing;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.Data;
import java.math.BigDecimal;
import java.util.List;
/**
* @author kazusa
* @description 消耗账单
*/
@Data
public class BillingUsage {
@JsonProperty("object")
private String object;
/** 账号金额消耗明细 */
@JsonProperty("daily_costs")
private List<DailyCost> dailyCosts;
/** 总使用金额/美分 */
@JsonProperty("total_usage")
private BigDecimal totalUsage;
}

View File

@ -0,0 +1,18 @@
package cn.kazusa.chatgpt.domain.billing;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.Data;
import java.util.List;
/**
* @author kazusa
* @description 金额消耗
*/
@Data
public class DailyCost {
@JsonProperty("timestamp")
private long timestamp;
@JsonProperty("line_items")
private List<LineItem> lineItems;
}

View File

@ -0,0 +1,17 @@
package cn.kazusa.chatgpt.domain.billing;
import lombok.Data;
import java.math.BigDecimal;
/**
* @author kazusa
* @description 消耗列表数据
*/
@Data
public class LineItem {
/** 模型 */
private String name;
/** 金额 */
private BigDecimal cost;
}

View File

@ -0,0 +1,13 @@
package cn.kazusa.chatgpt.domain.billing;
import lombok.Data;
/**
* @author kazusa
* @description 账单
*/
@Data
public class Plan {
private String title;
private String id;
}

View File

@ -0,0 +1,54 @@
package cn.kazusa.chatgpt.domain.billing;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 账户信息
*/
@Data
public class Subscription implements Serializable {
@JsonProperty("object")
private String object;
@JsonProperty("has_payment_method")
private boolean hasPaymentMethod;
@JsonProperty("canceled")
private boolean canceled;
@JsonProperty("canceled_at")
private Object canceledAt;
@JsonProperty("delinquent")
private Object delinquent;
@JsonProperty("access_until")
private long accessUntil;
@JsonProperty("soft_limit")
private long softLimit;
@JsonProperty("hard_limit")
private long hardLimit;
@JsonProperty("system_hard_limit")
private long systemHardLimit;
@JsonProperty("soft_limit_usd")
private double softLimitUsd;
@JsonProperty("hard_limit_usd")
private double hardLimitUsd;
@JsonProperty("system_hard_limit_usd")
private double systemHardLimitUsd;
@JsonProperty("plan")
private Plan plan;
@JsonProperty("account_name")
private String accountName;
@JsonProperty("po_number")
private Object poNumber;
@JsonProperty("billing_email")
private Object billingEmail;
@JsonProperty("tax_ids")
private Object taxIds;
@JsonProperty("billing_address")
private Object billingAddress;
@JsonProperty("business_address")
private Object businessAddress;
}

View File

@ -0,0 +1,27 @@
package cn.kazusa.chatgpt.domain.chat;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 对话信息
*/
@Data
@JsonIgnoreProperties(ignoreUnknown = true)
public class ChatChoice implements Serializable {
private long index;
/** stream = true 请求参数里返回的属性是 delta */
@JsonProperty("delta")
private Message delta;
/** stream = false 请求参数里返回的属性是 delta */
@JsonProperty("message")
private Message message;
@JsonProperty("finish_reason")
private String finishReason;
}

View File

@ -0,0 +1,74 @@
package cn.kazusa.chatgpt.domain.chat;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.Getter;
import lombok.NoArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import java.io.Serializable;
import java.util.List;
import java.util.Map;
/**
* @author kazusa
* @description 对话聊天请求信息依照OpenAI官网API构建参数https://platform.openai.com/playground
*/
@Data
@Builder
@Slf4j
@JsonInclude(JsonInclude.Include.NON_NULL)
@NoArgsConstructor
@AllArgsConstructor
@JsonIgnoreProperties(ignoreUnknown = true)
public class ChatCompletionRequest implements Serializable {
/** 默认模型 */
private String model = Model.GPT_3_5_TURBO.getCode();
/** 问题描述 */
private List<Message> messages;
/** 控制温度【随机性】0到2之间。较高的值(如0.8)将使输出更加随机,而较低的值(如0.2)将使输出更加集中和确定 */
private double temperature = 0.2;
/** 多样性控制使用温度采样的替代方法称为核心采样其中模型考虑具有top_p概率质量的令牌的结果。因此0.1 意味着只考虑包含前 10% 概率质量的代币 */
@JsonProperty("top_p")
private Double topP = 1d;
/** 为每个提示生成的完成次数 */
private Integer n = 1;
/** 是否为流式输出;就是一蹦一蹦的,出来结果 */
private boolean stream = false;
/** 停止输出标识 */
private List<String> stop;
/** 输出字符串限制0 ~ 4096 */
@JsonProperty("max_tokens")
private Integer maxTokens = 2048;
/** 频率惩罚;降低模型重复同一行的可能性 */
@JsonProperty("frequency_penalty")
private double frequencyPenalty = 0;
/** 存在惩罚;增强模型谈论新话题的可能性 */
@JsonProperty("presence_penalty")
private double presencePenalty = 0;
/** 生成多个调用结果,只显示最佳的。这样会更多的消耗你的 api token */
@JsonProperty("logit_bias")
private Map logitBias;
/** 调用标识,避免重复调用 */
private String user;
@Getter
@AllArgsConstructor
public enum Model {
/** gpt-3.5-turbo */
GPT_3_5_TURBO("gpt-3.5-turbo"),
/** GPT4.0 */
GPT_4("gpt-4"),
GPT_4o("gpt-4o"),
/** GPT4.0 超长上下文 */
GPT_4_32K("gpt-4-32k"),
;
private String code;
}
}

View File

@ -0,0 +1,49 @@
package cn.kazusa.chatgpt.domain.chat;
import cn.kazusa.chatgpt.domain.other.Usage;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.Data;
import java.io.Serializable;
import java.util.List;
/**
* @author kazusa
* @description 对话请求结果信息
*/
@Data
@JsonIgnoreProperties(ignoreUnknown = true)
public class ChatCompletionResponse implements Serializable {
/**
* ID
*/
private String id;
/**
* 对象
*/
private String object;
/**
* 模型
*/
private String model;
/**
* 对话
*/
private List<ChatChoice> choices;
/**
* 创建
*/
private long created;
/**
* 耗材
*/
private Usage usage;
/**
* 该指纹代表模型运行时使用的后端配置
* https://platform.openai.com/docs/api-reference/chat
*/
@JsonProperty("system_fingerprint")
private String systemFingerprint;
}

View File

@ -0,0 +1,68 @@
package cn.kazusa.chatgpt.domain.chat;
import cn.kazusa.chatgpt.common.Constants;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import com.fasterxml.jackson.annotation.JsonInclude;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 聊天消息对象
*/
@Data
@JsonInclude(JsonInclude.Include.NON_NULL)
@JsonIgnoreProperties(ignoreUnknown = true)
public class Message implements Serializable {
private String role;
private String content;
private String name;
public Message() {
}
private Message(Builder builder) {
this.role = builder.role;
this.content = builder.content;
this.name = builder.name;
}
public static Builder builder() {
return new Builder();
}
/**
* 建造者模式
*/
public static final class Builder {
private String role;
private String content;
private String name;
public Builder() {
}
public Builder role(Constants.Role role) {
this.role = role.getCode();
return this;
}
public Builder content(String content) {
this.content = content;
return this;
}
public Builder name(String name) {
this.name = name;
return this;
}
public Message build() {
return new Message(this);
}
}
}

View File

@ -0,0 +1,53 @@
package cn.kazusa.chatgpt.domain.edits;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Getter;
import lombok.NoArgsConstructor;
import lombok.NonNull;
import lombok.extern.slf4j.Slf4j;
import java.io.Serializable;
/**
* @author kazusa
* @description 文本编辑请求
*/
@Slf4j
@Getter
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class EditRequest implements Serializable {
/** 模型 */
@NonNull
private String model = Model.CODE_DAVINCI_EDIT_001.getCode();
/** 输入 */
@NonNull
private String input;
/** 修改描述 */
@NonNull
private String instruction;
/** 控制温度【随机性】0到2之间。较高的值(如0.8)将使输出更加随机,而较低的值(如0.2)将使输出更加集中和确定 */
@Builder.Default
private double temperature = 0.2;
/** 多样性控制使用温度采样的替代方法称为核心采样其中模型考虑具有top_p概率质量的令牌的结果。因此0.1 意味着只考虑包含前 10% 概率质量的代币 */
@JsonProperty("top_p")
private Double topP = 1d;
/** 为每个提示生成的完成次数 */
private Integer n = 1;
@Getter
@AllArgsConstructor
public enum Model{
/** text-davinci-edit-001 */
TEXT_DAVINCI_EDIT_001("text-davinci-edit-001"),
/** code-davinci-edit-001 */
CODE_DAVINCI_EDIT_001("code-davinci-edit-001"),
;
private String code;
}
}

View File

@ -0,0 +1,29 @@
package cn.kazusa.chatgpt.domain.edits;
import cn.kazusa.chatgpt.domain.other.Choice;
import cn.kazusa.chatgpt.domain.other.Usage;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 文本编辑响应
*/
@Data
public class EditResponse implements Serializable {
/** ID */
private String id;
/** 对象 */
private String object;
/** 模型 */
private String model;
/** 对话 */
private Choice[] choices;
/** 创建 */
private long created;
/** 耗材 */
private Usage usage;
}

View File

@ -0,0 +1,45 @@
package cn.kazusa.chatgpt.domain.embedd;
import com.fasterxml.jackson.annotation.JsonInclude;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Getter;
import lombok.NoArgsConstructor;
import lombok.NonNull;
import lombok.Setter;
import lombok.extern.slf4j.Slf4j;
import java.io.Serializable;
import java.util.List;
/**
* @author kazusa
* @description Creates a variation of a given image.
*/
@Slf4j
@Getter
@Builder
@JsonInclude(JsonInclude.Include.NON_NULL)
@NoArgsConstructor
@AllArgsConstructor
public class EmbeddingRequest implements Serializable {
/** 模型 */
@NonNull
@Builder.Default
private String model = Model.TEXT_EMBEDDING_ADA_002.getCode();
/** 输入信息 */
@NonNull
private List<String> input;
@Setter
private String user;
@Getter
@AllArgsConstructor
public enum Model {
TEXT_EMBEDDING_ADA_002("text-embedding-ada-002"),
;
private String code;
}
}

View File

@ -0,0 +1,21 @@
package cn.kazusa.chatgpt.domain.embedd;
import cn.kazusa.chatgpt.domain.other.Usage;
import lombok.Data;
import java.io.Serializable;
import java.util.List;
/**
* @author kazusa
* @description 反馈对象
*/
@Data
public class EmbeddingResponse implements Serializable {
private String object;
private List<Item> data;
private String model;
private Usage usage;
}

View File

@ -0,0 +1,20 @@
package cn.kazusa.chatgpt.domain.embedd;
import lombok.Data;
import java.io.Serializable;
import java.math.BigDecimal;
import java.util.List;
/**
* @author kazusa
* @description 条目信息
*/
@Data
public class Item implements Serializable {
private String object;
private List<BigDecimal> embedding;
private Integer index;
}

View File

@ -0,0 +1,26 @@
package cn.kazusa.chatgpt.domain.files;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 删除文件应答
* {
* "id": "file-XjGxS3KTG0uNmNOK362iJua3",
* "object": "file",
* "deleted": true
* }
*/
@Data
public class DeleteFileResponse implements Serializable {
/** 文件ID */
private String id;
/** 对象file */
private String object;
/** 删除true */
private boolean deleted;
}

View File

@ -0,0 +1,32 @@
package cn.kazusa.chatgpt.domain.files;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 文件信息
*/
@Data
public class File implements Serializable {
/** ID */
private String id;
/** 对象 */
private String object;
/** 字节 */
private long bytes;
/** 时间 */
private long created_at;
/** 名称 */
private String filename;
/** 类型search */
private String purpose;
/** 状态 */
private String status;
@JsonProperty("status_details")
private String statusDetails;
}

View File

@ -0,0 +1,13 @@
package cn.kazusa.chatgpt.domain.files;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 上传文件应答封装个空的文件应答结果对象
*/
@Data
public class UploadFileResponse extends File implements Serializable {
}

View File

@ -0,0 +1,43 @@
package cn.kazusa.chatgpt.domain.images;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Getter;
import lombok.NoArgsConstructor;
import lombok.NonNull;
import lombok.Setter;
import lombok.extern.slf4j.Slf4j;
import java.io.Serializable;
/**
* @author kazusa
* @description 修改图片
*/
@Slf4j
@Getter
@Builder
@JsonInclude(JsonInclude.Include.NON_NULL)
@NoArgsConstructor
@AllArgsConstructor
public class ImageEditRequest extends ImageEnum implements Serializable {
/** 问题描述 */
@NonNull
private String prompt;
/** 为每个提示生成的完成次数 */
@Builder.Default
private Integer n = 1;
/** 图片大小 */
@Builder.Default
private String size = Size.size_256.getCode();
/** 图片格式化方式URL、B64_JSON */
@JsonProperty("response_format")
@Builder.Default
private String responseFormat = ResponseFormat.URL.getCode();
@Setter
private String user;
}

View File

@ -0,0 +1,31 @@
package cn.kazusa.chatgpt.domain.images;
import lombok.AllArgsConstructor;
import lombok.Getter;
/**
* @author kazusa
* @description 图片枚举配置
*/
public class ImageEnum {
@Getter
@AllArgsConstructor
public enum Size {
size_256("256x256"),
size_512("512x512"),
size_1024("1024x1024"),
;
private String code;
}
@Getter
@AllArgsConstructor
public enum ResponseFormat {
URL("url"),
B64_JSON("b64_json"),
;
private String code;
}
}

View File

@ -0,0 +1,52 @@
package cn.kazusa.chatgpt.domain.images;
import com.fasterxml.jackson.annotation.JsonInclude;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Getter;
import lombok.NoArgsConstructor;
import lombok.NonNull;
import lombok.extern.slf4j.Slf4j;
import java.io.Serializable;
/**
* @author kazusa
* @description 图片请求
*/
@Slf4j
@Getter
@Builder
@JsonInclude(JsonInclude.Include.NON_NULL)
@NoArgsConstructor
@AllArgsConstructor
public class ImageRequest extends ImageEnum implements Serializable {
/** 模型 */
private String model = Model.DALL_E_3.code;
/** 问题描述 */
@NonNull
private String prompt;
/** 为每个提示生成的完成次数 */
@Builder.Default
private Integer n = 1;
/** 图片大小 */
@Builder.Default
private String size = Size.size_1024.getCode();
/** 图片格式化方式URLB64_JSON
@JsonProperty("response_format")
@Builder.Default
private String responseFormat = ResponseFormat.URL.getCode();
@Setter
private String user; */
@Getter
@AllArgsConstructor
public enum Model {
DALL_E_2("dall-e-2"),
DALL_E_3("dall-e-3"),
;
private final String code;
}
}

View File

@ -0,0 +1,18 @@
package cn.kazusa.chatgpt.domain.images;
import lombok.Data;
import java.io.Serializable;
import java.util.List;
/**
* @author kazusa
* @description 图片响应
*/
@Data
public class ImageResponse implements Serializable {
/** 条目数据 */
private List<Item> data;
/** 创建时间 */
private long created;
}

View File

@ -0,0 +1,20 @@
package cn.kazusa.chatgpt.domain.images;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 条目
*/
@Data
public class Item implements Serializable {
private String url;
// @JsonProperty("b64_json")
// private String b64Json;
@JsonProperty("revised_prompt")
private String revisedPrompt;
}

View File

@ -0,0 +1,23 @@
package cn.kazusa.chatgpt.domain.other;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 对话信息
*/
@Data
@JsonIgnoreProperties(ignoreUnknown = true)
public class Choice implements Serializable {
private long index;
private String text;
private Object logprobs;
@JsonProperty("finish_reason")
private String finishReason;
}

View File

@ -0,0 +1,30 @@
package cn.kazusa.chatgpt.domain.other;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import lombok.Data;
import java.io.Serializable;
import java.util.List;
/**
* @author kazusa
* @description OpenAI通用响应对象
*/
@Data
@JsonIgnoreProperties(ignoreUnknown = true)
public class OpenAiResponse<T> implements Serializable {
private String object;
private List<T> data;
private Error error;
@Data
public class Error {
private String message;
private String type;
private String param;
private String code;
}
}

View File

@ -0,0 +1,49 @@
package cn.kazusa.chatgpt.domain.other;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import com.fasterxml.jackson.annotation.JsonProperty;
import java.io.Serializable;
/**
* @author kazusa
* @description 使用量
*/
@JsonIgnoreProperties(ignoreUnknown = true)
public class Usage implements Serializable {
/** 提示令牌 */
@JsonProperty("prompt_tokens")
private long promptTokens;
/** 完成令牌 */
@JsonProperty("completion_tokens")
private long completionTokens;
/** 总量令牌 */
@JsonProperty("total_tokens")
private long totalTokens;
public long getPromptTokens() {
return promptTokens;
}
public void setPromptTokens(long promptTokens) {
this.promptTokens = promptTokens;
}
public long getCompletionTokens() {
return completionTokens;
}
public void setCompletionTokens(long completionTokens) {
this.completionTokens = completionTokens;
}
public long getTotalTokens() {
return totalTokens;
}
public void setTotalTokens(long totalTokens) {
this.totalTokens = totalTokens;
}
}

View File

@ -0,0 +1,81 @@
package cn.kazusa.chatgpt.domain.qa;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.Getter;
import lombok.NoArgsConstructor;
import lombok.NonNull;
import lombok.extern.slf4j.Slf4j;
import java.io.Serializable;
import java.util.List;
import java.util.Map;
/**
* @author kazusa
* @description 文本问答请求信息OpenAI官网API构建参数https://platform.openai.com/playground
*/
@Data
@Builder
@Slf4j
@JsonInclude(JsonInclude.Include.NON_NULL)
@NoArgsConstructor
@AllArgsConstructor
@JsonIgnoreProperties(ignoreUnknown = true)
public class QACompletionRequest implements Serializable {
/** 默认模型 */
@NonNull
@Builder.Default
private String model = Model.TEXT_DAVINCI_003.getCode();
/** 问题描述 */
@NonNull
private String prompt;
private String suffix;
/** 控制温度【随机性】0到2之间。较高的值(如0.8)将使输出更加随机,而较低的值(如0.2)将使输出更加集中和确定 */
private double temperature = 0.2;
/** 多样性控制使用温度采样的替代方法称为核心采样其中模型考虑具有top_p概率质量的令牌的结果。因此0.1 意味着只考虑包含前 10% 概率质量的代币 */
@JsonProperty("top_p")
private Double topP = 1d;
/** 为每个提示生成的完成次数 */
private Integer n = 1;
/** 是否为流式输出;就是一蹦一蹦的,出来结果 */
private boolean stream = false;
/** 停止输出标识 */
private List<String> stop;
/** 输出字符串限制0 ~ 4096 */
@JsonProperty("max_tokens")
private Integer maxTokens = 2048;
@Builder.Default
private boolean echo = false;
/** 频率惩罚;降低模型重复同一行的可能性 */
@JsonProperty("frequency_penalty")
private double frequencyPenalty = 0;
/** 存在惩罚;增强模型谈论新话题的可能性 */
@JsonProperty("presence_penalty")
private double presencePenalty = 0;
/** 生成多个调用结果,只显示最佳的。这样会更多的消耗你的 api token */
@JsonProperty("best_of")
@Builder.Default
private Integer bestOf = 1;
private Integer logprobs;
@JsonProperty("logit_bias")
private Map logitBias;
/** 调用标识,避免重复调用 */
private String user;
@Getter
@AllArgsConstructor
public enum Model {
TEXT_DAVINCI_003("text-davinci-003"),
TEXT_DAVINCI_002("text-davinci-002"),
DAVINCI("davinci"),
;
private String code;
}
}

View File

@ -0,0 +1,31 @@
package cn.kazusa.chatgpt.domain.qa;
import cn.kazusa.chatgpt.domain.other.Choice;
import cn.kazusa.chatgpt.domain.other.Usage;
import com.fasterxml.jackson.annotation.JsonIgnoreProperties;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 问答完成响应
*/
@Data
@JsonIgnoreProperties(ignoreUnknown = true)
public class QACompletionResponse implements Serializable {
/** ID */
private String id;
/** 对象 */
private String object;
/** 模型 */
private String model;
/** 对话 */
private Choice[] choices;
/** 创建 */
private long created;
/** 耗材 */
private Usage usage;
}

View File

@ -0,0 +1,39 @@
package cn.kazusa.chatgpt.domain.whisper;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import lombok.experimental.FieldNameConstants;
/**
* @author kazusa
* @description 语音转录请求
*/
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
@FieldNameConstants
@JsonInclude(JsonInclude.Include.NON_NULL)
public class TranscriptionsRequest {
/** 模型WHISPER_1 */
@Builder.Default
private String model = WhisperEnum.Model.WHISPER_1.getCode();
/** 提示语 */
private String prompt;
/** 输出格式 */
@JsonProperty("response_format")
@Builder.Default
private String responseFormat = WhisperEnum.ResponseFormat.JSON.getCode();
/** 控制温度【随机性】0到2之间。较高的值(如0.8)将使输出更加随机,而较低的值(如0.2)将使输出更加集中和确定 */
private double temperature = 0.2;
/**
* 音频语言ISO-639-1
*/
private String language;
}

View File

@ -0,0 +1,37 @@
package cn.kazusa.chatgpt.domain.whisper;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import lombok.experimental.FieldNameConstants;
import java.io.Serializable;
/**
* @author kazusa
* @description 翻译请求
*/
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
@FieldNameConstants
@JsonInclude(JsonInclude.Include.NON_NULL)
public class TranslationsRequest implements Serializable {
/** 模型WHISPER_1 */
@Builder.Default
private String model = WhisperEnum.Model.WHISPER_1.getCode();
/** 提示语 */
private String prompt;
/** 输出格式 */
@JsonProperty("response_format")
@Builder.Default
private String responseFormat = WhisperEnum.ResponseFormat.JSON.getCode();
/** 控制温度【随机性】0到2之间。较高的值(如0.8)将使输出更加随机,而较低的值(如0.2)将使输出更加集中和确定 */
private double temperature = 0.2;
}

View File

@ -0,0 +1,36 @@
package cn.kazusa.chatgpt.domain.whisper;
import lombok.AllArgsConstructor;
import lombok.Data;
import lombok.Getter;
import java.io.Serializable;
/**
* @author kazusa
* @description 请求信息语音转文字
*/
@Data
public class WhisperEnum implements Serializable {
@Getter
@AllArgsConstructor
public enum Model {
WHISPER_1("whisper-1"),
;
private String code;
}
@Getter
@AllArgsConstructor
public enum ResponseFormat {
JSON("json"),
TEXT("text"),
SRT("srt"),
VERBOSE_JSON("verbose_json"),
VTT("vtt"),
;
private String code;
}
}

View File

@ -0,0 +1,14 @@
package cn.kazusa.chatgpt.domain.whisper;
import lombok.Data;
import java.io.Serializable;
/**
* @author kazusa
* @description 语音转文字响应
*/
@Data
public class WhisperResponse implements Serializable {
private String text;
}

View File

@ -0,0 +1,48 @@
package cn.kazusa.chatgpt.interceptor;
import cn.hutool.http.ContentType;
import cn.hutool.http.Header;
import okhttp3.Interceptor;
import okhttp3.Request;
import okhttp3.Response;
import org.jetbrains.annotations.NotNull;
import java.io.IOException;
/**
* @author kazusa
* @description OpenAI API请求拦截器
*/
public class OpenAiInterceptor implements Interceptor {
/** OpenAi apiKey 需要在官网申请 */
private final String apiKeyBySystem;
/** 访问授权接口的认证 Token */
public OpenAiInterceptor(String apiKeyBySystem) {
this.apiKeyBySystem = apiKeyBySystem;
}
@NotNull
@Override
public Response intercept(Chain chain) throws IOException {
// 1. 获取原始 Request
Request original = chain.request();
// 2. 读取 apiKey优先使用自己传递的 apiKey
String apiKeyByUser = original.header("apiKey");
String apiKey = null == apiKeyByUser ? apiKeyBySystem : apiKeyByUser;
// 3. 构建 Request
Request request = original.newBuilder()
.url(original.url())
.header(Header.AUTHORIZATION.getValue(), "Bearer " + apiKey)
.header(Header.CONTENT_TYPE.getValue(), ContentType.JSON.getValue())
.method(original.method(), original.body())
.build();
// 4. 返回执行结果
return chain.proceed(request);
}
}

View File

@ -0,0 +1,51 @@
package cn.kazusa.chatgpt.session;
import cn.kazusa.chatgpt.IOpenAiApi;
import lombok.AllArgsConstructor;
import lombok.Data;
import lombok.Getter;
import lombok.NoArgsConstructor;
import lombok.Setter;
import lombok.extern.slf4j.Slf4j;
import okhttp3.OkHttpClient;
import okhttp3.sse.EventSource;
import okhttp3.sse.EventSources;
import org.jetbrains.annotations.NotNull;
/**
* @author kazusa
* @description 配置信息
*/
@Slf4j
@Data
@NoArgsConstructor
@AllArgsConstructor
public class Configuration {
@Getter
@Setter
private IOpenAiApi openAiApi;
@Getter
@Setter
private OkHttpClient okHttpClient;
@Getter
@NotNull
private String apiKey;
@Getter
private String apiHost;
/**
* 字段废弃不在使用
*/
@Getter
@Deprecated
private String authToken;
public EventSource.Factory createRequestFactory() {
return EventSources.createFactory(okHttpClient);
}
}

View File

@ -0,0 +1,254 @@
package cn.kazusa.chatgpt.session;
import cn.kazusa.chatgpt.domain.billing.BillingUsage;
import cn.kazusa.chatgpt.domain.billing.Subscription;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionRequest;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionResponse;
import cn.kazusa.chatgpt.domain.edits.EditRequest;
import cn.kazusa.chatgpt.domain.edits.EditResponse;
import cn.kazusa.chatgpt.domain.embedd.EmbeddingRequest;
import cn.kazusa.chatgpt.domain.embedd.EmbeddingResponse;
import cn.kazusa.chatgpt.domain.files.DeleteFileResponse;
import cn.kazusa.chatgpt.domain.files.UploadFileResponse;
import cn.kazusa.chatgpt.domain.images.ImageEditRequest;
import cn.kazusa.chatgpt.domain.images.ImageRequest;
import cn.kazusa.chatgpt.domain.images.ImageResponse;
import cn.kazusa.chatgpt.domain.other.OpenAiResponse;
import cn.kazusa.chatgpt.domain.qa.QACompletionRequest;
import cn.kazusa.chatgpt.domain.qa.QACompletionResponse;
import cn.kazusa.chatgpt.domain.whisper.TranscriptionsRequest;
import cn.kazusa.chatgpt.domain.whisper.TranslationsRequest;
import cn.kazusa.chatgpt.domain.whisper.WhisperResponse;
import com.fasterxml.jackson.core.JsonProcessingException;
import okhttp3.sse.EventSource;
import okhttp3.sse.EventSourceListener;
import org.jetbrains.annotations.NotNull;
import java.io.File;
import java.time.LocalDate;
import java.util.List;
import java.util.concurrent.CompletableFuture;
/**
* @author kazusa
* @description OpenAi 会话接口
*/
public interface OpenAiSession {
/**
* 文本问答简单请求
*
* @param question 请求信息
* @return 应答结果
*/
QACompletionResponse completions(String question);
/**
* 文本问答
*
* @param qaCompletionRequest 请求信息
* @return 应答结果
*/
QACompletionResponse completions(QACompletionRequest qaCompletionRequest);
/**
* 文本问答 & 流式反馈
*
* @param qaCompletionRequest 请求信息
* @param eventSourceListener 实现监听通过监听的 onEvent 方法接收数据
*/
EventSource completions(QACompletionRequest qaCompletionRequest, EventSourceListener eventSourceListener) throws JsonProcessingException;
/**
* 问答模型 GPT-3.5/4.0
*
* @param chatCompletionRequest 请求信息
* @return 应答结果
*/
ChatCompletionResponse completions(ChatCompletionRequest chatCompletionRequest);
/**
* 问答模型 GPT-3.5/4.0 & 流式反馈
*
* @param chatCompletionRequest 请求信息
* @param eventSourceListener 实现监听通过监听的 onEvent 方法接收数据
* @return 应答结果
*/
EventSource chatCompletions(ChatCompletionRequest chatCompletionRequest, EventSourceListener eventSourceListener) throws JsonProcessingException;
/**
* 问答模型 GPT-3.5/4.0 & 流式反馈 & 一次反馈
*
* @param chatCompletionRequest 请求信息
* @return 应答结果
*/
CompletableFuture<String> chatCompletions(ChatCompletionRequest chatCompletionRequest) throws InterruptedException, JsonProcessingException;
/**
* 问答模型 GPT-3.5/4.0 & 流式反馈
*
* @param apiHostByUser 自定义host
* @param apiKeyByUser 自定义Key
* @param chatCompletionRequest 请求信息
* @param eventSourceListener 实现监听通过监听的 onEvent 方法接收数据
* @return 应答结果
*/
EventSource chatCompletions(String apiHostByUser, String apiKeyByUser, ChatCompletionRequest chatCompletionRequest, EventSourceListener eventSourceListener) throws JsonProcessingException;
/**
* 文本修复
*
* @param editRequest 请求信息
* @return 应答结果
*/
EditResponse edit(EditRequest editRequest);
/**
* 生成图片
*
* @param prompt 图片描述
* @return 应答结果
*/
ImageResponse genImages(String prompt);
/**
* 生成图片
*
* @param imageRequest 图片描述
* @return 应答结果
*/
ImageResponse genImages(ImageRequest imageRequest);
/**
* 修改图片
*
* @param image 图片对象
* @param prompt 修改描述
* @return 应答结果
*/
ImageResponse editImages(File image, String prompt);
/**
* 修改图片
*
* @param image 图片对象
* @param imageEditRequest 图片参数
* @return 应答结果
*/
ImageResponse editImages(File image, ImageEditRequest imageEditRequest);
/**
* 修改图片
*
* @param image 图片对象小于4M的PNG图片
* @param mask 图片对象小于4M的PNG图片
* @param imageEditRequest 图片参数
* @return 应答结果
*/
ImageResponse editImages(File image, File mask, ImageEditRequest imageEditRequest);
/**
* 向量计算单个文本
* 文本向量计算是一种在自然语言处理NLP领域中用于测量和比较文本相似性的技术在这种方法中每个单词或短语都被转换为一个向量可以使用这些向量来比较不同文本之间的相似性并在需要时进行推荐或分类
*
* @param input 文本信息
* @return 应答结果
*/
EmbeddingResponse embeddings(String input);
/**
* 向量计算多个文本
* 文本向量计算是一种在自然语言处理NLP领域中用于测量和比较文本相似性的技术在这种方法中每个单词或短语都被转换为一个向量可以使用这些向量来比较不同文本之间的相似性并在需要时进行推荐或分类
*
* @param inputs 多个文本
* @return 应答结果
*/
EmbeddingResponse embeddings(String... inputs);
/**
* 向量计算多个文本
* 文本向量计算是一种在自然语言处理NLP领域中用于测量和比较文本相似性的技术在这种方法中每个单词或短语都被转换为一个向量可以使用这些向量来比较不同文本之间的相似性并在需要时进行推荐或分类
*
* @param inputs 多个文本
* @return 应答结果
*/
EmbeddingResponse embeddings(List<String> inputs);
/**
* 向量计算入参
* 文本向量计算是一种在自然语言处理NLP领域中用于测量和比较文本相似性的技术在这种方法中每个单词或短语都被转换为一个向量可以使用这些向量来比较不同文本之间的相似性并在需要时进行推荐或分类
*
* @param embeddingRequest 请求结果
* @return 应答结果
*/
EmbeddingResponse embeddings(EmbeddingRequest embeddingRequest);
/**
* 获取文件
*
* @return 应答结果
*/
OpenAiResponse<File> files();
/**
* 上传文件
*
* @param file 文件
* @return 应答结果
*/
UploadFileResponse uploadFile(File file);
/**
* 上传文件
*
* @param purpose Use "fine-tune" for Fine-tuning. This allows us to validate the format of the uploaded file.
* @param file 文件
* @return 应答结果
*/
UploadFileResponse uploadFile(String purpose, File file);
/**
* 删除文件
*
* @param fileId 文件ID
* @return 应答结果
*/
DeleteFileResponse deleteFile(String fileId);
/**
* 语音转文字
*
* @param file 语音文件
* @param transcriptionsRequest 请求信息
* @return 应答结果
*/
WhisperResponse speed2TextTranscriptions(File file, TranscriptionsRequest transcriptionsRequest);
/**
* 语音翻译
*
* @param file 语音文件
* @param translationsRequest 请求信息
* @return 应答结果
*/
WhisperResponse speed2TextTranslations(File file, TranslationsRequest translationsRequest);
/**
* 账单查询
*
* @return 应答结果
*/
Subscription subscription();
/**
* 消耗查询
*
* @param starDate 开始时间
* @param endDate 结束时间
* @return 应答数据
*/
BillingUsage billingUsage(@NotNull LocalDate starDate, @NotNull LocalDate endDate);
}

View File

@ -0,0 +1,11 @@
package cn.kazusa.chatgpt.session;
/**
* @author kazusa
* @description OpenAI会话工厂接口
*/
public interface OpenAiSessionFactory {
OpenAiSession openSession();
}

View File

@ -0,0 +1,349 @@
package cn.kazusa.chatgpt.session.defaults;
import cn.hutool.core.util.StrUtil;
import cn.hutool.http.ContentType;
import cn.kazusa.chatgpt.IOpenAiApi;
import cn.kazusa.chatgpt.common.Constants;
import cn.kazusa.chatgpt.domain.billing.BillingUsage;
import cn.kazusa.chatgpt.domain.billing.Subscription;
import cn.kazusa.chatgpt.domain.chat.ChatChoice;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionRequest;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionResponse;
import cn.kazusa.chatgpt.domain.chat.Message;
import cn.kazusa.chatgpt.domain.edits.EditRequest;
import cn.kazusa.chatgpt.domain.edits.EditResponse;
import cn.kazusa.chatgpt.domain.embedd.EmbeddingRequest;
import cn.kazusa.chatgpt.domain.embedd.EmbeddingResponse;
import cn.kazusa.chatgpt.domain.files.DeleteFileResponse;
import cn.kazusa.chatgpt.domain.files.UploadFileResponse;
import cn.kazusa.chatgpt.domain.images.ImageEditRequest;
import cn.kazusa.chatgpt.domain.images.ImageRequest;
import cn.kazusa.chatgpt.domain.images.ImageResponse;
import cn.kazusa.chatgpt.domain.other.OpenAiResponse;
import cn.kazusa.chatgpt.domain.qa.QACompletionRequest;
import cn.kazusa.chatgpt.domain.qa.QACompletionResponse;
import cn.kazusa.chatgpt.domain.whisper.TranscriptionsRequest;
import cn.kazusa.chatgpt.domain.whisper.TranslationsRequest;
import cn.kazusa.chatgpt.domain.whisper.WhisperResponse;
import cn.kazusa.chatgpt.session.Configuration;
import cn.kazusa.chatgpt.session.OpenAiSession;
import com.alibaba.fastjson.JSON;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import io.reactivex.Single;
import okhttp3.MediaType;
import okhttp3.MultipartBody;
import okhttp3.Request;
import okhttp3.RequestBody;
import okhttp3.Response;
import okhttp3.sse.EventSource;
import okhttp3.sse.EventSourceListener;
import org.jetbrains.annotations.NotNull;
import java.io.File;
import java.time.LocalDate;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.CompletableFuture;
/**
* @author kazusa
* @description OpenAI 会话
*/
public class DefaultOpenAiSession implements OpenAiSession {
/**
* 配置信息
*/
private final Configuration configuration;
/**
* OpenAI 接口
*/
private final IOpenAiApi openAiApi;
/**
* 工厂事件
*/
private final EventSource.Factory factory;
public DefaultOpenAiSession(Configuration configuration) {
this.configuration = configuration;
this.openAiApi = configuration.getOpenAiApi();
this.factory = configuration.createRequestFactory();
}
@Override
public QACompletionResponse completions(QACompletionRequest qaCompletionRequest) {
return this.openAiApi.completions(qaCompletionRequest).blockingGet();
}
@Override
public EventSource completions(QACompletionRequest qaCompletionRequest, EventSourceListener eventSourceListener) throws JsonProcessingException {
// 核心参数校验不对用户的传参做更改只返回错误信息
if (!qaCompletionRequest.isStream()) {
throw new RuntimeException("illegal parameter stream is false!");
}
// 构建请求信息
Request request = new Request.Builder()
.url(configuration.getApiHost().concat(IOpenAiApi.v1_completions))
.post(RequestBody.create(MediaType.parse(ContentType.JSON.getValue()), new ObjectMapper().writeValueAsString(qaCompletionRequest)))
.build();
// 返回事件结果
return factory.newEventSource(request, eventSourceListener);
}
@Override
public QACompletionResponse completions(String question) {
QACompletionRequest request = QACompletionRequest
.builder()
.prompt(question)
.build();
Single<QACompletionResponse> completions = this.openAiApi.completions(request);
return completions.blockingGet();
}
@Override
public ChatCompletionResponse completions(ChatCompletionRequest chatCompletionRequest) {
return this.openAiApi.completions(chatCompletionRequest).blockingGet();
}
@Override
public EventSource chatCompletions(ChatCompletionRequest chatCompletionRequest, EventSourceListener eventSourceListener) throws JsonProcessingException {
return chatCompletions(Constants.NULL, Constants.NULL, chatCompletionRequest, eventSourceListener);
}
@Override
public CompletableFuture<String> chatCompletions(ChatCompletionRequest chatCompletionRequest) throws InterruptedException, JsonProcessingException {
// 用于执行异步任务并获取结果
CompletableFuture<String> future = new CompletableFuture<>();
StringBuffer dataBuffer = new StringBuffer();
chatCompletions(chatCompletionRequest, new EventSourceListener(){
@Override
public void onEvent(EventSource eventSource, String id, String type, String data) {
if ("[DONE]".equalsIgnoreCase(data)) {
onClosed(eventSource);
future.complete(dataBuffer.toString());
}
ChatCompletionResponse chatCompletionResponse = JSON.parseObject(data, ChatCompletionResponse.class);
List<ChatChoice> choices = chatCompletionResponse.getChoices();
for (ChatChoice chatChoice : choices) {
Message delta = chatChoice.getDelta();
if (Constants.Role.ASSISTANT.getCode().equals(delta.getRole())) continue;
// 应答完成
String finishReason = chatChoice.getFinishReason();
if ("stop".equalsIgnoreCase(finishReason)) {
onClosed(eventSource);
return;
}
// 发送信息
try {
dataBuffer.append(delta.getContent());
} catch (Exception e) {
future.completeExceptionally(new RuntimeException("Request closed before completion"));
}
}
}
@Override
public void onClosed(EventSource eventSource) {
future.complete(dataBuffer.toString());
}
@Override
public void onFailure(EventSource eventSource, Throwable t, Response response) {
future.completeExceptionally(new RuntimeException("Request closed before completion"));
}
});
return future;
}
@Override
public EventSource chatCompletions(String apiHostByUser, String apiKeyByUser, ChatCompletionRequest chatCompletionRequest, EventSourceListener eventSourceListener) throws JsonProcessingException {
// 核心参数校验不对用户的传参做更改只返回错误信息
if (!chatCompletionRequest.isStream()) {
throw new RuntimeException("illegal parameter stream is false!");
}
// 动态设置 HostKey便于用户传递自己的信息
String apiHost = Constants.NULL.equals(apiHostByUser) ? configuration.getApiHost() : apiHostByUser;
String apiKey = Constants.NULL.equals(apiKeyByUser) ? configuration.getApiKey() : apiKeyByUser;
// 构建请求信息
Request request = new Request.Builder()
// url: https://api.openai.com/v1/chat/completions - 通过 IOpenAiApi 配置的 POST 接口用这样的方式从统一的地方获取配置信息
.url(apiHost.concat(IOpenAiApi.v1_chat_completions))
.addHeader("Authorization", "Bearer " + apiKey)
// 封装请求参数信息如果使用了 Fastjson 也可以替换 ObjectMapper 转换对象
.post(RequestBody.create(MediaType.parse(ContentType.JSON.getValue()), new ObjectMapper().writeValueAsString(chatCompletionRequest)))
.build();
// 返回结果信息EventSource 对象可以取消应答
return factory.newEventSource(request, eventSourceListener);
}
@Override
public EditResponse edit(EditRequest editRequest) {
return this.openAiApi.edits(editRequest).blockingGet();
}
@Override
public ImageResponse genImages(String prompt) {
ImageRequest imageRequest = ImageRequest.builder().prompt(prompt).build();
return this.genImages(imageRequest);
}
@Override
public ImageResponse genImages(ImageRequest imageRequest) {
return this.openAiApi.genImages(imageRequest).blockingGet();
}
@Override
public ImageResponse editImages(File image, String prompt) {
ImageEditRequest imageEditRequest = ImageEditRequest.builder().prompt(prompt).build();
return this.editImages(image, null, imageEditRequest);
}
@Override
public ImageResponse editImages(File image, ImageEditRequest imageEditRequest) {
return this.editImages(image, null, imageEditRequest);
}
@Override
public ImageResponse editImages(File image, File mask, ImageEditRequest imageEditRequest) {
// 1. imageMultipartBody
RequestBody imageBody = RequestBody.create(MediaType.parse("multipart/form-data"), image);
MultipartBody.Part imageMultipartBody = MultipartBody.Part.createFormData("image", image.getName(), imageBody);
// 2. maskMultipartBody
MultipartBody.Part maskMultipartBody = null;
if (Objects.nonNull(mask)) {
RequestBody maskBody = RequestBody.create(MediaType.parse("multipart/form-data"), mask);
maskMultipartBody = MultipartBody.Part.createFormData("mask", mask.getName(), maskBody);
}
// requestBodyMap
Map<String, RequestBody> requestBodyMap = new HashMap<>();
requestBodyMap.put("prompt", RequestBody.create(MediaType.parse("multipart/form-data"), imageEditRequest.getPrompt()));
requestBodyMap.put("n", RequestBody.create(MediaType.parse("multipart/form-data"), imageEditRequest.getN().toString()));
requestBodyMap.put("size", RequestBody.create(MediaType.parse("multipart/form-data"), imageEditRequest.getSize()));
requestBodyMap.put("response_format", RequestBody.create(MediaType.parse("multipart/form-data"), imageEditRequest.getResponseFormat()));
if (!(Objects.isNull(imageEditRequest.getUser()) || "".equals(imageEditRequest.getUser()))) {
requestBodyMap.put("user", RequestBody.create(MediaType.parse("multipart/form-data"), imageEditRequest.getUser()));
}
return this.openAiApi.editImages(imageMultipartBody, maskMultipartBody, requestBodyMap).blockingGet();
}
@Override
public EmbeddingResponse embeddings(String input) {
EmbeddingRequest embeddingRequest = EmbeddingRequest.builder().input(new ArrayList<String>() {{
add(input);
}}).build();
return this.embeddings(embeddingRequest);
}
@Override
public EmbeddingResponse embeddings(String... inputs) {
EmbeddingRequest embeddingRequest = EmbeddingRequest.builder().input(Arrays.asList(inputs)).build();
return this.embeddings(embeddingRequest);
}
@Override
public EmbeddingResponse embeddings(List<String> inputs) {
EmbeddingRequest embeddingRequest = EmbeddingRequest.builder().input(inputs).build();
return this.embeddings(embeddingRequest);
}
@Override
public EmbeddingResponse embeddings(EmbeddingRequest embeddingRequest) {
return this.openAiApi.embeddings(embeddingRequest).blockingGet();
}
@Override
public OpenAiResponse<File> files() {
return this.openAiApi.files().blockingGet();
}
@Override
public UploadFileResponse uploadFile(File file) {
return this.uploadFile("fine-tune", file);
}
@Override
public UploadFileResponse uploadFile(String purpose, File file) {
RequestBody fileBody = RequestBody.create(MediaType.parse("multipart/form-data"), file);
MultipartBody.Part multipartBody = MultipartBody.Part.createFormData("file", file.getName(), fileBody);
RequestBody purposeBody = RequestBody.create(MediaType.parse("multipart/form-data"), purpose);
return this.openAiApi.uploadFile(multipartBody, purposeBody).blockingGet();
}
@Override
public DeleteFileResponse deleteFile(String fileId) {
return this.openAiApi.deleteFile(fileId).blockingGet();
}
@Override
public WhisperResponse speed2TextTranscriptions(File file, TranscriptionsRequest transcriptionsRequest) {
// 1. 语音文件
RequestBody fileBody = RequestBody.create(MediaType.parse("multipart/form-data"), file);
MultipartBody.Part multipartBody = MultipartBody.Part.createFormData("file", file.getName(), fileBody);
// 2. 参数封装
Map<String, RequestBody> requestBodyMap = new HashMap<>();
if (StrUtil.isNotBlank(transcriptionsRequest.getLanguage())) {
requestBodyMap.put(TranscriptionsRequest.Fields.language, RequestBody.create(MediaType.parse("multipart/form-data"), transcriptionsRequest.getLanguage()));
}
if (StrUtil.isNotBlank(transcriptionsRequest.getModel())) {
requestBodyMap.put(TranscriptionsRequest.Fields.model, RequestBody.create(MediaType.parse("multipart/form-data"), transcriptionsRequest.getModel()));
}
if (StrUtil.isNotBlank(transcriptionsRequest.getPrompt())) {
requestBodyMap.put(TranscriptionsRequest.Fields.prompt, RequestBody.create(MediaType.parse("multipart/form-data"), transcriptionsRequest.getPrompt()));
}
if (StrUtil.isNotBlank(transcriptionsRequest.getResponseFormat())) {
requestBodyMap.put(TranscriptionsRequest.Fields.responseFormat, RequestBody.create(MediaType.parse("multipart/form-data"), transcriptionsRequest.getResponseFormat()));
}
requestBodyMap.put(TranscriptionsRequest.Fields.temperature, RequestBody.create(MediaType.parse("multipart/form-data"), String.valueOf(transcriptionsRequest.getTemperature())));
return this.openAiApi.speed2TextTranscriptions(multipartBody, requestBodyMap).blockingGet();
}
@Override
public WhisperResponse speed2TextTranslations(File file, TranslationsRequest translationsRequest) {
// 1. 语音文件
RequestBody fileBody = RequestBody.create(MediaType.parse("multipart/form-data"), file);
MultipartBody.Part multipartBody = MultipartBody.Part.createFormData("file", file.getName(), fileBody);
// 2. 参数封装
Map<String, RequestBody> requestBodyMap = new HashMap<>();
if (StrUtil.isNotBlank(translationsRequest.getModel())) {
requestBodyMap.put(TranslationsRequest.Fields.model, RequestBody.create(MediaType.parse("multipart/form-data"), translationsRequest.getModel()));
}
if (StrUtil.isNotBlank(translationsRequest.getPrompt())) {
requestBodyMap.put(TranslationsRequest.Fields.prompt, RequestBody.create(MediaType.parse("multipart/form-data"), translationsRequest.getPrompt()));
}
if (StrUtil.isNotBlank(translationsRequest.getResponseFormat())) {
requestBodyMap.put(TranslationsRequest.Fields.responseFormat, RequestBody.create(MediaType.parse("multipart/form-data"), translationsRequest.getResponseFormat()));
}
requestBodyMap.put(TranslationsRequest.Fields.temperature, RequestBody.create(MediaType.parse("multipart/form-data"), String.valueOf(translationsRequest.getTemperature())));
requestBodyMap.put(TranscriptionsRequest.Fields.temperature, RequestBody.create(MediaType.parse("multipart/form-data"), String.valueOf(translationsRequest.getTemperature())));
return this.openAiApi.speed2TextTranscriptions(multipartBody, requestBodyMap).blockingGet();
}
@Override
public Subscription subscription() {
return this.openAiApi.subscription().blockingGet();
}
@Override
public BillingUsage billingUsage(@NotNull LocalDate starDate, @NotNull LocalDate endDate) {
return this.openAiApi.billingUsage(starDate, endDate).blockingGet();
}
}

View File

@ -0,0 +1,57 @@
package cn.kazusa.chatgpt.session.defaults;
import cn.kazusa.chatgpt.IOpenAiApi;
import cn.kazusa.chatgpt.interceptor.OpenAiInterceptor;
import cn.kazusa.chatgpt.session.Configuration;
import cn.kazusa.chatgpt.session.OpenAiSession;
import cn.kazusa.chatgpt.session.OpenAiSessionFactory;
import okhttp3.OkHttpClient;
import okhttp3.logging.HttpLoggingInterceptor;
import retrofit2.Retrofit;
import retrofit2.adapter.rxjava2.RxJava2CallAdapterFactory;
import retrofit2.converter.jackson.JacksonConverterFactory;
import java.util.concurrent.TimeUnit;
/**
* @author kazusa
* @description OpenAi API Factory 会话工厂
*/
public class DefaultOpenAiSessionFactory implements OpenAiSessionFactory {
private final Configuration configuration;
public DefaultOpenAiSessionFactory(Configuration configuration) {
this.configuration = configuration;
}
@Override
public OpenAiSession openSession() {
// 1. 日志配置
HttpLoggingInterceptor httpLoggingInterceptor = new HttpLoggingInterceptor();
httpLoggingInterceptor.setLevel(HttpLoggingInterceptor.Level.BODY);
// 2. 开启 Http 客户端
OkHttpClient okHttpClient = new OkHttpClient
.Builder()
.addInterceptor(httpLoggingInterceptor)
.addInterceptor(new OpenAiInterceptor(configuration.getApiKey()))
.connectTimeout(450, TimeUnit.SECONDS)
.writeTimeout(450, TimeUnit.SECONDS)
.readTimeout(450, TimeUnit.SECONDS)
.build();
configuration.setOkHttpClient(okHttpClient);
// 3. 创建 API 服务
IOpenAiApi openAiApi = new Retrofit.Builder()
.baseUrl(configuration.getApiHost())
.client(okHttpClient)
.addCallAdapterFactory(RxJava2CallAdapterFactory.create())
.addConverterFactory(JacksonConverterFactory.create())
.build().create(IOpenAiApi.class);
configuration.setOpenAiApi(openAiApi);
return new DefaultOpenAiSession(configuration);
}
}

View File

@ -0,0 +1,305 @@
package cn.kazusa.chatgpt.test;
import cn.kazusa.chatgpt.common.Constants;
import cn.kazusa.chatgpt.domain.billing.BillingUsage;
import cn.kazusa.chatgpt.domain.billing.Subscription;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionRequest;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionResponse;
import cn.kazusa.chatgpt.domain.chat.Message;
import cn.kazusa.chatgpt.domain.edits.EditRequest;
import cn.kazusa.chatgpt.domain.edits.EditResponse;
import cn.kazusa.chatgpt.domain.embedd.EmbeddingResponse;
import cn.kazusa.chatgpt.domain.files.DeleteFileResponse;
import cn.kazusa.chatgpt.domain.files.UploadFileResponse;
import cn.kazusa.chatgpt.domain.images.ImageResponse;
import cn.kazusa.chatgpt.domain.other.OpenAiResponse;
import cn.kazusa.chatgpt.domain.qa.QACompletionRequest;
import cn.kazusa.chatgpt.domain.qa.QACompletionResponse;
import cn.kazusa.chatgpt.session.Configuration;
import cn.kazusa.chatgpt.session.OpenAiSession;
import cn.kazusa.chatgpt.session.OpenAiSessionFactory;
import cn.kazusa.chatgpt.session.defaults.DefaultOpenAiSessionFactory;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import lombok.extern.slf4j.Slf4j;
import okhttp3.Response;
import okhttp3.sse.EventSource;
import okhttp3.sse.EventSourceListener;
import org.junit.Before;
import org.junit.Test;
import java.io.File;
import java.io.IOException;
import java.time.LocalDate;
import java.util.ArrayList;
import java.util.Collections;
import java.util.concurrent.CountDownLatch;
/**
* @author kazusa
* @description OpenAI API单元测试类
*/
@Slf4j
public class ApiTest {
private OpenAiSession openAiSession;
@Before
public void test_OpenAiSessionFactory() {
// 1. 配置文件 [联系小傅哥获取key]
// 1.1 官网原始 apiHost https://api.openai.com/ - 官网的Key可直接使用
// 1.2 三方公司 apiHost https://pro-share-aws-api.zcyai.com/ - 需要找我获得 Key 支持3.5\4.0流式问答模型调用有些模型已废弃不对接使用
String gptSdkHost = System.getProperty("GPT_SDK_HOST");
String gptSdkKey = System.getProperty("GPT_SDK_KEY");
Configuration configuration = new Configuration();
configuration.setApiHost(gptSdkHost);
configuration.setApiKey(gptSdkKey);
// 2. 会话工厂
OpenAiSessionFactory factory = new DefaultOpenAiSessionFactory(configuration);
// 3. 开启会话
this.openAiSession = factory.openSession();
}
/**
* 常用对话模式推荐使用此模型进行测试
* 此对话模型 3.5/4.0 接近于官网体验 & 流式应答
*/
@Test
public void test_chat_completions_stream_channel() throws JsonProcessingException, InterruptedException {
// 1. 创建参数
ChatCompletionRequest chatCompletion = ChatCompletionRequest
.builder()
.stream(true)
.messages(Collections.singletonList(Message.builder().role(Constants.Role.USER).content("1+1").build()))
.model(ChatCompletionRequest.Model.GPT_3_5_TURBO.getCode())
.maxTokens(1024)
.build();
// 2. 用户配置 可选参数支持不同渠道的 apiHostapiKey- 方便给每个用户都分配了自己的key用于售卖场景
String apiHost = "";
String apiKey = "";
// 3. 发起请求
EventSource eventSource = openAiSession.chatCompletions(apiHost, apiKey, chatCompletion, new EventSourceListener() {
@Override
public void onEvent(EventSource eventSource, String id, String type, String data) {
log.info("测试结果 id:{} type:{} data:{}", id, type, data);
}
@Override
public void onFailure(EventSource eventSource, Throwable t, Response response) {
log.error("失败 throwable:{}", t != null ? t.getMessage() : "Unknown error", t);
if (response != null) {
log.error("响应 code:{} message:{}", response.code(), response.message());
}
}
});
// 等待
new CountDownLatch(1).await();
}
/**
* 常用对话模式推荐使用此模型进行测试
* 此对话模型 3.5/4.0 接近于官网体验 & 流式应答
*/
@Test
public void test_chat_completions_stream() throws JsonProcessingException, InterruptedException {
// 1. 创建参数
ChatCompletionRequest chatCompletion = ChatCompletionRequest
.builder()
.stream(true)
.messages(Collections.singletonList(Message.builder().role(Constants.Role.USER).content("1+1").build()))
.model(ChatCompletionRequest.Model.GPT_3_5_TURBO.getCode())
.maxTokens(1024)
.build();
// 2. 发起请求
EventSource eventSource = openAiSession.chatCompletions(chatCompletion, new EventSourceListener() {
@Override
public void onEvent(EventSource eventSource, String id, String type, String data) {
log.info("测试结果 id:{} type:{} data:{}", id, type, data);
}
@Override
public void onFailure(EventSource eventSource, Throwable t, Response response) {
log.error("失败 throwable:{}", t != null ? t.getMessage() : "Unknown error", t);
if (response != null) {
log.error("响应 code:{} message:{}", response.code(), response.message());
}
}
});
// 等待
new CountDownLatch(1).await();
}
/**
* 简单问答模式
*/
@Test
public void test_qa_completions() throws JsonProcessingException {
QACompletionResponse response01 = openAiSession.completions("写个java冒泡排序");
log.info("测试结果:{}", new ObjectMapper().writeValueAsString(response01.getChoices()));
}
/**
* 简单问答模式 * 流式应答
*/
@Test
public void test_qa_completions_stream() throws JsonProcessingException, InterruptedException {
// 1. 创建参数
QACompletionRequest request = QACompletionRequest
.builder()
.prompt("写个java冒泡排序")
.stream(true)
.build();
for (int i = 0; i < 1; i++) {
// 2. 发起请求
EventSource eventSource = openAiSession.completions(request, new EventSourceListener() {
@Override
public void onEvent(EventSource eventSource, String id, String type, String data) {
log.info("测试结果:{}", data);
}
});
}
// 等待
new CountDownLatch(1).await();
}
/**
* 此对话模型 3.5 接近于官网体验
*/
@Test
public void test_chat_completions() {
// 1. 创建参数
ChatCompletionRequest chatCompletion = ChatCompletionRequest
.builder()
.messages(Collections.singletonList(Message.builder().role(Constants.Role.USER).content("写一个java冒泡排序").build()))
.model(ChatCompletionRequest.Model.GPT_3_5_TURBO.getCode())
.build();
// 2. 发起请求
ChatCompletionResponse chatCompletionResponse = openAiSession.completions(chatCompletion);
// 3. 解析结果
chatCompletionResponse.getChoices().forEach(e -> {
log.info("测试结果:{}", e.getMessage());
});
}
/**
* 上下文对话
*/
@Test
public void test_chat_completions_context() {
// 1-1. 创建参数
ChatCompletionRequest chatCompletion = ChatCompletionRequest
.builder()
.messages(new ArrayList<>())
.model(ChatCompletionRequest.Model.GPT_3_5_TURBO.getCode())
.user("testUser01")
.build();
// 写入请求信息
chatCompletion.getMessages().add(Message.builder().role(Constants.Role.USER).content("写一个java冒泡排序").build());
// 1-2. 发起请求
ChatCompletionResponse chatCompletionResponse01 = openAiSession.completions(chatCompletion);
log.info("测试结果:{}", chatCompletionResponse01.getChoices());
// 写入请求信息
chatCompletion.getMessages().add(Message.builder().role(Constants.Role.USER).content(chatCompletionResponse01.getChoices().get(0).getMessage().getContent()).build());
chatCompletion.getMessages().add(Message.builder().role(Constants.Role.USER).content("换一种写法").build());
ChatCompletionResponse chatCompletionResponse02 = openAiSession.completions(chatCompletion);
log.info("测试结果:{}", chatCompletionResponse02.getChoices());
}
/**
* 文本修复
*/
@Test
public void test_edit() {
// 文本请求
EditRequest textRequest = EditRequest.builder()
.input("码农会锁")
.instruction("帮我修改错字")
.model(EditRequest.Model.TEXT_DAVINCI_EDIT_001.getCode()).build();
EditResponse textResponse = openAiSession.edit(textRequest);
log.info("测试结果:{}", textResponse);
// 代码请求
EditRequest codeRequest = EditRequest.builder()
// j <= 10 应该修改为 i <= 10
.input("for (int i = 1; j <= 10; i++) {\n" +
" System.out.println(i);\n" +
"}")
.instruction("这段代码执行时报错,请帮我修改").model(EditRequest.Model.CODE_DAVINCI_EDIT_001.getCode()).build();
EditResponse codeResponse = openAiSession.edit(codeRequest);
log.info("测试结果:{}", codeResponse);
}
/**
* 生成图片
*/
@Test
public void test_genImages() {
// 方式1简单调用
ImageResponse imageResponse01 = openAiSession.genImages("画一个996加班的程序员");
log.info("测试结果:{}", imageResponse01);
// // 方式2调参调用
// ImageResponse imageResponse02 = openAiSession.genImages(ImageRequest.builder()
// .prompt("画一个996加班的程序员")
// .size(ImageEnum.Size.size_256.getCode())
// .responseFormat(ImageEnum.ResponseFormat.B64_JSON.getCode()).build());
// log.info("测试结果:{}", imageResponse02);
}
/**
* 修改图片有3个方法入参不同
*/
@Test
public void test_editImages() throws IOException {
ImageResponse imageResponse = openAiSession.editImages(new File("/Users/fuzhengwei/1024/KnowledgePlanet/chatgpt/chatgpt-sdk-java/docs/images/996.png"), "去除图片中的文字");
log.info("测试结果:{}", imageResponse);
}
@Test
public void test_embeddings() {
EmbeddingResponse embeddingResponse = openAiSession.embeddings("哈喽", "", "hi!");
log.info("测试结果:{}", embeddingResponse);
}
@Test
public void test_files() {
OpenAiResponse<File> openAiResponse = openAiSession.files();
log.info("测试结果:{}", openAiResponse);
}
@Test
public void test_uploadFile() {
UploadFileResponse uploadFileResponse = openAiSession.uploadFile(new File("/Users/fuzhengwei/1024/KnowledgePlanet/chatgpt/chatgpt-sdk-java/docs/files/introduce.md"));
log.info("测试结果:{}", uploadFileResponse);
}
@Test
public void test_deleteFile() {
DeleteFileResponse deleteFileResponse = openAiSession.deleteFile("file id 上传后才能获得");
log.info("测试结果:{}", deleteFileResponse);
}
@Test
public void test_subscription() {
Subscription subscription = openAiSession.subscription();
log.info("测试结果:{}", subscription);
}
@Test
public void test_billingUsage() {
BillingUsage billingUsage = openAiSession.billingUsage(LocalDate.of(2023, 3, 20), LocalDate.now());
log.info("测试结果:{}", billingUsage.getTotalUsage());
}
}

View File

@ -0,0 +1,140 @@
package cn.kazusa.chatgpt.test;
import cn.hutool.http.ContentType;
import cn.hutool.http.Header;
import cn.kazusa.chatgpt.IOpenAiApi;
import cn.kazusa.chatgpt.common.Constants;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionRequest;
import cn.kazusa.chatgpt.domain.chat.ChatCompletionResponse;
import cn.kazusa.chatgpt.domain.chat.Message;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import io.reactivex.Single;
import lombok.extern.slf4j.Slf4j;
import okhttp3.HttpUrl;
import okhttp3.MediaType;
import okhttp3.OkHttpClient;
import okhttp3.Request;
import okhttp3.RequestBody;
import okhttp3.logging.HttpLoggingInterceptor;
import okhttp3.sse.EventSource;
import okhttp3.sse.EventSourceListener;
import okhttp3.sse.EventSources;
import org.junit.Test;
import retrofit2.Retrofit;
import retrofit2.adapter.rxjava2.RxJava2CallAdapterFactory;
import retrofit2.converter.jackson.JacksonConverterFactory;
import java.util.Collections;
import java.util.concurrent.CountDownLatch;
/**
* @author kazusa
* @description HTTP客户端测试类
*/
@Slf4j
public class HttpClientTest {
@Test
public void test_client() {
HttpLoggingInterceptor httpLoggingInterceptor = new HttpLoggingInterceptor();
httpLoggingInterceptor.setLevel(HttpLoggingInterceptor.Level.BODY);
OkHttpClient okHttpClient = new OkHttpClient
.Builder()
.addInterceptor(httpLoggingInterceptor)
.addInterceptor(chain -> {
Request original = chain.request();
// 从请求中获取 token 参数并将其添加到请求路径中
HttpUrl url = original.url().newBuilder()
.addQueryParameter("token", "eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiJ4ZmciLCJleHAiOjE2ODM5NDU3NDcsImlhdCI6MTY4Mzk0MjE0NywianRpIjoiM2QyMDExMTYtNmVjMS00Y2UzLWJhYzgtYzYxYmVmN2ZmNWE5IiwidXNlcm5hbWUiOiJ4ZmcifQ.3FDvUNuNoGemKLhcgagy8WH7xHwRU37t--BuH0N9skg")
.build();
Request request = original.newBuilder()
.url(url)
.header(Header.AUTHORIZATION.getValue(), "Bearer " + "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")
.header(Header.CONTENT_TYPE.getValue(), ContentType.JSON.getValue())
.method(original.method(), original.body())
.build();
return chain.proceed(request);
})
.build();
IOpenAiApi openAiApi = new Retrofit.Builder()
.baseUrl("https://api.xfg.im/b8b6/")
.client(okHttpClient)
.addCallAdapterFactory(RxJava2CallAdapterFactory.create())
.addConverterFactory(JacksonConverterFactory.create())
.build().create(IOpenAiApi.class);
Message message = Message.builder().role(Constants.Role.USER).content("写一个java冒泡排序").build();
ChatCompletionRequest chatCompletion = ChatCompletionRequest
.builder()
.messages(Collections.singletonList(message))
.model(ChatCompletionRequest.Model.GPT_3_5_TURBO.getCode())
.build();
Single<ChatCompletionResponse> chatCompletionResponseSingle = openAiApi.completions(chatCompletion);
ChatCompletionResponse chatCompletionResponse = chatCompletionResponseSingle.blockingGet();
chatCompletionResponse.getChoices().forEach(e -> {
System.out.println(e.getMessage());
});
}
@Test
public void test_client_stream() throws JsonProcessingException, InterruptedException {
HttpLoggingInterceptor httpLoggingInterceptor = new HttpLoggingInterceptor();
httpLoggingInterceptor.setLevel(HttpLoggingInterceptor.Level.BODY);
OkHttpClient okHttpClient = new OkHttpClient
.Builder()
.addInterceptor(httpLoggingInterceptor)
.addInterceptor(chain -> {
Request original = chain.request();
// 从请求中获取 token 参数并将其添加到请求路径中
HttpUrl url = original.url().newBuilder()
.addQueryParameter("token", "eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiJ4ZmciLCJleHAiOjE2ODM5NDU3NDcsImlhdCI6MTY4Mzk0MjE0NywianRpIjoiM2QyMDExMTYtNmVjMS00Y2UzLWJhYzgtYzYxYmVmN2ZmNWE5IiwidXNlcm5hbWUiOiJ4ZmcifQ.3FDvUNuNoGemKLhcgagy8WH7xHwRU37t--BuH0N9skg")
.build();
Request request = original.newBuilder()
.url(url)
.header(Header.AUTHORIZATION.getValue(), "Bearer " + "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")
.header(Header.CONTENT_TYPE.getValue(), ContentType.JSON.getValue())
.method(original.method(), original.body())
.build();
return chain.proceed(request);
})
.build();
Message message = Message.builder().role(Constants.Role.USER).content("写一个java冒泡排序").build();
ChatCompletionRequest chatCompletion = ChatCompletionRequest
.builder()
.messages(Collections.singletonList(message))
.model(ChatCompletionRequest.Model.GPT_3_5_TURBO.getCode())
.stream(true)
.build();
EventSource.Factory factory = EventSources.createFactory(okHttpClient);
String requestBody = new ObjectMapper().writeValueAsString(chatCompletion);
Request request = new Request.Builder()
.url("https://api.xfg.im/b8b6/v1/chat/completions")
.post(RequestBody.create(MediaType.parse(ContentType.JSON.getValue()), requestBody))
.build();
EventSource eventSource = factory.newEventSource(request, new EventSourceListener() {
@Override
public void onEvent(EventSource eventSource, String id, String type, String data) {
log.info("测试结果:{}", data);
}
});
// 等待
new CountDownLatch(1).await();
}
}