import io import base64 import json import os from pathlib import Path import subprocess import sys import tempfile import time import unittest from contextlib import redirect_stdout, redirect_stderr from unittest import mock sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from scripts import oracle_skill as skill from scripts.runtime_state import file_lock, read_config, update_config def response(status=200, payload=None): result = mock.Mock(status_code=status, ok=status < 400) result.json.return_value = payload or {"success": True, "data": "ok"} if status >= 400: result.raise_for_status.side_effect = skill.requests.HTTPError(str(status)) return result def worker(mode, directory, target): directory = Path(directory) skill.get_script_dir = lambda: str(directory) cfg = skill.get_config() (directory / (target + '.ready')).touch() deadline = time.monotonic() + 15 while not (directory / 'go').exists(): if time.monotonic() >= deadline: raise RuntimeError('worker barrier timeout') time.sleep(.01) if mode == 'renew': skill.prepare_command(['skill', 'query', '--client', target, 'SELECT 1 FROM DUAL']) def renew(values): with open(directory / 'renewals', 'a') as handle: handle.write('renew\n') time.sleep(.2) values.update(access_token='new-token', expires_at='2999-01-01T00:00:00+00:00') skill.save_config(values) return True skill._auto_device_login_unlocked = renew rejected = 'old-token' if cfg.get('expires_at', '').startswith('2999') else None assert skill._auto_device_login(skill.get_config(), rejected) def post(url, **kwargs): assert kwargs['headers']['Authorization'] == 'Bearer new-token' assert kwargs['json']['server_id'] == target return response() skill.requests.post = post assert skill.query(skill.get_server_id(), 'execute_query', sql='SELECT 1 FROM DUAL')['client_code'] == target elif mode == 'merge': cfg[target] = target skill.save_config(cfg) elif mode == 'crash': with file_lock(str(directory / 'crash.lock')): (directory / 'locked').touch() time.sleep(20) class RuntimeTests(unittest.TestCase): def setUp(self): self.temp = tempfile.TemporaryDirectory() self.addCleanup(self.temp.cleanup) self.directory = Path(self.temp.name) self.path = self.directory / 'config.json' self.path.write_text(json.dumps({ 'transit_url': 'https://transit.example', 'access_token': 'old-token', 'expires_at': '2999-01-01T00:00:00+00:00', 'client_code': 'HENLO', 'server_id': 'HENLO', 'user_name': 'test', 'client_list': [{'code': 'HENLO'}, {'code': 'WEIRUI'}], })) patch = mock.patch.object(skill, 'get_script_dir', return_value=str(self.directory)) patch.start() self.addCleanup(patch.stop) skill.EXECUTION_CONTEXT = None skill.JSON_OUTPUT = False self.addCleanup(self.reset_context) def reset_context(self): skill.EXECUTION_CONTEXT = None skill.JSON_OUTPUT = False def prepare(self, *args): return skill.prepare_command(['skill', *args]) def test_explicit_target_does_not_persist_and_survives_default_switch(self): self.prepare('query', '--client', 'WEIRUI', 'SELECT 1 FROM DUAL') update_config(str(self.path), {'client_code': 'OTHER', 'server_id': 'OTHER'}) cfg = skill.get_config() cfg['access_token'] = 'new-token' skill.save_config(cfg) self.assertEqual(skill.get_server_id(), 'WEIRUI') stored = read_config(str(self.path)) self.assertEqual(stored['client_code'], 'OTHER') self.assertEqual(stored['access_token'], 'new-token') def test_implicit_default_is_fixed_at_command_start(self): self.prepare('describe', 'BOSNDS3', 'M_PRODUCT') update_config(str(self.path), {'client_code': 'WEIRUI', 'server_id': 'WEIRUI', 'transit_url': 'https://other.example'}) self.assertEqual(skill.get_server_id(), 'HENLO') self.assertEqual(skill.get_transit_url(), 'https://transit.example') def test_missing_target_stops_before_login(self): update_config(str(self.path), {'client_code': '', 'server_id': ''}) with mock.patch.object(skill.requests, 'post') as post: with self.assertRaisesRegex(ValueError, '客户'): self.prepare('query', 'SELECT 1 FROM DUAL') post.assert_not_called() def test_json_missing_target_returns_structured_error(self): update_config(str(self.path), {'client_code': '', 'server_id': ''}) output = io.StringIO() with mock.patch.object(sys, 'argv', ['skill', 'query', '--json', 'SELECT 1 FROM DUAL']), redirect_stdout(output), redirect_stderr(io.StringIO()), self.assertRaises(SystemExit): skill.main() self.assertFalse(json.loads(output.getvalue())['success']) def test_empty_explicit_target_and_missing_arguments_are_rejected(self): for args in [('query', '--client', '', 'SELECT 1 FROM DUAL'), ('describe', '--client', 'HENLO', 'BOSNDS3')]: with self.subTest(args=args), self.assertRaises(ValueError): self.prepare(*args) def test_positional_conflict_and_legacy_layouts(self): with self.assertRaisesRegex(ValueError, '不一致'): self.prepare('ops_report', '--client', 'HENLO', 'WEIRUI') self.assertEqual(self.prepare('ops_report', 'weirui')[-1], 'WEIRUI') self.assertEqual(self.prepare('awr_download', '--client', 'HENLO', '20261008')[-2:], ['HENLO', '20261008']) self.assertEqual(self.prepare('agent_update', '--client', 'HENLO', '300')[-2:], ['HENLO', '300']) def test_sql_is_not_scanned_for_options(self): sql = "SELECT '--client WEIRUI --json' FROM DUAL" argv = self.prepare('query', '--client', 'HENLO', sql) self.assertEqual(argv[-1], sql) self.assertFalse(skill.JSON_OUTPUT) def test_permission_lookup_keeps_customer_between_steps(self): calls = [] def query(server, action, *args, **kwargs): calls.append(server) update_config(str(self.path), {'client_code': 'WEIRUI', 'server_id': 'WEIRUI'}) return {'success': True, 'data': '{}'} with mock.patch.object(skill, 'query', side_effect=query): skill.query_with_permission('1', '2', 'SELECT 1 FROM DUAL') self.assertEqual(calls, ['HENLO', 'HENLO']) def test_valid_token_requires_only_business_request(self): self.prepare('query', '--client', 'HENLO', 'SELECT 1 FROM DUAL') with mock.patch.object(skill.requests, 'get') as get, mock.patch.object(skill.requests, 'post', return_value=response()) as post: self.assertTrue(skill.query('HENLO', 'execute_query', sql='SELECT 1 FROM DUAL')['success']) get.assert_not_called() self.assertEqual(post.call_count, 1) def test_401_recovers_once_and_preserves_payload(self): self.prepare('query', '--client', 'WEIRUI', 'SELECT 1 FROM DUAL') def renew(cfg): cfg['access_token'] = 'new-token' skill.save_config(cfg) return True with mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew) as login, mock.patch.object(skill.requests, 'post', side_effect=[response(401), response()]) as post: result = skill.query('WEIRUI', 'execute_query', sql='SELECT 1 FROM DUAL') self.assertTrue(result['success']) self.assertEqual(login.call_count, 1) self.assertEqual(post.call_count, 2) self.assertEqual(post.call_args.kwargs['headers']['Authorization'], 'Bearer new-token') self.assertEqual(post.call_args_list[0].kwargs['json'], post.call_args_list[1].kwargs['json']) def test_second_401_and_403_do_not_sleep_or_loop(self): for statuses in ([401, 401], [403]): with self.subTest(statuses=statuses): self.reset_context() self.prepare('query', '--client', 'HENLO', 'SELECT 1 FROM DUAL') def renew(cfg): cfg['access_token'] += '-new' skill.save_config(cfg) return True with mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew), mock.patch.object(skill.requests, 'post', side_effect=[response(status) for status in statuses]) as post, mock.patch.object(skill.time, 'sleep') as sleep: self.assertFalse(skill.query('HENLO', 'execute_query')['success']) self.assertEqual(post.call_count, len(statuses)) sleep.assert_not_called() def test_rejected_old_token_reuses_already_renewed_token(self): update_config(str(self.path), {'access_token': 'new-token'}) with mock.patch.object(skill, '_auto_device_login_unlocked') as login: cfg = skill.get_config() self.assertTrue(skill._auto_device_login(cfg, rejected_token='old-token')) login.assert_not_called() self.assertEqual(cfg['access_token'], 'new-token') def test_revoked_device_returns_failure(self): with mock.patch.object(skill, '_auto_device_login_unlocked', return_value=False), mock.patch.object(skill.requests, 'post', return_value=response(401)) as post: self.assertFalse(skill.query('HENLO', 'execute_query')['success']) self.assertEqual(post.call_count, 1) def test_real_device_challenge_signing_preserves_default_customer(self): from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey from cryptography.hazmat.primitives.serialization import Encoding, PrivateFormat, NoEncryption private = Ed25519PrivateKey.generate() key = {'device_id': 'test-device', 'key_fingerprint': 'test-fingerprint', 'private_key': base64.b64encode(private.private_bytes(Encoding.Raw, PrivateFormat.Raw, NoEncryption())).decode('ascii')} self.prepare('query', '--client', 'WEIRUI', 'SELECT 1 FROM DUAL') update_config(str(self.path), {'expires_at': '2000-01-01T00:00:00+00:00'}) responses = [response(payload={'success': True, 'challenge': 'test-challenge'}), response(payload={'success': True, 'access_token': 'signed-token', 'expires_at': '2999-01-01T00:00:00+00:00', 'current_client': {'code': 'WEIRUI'}})] with mock.patch.object(skill, '_load_device_key', return_value=key), mock.patch.object(skill.requests, 'post', side_effect=responses) as post: self.assertTrue(skill.ensure_logged_in()) login = post.call_args.kwargs['json'] private.public_key().verify(base64.b64decode(login['signature']), b'test-challenge') self.assertEqual(login['client_code'], 'WEIRUI') cfg = read_config(str(self.path)) self.assertEqual(cfg['access_token'], 'signed-token') self.assertEqual(cfg['client_code'], 'HENLO') def test_device_denial_is_diagnostic_and_does_not_fallback(self): from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey from cryptography.hazmat.primitives.serialization import Encoding, PrivateFormat, NoEncryption private = Ed25519PrivateKey.generate() key = {'device_id': 'test-device', 'key_fingerprint': 'test-fingerprint', 'private_key': base64.b64encode(private.private_bytes(Encoding.Raw, PrivateFormat.Raw, NoEncryption())).decode('ascii')} cfg = skill.get_config() stderr = io.StringIO() with mock.patch.object(skill, '_load_device_key', return_value=key), mock.patch.object(skill.requests, 'post', return_value=response(403, {'success': False, 'error': 'device revoked'})) as post, redirect_stderr(stderr): self.assertFalse(skill._auto_device_login(cfg, rejected_token='old-token')) self.assertEqual(post.call_count, 1) self.assertIn('device revoked', stderr.getvalue()) self.assertEqual(cfg['access_token'], 'old-token') def test_management_transport_failure_is_not_replayed(self): cfg = skill.get_config() cfg['upgrade_admin_token'] = 'test-admin-token' skill.save_config(cfg) with mock.patch.object(skill.requests, 'post', side_effect=skill.requests.Timeout) as post: self.assertFalse(skill.agent_update('HENLO')['success']) self.assertEqual(post.call_count, 1) def test_admin_401_does_not_renew_user_session(self): cfg = skill.get_config() cfg['upgrade_admin_token'] = 'test-admin-token' skill.save_config(cfg) with mock.patch.object(skill.requests, 'post', return_value=response(401)) as post, mock.patch.object(skill, '_auto_device_login') as login: self.assertFalse(skill.agent_update('HENLO')['success']) self.assertEqual(post.call_count, 1) login.assert_not_called() def test_timeout_not_replayed_and_connection_failures_are_bounded(self): with mock.patch.object(skill.requests, 'post', side_effect=skill.requests.Timeout) as post, mock.patch.object(skill.time, 'sleep') as sleep: self.assertFalse(skill.query('HENLO', 'execute_query')['success']) self.assertEqual(post.call_count, 1) sleep.assert_not_called() with mock.patch.object(skill.requests, 'post', side_effect=skill.requests.ConnectionError) as post, mock.patch.object(skill.time, 'sleep'): self.assertFalse(skill.query('HENLO', 'execute_query', max_retries=1)['success']) self.assertEqual(post.call_count, 2) def test_json_query_with_renewal_has_clean_stdout(self): stdout, stderr = io.StringIO(), io.StringIO() def renew(cfg): print('renew progress', file=sys.stderr) cfg['access_token'] = 'new-token' skill.save_config(cfg) return True argv = ['skill', 'query', '--client', 'WEIRUI', '--json', 'SELECT 1 FROM DUAL'] with mock.patch.object(sys, 'argv', argv), mock.patch.object(skill.requests, 'post', side_effect=[response(401), response()]), mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew), redirect_stdout(stdout), redirect_stderr(stderr): skill.main() self.assertEqual(json.loads(stdout.getvalue())['client_code'], 'WEIRUI') self.assertIn('renew progress', stderr.getvalue()) def test_status_and_refresh_use_one_clients_request(self): with mock.patch.object(skill.requests, 'get', return_value=response(payload={'success': True, 'clients': [{'code': 'HENLO', 'online': True}]})) as get, redirect_stdout(io.StringIO()): skill.cmd_status() self.assertEqual(get.call_count, 1) self.assertTrue(get.call_args.args[0].endswith('/api/clients')) def test_brief_capabilities_preserve_full_contract(self): results = [] for brief in (False, True): output = io.StringIO() with redirect_stdout(output): skill.cmd_capabilities(type('Args', (), {'json': True, 'brief': brief})()) results.append(output.getvalue()) full, brief = map(json.loads, results) self.assertEqual([c['name'] for c in full['commands']], [c['name'] for c in brief['commands']]) self.assertIn('parameters', full['commands'][0]) self.assertNotIn('parameters', brief['commands'][0]) self.assertLess(len(results[1]), len(results[0]) // 2) def test_stale_snapshots_merge_only_changed_fields(self): a, b = skill.get_config(), skill.get_config() a['access_token'] = 'new-token' b['client_code'] = b['server_id'] = 'WEIRUI' skill.save_config(a) skill.save_config(b) result = read_config(str(self.path)) self.assertEqual(result['access_token'], 'new-token') self.assertEqual(result['client_code'], 'WEIRUI') def test_initialization_does_not_replace_existing_credentials(self): cfg = update_config(str(self.path), {'access_token': '', 'client_code': ''}, initialize=True) self.assertEqual(cfg['access_token'], 'old-token') self.assertEqual(read_config(str(self.path))['client_code'], 'HENLO') def test_template_initialization_keeps_all_fields(self): self.path.unlink() (self.directory / 'config.template.json').write_text(json.dumps({'transit_url': 'https://template.example', 'server_id': '', 'custom': 'keep'})) with mock.patch.object(skill, 'read_config', wraps=read_config) as read: # The CWD fallback may exist in the repository; make only that path absent. def read_without_cwd(path): if path == 'config.json': raise FileNotFoundError(path) return read_config(path) read.side_effect = read_without_cwd cfg = skill.get_config() self.assertEqual(cfg['custom'], 'keep') self.assertEqual(read_config(str(self.path))['transit_url'], 'https://template.example') def test_streaming_get_401_keeps_download_options(self): old, new = response(401), response() def renew(cfg): cfg['access_token'] = 'new-token' skill.save_config(cfg) return True with mock.patch.object(skill.requests, 'get', side_effect=[old, new]) as get, mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew): result = skill._http_get('https://transit.example/api/awr/download', params={'client_code': 'WEIRUI', 'date': '20261008'}, headers=skill.make_headers(), stream=True, timeout=120) self.assertIs(result, new) old.close.assert_called_once() self.assertTrue(get.call_args.kwargs['stream']) self.assertEqual(get.call_args.kwargs['params']['client_code'], 'WEIRUI') def test_lock_wait_is_bounded(self): path = str(self.directory / 'test.lock') with file_lock(path): with self.assertRaises(TimeoutError): with file_lock(path, timeout=.1): self.fail('lock acquired twice') def start_workers(self, mode): processes = [subprocess.Popen([sys.executable, __file__, '--worker', mode, str(self.directory), target], stdout=subprocess.PIPE, stderr=subprocess.PIPE) for target in ('HENLO', 'WEIRUI')] for process in processes: self.addCleanup(lambda p=process: p.kill() if p.poll() is None else None) deadline = time.monotonic() + 15 while not all((self.directory / (target + '.ready')).exists() for target in ('HENLO', 'WEIRUI')): if time.monotonic() > deadline or any(p.poll() is not None for p in processes): self.fail('workers failed to start') time.sleep(.01) (self.directory / 'go').touch() return processes def assert_workers(self, processes): for process in processes: stdout, stderr = process.communicate(timeout=15) self.assertEqual(process.returncode, 0, stderr.decode('utf-8', errors='replace')) def test_two_processes_share_one_expired_login_and_keep_targets(self): update_config(str(self.path), {'expires_at': '2000-01-01T00:00:00+00:00'}) processes = self.start_workers('renew') update_config(str(self.path), {'client_code': 'OTHER', 'server_id': 'OTHER'}) self.assert_workers(processes) self.assertEqual((self.directory / 'renewals').read_text().splitlines(), ['renew']) self.assertEqual(read_config(str(self.path))['client_code'], 'OTHER') def test_two_processes_share_one_401_renewal(self): self.assert_workers(self.start_workers('renew')) self.assertEqual((self.directory / 'renewals').read_text().splitlines(), ['renew']) def test_two_processes_merge_configuration_updates(self): self.assert_workers(self.start_workers('merge')) cfg = read_config(str(self.path)) self.assertEqual(cfg['HENLO'], 'HENLO') self.assertEqual(cfg['WEIRUI'], 'WEIRUI') def test_crashed_process_releases_lock(self): processes = self.start_workers('crash') deadline = time.monotonic() + 10 while not (self.directory / 'locked').exists(): if time.monotonic() > deadline: self.fail('lock not acquired') time.sleep(.01) for process in processes: process.kill() process.communicate(timeout=5) with file_lock(str(self.directory / 'crash.lock'), timeout=.5): pass if __name__ == '__main__': if len(sys.argv) > 1 and sys.argv[1] == '--worker': worker(*sys.argv[2:]) else: unittest.main()