mangsense/codet5-java-vulnerability-lora
09
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)