You can achieve that by utilizing the combination of ListIterator and Iterator. This approach allows you to free yourself from the necessity to maintain any positions.
The general idea is to create a list of iterators and a listIterator that would iterate over these lists.
Method hasNext() performs iteration over the list and checks whether it has at least one not exhausted iterator.
Method next() is a bit more elaborate. Firstly, it will try "reset" the listIterator from the position behind the last element to the position before the first element of the list.
Then inside the loop, it'll examine the next iterator in the list and if this iterator is exhausted it'll get removed. In order to make this operation fast, iterators are stored in a LinkedList. And if the element was found, it will be returned.
To treat the situation when "empty" iterator appear at the end of the list and listIterator reaches the last position, but there could a few more unvisited iterators in the list, an additional check is placed at the very end of the loop.
In the implementation below, I've used streams inside hasNext() and constructor for the purpose of conciseness. Even if you are not very comfortable with streams, the overall logic should clear, and you can easily substitute it with streams.
public class CustomListIterator<T> implements Iterator<T> {
private List<Iterator<T>> iterators;
private ListIterator<Iterator<T>> listIterator;
public CustomListIterator(List<List<T>> nestedList) {
this.iterators = nestedList.stream()
.map(List::iterator)
.collect(Collectors.toCollection(LinkedList::new));
this.listIterator = iterators.listIterator();
}
@Override
public boolean hasNext() {
return iterators.stream().anyMatch(Iterator::hasNext);
}
@Override
public T next() {
if (!iterators.isEmpty() && !listIterator.hasNext()) tryReset();
while (!iterators.isEmpty() && listIterator.hasNext()) {
Iterator<T> current = listIterator.next();
if (!current.hasNext()) {
listIterator.remove(); // removing exhausted iterator
} else {
return current.next();
}
if (!listIterator.hasNext()) tryReset();
}
throw new IllegalStateException();
}
private void tryReset() {
while (listIterator.hasPrevious()) {
listIterator.previous();
}
}
}
main() - demo
public static void main(String[] args) {
List<List<Integer>> numbers =
List.of(List.of(1, 4, 6, 7),
List.of(),
List.of(2, 5),
List.of(3));
CustomListIterator<Integer> nestedIterator = new CustomListIterator<>(numbers);
while (nestedIterator.hasNext()) {
System.out.print(nestedIterator.next() + "\t");
}
}
Output
1 2 3 4 5 6 7