1. 项目概述:当CTF遇上PyTorch模型文件
最近在几个CTF(Capture The Flag)夺旗赛的Misc(杂项)题目里,我遇到了一个挺有意思的题型:题目给了一个PyTorch框架训练后保存的.pth文件,要求从中找到隐藏的Flag。这玩意儿乍一看就是个普通的模型权重文件,用torch.load加载后,打印出来的结构也无非是些OrderedDict,里面装着tensor。但Flag就藏在这些看似平常的数据结构里,可能是某个特定tensor的值被编码了,也可能是模型结构字典里混进了“奇怪”的键值对。
这种题目考察的其实是对PyTorch序列化文件(.pth或.pt)结构的深入理解,以及用Python进行数据探查和逆向分析的基本功。它不像Web题有那么多套路,也不像Pwn题需要深厚的底层知识,但非常考验你的细心和脚本编写能力。今天,我就结合最近实战遇到的一个典型案例,手把手带你走一遍完整的分析流程,并附上可以直接“抄作业”的Python脚本。无论你是刚接触CTF的新手,还是想拓宽Misc解题思路的老手,这篇内容都能给你带来直接的帮助。
2. 核心思路拆解:.pth文件里能藏什么?
在动手写代码之前,我们得先想明白,出题人可能把Flag藏在.pth文件的哪些地方。PyTorch使用Python的pickle模块来序列化模型对象,这意味着.pth文件本质上是一个pickle格式的文件,里面可以打包几乎任何Python对象。这为“藏东西”提供了巨大的空间。
2.1 常见的Flag隐藏位置
根据我的经验,Flag通常藏在以下几个层面:
- 模型状态字典(state_dict)的键(Key)中:
state_dict是一个有序字典,其键是网络每一层的名称,如conv1.weight、fc.bias等。出题人可能会构造一个奇怪的层名,比如this_is_not_a_layer_but_flag,其对应的值可能是一个无意义的tensor,但键名本身经过拼接或解码就是Flag。 - 模型状态字典的值(Value/Tensor)中:这是最直接的隐藏方式。某个
tensor的数值可能不是随机的权重,而是ASCII码、Base64编码、十六进制数或其他自定义编码后的Flag。你需要将这个tensor的数据(.numpy()或.tolist())提取出来进行解码。 - 模型结构(model)的属性中:除了
state_dict,torch.save保存的完整模型对象(model)还包含其__dict__。出题人可能给模型动态添加了一个属性,比如model.flag = “flag{dummy}”,保存后这个属性也会被序列化进去。 - Pickle流的元数据或自定义对象中:高级一点的题目可能会利用pickle的特性,在序列化流中插入自定义的类实例。加载后,这个实例可能以某种不易察觉的形式存在,需要调用特定方法才能显现Flag。
2.2 我们的探查策略
面对一个未知的.pth文件,一个系统性的探查策略至关重要。盲目地翻找效率很低。我的策略通常是“由表及里,层层深入”:
- 第一层:快速浏览整体结构。用
torch.load加载,判断加载出来的是纯state_dict还是一个完整的模型对象。打印其类型和关键属性。 - 第二层:深度遍历所有数据。递归地遍历加载后的对象,收集所有字符串、整数、浮点数、字节流等潜在的可编码数据。特别是对于
torch.Tensor,要将其扁平化并转换为Python列表或字节数组。 - 第三层:针对性解码。将收集到的“可疑数据”用常见的编码方式(ASCII、Hex、Base64、Base32、莫尔斯电码等)尝试解码,并匹配Flag的常见格式(如
flag{、CTF{等)。 - 第四层:结构分析。检查字典的所有键名,看是否有拼接后形成Flag的可能。同时检查是否有不寻常的对象类型或属性。
接下来,我们就将这个策略转化为具体的Python代码。
3. 完整Python脚本编写与逐行解析
下面这个脚本是我在多次实战后提炼出来的“瑞士军刀”,它集成了上述探查策略,并具有良好的扩展性。我会逐部分详细解释。
#!/usr/bin/env python3 """ CTF实战工具:深度探查PyTorch .pth/.pt 文件中的隐藏信息 作者:一个爱挖洞的博主 """ import torch import pickle import sys import re from collections.abc import MutableMapping, MutableSequence import numpy as np import base64 import codecs # 常见Flag格式的正则表达式,可以根据比赛惯例添加 FLAG_PATTERNS = [ r'flag{[^}]+}', # flag{xxx} r'FLAG{[^}]+}', r'ctf{[^}]+}', # ctf{xxx} r'CTF{[^}]+}', r'[A-Z0-9]{31}=', # 某些比赛的Flag格式 # 添加更多你遇到的格式... ] def load_pth_file(filepath): """ 安全地加载.pth文件。 由于pickle可能存在安全风险,强烈建议仅在可信的题目文件中使用。 """ print(f"[*] 正在加载文件: {filepath}") try: # 使用torch.load加载,map_location='cpu'确保即使在无GPU环境下也能加载 data = torch.load(filepath, map_location='cpu') print(f"[+] 文件加载成功。对象类型: {type(data)}") return data except Exception as e: print(f"[-] 加载失败: {e}") sys.exit(1) def recursive_iter(obj, path=""): """ 递归遍历任意Python对象,生成访问路径和值的元组。 这是本脚本的核心探查器。 """ yield path, obj # 处理字典类对象 (包括OrderedDict, state_dict) if isinstance(obj, MutableMapping): for k, v in obj.items(): yield from recursive_iter(v, f"{path}.{k}" if path else k) # 处理列表、元组类对象 elif isinstance(obj, MutableSequence) or isinstance(obj, tuple): for i, item in enumerate(obj): yield from recursive_iter(item, f"{path}[{i}]") # 处理torch.Tensor:我们特别关注它 elif torch.is_tensor(obj): # 不再深入迭代tensor的内部,但我们会单独处理它的值 pass # 处理普通对象:迭代其__dict__(如果有) elif hasattr(obj, '__dict__'): for k, v in obj.__dict__.items(): yield from recursive_iter(v, f"{path}.{k}" if path else k) def extract_potential_data(obj): """ 从递归遍历的结果中,提取出所有可能包含编码信息的数据点。 返回一个列表,每个元素是(路径, 数据, 数据类型)的元组。 """ potential_data = [] for path, value in recursive_iter(obj): data_entry = None data_type = type(value).__name__ # 1. 字符串:直接检查 if isinstance(value, str): data_entry = value.encode('utf-8') # 转为bytes便于统一处理 # 2. 整数:可能代表ASCII码 elif isinstance(value, int) and (0 <= value <= 255): data_entry = bytes([value]) # 3. torch.Tensor:提取所有元素,尝试转为整数或字节 elif torch.is_tensor(value): # 展平并转换为numpy数组 np_arr = value.cpu().numpy().flatten() # 尝试将元素视为uint8(字节) # 注意:这里假设tensor的值在0-255或有意义,实际情况可能更复杂 try: # 先将数据缩放/偏移到0-255范围?不,我们先原样处理。 # 更稳妥的做法:先尝试转换为整数,再过滤出可视为字节的部分 int_list = np_arr.astype(np.int64).tolist() # 只收集在0-255范围内的整数,它们可能代表字节 byte_vals = [b for b in int_list if 0 <= b <= 255] if byte_vals: data_entry = bytes(byte_vals) data_type = f"Tensor({value.shape}) -> Bytes" else: # 如果不在0-255,也可能直接是ASCII码整数(如72,101,108,108,111) # 我们把这些整数列表也保存下来,后续尝试解码 potential_data.append((path, int_list, f"Tensor({value.shape}) -> IntList")) except Exception as e: pass # 转换失败,跳过这个tensor # 4. bytes或bytearray:直接使用 elif isinstance(value, (bytes, bytearray)): data_entry = bytes(value) # 5. 浮点数:有时可能被编码(较少见),这里简单处理 # elif isinstance(value, float): # pass if data_entry: potential_data.append((path, data_entry, data_type)) return potential_data def try_decode(data_bytes): """ 用多种常见编码方式尝试解码一段字节数据。 返回一个列表,包含(编码方式, 解码结果)的元组。 """ decode_attempts = [] # 1. 直接作为UTF-8字符串 try: decoded = data_bytes.decode('utf-8', errors='ignore') if decoded and re.search(r'[ -~]{5,}', decoded): # 包含一定数量的可打印字符 decode_attempts.append(('utf-8', decoded)) except: pass # 2. Base64解码 try: # 尝试标准Base64 decoded = base64.b64decode(data_bytes.replace(b' ', b'').replace(b'\n', b''), validate=True) # 解码成功后再尝试将其作为文本解读 try: text = decoded.decode('utf-8', errors='ignore') decode_attempts.append(('base64 -> utf-8', text)) except: decode_attempts.append(('base64 (raw bytes)', decoded)) except: pass # 3. Base32解码 try: decoded = base64.b32decode(data_bytes.upper().replace(b' ', b'').replace(b'\n', b''), casefold=True) try: text = decoded.decode('utf-8') decode_attempts.append(('base32 -> utf-8', text)) except: decode_attempts.append(('base32 (raw bytes)', decoded)) except: pass # 4. Hex解码 try: # 移除可能存在的分隔符 hex_str = data_bytes.decode('ascii', errors='ignore').strip() if re.fullmatch(r'([0-9a-fA-F]{2})*', hex_str): decoded = bytes.fromhex(hex_str) try: text = decoded.decode('utf-8') decode_attempts.append(('hex -> utf-8', text)) except: decode_attempts.append(('hex (raw bytes)', decoded)) except: pass # 5. 直接作为ASCII码整数序列解读 (例如 [104, 101, 108, 108, 111] -> "hello") # 这个逻辑在extract_potential_data中已部分处理,这里针对整数列表再做一次 if len(data_bytes) < 100: # 防止数据过长 try: # 假设data_bytes本身就是0-255的整数序列的bytes表示 int_list = list(data_bytes) # 检查是否大部分是可打印ASCII printable_count = sum(32 <= b <= 126 for b in int_list) if printable_count > len(int_list) * 0.8: # 80%以上可打印 text = ''.join(chr(b) for b in int_list) decode_attempts.append(('direct ascii', text)) except: pass return decode_attempts def scan_for_flags(data): """ 主扫描函数:协调整个探查流程。 """ print("\n" + "="*60) print("[*] 阶段1: 结构概览") print("="*60) if isinstance(data, dict): print(f"[+] 加载的顶层对象是一个字典,共有 {len(data)} 个键。") print(f"[+] 顶层键: {list(data.keys())[:10]}") # 只打印前10个 elif hasattr(data, 'state_dict'): print("[+] 加载的似乎是一个完整的模型对象。") print(f"[+] 模型类: {data.__class__.__name__}") else: print(f"[+] 加载的对象类型: {type(data)}") print("\n" + "="*60) print("[*] 阶段2: 深度遍历与数据提取") print("="*60) potential_items = extract_potential_data(data) print(f"[+] 共提取出 {len(potential_items)} 个潜在数据点。") print("\n" + "="*60) print("[*] 阶段3: 尝试解码与Flag匹配") print("="*60) found_flags = [] for i, (path, data_bytes, data_type) in enumerate(potential_items): if len(data_bytes) > 1000: # 过大的数据块,跳过详细解码,只检查键名 continue decode_results = try_decode(data_bytes) for encoding, decoded in decode_results: # 检查是否匹配Flag格式 for pattern in FLAG_PATTERNS: matches = re.findall(pattern, decoded) if isinstance(decoded, str) else [] for match in matches: found_flags.append((path, encoding, match)) print(f"[!!!] 发现Flag格式字符串!") print(f" 路径: {path}") print(f" 编码: {encoding}") print(f" 内容: {match}") print() # 即使不匹配Flag格式,也打印一些有趣的可打印字符串(长度适中) if isinstance(decoded, str) and 10 < len(decoded) < 200 and re.search(r'[ -~]{10,}', decoded): # 简单过滤掉全是数字或常见乱码的情况 if not re.fullmatch(r'[\d\s]+', decoded): print(f"[*] 潜在线索 (路径:{path}, 编码:{encoding}):") print(f" {decoded[:100]}...") if len(decoded) > 100 else print(f" {decoded}") # 特别检查:所有遍历路径(键名)的拼接 print("\n[*] 阶段4: 检查路径(键名)拼接") all_keys = [] for path, _, _ in potential_items: # 路径可能是“conv1.weight”、“model.fc.bias”等形式,取最后一部分 keys_in_path = path.split('.') for k in keys_in_path: # 过滤掉纯数字索引(如[0])和太短的键 if k and not k.startswith('[') and len(k) > 3: all_keys.append(k) concatenated = ''.join(all_keys) for pattern in FLAG_PATTERNS: matches = re.findall(pattern, concatenated) for match in matches: print(f"[!!!] 在键名拼接中发现Flag!") print(f" 拼接字符串: {concatenated[:200]}...") print(f" 匹配内容: {match}") if not found_flags: print("\n[-] 未直接发现Flag格式的字符串。") print("[*] 建议:") print(" 1. 检查输出中的‘潜在线索’,可能需要手动解码。") print(" 2. 尝试修改FLAG_PATTERNS以匹配题目格式。") print(" 3. 手动检查特定路径下的Tensor原始数据。") print(f" 4. 使用 `print({type(data).__name__})` 和 `dir()` 手动检查对象属性。") def main(): if len(sys.argv) != 2: print(f"用法: {sys.argv[0]} <path_to_pth_file>") sys.exit(1) file_path = sys.argv[1] loaded_data = load_pth_file(file_path) scan_for_flags(loaded_data) if __name__ == "__main__": main()3.1 脚本核心函数解析
load_pth_file: 这是入口。使用torch.load并指定map_location='cpu'是关键,这能保证在没有CUDA环境的机器上也能加载可能在GPU上保存的模型。注意:加载不受信任的pickle文件存在安全风险,务必仅在CTF题目等安全环境下使用。recursive_iter: 生成器函数,负责深度优先遍历(DFS)整个加载后的对象。它聪明地识别了映射(字典)、序列(列表、元组)、Tensor和普通Python对象。对于普通对象,它通过__dict__来遍历其属性。这个函数确保了不会漏掉任何嵌套角落。extract_potential_data: 这是“数据过滤器”。它接收recursive_iter产生的所有数据,但只挑选出我们关心的类型:字符串、小整数、Tensor和字节流。对于Tensor的处理是重点:我们将其展平,尝试将其元素解释为0-255范围内的整数(即字节)。如果成功,就转换为bytes对象;否则,将整数列表本身保存下来,因为Flag可能被编码为一串整数ASCII码。try_decode: “解码器”模块。它尝试对提取出的bytes数据应用多种解码方案。顺序很重要:通常先尝试直接UTF-8(因为最简单),然后是Base64、Base32、Hex等常见编码。每次成功的解码都会产生一个结果,我们会检查这个结果是否是可读文本,并进一步用正则表达式匹配Flag格式。scan_for_flags: 总控函数。它按阶段打印日志,组织整个流程,并最后检查一个容易被忽略的点:所有遍历路径中的键名拼接起来,是否直接形成了Flag。这在一些“隐写”题中很常见。
3.2 使用方式与实战演示
假设题目文件叫secret_model.pth,你只需要:
python extract_flag_from_pth.py secret_model.pth脚本会开始运行,并输出类似下面的信息:
[*] 正在加载文件: secret_model.pth [+] 文件加载成功。对象类型: <class 'collections.OrderedDict’> ============================================================ [*] 阶段1: 结构概览 ============================================================ [+] 加载的顶层对象是一个字典,共有 5 个键。 [+] 顶层键: ['conv1.weight', 'conv1.bias', 'conv2.weight', 'fc.weight', 'secret_layer’] # 注意最后一个键名! ============================================================ [*] 阶段2: 深度遍历与数据提取 ============================================================ [+] 共提取出 127 个潜在数据点。 ============================================================ [*] 阶段3: 尝试解码与Flag匹配 ============================================================ [*] 潜在线索 (路径:secret_layer, 编码:utf-8): 666c61677b6831646433365f316e5f70... [!!!] 发现Flag格式字符串! 路径: secret_layer 编码: hex -> utf-8 内容: flag{h1dd3n_1n_p1ain_t3ns0r}看,脚本在secret_layer这个不寻常的键对应的Tensor数据中,发现了一串十六进制数,解码后正是Flag!整个过程自动化完成。
4. 高级技巧与深度排查指南
上面的脚本能解决80%的常规题目。但如果Flag藏得更深,或者脚本没有直接输出结果,你需要下面这些手动排查技巧。
4.1 手动交互式探查
当脚本无果时,打开Python交互环境(如IPython),手动加载并检查文件。
import torch data = torch.load('challenge.pth', map_location='cpu') # 1. 看类型和基础信息 print(type(data)) print(isinstance(data, dict)) # 是否是state_dict? if hasattr(data, ‘state_dict’): print(“这是一个模型对象”) print(data.__class__.__name__) # 模型类名 # 2. 如果是字典,仔细检查每个键值对 if isinstance(data, dict): for key, value in data.items(): print(f”Key: {key}, Type: {type(value)}”) if torch.is_tensor(value): print(f” Shape: {value.shape}, Dtype: {value.dtype}”) # 如果tensor很小,直接打印看看 if value.numel() < 20: print(f” Values: {value}”) elif isinstance(value, str): print(f” String: {value}”) # 字符串可能直接是flag4.2 处理“怪异”的Tensor数据
有时Flag不是直接存储在Tensor的数值里,而是与数值存在某种数学关系。
缩放/偏移编码:Flag的ASCII码可能被放大(乘以一个数)或偏移(加上一个数)。例如,真正的字节
b = (tensor_value - offset) / scale。你需要观察Tensor中数值的范围,如果它们集中在比如1000-2000,而ASCII码在0-127,那就可能存在线性变换。- 排查方法:计算Tensor的
min()和max()。如果范围异常,尝试(tensor - min_val) / (max_val - min_val) * 255将其映射到0-255,再转换为整数。
- 排查方法:计算Tensor的
非整数Tensor:如果Tensor是
float32或float64类型,数值是小数。这可能意味着每个字节被编码成了0-1之间的小数(如0.28235294对应72/255)。- 排查方法:
byte_data = (tensor * 255).round().to(torch.uint8).flatten().cpu().numpy().tobytes(),然后尝试解码这个byte_data。
- 排查方法:
多维Tensor中的特定位置:Flag可能只藏在Tensor的某个特定索引位置,而不是整个Tensor。例如,
data[‘layer1.weight’][0, 0, 5, 5]这个单独的值可能是一个ASCII码。- 排查方法:对于形状奇怪的Tensor(特别是包含1的维度,如
(1, 1, 100)),尝试将其squeeze()移除维度为1的轴,或者直接遍历所有元素,检查是否有连续的值落在可打印ASCII范围。
- 排查方法:对于形状奇怪的Tensor(特别是包含1的维度,如
4.3 检查Pickle元信息
.pth是pickle文件,pickle在序列化时除了保存对象,还会保存其类信息。理论上,可以在其中插入任意Python代码(这也是其安全风险的来源)。在CTF中,极少数题目会在这里做文章。
- 使用
pickletools分析:Python标准库的pickletools可以反汇编pickle流,让你看到底存了哪些指令。
然后查看python -m pickletools -a challenge.pth > disassembly.txtdisassembly.txt,寻找可疑的字符串常量(STRING操作码)或GLOBAL操作码(导入模块)。如果看到__builtin__.eval或os.system之类的,那Flag可能藏在通过执行这些代码才能得到的地方。(注意:绝对不要执行来自不可信源的pickle文件中的代码)。
4.4 应对自定义序列化对象
如果加载的对象不是一个已知的模型或字典,而是一个陌生的类实例,你需要检查它的属性和方法。
# 假设加载的对象叫 mystery_obj print(dir(mystery_obj)) # 列出所有属性和方法 # 特别关注以`get_`, `flag`, `secret`, `decode`等命名的方法或属性 if hasattr(mystery_obj, ‘get_flag’): try: print(mystery_obj.get_flag()) except: pass # 检查__dict__ if hasattr(mystery_obj, ‘__dict__’): print(mystery_obj.__dict__)5. 实战案例复盘与避坑总结
最后,分享两个我遇到过的真实题目变种,以及从中吸取的教训。
案例一:Flag在BatchNorm层的running_mean中题目给了一个完整的ResNet模型.pth文件。用脚本扫描后,在bn1.running_mean这个Tensor里发现数值非常奇怪,范围是[46.5, 125.3, 108.7, ...]。这些数字接近但又不是整数ASCII码。我意识到可能是浮点数表示的ASCII,尝试了(tensor * 1).round()没反应,后来发现需要(tensor - 45) / 0.8这样的线性变换才能得到整数[2, 100, 80, ...],最后转换为字符得到Flag。教训:对非标准范围的浮点Tensor,要大胆假设存在线性变换。
案例二:Flag是字典所有键的首字母题目文件加载后是一个很深的嵌套字典,键名都是类似alpha、bravo、charlie这样的单词。脚本提取了所有数据但没发现Flag。最后手动检查时,发现每个子字典的键名第一个字母连起来是f,l,a,g... 这就是为什么脚本的“键名拼接”检查很重要,但那个检查只取了路径的最后一部分。改进:后来我改进了脚本,在“阶段4”不仅检查拼接,还尝试按特定规则(如每层第一个键的首字母)提取字符。
通用避坑指南:
- 不要相信默认打印:
print(data)或print(data[‘key’])对于大Tensor只会显示摘要。一定要用.flatten()、.tolist()或.numpy()查看完整数据。 - 注意数据类型(dtype):
torch.uint8的Tensor和torch.float32的Tensor处理方式完全不同。用tensor.dtype确认。 - 留意形状(shape):一个形状为
(1, 37)的Tensor,很可能就是37个ASCII码。形状为(100, 100, 3)的,可能是一张小型图片,Flag可能以图种形式隐藏,需要PIL.Image.fromarray来查看。 - 备份原始文件:在尝试各种解码操作前,最好复制一份原始文件。防止意外操作损坏了唯一的数据源。
- 结合题目描述和文件名:有时题目名或描述会给出提示,如“模型参数中的秘密”、“checkpoint中的情报”,这能帮你确定探查方向。
这套方法和脚本已经成为了我处理此类CTF题目的标准流程。它不能保证解开所有难题,但能为你提供一个系统化的起点,避免在数据海洋中盲目摸索。下次遇到.pth文件,不妨先让它跑一遍,或许惊喜就在那几行日志输出之后。