From 6da2be0da0bb8fdca08ec9786b0233e6927897dc Mon Sep 17 00:00:00 2001 From: Arison591 <2234819536@qq.com> Date: Fri, 17 Jul 2026 15:46:51 +0800 Subject: [PATCH] fix: replace unsafe reward pickle protocol --- flow_grpo/rewards.py | 25 ++-- reward_server/protocol.py | 115 ++++++++++++++++++ scripts/serve_general_reward.py | 32 +++-- scripts/serve_reward_3d.py | 56 +++++---- tests/test_reward_protocol.py | 209 ++++++++++++++++++++++++++++++++ 5 files changed, 384 insertions(+), 53 deletions(-) create mode 100644 reward_server/protocol.py create mode 100644 tests/test_reward_protocol.py diff --git a/flow_grpo/rewards.py b/flow_grpo/rewards.py index 0e3b8f4..30f06fa 100644 --- a/flow_grpo/rewards.py +++ b/flow_grpo/rewards.py @@ -7,7 +7,6 @@ """ import os -import pickle from io import BytesIO import numpy as np @@ -16,6 +15,8 @@ from PIL import Image from requests.adapters import HTTPAdapter, Retry +from reward_server.protocol import encode_general_request, encode_reward_3d_request + REWARD_3D = "reward_3d" REWARD_GENERAL = "reward_general" REWARD_TOTAL = "reward_total" @@ -117,13 +118,14 @@ def _fn(images, prompts, metadata): for video_batch, prompt_batch, trajectory_batch in zip( videos_batched, prompts_batched, camera_trajectories_batched ): - data = { - "videos": video_batch, - "prompts": prompt_batch, - "camera_trajectories": trajectory_batch, - } - response = sess.post(url, data=pickle.dumps(data), timeout=2000) - response_data = pickle.loads(response.content) + payload = encode_reward_3d_request( + video_batch, + prompt_batch, + trajectory_batch, + ) + response = sess.post(url, json=payload, timeout=2000) + response.raise_for_status() + response_data = response.json() all_scores += response_data["outputs"] if response_data.get("details"): all_reconstruction_scores += [float(item["gs_score"]) for item in response_data["details"]] @@ -195,9 +197,10 @@ def _fn(images, prompts, metadata): all_scores = [] for image_batch, prompt_batch in zip(images_batched, prompts_batched): - data = {"images": image_batch, "prompts": prompt_batch} - response = sess.post(url, data=pickle.dumps(data), timeout=1000) - response_data = pickle.loads(response.content) + payload = encode_general_request(image_batch, prompt_batch) + response = sess.post(url, json=payload, timeout=1000) + response.raise_for_status() + response_data = response.json() all_scores += response_data["outputs"] return all_scores, {} diff --git a/reward_server/protocol.py b/reward_server/protocol.py new file mode 100644 index 0000000..325aec8 --- /dev/null +++ b/reward_server/protocol.py @@ -0,0 +1,115 @@ +"""JSON protocol helpers shared by reward clients and servers.""" + +import base64 +import binascii +from typing import Any + + +def _require_list(value: Any, field: str) -> list: + if not isinstance(value, list): + raise ValueError(f"'{field}' must be a list") + return value + + +def _encode_bytes(value: Any, field: str) -> str: + if not isinstance(value, (bytes, bytearray, memoryview)): + raise ValueError(f"'{field}' entries must be bytes") + return base64.b64encode(bytes(value)).decode("ascii") + + +def _decode_bytes(value: Any, field: str) -> bytes: + if not isinstance(value, str): + raise ValueError(f"'{field}' entries must be base64 strings") + try: + return base64.b64decode(value, validate=True) + except (binascii.Error, ValueError) as exc: + raise ValueError(f"'{field}' contains invalid base64 data") from exc + + +def _validate_prompts(prompts: Any, batch_size: int) -> list[str]: + prompts = _require_list(prompts, "prompts") + if len(prompts) != batch_size: + raise ValueError("media and prompt batch sizes must match") + if not all(isinstance(prompt, str) for prompt in prompts): + raise ValueError("'prompts' entries must be strings") + return prompts + + +def _to_json_compatible(value: Any) -> Any: + if value is None or isinstance(value, (bool, int, float, str)): + return value + if isinstance(value, dict): + if not all(isinstance(key, str) for key in value): + raise ValueError("camera trajectory keys must be strings") + return {key: _to_json_compatible(item) for key, item in value.items()} + if isinstance(value, (list, tuple)): + return [_to_json_compatible(item) for item in value] + if hasattr(value, "tolist"): + return _to_json_compatible(value.tolist()) + if hasattr(value, "item"): + return _to_json_compatible(value.item()) + raise ValueError(f"camera trajectory value is not JSON serializable: {type(value).__name__}") + + +def encode_general_request(images: list[bytes], prompts: list[str]) -> dict[str, Any]: + """Encode a general-reward batch as a JSON-compatible object.""" + images = _require_list(images, "images") + prompts = _validate_prompts(prompts, len(images)) + return { + "images": [_encode_bytes(image, "images") for image in images], + "prompts": prompts, + } + + +def decode_general_request(payload: Any) -> tuple[list[bytes], list[str]]: + """Validate and decode a general-reward request.""" + if not isinstance(payload, dict): + raise ValueError("request payload must be a JSON object") + images = _require_list(payload.get("images"), "images") + prompts = _validate_prompts(payload.get("prompts"), len(images)) + return [_decode_bytes(image, "images") for image in images], prompts + + +def encode_reward_3d_request( + videos: list[list[bytes]], + prompts: list[str], + camera_trajectories: list[Any] | None = None, +) -> dict[str, Any]: + """Encode a 3D-reward batch as a JSON-compatible object.""" + videos = _require_list(videos, "videos") + prompts = _validate_prompts(prompts, len(videos)) + encoded_videos = [] + for video in videos: + video = _require_list(video, "videos") + encoded_videos.append([_encode_bytes(frame, "videos") for frame in video]) + + payload = {"videos": encoded_videos, "prompts": prompts} + if camera_trajectories is not None: + camera_trajectories = _require_list(camera_trajectories, "camera_trajectories") + if len(camera_trajectories) != len(videos): + raise ValueError("video and camera trajectory batch sizes must match") + payload["camera_trajectories"] = _to_json_compatible(camera_trajectories) + return payload + + +def decode_reward_3d_request( + payload: Any, +) -> tuple[list[list[bytes]], list[str], list[Any] | None]: + """Validate and decode a 3D-reward request.""" + if not isinstance(payload, dict): + raise ValueError("request payload must be a JSON object") + videos = _require_list(payload.get("videos"), "videos") + prompts = _validate_prompts(payload.get("prompts"), len(videos)) + + decoded_videos = [] + for video in videos: + video = _require_list(video, "videos") + decoded_videos.append([_decode_bytes(frame, "videos") for frame in video]) + + camera_trajectories = payload.get("camera_trajectories") + if camera_trajectories is not None: + camera_trajectories = _require_list(camera_trajectories, "camera_trajectories") + if len(camera_trajectories) != len(videos): + raise ValueError("video and camera trajectory batch sizes must match") + + return decoded_videos, prompts, camera_trajectories diff --git a/scripts/serve_general_reward.py b/scripts/serve_general_reward.py index 654f86b..6eb1b96 100644 --- a/scripts/serve_general_reward.py +++ b/scripts/serve_general_reward.py @@ -1,21 +1,23 @@ #!/usr/bin/env python3 import argparse import os -import pickle -import traceback -from flask import Blueprint, Flask, request +from flask import Blueprint, Flask, current_app, jsonify, request -from reward_server.general_reward import MultiGPUGeneralRewardManager +from reward_server.protocol import decode_general_request root = Blueprint("root", __name__) general_reward_manager = None -def create_app(): +def create_app(manager=None): global general_reward_manager - general_reward_manager = MultiGPUGeneralRewardManager() - general_reward_manager.initialize() + if manager is None: + from reward_server.general_reward import MultiGPUGeneralRewardManager + + manager = MultiGPUGeneralRewardManager() + manager.initialize() + general_reward_manager = manager app = Flask(__name__) app.register_blueprint(root) return app @@ -23,22 +25,26 @@ def create_app(): @root.route("/", methods=["POST"]) def inference(): - data = request.get_data() + if not request.is_json: + return jsonify({"error": "Request body must be JSON"}), 400 + try: - payload = pickle.loads(data) - batch_images = payload["images"] - batch_prompts = payload["prompts"] + batch_images, batch_prompts = decode_general_request(request.get_json(silent=True)) batch_size = len(batch_images) + except ValueError as exc: + return jsonify({"error": str(exc)}), 400 + try: global general_reward_manager if general_reward_manager is None: outputs = [0.5] * batch_size else: outputs = general_reward_manager.compute_batch_scores(batch_images, batch_prompts) - return pickle.dumps({"outputs": outputs}), 200 + return jsonify({"outputs": [float(output) for output in outputs]}), 200 except Exception: - return traceback.format_exc().encode("utf-8"), 500 + current_app.logger.exception("General reward computation failed") + return jsonify({"error": "General reward computation failed"}), 500 HOST = "127.0.0.1" diff --git a/scripts/serve_reward_3d.py b/scripts/serve_reward_3d.py index 4e1f7d9..3614967 100644 --- a/scripts/serve_reward_3d.py +++ b/scripts/serve_reward_3d.py @@ -1,10 +1,8 @@ -import pickle import os import argparse -import traceback import signal import sys -from reward_server.reward_3d import MultiGPUReward3DManager +from reward_server.protocol import decode_reward_3d_request # import debugpy # try: @@ -15,7 +13,7 @@ # except Exception as e: # pass -from flask import Flask, request, Blueprint +from flask import Blueprint, Flask, current_app, jsonify, request root = Blueprint("root", __name__) @@ -29,15 +27,20 @@ def signal_handler(sig, frame): reward_3d_manager.shutdown() sys.exit(0) -def create_app(scorer_type='qwen', use_lpips=True): +def create_app(scorer_type='qwen', use_lpips=True, manager=None, install_signal_handlers=True): global reward_3d_manager - print(f"Initializing multi-GPU 3D reward server (scorer: {scorer_type}, lpips: {use_lpips})...") - reward_3d_manager = MultiGPUReward3DManager(scorer_type=scorer_type, use_lpips=use_lpips) - reward_3d_manager.initialize() + if manager is None: + from reward_server.reward_3d import MultiGPUReward3DManager + + print(f"Initializing multi-GPU 3D reward server (scorer: {scorer_type}, lpips: {use_lpips})...") + manager = MultiGPUReward3DManager(scorer_type=scorer_type, use_lpips=use_lpips) + manager.initialize() + reward_3d_manager = manager # Register signal handlers for graceful shutdown - signal.signal(signal.SIGINT, signal_handler) - signal.signal(signal.SIGTERM, signal_handler) + if install_signal_handlers: + signal.signal(signal.SIGINT, signal_handler) + signal.signal(signal.SIGTERM, signal_handler) app = Flask(__name__) app.register_blueprint(root) @@ -46,20 +49,22 @@ def create_app(scorer_type='qwen', use_lpips=True): @root.route("/", methods=["POST"]) def inference(): print(f"received POST request from {request.remote_addr}") - data = request.get_data() + if not request.is_json: + return jsonify({"error": "Request body must be JSON"}), 400 try: # expects a dict with "videos" and "prompts" # videos: List[List[bytes]] - outer list is batch_size, inner list is frames per video # prompts: List[str] - text prompts for each video - data = pickle.loads(data) - - batch_videos = data["videos"] # List[List[bytes]] - batch_prompts = data["prompts"] # List[str] - batch_camera_trajectories = data.get("camera_trajectories") + batch_videos, batch_prompts, batch_camera_trajectories = decode_reward_3d_request( + request.get_json(silent=True) + ) batch_size = len(batch_videos) print(f"Got batch of size {batch_size} for 3D reward evaluation") + except ValueError as exc: + return jsonify({"error": str(exc)}), 400 + try: global reward_3d_manager if reward_3d_manager is None: print("Error: 3D reward server is not initialized") @@ -75,20 +80,13 @@ def inference(): print(f"3D reward batch processing results: {outputs}") - response = {"outputs": outputs, "details": details} - - # returns: a dict with "outputs" - # outputs: List of scores (float values) with length = batch_size - response = pickle.dumps(response) - - returncode = 200 + return jsonify({ + "outputs": [float(output) for output in outputs], + "details": details, + }), 200 except Exception: - response = traceback.format_exc() - print(response) - response = response.encode("utf-8") - returncode = 500 - - return response, returncode + current_app.logger.exception("3D reward computation failed") + return jsonify({"error": "3D reward computation failed"}), 500 HOST = "127.0.0.1" diff --git a/tests/test_reward_protocol.py b/tests/test_reward_protocol.py new file mode 100644 index 0000000..c3299f9 --- /dev/null +++ b/tests/test_reward_protocol.py @@ -0,0 +1,209 @@ +import unittest +from unittest.mock import patch + +import torch + +from flow_grpo import rewards +from reward_server.protocol import ( + decode_general_request, + decode_reward_3d_request, + encode_general_request, + encode_reward_3d_request, +) +from scripts import serve_general_reward, serve_reward_3d + + +class FakeResponse: + def __init__(self, payload): + self.payload = payload + + def raise_for_status(self): + return None + + def json(self): + return self.payload + + +class FakeSession: + def __init__(self, response): + self.response = response + self.requests = [] + self.trust_env = True + + def mount(self, *args, **kwargs): + return None + + def post(self, url, **kwargs): + self.requests.append((url, kwargs)) + return FakeResponse(self.response) + + +class FakeGeneralRewardManager: + def __init__(self): + self.calls = [] + + def compute_batch_scores(self, images, prompts): + self.calls.append((images, prompts)) + return [0.75] * len(images) + + +class FakeReward3DManager: + def __init__(self): + self.calls = [] + self.last_results = { + "per_video_results": [ + { + "gs_score": 0.25, + "meta_score": 0.5, + "camera_motion_score": 0.75, + "trajectory_comparison_path": "trajectory.png", + } + ] + } + + def compute_batch_scores(self, videos, prompts, camera_trajectories=None): + self.calls.append((videos, prompts, camera_trajectories)) + return [1.5] * len(videos) + + +class RewardProtocolTest(unittest.TestCase): + def test_general_request_round_trip(self): + images = [b"first-image", b"\x00\xffsecond-image"] + prompts = ["first prompt", "second prompt"] + + payload = encode_general_request(images, prompts) + + self.assertEqual(decode_general_request(payload), (images, prompts)) + self.assertTrue(all(isinstance(image, str) for image in payload["images"])) + + def test_reward_3d_request_round_trip(self): + videos = [[b"frame-1", b"frame-2"], [b"frame-3"]] + prompts = ["first prompt", "second prompt"] + trajectories = [{"frame_0": [[1.0, 0.0], [0.0, 1.0]]}, None] + + payload = encode_reward_3d_request(videos, prompts, trajectories) + + self.assertEqual( + decode_reward_3d_request(payload), + (videos, prompts, trajectories), + ) + self.assertTrue( + all(isinstance(frame, str) for video in payload["videos"] for frame in video) + ) + + def test_decoders_reject_invalid_payloads(self): + with self.assertRaises(ValueError): + decode_general_request({"images": ["not-base64"], "prompts": ["prompt"]}) + + with self.assertRaises(ValueError): + decode_general_request({"images": [], "prompts": ["extra prompt"]}) + + with self.assertRaises(ValueError): + decode_reward_3d_request( + { + "videos": [["ZnJhbWU="]], + "prompts": ["prompt"], + "camera_trajectories": [], + } + ) + + def test_general_server_accepts_json_and_rejects_binary_payloads(self): + manager = FakeGeneralRewardManager() + app = serve_general_reward.create_app(manager=manager) + client = app.test_client() + + response = client.post( + "/", + json=encode_general_request([b"image"], ["prompt"]), + ) + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.get_json(), {"outputs": [0.75]}) + self.assertEqual(manager.calls, [([b"image"], ["prompt"])]) + + response = client.post( + "/", + data=b"\x80\x04untrusted-pickle-payload", + content_type="application/octet-stream", + ) + self.assertEqual(response.status_code, 400) + self.assertEqual(response.mimetype, "application/json") + self.assertEqual(len(manager.calls), 1) + + def test_reward_3d_server_accepts_json_and_rejects_binary_payloads(self): + manager = FakeReward3DManager() + app = serve_reward_3d.create_app( + manager=manager, + install_signal_handlers=False, + ) + client = app.test_client() + trajectories = [{"frame_0": "1 0 0 0"}] + + response = client.post( + "/", + json=encode_reward_3d_request([[b"frame"]], ["prompt"], trajectories), + ) + + self.assertEqual(response.status_code, 200) + self.assertEqual( + response.get_json(), + { + "outputs": [1.5], + "details": manager.last_results["per_video_results"], + }, + ) + self.assertEqual(manager.calls, [([[b"frame"]], ["prompt"], trajectories)]) + + response = client.post( + "/", + data=b"\x80\x04untrusted-pickle-payload", + content_type="application/octet-stream", + ) + self.assertEqual(response.status_code, 400) + self.assertEqual(response.mimetype, "application/json") + self.assertEqual(len(manager.calls), 1) + + def test_reward_clients_send_json(self): + general_session = FakeSession({"outputs": [0.5]}) + with patch.object(rewards.requests, "Session", return_value=general_session): + score_fn = rewards.remote_reward_general("cpu") + scores, metadata = score_fn( + torch.zeros(1, 3, 2, 2), + ["prompt"], + [{}], + ) + + self.assertEqual(scores, [0.5]) + self.assertEqual(metadata, {}) + _, general_kwargs = general_session.requests[0] + self.assertNotIn("data", general_kwargs) + images, prompts = decode_general_request(general_kwargs["json"]) + self.assertEqual(len(images), 1) + self.assertTrue(images[0].startswith(b"\xff\xd8")) + self.assertEqual(prompts, ["prompt"]) + + reward_3d_session = FakeSession({"outputs": [1.0], "details": None}) + with patch.object(rewards.requests, "Session", return_value=reward_3d_session): + score_fn = rewards.remote_reward_3d("cpu") + scores, metadata = score_fn( + torch.zeros(1, 1, 3, 2, 2), + ["prompt"], + [{"camera_trajectory": {"frame_0": "1 0 0 0"}}], + ) + + self.assertEqual(scores, [1.0]) + self.assertEqual(metadata, {}) + _, reward_3d_kwargs = reward_3d_session.requests[0] + self.assertNotIn("data", reward_3d_kwargs) + videos, prompts, trajectories = decode_reward_3d_request( + reward_3d_kwargs["json"] + ) + self.assertEqual(len(videos), 1) + self.assertEqual(len(videos[0]), 1) + self.assertTrue(videos[0][0].startswith(b"\xff\xd8")) + self.assertEqual(prompts, ["prompt"]) + self.assertEqual(trajectories, [{"frame_0": "1 0 0 0"}]) + + +if __name__ == "__main__": + unittest.main()