Add VULNERABLE_VIA edges, taint analysis improvements, and UI fixes

Add create_vulnerable_via_edges() to TaintAnalyzer for BFS-based
reachability from entry points to vulnerable nodes. Improve source/sink
detection with tiered matching: function name normalization, callee
name fallback, and deduplication of taint edges.

Fix orphaned tool results in query_controller by reconstructing missing
assistant messages with tool_calls when native storage callback was skipped.

Add ESC key shortcut to cancel edit mode without saving in QueryTab and
ExplainTab. Wire SecurityAnalysisWorker to emit VULNERABLE_VIA edge count.
This commit is contained in:
Jason Tang
2026-03-13 21:15:39 -04:00
parent 4dd2d71e5a
commit c54ab1b959
5 changed files with 257 additions and 3 deletions
+46 -1
View File
@@ -2550,7 +2550,52 @@ Tool Usage Guidelines:
# NOTE: Assistant message with tool calls and tool results are now loaded from native storage above
# No need to manually add them again - that would cause duplicates
# Validate: ensure tool results have a preceding assistant message with tool_calls.
# If the native_message_callback didn't fire (e.g. non-streaming fallback for
# thinking+tools models), the assistant message may be missing from storage.
has_assistant_with_tools = any(
msg.get("role") == "assistant" and (
msg.get("tool_calls") or
(isinstance(msg.get("content"), list) and
any(isinstance(b, dict) and b.get("type") == "tool_use" for b in msg["content"]))
)
for msg in conversation_messages
)
has_tool_results = any(
msg.get("role") == "tool" or
(msg.get("role") == "user" and isinstance(msg.get("content"), list) and
any(isinstance(b, dict) and b.get("type") == "tool_result" for b in msg["content"]))
for msg in conversation_messages
)
if has_tool_results and not has_assistant_with_tools and tool_calls:
log.log_warn(
"Native storage missing assistant message with tool_calls. "
"Reconstructing to prevent API error."
)
assistant_tool_calls = []
for tc in tool_calls:
args = tc.arguments if isinstance(tc.arguments, str) else json.dumps(tc.arguments)
assistant_tool_calls.append({
"id": tc.id,
"type": "function",
"function": {"name": tc.name, "arguments": args}
})
# Insert before the first tool result message
insert_idx = next(
(i for i, msg in enumerate(conversation_messages)
if msg.get("role") == "tool" or
(msg.get("role") == "user" and isinstance(msg.get("content"), list) and
any(isinstance(b, dict) and b.get("type") == "tool_result" for b in msg["content"]))),
len(conversation_messages)
)
conversation_messages.insert(insert_idx, {
"role": "assistant",
"content": "",
"tool_calls": assistant_tool_calls
})
return conversation_messages
except Exception as e:
+5 -2
View File
@@ -873,8 +873,11 @@ class SecurityAnalysisWorker(QThread):
if self._cancelled or (self._analyzer and self._analyzer.cancelled):
self.cancelled.emit()
return
edges_created = len(paths)
self.completed.emit(len(paths), edges_created)
vuln_via_edges = self._analyzer.create_vulnerable_via_edges()
if self._cancelled or (self._analyzer and self._analyzer.cancelled):
self.cancelled.emit()
return
self.completed.emit(len(paths), vuln_via_edges)
except Exception as exc:
self.failed.emit(str(exc))
+161
View File
@@ -70,6 +70,15 @@ class TaintAnalyzer:
"CALLS_VULNERABLE_FUNCTION",
}
ENTRY_POINT_NAMES = {
"main", "_main", "wmain", "_wmain",
"WinMain", "wWinMain", "_WinMain@16", "_wWinMain@16",
"DllMain", "_DllMain@12", "DllEntryPoint",
"start", "_start", "entry", "_entry",
"mainCRTStartup", "wmainCRTStartup",
"WinMainCRTStartup", "wWinMainCRTStartup",
}
def __init__(self, graph_store: GraphStore, binary_hash: str):
self.graph_store = graph_store
self.binary_hash = binary_hash
@@ -133,6 +142,8 @@ class TaintAnalyzer:
def _create_taint_edges(self, path: List[int]) -> None:
for idx in range(len(path) - 1):
if self.graph_store.has_edge(path[idx], path[idx + 1], EdgeType.TAINT_FLOWS_TO.value):
continue
self.graph_store.add_edge(GraphEdge(
binary_hash=self.binary_hash,
source_id=path[idx],
@@ -154,32 +165,182 @@ class TaintAnalyzer:
nodes = self.graph_store.get_nodes_by_type(self.binary_hash, "FUNCTION")
results = []
for node in nodes:
# Tier 1: Security flags
if self._has_any_flag(node.security_flags, self.SOURCE_FLAGS):
results.append(node)
continue
# Tier 2: Network APIs
if self._has_any_api(node.network_apis, self.TAINT_SOURCES):
results.append(node)
continue
# Tier 3 (part a): File I/O APIs
if self._has_any_api(node.file_io_apis, self.TAINT_SOURCES):
results.append(node)
continue
# Tier 3 (part b): Function name check against TAINT_SOURCES
if node.name:
if node.name in self.TAINT_SOURCES or self._normalize_function_name(node.name) in self.TAINT_SOURCES:
results.append(node)
continue
# Tier 4: Callees fallback
callee_ids = self._get_callees(node.id)
matched = False
for callee_id in callee_ids:
callee = self.graph_store.get_node_by_id(callee_id)
if callee and callee.name:
if callee.name in self.TAINT_SOURCES or self._normalize_function_name(callee.name) in self.TAINT_SOURCES:
matched = True
break
if matched:
results.append(node)
continue
return results
def _find_sink_nodes(self) -> List[GraphNode]:
nodes = self.graph_store.get_nodes_by_type(self.binary_hash, "FUNCTION")
results = []
for node in nodes:
# Tier 1: Security flags
if self._has_any_flag(node.security_flags, self.SINK_FLAGS):
results.append(node)
continue
# Tier 2: Network APIs
if self._has_any_api(node.network_apis, self.TAINT_SINKS):
results.append(node)
continue
# Tier 3 (part a): File I/O APIs
if self._has_any_api(node.file_io_apis, self.TAINT_SINKS):
results.append(node)
continue
# Tier 3 (part b): Function name check against TAINT_SINKS
if node.name:
if node.name in self.TAINT_SINKS or self._normalize_function_name(node.name) in self.TAINT_SINKS:
results.append(node)
continue
# Tier 4: Callees fallback
callee_ids = self._get_callees(node.id)
matched = False
for callee_id in callee_ids:
callee = self.graph_store.get_node_by_id(callee_id)
if callee and callee.name:
if callee.name in self.TAINT_SINKS or self._normalize_function_name(callee.name) in self.TAINT_SINKS:
matched = True
break
if matched:
results.append(node)
continue
return results
def create_vulnerable_via_edges(self) -> int:
"""Create VULNERABLE_VIA edges from entry points to vulnerable nodes."""
entry_points = self._find_entry_points()
vulnerable_nodes = self._find_vulnerable_nodes()
if not entry_points or not vulnerable_nodes:
return 0
edges_created = 0
for entry in entry_points:
if self.cancelled:
break
for vuln_node in vulnerable_nodes:
if entry.id == vuln_node.id:
continue
if self.graph_store.has_edge(entry.id, vuln_node.id, EdgeType.VULNERABLE_VIA.value):
continue
path_length = self._bfs_path_length(entry.id, vuln_node.id)
if path_length is not None:
vuln_type = self._get_vulnerability_type(vuln_node)
metadata = f'{{"path_length":{path_length},"vuln_type":"{vuln_type}"}}'
self.graph_store.add_edge(GraphEdge(
binary_hash=self.binary_hash,
source_id=entry.id,
target_id=vuln_node.id,
edge_type=EdgeType.VULNERABLE_VIA,
weight=1.0,
metadata=metadata,
))
edges_created += 1
return edges_created
def _find_entry_points(self) -> List[GraphNode]:
nodes = self.graph_store.get_nodes_by_type(self.binary_hash, "FUNCTION")
entry_points = []
seen_ids: Set = set()
for node in nodes:
is_entry = False
flags = node.security_flags or []
if "ENTRY_POINT" in flags or "EXPORTED" in flags:
is_entry = True
if node.name and node.name in self.ENTRY_POINT_NAMES:
is_entry = True
if is_entry and node.id not in seen_ids:
seen_ids.add(node.id)
entry_points.append(node)
return entry_points
def _find_vulnerable_nodes(self) -> List[GraphNode]:
nodes = self.graph_store.get_nodes_by_type(self.binary_hash, "FUNCTION")
vulnerable = []
for node in nodes:
flags = node.security_flags or []
for flag in flags:
if flag.endswith("_RISK") or flag.startswith("VULN_"):
vulnerable.append(node)
break
return vulnerable
def _get_vulnerability_type(self, node: GraphNode) -> str:
flags = node.security_flags or []
for flag in flags:
if flag.startswith("VULN_"):
return flag[5:]
for flag in flags:
if flag.endswith("_RISK"):
return flag.replace("_RISK", "")
return "UNKNOWN"
def _bfs_path_length(self, source_id, target_id) -> Optional[int]:
if source_id == target_id:
return 0
visited = {source_id}
queue = [(source_id, 0)]
while queue:
current_id, depth = queue.pop(0)
if depth >= self.MAX_PATH_LENGTH:
continue
for neighbor_id in self._get_callees(current_id):
if neighbor_id == target_id:
return depth + 1
if neighbor_id not in visited:
visited.add(neighbor_id)
queue.append((neighbor_id, depth + 1))
return None
@staticmethod
def _normalize_function_name(name: str) -> str:
if not name:
return name
normalized = name
if "::" in normalized:
normalized = normalized.split("::")[-1]
if ".DLL_" in normalized.upper():
idx = normalized.upper().find(".DLL_")
if idx > 0:
normalized = normalized[idx + 5:]
if normalized.startswith("<EXTERNAL>::"):
normalized = normalized[12:]
if normalized.startswith("__imp_"):
normalized = normalized[6:]
while normalized.startswith("_") and len(normalized) > 1:
normalized = normalized[1:]
at_idx = normalized.rfind("@")
if at_idx > 0:
suffix = normalized[at_idx + 1:]
if suffix.isdigit():
normalized = normalized[:at_idx]
return normalized
@staticmethod
def _has_any_flag(flags: Iterable[str], targets: Set[str]) -> bool:
return any(flag in targets for flag in (flags or []))
+24
View File
@@ -175,6 +175,12 @@ class ExplainTabView(QWidget):
self.explain_editor.setPlainText(self.markdown_content)
self.explain_editor.hide() # Hidden by default
# ESC key discards edits and returns to view mode
from PySide6.QtGui import QShortcut
esc_shortcut = QShortcut(QKeySequence(Qt.Key_Escape), self.explain_editor)
esc_shortcut.setContext(Qt.ShortcutContext.WidgetShortcut)
esc_shortcut.activated.connect(self.cancel_edit_mode)
def create_line_explanation_panel(self):
"""Create the line explanation panel (collapsible, dismissable)"""
self.line_explanation_group = QGroupBox("Line Explanation")
@@ -304,6 +310,24 @@ class ExplainTabView(QWidget):
self.edit_mode_changed.emit(self.is_edit_mode)
def cancel_edit_mode(self):
"""Discard edits and return to view mode without saving"""
if not self.is_edit_mode:
return
self.is_edit_mode = False
old_sizes = self.main_splitter.sizes()
security_size = old_sizes[3] if len(old_sizes) >= 4 else 0
line_size = old_sizes[2] if len(old_sizes) >= 3 else 0
self.explain_editor.hide()
self.explain_browser.show()
self.edit_save_button.setText("Edit")
from PySide6.QtCore import QTimer
def _restore_sizes():
total = sum(self.main_splitter.sizes())
active_size = total - line_size - security_size
self.main_splitter.setSizes([active_size, 0, line_size, security_size])
QTimer.singleShot(0, _restore_sizes)
def set_current_offset(self, offset_hex):
"""Update the displayed current offset"""
self.current_offset_label.setText(offset_hex)
+21
View File
@@ -186,6 +186,12 @@ class QueryTabView(QWidget):
self.query_editor.setPlainText(self.markdown_content)
self.query_editor.hide() # Hidden by default
# ESC key discards edits and returns to view mode
from PySide6.QtGui import QShortcut
esc_shortcut = QShortcut(QKeySequence(Qt.Key_Escape), self.query_editor)
esc_shortcut.setContext(Qt.ShortcutContext.WidgetShortcut)
esc_shortcut.activated.connect(self.cancel_edit_mode)
def create_history_table(self):
self.history_table = QTableWidget()
self.history_table.setColumnCount(2)
@@ -293,6 +299,21 @@ class QueryTabView(QWidget):
self.edit_mode_changed.emit(self.is_edit_mode)
def cancel_edit_mode(self):
"""Discard edits and return to view mode without saving"""
if not self.is_edit_mode:
return
self.is_edit_mode = False
old_sizes = self.splitter.sizes()
history_size = old_sizes[2] if len(old_sizes) >= 3 else 80
input_size = old_sizes[3] if len(old_sizes) >= 4 else 100
self.query_editor.hide()
self.query_browser.show()
self.edit_save_button.setText("Edit")
total = sum(self.splitter.sizes())
active_size = total - history_size - input_size
self.splitter.setSizes([active_size, 0, history_size, input_size])
def on_submit_clicked(self):
"""Handle submit button click - toggles between submit and stop"""
if self.query_running: