Skip to content
Open
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
16 changes: 8 additions & 8 deletions mmengine/fileio/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,9 +8,9 @@
from .io import (copy_if_symlink_fails, copyfile, copyfile_from_local,
copyfile_to_local, copytree, copytree_from_local,
copytree_to_local, dump, exists, generate_presigned_url, get,
get_file_backend, get_local_path, get_text, isdir, isfile,
join_path, list_dir_or_file, load, put, put_text, remove,
rmtree)
get_file_backend, get_local_path, get_text, glob, iglob,
isdir, isfile, join_path, list_dir_or_file, load, put,
put_text, remove, rmtree)
from .parse import dict_from_file, list_from_file

__all__ = [
Expand All @@ -19,9 +19,9 @@
'copy_if_symlink_fails', 'copyfile', 'copyfile_from_local',
'copyfile_to_local', 'copytree', 'copytree_from_local',
'copytree_to_local', 'exists', 'generate_presigned_url', 'get',
'get_file_backend', 'get_local_path', 'get_text', 'isdir', 'isfile',
'join_path', 'list_dir_or_file', 'put', 'put_text', 'remove', 'rmtree',
'load', 'dump', 'register_handler', 'BaseFileHandler', 'JsonHandler',
'PickleHandler', 'YamlHandler', 'list_from_file', 'dict_from_file',
'register_backend'
'get_file_backend', 'get_local_path', 'get_text', 'glob', 'iglob', 'isdir',
'isfile', 'join_path', 'list_dir_or_file', 'put', 'put_text', 'remove',
'rmtree', 'load', 'dump', 'register_handler', 'BaseFileHandler',
'JsonHandler', 'PickleHandler', 'YamlHandler', 'list_from_file',
'dict_from_file', 'register_backend'
]
115 changes: 115 additions & 0 deletions mmengine/fileio/io.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
>>> # Directory call unified I/O functions
>>> fileio.get('s3://path/of/your/file')
"""
import fnmatch
import json
import warnings
from contextlib import contextmanager
Expand Down Expand Up @@ -707,6 +708,120 @@ def copy_if_symlink_fails(
return backend.copy_if_symlink_fails(src, dst)


def _split_glob_pattern(pattern: str) -> Tuple[str, str]:
"""分离通配模式的静态目录和相对模式。"""
magic_index = next(
(index for index, char in enumerate(pattern) if char in '*?['), -1)
if magic_index < 0:
return '', pattern

separator_index = max(
pattern.rfind('/', 0, magic_index),
pattern.rfind('\\', 0, magic_index))
if separator_index < 0:
return '', pattern
if separator_index == 0:
return pattern[:1], pattern[1:]

root = pattern[:separator_index]
# 保留 Windows 驱动器根目录(例如 ``C:\\``)的分隔符。
if root.endswith(':'):
root += pattern[separator_index]
return root, pattern[separator_index + 1:]


def _split_glob_path(path: str) -> list:
"""按统一的路径分隔符拆分相对路径。"""
return [
part for part in path.replace('\\', '/').split('/')
if part not in ('', '.')
]


def _match_glob_path(path: str, pattern: str, recursive: bool) -> bool:
"""匹配相对路径,避免普通 ``*`` 跨越目录分隔符。"""
path_parts = _split_glob_path(path)
pattern_parts = _split_glob_path(pattern)

def _match(path_index: int, pattern_index: int) -> bool:
if pattern_index == len(pattern_parts):
return path_index == len(path_parts)

current_pattern = pattern_parts[pattern_index]
if recursive and current_pattern == '**':
skip_current = _match(path_index, pattern_index + 1)
consume_current = False
if path_index < len(path_parts):
consume_current = _match(path_index + 1, pattern_index)
return skip_current or consume_current

if path_index >= len(path_parts):
return False
if not fnmatch.fnmatchcase(path_parts[path_index], current_pattern):
return False
return _match(path_index + 1, pattern_index + 1)

return _match(0, 0)


def iglob(
pattern: Union[str, Path],
*,
recursive: bool = False,
backend_args: Optional[dict] = None,
) -> Iterator[str]:
"""返回匹配文件后端路径模式的迭代器。

