handle network switching

This commit is contained in:
mertalev
2026-02-20 19:40:07 -05:00
parent 13f76295ca
commit 04cfb7fa2f
10 changed files with 80 additions and 20 deletions

View File

@@ -184,7 +184,7 @@ interface NetworkApi {
fun removeCertificate(callback: (Result<Unit>) -> Unit)
fun hasCertificate(): Boolean
fun getClientPointer(): Long
fun setRequestHeaders(headers: Map<String, String>)
fun setRequestHeaders(headers: Map<String, String>, serverUrls: List<String>)
companion object {
/** The codec used by NetworkApi. */
@@ -286,8 +286,9 @@ interface NetworkApi {
channel.setMessageHandler { message, reply ->
val args = message as List<Any?>
val headersArg = args[0] as Map<String, String>
val serverUrlsArg = args[1] as List<String>
val wrapped: List<Any?> = try {
api.setRequestHeaders(headersArg)
api.setRequestHeaders(headersArg, serverUrlsArg)
listOf(null)
} catch (exception: Throwable) {
NetworkPigeonUtils.wrapError(exception)

View File

@@ -79,7 +79,7 @@ private class NetworkApiImpl() : NetworkApi {
return HttpClientManager.getClientPointer()
}
override fun setRequestHeaders(headers: Map<String, String>) {
override fun setRequestHeaders(headers: Map<String, String>, serverUrls: List<String>) {
HttpClientManager.setRequestHeaders(headers)
}
}

View File

@@ -225,7 +225,7 @@ protocol NetworkApi {
func removeCertificate(completion: @escaping (Result<Void, Error>) -> Void)
func hasCertificate() throws -> Bool
func getClientPointer() throws -> Int64
func setRequestHeaders(headers: [String: String]) throws
func setRequestHeaders(headers: [String: String], serverUrls: [String]) throws
}
/// Generated setup class from Pigeon to handle messages through the `binaryMessenger`.
@@ -314,8 +314,9 @@ class NetworkApiSetup {
setRequestHeadersChannel.setMessageHandler { message, reply in
let args = message as! [Any?]
let headersArg = args[0] as! [String: String]
let serverUrlsArg = args[1] as! [String]
do {
try api.setRequestHeaders(headers: headersArg)
try api.setRequestHeaders(headers: headersArg, serverUrls: serverUrlsArg)
reply(wrapResult(nil))
} catch {
reply(wrapError(error))

View File

@@ -59,13 +59,43 @@ class NetworkApiImpl: NetworkApi {
return Int64(Int(bitPattern: pointer))
}
func setRequestHeaders(headers: [String : String]) throws {
var filtered = headers
filtered.removeValue(forKey: "x-immich-user-token") // the session uses cookie auth
var current = URLSessionManager.shared.session.configuration.httpAdditionalHeaders as? [String: String] ?? [:]
current.removeValue(forKey: "User-Agent")
if filtered != current {
UserDefaults.standard.set(filtered, forKey: HEADERS_KEY)
func setRequestHeaders(headers: [String : String], serverUrls: [String]) throws {
var headers = headers
if let token = headers.removeValue(forKey: "x-immich-user-token") {
for serverUrl in serverUrls {
guard let url = URL(string: serverUrl), let domain = url.host else { continue }
let isSecure = serverUrl.hasPrefix("https")
let cookies: [(String, String, Bool)] = [
("immich_access_token", token, true),
("immich_is_authenticated", "true", false),
("immich_auth_type", "password", true),
]
let expiry = Date().addingTimeInterval(400 * 24 * 60 * 60)
for (name, value, httpOnly) in cookies {
var properties: [HTTPCookiePropertyKey: Any] = [
.name: name,
.value: value,
.domain: domain,
.path: "/",
.expires: expiry,
]
if isSecure { properties[.secure] = "TRUE" }
if httpOnly { properties[.init("HttpOnly")] = "TRUE" }
if let cookie = HTTPCookie(properties: properties) {
URLSessionManager.cookieStorage.setCookie(cookie)
}
}
}
} else {
URLSessionManager.cookieStorage.removeCookies(since: .distantPast)
}
guard let groupDefaults = UserDefaults(suiteName: APP_GROUP) else { return }
if headers != groupDefaults.dictionary(forKey: HEADERS_KEY) as? [String: String] {
groupDefaults.set(headers, forKey: HEADERS_KEY)
}
if serverUrls.first != groupDefaults.string(forKey: SERVER_URL_KEY) {
groupDefaults.set(serverUrls.first, forKey: SERVER_URL_KEY)
}
}
}

View File

@@ -3,7 +3,8 @@ import native_video_player
let CLIENT_CERT_LABEL = "app.alextran.immich.client_identity"
let HEADERS_KEY = "immich.request_headers"
private let APP_GROUP = "group.app.immich.share"
let SERVER_URL_KEY = "immich.server_url"
let APP_GROUP = "group.app.immich.share"
/// Manages a shared URLSession with SSL configuration support.
class URLSessionManager: NSObject {
@@ -11,6 +12,7 @@ class URLSessionManager: NSObject {
let session: URLSession
let delegate: URLSessionManagerDelegate
static let cookieStorage = HTTPCookieStorage.sharedCookieStorage(forGroupContainerIdentifier: APP_GROUP)
private let configuration = {
let config = URLSessionConfiguration.default
@@ -25,14 +27,14 @@ class URLSessionManager: NSObject {
directory: cacheDir
)
config.httpCookieStorage = HTTPCookieStorage.sharedCookieStorage(forGroupContainerIdentifier: APP_GROUP)
config.httpCookieStorage = cookieStorage
config.httpMaximumConnectionsPerHost = 64
config.timeoutIntervalForRequest = 60
config.timeoutIntervalForResource = 300
let version = Bundle.main.object(forInfoDictionaryKey: "CFBundleShortVersionString") as? String ?? "unknown"
var headers: [String: String] = ["User-Agent": "Immich_iOS_\(version)"]
if let saved = UserDefaults.standard.dictionary(forKey: HEADERS_KEY) as? [String: String] {
if let saved = UserDefaults(suiteName: APP_GROUP)?.dictionary(forKey: HEADERS_KEY) as? [String: String] {
headers.merge(saved) { _, new in new }
}
config.httpAdditionalHeaders = headers

View File

@@ -100,7 +100,7 @@ class HeaderSettingsPage extends HookConsumerWidget {
var encoded = jsonEncode(headersMap);
await Store.put(StoreKey.customHeaders, encoded);
await networkApi.setRequestHeaders(ApiService.getRequestHeaders());
await networkApi.setRequestHeaders(ApiService.getRequestHeaders(), ApiService.getServerUrls());
}
}

View File

@@ -281,7 +281,7 @@ class NetworkApi {
}
}
Future<void> setRequestHeaders(Map<String, String> headers) async {
Future<void> setRequestHeaders(Map<String, String> headers, List<String> serverUrls) async {
final String pigeonVar_channelName =
'dev.flutter.pigeon.immich_mobile.NetworkApi.setRequestHeaders$pigeonVar_messageChannelSuffix';
final BasicMessageChannel<Object?> pigeonVar_channel = BasicMessageChannel<Object?>(
@@ -289,7 +289,7 @@ class NetworkApi {
pigeonChannelCodec,
binaryMessenger: pigeonVar_binaryMessenger,
);
final Future<Object?> pigeonVar_sendFuture = pigeonVar_channel.send(<Object?>[headers]);
final Future<Object?> pigeonVar_sendFuture = pigeonVar_channel.send(<Object?>[headers, serverUrls]);
final List<Object?>? pigeonVar_replyList = await pigeonVar_sendFuture as List<Object?>?;
if (pigeonVar_replyList == null) {
throw _createConnectionError(pigeonVar_channelName);

View File

@@ -91,6 +91,7 @@ class AuthNotifier extends StateNotifier<AuthState> {
await _widgetService.clearCredentials();
await _authService.logout();
await networkApi.setRequestHeaders(const {}, const []);
await _ref.read(backgroundUploadServiceProvider).cancel();
_ref.read(foregroundUploadServiceProvider).cancel();
} finally {
@@ -125,7 +126,7 @@ class AuthNotifier extends StateNotifier<AuthState> {
Future<bool> saveAuthInfo({required String accessToken}) async {
await _apiService.setAccessToken(accessToken);
await networkApi.setRequestHeaders(ApiService.getRequestHeaders());
await networkApi.setRequestHeaders(ApiService.getRequestHeaders(), ApiService.getServerUrls());
final serverEndpoint = Store.get(StoreKey.serverEndpoint);
final customHeaders = Store.tryGet(StoreKey.customHeaders);

View File

@@ -174,6 +174,31 @@ class ApiService implements Authentication {
}
}
static List<String> getServerUrls() {
final urls = <String>[];
final serverEndpoint = Store.tryGet(StoreKey.serverEndpoint);
if (serverEndpoint != null && serverEndpoint.isNotEmpty) {
urls.add(serverEndpoint);
}
final serverUrl = Store.tryGet(StoreKey.serverUrl);
if (serverUrl != null && serverUrl.isNotEmpty) {
urls.add(serverUrl);
}
final localEndpoint = Store.tryGet(StoreKey.localEndpoint);
if (localEndpoint != null && localEndpoint.isNotEmpty) {
urls.add(localEndpoint);
}
final externalJson = Store.tryGet(StoreKey.externalEndpointList);
if (externalJson != null) {
final List<dynamic> list = jsonDecode(externalJson);
for (final entry in list) {
final url = entry['url'] as String?;
if (url != null && url.isNotEmpty) urls.add(url);
}
}
return urls;
}
static Map<String, String> getRequestHeaders() {
var accessToken = Store.get(StoreKey.accessToken, "");
var customHeadersStr = Store.get(StoreKey.customHeaders, "");

View File

@@ -43,5 +43,5 @@ abstract class NetworkApi {
int getClientPointer();
void setRequestHeaders(Map<String, String> headers);
void setRequestHeaders(Map<String, String> headers, List<String> serverUrls);
}