Skip to content

Commit 2a6d129

Browse files
committed
refactor: provide separate builders for VertexAI / AiStudio generative modules
1 parent e13d3ad commit 2a6d129

4 files changed

Lines changed: 113 additions & 29 deletions

File tree

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

Lines changed: 20 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -186,8 +186,8 @@ public static Generative friendliai(Function<FriendliaiGenerative.Builder, Objec
186186
}
187187

188188
/** Configure a default {@code generative-palm} module. */
189-
public static Generative google(String projectId) {
190-
return GoogleGenerative.of(projectId);
189+
public static Generative googleVertex(String projectId) {
190+
return GoogleGenerative.vertex(projectId);
191191
}
192192

193193
/**
@@ -196,9 +196,24 @@ public static Generative google(String projectId) {
196196
* @param projectId Project ID.
197197
* @param fn Lambda expression for optional parameters.
198198
*/
199-
public static Generative google(String projectId,
200-
Function<GoogleGenerative.Builder, ObjectBuilder<GoogleGenerative>> fn) {
201-
return GoogleGenerative.of(projectId, fn);
199+
public static Generative googleVertex(String projectId,
200+
Function<GoogleGenerative.VertexBuilder, ObjectBuilder<GoogleGenerative>> fn) {
201+
return GoogleGenerative.vertex(projectId, fn);
202+
}
203+
204+
/** Configure a default {@code generative-palm} module. */
205+
public static Generative googleAiStudio() {
206+
return GoogleGenerative.aiStudio();
207+
}
208+
209+
/**
210+
* Configure a {@code generative-palm} module.
211+
*
212+
* @param fn Lambda expression for optional parameters.
213+
*/
214+
public static Generative googleAiStudio(
215+
Function<GoogleGenerative.AiStudioBuilder, ObjectBuilder<GoogleGenerative>> fn) {
216+
return GoogleGenerative.aiStudio(fn);
202217
}
203218

204219
/** Configure a default {@code generative-mistral} module. */

src/main/java/io/weaviate/client6/v1/api/collections/generate/DynamicProvider.java

Lines changed: 15 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -104,9 +104,21 @@ public static DynamicProvider friendliai(
104104
*
105105
* @param fn Lambda expression for optional parameters.
106106
*/
107-
public static DynamicProvider google(
108-
Function<GoogleGenerative.Provider.Builder, ObjectBuilder<GoogleGenerative.Provider>> fn) {
109-
return GoogleGenerative.Provider.of(fn);
107+
public static DynamicProvider googleAiStudio(
108+
Function<GoogleGenerative.Provider.AiStudioBuilder, ObjectBuilder<GoogleGenerative.Provider>> fn) {
109+
return GoogleGenerative.Provider.aiStudio(fn);
110+
}
111+
112+
/**
113+
* Configure {@code generative-palm} as a dynamic provider.
114+
*
115+
* @param projectId Google project ID.
116+
* @param fn Lambda expression for optional parameters.
117+
*/
118+
public static DynamicProvider googleVertex(
119+
String projectId,
120+
Function<GoogleGenerative.Provider.VertexBuilder, ObjectBuilder<GoogleGenerative.Provider>> fn) {
121+
return GoogleGenerative.Provider.vertex(projectId, fn);
110122
}
111123

112124
/**

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

Lines changed: 72 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99

1010
import io.weaviate.client6.v1.api.collections.Generative;
1111
import io.weaviate.client6.v1.api.collections.generate.DynamicProvider;
12+
import io.weaviate.client6.v1.api.collections.vectorizers.Text2VecGoogleVectorizer;
1213
import io.weaviate.client6.v1.internal.ObjectBuilder;
1314
import io.weaviate.client6.v1.internal.grpc.protocol.WeaviateProtoBase;
1415
import io.weaviate.client6.v1.internal.grpc.protocol.WeaviateProtoGenerative;
@@ -32,12 +33,20 @@ public Object _self() {
3233
return this;
3334
}
3435

35-
public static GoogleGenerative of(String projectId) {
36-
return of(projectId, ObjectBuilder.identity());
36+
public static GoogleGenerative aiStudio() {
37+
return aiStudio(ObjectBuilder.identity());
3738
}
3839

39-
public static GoogleGenerative of(String projectId, Function<Builder, ObjectBuilder<GoogleGenerative>> fn) {
40-
return fn.apply(new Builder(projectId)).build();
40+
public static GoogleGenerative aiStudio(Function<AiStudioBuilder, ObjectBuilder<GoogleGenerative>> fn) {
41+
return fn.apply(new AiStudioBuilder()).build();
42+
}
43+
44+
public static GoogleGenerative vertex(String projectId) {
45+
return vertex(projectId, ObjectBuilder.identity());
46+
}
47+
48+
public static GoogleGenerative vertex(String projectId, Function<VertexBuilder, ObjectBuilder<GoogleGenerative>> fn) {
49+
return fn.apply(new VertexBuilder(projectId)).build();
4150
}
4251

4352
public GoogleGenerative(Builder builder) {
@@ -51,22 +60,23 @@ public GoogleGenerative(Builder builder) {
5160
builder.temperature);
5261
}
5362

54-
public static class Builder implements ObjectBuilder<GoogleGenerative> {
63+
public abstract static class Builder implements ObjectBuilder<GoogleGenerative> {
64+
private String baseUrl;
5565
private final String projectId;
5666

57-
private String baseUrl;
5867
private String model;
5968
private Integer maxTokens;
6069
private Integer topK;
6170
private Float topP;
6271
private Float temperature;
6372

64-
public Builder(String projectId) {
73+
public Builder(String baseUrl, String projectId) {
6574
this.projectId = projectId;
75+
this.baseUrl = baseUrl;
6676
}
6777

6878
/** Base URL of the generative provider. */
69-
public Builder baseUrl(String baseUrl) {
79+
protected Builder baseUrl(String baseUrl) {
7080
this.baseUrl = baseUrl;
7181
return this;
7282
}
@@ -110,6 +120,24 @@ public GoogleGenerative build() {
110120
}
111121
}
112122

123+
public static class AiStudioBuilder extends Builder {
124+
public AiStudioBuilder() {
125+
super(Text2VecGoogleVectorizer.AiStudioBuilder.BASE_URL, null);
126+
}
127+
}
128+
129+
public static class VertexBuilder extends Builder {
130+
public VertexBuilder(String projectId) {
131+
super(Text2VecGoogleVectorizer.VertexBuilder.DEFAULT_BASE_URL, projectId);
132+
}
133+
134+
/** Base URL of the generative provider. */
135+
public VertexBuilder baseUrl(String baseUrl) {
136+
super.baseUrl(baseUrl);
137+
return this;
138+
}
139+
}
140+
113141
public static record Metadata(TokenMetadata tokens, Usage usage) implements ProviderMetadata {
114142

115143
public static record TokenCount(Long totalBillableCharacters, Long totalTokens) {
@@ -138,9 +166,15 @@ public static record Provider(
138166
List<String> images,
139167
List<String> imageProperties) implements DynamicProvider {
140168

141-
public static Provider of(
142-
Function<GoogleGenerative.Provider.Builder, ObjectBuilder<GoogleGenerative.Provider>> fn) {
143-
return fn.apply(new Builder()).build();
169+
public static Provider vertex(
170+
String projectId,
171+
Function<GoogleGenerative.Provider.VertexBuilder, ObjectBuilder<GoogleGenerative.Provider>> fn) {
172+
return fn.apply(new VertexBuilder(projectId)).build();
173+
}
174+
175+
public static Provider aiStudio(
176+
Function<GoogleGenerative.Provider.AiStudioBuilder, ObjectBuilder<GoogleGenerative.Provider>> fn) {
177+
return fn.apply(new AiStudioBuilder()).build();
144178
}
145179

146180
@Override
@@ -205,24 +239,30 @@ public Provider(Builder builder) {
205239
builder.imageProperties);
206240
}
207241

208-
public static class Builder implements ObjectBuilder<GoogleGenerative.Provider> {
242+
public abstract static class Builder implements ObjectBuilder<GoogleGenerative.Provider> {
243+
private final String projectId;
209244
private String baseUrl;
245+
210246
private Integer topK;
211247
private Float topP;
212248
private String model;
213249
private Integer maxTokens;
214250
private Float temperature;
215251
private Float frequencyPenalty;
216252
private Float presencePenalty;
217-
private String projectId;
218253
private String endpointId;
219254
private String region;
220255
private final List<String> stopSequences = new ArrayList<>();
221256
private final List<String> images = new ArrayList<>();
222257
private final List<String> imageProperties = new ArrayList<>();
223258

259+
public Builder(String baseUrl, String projectId) {
260+
this.projectId = projectId;
261+
this.baseUrl = baseUrl;
262+
}
263+
224264
/** Base URL of the generative provider. */
225-
public Builder baseUrl(String baseUrl) {
265+
protected Builder baseUrl(String baseUrl) {
226266
this.baseUrl = baseUrl;
227267
return this;
228268
}
@@ -276,11 +316,6 @@ public Builder stopSequences(List<String> stopSequences) {
276316
return this;
277317
}
278318

279-
public Builder projectId(String projectId) {
280-
this.projectId = projectId;
281-
return this;
282-
}
283-
284319
public Builder endpointId(String endpointId) {
285320
this.endpointId = endpointId;
286321
return this;
@@ -323,5 +358,23 @@ public GoogleGenerative.Provider build() {
323358
return new GoogleGenerative.Provider(this);
324359
}
325360
}
361+
362+
public static class AiStudioBuilder extends Builder {
363+
public AiStudioBuilder() {
364+
super(Text2VecGoogleVectorizer.AiStudioBuilder.BASE_URL, null);
365+
}
366+
}
367+
368+
public static class VertexBuilder extends Builder {
369+
public VertexBuilder(String projectId) {
370+
super(Text2VecGoogleVectorizer.VertexBuilder.DEFAULT_BASE_URL, projectId);
371+
}
372+
373+
/** Base URL of the generative provider. */
374+
public VertexBuilder baseUrl(String baseUrl) {
375+
super.baseUrl(baseUrl);
376+
return this;
377+
}
378+
}
326379
}
327380
}

src/main/java/io/weaviate/client6/v1/api/collections/vectorizers/Text2VecGoogleVectorizer.java

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -203,14 +203,18 @@ public Text2VecGoogleVectorizer build() {
203203
}
204204

205205
public static class AiStudioBuilder extends Builder {
206+
public static final String BASE_URL = "generativelanguage.googleapis.com";
207+
206208
public AiStudioBuilder() {
207-
super("generativelanguage.googleapis.com", null);
209+
super(BASE_URL, null);
208210
}
209211
}
210212

211213
public static class VertexBuilder extends Builder {
214+
public static final String DEFAULT_BASE_URL = "us-central1-aiplatform.googleapis.com";
215+
212216
public VertexBuilder(String projectId) {
213-
super("us-central1-aiplatform.googleapis.com", projectId);
217+
super(DEFAULT_BASE_URL, projectId);
214218
}
215219

216220
public VertexBuilder baseUrl(String baseUrl) {

0 commit comments

Comments
 (0)