Team Ai
Modelpublic

mangsense/codet5-java-vulnerability-lora

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes9downloads
handler.py197 linesDownload Raw Back to root
1import os
2import torch
3from transformers import AutoTokenizer, T5ForSequenceClassification
4from typing import Dict, List, Any
5
6class EndpointHandler:
7    """
8    HuggingFace Inference Endpoint Handler for Java Vulnerability Detection
9    CodeT5 기반 분류 모델 (LoRA fine-tuned)
10    """
11    
12    def __init__(self, path="."):
13        """
14        모델과 토크나이저를 초기화합니다.
15        
16        Args:
17            path (str): 모델이 저장된 경로 (HuggingFace Hub에서 자동으로 설정됨)
18        """
19        print(f"🚀 Loading Java Vulnerability Detection Model from {path}")
20        
21        # 디바이스 설정
22        self.device = "cuda" if torch.cuda.is_available() else "cpu"
23        print(f"📍 Device: {self.device}")
24        
25        # 토크나이저 로드
26        self.tokenizer = AutoTokenizer.from_pretrained(path)
27        
28        # T5ForSequenceClassification 모델 로드
29        self.model = T5ForSequenceClassification.from_pretrained(
30            path,
31            torch_dtype=torch.float16 if self.device == "cuda" else torch.float32
32        )
33        
34        # 모델을 평가 모드로 설정하고 디바이스로 이동
35        self.model.to(self.device)
36        self.model.eval()
37        
38        print("✅ Model loaded successfully!")
39    
40    def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]:
41        """
42        메인 추론 메서드 (HuggingFace Inference API가 호출)
43        
44        Args:
45            data (dict): 입력 데이터
46                - "inputs" (str): Java 코드 또는
47                - "code" (str): Java 코드
48        
49        Returns:
50            list: 예측 결과 리스트
51        """
52        # 1. 전처리
53        inputs = self.preprocess(data)
54        
55        # 2. 추론
56        outputs = self.inference(inputs)
57        
58        # 3. 후처리
59        result = self.postprocess(outputs)
60        
61        return result
62    
63    def preprocess(self, request: Dict[str, Any]) -> Dict[str, torch.Tensor]:
64        """
65        입력 데이터를 전처리합니다.
66        
67        Args:
68            request (dict): API 요청 데이터
69        
70        Returns:
71            dict: 토크나이즈된 입력 텐서
72        """
73        # 입력 텍스트 추출
74        if isinstance(request, dict):
75            # "inputs" 또는 "code" 키에서 Java 코드 추출
76            code = request.get("inputs") or request.get("code")
77        elif isinstance(request, list) and len(request) > 0:
78            code = request[0].get("inputs") or request[0].get("code")
79        elif isinstance(request, str):
80            code = request
81        else:
82            raise ValueError(
83                "Invalid request format. Expected {'inputs': 'Java code here'} "
84                "or {'code': 'Java code here'}"
85            )
86        
87        if not code:
88            raise ValueError("No code provided in request")
89        
90        # 프롬프트 템플릿 적용
91        input_text = f"Is this Java code vulnerable?:\n{code}"
92        
93        # 토크나이징
94        inputs = self.tokenizer(
95            input_text,
96            max_length=512,
97            truncation=True,
98            padding="max_length",
99            return_tensors="pt"
100        )
101        
102        # 디바이스로 이동
103        inputs = {k: v.to(self.device) for k, v in inputs.items()}
104        
105        return inputs
106    
107    def inference(self, inputs: Dict[str, torch.Tensor]) -> torch.Tensor:
108        """
109        모델 추론을 수행합니다.
110        
111        Args:
112            inputs (dict): 전처리된 입력 텐서
113        
114        Returns:
115            torch.Tensor: 모델 출력 로짓
116        """
117        with torch.no_grad():
118            outputs = self.model(**inputs)
119            logits = outputs.logits
120        
121        return logits
122    
123    def postprocess(self, logits: torch.Tensor) -> List[Dict[str, Any]]:
124        """
125        모델 출력을 사람이 읽을 수 있는 형태로 변환합니다.
126        
127        Args:
128            logits (torch.Tensor): 모델 출력 로짓
129        
130        Returns:
131            list: 예측 결과 리스트
132        """
133        # 로짓 처리 (단일 출력 vs 다중 클래스)
134        if logits.shape[-1] == 1:
135            # Binary classification with single output
136            prob = torch.sigmoid(logits).item()
137            predicted_class = 1 if prob > 0.5 else 0
138            confidence = prob if predicted_class == 1 else (1 - prob)
139            probabilities = {
140                "LABEL_0": 1 - prob,
141                "LABEL_1": prob
142            }
143        else:
144            # Multi-class classification
145            probs = torch.softmax(logits, dim=1)[0]
146            predicted_class = torch.argmax(logits, dim=1).item()
147            confidence = probs[predicted_class].item()
148            probabilities = {
149                f"LABEL_{i}": probs[i].item() 
150                for i in range(len(probs))
151            }
152        
153        # 레이블 매핑
154        label_map = {
155            0: "safe",
156            1: "vulnerable"
157        }
158        
159        # 결과 포맷팅
160        result = {
161            "label": label_map.get(predicted_class, f"LABEL_{predicted_class}"),
162            "score": confidence,
163            "probabilities": probabilities,
164            "details": {
165                "is_vulnerable": predicted_class == 1,
166                "confidence_percentage": f"{confidence * 100:.2f}%",
167                "safe_probability": probabilities.get("LABEL_0", 0),
168                "vulnerable_probability": probabilities.get("LABEL_1", 0)
169            }
170        }
171        
172        return [result]
173
174
175# 로컬 테스트용 코드
176if __name__ == "__main__":
177    # 로컬에서 테스트할 때 사용
178    handler = EndpointHandler(path=".")
179    
180    # 테스트 케이스
181    test_code = """
182import java.sql.*;
183public class SQLInjectionVulnerable {
184    public void getUser(String userInput) {
185        String query = "SELECT * FROM users WHERE username = '" + userInput + "'";
186        Statement statement = connection.createStatement();
187        ResultSet resultSet = statement.executeQuery(query);
188    }
189}
190"""
191    
192    # 추론 실행
193    request = {"inputs": test_code}
194    result = handler(request)
195    
196    print("\n📊 Test Result:")
197    print(result)