class Solution {
public:
int shortestPathAllKeys(vector<string>& grid) {
int keys = accumulate(grid.begin(), grid.end(), 0, [] (auto acc, auto &s) {
return acc + accumulate(s.begin(), s.end(), 0, [] (auto acc, auto &c) {
return islower(c) ? acc | (1 << (c - 'a')) : acc;
});
});
auto [start_i, start_j] = getStart(grid);
int init_state = getNextState(0, start_i, start_j);
unordered_set<int> visited = {init_state};
vector<int> q = {init_state};
int ret = 0;
while (!q.empty()) {
vector<int> tmp_q;
ret += 1;
for (auto state : q) {
int adj[][2] = {{-1,0}, {1,0}, {0,-1}, {0,1}};
auto [x, y] = getPos(state);
for (auto [m, n] : adj) {
int i = x + m;
int j = y + n;
if (i < 0 || i >= grid.size() || j < 0 || j >= grid[0].length() || grid[i][j] == '#') continue;
if (isupper(grid[i][j]) && !hasKey(state, grid[i][j])) continue;
int next_state = getNextState(state, i, j);
if (islower(grid[i][j]) && !hasKey(next_state, grid[i][j])) {
next_state = getNextState(next_state, grid[i][j]);
if (keys == getKey(next_state)) return ret;
}
if (visited.count(next_state)) {
continue;
}
visited.insert(next_state);
tmp_q.push_back(next_state);
}
}
q.swap(tmp_q);
}
return -1;
}
pair<int, int> getStart(vector<string>& grid) {
for (int i = 0; i < grid.size(); ++i) {
for (int j = 0; j < grid[0].length(); ++j) {
if (grid[i][j] == '@') return {i, j};
}
}
return {};
}
int getNextState(int state, int i, int j) {
return (i << 16) | (j << 8) | (0xff & state);
}
int getNextState(int state, char key) {
return state | (1 << (key - 'a'));
}
pair<int, int> getPos(int state) {
return {state >> 16, (state & 0xff00) >> 8};
}
int getKey(int state) {
return state & 0xff;
}
bool hasKey(int state, char lock) {
return state & (1 << (tolower(lock) - 'a'));
}
};