Coverage for /usr/lib/python3/dist-packages/sympy/strategies/core.py: 32%

80 statements  

« prev     ^ index     » next       coverage.py v7.9.1, created at 2025-06-14 15:55 +0200

1""" Generic SymPy-Independent Strategies """ 

2from __future__ import annotations 

3from collections.abc import Callable, Mapping 

4from typing import TypeVar 

5from sys import stdout 

6 

7 

8_S = TypeVar('_S') 

9_T = TypeVar('_T') 

10 

11 

12def identity(x: _T) -> _T: 

13 return x 

14 

15 

16def exhaust(rule: Callable[[_T], _T]) -> Callable[[_T], _T]: 

17 """ Apply a rule repeatedly until it has no effect """ 

18 def exhaustive_rl(expr: _T) -> _T: 

19 new, old = rule(expr), expr 

20 while new != old: 

21 new, old = rule(new), new 

22 return new 

23 return exhaustive_rl 

24 

25 

26def memoize(rule: Callable[[_S], _T]) -> Callable[[_S], _T]: 

27 """Memoized version of a rule 

28 

29 Notes 

30 ===== 

31 

32 This cache can grow infinitely, so it is not recommended to use this 

33 than ``functools.lru_cache`` unless you need very heavy computation. 

34 """ 

35 cache: dict[_S, _T] = {} 

36 

37 def memoized_rl(expr: _S) -> _T: 

38 if expr in cache: 

39 return cache[expr] 

40 else: 

41 result = rule(expr) 

42 cache[expr] = result 

43 return result 

44 return memoized_rl 

45 

46 

47def condition( 

48 cond: Callable[[_T], bool], rule: Callable[[_T], _T] 

49) -> Callable[[_T], _T]: 

50 """ Only apply rule if condition is true """ 

51 def conditioned_rl(expr: _T) -> _T: 

52 if cond(expr): 

53 return rule(expr) 

54 return expr 

55 return conditioned_rl 

56 

57 

58def chain(*rules: Callable[[_T], _T]) -> Callable[[_T], _T]: 

59 """ 

60 Compose a sequence of rules so that they apply to the expr sequentially 

61 """ 

62 def chain_rl(expr: _T) -> _T: 

63 for rule in rules: 

64 expr = rule(expr) 

65 return expr 

66 return chain_rl 

67 

68 

69def debug(rule, file=None): 

70 """ Print out before and after expressions each time rule is used """ 

71 if file is None: 

72 file = stdout 

73 

74 def debug_rl(*args, **kwargs): 

75 expr = args[0] 

76 result = rule(*args, **kwargs) 

77 if result != expr: 

78 file.write("Rule: %s\n" % rule.__name__) 

79 file.write("In: %s\nOut: %s\n\n" % (expr, result)) 

80 return result 

81 return debug_rl 

82 

83 

84def null_safe(rule: Callable[[_T], _T | None]) -> Callable[[_T], _T]: 

85 """ Return original expr if rule returns None """ 

86 def null_safe_rl(expr: _T) -> _T: 

87 result = rule(expr) 

88 if result is None: 

89 return expr 

90 return result 

91 return null_safe_rl 

92 

93 

94def tryit(rule: Callable[[_T], _T], exception) -> Callable[[_T], _T]: 

95 """ Return original expr if rule raises exception """ 

96 def try_rl(expr: _T) -> _T: 

97 try: 

98 return rule(expr) 

99 except exception: 

100 return expr 

101 return try_rl 

102 

103 

104def do_one(*rules: Callable[[_T], _T]) -> Callable[[_T], _T]: 

105 """ Try each of the rules until one works. Then stop. """ 

106 def do_one_rl(expr: _T) -> _T: 

107 for rl in rules: 

108 result = rl(expr) 

109 if result != expr: 

110 return result 

111 return expr 

112 return do_one_rl 

113 

114 

115def switch( 

116 key: Callable[[_S], _T], 

117 ruledict: Mapping[_T, Callable[[_S], _S]] 

118) -> Callable[[_S], _S]: 

119 """ Select a rule based on the result of key called on the function """ 

120 def switch_rl(expr: _S) -> _S: 

121 rl = ruledict.get(key(expr), identity) 

122 return rl(expr) 

123 return switch_rl 

124 

125 

126# XXX Untyped default argument for minimize function 

127# where python requires SupportsRichComparison type 

128def _identity(x): 

129 return x 

130 

131 

132def minimize( 

133 *rules: Callable[[_S], _T], 

134 objective=_identity 

135) -> Callable[[_S], _T]: 

136 """ Select result of rules that minimizes objective 

137 

138 >>> from sympy.strategies import minimize 

139 >>> inc = lambda x: x + 1 

140 >>> dec = lambda x: x - 1 

141 >>> rl = minimize(inc, dec) 

142 >>> rl(4) 

143 3 

144 

145 >>> rl = minimize(inc, dec, objective=lambda x: -x) # maximize 

146 >>> rl(4) 

147 5 

148 """ 

149 def minrule(expr: _S) -> _T: 

150 return min([rule(expr) for rule in rules], key=objective) 

151 return minrule