|
1 | 1 | from __future__ import annotations |
2 | 2 |
|
| 3 | +import ast |
| 4 | +import re |
3 | 5 | import textwrap |
4 | 6 | from pathlib import Path |
5 | | -from typing import Dict, List, Optional, Tuple, Union |
| 7 | +from typing import Dict, List, Optional, Set, Tuple, Union |
6 | 8 |
|
7 | 9 | from frida_bindgen_core import Procedure, Type |
| 10 | +from frida_bindgen_core.naming import to_pascal_case |
8 | 11 |
|
9 | 12 | from .model import ( |
10 | 13 | Enumeration, |
@@ -32,7 +35,8 @@ def read_asset(name: str) -> str: |
32 | 35 |
|
33 | 36 |
|
34 | 37 | FACADE_TYPING_IMPORTS = ( |
35 | | - "from typing import Any, Callable, Dict, List, Literal, Mapping, NotRequired, Optional, Tuple, TypedDict, Union" |
| 38 | + "from typing import (Any, Callable, Dict, List, Literal, Mapping, NotRequired, Optional, Tuple, TypedDict, " |
| 39 | + "Union, overload)" |
36 | 40 | ) |
37 | 41 |
|
38 | 42 |
|
@@ -366,6 +370,8 @@ def generate_aio_class(otype: ObjectType, model: Model) -> str: |
366 | 370 | if not members: |
367 | 371 | members.append(" pass") |
368 | 372 |
|
| 373 | + members = apply_signal_overloads(members, otype, model) |
| 374 | + |
369 | 375 | return f"class {otype.py_name}:\n" + "\n\n".join(members) |
370 | 376 |
|
371 | 377 |
|
@@ -461,6 +467,8 @@ def generate_py_class(otype: ObjectType, model: Model) -> str: |
461 | 467 | if not members: |
462 | 468 | members.append(" pass") |
463 | 469 |
|
| 470 | + members = apply_signal_overloads(members, otype, model) |
| 471 | + |
464 | 472 | return f"class {otype.py_name}:\n" + "\n\n".join(members) |
465 | 473 |
|
466 | 474 |
|
@@ -496,13 +504,54 @@ def generate_facade_init(otype: ObjectType) -> str: |
496 | 504 |
|
497 | 505 |
|
498 | 506 | def generate_facade_signals() -> str: |
499 | | - return """ def on(self, signal, callback): |
| 507 | + return """ def on(self, signal: str, callback: Callable[..., Any]) -> None: |
500 | 508 | self._impl.on(signal, _make_signal_handler(callback)) |
501 | 509 |
|
502 | | - def off(self, signal, callback): |
| 510 | + def off(self, signal: str, callback: Callable[..., Any]) -> None: |
503 | 511 | self._impl.off(signal, callback)""" |
504 | 512 |
|
505 | 513 |
|
| 514 | +def facade_prelude_names(model: Model) -> Set[str]: |
| 515 | + names = set() |
| 516 | + for asset in model.customizations.facade_preludes: |
| 517 | + for node in ast.parse(read_asset(asset)).body: |
| 518 | + if isinstance(node, ast.ClassDef): |
| 519 | + names.add(node.name) |
| 520 | + elif isinstance(node, ast.Assign): |
| 521 | + names.update(t.id for t in node.targets if isinstance(t, ast.Name)) |
| 522 | + return names |
| 523 | + |
| 524 | + |
| 525 | +def signal_callback_alias(otype: ObjectType, signal) -> str: |
| 526 | + return f"{otype.py_name}{to_pascal_case(signal.name.replace('-', '_'))}Callback" |
| 527 | + |
| 528 | + |
| 529 | +def signal_overload_block(otype: ObjectType, model: Model, name: str) -> Optional[str]: |
| 530 | + aliases = facade_prelude_names(model) |
| 531 | + lines = [] |
| 532 | + for signal in otype.signals: |
| 533 | + alias = signal_callback_alias(otype, signal) |
| 534 | + if alias in aliases: |
| 535 | + lines.append(" @overload") |
| 536 | + lines.append(f' def {name}(self, signal: Literal["{signal.name}"], callback: {alias}) -> None: ...') |
| 537 | + return "\n".join(lines) if len(lines) > 2 else None |
| 538 | + |
| 539 | + |
| 540 | +def apply_signal_overloads(members: List[str], otype: ObjectType, model: Model) -> List[str]: |
| 541 | + result = [] |
| 542 | + for member in members: |
| 543 | + for name in ("on", "off"): |
| 544 | + block = signal_overload_block(otype, model, name) |
| 545 | + if block is None: |
| 546 | + continue |
| 547 | + pattern = re.compile(rf"^ (?:async )?def {name}\(self, signal", re.MULTILINE) |
| 548 | + match = pattern.search(member) |
| 549 | + if match is not None: |
| 550 | + member = member[: match.start()] + block + "\n" + member[match.start() :] |
| 551 | + result.append(member) |
| 552 | + return result |
| 553 | + |
| 554 | + |
506 | 555 | def facade_repr_property_names(otype: ObjectType) -> List[str]: |
507 | 556 | names = [] |
508 | 557 | for method in otype.methods: |
|
0 commit comments