Skip to content

Commit f18ddf2

Browse files
committed
test: add JSON tests for Generative.CustomTypeAdapterFactory
1 parent aee94ba commit f18ddf2

8 files changed

Lines changed: 299 additions & 59 deletions

File tree

src/main/java/io/weaviate/client6/v1/api/collections/Generative.java

Lines changed: 21 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -6,11 +6,11 @@
66
import java.util.function.Function;
77

88
import com.google.gson.Gson;
9+
import com.google.gson.JsonParser;
910
import com.google.gson.TypeAdapter;
1011
import com.google.gson.TypeAdapterFactory;
1112
import com.google.gson.reflect.TypeToken;
1213
import com.google.gson.stream.JsonReader;
13-
import com.google.gson.stream.JsonToken;
1414
import com.google.gson.stream.JsonWriter;
1515

1616
import io.weaviate.client6.v1.api.collections.generative.AnthropicGenerative;
@@ -161,7 +161,7 @@ public static Generative frienliai() {
161161
*
162162
* @param fn Lambda expression for optional parameters.
163163
*/
164-
public static Generative frienliai(Function<FriendliaiGenerative.Builder, ObjectBuilder<FriendliaiGenerative>> fn) {
164+
public static Generative friendliai(Function<FriendliaiGenerative.Builder, ObjectBuilder<FriendliaiGenerative>> fn) {
165165
return FriendliaiGenerative.of(fn);
166166
}
167167

@@ -508,35 +508,39 @@ public <T> TypeAdapter<T> create(Gson gson, TypeToken<T> type) {
508508
init(gson);
509509
}
510510

511-
final TypeAdapter<Generative> writeAdapter = (TypeAdapter<Generative>) gson.getDelegateAdapter(this,
511+
final TypeAdapter<T> writeAdapter = (TypeAdapter<T>) gson.getDelegateAdapter(this,
512512
TypeToken.get(rawType));
513513
return (TypeAdapter<T>) new TypeAdapter<Generative>() {
514514

515515
@Override
516516
public void write(JsonWriter out, Generative value) throws IOException {
517517
out.beginObject();
518518
out.name(value._kind().jsonValue());
519-
writeAdapter.write(out, value._self());
519+
writeAdapter.write(out, (T) value._self());
520520
out.endObject();
521521
}
522522

523523
@Override
524524
public Generative read(JsonReader in) throws IOException {
525-
in.beginObject();
526-
var moduleName = in.nextName();
527-
try {
528-
var kind = Generative.Kind.valueOfJson(moduleName);
529-
var adapter = readAdapters.get(kind);
530-
assert adapter != null : "no generative adapter for kind " + kind;
531-
return adapter.read(in);
532-
} catch (IllegalArgumentException e) {
533-
return null;
534-
} finally {
535-
if (in.peek() == JsonToken.BEGIN_OBJECT) {
536-
in.beginObject();
525+
var jsonObject = JsonParser.parseReader(in).getAsJsonObject();
526+
var provider = jsonObject.keySet().iterator().next();
527+
528+
var generative = jsonObject.get(provider).getAsJsonObject();
529+
Generative.Kind kind;
530+
if (provider.equals(Generative.Kind.OPENAI.jsonValue())) {
531+
kind = generative.has("deploymentId") && generative.has("resourceName")
532+
? Generative.Kind.AZURE_OPENAI
533+
: Generative.Kind.OPENAI;
534+
} else {
535+
try {
536+
kind = Generative.Kind.valueOfJson(provider);
537+
} catch (IllegalArgumentException e) {
538+
return null;
537539
}
538-
in.endObject();
539540
}
541+
var adapter = readAdapters.get(kind);
542+
assert adapter != null : "no generative adapter for kind " + kind;
543+
return adapter.fromJsonTree(generative);
540544
}
541545
}.nullSafe();
542546
}

src/main/java/io/weaviate/client6/v1/api/collections/generative/AnthropicGenerative.java

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,6 @@
99

1010
import io.weaviate.client6.v1.api.collections.Generative;
1111
import io.weaviate.client6.v1.internal.ObjectBuilder;
12-
import io.weaviate.client6.v1.internal.TaggedUnion;
1312

1413
public record AnthropicGenerative(
1514
@SerializedName("model") String model,

src/main/java/io/weaviate/client6/v1/api/collections/generative/AwsGenerative.java

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010
public record AwsGenerative(
1111
@SerializedName("region") String region,
1212
@SerializedName("service") String service,
13-
@SerializedName("endpoint") String baseURL,
13+
@SerializedName("endpoint") String baseUrl,
1414
@SerializedName("model") String model) implements Generative {
1515

1616
@Override

src/main/java/io/weaviate/client6/v1/api/collections/generative/CohereGenerative.java

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@
1212

1313
public record CohereGenerative(
1414
@SerializedName("baseURL") String baseUrl,
15-
@SerializedName("kProperty") Integer k,
15+
@SerializedName("kProperty") Integer topK,
1616
@SerializedName("model") String model,
1717
@SerializedName("maxTokensProperty") Integer maxTokens,
1818
@SerializedName("temperatureProperty") Float temperature,
@@ -40,7 +40,7 @@ public static CohereGenerative of(Function<Builder, ObjectBuilder<CohereGenerati
4040
public CohereGenerative(Builder builder) {
4141
this(
4242
builder.baseUrl,
43-
builder.k,
43+
builder.topK,
4444
builder.model,
4545
builder.maxTokens,
4646
builder.temperature,
@@ -50,7 +50,7 @@ public CohereGenerative(Builder builder) {
5050

5151
public static class Builder implements ObjectBuilder<CohereGenerative> {
5252
private String baseUrl;
53-
private Integer k;
53+
private Integer topK;
5454
private String model;
5555
private Integer maxTokens;
5656
private Float temperature;
@@ -63,8 +63,8 @@ public Builder baseUrl(String baseUrl) {
6363
return this;
6464
}
6565

66-
public Builder k(int k) {
67-
this.k = k;
66+
public Builder topK(int topK) {
67+
this.topK = topK;
6868
return this;
6969
}
7070

src/main/java/io/weaviate/client6/v1/api/collections/generative/DatabricksGenerative.java

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -34,23 +34,23 @@ public static DatabricksGenerative of(String baseURL, Function<Builder, ObjectBu
3434

3535
public DatabricksGenerative(Builder builder) {
3636
this(
37-
builder.baseURL,
37+
builder.baseUrl,
3838
builder.maxTokens,
3939
builder.topK,
4040
builder.topP,
4141
builder.temperature);
4242
}
4343

4444
public static class Builder implements ObjectBuilder<DatabricksGenerative> {
45-
private final String baseURL;
45+
private final String baseUrl;
4646

4747
private Integer maxTokens;
4848
private Integer topK;
4949
private Float topP;
5050
private Float temperature;
5151

52-
public Builder(String baseURL) {
53-
this.baseURL = baseURL;
52+
public Builder(String baseUrl) {
53+
this.baseUrl = baseUrl;
5454
}
5555

5656
/** Limit the number of tokens to generate in the response. */

src/main/java/io/weaviate/client6/v1/api/collections/generative/DummyGenerative.java

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
package io.weaviate.client6.v1.api.collections.generative;
22

33
import io.weaviate.client6.v1.api.collections.Generative;
4-
import io.weaviate.client6.v1.api.collections.generate.ProviderMetadata;
54

65
public record DummyGenerative() implements Generative {
76
@Override

src/main/java/io/weaviate/client6/v1/api/collections/generative/OllamaGenerative.java

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
import io.weaviate.client6.v1.internal.ObjectBuilder;
99

1010
public record OllamaGenerative(
11-
@SerializedName("apiEndpoint") String apiEndpoint,
11+
@SerializedName("apiEndpoint") String baseUrl,
1212
@SerializedName("model") String model) implements Generative {
1313

1414
@Override
@@ -31,17 +31,17 @@ public static OllamaGenerative of(Function<Builder, ObjectBuilder<OllamaGenerati
3131

3232
public OllamaGenerative(Builder builder) {
3333
this(
34-
builder.apiEndpoint,
34+
builder.baseUrl,
3535
builder.model);
3636
}
3737

3838
public static class Builder implements ObjectBuilder<OllamaGenerative> {
39-
private String apiEndpoint;
39+
private String baseUrl;
4040
private String model;
4141

42-
/** Destination endpoint of the generative provider. */
43-
public Builder apiEndpoint(String apiEndpoint) {
44-
this.apiEndpoint = apiEndpoint;
42+
/** Base URL of the generative model. */
43+
public Builder baseUrl(String baseUrl) {
44+
this.baseUrl = baseUrl;
4545
return this;
4646
}
4747

0 commit comments

Comments
 (0)