Skip to content

Commit e4eb52b

Browse files
committed
feat: provide separate builder for AWS sagemaker/bedrock
1 parent 5ca18f9 commit e4eb52b

5 files changed

Lines changed: 176 additions & 58 deletions

File tree

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

Lines changed: 31 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -94,27 +94,47 @@ public static Generative anyscale(Function<AnyscaleGenerative.Builder, ObjectBui
9494
}
9595

9696
/**
97-
* Configure a default {@code generative-aws} module.
97+
* Configure a default {@code generative-aws} module with Bedrock integration.
98+
*
99+
* @param region AWS region.
100+
* @param model Model to use with Bedrock service.
101+
*/
102+
public static Generative awsBedrock(String region, String model) {
103+
return AwsGenerative.bedrock(region, model);
104+
}
105+
106+
/**
107+
* Configure a {@code generative-aws} module with Bedrock integration.
108+
*
109+
* @param region AWS region.
110+
* @param model Model to use with Bedrock service.
111+
* @param fn Lambda expression for optional parameters.
112+
*/
113+
public static Generative awsBedrock(String region, String model,
114+
Function<AwsGenerative.BedrockBuilder, ObjectBuilder<AwsGenerative>> fn) {
115+
return AwsGenerative.bedrock(region, model, fn);
116+
}
117+
118+
/**
119+
* Configure a default {@code generative-aws} module with Sagemaker integration.
98120
*
99121
* @param region AWS region.
100-
* @param service AWS service to use, e.g. {@code "bedrock"} or
101-
* {@code "sagemaker"}.
122+
* @param baseUrl Base inference URL.
102123
*/
103-
public static Generative aws(String region, String service) {
104-
return AwsGenerative.of(region, service);
124+
public static Generative awsSagemaker(String region, String baseUrl) {
125+
return AwsGenerative.sagemaker(region, baseUrl);
105126
}
106127

107128
/**
108-
* Configure a {@code generative-aws} module.
129+
* Configure a {@code generative-aws} module with Sagemaker integration.
109130
*
110131
* @param region AWS region.
111-
* @param service AWS service to use, e.g. {@code "bedrock"} or
112-
* {@code "sagemaker"}.
132+
* @param baseUrl Base inference URL.
113133
* @param fn Lambda expression for optional parameters.
114134
*/
115-
public static Generative aws(String region, String service,
116-
Function<AwsGenerative.Builder, ObjectBuilder<AwsGenerative>> fn) {
117-
return AwsGenerative.of(region, service, fn);
135+
public static Generative awsSagemaker(String region, String baseUrl,
136+
Function<AwsGenerative.SagemakerBuilder, ObjectBuilder<AwsGenerative>> fn) {
137+
return AwsGenerative.sagemaker(region, baseUrl, fn);
118138
}
119139

120140
/** Configure a default {@code generative-cohere} module. */

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

Lines changed: 22 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -44,11 +44,29 @@ public static DynamicProvider anyscale(
4444
/**
4545
* Configure {@code generative-aws} as a dynamic provider.
4646
*
47-
* @param fn Lambda expression for optional parameters.
47+
* @param region AWS region.
48+
* @param model Inference model.
49+
* @param fn Lambda expression for optional parameters.
50+
*/
51+
public static DynamicProvider awsBedrock(
52+
String region,
53+
String model,
54+
Function<AwsGenerative.Provider.BedrockBuilder, ObjectBuilder<AwsGenerative.Provider>> fn) {
55+
return AwsGenerative.Provider.bedrock(region, model, fn);
56+
}
57+
58+
/**
59+
* Configure {@code generative-aws} as a dynamic provider.
60+
*
61+
* @param region AWS region.
62+
* @param baseUrl Base inference URL.
63+
* @param fn Lambda expression for optional parameters.
4864
*/
49-
public static DynamicProvider aws(
50-
Function<AwsGenerative.Provider.Builder, ObjectBuilder<AwsGenerative.Provider>> fn) {
51-
return AwsGenerative.Provider.of(fn);
65+
public static DynamicProvider awsSagemaker(
66+
String region,
67+
String baseUrl,
68+
Function<AwsGenerative.Provider.SagemakerBuilder, ObjectBuilder<AwsGenerative.Provider>> fn) {
69+
return AwsGenerative.Provider.sagemaker(region, baseUrl, fn);
5270
}
5371

5472
/**

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

Lines changed: 95 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -9,13 +9,14 @@
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.Text2VecAwsVectorizer.Service;
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;
1516

1617
public record AwsGenerative(
1718
@SerializedName("region") String region,
18-
@SerializedName("service") String service,
19+
@SerializedName("service") Service service,
1920
@SerializedName("endpoint") String baseUrl,
2021
@SerializedName("model") String model) implements Generative {
2122

@@ -29,27 +30,37 @@ public Object _self() {
2930
return this;
3031
}
3132

32-
public static AwsGenerative of(String region, String service) {
33-
return of(region, service, ObjectBuilder.identity());
33+
public static AwsGenerative bedrock(String region, String model) {
34+
return bedrock(region, model, ObjectBuilder.identity());
3435
}
3536

36-
public static AwsGenerative of(String region, String service, Function<Builder, ObjectBuilder<AwsGenerative>> fn) {
37-
return fn.apply(new Builder(region, service)).build();
37+
public static AwsGenerative bedrock(String region, String model,
38+
Function<BedrockBuilder, ObjectBuilder<AwsGenerative>> fn) {
39+
return fn.apply(new BedrockBuilder(region, model)).build();
40+
}
41+
42+
public static AwsGenerative sagemaker(String region, String baseUrl) {
43+
return sagemaker(region, baseUrl, ObjectBuilder.identity());
44+
}
45+
46+
public static AwsGenerative sagemaker(String region, String baseUrl,
47+
Function<SagemakerBuilder, ObjectBuilder<AwsGenerative>> fn) {
48+
return fn.apply(new SagemakerBuilder(region, baseUrl)).build();
3849
}
3950

4051
public AwsGenerative(Builder builder) {
4152
this(
42-
builder.service,
4353
builder.region,
54+
builder.service,
4455
builder.baseUrl,
4556
builder.model);
4657
}
4758

4859
public static class Builder implements ObjectBuilder<AwsGenerative> {
4960
private final String region;
50-
private final String service;
61+
private final Service service;
5162

52-
public Builder(String service, String region) {
63+
public Builder(Service service, String region) {
5364
this.service = service;
5465
this.region = region;
5566
}
@@ -58,13 +69,13 @@ public Builder(String service, String region) {
5869
private String model;
5970

6071
/** Base URL of the generative provider. */
61-
public Builder baseUrl(String baseUrl) {
72+
protected Builder baseUrl(String baseUrl) {
6273
this.baseUrl = baseUrl;
6374
return this;
6475
}
6576

6677
/** Select generative model. */
67-
public Builder model(String model) {
78+
protected Builder model(String model) {
6879
this.model = model;
6980
return this;
7081
}
@@ -75,12 +86,37 @@ public AwsGenerative build() {
7586
}
7687
}
7788

89+
public static class BedrockBuilder extends Builder {
90+
public BedrockBuilder(String region, String model) {
91+
super(Service.BEDROCK, region);
92+
super.model(model);
93+
}
94+
95+
@Override
96+
/** Required for {@link Service#BEDROCK}. */
97+
public Builder model(String model) {
98+
return super.model(model);
99+
}
100+
}
101+
102+
public static class SagemakerBuilder extends Builder {
103+
public SagemakerBuilder(String region, String baseUrl) {
104+
super(Service.SAGEMAKER, region);
105+
super.baseUrl(baseUrl);
106+
}
107+
108+
/** Required for {@link Service#SAGEMAKER}. */
109+
public Builder baseUrl(String baseUrl) {
110+
return super.baseUrl(baseUrl);
111+
}
112+
}
113+
78114
public static record Metadata() implements ProviderMetadata {
79115
}
80116

81117
public static record Provider(
82118
String region,
83-
String service,
119+
Service service,
84120
String baseUrl,
85121
String model,
86122
String targetModel,
@@ -89,9 +125,18 @@ public static record Provider(
89125
List<String> images,
90126
List<String> imageProperties) implements DynamicProvider {
91127

92-
public static Provider of(
93-
Function<AwsGenerative.Provider.Builder, ObjectBuilder<AwsGenerative.Provider>> fn) {
94-
return fn.apply(new Builder()).build();
128+
public static Provider bedrock(
129+
String region,
130+
String model,
131+
Function<AwsGenerative.Provider.BedrockBuilder, ObjectBuilder<AwsGenerative.Provider>> fn) {
132+
return fn.apply(new BedrockBuilder(region, model)).build();
133+
}
134+
135+
public static Provider sagemaker(
136+
String region,
137+
String baseUrl,
138+
Function<AwsGenerative.Provider.SagemakerBuilder, ObjectBuilder<AwsGenerative.Provider>> fn) {
139+
return fn.apply(new SagemakerBuilder(region, baseUrl)).build();
95140
}
96141

97142
@Override
@@ -102,7 +147,10 @@ public void appendTo(
102147
provider.setRegion(region);
103148
}
104149
if (service != null) {
105-
provider.setService(service);
150+
provider.setService(
151+
service == Service.BEDROCK ? "bedrock"
152+
: service == Service.SAGEMAKER ? "sagemaker"
153+
: "unknown");
106154
}
107155
if (baseUrl != null) {
108156
provider.setEndpoint(baseUrl);
@@ -143,9 +191,9 @@ public Provider(Builder builder) {
143191
builder.imageProperties);
144192
}
145193

146-
public static class Builder implements ObjectBuilder<AwsGenerative.Provider> {
147-
private String region;
148-
private String service;
194+
public abstract static class Builder implements ObjectBuilder<AwsGenerative.Provider> {
195+
private final Service service;
196+
private final String region;
149197
private String baseUrl;
150198
private String model;
151199
private String targetModel;
@@ -154,24 +202,19 @@ public static class Builder implements ObjectBuilder<AwsGenerative.Provider> {
154202
private final List<String> images = new ArrayList<>();
155203
private final List<String> imageProperties = new ArrayList<>();
156204

157-
public Builder region(String region) {
158-
this.region = region;
159-
return this;
160-
}
161-
162-
public Builder service(String service) {
205+
protected Builder(Service service, String region) {
163206
this.service = service;
164-
return this;
207+
this.region = region;
165208
}
166209

167210
/** Base URL of the generative provider. */
168-
public Builder baseUrl(String baseUrl) {
211+
protected Builder baseUrl(String baseUrl) {
169212
this.baseUrl = baseUrl;
170213
return this;
171214
}
172215

173216
/** Select generative model. */
174-
public Builder model(String model) {
217+
protected Builder model(String model) {
175218
this.model = model;
176219
return this;
177220
}
@@ -218,5 +261,30 @@ public AwsGenerative.Provider build() {
218261
return new AwsGenerative.Provider(this);
219262
}
220263
}
264+
265+
public static class BedrockBuilder extends Builder {
266+
public BedrockBuilder(String region, String model) {
267+
super(Service.BEDROCK, region);
268+
super.model(model);
269+
}
270+
271+
@Override
272+
/** Required for {@link Service#BEDROCK}. */
273+
public Builder model(String model) {
274+
return super.model(model);
275+
}
276+
}
277+
278+
public static class SagemakerBuilder extends Builder {
279+
public SagemakerBuilder(String region, String baseUrl) {
280+
super(Service.SAGEMAKER, region);
281+
super.baseUrl(baseUrl);
282+
}
283+
284+
/** Required for {@link Service#SAGEMAKER}. */
285+
public Builder baseUrl(String baseUrl) {
286+
return super.baseUrl(baseUrl);
287+
}
288+
}
221289
}
222290
}

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

Lines changed: 8 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -109,16 +109,20 @@ public Text2VecAwsVectorizer(Builder builder) {
109109
builder.quantization);
110110
}
111111

112-
private abstract static class Builder implements ObjectBuilder<Text2VecAwsVectorizer> {
112+
public abstract static class Builder implements ObjectBuilder<Text2VecAwsVectorizer> {
113113
private final boolean vectorizeCollectionName = false;
114114
private Quantization quantization;
115115
private List<String> sourceProperties = new ArrayList<>();
116116
private VectorIndex vectorIndex = VectorIndex.DEFAULT_VECTOR_INDEX;
117117

118+
private final Service service;
118119
private String baseUrl;
119120
private String model;
120121
private String region;
121-
private Service service;
122+
123+
protected Builder(Service service) {
124+
this.service = service;
125+
}
122126

123127
/** Required for {@link Service#SAGEMAKER}. */
124128
protected Builder baseUrl(String baseUrl) {
@@ -137,11 +141,6 @@ public Builder region(String region) {
137141
return this;
138142
}
139143

140-
public Builder service(Service service) {
141-
this.service = service;
142-
return this;
143-
}
144-
145144
/** Add properties to include in the embedding. */
146145
public Builder sourceProperties(String... properties) {
147146
return sourceProperties(Arrays.asList(properties));
@@ -177,8 +176,7 @@ public Text2VecAwsVectorizer build() {
177176

178177
public static class BedrockBuilder extends Builder {
179178
public BedrockBuilder(String model) {
180-
super();
181-
super.service(Service.BEDROCK);
179+
super(Service.BEDROCK);
182180
super.model(model);
183181
}
184182

@@ -191,8 +189,7 @@ public Builder model(String model) {
191189

192190
public static class SagemakerBuilder extends Builder {
193191
public SagemakerBuilder(String baseUrl) {
194-
super();
195-
super.service(Service.SAGEMAKER);
192+
super(Service.SAGEMAKER);
196193
super.baseUrl(baseUrl);
197194
}
198195

0 commit comments

Comments
 (0)