Args:
pattern (str or Path): 包含 ``*``、``?`` 或字符组模式的路径。
recursive (bool): 是否让独立的 ``**`` 匹配多级目录。默认为 ``False``。
backend_args (dict, optional): 初始化文件后端的参数。默认为 ``None``。

Yields:
str: 匹配到的文件或目录路径。
"""
pattern = str(pattern)
if not any(char in pattern for char in '*?['):
if exists(pattern, backend_args=backend_args):
yield pattern
return

root, relative_pattern = _split_glob_pattern(pattern)
search_root = root or '.'
try:
entries = list_dir_or_file(
search_root,
list_dir=True,
list_file=True,
recursive=True,
backend_args=backend_args)
for relative_path in entries:
relative_path = str(relative_path)
if not _match_glob_path(relative_path, relative_pattern,
recursive):
continue
if root:
yield str(
join_path(root, relative_path, backend_args=backend_args))
else:
yield relative_path
except FileNotFoundError:
return


def glob(
pattern: Union[str, Path],
*,
recursive: bool = False,
backend_args: Optional[dict] = None,
) -> list:
"""返回匹配文件后端路径模式的列表。

参数与 :func:`iglob` 相同;结果按后端枚举顺序返回。
"""
return list(iglob(pattern, recursive=recursive, backend_args=backend_args))


def list_dir_or_file(
dir_path: Union[str, Path],
list_dir: bool = True,
Expand Down
49 changes: 49 additions & 0 deletions tests/test_fileio/test_io.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
import pytest

import mmengine.fileio as fileio
import mmengine.fileio.io as fileio_io

sys.modules['petrel_client'] = MagicMock()
sys.modules['petrel_client.client'] = MagicMock()
Expand Down Expand Up @@ -533,3 +534,51 @@ def test_list_dir_or_file():
osp.join('dir2', 'dir3', 'text4.txt'),
osp.join('dir2', 'img.jpg'), 'text1.txt', 'text2.txt'
}


def test_glob_and_iglob():
# 使用本地后端验证单层和递归模式。
with build_temporary_directory() as tmp_dir:
assert set(fileio.glob(osp.join(tmp_dir, '*.txt'))) == {
osp.join(tmp_dir, 'text1.txt'),
osp.join(tmp_dir, 'text2.txt')
}
assert set(fileio.glob(
osp.join(tmp_dir, '*',
'*.txt'))) == {osp.join(tmp_dir, 'dir1', 'text3.txt')}
expected = {
osp.join(tmp_dir, 'text1.txt'),
osp.join(tmp_dir, 'text2.txt'),
osp.join(tmp_dir, 'dir1', 'text3.txt'),
osp.join(tmp_dir, 'dir2', 'dir3', 'text4.txt')
}
assert set(
fileio.iglob(osp.join(tmp_dir, '**', '*.txt'),
recursive=True)) == expected
assert fileio.glob(
osp.join(tmp_dir, 'missing', '*.txt'),
backend_args={'backend': 'local'}) == []


def test_iglob_passes_backend_args():
with patch.object(
fileio_io,
'list_dir_or_file',
return_value=iter([osp.join('dir', 'a.txt'),
'b.txt'])) as list_mock:
result = list(
fileio.iglob(
'virtual/**/*.txt',
recursive=True,
backend_args={'backend': 'local'}))

assert result == [
osp.join('virtual', 'dir', 'a.txt'),
osp.join('virtual', 'b.txt')
]
list_mock.assert_called_once_with(
'virtual',
list_dir=True,
list_file=True,
recursive=True,
backend_args={'backend': 'local'})