From 5304ed2949dd4fe947469f51572a91f89d815bf8 Mon Sep 17 00:00:00 2001 From: bifrost0x Date: Mon, 14 Sep 2026 12:29:20 +0200 Subject: [PATCH 1/2] Fix file metadata escaping and bound key and SMB work --- .env.example | 3 + app/smb_protocol.py | 67 +++++++---- app/socket_events.py | 45 +++++++- config.py | 3 + docs/wiki/Configuration-Reference.md | 7 ++ static/js/sftp-file-manager.js | 10 +- tests/e2e/file-workspace.spec.js | 45 ++++++++ tests/js/sftp-transfer-queue.test.js | 7 -- tests/test_key_socket_events.py | 164 ++++++++++++++++++++++++++- tests/test_smb_pool.py | 15 ++- tests/test_smb_protocol_contract.py | 138 ++++++++++++++++++++-- 11 files changed, 454 insertions(+), 50 deletions(-) diff --git a/.env.example b/.env.example index 598bd331..646cfa5d 100644 --- a/.env.example +++ b/.env.example @@ -218,6 +218,9 @@ RATELIMIT_REAUTH=5 per minute SSH_CONNECT_RATELIMIT=10 per minute # Per-user upload/replacement rate for encrypted SSH keys. SSH_KEY_WRITE_RATELIMIT=30 per minute +# Shared per-user budget for key listings and rename/replace/delete refreshes. +# Checked before mutations so accepted changes still return an updated list. +SSH_KEY_LIST_RATELIMIT=30 per minute # Per-account encrypted SSH-key limits. Existing over-limit stores remain # readable and can still be renamed, deleted, or replaced with smaller keys. SSH_KEY_MAX_RECORDS=100 diff --git a/app/smb_protocol.py b/app/smb_protocol.py index 2b26f7e1..66986a1f 100644 --- a/app/smb_protocol.py +++ b/app/smb_protocol.py @@ -29,6 +29,8 @@ DirectoryNotEmpty, IOTimeout, LogonFailure, + NoMoreFiles, + NoSuchFile, ObjectNameCollision, ObjectNameNotFound, ObjectPathNotFound, @@ -50,6 +52,7 @@ FileAttributes, FileInformationClass, FilePipePrinterAccessMask, + QueryDirectoryFlags, ) from smbprotocol.session import Session, SessionFlags @@ -242,21 +245,51 @@ def _entry_from_directory_info(raw_info): ) +def _directory_entries(raw, pattern): + """Enumerate bounded pages on the already-verified handle, without DFS. + + Keep progress checks below the filtering boundary: the high-level client + iterator can request another page before returning control to our caller. + Each directory may contain one '.' and one '..'; neither consumes the + caller's budget for ordinary entries. + """ + flags = QueryDirectoryFlags.SMB2_RESTART_SCANS + special_entries = set() + while True: + try: + page = raw.fd.query_directory( + pattern, + FileInformationClass.FILE_ID_FULL_DIRECTORY_INFORMATION, + flags=flags, + ) + except (NoMoreFiles, NoSuchFile): + return + if not page: + raise SMBProtocolError('OPERATION_FAILED') + flags = 0 + for raw_info in page: + entry = _entry_from_directory_info(raw_info) + if entry.name in {'.', '..'}: + if entry.name in special_entries: + raise SMBProtocolError('OPERATION_FAILED') + special_entries.add(entry.name) + continue + yield entry + + def _query_exact_child(directory, name): """Resolve one exact child through an already-open parent handle.""" matches = [] - for raw_info in directory.query_directory( - name, - FileInformationClass.FILE_ID_FULL_DIRECTORY_INFORMATION, - ): - entry = _entry_from_directory_info(raw_info) - if entry.name in {'.', '..'}: - continue - if entry.name.casefold() != name.casefold(): - raise SMBProtocolError('CONFLICT') - matches.append(entry) - if len(matches) > 1: - raise SMBProtocolError('CONFLICT') + entries = _directory_entries(directory, name) + try: + for entry in entries: + if entry.name.casefold() != name.casefold(): + raise SMBProtocolError('CONFLICT') + matches.append(entry) + if len(matches) > 1: + raise SMBProtocolError('CONFLICT') + finally: + entries.close() if not matches: raise SMBProtocolError('NOT_FOUND') return matches[0] @@ -795,10 +828,7 @@ def __init__( None if connection_kwargs is None else dict(connection_kwargs) ) try: - self._iterator = raw.query_directory( - '*', - FileInformationClass.FILE_ID_FULL_DIRECTORY_INFORMATION, - ) + self._iterator = _directory_entries(raw, '*') except BaseException: try: raw.close() @@ -816,10 +846,7 @@ def __next__(self): if self.closed: raise StopIteration try: - while True: - entry = _entry_from_directory_info(next(self._iterator)) - if entry.name not in {'.', '..'}: - return entry + return next(self._iterator) except StopIteration: self.close() raise diff --git a/app/socket_events.py b/app/socket_events.py index e44ddbcf..f5284049 100644 --- a/app/socket_events.py +++ b/app/socket_events.py @@ -2573,13 +2573,39 @@ def handle_delete_jump_host(data, current_user=None): log_error("Failed to delete jump host", error=str(e)) emit('error', {'error': 'Failed to delete jump host'}) +def _key_summary_rate_limit(current_user): + if check_socket_rate_limit( + current_user.id, + 'ssh_key_list', + config.RATELIMIT_SSH_KEY_LIST, + ): + return _key_mutation_error( + 'Too many SSH key list requests. Please wait a moment.' + ) + return None + + +def _emit_key_summaries(current_user): + """Emit fresh usability after the caller reserves one summary operation.""" + try: + keys = key_manager.load_key_summaries(current_user.id) + emit('keys_list', {'keys': keys}) + except StorageCorruptionError as error: + return _emit_storage_error(error, current_user) + except Exception as e: + log_error("Failed to load keys", error=str(e)) + emit('error', {'error': 'Failed to load keys'}) + + @socketio.on('list_keys') @socket_login_required def handle_list_keys(current_user=None): """Return list of stored SSH keys for this user.""" try: - keys = key_manager.load_key_summaries(current_user.id) - emit('keys_list', {'keys': keys}) + limited = _key_summary_rate_limit(current_user) + if limited: + return limited + return _emit_key_summaries(current_user) except StorageCorruptionError as error: return _emit_storage_error(error, current_user) except Exception as e: @@ -2639,6 +2665,9 @@ def handle_upload_key(data, current_user=None): def handle_rename_key(data, current_user=None): """Rename one owned SSH key without exposing its encrypted contents.""" try: + limited = _key_summary_rate_limit(current_user) + if limited: + return limited data = data if isinstance(data, dict) else {} result, error = key_manager.rename_key( current_user.id, @@ -2656,7 +2685,7 @@ def handle_rename_key(data, current_user=None): ) payload = {'success': True, 'key': result['key']} emit('key_renamed', payload) - handle_list_keys(current_user=current_user) + _emit_key_summaries(current_user) return payload except StorageCorruptionError as error: return _emit_storage_error(error, current_user) @@ -2669,6 +2698,9 @@ def handle_rename_key(data, current_user=None): def handle_replace_key(data, current_user=None): """Replace one owned SSH key without changing its stable identity.""" try: + limited = _key_summary_rate_limit(current_user) + if limited: + return limited data = data if isinstance(data, dict) else {} key_id = data.get('key_id') key_content = data.get('key_content') @@ -2715,7 +2747,7 @@ def handle_replace_key(data, current_user=None): ) payload = {'success': True, 'key': key} emit('key_replaced', payload) - handle_list_keys(current_user=current_user) + _emit_key_summaries(current_user) return payload except StorageCorruptionError as error: return _emit_storage_error(error, current_user) @@ -2727,6 +2759,9 @@ def handle_replace_key(data, current_user=None): def handle_delete_key(data, current_user=None): """Delete an SSH key for this user.""" try: + limited = _key_summary_rate_limit(current_user) + if limited: + return limited key_id = data.get('key_id') if not key_id: emit('error', {'error': 'Key ID required'}) @@ -2736,7 +2771,7 @@ def handle_delete_key(data, current_user=None): if success: log_key_delete(current_user.username, key_id, request.remote_addr) emit('key_deleted', {'key_id': key_id}) - handle_list_keys(current_user=current_user) + _emit_key_summaries(current_user) else: emit('error', {'error': 'Failed to delete key'}) diff --git a/config.py b/config.py index b13dd708..ebac86b5 100644 --- a/config.py +++ b/config.py @@ -754,6 +754,9 @@ def _validate_quota_pair(kind, global_limit, per_user_limit, fair_slots): RATELIMIT_SSH_KEY_WRITE = os.environ.get( 'SSH_KEY_WRITE_RATELIMIT', '30 per minute' ) +RATELIMIT_SSH_KEY_LIST = os.environ.get( + 'SSH_KEY_LIST_RATELIMIT', '30 per minute' +) RATELIMIT_COMMAND_MUTATION = os.environ.get( 'COMMAND_MUTATION_RATELIMIT', '60 per minute', diff --git a/docs/wiki/Configuration-Reference.md b/docs/wiki/Configuration-Reference.md index 17046bcd..2443b7dd 100644 --- a/docs/wiki/Configuration-Reference.md +++ b/docs/wiki/Configuration-Reference.md @@ -98,8 +98,15 @@ Connection, transfer, background-work, and thread limits form one capacity model | `RATELIMIT_REAUTH` | `5 per minute` | | `SSH_CONNECT_RATELIMIT` | `10 per minute` | | `SSH_KEY_WRITE_RATELIMIT` | `30 per minute` | +| `SSH_KEY_LIST_RATELIMIT` | `30 per minute` | | `CONNECTION_MUTATION_RATELIMIT` | `60 per minute` | +Key listings and the refresh after renaming, replacing, or deleting a key share +`SSH_KEY_LIST_RATELIMIT` per user across browser connections. Admission is checked +before the mutation; when exhausted, the change is rejected without altering the +key. Accepted changes retain their acknowledgement and updated key list. Usability +is still checked against the current key files, without caching decrypted keys. + ## SSH key and live-output limits | Variable | Default | diff --git a/static/js/sftp-file-manager.js b/static/js/sftp-file-manager.js index 5535f200..dca9f7c7 100644 --- a/static/js/sftp-file-manager.js +++ b/static/js/sftp-file-manager.js @@ -6008,9 +6008,13 @@ class SFTPFileManager { } escapeHtml(text) { - const div = document.createElement('div'); - div.textContent = text; - return div.innerHTML; + // Used in both text and quoted attributes in the file workspace. + return String(text ?? '') + .replaceAll('&', '&') + .replaceAll('<', '<') + .replaceAll('>', '>') + .replaceAll('"', '"') + .replaceAll("'", '''); } showUploadProgress(batch = this.currentUploadBatch) { diff --git a/tests/e2e/file-workspace.spec.js b/tests/e2e/file-workspace.spec.js index f2c62cda..92ee8323 100644 --- a/tests/e2e/file-workspace.spec.js +++ b/tests/e2e/file-workspace.spec.js @@ -75,6 +75,51 @@ async function openWorkspaceWithSources(page) { }); } +for (const kind of ['sftp', 'smb']) { + test(`${kind} filenames preserve text, checkbox attributes and sorted selection`, async ({ page }) => { + await openWorkspaceWithSources(page); + // Ordinary punctuation and entity-like text, without executable markup. + const filename = 'Über "quotes" and \'apostrophes\' & ".txt'; + await page.evaluate(({ kind, filename }) => { + const manager = window.sftpFileManager; + manager.closeSourceLauncher(); + const sourceId = kind === 'sftp' + ? 'sftp-session:workspace-source' + : `smb-quick:${'a'.repeat(32)}`; + Object.assign(manager.panes.left, manager.createEmptyPaneState(), { + source: { + sourceId, kind, label: 'Punctuation files', + capabilities: ['list', 'read'], security: {}, access: {}, + }, + path: '/', + files: [ + { name: filename, is_dir: false, size: 12 }, + { name: 'Folder', is_dir: true, size: 0 }, + { name: 'alpha.txt', is_dir: false, size: 1 }, + ], + }); + manager.renderPane('left'); + }, { kind, filename }); + + const rows = page.locator('#fmLeftList .fm-file-item'); + await expect(rows).toHaveCount(3); + await expect(rows.first().locator('.fm-file-name')).toHaveText('Folder'); + const row = page.locator('#fmLeftList .fm-file-item[data-index="0"]'); + await expect(row.locator('.fm-file-name')).toHaveText(filename); + const checkbox = row.getByRole('checkbox'); + await expect(checkbox).toHaveAttribute('aria-label', `Select: ${filename}`); + expect(await checkbox.evaluate(element => element.getAttributeNames().sort())).toEqual([ + 'aria-checked', 'aria-label', 'class', 'role', 'type', + ]); + await checkbox.click(); + await expect(checkbox).toHaveAttribute('aria-checked', 'true'); + expect(await page.evaluate(() => [...window.sftpFileManager.panes.left.selected])).toEqual([0]); + await checkbox.click(); + await expect(checkbox).toHaveAttribute('aria-checked', 'false'); + await assertNoExternalRequests(page); + }); +} + test('source-first workspace preserves panes and exposes only functional SFTP actions', async ({ page }) => { await openWorkspaceWithSources(page); diff --git a/tests/js/sftp-transfer-queue.test.js b/tests/js/sftp-transfer-queue.test.js index b9e02f37..66d6132e 100644 --- a/tests/js/sftp-transfer-queue.test.js +++ b/tests/js/sftp-transfer-queue.test.js @@ -4064,13 +4064,6 @@ test('remote filenames never enter attributes even when they contain quote and e }, workspace: { layout: 'single' }, displayMode: 'embedded', - escapeHtml(value) { - return String(value) - .replaceAll('&', '&') - .replaceAll('"', '"') - .replaceAll('<', '<') - .replaceAll('>', '>'); - }, updatePaneStatus() {}, t(_key, fallback) { return fallback; }, }); diff --git a/tests/test_key_socket_events.py b/tests/test_key_socket_events.py index 7ea63213..0b51aeb5 100644 --- a/tests/test_key_socket_events.py +++ b/tests/test_key_socket_events.py @@ -21,7 +21,10 @@ def create_socket_user(app, username): return user.id, sid -def call_socket_handler(app, monkeypatch, handler, sid, payload): +_NO_PAYLOAD = object() + + +def call_socket_handler(app, monkeypatch, handler, sid, payload=_NO_PAYLOAD): import app.socket_events as socket_events emitted = [] @@ -32,7 +35,9 @@ def call_socket_handler(app, monkeypatch, handler, sid, payload): ) with app.test_request_context('/socket.io'): request.sid = sid - acknowledgement = handler(payload) + acknowledgement = ( + handler() if payload is _NO_PAYLOAD else handler(payload) + ) return acknowledgement, emitted @@ -371,3 +376,158 @@ def test_key_replace_audit_sanitizes_values(monkeypatch): assert messages[0].startswith('KEY_REPLACE_FAILED | ') assert '\n' not in messages[0] assert '\r' not in messages[0] + + +@pytest.mark.parametrize('operation', ('rename', 'replace', 'delete')) +def test_key_mutation_reserves_refresh_before_changing_storage( + app, monkeypatch, rsa_private_key_pem, operation, +): + from app import key_manager + import app.socket_events as socket_events + + user_id, sid = create_socket_user(app, 'summary_limited') + monkeypatch.setattr(socket_events.config, 'RATELIMIT_SSH_KEY_LIST', '1 per minute') + with app.app_context(): + key, error = key_manager.save_key(user_id, 'Original', rsa_private_key_pem) + assert error is None + before = key_manager.get_user_keys_file(user_id).read_bytes() + + calls = [] + original_loader = key_manager.load_key_summaries + + def summaries(user_id): + calls.append(user_id) + return original_loader(user_id) + + monkeypatch.setattr(key_manager, 'load_key_summaries', summaries) + _, emitted = call_socket_handler( + app, monkeypatch, socket_events.handle_list_keys, sid, + ) + assert emitted[0][0] == 'keys_list' + assert emitted[0][1]['keys'][0]['usable'] is True + + denied, emitted = call_socket_handler( + app, monkeypatch, getattr(socket_events, f'handle_{operation}_key'), sid, + {'key_id': key['id'], 'name': 'Changed', 'key_content': rsa_private_key_pem}, + ) + assert denied['success'] is False + assert emitted == [('error', {'error': denied['error']})] + assert calls == [user_id] + with app.app_context(): + assert key_manager.get_user_keys_file(user_id).read_bytes() == before + content, error = key_manager.read_key_content(user_id, key['id']) + assert error is None + assert content == rsa_private_key_pem + + +@pytest.mark.parametrize('operation', ('rename', 'replace', 'delete')) +def test_key_mutation_last_admission_still_delivers_fresh_list( + app, monkeypatch, rsa_private_key_pem, operation, +): + from app import key_manager + import app.socket_events as socket_events + + user_id, sid = create_socket_user(app, 'summary_last_slot') + monkeypatch.setattr(socket_events.config, 'RATELIMIT_SSH_KEY_LIST', '1 per minute') + with app.app_context(): + key, error = key_manager.save_key(user_id, 'Original', rsa_private_key_pem) + assert error is None + + _, emitted = call_socket_handler( + app, monkeypatch, getattr(socket_events, f'handle_{operation}_key'), sid, + {'key_id': key['id'], 'name': 'Changed', 'key_content': rsa_private_key_pem}, + ) + expected_event = {'rename': 'key_renamed', 'replace': 'key_replaced', 'delete': 'key_deleted'}[operation] + assert [event for event, _ in emitted] == [expected_event, 'keys_list'] + listed = emitted[-1][1]['keys'] + if operation == 'delete': + assert listed == [] + else: + assert listed[0]['id'] == key['id'] + assert listed[0]['usable'] is True + assert listed[0]['name'] == ('Changed' if operation == 'rename' else 'Original') + denied, emitted = call_socket_handler( + app, monkeypatch, socket_events.handle_list_keys, sid, + ) + assert denied['success'] is False + assert [event for event, _ in emitted] == ['error'] + + +def test_key_summary_budget_is_shared_between_sockets_but_not_users(app, monkeypatch): + from app.auth import register_socket_session + from app.models import db + import app.socket_events as socket_events + + user_id, sid = create_socket_user(app, 'summary_owner') + other_id, other_sid = create_socket_user(app, 'summary_other') + with app.app_context(): + register_socket_session(user_id, 'summary-second-socket') + db.session.commit() + monkeypatch.setattr(socket_events.config, 'RATELIMIT_SSH_KEY_LIST', '1 per minute') + calls = [] + + def summaries(user_id): + calls.append(user_id) + return [] + + monkeypatch.setattr(socket_events.key_manager, 'load_key_summaries', summaries) + call_socket_handler(app, monkeypatch, socket_events.handle_list_keys, sid) + denied, emitted = call_socket_handler( + app, monkeypatch, socket_events.handle_list_keys, 'summary-second-socket', + ) + assert denied['success'] is False + assert [event for event, _ in emitted] == ['error'] + _, emitted = call_socket_handler(app, monkeypatch, socket_events.handle_list_keys, other_sid) + assert emitted == [('keys_list', {'keys': []})] + assert calls == [user_id, other_id] + + +def test_key_listing_preserves_storage_error_acknowledgement(app, monkeypatch): + from app import key_manager + import app.socket_events as socket_events + + user_id, sid = create_socket_user(app, 'summary_corrupt') + with app.app_context(): + path = key_manager.get_user_keys_file(user_id) + path.write_text('{broken', encoding='utf-8') + before = path.read_bytes() + acknowledgement, emitted = call_socket_handler( + app, monkeypatch, socket_events.handle_list_keys, sid, + ) + assert acknowledgement['success'] is False + assert acknowledgement['code'] == 'storage_error' + assert emitted == [('error', acknowledgement)] + assert path.read_bytes() == before + + +@pytest.mark.parametrize('operation', ('rename', 'replace', 'delete')) +def test_key_refresh_read_failure_does_not_reverse_successful_mutation( + app, monkeypatch, rsa_private_key_pem, operation, +): + from app import key_manager + import app.socket_events as socket_events + + user_id, sid = create_socket_user(app, 'summary_read_failure') + with app.app_context(): + key, error = key_manager.save_key(user_id, 'Original', rsa_private_key_pem) + assert error is None + + def failed_read(_user_id): + raise OSError('read failed') + + monkeypatch.setattr(key_manager, 'load_key_summaries', failed_read) + acknowledgement, emitted = call_socket_handler( + app, monkeypatch, getattr(socket_events, f'handle_{operation}_key'), sid, + {'key_id': key['id'], 'name': 'Changed', 'key_content': rsa_private_key_pem}, + ) + if operation != 'delete': + assert acknowledgement['success'] is True + expected_event = {'rename': 'key_renamed', 'replace': 'key_replaced', 'delete': 'key_deleted'}[operation] + assert [event for event, _ in emitted] == [expected_event, 'error'] + assert emitted[-1][1] == {'error': 'Failed to load keys'} + with app.app_context(): + keys = key_manager.load_keys(user_id) + if operation == 'delete': + assert keys == [] + else: + assert keys[0]['name'] == ('Changed' if operation == 'rename' else 'Original') diff --git a/tests/test_smb_pool.py b/tests/test_smb_pool.py index 4d85abff..776704aa 100644 --- a/tests/test_smb_pool.py +++ b/tests/test_smb_pool.py @@ -251,8 +251,13 @@ def test_share_root_is_validated_before_descriptor_is_published(): assert pool.get_source(descriptor.source_id, '1') is not None -@pytest.mark.parametrize('code', ('SHARE_UNAVAILABLE', 'PERMISSION_DENIED')) -def test_share_root_failure_closes_session_releases_quota_and_is_not_published(code): +@pytest.mark.parametrize('code', ( + 'SHARE_UNAVAILABLE', 'PERMISSION_DENIED', 'OPERATION_FAILED', +)) +@pytest.mark.parametrize('failed_lane', (1, 2)) +def test_share_root_failure_closes_session_releases_quota_and_is_not_published( + code, failed_lane, +): pool, protocol, quota, _events = _pool() original_connect = protocol.connect @@ -262,7 +267,8 @@ def connect(**kwargs): failure.diagnostic_phase = 'file_operation' failure.diagnostic_exception_type = 'SMBOSError' failure.diagnostic_nt_status = '0xC0000022' - session.inspect_error = failure + if len(protocol.sessions) == failed_lane: + session.inspect_error = failure return session protocol.connect = connect @@ -274,7 +280,8 @@ def connect(**kwargs): assert exc.value.diagnostic_phase == 'share_access' assert exc.value.diagnostic_exception_type == 'SMBOSError' assert exc.value.diagnostic_nt_status == '0xC0000022' - assert protocol.sessions[0].closed == 1 + assert len(protocol.sessions) == failed_lane + assert all(session.closed == 1 for session in protocol.sessions) assert quota.reservations[0].released is True assert pool.source_count == 0 diff --git a/tests/test_smb_protocol_contract.py b/tests/test_smb_protocol_contract.py index 3f37e939..5fb18c5d 100644 --- a/tests/test_smb_protocol_contract.py +++ b/tests/test_smb_protocol_contract.py @@ -1780,17 +1780,127 @@ def test_explicit_empty_expected_identity_chain_rejects_non_root_path( assert error.value.public_code == 'OPERATION_FAILED' +class _DirectoryPages: + """Finite decoded pages; no network or protocol payload generation.""" + + def __init__(self, pages, terminal=None): + from app import smb_protocol + + self.fd = self + self.pages = iter(pages) + self.terminal = terminal or smb_protocol.NoMoreFiles + self.calls = [] + self.closed = False + + def query_directory(self, pattern, info_class, *, flags): + self.calls.append((pattern, info_class, flags)) + try: + return next(self.pages) + except StopIteration: + raise self.terminal(None) + + def close(self): + self.closed = True + + +@pytest.mark.parametrize('terminal_name', ('NoMoreFiles', 'NoSuchFile')) +@pytest.mark.parametrize('names', ([], ['.', '..'], ['.', '..', 'first', 'second'], ['first', 'second'])) +def test_directory_pages_preserve_entries_flags_and_normal_completion( + monkeypatch, terminal_name, names, +): + from app import smb_protocol + + monkeypatch.setattr(smb_protocol, '_entry_from_directory_info', lambda entry: entry) + raw = _DirectoryPages( + [[_object_info(name, index + 1, directory=True)] for index, name in enumerate(names)], + terminal=getattr(smb_protocol, terminal_name), + ) + iterator = smb_protocol._VerifiedDirectoryIterator(raw, _object_info('', 10, directory=True)) + assert [entry.name for entry in iterator] == [name for name in names if name not in {'.', '..'}] + assert [call[2] for call in raw.calls] == [ + smb_protocol.QueryDirectoryFlags.SMB2_RESTART_SCANS, *([0] * len(names)), + ] + assert all(call[0] == '*' for call in raw.calls) + assert all(call[1] == smb_protocol.FileInformationClass.FILE_ID_FULL_DIRECTORY_INFORMATION for call in raw.calls) + assert iterator.closed is True + assert raw.closed is True + + +@pytest.mark.parametrize('pages', ((['.'], ['.']), (['..'], ['..']), (['.', '..'], ['.']), ([],))) +@pytest.mark.parametrize('exact_child', (False, True)) +def test_directory_progress_rejects_repeated_special_entries_and_empty_pages( + monkeypatch, pages, exact_child, +): + from app import smb_protocol + + monkeypatch.setattr(smb_protocol, '_entry_from_directory_info', lambda entry: entry) + raw = _DirectoryPages([ + [_object_info(name, 1, directory=True) for name in page] for page in pages + ]) + iterator = None + with pytest.raises(SMBProtocolError) as error: + if exact_child: + smb_protocol._query_exact_child(raw, 'normal.txt') + else: + iterator = smb_protocol._VerifiedDirectoryIterator(raw, _object_info('', 10, directory=True)) + next(iterator) + assert error.value.public_code == 'OPERATION_FAILED' + assert len(raw.calls) == len(pages) + if iterator is not None: + assert iterator.closed is True + assert raw.closed is True + else: + # Exact lookup borrows the parent's handle; verified-path cleanup owns it. + assert raw.closed is False + + +@pytest.mark.parametrize('names, expected', ( + (['.', '..', 'Normal.txt'], None), + (['.', '..'], 'NOT_FOUND'), + (['Normal.txt', 'normal.txt'], 'CONFLICT'), + (['other.txt'], 'CONFLICT'), +)) +def test_exact_child_preserves_matching_and_conflict_semantics(monkeypatch, names, expected): + from app import smb_protocol + + monkeypatch.setattr(smb_protocol, '_entry_from_directory_info', lambda entry: entry) + raw = _DirectoryPages([[_object_info(name, i + 1) for i, name in enumerate(names)]]) + if expected is not None: + with pytest.raises(SMBProtocolError) as error: + smb_protocol._query_exact_child(raw, 'normal.txt') + assert error.value.public_code == expected + else: + assert smb_protocol._query_exact_child(raw, 'normal.txt').name == 'Normal.txt' + assert raw.closed is False + + +def test_directory_special_entry_allowance_is_per_iterator(monkeypatch): + from app import smb_protocol + + monkeypatch.setattr(smb_protocol, '_entry_from_directory_info', lambda entry: entry) + for _ in range(2): + raw = _DirectoryPages([[_object_info(name, 1, directory=True) for name in ('.', '..')]]) + iterator = smb_protocol._VerifiedDirectoryIterator(raw, _object_info('', 10, directory=True)) + assert list(iterator) == [] + assert raw.closed is True + + def test_exact_child_no_result_fails_closed(): from app import smb_protocol class Directory: - def query_directory(self, pattern, info_class): + @property + def fd(self): + return self + + def query_directory(self, pattern, info_class, *, flags): assert pattern == 'missing.txt' assert info_class == ( smb_protocol.FileInformationClass .FILE_ID_FULL_DIRECTORY_INFORMATION ) - return iter(()) + assert flags == smb_protocol.QueryDirectoryFlags.SMB2_RESTART_SCANS + raise smb_protocol.NoSuchFile(None) with pytest.raises(SMBProtocolError) as error: smb_protocol._query_exact_child(Directory(), 'missing.txt') @@ -1798,13 +1908,17 @@ def query_directory(self, pattern, info_class): assert error.value.public_code == 'NOT_FOUND' -def test_verified_iterator_closes_raw_when_enumeration_setup_fails(): +def test_verified_iterator_closes_raw_when_first_page_fails(): from app import smb_protocol class Raw: closed = False - def query_directory(self, *_args): + @property + def fd(self): + return self + + def query_directory(self, *_args, **_kwargs): raise RuntimeError('query setup failed') def close(self): @@ -1812,10 +1926,11 @@ def close(self): raw = Raw() with pytest.raises(RuntimeError, match='query setup failed'): - smb_protocol._VerifiedDirectoryIterator( + iterator = smb_protocol._VerifiedDirectoryIterator( raw, _object_info('', 1, directory=True), ) + next(iterator) assert raw.closed is True @@ -1831,9 +1946,13 @@ class Raw: def __init__(self, entries): self.entries = entries self.closed = False + self.fd = self - def query_directory(self, *_args): - return iter(self.entries) + def query_directory(self, *_args, **_kwargs): + if not self.entries: + raise smb_protocol.NoMoreFiles(None) + entries, self.entries = self.entries, [] + return entries def close(self): self.closed = True @@ -1893,9 +2012,10 @@ def test_verified_iterator_child_identity_mismatch_closes_child( class Raw: def __init__(self): self.closed = False + self.fd = self - def query_directory(self, *_args): - return iter(()) + def query_directory(self, *_args, **_kwargs): + raise smb_protocol.NoMoreFiles(None) def close(self): self.closed = True From e6f00266811dc5c68dab92cf029c1cc8621fc7a9 Mon Sep 17 00:00:00 2001 From: bifrost0x Date: Mon, 14 Sep 2026 12:35:53 +0200 Subject: [PATCH 2/2] Give shared browser fixtures a bounded test request budget --- tests/e2e/run_app.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/e2e/run_app.py b/tests/e2e/run_app.py index d1cabf8c..cfc3d6e0 100644 --- a/tests/e2e/run_app.py +++ b/tests/e2e/run_app.py @@ -271,6 +271,9 @@ def main(): 'REGISTRATION_ENABLED': 'True', 'RATELIMIT_STORAGE_URL': 'memory://', 'RATELIMIT_LOGIN_LIMIT': '100 per minute', + # Browser cases share seeded accounts and open them in bursts. + # Budget enforcement is exercised by the socket contract tests. + 'SSH_KEY_LIST_RATELIMIT': '100 per minute', 'CORS_ORIGINS': ( 'http://127.0.0.1:' + os.environ.get('WEBSSH_E2E_PORT', '4173')