class Solution:
def bstToGst(self, root: Optional[TreeNode]) -> Optional[TreeNode]:
self.dfs(root, self.getSum(root))
return(root)
def getSum(self, node: Optional[TreeNode]) -> int:
if not node:
return(0)
totalSum = 0
stack = [node]
while stack:
currNode = stack.pop()
totalSum += currNode.val
if currNode.left:
stack.append(currNode.left)
if currNode.right:
stack.append(currNode.right)
return(totalSum)
def dfs(self, node: Optional[TreeNode], currSum: int) -> Optional[TreeNode]:
if not node:
return(None)
currVal = node.val
leftSum = self.getSum(node.left)
node.val = currSum - leftSum
node.left = self.dfs(node.left, currSum)
node.right = self.dfs(node.right, currSum - leftSum - currVal)
return(node)