11package io .weaviate .containers ;
22
33import java .io .IOException ;
4+ import java .time .Duration ;
5+ import java .util .ArrayList ;
46import java .util .Arrays ;
57import java .util .HashMap ;
68import java .util .HashSet ;
9+ import java .util .List ;
710import java .util .Map ;
811import java .util .Set ;
912import java .util .function .Function ;
1013
14+ import org .testcontainers .containers .Network ;
15+ import org .testcontainers .containers .wait .strategy .Wait ;
16+ import org .testcontainers .lifecycle .Startable ;
1117import org .testcontainers .weaviate .WeaviateContainer ;
1218
1319import io .weaviate .client6 .v1 .api .Config ;
@@ -20,6 +26,22 @@ public class Weaviate extends WeaviateContainer {
2026 public static String OIDC_ISSUER = "https://auth.wcs.api.weaviate.io/auth/realms/SeMI" ;
2127
2228 private volatile SharedClient clientInstance ;
29+ private final String containerName ;
30+
31+ /**
32+ * By default, testcontainer's name is only available after calling
33+ * {@link #start}.
34+ * We need to know each container's name in advance to run a cluster
35+ * of several nodes, in which case we alse set the name manually.
36+ *
37+ * @see Builder#build()
38+ */
39+ @ Override
40+ public String getContainerName () {
41+ return containerName != null
42+ ? containerName
43+ : super .getContainerName ();
44+ }
2345
2446 /**
2547 * Create a new instance of WeaviateClient connected to this container if none
@@ -85,17 +107,22 @@ public static Weaviate.Builder custom() {
85107 }
86108
87109 public static class Builder {
88- private String versionTag ;
110+ private String versionTag = VERSION ;
111+ private String containerName = "weaviate" ;
89112 private Set <String > enableModules = new HashSet <>();
90113 private Set <String > adminUsers = new HashSet <>();
91114 private Set <String > viewerUsers = new HashSet <>();
92115 private Map <String , String > environment = new HashMap <>();
93116
94117 public Builder () {
95- this .versionTag = VERSION ;
96118 enableAutoSchema (false );
97119 }
98120
121+ public Builder withContainerName (String containerName ) {
122+ this .containerName = containerName ;
123+ return this ;
124+ }
125+
99126 public Builder withVersion (String version ) {
100127 this .versionTag = version ;
101128 return this ;
@@ -138,6 +165,7 @@ public Builder withFilesystemBackup(String fsPath) {
138165 environment .put ("BACKUP_FILESYSTEM_PATH" , fsPath );
139166 return this ;
140167 }
168+
141169 public Builder withAdminUsers (String ... admins ) {
142170 adminUsers .addAll (Arrays .asList (admins ));
143171 return this ;
@@ -195,7 +223,7 @@ public Builder withOIDC(String clientId, String issuer, String usernameClaim, St
195223 }
196224
197225 public Weaviate build () {
198- var c = new Weaviate (DOCKER_IMAGE + ":" + versionTag );
226+ var c = new Weaviate (containerName , DOCKER_IMAGE + ":" + versionTag );
199227
200228 if (!enableModules .isEmpty ()) {
201229 c .withEnv ("ENABLE_API_BASED_MODULES" , Boolean .TRUE .toString ());
@@ -217,13 +245,18 @@ public Weaviate build() {
217245 }
218246
219247 environment .forEach ((name , value ) -> c .withEnv (name , value ));
220- c .withCreateContainerCmdModifier (cmd -> cmd .withHostName ("weaviate" ));
248+ c .withCreateContainerCmdModifier (cmd -> cmd .withHostName (containerName ));
221249 return c ;
222250 }
223251 }
224252
225- private Weaviate (String dockerImageName ) {
253+ private Weaviate () {
254+ this ("weaviate" , DOCKER_IMAGE + ":" + VERSION );
255+ }
256+
257+ private Weaviate (String containerName , String dockerImageName ) {
226258 super (dockerImageName );
259+ this .containerName = containerName ;
227260 }
228261
229262 @ Override
@@ -264,4 +297,92 @@ private void close(Weaviate caller) throws Exception {
264297 public void close () throws IOException {
265298 }
266299 }
300+
301+ public static Weaviate cluster (int replicas ) {
302+ List <Weaviate > nodes = new ArrayList <>();
303+ for (var i = 0 ; i < replicas ; i ++) {
304+ nodes .add (Weaviate .custom ()
305+ .withContainerName ("weaviate-" + i )
306+ .build ());
307+ }
308+ return new Cluster (nodes );
309+ }
310+
311+ public static class Cluster extends Weaviate {
312+ /** Leader and followers combined. */
313+ private final List <Weaviate > nodes ;
314+
315+ private final Weaviate leader ;
316+ private final List <Weaviate > followers ;
317+
318+ private Cluster (List <Weaviate > nodes ) {
319+ assert nodes .size () > 1 : "cluster must have 1+ nodes" ;
320+
321+ this .nodes = List .copyOf (nodes );
322+ this .leader = nodes .remove (0 );
323+ this .followers = List .copyOf (nodes );
324+
325+ for (var follower : followers ) {
326+ follower .dependsOn (leader );
327+ }
328+ setNetwork (Network .SHARED );
329+ bindNodes (7110 , 7111 , 8300 );
330+ }
331+
332+ @ Override
333+ public WeaviateContainer dependsOn (List <? extends Startable > startables ) {
334+ leader .dependsOn (startables );
335+ return this ;
336+ }
337+
338+ @ Override
339+ public void setNetwork (Network network ) {
340+ nodes .forEach (node -> node .setNetwork (network ));
341+ }
342+
343+ @ Override
344+ public WeaviateClient getClient () {
345+ if (!isRunning ()) {
346+ start ();
347+ }
348+ return leader .getClient ();
349+ }
350+
351+ @ Override
352+ public void start () {
353+ followers .forEach (Startable ::start ); // testcontainers will resolve dependencies
354+ }
355+
356+ @ Override
357+ public void stop () {
358+ followers .forEach (Startable ::stop );
359+ leader .stop ();
360+ }
361+
362+ /**
363+ * Set environment variables for inter-cluster communication.
364+ *
365+ * @param gossip Gossip bind port.
366+ * @param data Data bind port.
367+ * @param raft RAFT port.
368+ */
369+ private void bindNodes (int gossip , int data , int raft ) {
370+ var publicPort = leader .getExposedPorts ().get (0 ); // see WeaviateContainer Testcontainer.
371+
372+ nodes .forEach (node -> node
373+ .withEnv ("CLUSTER_GOSSIP_BIND_PORT" , String .valueOf (gossip ))
374+ .withEnv ("CLUSTER_DATA_BIND_PORT" , String .valueOf (data ))
375+ .withEnv ("RAFT_PORT" , String .valueOf (raft ))
376+ .withEnv ("RAFT_BOOTSTRAP_EXPECT" , "1" ));
377+
378+ followers .forEach (node -> node
379+ .withEnv ("CLUSTER_JOIN" , leader .containerName + ":" + gossip )
380+ .withEnv ("RAFT_JOIN" , leader .containerName )
381+ .waitingFor (
382+ Wait .forHttp ("/v1/.well-known/ready" )
383+ .forPort (publicPort )
384+ .forStatusCode (200 )
385+ .withStartupTimeout (Duration .ofSeconds (10 ))));
386+ }
387+ }
267388}
0 commit comments