mirror of
https://github.com/jtang613/IDAssist
synced 2026-08-09 12:43:55 +00:00
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:
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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 []))
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user