Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -114,7 +114,7 @@ public class WorkflowApplication implements AutoCloseable {
private final WorkflowLifeCycleCloudEventFactory lifeCycleCloudEventFactory;
private final ScheduledExecutorService schedulerExecutorService;
private final Set<String> allowedCommands;
private final Map<Class<?>, List<?>> serviceLoadedClasses = new ConcurrentHashMap<>();
private final Map<Class<?>, ServiceLoader<?>> servicesLoaded = new ConcurrentHashMap<>();
Comment thread
fjtirado marked this conversation as resolved.

private WorkflowApplication(Builder builder) {
this.taskFactory = builder.taskFactory;
Expand Down Expand Up @@ -712,21 +712,19 @@ public Set<String> allowedCommands() {

@SuppressWarnings("unchecked")
public <T extends Comparable<?>> List<T> serviceLoadedClasses(Class<T> clazz) {
return (List<T>)
serviceLoadedClasses.computeIfAbsent(
clazz,
c ->
ServiceLoader.load(clazz).stream()
.map(ServiceLoader.Provider::get)
.sorted()
.toList());
ServiceLoader<?> serviceLoader = servicesLoaded.computeIfAbsent(clazz, ServiceLoader::load);
return (List<T>) serviceLoader.stream().map(ServiceLoader.Provider::get).sorted().toList();
}

public <T extends Comparable<?>> T serviceLoadedClass(Class<T> serviceClass) {
List<T> list = serviceLoadedClasses(serviceClass);
if (list.isEmpty()) {
throw new IllegalStateException("No " + serviceClass + " implementation found");
}
return list.get(0);
ServiceLoader<?> serviceLoader =
servicesLoaded.computeIfAbsent(serviceClass, ServiceLoader::load);
return (T)
serviceLoader.stream()
.map(ServiceLoader.Provider::get)
.sorted()
.findFirst()
.orElseThrow(
() -> new IllegalStateException("No " + serviceClass + " implementation found"));
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
import io.serverlessworkflow.api.types.CallTask;
import io.serverlessworkflow.api.types.Task;
import io.serverlessworkflow.api.types.TaskBase;
import io.serverlessworkflow.impl.WorkflowApplication;
import io.serverlessworkflow.impl.WorkflowDefinition;
import io.serverlessworkflow.impl.WorkflowMutablePosition;
import io.serverlessworkflow.impl.executors.CallTaskExecutor.CallTaskExecutorBuilder;
Expand All @@ -32,9 +33,7 @@
import io.serverlessworkflow.impl.executors.SwitchExecutor.SwitchExecutorBuilder;
import io.serverlessworkflow.impl.executors.TryExecutor.TryExecutorBuilder;
import io.serverlessworkflow.impl.executors.WaitExecutor.WaitExecutorBuilder;
import java.util.Collection;
import java.util.ServiceLoader;
import java.util.ServiceLoader.Provider;
import java.util.List;

public class DefaultTaskExecutorFactory implements TaskExecutorFactory {

Expand All @@ -46,9 +45,6 @@ public static TaskExecutorFactory get() {

protected DefaultTaskExecutorFactory() {}

private Collection<CallableTaskBuilder> callTasks =
ServiceLoader.load(CallableTaskBuilder.class).stream().map(Provider::get).sorted().toList();

@Override
public TaskExecutorBuilder<? extends TaskBase> getTaskExecutor(
WorkflowMutablePosition position, Task task, WorkflowDefinition definition) {
Expand All @@ -57,7 +53,10 @@ public TaskExecutorBuilder<? extends TaskBase> getTaskExecutor(
TaskBase taskBase = (TaskBase) callTask.get();
if (taskBase != null) {
return new CallTaskExecutorBuilder(
position, taskBase, definition, findCallTask(taskBase.getClass()));
position,
taskBase,
definition,
findCallTask(taskBase.getClass(), definition.application()));
}
} else if (task.getSwitchTask() != null) {
return new SwitchExecutorBuilder(position, task.getSwitchTask(), definition);
Expand Down Expand Up @@ -86,7 +85,9 @@ public TaskExecutorBuilder<? extends TaskBase> getTaskExecutor(
}

@SuppressWarnings("unchecked")
private <T extends TaskBase> CallableTaskBuilder<T> findCallTask(Class<T> clazz) {
private <T extends TaskBase> CallableTaskBuilder<T> findCallTask(
Class<T> clazz, WorkflowApplication app) {
List<CallableTaskBuilder> callTasks = app.serviceLoadedClasses(CallableTaskBuilder.class);
return (CallableTaskBuilder<T>)
callTasks.stream()
.filter(s -> s.accept(clazz))
Expand Down
Loading