diff --git a/upath/_chain.py b/upath/_chain.py index 23a25741..b4304775 100644 --- a/upath/_chain.py +++ b/upath/_chain.py @@ -175,6 +175,18 @@ def _iter_fileobject_protocol_options( yield fileobject, protocol, so +def _merge_protocol_options( + protocol: str, + storage_options: dict[str, Any], + /, +) -> dict[str, Any]: + """merges storage options nested under protocol, flat options take precedence""" + if protocol not in storage_options: + return storage_options + so = storage_options.copy() + return {**so.pop(protocol), **so} + + class FSSpecChainParser: """parse an fsspec chained urlpath""" @@ -260,7 +272,7 @@ def unchain( _iter_fileobject_protocol_options( path_bit if segments else None, protocol or "", - storage_options, + _merge_protocol_options(segments[0].protocol, storage_options), ), ): t_fo, t_proto, t_so = proto_fo_so or (None, "", {}) diff --git a/upath/tests/test_chain.py b/upath/tests/test_chain.py index 8d91be4f..14635491 100644 --- a/upath/tests/test_chain.py +++ b/upath/tests/test_chain.py @@ -1,7 +1,9 @@ +import copy import os from pathlib import Path import pytest +from fsspec.core import url_to_fs from fsspec.implementations.memory import MemoryFileSystem from upath import UPath @@ -133,3 +135,71 @@ def test_chain_parser_roundtrip(urlpath: str): rechained, kw = parser.chain(segments) assert rechained == urlpath assert kw == {} + + +@pytest.mark.parametrize( + "urlpath,storage_options", + [ + ("simplecache::memory://a/b", {"simplecache": {"same_names": True}}), + ("simplecache::memory://a/b", {"same_names": True}), + ( + "simplecache::memory://a/b", + {"simplecache": {"same_names": True}, "same_names": False}, + ), + ( + "simplecache::memory://a/b", + {"simplecache": {"same_names": True}, "memory": {"key": "value"}}, + ), + ("memory://a/b", {"memory": {"key": "value"}}), + ], +) +def test_chaining_upath_storage_options_match_fsspec(urlpath, storage_options): + # fsspec mutates the nested dicts + fs, _ = url_to_fs(urlpath, **copy.deepcopy(storage_options)) + pth = UPath(urlpath, **copy.deepcopy(storage_options)) + assert dict(pth.storage_options) == fs.storage_options + assert pth.fs.storage_options == fs.storage_options + + +def test_chaining_upath_storage_options_not_mutated(): + storage_options = { + "simplecache": {"same_names": True}, + "same_names": False, + "memory": {"key": "value"}, + } + expected = copy.deepcopy(storage_options) + UPath("simplecache::memory://a/b", **storage_options) + assert storage_options == expected + + +def test_read_file_outermost_storage_options(memory_file_urlpath, tmp_path): + pth = UPath( + f"simplecache::{memory_file_urlpath}", + simplecache={"cache_storage": str(tmp_path)}, + ) + assert pth.read_bytes() == b"hello world" + assert list(tmp_path.iterdir()) + + +def test_chain_parser_storage_options(): + parser = FSSpecChainParser() + segments = parser.unchain( + "zip://file.txt::memory:///tmp.zip", + protocol=None, + storage_options={"zip": {"mode": "r"}, "memory": {"key": "value"}}, + ) + assert [s.storage_options for s in segments] == [ + {"mode": "r"}, + {"key": "value"}, + ] + + +def test_chain_parser_roundtrip_storage_options(): + parser = FSSpecChainParser() + segments = parser.unchain( + "zip://file.txt::memory:///tmp.zip", + protocol=None, + storage_options={"zip": {"mode": "r"}, "memory": {"key": "value"}}, + ) + rechained, kw = parser.chain(segments) + assert parser.unchain(rechained, protocol=None, storage_options=kw) == segments