99
1010import io .weaviate .client6 .v1 .api .collections .Generative ;
1111import io .weaviate .client6 .v1 .api .collections .generate .DynamicProvider ;
12+ import io .weaviate .client6 .v1 .api .collections .vectorizers .Text2VecAwsVectorizer .Service ;
1213import io .weaviate .client6 .v1 .internal .ObjectBuilder ;
1314import io .weaviate .client6 .v1 .internal .grpc .protocol .WeaviateProtoBase ;
1415import io .weaviate .client6 .v1 .internal .grpc .protocol .WeaviateProtoGenerative ;
1516
1617public 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}
0 commit comments