1#
2# Licensed to the Apache Software Foundation (ASF) under one
3# or more contributor license agreements. See the NOTICE file
4# distributed with this work for additional information
5# regarding copyright ownership. The ASF licenses this file
6# to you under the Apache License, Version 2.0 (the
7# "License"); you may not use this file except in compliance
8# with the License. You may obtain a copy of the License at
9#
10# http://www.apache.org/licenses/LICENSE-2.0
11#
12# Unless required by applicable law or agreed to in writing,
13# software distributed under the License is distributed on an
14# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15# KIND, either express or implied. See the License for the
16# specific language governing permissions and limitations
17# under the License.
18from __future__ import annotations
19
20import functools
21import inspect
22import logging
23import pkgutil
24import sys
25from collections import defaultdict
26from collections.abc import Callable, Iterator
27from importlib import import_module
28from typing import TYPE_CHECKING
29
30from .dag_file import (
31 MODIFIED_DAG_MODULE_NAME as MODIFIED_DAG_MODULE_NAME,
32 UNUSUAL_MODULE_PREFIX as UNUSUAL_MODULE_PREFIX,
33 accepts_dag_definition as accepts_dag_definition,
34 get_unique_dag_module_name as get_unique_dag_module_name,
35 might_contain_dag as might_contain_dag,
36 might_contain_dag_via_default_heuristic as might_contain_dag_via_default_heuristic,
37)
38from .file_discovery import (
39 find_path_from_directory as find_path_from_directory,
40)
41
42if sys.version_info >= (3, 12):
43 from importlib import metadata
44else:
45 import importlib_metadata as metadata
46
47log = logging.getLogger(__name__)
48
49EPnD = tuple[metadata.EntryPoint, metadata.Distribution]
50
51if TYPE_CHECKING:
52 from types import ModuleType
53
54
55def accepts_context(callback: Callable) -> bool:
56 """Check if callback accepts a 'context' parameter or **kwargs."""
57 try:
58 sig = inspect.signature(callback)
59 except (ValueError, TypeError):
60 return True
61 params = sig.parameters
62 return "context" in params or any(p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values())
63
64
65def accepts_keyword_args(func: Callable) -> bool:
66 """Check if a callable accepts any keyword arguments (named params or **kwargs)."""
67 try:
68 sig = inspect.signature(func)
69 except (ValueError, TypeError):
70 return True
71 return any(
72 p.kind
73 in (
74 inspect.Parameter.POSITIONAL_OR_KEYWORD,
75 inspect.Parameter.KEYWORD_ONLY,
76 inspect.Parameter.VAR_KEYWORD,
77 )
78 for p in sig.parameters.values()
79 )
80
81
82def import_string(dotted_path: str):
83 """
84 Import a dotted module path and return the attribute/class designated by the last name in the path.
85
86 Raise ImportError if the import failed.
87 """
88 # TODO: Add support for nested classes. Currently, it only works for top-level classes.
89 try:
90 module_path, class_name = dotted_path.rsplit(".", 1)
91 except ValueError:
92 raise ImportError(f"{dotted_path} doesn't look like a module path")
93
94 module = import_module(module_path)
95
96 try:
97 return getattr(module, class_name)
98 except AttributeError:
99 raise ImportError(f'Module "{module_path}" does not define a "{class_name}" attribute/class')
100
101
102def qualname(o: object | Callable, use_qualname: bool = False, exclude_module: bool = False) -> str:
103 """
104 Convert an attribute/class/callable to a string.
105
106 By default, returns a string importable by ``import_string`` (includes module path).
107 With exclude_module=True, returns only the qualified name without module prefix,
108 useful for stable identification across deployments where module paths may vary.
109 """
110 if callable(o) and hasattr(o, "__module__"):
111 if exclude_module:
112 if hasattr(o, "__qualname__"):
113 return o.__qualname__
114 if hasattr(o, "__name__"):
115 return o.__name__
116 # Handle functools.partial objects specifically (not just any object with 'func' attr)
117 if isinstance(o, functools.partial):
118 return qualname(o.func, exclude_module=True)
119 return type(o).__qualname__
120 if use_qualname and hasattr(o, "__qualname__"):
121 return f"{o.__module__}.{o.__qualname__}"
122 if hasattr(o, "__name__"):
123 return f"{o.__module__}.{o.__name__}"
124
125 cls = o
126
127 if not isinstance(cls, type): # instance or class
128 cls = type(cls)
129
130 name = cls.__qualname__
131 module = cls.__module__
132
133 if exclude_module:
134 return name
135
136 if module and module != "__builtin__":
137 return f"{module}.{name}"
138
139 return name
140
141
142def iter_namespace(ns: ModuleType):
143 return pkgutil.iter_modules(ns.__path__, ns.__name__ + ".")
144
145
146def is_valid_dotpath(path: str) -> bool:
147 """
148 Check if a string follows valid dotpath format (ie: 'package.subpackage.module').
149
150 :param path: String to check
151 """
152 import re
153
154 if not isinstance(path, str):
155 return False
156
157 # Pattern explanation:
158 # ^ - Start of string
159 # [a-zA-Z_] - Must start with letter or underscore
160 # [a-zA-Z0-9_] - Following chars can be letters, numbers, or underscores
161 # (\.[a-zA-Z_][a-zA-Z0-9_]*)* - Can be followed by dots and valid identifiers
162 # $ - End of string
163 pattern = r"^[a-zA-Z_][a-zA-Z0-9_]*(\.[a-zA-Z_][a-zA-Z0-9_]*)*$"
164
165 return bool(re.match(pattern, path))
166
167
168@functools.cache
169def _get_grouped_entry_points() -> dict[str, list[EPnD]]:
170 mapping: dict[str, list[EPnD]] = defaultdict(list)
171 for dist in metadata.distributions():
172 try:
173 for e in dist.entry_points:
174 mapping[e.group].append((e, dist))
175 except Exception as e:
176 log.warning("Error when retrieving package metadata (skipping it): %s, %s", dist, e)
177 return mapping
178
179
180def entry_points_with_dist(group: str) -> Iterator[EPnD]:
181 """
182 Retrieve entry points of the given group.
183
184 This is like the ``entry_points()`` function from ``importlib.metadata``,
185 except it also returns the distribution the entry point was loaded from.
186
187 Note that this may return multiple distributions to the same package if they
188 are loaded from different ``sys.path`` entries. The caller site should
189 implement appropriate deduplication logic if needed.
190
191 :param group: Filter results to only this entrypoint group
192 :return: Generator of (EntryPoint, Distribution) objects for the specified groups
193 """
194 return iter(_get_grouped_entry_points()[group])