Skip to content

Commit fc5916a

Browse files
committed
[grid] Provide a command line flag for creating SessionFactory instances
1 parent cf471a3 commit fc5916a

5 files changed

Lines changed: 152 additions & 22 deletions

File tree

java/server/src/org/openqa/selenium/grid/config/TomlConfig.java

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,7 @@ public class TomlConfig implements Config {
3939

4040
private final Toml toml;
4141

42-
TomlConfig(Reader reader) {
42+
public TomlConfig(Reader reader) {
4343
try {
4444
toml = JToml.parse(reader);
4545
} catch (IOException e) {

java/server/src/org/openqa/selenium/grid/node/config/NodeOptions.java

Lines changed: 74 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,8 @@
3232
import org.openqa.selenium.json.JsonOutput;
3333
import org.openqa.selenium.remote.service.DriverService;
3434

35+
import java.lang.reflect.Method;
36+
import java.lang.reflect.Modifier;
3537
import java.net.URI;
3638
import java.net.URISyntaxException;
3739
import java.util.ArrayList;
@@ -81,32 +83,90 @@ public Map<Capabilities, Collection<SessionFactory>> getSessionFactories(
8183

8284
Map<WebDriverInfo, Collection<SessionFactory>> allDrivers = discoverDrivers(maxSessions, factoryFactory);
8385

84-
// If drivers have been specified, use those.
85-
List<String> drivers = config.getAll("node", "drivers").orElse(new ArrayList<>()).stream()
86-
.map(String::toLowerCase)
87-
.collect(Collectors.toList());
88-
8986
ImmutableMultimap.Builder<Capabilities, SessionFactory> sessionFactories = ImmutableMultimap.builder();
9087

91-
if (!drivers.isEmpty()) {
92-
allDrivers.entrySet().stream()
93-
.filter(entry -> drivers.contains(entry.getKey().getDisplayName().toLowerCase()))
94-
.sorted(Comparator.comparing(entry -> entry.getKey().getDisplayName().toLowerCase()))
95-
.peek(this::report)
96-
.forEach(entry -> sessionFactories.putAll(entry.getKey().getCanonicalCapabilities(), entry.getValue()));
88+
addDriverFactoriesFromConfig(sessionFactories);
89+
addSpecificDrivers(allDrivers, sessionFactories);
90+
addDetectedDrivers(allDrivers, sessionFactories);
91+
92+
return sessionFactories.build().asMap();
93+
}
94+
95+
private void addDriverFactoriesFromConfig(ImmutableMultimap.Builder<Capabilities, SessionFactory> sessionFactories) {
96+
Optional<List<String>> additionalDriverFactories = config.getAll("node", "driver-factories");
97+
if (!additionalDriverFactories.isPresent()) {
98+
return;
99+
}
100+
101+
List<String> allConfigs = additionalDriverFactories.get();
102+
if (allConfigs.size() % 2 != 0) {
103+
throw new ConfigException("Expected each driver class to be mapped to a config");
104+
}
105+
106+
for (int i = 0; i < allConfigs.size(); i++) {
107+
SessionFactory sessionFactory = createSessionFactory(allConfigs.get(i));
108+
i++;
109+
if (i == allConfigs.size()) {
110+
throw new ConfigException("Unable to find JSON config");
111+
}
112+
Capabilities stereotype = JSON.toType(allConfigs.get(i), Capabilities.class);
113+
114+
sessionFactories.put(stereotype, sessionFactory);
115+
}
116+
}
117+
118+
private SessionFactory createSessionFactory(String clazz) {
119+
LOG.fine(String.format("Creating %s as instance of %s", clazz, SessionFactory.class));
120+
121+
try {
122+
// Use the context class loader since this is what the `--ext`
123+
// flag modifies.
124+
Class<?> ClassClazz = Class.forName(clazz, true, Thread.currentThread().getContextClassLoader());
125+
Method create = ClassClazz.getMethod("create", Config.class);
126+
127+
if (!Modifier.isStatic(create.getModifiers())) {
128+
throw new IllegalArgumentException(String.format(
129+
"Class %s's `create(Config)` method must be static", clazz));
130+
}
131+
132+
if (!SessionFactory.class.isAssignableFrom(create.getReturnType())) {
133+
throw new IllegalArgumentException(String.format(
134+
"Class %s's `create(Config)` method must be static", clazz));
135+
}
97136

98-
return sessionFactories.build().asMap();
137+
return (SessionFactory) create.invoke(null, config);
138+
} catch (NoSuchMethodException e) {
139+
throw new IllegalArgumentException(String.format(
140+
"Class %s must have a static `create(Config)` method", clazz));
141+
} catch (ReflectiveOperationException e) {
142+
throw new IllegalArgumentException("Unable to find class: " + clazz, e);
99143
}
144+
}
100145

146+
private void addDetectedDrivers(
147+
Map<WebDriverInfo, Collection<SessionFactory>> allDrivers,
148+
ImmutableMultimap.Builder<Capabilities, SessionFactory> sessionFactories) {
101149
if (!config.getBool("node", "detect-drivers").orElse(false)) {
102-
return sessionFactories.build().asMap();
150+
return;
103151
}
104152

105153
allDrivers.entrySet().stream()
106154
.peek(this::report)
107155
.forEach(entry -> sessionFactories.putAll(entry.getKey().getCanonicalCapabilities(), entry.getValue()));
156+
}
108157

109-
return sessionFactories.build().asMap();
158+
private void addSpecificDrivers(
159+
Map<WebDriverInfo, Collection<SessionFactory>> allDrivers,
160+
ImmutableMultimap.Builder<Capabilities, SessionFactory> sessionFactories) {
161+
List<String> drivers = config.getAll("node", "drivers").orElse(new ArrayList<>()).stream()
162+
.map(String::toLowerCase)
163+
.collect(Collectors.toList());
164+
165+
allDrivers.entrySet().stream()
166+
.filter(entry -> drivers.contains(entry.getKey().getDisplayName().toLowerCase()))
167+
.sorted(Comparator.comparing(entry -> entry.getKey().getDisplayName().toLowerCase()))
168+
.peek(this::report)
169+
.forEach(entry -> sessionFactories.putAll(entry.getKey().getCanonicalCapabilities(), entry.getValue()));
110170
}
111171

112172
private Map<WebDriverInfo, Collection<SessionFactory>> discoverDrivers(

java/server/src/org/openqa/selenium/grid/node/httpd/NodeFlags.java

Lines changed: 19 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -26,36 +26,48 @@
2626
import java.net.URL;
2727
import java.util.Collections;
2828
import java.util.HashSet;
29+
import java.util.List;
2930
import java.util.Set;
3031

3132
import static org.openqa.selenium.grid.config.StandardGridRoles.NODE_ROLE;
3233

3334
@AutoService(HasRoles.class)
3435
public class NodeFlags implements HasRoles {
3536

36-
@Parameter(
37-
names = {"--detect-drivers"}, arity = 1,
38-
description = "Autodetect which drivers are available on the current system, and add them to the node.")
39-
@ConfigValue(section = "node", name = "detect-drivers", example = "true")
40-
public Boolean autoconfigure = true;
41-
4237
@Parameter(
4338
names = "--max-sessions",
4439
description = "Maximum number of concurrent sessions.")
4540
@ConfigValue(section = "node", name = "max-concurrent-sessions", example = "8")
4641
public int maxSessions = Runtime.getRuntime().availableProcessors();
4742

43+
@Parameter(
44+
names = {"--detect-drivers"}, arity = 1,
45+
description = "Autodetect which drivers are available on the current system, and add them to the node.")
46+
@ConfigValue(section = "node", name = "detect-drivers", example = "true")
47+
public Boolean autoconfigure = true;
48+
4849
@Parameter(
4950
names = {"-I", "--driver-implementation"},
5051
description = "Drivers that should be checked. If specified, will skip autoconfiguration. Example: -I \"firefox\" -I \"chrome\"")
5152
@ConfigValue(section = "node", name = "drivers", example = "[\"firefox\", \"chrome\"]")
5253
public Set<String> driverNames = new HashSet<>();
5354

55+
@Parameter(
56+
names = {"--driver-factory"},
57+
description = "Mapping of fully qualified class name to a browser configuration that this matches against. " +
58+
"`--driver-factory org.openqa.selenium.example.LynxDriverFactory '{\"browserName\": \"lynx\"}')",
59+
arity = 2,
60+
variableArity = true)
61+
@ConfigValue(
62+
section = "node",
63+
name = "driver-factories",
64+
example = "[\"org.openqa.selenium.example.LynxDriverFactory '{\"browserName\": \"lynx\"}']")
65+
public List<String> driverFactory2Config;
66+
5467
@Parameter(
5568
names = {"--public-url"},
5669
description = "Public URL of the Grid as a whole (typically the address of the hub or the router)")
5770
@ConfigValue(section = "node", name = "grid-url", example = "\"https://grid.example.com\"")
58-
5971
public URL gridUri;
6072

6173
@Override

java/server/test/org/openqa/selenium/grid/node/config/BUILD.bazel

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@ java_test_suite(
1111
"//java/client/src/org/openqa/selenium/edge",
1212
"//java/client/src/org/openqa/selenium/firefox",
1313
"//java/client/src/org/openqa/selenium/ie",
14+
"//java/client/src/org/openqa/selenium/json",
1415
"//java/client/src/org/openqa/selenium/remote",
1516
"//java/client/src/org/openqa/selenium/remote/http",
1617
"//java/client/src/org/openqa/selenium/safari",

java/server/test/org/openqa/selenium/grid/node/config/NodeOptionsTest.java

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,30 +17,45 @@
1717

1818
package org.openqa.selenium.grid.node.config;
1919

20+
import com.google.common.collect.ImmutableList;
21+
import com.google.common.collect.ImmutableMap;
2022
import org.assertj.core.api.Condition;
2123
import org.junit.Before;
2224
import org.junit.Test;
25+
import org.openqa.selenium.Capabilities;
26+
import org.openqa.selenium.ImmutableCapabilities;
2327
import org.openqa.selenium.Platform;
2428
import org.openqa.selenium.WebDriverInfo;
2529
import org.openqa.selenium.chrome.ChromeDriverInfo;
2630
import org.openqa.selenium.events.EventBus;
2731
import org.openqa.selenium.events.local.GuavaEventBus;
2832
import org.openqa.selenium.grid.config.Config;
2933
import org.openqa.selenium.grid.config.MapConfig;
34+
import org.openqa.selenium.grid.config.TomlConfig;
35+
import org.openqa.selenium.grid.data.CreateSessionRequest;
36+
import org.openqa.selenium.grid.node.ActiveSession;
37+
import org.openqa.selenium.grid.node.SessionFactory;
3038
import org.openqa.selenium.grid.node.local.LocalNode;
39+
import org.openqa.selenium.json.Json;
3140
import org.openqa.selenium.remote.http.HttpClient;
3241
import org.openqa.selenium.remote.tracing.DefaultTestTracer;
3342
import org.openqa.selenium.remote.tracing.Tracer;
3443

44+
import java.io.StringReader;
3545
import java.net.URI;
3646
import java.net.URISyntaxException;
3747
import java.util.ArrayList;
48+
import java.util.Collection;
3849
import java.util.Collections;
3950
import java.util.List;
51+
import java.util.Map;
52+
import java.util.Optional;
4053

4154
import static java.util.Collections.emptyMap;
55+
import static java.util.Collections.emptySet;
4256
import static java.util.Collections.singletonMap;
4357
import static org.assertj.core.api.Assertions.assertThat;
58+
import static org.assertj.core.api.Assertions.fail;
4459
import static org.junit.Assume.assumeFalse;
4560
import static org.junit.Assume.assumeTrue;
4661
import static org.mockito.Mockito.spy;
@@ -155,10 +170,52 @@ public void doNotDetectDriversByDefault() {
155170
assertThat(reported).isEmpty();
156171
}
157172

173+
@Test
174+
public void canBeConfiguredToUseHelperClassesToCreateSessionFactories() {
175+
Capabilities caps = new ImmutableCapabilities("browserName", "cheese");
176+
StringBuilder capsString = new StringBuilder();
177+
new Json().newOutput(capsString).setPrettyPrint(false).write(caps);
178+
179+
Config config = new TomlConfig(new StringReader(String.format(
180+
"[node]\n" +
181+
"detect-drivers = false\n" +
182+
"driver-factories = [" +
183+
" \"%s\",\n" +
184+
" \"%s\"\n" +
185+
"]",
186+
HelperFactory.class.getName(),
187+
capsString.toString().replace("\"", "\\\""))));
188+
189+
190+
NodeOptions options = new NodeOptions(config);
191+
Map<Capabilities, Collection<SessionFactory>> factories = options.getSessionFactories(info -> emptySet());
192+
193+
Collection<SessionFactory> sessionFactories = factories.get(caps);
194+
assertThat(sessionFactories).size().isEqualTo(1);
195+
assertThat(sessionFactories.iterator().next()).isInstanceOf(SessionFactory.class);
196+
}
197+
158198
private Condition<? super List<? extends WebDriverInfo>> supporting(String name) {
159199
return new Condition<>(
160200
infos -> infos.stream().anyMatch(info -> name.equals(info.getCanonicalCapabilities().getBrowserName())),
161201
"supporting %s",
162202
name);
163203
}
204+
205+
public static class HelperFactory {
206+
207+
public static SessionFactory create(Config config) {
208+
return new SessionFactory() {
209+
@Override
210+
public Optional<ActiveSession> apply(CreateSessionRequest createSessionRequest) {
211+
return Optional.empty();
212+
}
213+
214+
@Override
215+
public boolean test(Capabilities capabilities) {
216+
return true;
217+
}
218+
};
219+
}
220+
}
164221
}

0 commit comments

Comments
 (0)